Compare commits

..

62 Commits

Author SHA1 Message Date
henrygd
3534552d37 updates 2026-04-29 20:06:51 -04:00
henrygd
723401819f update 2026-04-29 18:41:42 -04:00
henrygd
2ea576c989 updates 2026-04-29 18:38:09 -04:00
henrygd
526a2c6aab updates 2026-04-29 18:21:39 -04:00
henrygd
aaa8eb773f updates 2026-04-29 18:05:40 -04:00
henrygd
099935e78e updates 2026-04-29 17:59:30 -04:00
henrygd
d2eb3b259a updates 2026-04-29 15:49:43 -04:00
henrygd
b89314889d update collections 2026-04-28 19:20:27 -04:00
henrygd
04e2b8b974 updates 2026-04-28 18:29:41 -04:00
henrygd
891b03426f updates 2026-04-28 17:46:56 -04:00
henrygd
b182b699d7 update 2026-04-27 10:05:58 -04:00
henrygd
e65a4a515e updates 2026-04-26 22:40:18 -04:00
henrygd
df249b24f6 updates 2026-04-26 19:25:57 -04:00
henrygd
788483ac56 updates 2026-04-26 19:03:21 -04:00
henrygd
f830665984 updates 2026-04-26 17:19:15 -04:00
henrygd
af49ebf2df updates 2026-04-26 15:37:00 -04:00
henrygd
0378023b6f update 2026-04-26 13:37:33 -04:00
henrygd
89ac8dc585 updates 2026-04-25 18:43:47 -04:00
henrygd
9896bcdf43 updates 2026-04-25 15:27:24 -04:00
henrygd
ddd47e67ac update 2026-04-25 14:39:04 -04:00
henrygd
027159420c update 2026-04-24 01:50:27 -04:00
henrygd
e154123511 updates 2026-04-23 21:34:56 -04:00
henrygd
9f7c1b22bb updates 2026-04-23 02:33:35 -04:00
henrygd
0d440e5fb9 updates 2026-04-23 01:13:01 -04:00
henrygd
5fc774666f updates 2026-04-22 21:40:52 -04:00
henrygd
8f03cbf11c updates 2026-04-22 19:40:21 -04:00
henrygd
1c5808f430 update 2026-04-22 19:29:36 -04:00
henrygd
a35cc6ef39 upupdate 2026-04-22 18:03:31 -04:00
henrygd
16e0f6c4a2 updates 2026-04-22 17:42:11 -04:00
henrygd
6472af1ba4 updates 2026-04-21 21:57:24 -04:00
henrygd
e931165566 updates 2026-04-21 15:44:08 -04:00
henrygd
48fe407292 use network probes 2026-04-21 15:29:46 -04:00
henrygd
a95376b4a2 updates 2026-04-21 12:33:16 -04:00
henrygd
732983493a update 2026-04-20 21:28:09 -04:00
henrygd
264b17f429 updte 2026-04-20 21:27:16 -04:00
henrygd
cef5ab10a5 updates 2026-04-20 21:24:46 -04:00
henrygd
3a881e1d5e add probes page 2026-04-20 11:52:37 -04:00
henrygd
209bb4ebb4 update 2026-04-20 10:48:05 -04:00
henrygd
e71ffd4d2a updates 2026-04-19 21:44:21 -04:00
henrygd
ea19ef6334 updates 2026-04-19 19:12:04 -04:00
henrygd
40da2b4358 updates 2026-04-18 20:28:22 -04:00
henrygd
d0d5912d85 updates 2026-04-18 18:09:45 -04:00
Claude
4162186ae0 Merge remote-tracking branch 'upstream/main' into feat/network-probes
# Conflicts:
#	agent/connection_manager.go
2026-04-18 01:19:49 +00:00
xiaomiku01
578ba985e9 Merge branch 'main' into feat/network-probes
Resolved conflict in internal/records/records.go:
- Upstream refactor moved deletion code to records_deletion.go and
  switched averaging functions from package-level globals to local
  variables (var row StatsRecord / params := make(dbx.Params, 1)).
- Kept AverageProbeStats and rewrote it to match the new local-variable
  pattern.
- Dropped duplicated deletion helpers from records.go (they now live in
  records_deletion.go).
- Added "network_probe_stats" to the collections list in
  records_deletion.go:deleteOldSystemStats so probe stats keep the same
  retention policy.
2026-04-17 13:49:18 +08:00
xiaomiku01
485830452e fix(agent): exclude DNS resolution from TCP probe latency
Resolve the target hostname before starting the timer so the
measurement reflects pure TCP handshake time only.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 21:21:15 +08:00
xiaomiku01
2fd00cd0b5 feat(agent): use native ICMP sockets with fallback to system ping
Replace the ping-command-only implementation with a three-tier
approach using golang.org/x/net/icmp:

1. Raw socket (ip4:icmp) — works with root or CAP_NET_RAW
2. Unprivileged datagram socket (udp4) — works on Linux/macOS
   without special privileges
3. System ping command — fallback when neither socket works

The method is auto-detected on first probe and cached for all
subsequent calls, avoiding repeated failed attempts.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 21:09:12 +08:00
xiaomiku01
853a294157 fix(ui): add gap detection to probe chart and fix color limit
- Apply appendData() for gap detection in both realtime and non-realtime
  modes, so the latency chart shows breaks instead of smooth lines when
  data is missing during service interruptions
- Handle null stats in gap marker entries to prevent runtime crashes
- Fix color assignment: use CSS variables (--chart-1..5) for ≤5 probes,
  switch to dynamic HSL distribution for >5 probes so all lines are
  visible with distinct colors

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 19:46:03 +08:00
xiaomiku01
aa9ab49654 fix(ui): auto-refresh probe stats when system data updates
Pass system record to NetworkProbes component and use it as a
dependency in the non-realtime fetch effect, matching the pattern
used by system_stats and container_stats in use-system-data.ts.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 18:44:09 +08:00
xiaomiku01
9a5959b57e fix: address network probe code quality issues
- Use shared http.Client in ProbeManager to avoid connection/transport leak
- Skip probe goroutine and agent request when system has no enabled probes
- Validate HTTP probe target URL scheme (http:// or https://) on creation

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 18:40:27 +08:00
xiaomiku01
50f8548479 fix: add migration for network probe collections on existing databases
Existing databases from main branch lack the network_probes and
network_probe_stats collections, which were only in the initial snapshot.
This separate migration ensures they are created on upgrade.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 17:26:58 +08:00
xiaomiku01
bc0581ea61 feat: add network probe data to realtime mode
Include probe results in the 1-second realtime WebSocket broadcast so
the frontend can update probe latency/loss every second, matching the
behavior of system and container metrics.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 11:54:22 +08:00
xiaomiku01
fab5e8a656 fix(ui): filter deleted probes from latency chart stats
Stats records in the DB contain historical data for all probes including
deleted ones. Now filters stats by active probe keys and clears state
when all probes are removed.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 01:21:38 +08:00
xiaomiku01
3a0896e57e fix(ui): address code quality review findings for network probes
- Rename setInterval to setProbeInterval to avoid shadowing global
- Move probeKey function outside component (pure function)
- Fix probes.length dependency to use probes directly
- Use proper type for stats fetch instead of any
- Fix name column fallback to show target instead of dash

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 01:21:38 +08:00
xiaomiku01
7fdc403470 feat(ui): integrate network probes into system detail page
Lazy-load the NetworkProbes component in both default and tabbed
layouts so the probes table and latency chart appear on the system
detail page.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 01:21:38 +08:00
xiaomiku01
e833d44c43 feat(ui): add network probes table and latency chart section
Displays probe list with protocol badges, latency/loss stats, and
delete functionality. Includes a latency line chart using ChartCard
with data sourced from the network-probe-stats API.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 01:21:38 +08:00
xiaomiku01
77dd4bdaf5 feat(ui): add network probe creation dialog
Dialog component for adding ICMP/TCP/HTTP network probes with
protocol selection, target, port, interval, and name fields.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 01:21:38 +08:00
xiaomiku01
ecba63c4bb feat(ui): add NetworkProbeRecord and NetworkProbeStatsRecord types
Add TypeScript interfaces for the network probes feature API responses.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 01:21:38 +08:00
xiaomiku01
f9feaf5343 feat(hub): add network probe API, sync, result collection, and aggregation
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 01:21:38 +08:00
xiaomiku01
ddf5e925c8 feat: add network_probes and network_probe_stats PocketBase collections 2026-04-11 01:21:38 +08:00
xiaomiku01
865e6db90f feat(agent): add ProbeManager with ICMP/TCP/HTTP probes and handlers
Implements the core probe execution engine (ProbeManager) that runs
network probes on configurable intervals, collects latency samples,
and aggregates results over a 60s sliding window. Adds two new
WebSocket handlers (SyncNetworkProbes, GetNetworkProbeResults) for
hub-agent communication and integrates probe lifecycle into the agent.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-11 01:21:38 +08:00
xiaomiku01
a42d899e64 feat: add shared probe entity types (Config, Result) 2026-04-11 01:21:38 +08:00
xiaomiku01
3eaf12a7d5 feat: add SyncNetworkProbes and GetNetworkProbeResults action types 2026-04-11 01:21:38 +08:00
433 changed files with 8579 additions and 66074 deletions

View File

@@ -1,6 +1,6 @@
# Node.js dependencies
node_modules/
**/node_modules/
node_modules
internalsite/node_modules
# Go build artifacts and binaries
build

View File

@@ -1,12 +0,0 @@
version: 2
updates:
- package-ecosystem: gomod
directory: /
schedule:
interval: weekly
- package-ecosystem: github-actions
directory: /
schedule:
interval: weekly

View File

@@ -29,7 +29,6 @@ jobs:
# henrygd/beszel-agent:alpine
- image: henrygd/beszel-agent
dockerfile: ./internal/dockerfile_agent_alpine
flavor: latest=false
registry: docker.io
username_secret: DOCKERHUB_USERNAME
password_secret: DOCKERHUB_TOKEN
@@ -42,7 +41,7 @@ jobs:
# henrygd/beszel-agent-nvidia
- image: henrygd/beszel-agent-nvidia
dockerfile: ./internal/dockerfile_agent_nvidia
platforms: linux/amd64,linux/arm64
platforms: linux/amd64
registry: docker.io
username_secret: DOCKERHUB_USERNAME
password_secret: DOCKERHUB_TOKEN
@@ -53,20 +52,6 @@ jobs:
type=semver,pattern={{major}}
type=raw,value={{sha}},enable=${{ github.ref_type != 'tag' }}
# henrygd/beszel-agent-nvidia:slim
- image: henrygd/beszel-agent-nvidia
dockerfile: ./internal/dockerfile_agent_nvidia_slim
flavor: latest=false
platforms: linux/amd64,linux/arm64
registry: docker.io
username_secret: DOCKERHUB_USERNAME
password_secret: DOCKERHUB_TOKEN
tags: |
type=raw,value=slim
type=semver,pattern={{version}}-slim
type=semver,pattern={{major}}.{{minor}}-slim
type=semver,pattern={{major}}-slim
# henrygd/beszel-agent-intel
- image: henrygd/beszel-agent-intel
dockerfile: ./internal/dockerfile_agent_intel
@@ -111,7 +96,7 @@ jobs:
# ghcr.io/henrygd/beszel-agent-nvidia
- image: ghcr.io/${{ github.repository }}/beszel-agent-nvidia
dockerfile: ./internal/dockerfile_agent_nvidia
platforms: linux/amd64,linux/arm64
platforms: linux/amd64
registry: ghcr.io
username: ${{ github.actor }}
password_secret: GITHUB_TOKEN
@@ -122,20 +107,6 @@ jobs:
type=semver,pattern={{major}}
type=raw,value={{sha}},enable=${{ github.ref_type != 'tag' }}
# ghcr.io/henrygd/beszel-agent-nvidia:slim
- image: ghcr.io/${{ github.repository }}/beszel-agent-nvidia
dockerfile: ./internal/dockerfile_agent_nvidia_slim
flavor: latest=false
platforms: linux/amd64,linux/arm64
registry: ghcr.io
username: ${{ github.actor }}
password_secret: GITHUB_TOKEN
tags: |
type=raw,value=slim
type=semver,pattern={{version}}-slim
type=semver,pattern={{major}}.{{minor}}-slim
type=semver,pattern={{major}}-slim
# ghcr.io/henrygd/beszel-agent-intel
- image: ghcr.io/${{ github.repository }}/beszel-agent-intel
dockerfile: ./internal/dockerfile_agent_intel
@@ -153,7 +124,6 @@ jobs:
# ghcr.io/henrygd/beszel-agent:alpine
- image: ghcr.io/${{ github.repository }}/beszel-agent
dockerfile: ./internal/dockerfile_agent_alpine
flavor: latest=false
registry: ghcr.io
username: ${{ github.actor }}
password_secret: GITHUB_TOKEN
@@ -163,7 +133,7 @@ jobs:
type=semver,pattern={{major}}.{{minor}}-alpine
type=semver,pattern={{major}}-alpine
# henrygd/beszel-agent
# henrygd/beszel-agent (keep at bottom so it gets built after :alpine and gets the latest tag)
- image: henrygd/beszel-agent
dockerfile: ./internal/dockerfile_agent
registry: docker.io
@@ -182,7 +152,7 @@ jobs:
steps:
- name: Checkout
uses: actions/checkout@v7
uses: actions/checkout@v4
- name: Set up bun
uses: oven-sh/setup-bun@v2
@@ -194,18 +164,16 @@ jobs:
run: bun run --cwd ./internal/site build
- name: Set up QEMU
uses: docker/setup-qemu-action@v4
uses: docker/setup-qemu-action@v3
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v4
uses: docker/setup-buildx-action@v3
- name: Docker metadata
id: metadata
uses: docker/metadata-action@v6
uses: docker/metadata-action@v5
with:
images: ${{ matrix.image }}
# Variant images must not overwrite the standard image's latest tag.
flavor: ${{ matrix.flavor || 'latest=auto' }}
tags: ${{ matrix.tags }}
# https://github.com/docker/login-action
@@ -213,7 +181,7 @@ jobs:
env:
password_secret_exists: ${{ secrets[matrix.password_secret] != '' && 'true' || 'false' }}
if: github.event_name != 'pull_request' && env.password_secret_exists == 'true'
uses: docker/login-action@v4
uses: docker/login-action@v3
with:
username: ${{ matrix.username || secrets[matrix.username_secret] }}
password: ${{ secrets[matrix.password_secret] }}
@@ -222,13 +190,11 @@ jobs:
# Build and push Docker image with Buildx (don't push on PR)
# https://github.com/docker/build-push-action
- name: Build and push Docker image
uses: docker/build-push-action@v7
uses: docker/build-push-action@v5
with:
context: ./
file: ${{ matrix.dockerfile }}
platforms: ${{ matrix.platforms || 'linux/amd64,linux/arm64,linux/arm/v6,linux/arm/v7' }}
platforms: ${{ matrix.platforms || 'linux/amd64,linux/arm64,linux/arm/v7' }}
push: ${{ github.ref_type == 'tag' && secrets[matrix.password_secret] != '' }}
provenance: mode=max
sbom: true
tags: ${{ steps.metadata.outputs.tags }}
labels: ${{ steps.metadata.outputs.labels }}

View File

@@ -1,109 +0,0 @@
name: Helm charts
on:
pull_request:
paths:
- "supplemental/helm/**"
push:
branches:
- main
paths:
- "supplemental/helm/**"
permissions:
contents: read
packages: write
env:
OCI_REGISTRY: ghcr.io/henrygd/beszel-charts
jobs:
changes:
name: Detect changed charts
runs-on: ubuntu-latest
outputs:
charts: ${{ steps.changes.outputs.charts }}
steps:
- name: Checkout repository
uses: actions/checkout@v7
with:
fetch-depth: 0
- name: Detect changed charts
id: changes
env:
BASE_SHA: ${{ github.event_name == 'pull_request' && github.event.pull_request.base.sha || github.event.before }}
run: |
charts=()
for name in beszel-agent beszel-hub; do
path="supplemental/helm/$name"
if ! git diff --quiet "$BASE_SHA" "$GITHUB_SHA" -- "$path"; then
charts+=("$name|$path")
fi
done
printf '%s\n' "${charts[@]}" \
| jq -Rsc 'split("\n") | map(select(length > 0) | split("|") | {name: .[0], path: .[1]})' \
| xargs -0 printf 'charts=%s\n' >> "$GITHUB_OUTPUT"
validate-and-publish:
name: ${{ github.event_name == 'push' && 'Publish' || 'Validate' }} ${{ matrix.chart.name }}
needs: changes
if: needs.changes.outputs.charts != '[]'
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
chart: ${{ fromJSON(needs.changes.outputs.charts) }}
steps:
- name: Checkout repository
uses: actions/checkout@v7
- name: Set up Helm
uses: azure/setup-helm@v5
- name: Lint chart
run: helm lint "${{ matrix.chart.path }}" --set env.KEY=ci-placeholder
- name: Render chart
run: helm template "${{ matrix.chart.name }}" "${{ matrix.chart.path }}" --set env.KEY=ci-placeholder > /dev/null
- name: Package chart
id: package
env:
CHART_NAME: ${{ matrix.chart.name }}
CHART_PATH: ${{ matrix.chart.path }}
run: |
version=$(awk '/^version:/ { print $2 }' "$CHART_PATH/Chart.yaml")
test -n "$version"
mkdir -p .helm-packages
helm package "$CHART_PATH" --destination .helm-packages
package=".helm-packages/${CHART_NAME}-${version}.tgz"
test -f "$package"
echo "version=$version" >> "$GITHUB_OUTPUT"
echo "package=$package" >> "$GITHUB_OUTPUT"
- name: Log in to GHCR
env:
GITHUB_TOKEN: ${{ github.token }}
run: echo "$GITHUB_TOKEN" | helm registry login ghcr.io --username "$GITHUB_ACTOR" --password-stdin
- name: Check chart version is unpublished
env:
CHART_NAME: ${{ matrix.chart.name }}
CHART_VERSION: ${{ steps.package.outputs.version }}
run: |
chart="oci://${OCI_REGISTRY}/${CHART_NAME}"
if helm show chart "$chart" --version "$CHART_VERSION" > /dev/null 2>&1; then
echo "${CHART_NAME} ${CHART_VERSION} is already published. Bump version in Chart.yaml." >&2
exit 1
fi
- name: Publish chart
if: github.event_name == 'push'
run: helm push "${{ steps.package.outputs.package }}" "oci://${OCI_REGISTRY}"

View File

@@ -15,7 +15,7 @@ jobs:
name: Lock Inactive Issues
runs-on: ubuntu-24.04
steps:
- uses: klaasnicolaas/action-inactivity-lock@v2.0.1
- uses: klaasnicolaas/action-inactivity-lock@v1.1.3
id: lock
with:
days-inactive-issues: 14
@@ -29,7 +29,7 @@ jobs:
runs-on: ubuntu-24.04
steps:
- name: Close Stale Issues
uses: actions/stale@v11
uses: actions/stale@v10
with:
repo-token: ${{ secrets.GITHUB_TOKEN }}

View File

@@ -13,7 +13,7 @@ jobs:
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@v7
uses: actions/checkout@v4
with:
fetch-depth: 0
@@ -27,12 +27,12 @@ jobs:
run: bun run --cwd ./internal/site build
- name: Set up Go
uses: actions/setup-go@v7
uses: actions/setup-go@v5
with:
go-version: stable
go-version: "^1.22.1"
- name: Set up .NET
uses: actions/setup-dotnet@v6
uses: actions/setup-dotnet@v4
with:
dotnet-version: "9.0.x"
@@ -42,7 +42,7 @@ jobs:
shell: bash
- name: GoReleaser beszel
uses: goreleaser/goreleaser-action@v7
uses: goreleaser/goreleaser-action@v6
with:
workdir: ./
distribution: goreleaser

View File

@@ -1,101 +0,0 @@
name: Update Helm charts
on:
release:
types:
- published
permissions:
contents: write
pull-requests: write
concurrency:
group: update-helm-charts
cancel-in-progress: false
jobs:
update:
name: Propose chart update
if: ${{ github.repository_owner == 'henrygd' && startsWith(github.event.release.tag_name, 'v') && !github.event.release.prerelease }}
runs-on: ubuntu-latest
env:
BRANCH: automation/update-helm-app-version
RELEASE_TAG: ${{ github.event.release.tag_name }}
AUTOMATION_TOKEN: ${{ secrets.CR_TOKEN || github.token }}
steps:
- name: Checkout main
uses: actions/checkout@v7
with:
ref: main
token: ${{ env.AUTOMATION_TOKEN }}
- name: Update chart versions
id: update
run: |
version="${RELEASE_TAG#v}"
if [[ ! "$version" =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]]; then
echo "Unsupported software release version: $version" >&2
exit 1
fi
changed=false
for chart in supplemental/helm/beszel-agent supplemental/helm/beszel-hub; do
current_app_version=$(awk -F '"' '/^appVersion:/ { print $2 }' "$chart/Chart.yaml")
if [[ "$current_app_version" == "$version" ]]; then
echo "$chart already uses appVersion $version"
continue
fi
newest_version=$(printf '%s\n' "$current_app_version" "$version" | sort -V | tail -n 1)
if [[ "$newest_version" != "$version" ]]; then
echo "Skipping stale update of $chart from $current_app_version to $version"
continue
fi
chart_version=$(awk '/^version:/ { print $2 }' "$chart/Chart.yaml")
if [[ ! "$chart_version" =~ ^([0-9]+)\.([0-9]+)\.([0-9]+)$ ]]; then
echo "Unsupported chart version in $chart/Chart.yaml: $chart_version" >&2
exit 1
fi
next_chart_version="${BASH_REMATCH[1]}.${BASH_REMATCH[2]}.$((BASH_REMATCH[3] + 1))"
NEW_APP_VERSION="$version" NEW_CHART_VERSION="$next_chart_version" \
perl -pi -e 's/^appVersion:.*$/appVersion: "$ENV{NEW_APP_VERSION}"/; s/^version:.*$/version: $ENV{NEW_CHART_VERSION}/' \
"$chart/Chart.yaml"
OLD_APP_VERSION="$current_app_version" NEW_APP_VERSION="$version" \
perl -pi -e 's/\Q$ENV{OLD_APP_VERSION}\E/$ENV{NEW_APP_VERSION}/g' "$chart/README.md"
echo "$chart: appVersion $current_app_version -> $version, chart $chart_version -> $next_chart_version"
changed=true
done
echo "changed=$changed" >> "$GITHUB_OUTPUT"
- name: Open or update pull request
if: steps.update.outputs.changed == 'true'
env:
GH_TOKEN: ${{ env.AUTOMATION_TOKEN }}
run: |
version="${RELEASE_TAG#v}"
title="chore(helm): update app version to ${version}"
body="Updates the Helm charts for [Beszel ${version}](${GITHUB_SERVER_URL}/${GITHUB_REPOSITORY}/releases/tag/${RELEASE_TAG}) and bumps their chart patch versions. Merging this pull request publishes the updated charts to GHCR."
git config user.name "github-actions[bot]"
git config user.email "41898282+github-actions[bot]@users.noreply.github.com"
git checkout -B "$BRANCH"
git add supplemental/helm/beszel-agent/Chart.yaml \
supplemental/helm/beszel-agent/README.md \
supplemental/helm/beszel-hub/Chart.yaml \
supplemental/helm/beszel-hub/README.md
git commit -m "$title"
git fetch origin "$BRANCH" || true
git push --force-with-lease origin "HEAD:refs/heads/${BRANCH}"
pr_number=$(gh pr list --head "$BRANCH" --base main --state open --json number --jq '.[0].number')
if [[ -n "$pr_number" ]]; then
gh pr edit "$pr_number" --title "$title" --body "$body"
else
gh pr create --base main --head "$BRANCH" --title "$title" --body "$body"
fi

View File

@@ -2,6 +2,10 @@
name: VulnCheck
on:
pull_request:
branches:
- main
push:
branches:
- main
@@ -15,11 +19,11 @@ jobs:
runs-on: ubuntu-latest
steps:
- name: Check out code into the Go module directory
uses: actions/checkout@v7
uses: actions/checkout@v6
- name: Set up Go
uses: actions/setup-go@v7
uses: actions/setup-go@v6
with:
go-version: stable
go-version: 1.26.x
# cached: false
- name: Get official govulncheck
run: go install golang.org/x/vuln/cmd/govulncheck@latest

3
.gitignore vendored
View File

@@ -3,6 +3,7 @@ pb_data
data
temp
.vscode
beszel-agent
beszel_data
beszel_data*
dist
@@ -20,5 +21,3 @@ __debug_*
agent/lhm/obj
agent/lhm/bin
dockerfile_agent_dev
.cr-release-packages
.tmp

View File

@@ -31,16 +31,12 @@ builds:
goarch: arm64
- goos: freebsd
goarch: arm
- goos: darwin
goarch: arm
- id: beszel-agent
binary: beszel-agent
main: internal/cmd/agent/agent.go
env:
- CGO_ENABLED=0
ldflags:
- -s -w -X github.com/henrygd/beszel/internal/ghupdate.buildGOARM={{ .Arm }}
goos:
- linux
- darwin
@@ -56,10 +52,6 @@ builds:
- mipsle
- mips
- ppc64le
goarm:
- "5"
- "6"
- "7"
gomips:
- hardfloat
- softfloat
@@ -79,8 +71,6 @@ builds:
gomips: hardfloat
- goos: windows
goarch: arm
- goos: darwin
goarch: arm
- goos: darwin
goarch: riscv64
- goos: windows
@@ -107,7 +97,6 @@ archives:
{{ .Binary }}_
{{- .Os }}_
{{- .Arch }}
{{- if ne .Arm "6" }}{{ with .Arm }}v{{ . }}{{ end }}{{ end }}
format_overrides:
- goos: windows
formats: [zip]

View File

@@ -52,7 +52,7 @@ lint:
golangci-lint run
test:
go test -tags='testing no_ui' ./...
go test -tags=testing ./...
tidy:
go mod tidy

View File

@@ -2,8 +2,6 @@
## Reporting a Vulnerability
**PLEASE ONLY USE SECURITY ADVISORIES FOR REAL HIGH SEVERITY VULNERABILITIES.**
If you find a vulnerability in the latest version, please [submit a private advisory](https://github.com/henrygd/beszel/security/advisories/new).
If you find a vulnerability in the latest version, and it is not high severity, open an issue instead of an advisory.
I am overwhelmed with advisories, often erroneous, which are clearly found and written by AI. I don't have the capacity to review all of them.
If it's low severity (use best judgement) you may open an issue instead of an advisory.

View File

@@ -6,7 +6,6 @@ package agent
import (
"log/slog"
"net"
"strings"
"sync"
"time"
@@ -30,10 +29,9 @@ type Agent struct {
fsNames []string // List of filesystem device names being monitored
fsStats map[string]*system.FsStats // Keeps track of disk stats for each filesystem
diskPrev map[uint16]map[string]prevDisk // Previous disk I/O counters per cache interval
diskBaseline map[string]prevDisk // Latest disk I/O counters of any interval, seeds a new interval
diskUsageCacheDuration time.Duration // How long to cache disk usage (to avoid waking sleeping disks)
lastDiskUsageUpdate time.Time // Last time disk usage was collected
netInterfaces map[string]bool // Valid network interfaces; true if byte counters come from MAC stats (Jetson nvethernet)
netInterfaces map[string]struct{} // Stores all valid network interfaces
netIoStats map[uint16]system.NetIoStats // Keeps track of bandwidth usage per cache interval
netInterfaceDeltaTrackers map[uint16]*deltatracker.DeltaTracker[string, uint64] // Per-cache-time NIC delta trackers
dockerManager *dockerManager // Manages Docker API requests
@@ -46,15 +44,11 @@ type Agent struct {
connectionManager *ConnectionManager // Channel to signal connection events
handlerRegistry *HandlerRegistry // Registry for routing incoming messages
server *ssh.Server // SSH server
serverListener net.Listener // SSH listener, also closed if Serve has not started yet
serverMu sync.Mutex // Guards server and serverListener
dataDir string // Directory for persisting data
keys []gossh.PublicKey // SSH public keys
smartManager *SmartManager // Manages SMART data
systemdManager *systemdManager // Manages systemd services
monitorManager *MonitorManager // Manages network monitors
storagePoolManager *StoragePoolManager // Manages storage pool and dataset data
packageUpdates *packageUpdatesManager // Checks for pending package updates
probeManager *ProbeManager // Manages network probes
}
// NewAgent creates a new agent with the given data directory for persisting data.
@@ -128,21 +122,8 @@ func NewAgent(dataDir ...string) (agent *Agent, err error) {
// initialize handler registry
agent.handlerRegistry = NewHandlerRegistry()
// initialize monitor manager
agent.monitorManager = newMonitorManager()
agent.storagePoolManager = newStoragePoolManager()
// Retain ZFS_INTERVAL for the shared storage pool detail refresh interval.
if zfsIntervalEnv, exists := utils.GetEnv("ZFS_INTERVAL"); exists {
if duration, err := time.ParseDuration(zfsIntervalEnv); err == nil && duration > 0 {
agent.storagePoolManager.detailInterval = duration
agent.systemDetails.ZfsInterval = duration
slog.Info("ZFS_INTERVAL", "duration", duration)
} else {
slog.Warn("Invalid ZFS_INTERVAL", "err", err)
}
}
// initialize probe manager
agent.probeManager = newProbeManager()
// initialize disk info
agent.initializeDiskInfo()
@@ -154,17 +135,12 @@ func NewAgent(dataDir ...string) (agent *Agent, err error) {
if err != nil {
slog.Debug("Systemd", "err", err)
}
if agent.systemdManager != nil {
agent.systemInfo.SystemdLogs = agent.systemdManager.logsEnabled
}
agent.smartManager, err = NewSmartManager()
if err != nil {
slog.Debug("SMART", "err", err)
}
agent.packageUpdates = newPackageUpdatesManager(agent.dataDir)
// initialize GPU manager
agent.gpuManager, err = NewGPUManager()
if err != nil {
@@ -206,9 +182,9 @@ func (a *Agent) gatherStats(options common.DataRequestOptions) *system.CombinedD
}
}
if a.monitorManager != nil {
data.Monitors = a.monitorManager.GetResults(cacheTimeMs)
slog.Debug("Monitors", "data", data.Monitors)
if a.probeManager != nil {
data.Probes = a.probeManager.GetResults(cacheTimeMs)
slog.Debug("Probes", "data", data.Probes)
}
// skip updating systemd services if cache time is not the default 60sec interval
@@ -220,29 +196,13 @@ func (a *Agent) gatherStats(options common.DataRequestOptions) *system.CombinedD
}
if a.systemdManager.hasFreshStats {
data.SystemdServices = a.systemdManager.getServiceStats(nil, false)
data.SystemdServicesUpdated = true
// Preserve an explicit zero count so the hub can distinguish a fresh
// empty snapshot from a response that omitted systemd data.
if totalCount == 0 {
data.Info.Services = []uint16{0, 0}
}
}
}
if a.packageUpdates != nil {
data.Info.PackageUpdates = a.packageUpdates.get(time.Now())
}
data.Stats.ExtraFs = make(map[string]*system.FsStats)
data.Info.ExtraFsPct = make(map[string]float64)
for name, stats := range a.fsStats {
if stats.Root {
if stats.Name != "" {
data.Info.RootDiskName = stats.Name
}
continue
}
if stats.DiskTotal > 0 {
if !stats.Root && stats.DiskTotal > 0 {
// Use custom name if available, otherwise use device name
key := name
if stats.Name != "" {
@@ -266,11 +226,7 @@ func (a *Agent) gatherStats(options common.DataRequestOptions) *system.CombinedD
// Start initializes and starts the agent with optional WebSocket connection
func (a *Agent) Start(serverOptions ServerOptions) error {
a.keys = serverOptions.Keys
err := a.connectionManager.Start(serverOptions)
if err != nil {
a.cleanupSensorShadow()
}
return err
return a.connectionManager.Start(serverOptions)
}
func (a *Agent) getFingerprint() string {

View File

@@ -1,13 +1,6 @@
// Package battery provides battery information for the host and connected devices.
// Package battery provides functions to check if the system has a battery and return the charge state and percentage.
package battery
import (
"errors"
"sort"
"strconv"
"strings"
)
const (
stateUnknown uint8 = iota
stateEmpty
@@ -16,58 +9,3 @@ const (
stateDischarging
stateIdle
)
// Battery is a readable battery reported by the operating system.
type Battery struct {
Name string
Percent uint8
State uint8
FullChargeCapacity uint64
HasFullChargeCapacity bool
System bool
}
var errNoBatteries = errors.New("no readable batteries")
// normalizeBatteries supplies stable fallback names and disambiguates duplicates.
func normalizeBatteries(batteries []Battery) []Battery {
nameCounts := make(map[string]int, len(batteries))
for i := range batteries {
// Names come from firmware (e.g. sysfs model_name) and are not guaranteed to
// be valid UTF-8. Invalid bytes are rejected when the hub decodes the CBOR
// payload, which drops every metric for the system, so strip them here.
name := strings.TrimSpace(strings.ToValidUTF8(batteries[i].Name, ""))
if name == "" {
name = "Battery " + strconv.Itoa(i+1)
}
nameCounts[name]++
if nameCounts[name] > 1 {
name += " (" + strconv.Itoa(nameCounts[name]) + ")"
}
batteries[i].Name = name
}
return batteries
}
// Primary returns the representative battery. Reported full-charge capacity wins,
// then system-scoped devices, then name for deterministic ties.
func Primary(batteries []Battery) (Battery, bool) {
if len(batteries) == 0 {
return Battery{}, false
}
ordered := append([]Battery(nil), batteries...)
sort.SliceStable(ordered, func(i, j int) bool {
a, b := ordered[i], ordered[j]
if a.HasFullChargeCapacity != b.HasFullChargeCapacity {
return a.HasFullChargeCapacity
}
if a.HasFullChargeCapacity && a.FullChargeCapacity != b.FullChargeCapacity {
return a.FullChargeCapacity > b.FullChargeCapacity
}
if a.System != b.System {
return a.System
}
return a.Name < b.Name
})
return ordered[0], true
}

View File

@@ -3,7 +3,11 @@
package battery
import (
"errors"
"log/slog"
"math"
"os/exec"
"sync"
"howett.net/plist"
)
@@ -31,46 +35,62 @@ func readMacBatteries() ([]macBattery, error) {
return batteries, nil
}
func HasReadableBattery() bool {
batteries, _ := GetBatteryStats()
return len(batteries) > 0
}
// GetBatteryStats returns every readable battery reported by macOS.
func GetBatteryStats() ([]Battery, error) {
// HasReadableBattery checks if the system has a battery and returns true if it does.
var HasReadableBattery = sync.OnceValue(func() bool {
systemHasBattery := false
batteries, err := readMacBatteries()
if err != nil {
return nil, err
}
if len(batteries) == 0 {
return nil, errNoBatteries
}
result := make([]Battery, 0, len(batteries))
slog.Debug("Batteries", "batteries", batteries, "err", err)
for _, bat := range batteries {
if bat.MaxCapacity <= 0 {
if bat.MaxCapacity > 0 {
systemHasBattery = true
break
}
}
return systemHasBattery
})
// GetBatteryStats returns the current battery percent and charge state.
// Uses CurrentCapacity/MaxCapacity to match the value macOS displays.
func GetBatteryStats() (batteryPercent uint8, batteryState uint8, err error) {
if !HasReadableBattery() {
return batteryPercent, batteryState, errors.ErrUnsupported
}
batteries, err := readMacBatteries()
if len(batteries) == 0 {
return batteryPercent, batteryState, errors.New("no batteries")
}
totalCapacity := 0
totalCharge := 0
batteryState = math.MaxUint8
for _, bat := range batteries {
if bat.MaxCapacity == 0 {
// skip ghost batteries with 0 capacity
// https://github.com/distatus/battery/issues/34
continue
}
percent := min(max(float64(bat.CurrentCapacity)/float64(bat.MaxCapacity)*100, 0), 100)
state := stateUnknown
totalCapacity += bat.MaxCapacity
totalCharge += min(bat.CurrentCapacity, bat.MaxCapacity)
switch {
case !bat.ExternalConnected:
state = stateDischarging
batteryState = stateDischarging
case bat.IsCharging:
state = stateCharging
batteryState = stateCharging
case bat.CurrentCapacity == 0:
state = stateEmpty
batteryState = stateEmpty
case !bat.FullyCharged:
state = stateIdle
batteryState = stateIdle
default:
state = stateFull
batteryState = stateFull
}
result = append(result, Battery{Name: "Primary", Percent: uint8(percent), State: state,
FullChargeCapacity: uint64(bat.MaxCapacity), HasFullChargeCapacity: true, System: true})
}
if len(result) == 0 {
return nil, errNoBatteries
if totalCapacity == 0 || batteryState == math.MaxUint8 {
return batteryPercent, batteryState, errors.New("no battery capacity")
}
return normalizeBatteries(result), nil
batteryPercent = uint8(float64(totalCharge) / float64(totalCapacity) * 100)
return batteryPercent, batteryState, nil
}

View File

@@ -3,19 +3,58 @@
package battery
import (
"errors"
"log/slog"
"math"
"os"
"path/filepath"
"strconv"
"sync"
"github.com/henrygd/beszel/agent/utils"
)
var batteryRoot = "/sys/class/power_supply"
// getBatteryPaths returns the paths of all batteries in /sys/class/power_supply
var getBatteryPaths func() ([]string, error)
// HasReadableBattery reports whether collection currently finds a readable battery.
func HasReadableBattery() bool {
batteries, _ := GetBatteryStats()
return len(batteries) > 0
// HasReadableBattery checks if the system has a battery and returns true if it does.
var HasReadableBattery func() bool
func init() {
resetBatteryState("/sys/class/power_supply")
}
// resetBatteryState resets the sync.Once functions to a fresh state.
// Tests call this after swapping sysfsPowerSupply so the new path is picked up.
func resetBatteryState(sysfsPowerSupplyPath string) {
getBatteryPaths = sync.OnceValues(func() ([]string, error) {
entries, err := os.ReadDir(sysfsPowerSupplyPath)
if err != nil {
return nil, err
}
var paths []string
for _, e := range entries {
path := filepath.Join(sysfsPowerSupplyPath, e.Name())
if utils.ReadStringFile(filepath.Join(path, "type")) == "Battery" {
paths = append(paths, path)
}
}
return paths, nil
})
HasReadableBattery = sync.OnceValue(func() bool {
systemHasBattery := false
paths, err := getBatteryPaths()
for _, path := range paths {
if _, ok := utils.ReadStringFileOK(filepath.Join(path, "capacity")); ok {
systemHasBattery = true
break
}
}
if !systemHasBattery {
slog.Debug("No battery found", "err", err)
}
return systemHasBattery
})
}
func parseSysfsState(status string) uint8 {
@@ -35,18 +74,26 @@ func parseSysfsState(status string) uint8 {
}
}
// GetBatteryStats re-enumerates power supplies and returns every readable battery.
func GetBatteryStats() ([]Battery, error) {
entries, err := os.ReadDir(batteryRoot)
if err != nil {
return nil, err
// GetBatteryStats returns the current battery percent and charge state.
// Reads /sys/class/power_supply/*/capacity directly so the kernel-reported
// value is used, which is always 0-100 and matches what the OS displays.
func GetBatteryStats() (batteryPercent uint8, batteryState uint8, err error) {
if !HasReadableBattery() {
return batteryPercent, batteryState, errors.ErrUnsupported
}
batteries := make([]Battery, 0, len(entries))
for _, entry := range entries {
path := filepath.Join(batteryRoot, entry.Name())
if utils.ReadStringFile(filepath.Join(path, "type")) != "Battery" {
continue
}
paths, err := getBatteryPaths()
if err != nil {
return batteryPercent, batteryState, err
}
if len(paths) == 0 {
return batteryPercent, batteryState, errors.New("no batteries")
}
batteryState = math.MaxUint8
totalPercent := 0
count := 0
for _, path := range paths {
capStr, ok := utils.ReadStringFileOK(filepath.Join(path, "capacity"))
if !ok {
continue
@@ -55,31 +102,19 @@ func GetBatteryStats() ([]Battery, error) {
if parseErr != nil {
continue
}
cap = min(max(cap, 0), 100)
name := utils.ReadStringFile(filepath.Join(path, "model_name"))
if name == "" {
name = utils.ReadStringFile(filepath.Join(path, "model"))
totalPercent += cap
count++
state := parseSysfsState(utils.ReadStringFile(filepath.Join(path, "status")))
if state != stateUnknown {
batteryState = state
}
if name == "" {
name = entry.Name()
}
battery := Battery{
Name: name,
Percent: uint8(cap),
State: parseSysfsState(utils.ReadStringFile(filepath.Join(path, "status"))),
System: utils.ReadStringFile(filepath.Join(path, "scope")) != "Device",
}
for _, fullName := range []string{"charge_full", "energy_full"} {
if parsed, ok := utils.ReadUintFile(filepath.Join(path, fullName)); ok && parsed > 0 {
battery.FullChargeCapacity = parsed
battery.HasFullChargeCapacity = true
break
}
}
batteries = append(batteries, battery)
}
if len(batteries) == 0 {
return nil, errNoBatteries
if count == 0 || batteryState == math.MaxUint8 {
return batteryPercent, batteryState, errors.New("no battery capacity")
}
return normalizeBatteries(batteries), nil
batteryPercent = uint8(totalPercent / count)
return batteryPercent, batteryState, nil
}

View File

@@ -8,102 +8,194 @@ import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type fakeBattery struct{ id, name, capacity, status, full, scope string }
func setupFakeSysfs(t *testing.T) (string, func(fakeBattery)) {
// setupFakeSysfs creates a temporary sysfs-like tree under t.TempDir(),
// swaps sysfsPowerSupply, resets the sync.Once caches, and restores
// everything on cleanup. Returns a helper to create battery directories.
func setupFakeSysfs(t *testing.T) (tmpDir string, addBattery func(name, capacity, status string)) {
t.Helper()
root := t.TempDir()
previousRoot := batteryRoot
batteryRoot = root
t.Cleanup(func() { batteryRoot = previousRoot })
write := func(path, value string) {
tmp := t.TempDir()
resetBatteryState(tmp)
write := func(path, content string) {
t.Helper()
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o755))
require.NoError(t, os.WriteFile(path, []byte(value), 0o644))
}
add := func(b fakeBattery) {
t.Helper()
dir := filepath.Join(root, b.id)
write(filepath.Join(dir, "type"), "Battery")
if b.capacity != "" {
write(filepath.Join(dir, "capacity"), b.capacity)
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o755); err != nil {
t.Fatal(err)
}
write(filepath.Join(dir, "status"), b.status)
if b.name != "" {
write(filepath.Join(dir, "model_name"), b.name)
}
if b.full != "" {
write(filepath.Join(dir, "energy_full"), b.full)
}
if b.scope != "" {
write(filepath.Join(dir, "scope"), b.scope)
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
return root, add
addBattery = func(name, capacity, status string) {
t.Helper()
batDir := filepath.Join(tmp, name)
write(filepath.Join(batDir, "type"), "Battery")
write(filepath.Join(batDir, "capacity"), capacity)
write(filepath.Join(batDir, "status"), status)
}
return tmp, addBattery
}
func TestParseSysfsState(t *testing.T) {
assert.Equal(t, stateEmpty, parseSysfsState("Empty"))
assert.Equal(t, stateFull, parseSysfsState("Full"))
assert.Equal(t, stateCharging, parseSysfsState("Charging"))
assert.Equal(t, stateDischarging, parseSysfsState("Discharging"))
assert.Equal(t, stateIdle, parseSysfsState("Not charging"))
assert.Equal(t, stateUnknown, parseSysfsState("SomethingElse"))
tests := []struct {
input string
want uint8
}{
{"Empty", stateEmpty},
{"Full", stateFull},
{"Charging", stateCharging},
{"Discharging", stateDischarging},
{"Not charging", stateIdle},
{"", stateUnknown},
{"SomethingElse", stateUnknown},
}
for _, tt := range tests {
assert.Equal(t, tt.want, parseSysfsState(tt.input), "parseSysfsState(%q)", tt.input)
}
}
func TestGetBatteryStatsMultipleNamedAndPrimary(t *testing.T) {
_, add := setupFakeSysfs(t)
add(fakeBattery{id: "BAT0", name: "Primary", capacity: "105", status: "Charging", full: "5000", scope: "System"})
add(fakeBattery{id: "hidpp_battery_0", name: "MX Keys S", capacity: "55", status: "Unknown", full: "900", scope: "Device"})
batteries, err := GetBatteryStats()
require.NoError(t, err)
require.Len(t, batteries, 2)
assert.Equal(t, "Primary", batteries[0].Name)
assert.Equal(t, uint8(100), batteries[0].Percent)
assert.Equal(t, stateUnknown, batteries[1].State)
primary, ok := Primary(batteries)
require.True(t, ok)
assert.Equal(t, "Primary", primary.Name)
func TestGetBatteryStats_SingleBattery(t *testing.T) {
_, addBattery := setupFakeSysfs(t)
addBattery("BAT0", "72", "Discharging")
pct, state, err := GetBatteryStats()
assert.NoError(t, err)
assert.Equal(t, uint8(72), pct)
assert.Equal(t, stateDischarging, state)
}
func TestGetBatteryStatsFallbackDuplicatesAndUnreadable(t *testing.T) {
root, add := setupFakeSysfs(t)
add(fakeBattery{id: "BAT0", name: "Keyboard", capacity: "80", status: "Discharging"})
add(fakeBattery{id: "BAT1", name: "Keyboard", capacity: "-4", status: "SomethingWeird"})
add(fakeBattery{id: "BAT2", capacity: "not-a-number", status: "Charging"})
add(fakeBattery{id: "BAT3", capacity: "42", status: "Full"})
ac := filepath.Join(root, "AC0")
require.NoError(t, os.MkdirAll(ac, 0o755))
require.NoError(t, os.WriteFile(filepath.Join(ac, "type"), []byte("Mains"), 0o644))
batteries, err := GetBatteryStats()
require.NoError(t, err)
require.Len(t, batteries, 3)
assert.Equal(t, "Keyboard", batteries[0].Name)
assert.Equal(t, "Keyboard (2)", batteries[1].Name)
assert.Equal(t, uint8(0), batteries[1].Percent)
assert.Equal(t, "BAT3", batteries[2].Name)
func TestGetBatteryStats_MultipleBatteries(t *testing.T) {
_, addBattery := setupFakeSysfs(t)
addBattery("BAT0", "80", "Charging")
addBattery("BAT1", "40", "Charging")
pct, state, err := GetBatteryStats()
assert.NoError(t, err)
// average of 80 and 40 = 60
assert.EqualValues(t, 60, pct)
assert.Equal(t, stateCharging, state)
}
func TestGetBatteryStatsHotPlugReenumerates(t *testing.T) {
_, add := setupFakeSysfs(t)
_, err := GetBatteryStats()
func TestGetBatteryStats_FullBattery(t *testing.T) {
_, addBattery := setupFakeSysfs(t)
addBattery("BAT0", "100", "Full")
pct, state, err := GetBatteryStats()
assert.NoError(t, err)
assert.Equal(t, uint8(100), pct)
assert.Equal(t, stateFull, state)
}
func TestGetBatteryStats_EmptyBattery(t *testing.T) {
_, addBattery := setupFakeSysfs(t)
addBattery("BAT0", "0", "Empty")
pct, state, err := GetBatteryStats()
assert.NoError(t, err)
assert.Equal(t, uint8(0), pct)
assert.Equal(t, stateEmpty, state)
}
func TestGetBatteryStats_NotCharging(t *testing.T) {
_, addBattery := setupFakeSysfs(t)
addBattery("BAT0", "80", "Not charging")
pct, state, err := GetBatteryStats()
assert.NoError(t, err)
assert.Equal(t, uint8(80), pct)
assert.Equal(t, stateIdle, state)
}
func TestGetBatteryStats_NoBatteries(t *testing.T) {
setupFakeSysfs(t) // empty directory, no batteries
_, _, err := GetBatteryStats()
assert.Error(t, err)
assert.False(t, HasReadableBattery())
add(fakeBattery{id: "BAT0", capacity: "64", status: "Discharging"})
batteries, err := GetBatteryStats()
require.NoError(t, err)
}
func TestGetBatteryStats_NonBatterySupplyIgnored(t *testing.T) {
tmp, addBattery := setupFakeSysfs(t)
// Add a real battery
addBattery("BAT0", "55", "Charging")
// Add an AC adapter (type != Battery) - should be ignored
acDir := filepath.Join(tmp, "AC0")
if err := os.MkdirAll(acDir, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(acDir, "type"), []byte("Mains"), 0o644); err != nil {
t.Fatal(err)
}
pct, state, err := GetBatteryStats()
assert.NoError(t, err)
assert.Equal(t, uint8(55), pct)
assert.Equal(t, stateCharging, state)
}
func TestGetBatteryStats_InvalidCapacitySkipped(t *testing.T) {
tmp, addBattery := setupFakeSysfs(t)
// One battery with valid capacity
addBattery("BAT0", "90", "Discharging")
// Another with invalid capacity text
badDir := filepath.Join(tmp, "BAT1")
if err := os.MkdirAll(badDir, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(badDir, "type"), []byte("Battery"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(badDir, "capacity"), []byte("not-a-number"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(badDir, "status"), []byte("Discharging"), 0o644); err != nil {
t.Fatal(err)
}
pct, _, err := GetBatteryStats()
assert.NoError(t, err)
// Only BAT0 counted
assert.Equal(t, uint8(90), pct)
}
func TestGetBatteryStats_UnknownStatusOnly(t *testing.T) {
_, addBattery := setupFakeSysfs(t)
addBattery("BAT0", "50", "SomethingWeird")
_, _, err := GetBatteryStats()
assert.Error(t, err)
}
func TestHasReadableBattery_True(t *testing.T) {
_, addBattery := setupFakeSysfs(t)
addBattery("BAT0", "50", "Charging")
assert.True(t, HasReadableBattery())
require.Len(t, batteries, 1)
assert.Equal(t, uint8(64), batteries[0].Percent)
}
func TestGetBatteryStatsNoReadableCapacity(t *testing.T) {
_, add := setupFakeSysfs(t)
add(fakeBattery{id: "BAT0", status: "Charging"})
_, err := GetBatteryStats()
assert.Error(t, err)
func TestHasReadableBattery_False(t *testing.T) {
setupFakeSysfs(t) // no batteries
assert.False(t, HasReadableBattery())
}
func TestHasReadableBattery_NoCapacityFile(t *testing.T) {
tmp, _ := setupFakeSysfs(t)
// Battery dir with type file but no capacity file
batDir := filepath.Join(tmp, "BAT0")
err := os.MkdirAll(batDir, 0o755)
assert.NoError(t, err)
err = os.WriteFile(filepath.Join(batDir, "type"), []byte("Battery"), 0o644)
assert.NoError(t, err)
assert.False(t, HasReadableBattery())
}

View File

@@ -8,6 +8,6 @@ func HasReadableBattery() bool {
return false
}
func GetBatteryStats() ([]Battery, error) {
return nil, errors.ErrUnsupported
func GetBatteryStats() (uint8, uint8, error) {
return 0, 0, errors.ErrUnsupported
}

View File

@@ -1,48 +0,0 @@
package battery
import (
"testing"
"unicode/utf8"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestPrimarySelection(t *testing.T) {
tests := []struct {
name string
bats []Battery
want string
}{
{"largest reported capacity", []Battery{{Name: "Small", FullChargeCapacity: 20, HasFullChargeCapacity: true, System: true}, {Name: "Large", FullChargeCapacity: 80, HasFullChargeCapacity: true}}, "Large"},
{"reported ranks over missing", []Battery{{Name: "Unknown", System: true}, {Name: "Known", FullChargeCapacity: 1, HasFullChargeCapacity: true}}, "Known"},
{"system wins capacity tie", []Battery{{Name: "Peripheral", FullChargeCapacity: 50, HasFullChargeCapacity: true}, {Name: "System", FullChargeCapacity: 50, HasFullChargeCapacity: true, System: true}}, "System"},
{"name resolves final tie", []Battery{{Name: "Zed"}, {Name: "Alpha"}}, "Alpha"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, ok := Primary(tt.bats)
require.True(t, ok)
assert.Equal(t, tt.want, got.Name)
})
}
_, ok := Primary(nil)
assert.False(t, ok)
}
func TestNormalizeBatteriesFallbackNames(t *testing.T) {
bats := normalizeBatteries([]Battery{{}, {}, {Name: "Mouse"}, {Name: "Mouse"}})
assert.Equal(t, []string{"Battery 1", "Battery 2", "Mouse", "Mouse (2)"}, []string{bats[0].Name, bats[1].Name, bats[2].Name, bats[3].Name})
}
func TestNormalizeBatteriesStripsInvalidUTF8(t *testing.T) {
// Firmware occasionally reports names that are not valid UTF-8 (a ThinkPad
// reporting "LNV-5B11K63024@\xd0" in model_name is a real example).
bats := normalizeBatteries([]Battery{{Name: "LNV-5B11K63024@\xd0"}, {Name: "\xff\xfe"}})
assert.Equal(t, "LNV-5B11K63024@", bats[0].Name)
// A name made up entirely of invalid bytes falls back to the generic name.
assert.Equal(t, "Battery 2", bats[1].Name)
for _, b := range bats {
assert.True(t, utf8.ValidString(b.Name))
}
}

View File

@@ -7,6 +7,9 @@ package battery
import (
"errors"
"log/slog"
"math"
"sync"
"syscall"
"unsafe"
@@ -76,7 +79,7 @@ var (
setupDiDestroyDeviceInfoList = setupapi.NewProc("SetupDiDestroyDeviceInfoList")
)
// winBatteryGet reads one battery by index.
// winBatteryGet reads one battery by index. Returns (fullCapacity, currentCapacity, state, error).
// Returns error == errNotFound when there are no more batteries.
var errNotFound = errors.New("no more batteries")
@@ -119,7 +122,7 @@ func readWinBatteryState(powerState uint32) uint8 {
}
}
func winBatteryGet(idx int) (Battery, error) {
func winBatteryGet(idx int) (full, current uint32, state uint8, err error) {
hdev, err := setupDiSetup(
setupDiGetClassDevsW,
4,
@@ -129,7 +132,7 @@ func winBatteryGet(idx int) (Battery, error) {
0, 0,
)
if err != nil {
return Battery{}, err
return 0, 0, stateUnknown, err
}
defer syscall.SyscallN(setupDiDestroyDeviceInfoList.Addr(), hdev)
@@ -145,10 +148,10 @@ func winBatteryGet(idx int) (Battery, error) {
0,
)
if errno == 259 { // ERROR_NO_MORE_ITEMS
return Battery{}, errNotFound
return 0, 0, stateUnknown, errNotFound
}
if errno != 0 {
return Battery{}, errno
return 0, 0, stateUnknown, errno
}
var cbRequired uint32
@@ -162,7 +165,7 @@ func winBatteryGet(idx int) (Battery, error) {
0,
)
if errno != 0 && errno != 122 { // ERROR_INSUFFICIENT_BUFFER
return Battery{}, errno
return 0, 0, stateUnknown, errno
}
didd := make([]uint16, cbRequired/2)
cbSize := (*uint32)(unsafe.Pointer(&didd[0]))
@@ -182,7 +185,7 @@ func winBatteryGet(idx int) (Battery, error) {
0,
)
if errno != 0 {
return Battery{}, errno
return 0, 0, stateUnknown, errno
}
devicePath := &didd[2:][0]
@@ -196,7 +199,7 @@ func winBatteryGet(idx int) (Battery, error) {
0,
)
if err != nil {
return Battery{}, err
return 0, 0, stateUnknown, err
}
defer windows.CloseHandle(handle)
@@ -213,7 +216,7 @@ func winBatteryGet(idx int) (Battery, error) {
&dwOut, nil,
)
if err != nil || bqi.BatteryTag == 0 {
return Battery{}, errors.New("battery tag not returned")
return 0, 0, stateUnknown, errors.New("battery tag not returned")
}
var bi batteryInformation
@@ -226,21 +229,7 @@ func winBatteryGet(idx int) (Battery, error) {
uint32(unsafe.Sizeof(bi)),
&dwOut, nil,
); err != nil {
return Battery{}, err
}
// BatteryDeviceName is optional, so retain the deterministic fallback on error.
name := ""
nameQuery := bqi
nameQuery.InformationLevel = 4 // BatteryDeviceName
nameBuffer := make([]uint16, 128)
if err := windows.DeviceIoControl(
handle, 2703428,
(*byte)(unsafe.Pointer(&nameQuery)), uint32(unsafe.Sizeof(nameQuery)),
(*byte)(unsafe.Pointer(&nameBuffer[0])), uint32(len(nameBuffer)*2),
&dwOut, nil,
); err == nil {
name = windows.UTF16ToString(nameBuffer)
return 0, 0, stateUnknown, err
}
bws := batteryWaitStatus{BatteryTag: bqi.BatteryTag}
@@ -254,38 +243,56 @@ func winBatteryGet(idx int) (Battery, error) {
uint32(unsafe.Sizeof(bs)),
&dwOut, nil,
); err != nil {
return Battery{}, err
return 0, 0, stateUnknown, err
}
if bs.Capacity == 0xffffffff || bi.FullChargedCapacity == 0 || bi.FullChargedCapacity == 0xffffffff {
return Battery{}, errors.New("battery capacity unknown")
if bs.Capacity == 0xffffffff { // BATTERY_UNKNOWN_CAPACITY
return 0, 0, stateUnknown, errors.New("battery capacity unknown")
}
percent := min(float64(bs.Capacity)/float64(bi.FullChargedCapacity)*100, 100)
return Battery{Name: name, Percent: uint8(percent), State: readWinBatteryState(bs.PowerState),
FullChargeCapacity: uint64(bi.FullChargedCapacity), HasFullChargeCapacity: true, System: true}, nil
return bi.FullChargedCapacity, bs.Capacity, readWinBatteryState(bs.PowerState), nil
}
// HasReadableBattery checks if the system has a battery and returns true if it does.
func HasReadableBattery() bool {
batteries, _ := GetBatteryStats()
return len(batteries) > 0
}
var HasReadableBattery = sync.OnceValue(func() bool {
systemHasBattery := false
full, _, _, err := winBatteryGet(0)
if err == nil && full > 0 {
systemHasBattery = true
}
if !systemHasBattery {
slog.Debug("No battery found", "err", err)
}
return systemHasBattery
})
// GetBatteryStats returns the current battery percent and charge state.
func GetBatteryStats() (batteryPercent uint8, batteryState uint8, err error) {
if !HasReadableBattery() {
return batteryPercent, batteryState, errors.ErrUnsupported
}
totalFull := uint32(0)
totalCurrent := uint32(0)
batteryState = math.MaxUint8
// GetBatteryStats returns every readable battery reported by Windows.
func GetBatteryStats() ([]Battery, error) {
batteries := make([]Battery, 0, 2)
for i := 0; ; i++ {
battery, bErr := winBatteryGet(i)
full, current, state, bErr := winBatteryGet(i)
if errors.Is(bErr, errNotFound) {
break
}
if bErr != nil {
if bErr != nil || full == 0 {
continue
}
batteries = append(batteries, battery)
totalFull += full
totalCurrent += min(current, full)
batteryState = state
}
if len(batteries) == 0 {
return nil, errNoBatteries
if totalFull == 0 || batteryState == math.MaxUint8 {
return batteryPercent, batteryState, errors.New("no battery capacity")
}
return normalizeBatteries(batteries), nil
batteryPercent = uint8(float64(totalCurrent) / float64(totalFull) * 100)
return batteryPercent, batteryState, nil
}

View File

@@ -1,26 +0,0 @@
// Package btrfs reads btrfs filesystem state from sysfs.
package btrfs
// Filesystem is a mounted btrfs filesystem read from /sys/fs/btrfs/<uuid>.
type Filesystem struct {
UUID string // stable filesystem UUID from sysfs
MountID string // kernel filesystem identity for matching monitored mounts
IODevice string // sole member block-device name, empty for multi-device/unknown pools
Name string // label, else first mountpoint, else UUID
Size uint64 // effective usable capacity, or raw member capacity when Raw
Raw bool // capacity and usage are physical bytes, unsuitable for disk alerts
Alloc uint64 // raw bytes allocated to data, metadata and system chunks
Health string // ONLINE, or DEGRADED when a device is missing
NRead uint64 // cumulative bytes read across member devices
NWrite uint64 // cumulative bytes written across member devices
Devices []Device
}
// Device is one member device (devinfo/<devid>) with its error counters.
type Device struct {
Name string // "devid N"; sysfs does not expose the block device path
State string // ONLINE or MISSING
ReadErrs uint64
WriteErrs uint64
CorruptionErrs uint64
}

View File

@@ -1,285 +0,0 @@
//go:build linux
package btrfs
import (
"errors"
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
"unsafe"
"github.com/henrygd/beszel/agent/utils"
"golang.org/x/sys/unix"
)
var (
sysfsPath = "/sys/fs/btrfs"
mountsPath = "/proc/self/mounts"
mountinfoPath = "/proc/self/mountinfo"
mountUUID = MountID
deviceSize = ioctlDeviceSize
filesystemUsage = statfsUsage
)
// Filesystems returns all mounted btrfs filesystems, or nil when there are none.
func Filesystems() ([]Filesystem, error) {
entries, err := os.ReadDir(sysfsPath)
if errors.Is(err, os.ErrNotExist) {
return nil, nil
}
if err != nil {
return nil, err
}
mounts := mountpointsByDevice()
var filesystems []Filesystem
for _, entry := range entries {
if !entry.IsDir() || entry.Name() == "features" {
continue
}
fs, err := readFilesystem(filepath.Join(sysfsPath, entry.Name()), mounts)
if err != nil {
return nil, fmt.Errorf("btrfs %s: %w", entry.Name(), err)
}
filesystems = append(filesystems, fs)
}
return filesystems, nil
}
func readFilesystem(dir string, mounts map[string]string) (Filesystem, error) {
fs := Filesystem{UUID: filepath.Base(dir), Name: utils.ReadStringFile(filepath.Join(dir, "label")), Health: "UNKNOWN"}
for _, kind := range []string{"data", "metadata", "system"} {
if value, ok := utils.ReadUintFile(filepath.Join(dir, "allocation", kind, "disk_used")); ok {
fs.Alloc += value
}
}
// devices/<name> links to the block device's sysfs directory.
devices, err := os.ReadDir(filepath.Join(dir, "devices"))
if err != nil && !errors.Is(err, os.ErrNotExist) {
return fs, err
}
mountpoint := mounts["uuid:"+fs.UUID]
if fs.Name == "" {
fs.Name = mountpoint
}
var backingSize uint64
for _, dev := range devices {
if mountpoint == "" {
mountpoint = mounts[dev.Name()]
}
if fs.Name == "" {
fs.Name = mountpoint
}
devDir := filepath.Join(dir, "devices", dev.Name())
if size, ok := utils.ReadUintFile(filepath.Join(devDir, "size")); ok {
backingSize += size * 512
}
if stat := strings.Fields(utils.ReadStringFile(filepath.Join(devDir, "stat"))); len(stat) >= 7 {
fs.NRead += parseUint(stat[2]) * 512
fs.NWrite += parseUint(stat[6]) * 512
}
}
devids, err := os.ReadDir(filepath.Join(dir, "devinfo"))
if err != nil && !errors.Is(err, os.ErrNotExist) {
return fs, err
}
capacityAvailable := len(devids) > 0
healthKnown := len(devids) > 0
for _, devid := range devids {
devDir := filepath.Join(dir, "devinfo", devid.Name())
// Replacement targets do not add filesystem capacity.
replaceTarget, _ := utils.ReadUintFile(filepath.Join(devDir, "replace_target"))
if replaceTarget != 1 {
devid, err := strconv.ParseUint(devid.Name(), 10, 64)
if err != nil {
return fs, err
}
size, err := deviceSize(mountpoint, devid)
if err != nil {
capacityAvailable = false
}
fs.Size += size
}
dev := Device{Name: "devid " + devid.Name(), State: "ONLINE"}
missing := utils.ReadStringFile(filepath.Join(devDir, "missing"))
if missing != "0" && missing != "1" {
healthKnown = false
dev.State = "UNKNOWN"
}
if missing == "1" {
dev.State = "MISSING"
fs.Health = "DEGRADED"
}
for line := range strings.Lines(utils.ReadStringFile(filepath.Join(devDir, "error_stats"))) {
if fields := strings.Fields(line); len(fields) == 2 {
switch fields[0] {
case "read_errs":
dev.ReadErrs = parseUint(fields[1])
case "write_errs":
dev.WriteErrs = parseUint(fields[1])
case "corruption_errs":
dev.CorruptionErrs = parseUint(fields[1])
}
}
}
fs.Devices = append(fs.Devices, dev)
}
// Use one capacity source for the whole filesystem: device IDs cannot be
// reliably matched to block-device names in sysfs. A partial ioctl result
// must not be added to the complete backing-device total.
if !capacityAvailable {
fs.Size = backingSize
}
if fs.Health != "DEGRADED" && healthKnown {
fs.Health = "ONLINE"
}
fs.MountID = mountUUID(mountpoint)
if len(devices) == 1 && len(devids) == 1 && fs.Health == "ONLINE" {
fs.IODevice = devices[0].Name()
}
fs.Raw = true
if used, available, err := filesystemUsage(mountpoint); err == nil {
// Effective capacity excludes reserved/unavailable space, so Size-Alloc
// is available to applications and the usage ratio matches df.
fs.Size, fs.Alloc, fs.Raw = used+available, used, false
}
if fs.Name == "" {
fs.Name = filepath.Base(dir)
}
return fs, nil
}
// mountpointsByDevice prefers UUID matches from mountinfo and retains source
// device names as a fallback for environments where FS_INFO is unavailable.
func mountpointsByDevice() map[string]string {
mounts := mountpointsByUUID(utils.ReadStringFile(mountinfoPath), mountUUID)
for line := range strings.Lines(utils.ReadStringFile(mountsPath)) {
fields := strings.Fields(line)
if len(fields) < 3 || fields[2] != "btrfs" {
continue
}
device := fields[0]
if resolved, err := filepath.EvalSymlinks(device); err == nil {
device = resolved
}
if _, seen := mounts[filepath.Base(device)]; !seen {
mounts[filepath.Base(device)] = unescapeMountPath(fields[1])
}
}
return mounts
}
func parseUint(s string) uint64 {
n, _ := strconv.ParseUint(s, 10, 64)
return n
}
// ioctlDeviceSize reads Btrfs's recorded device size, which can be smaller
// than the block device after a filesystem resize. BTRFS_IOC_DEV_INFO is
// _IOWR(0x94, 30, struct btrfs_ioctl_dev_info_args), a 4096-byte ABI structure.
func ioctlDeviceSize(mountpoint string, devid uint64) (uint64, error) {
if mountpoint == "" {
return 0, errors.New("no accessible mountpoint")
}
f, err := os.Open(mountpoint)
if err != nil {
return 0, err
}
defer f.Close()
args := struct {
Devid uint64
UUID [16]byte
BytesUsed uint64
TotalBytes uint64
Reserved [4096 - 40]byte
}{Devid: devid}
_, _, errno := unix.Syscall(unix.SYS_IOCTL, f.Fd(), 0xd000941e, uintptr(unsafe.Pointer(&args)))
if errno != 0 {
return 0, errno
}
return args.TotalBytes, nil
}
// The filesystem magic is unsigned even when Statfs_t.Type is int32.
func isBtrfs(stat *unix.Statfs_t) bool {
return uint32(stat.Type) == unix.BTRFS_SUPER_MAGIC
}
func statfsUsage(path string) (used, available uint64, err error) {
if path == "" {
return 0, 0, errors.New("no accessible mountpoint")
}
var stat unix.Statfs_t
if err = unix.Statfs(path, &stat); err != nil {
return
}
if !isBtrfs(&stat) {
return 0, 0, errors.New("mountpoint is not Btrfs")
}
blockSize := uint64(stat.Bsize)
return (stat.Blocks - min(stat.Blocks, stat.Bfree)) * blockSize, min(stat.Blocks, stat.Bavail) * blockSize, nil
}
// MountID returns the filesystem UUID via BTRFS_IOC_FS_INFO. Unlike statfs
// f_fsid, this identity is shared by all subvolumes and bind mounts.
func MountID(path string) string {
if path == "" {
return ""
}
var stat unix.Statfs_t
if unix.Statfs(path, &stat) != nil || !isBtrfs(&stat) {
return ""
}
f, err := os.Open(path)
if err != nil {
return ""
}
defer f.Close()
args := struct {
MaxID uint64
NumDevices uint64
FSID [16]byte
Reserved [992]byte
}{}
// _IOR(0x94, 31, 1024). Reuse the platform's read-direction bits;
// MIPS/PowerPC use a different encoding than asm-generic.
request := uintptr(unix.FS_IOC_GETFLAGS&0xe0000000) | 0x0400941f
_, _, errno := unix.Syscall(unix.SYS_IOCTL, f.Fd(), request, uintptr(unsafe.Pointer(&args)))
if errno != 0 {
return ""
}
id := args.FSID
return fmt.Sprintf("%x-%x-%x-%x-%x", id[:4], id[4:6], id[6:8], id[8:10], id[10:])
}
// Btrfs mountinfo device numbers can be virtual (0:N), so query the UUID
// through the mount instead of comparing those numbers with sysfs block devs.
// Retry another path when a bind mount is inaccessible. Once resolved, reuse
// the result for that mount device to avoid opening every Docker bind mount.
func mountpointsByUUID(mountinfo string, identify func(string) string) map[string]string {
mounts := make(map[string]string)
resolved := make(map[string]bool)
for line := range strings.Lines(mountinfo) {
before, after, ok := strings.Cut(line, " - ")
fields, fs := strings.Fields(before), strings.Fields(after)
if !ok || len(fields) < 6 || len(fs) < 3 || fs[0] != "btrfs" || resolved[fields[2]] {
continue
}
path := unescapeMountPath(fields[4])
uuid := identify(path)
if uuid == "" {
continue
}
resolved[fields[2]] = true
if mounts["uuid:"+uuid] == "" {
mounts["uuid:"+uuid] = path
}
}
return mounts
}
func unescapeMountPath(path string) string {
return strings.NewReplacer(`\040`, " ", `\011`, "\t", `\012`, "\n", `\134`, `\`).Replace(path)
}

View File

@@ -1,274 +0,0 @@
//go:build testing && linux
package btrfs
import (
"os"
"path/filepath"
"strconv"
"testing"
"github.com/henrygd/beszel/agent/utils"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/sys/unix"
)
func TestFilesystems(t *testing.T) {
root := t.TempDir()
oldSysfs, oldMounts := sysfsPath, mountsPath
sysfsPath, mountsPath = root, filepath.Join(root, "mounts")
t.Cleanup(func() { sysfsPath, mountsPath = oldSysfs, oldMounts })
fsDir := filepath.Join(root, "1b2c3d4e-0000-0000-0000-000000000000")
write := func(rel, content string) {
path := filepath.Join(fsDir, rel)
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o755))
require.NoError(t, os.WriteFile(path, []byte(content), 0o644))
}
require.NoError(t, os.MkdirAll(filepath.Join(root, "features"), 0o755))
oldUsage := filesystemUsage
filesystemUsage = func(string) (uint64, uint64, error) { return 0, 0, os.ErrNotExist }
t.Cleanup(func() { filesystemUsage = oldUsage })
oldDeviceSize := deviceSize
t.Cleanup(func() { deviceSize = oldDeviceSize })
deviceSize = func(_ string, devid uint64) (uint64, error) {
value, _ := utils.ReadUintFile(filepath.Join(fsDir, "recorded-size", strconv.FormatUint(devid, 10)))
return value, nil
}
// Recorded member capacities differ from the unchanged backing devices.
write("recorded-size/1", "256000\n")
write("recorded-size/2", "128000\n")
write("label", "tank\n")
write("allocation/data/disk_used", "4096\n")
write("allocation/metadata/disk_used", "2048\n")
write("allocation/system/disk_used", "1024\n")
write("devices/sda/size", "1000\n")
write("devices/sda/stat", "10 0 200 0 20 0 400 0 0 0 0\n")
write("devices/sdb/size", "1000\n")
write("devices/sdb/stat", "10 0 100 0 20 0 100 0 0 0 0\n")
write("devinfo/1/missing", "0\n")
write("devinfo/1/error_stats", "write_errs 1\nread_errs 2\nflush_errs 0\ncorruption_errs 3\ngeneration_errs 0\n")
write("devinfo/2/missing", "1\n")
filesystems, err := Filesystems()
require.NoError(t, err)
require.Len(t, filesystems, 1)
assert.Equal(t, Filesystem{
UUID: "1b2c3d4e-0000-0000-0000-000000000000", Raw: true, Name: "tank", Size: 384000, Alloc: 7168, Health: "DEGRADED", NRead: 153600, NWrite: 256000,
Devices: []Device{
{Name: "devid 1", State: "ONLINE", ReadErrs: 2, WriteErrs: 1, CorruptionErrs: 3},
{Name: "devid 2", State: "MISSING"},
},
}, filesystems[0])
// Unlabeled filesystems fall back to the first mountpoint, then the UUID.
write("label", "\n")
require.NoError(t, os.WriteFile(mountsPath, []byte(
"/dev/sdz1 /other btrfs rw 0 0\n/dev/sdb /mnt/storage btrfs rw 0 0\n/dev/sdb /mnt/storage/sub btrfs rw,subvol=/sub 0 0\n",
), 0o644))
filesystems, err = Filesystems()
require.NoError(t, err)
assert.Equal(t, "/mnt/storage", filesystems[0].Name)
require.NoError(t, os.Remove(mountsPath))
filesystems, err = Filesystems()
require.NoError(t, err)
assert.Equal(t, "1b2c3d4e-0000-0000-0000-000000000000", filesystems[0].Name)
write("devinfo/3/replace_target", "1\n")
write("recorded-size/3", "512000\n")
filesystems, err = Filesystems()
require.NoError(t, err)
assert.Equal(t, uint64(384000), filesystems[0].Size, "replacement target must not inflate capacity")
deviceSize = func(string, uint64) (uint64, error) { return 0, os.ErrPermission }
filesystems, err = Filesystems()
require.NoError(t, err)
require.Len(t, filesystems, 1)
assert.Equal(t, uint64(1024000), filesystems[0].Size)
assert.Equal(t, "DEGRADED", filesystems[0].Health)
assert.Equal(t, uint64(153600), filesystems[0].NRead)
// A partial ioctl result must not be mixed with the backing-device total.
deviceSize = func(_ string, devid uint64) (uint64, error) {
if devid == 2 {
return 0, os.ErrPermission
}
return 256000, nil
}
filesystems, err = Filesystems()
require.NoError(t, err)
assert.Equal(t, uint64(1024000), filesystems[0].Size)
// With no mount visible (e.g. Docker), the real lookup falls back too.
deviceSize = ioctlDeviceSize
filesystems, err = Filesystems()
require.NoError(t, err)
require.Len(t, filesystems, 1)
assert.Equal(t, uint64(1024000), filesystems[0].Size)
filesystemUsage = func(string) (uint64, uint64, error) { return 100, 900, nil }
filesystems, err = Filesystems()
require.NoError(t, err)
assert.Equal(t, uint64(1000), filesystems[0].Size)
assert.Equal(t, uint64(100), filesystems[0].Alloc)
assert.False(t, filesystems[0].Raw)
}
func TestFilesystemsNoBtrfs(t *testing.T) {
oldPath := sysfsPath
sysfsPath = filepath.Join(t.TempDir(), "missing")
t.Cleanup(func() { sysfsPath = oldPath })
filesystems, err := Filesystems()
require.NoError(t, err)
assert.Nil(t, filesystems)
}
func TestIoctlDeviceSizeFailure(t *testing.T) {
_, err := ioctlDeviceSize("", 1)
require.Error(t, err)
_, err = ioctlDeviceSize(t.TempDir(), 1)
require.Error(t, err)
assert.ErrorIs(t, err, unix.ENOTTY)
}
func TestMountpointsDecodeEscapes(t *testing.T) {
oldMounts := mountsPath
mountsPath = filepath.Join(t.TempDir(), "mounts")
t.Cleanup(func() { mountsPath = oldMounts })
require.NoError(t, os.WriteFile(mountsPath, []byte("/dev/test-btrfs /mnt/my\\040data btrfs rw 0 0\n"), 0o644))
assert.Equal(t, "/mnt/my data", mountpointsByDevice()["test-btrfs"])
}
func TestFilesystemWithoutDevinfo(t *testing.T) {
root := t.TempDir()
require.NoError(t, os.MkdirAll(filepath.Join(root, "devices", "sda"), 0755))
require.NoError(t, os.WriteFile(filepath.Join(root, "devices", "sda", "size"), []byte("1000"), 0644))
fs, err := readFilesystem(root, nil)
require.NoError(t, err)
assert.Equal(t, uint64(512000), fs.Size)
assert.True(t, fs.Raw)
assert.Equal(t, "UNKNOWN", fs.Health)
assert.Empty(t, fs.Devices)
require.NoError(t, os.MkdirAll(filepath.Join(root, "devinfo", "1"), 0755))
fs, err = readFilesystem(root, nil)
require.NoError(t, err)
assert.Equal(t, "UNKNOWN", fs.Health)
require.Len(t, fs.Devices, 1)
assert.Equal(t, "UNKNOWN", fs.Devices[0].State)
// Some older interfaces lack the devices directory too.
fs, err = readFilesystem(t.TempDir(), nil)
require.NoError(t, err)
assert.Equal(t, "UNKNOWN", fs.Health)
}
func TestLocalBtrfsUsage(t *testing.T) {
path := os.Getenv("BESZEL_TEST_BTRFS_MOUNT")
if path == "" {
t.Skip("set BESZEL_TEST_BTRFS_MOUNT for read-only live validation")
}
used, available, err := statfsUsage(path)
require.NoError(t, err)
filesystems, err := Filesystems()
require.NoError(t, err)
for _, fs := range filesystems {
if !fs.Raw && fs.Alloc == used && fs.Size == used+available {
t.Logf("pool=%s used=%d available=%d effective_capacity=%d", fs.Name, used, available, fs.Size)
return
}
}
t.Fatal("collector did not report the mounted filesystem's usable capacity")
}
func TestMountID(t *testing.T) {
assert.Empty(t, MountID(""))
assert.Empty(t, MountID(filepath.Join(t.TempDir(), "missing")))
path := os.Getenv("BESZEL_TEST_BTRFS_MOUNT")
if path == "" {
t.Skip("set BESZEL_TEST_BTRFS_MOUNT for live identity validation")
}
id := MountID(path)
require.NotEmpty(t, id)
assert.Equal(t, id, MountID(filepath.Join(path, ".")))
}
func TestMountinfoUUIDLookup(t *testing.T) {
info := `1 0 0:40 /@ /inaccessible ro shared:1 - btrfs /dev/mapper/unavailable rw
2 0 0:40 /@/docker/hosts /etc/hosts ro - btrfs /dev/mapper/unavailable rw
3 0 0:40 /@/docker/hostname /etc/hostname ro - btrfs /dev/mapper/unavailable rw
4 0 0:41 /subvol /extra-filesystems/my\040disk ro master:2 - btrfs /dev/missing rw
5 0 0:42 / /ext4 ro - ext4 /dev/mapper/unavailable rw
malformed
6 0 0:43 / /bad ro - btrfs
`
var calls []string
mounts := mountpointsByUUID(info, func(path string) string {
calls = append(calls, path)
switch path {
case "/etc/hosts":
return "root-uuid"
case "/extra-filesystems/my disk":
return "extra-uuid"
}
return ""
})
assert.Equal(t, map[string]string{"uuid:root-uuid": "/etc/hosts", "uuid:extra-uuid": "/extra-filesystems/my disk"}, mounts)
assert.Equal(t, []string{"/inaccessible", "/etc/hosts", "/extra-filesystems/my disk"}, calls)
}
func TestDockerFilesystemWithoutDeviceNodes(t *testing.T) {
root := t.TempDir()
oldSysfs, oldMounts, oldInfo, oldUUID, oldUsage := sysfsPath, mountsPath, mountinfoPath, mountUUID, filesystemUsage
t.Cleanup(func() {
sysfsPath, mountsPath, mountinfoPath, mountUUID, filesystemUsage = oldSysfs, oldMounts, oldInfo, oldUUID, oldUsage
})
sysfsPath = filepath.Join(root, "sysfs")
mountsPath = filepath.Join(root, "missing-mounts")
mountinfoPath = filepath.Join(root, "mountinfo")
uuid := "11111111-1111-4111-8111-111111111111"
dir := filepath.Join(sysfsPath, uuid)
for path, content := range map[string]string{"devices/dm-0/size": "1000", "devinfo/1/missing": "0"} {
target := filepath.Join(dir, path)
require.NoError(t, os.MkdirAll(filepath.Dir(target), 0755))
require.NoError(t, os.WriteFile(target, []byte(content), 0644))
}
require.NoError(t, os.WriteFile(mountinfoPath, []byte("2 1 0:40 /@/docker/hosts /etc/hosts ro - btrfs /dev/mapper/not-in-container rw\n"), 0644))
mountUUID = func(path string) string {
if path == "/etc/hosts" {
return uuid
}
return ""
}
filesystemUsage = func(path string) (uint64, uint64, error) { require.Equal(t, "/etc/hosts", path); return 100, 900, nil }
fs, err := Filesystems()
require.NoError(t, err)
require.Len(t, fs, 1)
assert.Equal(t, uuid, fs[0].MountID)
assert.Equal(t, "dm-0", fs[0].IODevice)
assert.False(t, fs[0].Raw)
assert.Equal(t, uint64(1000), fs[0].Size)
}
func TestLivePoolMountIdentity(t *testing.T) {
path := os.Getenv("BESZEL_TEST_BTRFS_MOUNT")
if path == "" {
t.Skip("set BESZEL_TEST_BTRFS_MOUNT for live validation")
}
id := MountID(path)
require.NotEmpty(t, id)
pools, err := Filesystems()
require.NoError(t, err)
for _, pool := range pools {
if pool.UUID != id {
continue
}
assert.Equal(t, id, pool.MountID)
assert.False(t, pool.Raw)
t.Logf("uuid=%s mount_identity=%s io_device=%s raw=%v", pool.UUID, pool.MountID, pool.IODevice, pool.Raw)
return
}
t.Fatal("mounted Btrfs filesystem was not discovered")
}

View File

@@ -1,11 +0,0 @@
//go:build !linux
package btrfs
import "errors"
func Filesystems() ([]Filesystem, error) {
return nil, errors.ErrUnsupported
}
func MountID(string) string { return "" }

View File

@@ -2,7 +2,6 @@ package agent
import (
"crypto/tls"
"crypto/x509"
"errors"
"fmt"
"log/slog"
@@ -12,13 +11,11 @@ import (
"os"
"path"
"strings"
"sync"
"time"
"github.com/henrygd/beszel"
"github.com/henrygd/beszel/agent/utils"
"github.com/henrygd/beszel/internal/common"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/fxamacker/cbor/v2"
"github.com/lxzan/gws"
@@ -27,35 +24,15 @@ import (
)
const (
// Keep the connection alive long enough for a slow collection cycle to
// finish before the hub considers the agent disconnected.
wsDeadline = 120 * time.Second
wsDeadline = 70 * time.Second
)
// errNoHubURL is returned when HUB_URL is unset. This is not a failure
// condition: an agent configured with only a public key runs in SSH-only mode,
// where the hub dials the agent and no outbound WebSocket client is expected.
var errNoHubURL = errors.New("HUB_URL environment variable not set")
type caCertFileError struct {
err error
}
func (e *caCertFileError) Error() string {
return e.err.Error()
}
func (e *caCertFileError) Unwrap() error {
return e.err
}
// WebSocketClient manages the WebSocket connection between the agent and hub.
// It handles authentication, message routing, and connection lifecycle management.
type WebSocketClient struct {
gws.BuiltinEventHandler
options *gws.ClientOption // WebSocket client configuration options
agent *Agent // Reference to the parent agent
connMu sync.RWMutex // Guards Conn and hubVerified across callbacks
Conn *gws.Conn // Active WebSocket connection
hubURL *url.URL // Parsed hub URL for connection
token string // Authentication token for hub registration
@@ -63,7 +40,6 @@ type WebSocketClient struct {
hubRequest *common.HubRequest[cbor.RawMessage] // Reusable request structure for message parsing
lastConnectAttempt time.Time // Timestamp of last connection attempt
hubVerified bool // Whether the hub has been cryptographically verified
tlsConfig *tls.Config // Optional TLS configuration with custom CA certificates
}
// newWebSocketClient creates a new WebSocket client for the given agent.
@@ -71,24 +47,20 @@ type WebSocketClient struct {
func newWebSocketClient(agent *Agent) (client *WebSocketClient, err error) {
hubURLStr, exists := utils.GetEnv("HUB_URL")
if !exists {
return nil, errNoHubURL
return nil, errors.New("HUB_URL environment variable not set")
}
client = &WebSocketClient{}
client.hubURL, err = url.Parse(hubURLStr)
if err != nil || client.hubURL.Host == "" {
return nil, fmt.Errorf("invalid HUB_URL %q: must include scheme and host (e.g. http://hub.example.com:8090)", hubURLStr)
if err != nil {
return nil, errors.New("invalid hub URL")
}
// get registration token
client.token, err = getToken()
if err != nil {
return nil, err
}
client.tlsConfig, err = getTLSConfig()
if err != nil {
return nil, err
}
client.agent = agent
client.hubRequest = &common.HubRequest[cbor.RawMessage]{}
@@ -115,52 +87,7 @@ func getToken() (string, error) {
if err != nil {
return "", err
}
return parseTokenFile(string(tokenBytes), tokenFile)
}
// parseTokenFile reads a single token from TOKEN_FILE.
// Blank lines and comments are ignored. Multiple tokens are rejected because
// the agent supports only one outbound hub connection.
func parseTokenFile(contents, path string) (string, error) {
var token string
for line := range strings.Lines(contents) {
line = strings.TrimSpace(line)
if len(line) == 0 || strings.HasPrefix(line, "#") {
continue
}
if token != "" {
return "", fmt.Errorf("%s must contain a single token", path)
}
token = line
}
// An empty file keeps returning an empty token, as before: the caller decides
// what to do about it.
return token, nil
}
// getTLSConfig returns a TLS configuration containing the system certificate
// pool plus any certificates configured through CA_CERT_FILE. A nil config lets
// gws use Go's default TLS configuration and system roots.
func getTLSConfig() (*tls.Config, error) {
caCertFile, _ := utils.GetEnv("CA_CERT_FILE")
if caCertFile == "" {
return nil, nil
}
caCertPEM, err := os.ReadFile(caCertFile)
if err != nil {
return nil, &caCertFileError{fmt.Errorf("read CA_CERT_FILE %q: %w", caCertFile, err)}
}
rootCAs, err := x509.SystemCertPool()
if err != nil {
return nil, &caCertFileError{fmt.Errorf("load system CA certificate pool: %w", err)}
}
if !rootCAs.AppendCertsFromPEM(caCertPEM) {
return nil, &caCertFileError{fmt.Errorf("CA_CERT_FILE %q does not contain any valid PEM certificates", caCertFile)}
}
return &tls.Config{RootCAs: rootCAs}, nil
return strings.TrimSpace(string(tokenBytes)), nil
}
// getOptions returns the WebSocket client options, creating them if necessary.
@@ -185,7 +112,7 @@ func (client *WebSocketClient) getOptions() *gws.ClientOption {
client.options = &gws.ClientOption{
Addr: client.hubURL.String(),
TlsConfig: client.tlsConfig,
TlsConfig: &tls.Config{InsecureSkipVerify: true},
RequestHeader: http.Header{
"User-Agent": []string{getUserAgent()},
"X-Token": []string{client.token},
@@ -206,16 +133,12 @@ func (client *WebSocketClient) Connect() (err error) {
// make sure previous connection is closed
client.Close()
conn, _, err := gws.NewClient(client, client.getOptions())
client.Conn, _, err = gws.NewClient(client, client.getOptions())
if err != nil {
return err
}
client.connMu.Lock()
client.Conn = conn
client.hubVerified = false
client.connMu.Unlock()
go conn.ReadLoop()
go client.Conn.ReadLoop()
return nil
}
@@ -229,14 +152,6 @@ func (client *WebSocketClient) OnOpen(conn *gws.Conn) {
// OnClose handles WebSocket connection closure.
// It logs the closure reason and notifies the connection manager.
func (client *WebSocketClient) OnClose(conn *gws.Conn, err error) {
client.connMu.Lock()
if client.Conn != conn {
client.connMu.Unlock()
return
}
client.Conn = nil
client.hubVerified = false
client.connMu.Unlock()
if err != nil {
slog.Warn("Connection closed", "err", strings.TrimPrefix(err.Error(), "gws: "))
}
@@ -247,9 +162,6 @@ func (client *WebSocketClient) OnClose(conn *gws.Conn, err error) {
// It decodes CBOR messages and routes them to appropriate handlers.
func (client *WebSocketClient) OnMessage(conn *gws.Conn, message *gws.Message) {
defer message.Close()
if client.getConn() != conn {
return
}
conn.SetDeadline(time.Now().Add(wsDeadline))
if message.Opcode != gws.OpcodeBinary {
@@ -264,7 +176,7 @@ func (client *WebSocketClient) OnMessage(conn *gws.Conn, message *gws.Message) {
return
}
if err := client.handleHubRequest(&HubRequest, HubRequest.Id, conn); err != nil {
if err := client.handleHubRequest(&HubRequest, HubRequest.Id); err != nil {
slog.Error("Error handling message", "err", err)
}
}
@@ -277,7 +189,7 @@ func (client *WebSocketClient) OnPing(conn *gws.Conn, message []byte) {
}
// handleAuthChallenge verifies the authenticity of the hub and returns the system's fingerprint.
func (client *WebSocketClient) handleAuthChallenge(msg *common.HubRequest[cbor.RawMessage], requestID *uint32, conn *gws.Conn) (err error) {
func (client *WebSocketClient) handleAuthChallenge(msg *common.HubRequest[cbor.RawMessage], requestID *uint32) (err error) {
var authRequest common.FingerprintRequest
if err := cbor.Unmarshal(msg.Data, &authRequest); err != nil {
return err
@@ -287,13 +199,7 @@ func (client *WebSocketClient) handleAuthChallenge(msg *common.HubRequest[cbor.R
return err
}
client.connMu.Lock()
if conn != nil && client.Conn != conn {
client.connMu.Unlock()
return gws.ErrConnClosed
}
client.hubVerified = true
client.connMu.Unlock()
client.agent.connectionManager.eventChan <- WebSocketConnect
response := &common.FingerprintResponse{
@@ -307,9 +213,6 @@ func (client *WebSocketClient) handleAuthChallenge(msg *common.HubRequest[cbor.R
_, response.Port, _ = net.SplitHostPort(serverAddr)
}
if conn != nil {
return client.sendResponseOnConn(conn, response, requestID)
}
return client.sendResponse(response, requestID)
}
@@ -330,65 +233,35 @@ func (client *WebSocketClient) verifySignature(signature []byte) (err error) {
// Close closes the WebSocket connection gracefully.
// This method is safe to call multiple times.
func (client *WebSocketClient) Close() {
if conn := client.getConn(); conn != nil {
_ = conn.WriteClose(1000, nil)
if client.Conn != nil {
_ = client.Conn.WriteClose(1000, nil)
}
}
func (client *WebSocketClient) getConn() *gws.Conn {
client.connMu.RLock()
defer client.connMu.RUnlock()
return client.Conn
}
func (client *WebSocketClient) isVerified() bool {
client.connMu.RLock()
defer client.connMu.RUnlock()
return client.Conn != nil && client.hubVerified
}
// handleHubRequest routes the request to the appropriate handler using the handler registry.
func (client *WebSocketClient) handleHubRequest(msg *common.HubRequest[cbor.RawMessage], requestID *uint32, conn *gws.Conn) error {
client.connMu.RLock()
verified := client.hubVerified
client.connMu.RUnlock()
sendResponse := client.sendResponse
if conn != nil {
sendResponse = func(data any, requestID *uint32) error {
return client.sendResponseOnConn(conn, data, requestID)
}
}
func (client *WebSocketClient) handleHubRequest(msg *common.HubRequest[cbor.RawMessage], requestID *uint32) error {
ctx := &HandlerContext{
Client: client,
Conn: conn,
Agent: client.agent,
Request: msg,
RequestID: requestID,
HubVerified: verified,
ConnectionType: system.ConnectionTypeWebSocket,
SendResponse: sendResponse,
Client: client,
Agent: client.agent,
Request: msg,
RequestID: requestID,
HubVerified: client.hubVerified,
SendResponse: client.sendResponse,
}
return client.agent.handlerRegistry.Handle(ctx)
}
// sendMessage encodes the given data to CBOR and sends it as a binary message over the WebSocket connection to the hub.
func (client *WebSocketClient) sendMessage(data any) error {
return client.sendMessageOnConn(client.getConn(), data)
}
func (client *WebSocketClient) sendMessageOnConn(conn *gws.Conn, data any) error {
bytes, err := cbor.Marshal(data)
if err != nil {
return err
}
if conn == nil {
return gws.ErrConnClosed
}
err = conn.WriteMessage(gws.OpcodeBinary, bytes)
err = client.Conn.WriteMessage(gws.OpcodeBinary, bytes)
if err != nil {
// If writing fails (e.g., broken pipe due to network issues),
// close the connection to trigger reconnection logic (#1263)
_ = conn.WriteClose(1000, nil)
client.Close()
}
return err
}
@@ -397,16 +270,12 @@ func (client *WebSocketClient) sendMessageOnConn(conn *gws.Conn, data any) error
// For ID-based requests, we must populate legacy typed fields for backward
// compatibility with older hubs (<= 0.17) that don't read the generic Data field.
func (client *WebSocketClient) sendResponse(data any, requestID *uint32) error {
return client.sendResponseOnConn(client.getConn(), data, requestID)
}
func (client *WebSocketClient) sendResponseOnConn(conn *gws.Conn, data any, requestID *uint32) error {
if requestID != nil {
response := newAgentResponse(data, requestID)
return client.sendMessageOnConn(conn, response)
return client.sendMessage(response)
}
// Legacy format - send data directly
return client.sendMessageOnConn(conn, data)
return client.sendMessage(data)
}
// getUserAgent returns one of two User-Agent strings based on current time.

View File

@@ -4,20 +4,8 @@ package agent
import (
"crypto/ed25519"
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"math/big"
"net"
"net/http"
"net/http/httptest"
"net/url"
"os"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
@@ -27,34 +15,11 @@ import (
"github.com/henrygd/beszel/internal/common"
"github.com/fxamacker/cbor/v2"
"github.com/lxzan/gws"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/crypto/ssh"
)
// TestNewWebSocketClientNoHubURL verifies that an unset HUB_URL returns the
// errNoHubURL sentinel rather than an opaque error. Callers rely on this to
// distinguish SSH-only mode -- a supported configuration in which the hub dials
// the agent -- from an actual misconfiguration.
func TestNewWebSocketClientNoHubURL(t *testing.T) {
agent := createTestAgent(t)
// t.Setenv registers restoration of the original value; unset afterwards so
// GetEnv's LookupEnv reports the variable as absent rather than empty.
t.Setenv("BESZEL_AGENT_HUB_URL", "")
os.Unsetenv("BESZEL_AGENT_HUB_URL")
t.Setenv("HUB_URL", "")
os.Unsetenv("HUB_URL")
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
client, err := newWebSocketClient(agent)
require.Error(t, err)
assert.Nil(t, client)
assert.ErrorIs(t, err, errNoHubURL)
}
// TestNewWebSocketClient tests WebSocket client creation
func TestNewWebSocketClient(t *testing.T) {
agent := createTestAgent(t)
@@ -86,18 +51,11 @@ func TestNewWebSocketClient(t *testing.T) {
errorMsg: "HUB_URL environment variable not set",
},
{
name: "malformed URL",
name: "invalid URL",
hubURL: "ht\ttp://invalid",
token: "test-token",
expectError: true,
errorMsg: "invalid HUB_URL",
},
{
name: "URL without host",
hubURL: "http:/api",
token: "test-token",
expectError: true,
errorMsg: "invalid HUB_URL",
errorMsg: "invalid hub URL",
},
{
name: "missing token",
@@ -199,158 +157,6 @@ func TestWebSocketClient_GetOptions(t *testing.T) {
}
}
func TestWebSocketClient_TLSVerification(t *testing.T) {
agent := createTestAgent(t)
serverCert, serverCertPEM := newSelfSignedServerCertificate(t)
upgrader := gws.NewUpgrader(&gws.BuiltinEventHandler{}, nil)
server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r)
if err == nil {
go conn.ReadLoop()
}
}))
server.TLS = &tls.Config{Certificates: []tls.Certificate{serverCert}}
server.StartTLS()
t.Cleanup(server.Close)
caCertFile := filepath.Join(t.TempDir(), "hub-ca.crt")
require.NoError(t, os.WriteFile(caCertFile, serverCertPEM, 0600))
newClient := func(t *testing.T, caCertFile string) *WebSocketClient {
t.Helper()
t.Setenv("BESZEL_AGENT_HUB_URL", server.URL)
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
t.Setenv("BESZEL_AGENT_CA_CERT_FILE", caCertFile)
client, err := newWebSocketClient(agent)
require.NoError(t, err)
return client
}
t.Run("system roots are used by default", func(t *testing.T) {
client := newClient(t, "")
assert.Nil(t, client.getOptions().TlsConfig)
_, _, err := gws.NewClient(&gws.BuiltinEventHandler{}, client.getOptions())
require.Error(t, err)
})
t.Run("custom CA trusts self-signed certificate", func(t *testing.T) {
systemRoots, err := x509.SystemCertPool()
require.NoError(t, err)
require.True(t, systemRoots.AppendCertsFromPEM(serverCertPEM))
client := newClient(t, caCertFile)
tlsConfig := client.getOptions().TlsConfig
require.NotNil(t, tlsConfig)
assert.True(t, tlsConfig.RootCAs.Equal(systemRoots))
conn, _, err := gws.NewClient(&gws.BuiltinEventHandler{}, client.getOptions())
require.NoError(t, err)
require.NoError(t, conn.NetConn().Close())
})
t.Run("custom CA does not bypass hostname verification", func(t *testing.T) {
client := newClient(t, caCertFile)
client.getOptions().TlsConfig.ServerName = "wrong.example.com"
_, _, err := gws.NewClient(&gws.BuiltinEventHandler{}, client.getOptions())
require.Error(t, err)
})
}
func TestWebSocketClient_NonTLSConnection(t *testing.T) {
agent := createTestAgent(t)
upgrader := gws.NewUpgrader(&gws.BuiltinEventHandler{}, nil)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r)
if err == nil {
go conn.ReadLoop()
}
}))
t.Cleanup(server.Close)
t.Setenv("BESZEL_AGENT_HUB_URL", server.URL)
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
t.Setenv("BESZEL_AGENT_CA_CERT_FILE", "")
client, err := newWebSocketClient(agent)
require.NoError(t, err)
assert.Nil(t, client.getOptions().TlsConfig)
conn, _, err := gws.NewClient(&gws.BuiltinEventHandler{}, client.getOptions())
require.NoError(t, err)
require.NoError(t, conn.NetConn().Close())
}
func TestGetTLSConfigErrors(t *testing.T) {
tempDir := t.TempDir()
testCases := []struct {
name string
path string
contents []byte
errorMatch string
}{
{
name: "missing file",
path: filepath.Join(tempDir, "missing.pem"),
errorMatch: "read CA_CERT_FILE",
},
{
name: "unreadable path",
path: tempDir,
errorMatch: "read CA_CERT_FILE",
},
{
name: "empty file",
path: filepath.Join(tempDir, "empty.pem"),
contents: []byte{},
errorMatch: "does not contain any valid PEM certificates",
},
{
name: "malformed file",
path: filepath.Join(tempDir, "malformed.pem"),
contents: []byte("not a PEM certificate"),
errorMatch: "does not contain any valid PEM certificates",
},
}
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
if tc.contents != nil {
require.NoError(t, os.WriteFile(tc.path, tc.contents, 0600))
}
t.Setenv("BESZEL_AGENT_CA_CERT_FILE", tc.path)
tlsConfig, err := getTLSConfig()
require.Error(t, err)
assert.Nil(t, tlsConfig)
assert.Contains(t, err.Error(), tc.errorMatch)
assert.Contains(t, err.Error(), strconv.Quote(tc.path))
})
}
}
func newSelfSignedServerCertificate(t *testing.T) (tls.Certificate, []byte) {
t.Helper()
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoError(t, err)
template := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "127.0.0.1"},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
IPAddresses: []net.IP{net.ParseIP("127.0.0.1")},
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment | x509.KeyUsageCertSign,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
BasicConstraintsValid: true,
IsCA: true,
}
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
require.NoError(t, err)
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)})
certificate, err := tls.X509KeyPair(certPEM, keyPEM)
require.NoError(t, err)
return certificate, certPEM
}
// TestWebSocketClient_VerifySignature tests signature verification
func TestWebSocketClient_VerifySignature(t *testing.T) {
agent := createTestAgent(t)
@@ -474,7 +280,7 @@ func TestWebSocketClient_HandleHubRequest(t *testing.T) {
Data: cbor.RawMessage{},
}
err := client.handleHubRequest(hubRequest, nil, nil)
err := client.handleHubRequest(hubRequest, nil)
if tc.expectError {
assert.Error(t, err)
@@ -536,18 +342,6 @@ func TestWebSocketClient_Close(t *testing.T) {
})
}
func TestWebSocketClient_IgnoresStaleClose(t *testing.T) {
agent := createTestAgent(t)
agent.connectionManager.eventChan = make(chan ConnectionEvent, 1)
current := &gws.Conn{}
client := &WebSocketClient{agent: agent, Conn: current, hubVerified: true}
client.OnClose(&gws.Conn{}, nil)
assert.Same(t, current, client.getConn())
assert.True(t, client.hubVerified)
assert.Empty(t, agent.connectionManager.eventChan)
}
// TestWebSocketClient_ConnectRateLimit tests connection rate limiting
func TestWebSocketClient_ConnectRateLimit(t *testing.T) {
agent := createTestAgent(t)
@@ -608,41 +402,6 @@ func TestGetToken(t *testing.T) {
assert.Equal(t, expectedToken, token)
})
t.Run("TOKEN_FILE with surrounding blank lines and comments", func(t *testing.T) {
expectedToken := "test-token-with-noise"
tokenFile := filepath.Join(t.TempDir(), "token")
require.NoError(t, os.WriteFile(tokenFile, []byte("# hub token\n\n"+expectedToken+"\n\n"), 0o600))
t.Setenv("TOKEN_FILE", tokenFile)
token, err := getToken()
assert.NoError(t, err)
assert.Equal(t, expectedToken, token)
})
t.Run("TOKEN_FILE with multiple tokens is rejected", func(t *testing.T) {
tokenFile := filepath.Join(t.TempDir(), "token")
require.NoError(t, os.WriteFile(tokenFile, []byte("11111111-1111-1111-1111-111111111111\n22222222-2222-2222-2222-222222222222\n"), 0o600))
t.Setenv("TOKEN_FILE", tokenFile)
token, err := getToken()
require.Error(t, err)
assert.Empty(t, token)
assert.Contains(t, err.Error(), "must contain a single token")
})
t.Run("TOKEN_FILE holding only comments behaves like an empty file", func(t *testing.T) {
tokenFile := filepath.Join(t.TempDir(), "token")
require.NoError(t, os.WriteFile(tokenFile, []byte("\n# only a comment\n"), 0o600))
t.Setenv("TOKEN_FILE", tokenFile)
token, err := getToken()
assert.NoError(t, err)
assert.Equal(t, "", token)
})
t.Run("token from BESZEL_AGENT_TOKEN_FILE", func(t *testing.T) {
// Create a temporary token file
expectedToken := "test-token-from-beszel-file"
@@ -697,12 +456,12 @@ func TestGetToken(t *testing.T) {
t.Run("error when TOKEN_FILE points to non-existent file", func(t *testing.T) {
// Set TOKEN_FILE to a non-existent file
t.Setenv("TOKEN_FILE", filepath.Join(t.TempDir(), "missing.txt"))
t.Setenv("TOKEN_FILE", "/non/existent/file.txt")
token, err := getToken()
assert.Error(t, err)
assert.Equal(t, "", token)
assert.ErrorIs(t, err, os.ErrNotExist)
assert.Contains(t, err.Error(), "no such file or directory")
})
t.Run("handles empty token file", func(t *testing.T) {
@@ -738,11 +497,3 @@ func TestGetToken(t *testing.T) {
assert.Equal(t, expectedToken, token, "Whitespace should be stripped from token file content")
})
}
func TestWebSocketDeadlineCoversSlowCollection(t *testing.T) {
const minimumDeadline = 120 * time.Second
if wsDeadline < minimumDeadline {
t.Fatalf("WebSocket deadline %s is shorter than the slow-collection window of %s", wsDeadline, minimumDeadline)
}
}

View File

@@ -8,30 +8,26 @@ import (
"os"
"os/signal"
"strings"
"sync"
"syscall"
"time"
"github.com/gliderlabs/ssh"
"github.com/henrygd/beszel/agent/health"
"github.com/henrygd/beszel/agent/utils"
"github.com/henrygd/beszel/internal/entities/system"
)
// ConnectionManager manages the connection state and events for the agent.
// It handles both WebSocket and SSH connections, automatically switching between
// them based on availability and managing reconnection attempts.
type ConnectionManager struct {
agent *Agent // Reference to the parent agent
// mu guards state shared by the event loop, connection attempts and SSH callbacks.
mu sync.Mutex
agent *Agent // Reference to the parent agent
State ConnectionState // Current connection state
eventChan chan ConnectionEvent // Channel for connection events
sshChanged chan struct{} // Coalesced, nonblocking SSH connection notifications
wsClient *WebSocketClient // WebSocket client for hub communication
serverOptions ServerOptions // Configuration for SSH server
wsTicker *time.Ticker // Ticker for WebSocket connection attempts
isConnecting bool // Prevents multiple simultaneous reconnection attempts
sshConnections int // Authenticated SSH TCP connections, not sessions
ConnectionType system.ConnectionType
}
// ConnectionState represents the current connection state of the agent.
@@ -60,9 +56,8 @@ const wsTickerInterval = 10 * time.Second
// newConnectionManager creates a new connection manager for the given agent.
func newConnectionManager(agent *Agent) *ConnectionManager {
cm := &ConnectionManager{
agent: agent,
State: Disconnected,
sshChanged: make(chan struct{}, 1),
agent: agent,
State: Disconnected,
}
return cm
}
@@ -83,81 +78,6 @@ func (c *ConnectionManager) stopWsTicker() {
}
}
// getState returns the current connection state.
func (c *ConnectionManager) getState() ConnectionState {
c.mu.Lock()
defer c.mu.Unlock()
return c.State
}
func (c *ConnectionManager) hasSSHConnection() bool {
c.mu.Lock()
defer c.mu.Unlock()
return c.sshConnections > 0
}
func (c *ConnectionManager) notifySSHChange() {
select {
case c.sshChanged <- struct{}{}:
default:
}
}
// sshConnectionTrackedKey marks an SSH connection context as already counted.
type sshConnectionTrackedKey struct{}
// sshConnectionOpened tracks the authenticated TCP connection. Individual SSH
// sessions are short-lived and must not trigger a return to WebSocket.
//
// It is called from the session handler rather than the public key handler,
// which runs when a key is offered and before the client has proven it holds
// the private key. A connection is counted once however many sessions it opens.
func (c *ConnectionManager) sshConnectionOpened(ctx ssh.Context) {
ctx.Lock()
tracked := ctx.Value(sshConnectionTrackedKey{}) != nil
if !tracked {
ctx.SetValue(sshConnectionTrackedKey{}, true)
}
ctx.Unlock()
if tracked {
return
}
c.mu.Lock()
c.sshConnections++
first := c.sshConnections == 1
c.mu.Unlock()
if first {
c.notifySSHChange()
}
go func() {
<-ctx.Done()
c.mu.Lock()
c.sshConnections--
last := c.sshConnections == 0
c.mu.Unlock()
if last {
c.notifySSHChange()
}
}()
}
// setConnecting sets the isConnecting flag and reports its previous value.
func (c *ConnectionManager) setConnecting(v bool) (previous bool) {
c.mu.Lock()
defer c.mu.Unlock()
previous = c.isConnecting
c.isConnecting = v
return previous
}
// isConnectingNow reports whether a reconnection attempt is currently in flight.
func (c *ConnectionManager) isConnectingNow() bool {
c.mu.Lock()
defer c.mu.Unlock()
return c.isConnecting
}
// Start begins connection attempts and enters the main event loop.
// It handles connection events, periodic health updates, and graceful shutdown.
func (c *ConnectionManager) Start(serverOptions ServerOptions) error {
@@ -167,19 +87,7 @@ func (c *ConnectionManager) Start(serverOptions ServerOptions) error {
wsClient, err := newWebSocketClient(c.agent)
if err != nil {
var caCertErr *caCertFileError
if errors.As(err, &caCertErr) {
return err
}
disableSSH, _ := utils.GetEnv("DISABLE_SSH")
if errors.Is(err, errNoHubURL) && disableSSH != "true" {
// SSH-only mode: the hub dials the agent, so there is nothing to warn
// about. With SSH also disabled there is no connection method at all,
// so that case still warns.
slog.Debug("WebSocket client not configured", "err", err)
} else {
slog.Warn("Error creating WebSocket client", "err", err)
}
slog.Warn("Error creating WebSocket client", "err", err)
}
c.wsClient = wsClient
@@ -201,15 +109,8 @@ func (c *ConnectionManager) Start(serverOptions ServerOptions) error {
select {
case connectionEvent := <-c.eventChan:
c.handleEvent(connectionEvent)
case <-c.sshChanged:
c.handleSSHChange()
case <-c.wsTicker.C:
// skip if connect() is still running its own attempt
if !c.isConnectingNow() {
if err := c.startWebSocketConnection(); err != nil {
c.startSSHServer()
}
}
_ = c.startWebSocketConnection()
case <-healthTicker:
_ = health.Update()
case <-sigCtx.Done():
@@ -219,14 +120,6 @@ func (c *ConnectionManager) Start(serverOptions ServerOptions) error {
}
}
func (c *ConnectionManager) handleSSHChange() {
if c.hasSSHConnection() {
c.handleEvent(SSHConnect)
} else {
c.handleEvent(SSHDisconnect)
}
}
// stop does not stop the connection manager itself, just any active connections. The manager will attempt to reconnect after stopping, so this should only be called immediately before shutting down the entire agent.
//
// If we need or want to expose a graceful Stop method in the future, do something like this to actually stop the manager:
@@ -248,9 +141,8 @@ func (c *ConnectionManager) handleSSHChange() {
// }
func (c *ConnectionManager) stop() error {
_ = c.agent.StopServer()
c.agent.monitorManager.Stop()
c.agent.probeManager.Stop()
c.closeWebSocket()
c.agent.cleanupSensorShadow()
return health.CleanUp()
}
@@ -258,29 +150,15 @@ func (c *ConnectionManager) stop() error {
func (c *ConnectionManager) handleEvent(event ConnectionEvent) {
switch event {
case WebSocketConnect:
if c.wsClient == nil || !c.wsClient.isVerified() {
return // a superseded connection authenticated after a new attempt began
}
// WebSocket is preferred, so it takes over even if an attempt that was
// already in flight authenticates after SSH has connected.
c.handleStateChange(WebSocketConnected)
case SSHConnect:
if c.getState() == Disconnected && c.hasSSHConnection() {
c.handleStateChange(SSHConnected)
}
c.handleStateChange(SSHConnected)
case WebSocketDisconnect:
if c.wsClient != nil && c.wsClient.getConn() != nil {
return // an older connection closed after its replacement was installed
}
if c.getState() == WebSocketConnected {
if c.State == WebSocketConnected {
c.handleStateChange(Disconnected)
} else if c.getState() == Disconnected {
// The WebSocket upgrade can succeed before authentication fails.
// In that case Connect returned nil, so its error path cannot start SSH.
c.startSSHServer()
}
case SSHDisconnect:
if c.getState() == SSHConnected && !c.hasSSHConnection() {
if c.State == SSHConnected {
c.handleStateChange(Disconnected)
}
}
@@ -289,39 +167,30 @@ func (c *ConnectionManager) handleEvent(event ConnectionEvent) {
// handleStateChange updates the connection state and performs necessary actions
// based on the new state, including stopping services and initiating reconnections.
func (c *ConnectionManager) handleStateChange(newState ConnectionState) {
c.mu.Lock()
if c.State == newState {
c.mu.Unlock()
return
}
c.State = newState
c.mu.Unlock()
switch newState {
case WebSocketConnected:
slog.Info("WebSocket connected", "host", c.wsClient.hubURL.Host)
c.ConnectionType = system.ConnectionTypeWebSocket
c.stopWsTicker()
_ = c.agent.StopServer()
c.isConnecting = false
case SSHConnected:
// stop new ws connection attempts
slog.Info("SSH connection established")
c.ConnectionType = system.ConnectionTypeSSH
c.stopWsTicker()
c.isConnecting = false
case Disconnected:
// Listen for SSH whenever disconnected so the hub can fall back to it
// or redial straight away. WebSocket is still tried first below and
// stops the server if it connects.
c.startSSHServer()
// Always keep the ticker running while disconnected. A pending WebSocket
// handshake started by connect() can fail asynchronously (e.g. the hub
// closes the socket, or the deadline set in OnOpen expires) after
// connect() has already returned with a nil error, in which case the
// ticker would otherwise never get re-armed and the agent would stop
// retrying entirely (#2326).
c.startWsTicker()
if c.setConnecting(true) {
c.ConnectionType = system.ConnectionTypeNone
if c.isConnecting {
// Already handling reconnection, avoid duplicate attempts
return
}
c.isConnecting = true
slog.Warn("Disconnected from hub")
// make sure old ws connection is closed
c.closeWebSocket()
@@ -333,8 +202,10 @@ func (c *ConnectionManager) handleStateChange(newState ConnectionState) {
// connect handles the connection logic with proper delays and priority.
// It attempts WebSocket connection first, falling back to SSH server if needed.
func (c *ConnectionManager) connect() {
c.setConnecting(true)
defer c.setConnecting(false)
c.isConnecting = true
defer func() {
c.isConnecting = false
}()
if c.wsClient != nil && time.Since(c.wsClient.lastConnectAttempt) < 5*time.Second {
time.Sleep(5 * time.Second)
@@ -348,15 +219,16 @@ func (c *ConnectionManager) connect() {
_ = c.stop()
os.Exit(1)
}
if c.getState() == Disconnected {
if c.State == Disconnected {
c.startSSHServer()
c.startWsTicker()
}
}
}
// startWebSocketConnection attempts to establish a WebSocket connection to the hub.
func (c *ConnectionManager) startWebSocketConnection() error {
if c.getState() != Disconnected {
if c.State != Disconnected {
return errors.New("already connected")
}
if c.wsClient == nil {
@@ -376,28 +248,9 @@ func (c *ConnectionManager) startWebSocketConnection() error {
// startSSHServer starts the SSH server if the agent is currently disconnected.
func (c *ConnectionManager) startSSHServer() {
c.mu.Lock()
if c.State != Disconnected {
c.mu.Unlock()
return
if c.State == Disconnected {
go c.agent.StartServer(c.serverOptions)
}
if disabled, _ := utils.GetEnv("DISABLE_SSH"); disabled == "true" {
c.mu.Unlock()
return
}
server, listener, err := c.agent.prepareSSHServer(c.serverOptions)
c.mu.Unlock()
if err != nil {
if !errors.Is(err, errSSHServerRunning) {
slog.Warn("SSH server failed to start", "err", err)
}
return
}
go func() {
if err := c.agent.serveSSHServer(server, listener); err != nil && !errors.Is(err, ssh.ErrServerClosed) && !errors.Is(err, net.ErrClosed) {
slog.Warn("SSH server stopped", "err", err)
}
}()
}
// closeWebSocket closes the WebSocket connection if it exists.

View File

@@ -9,10 +9,8 @@ import (
"net"
"net/url"
"testing"
"testing/synctest"
"time"
"github.com/lxzan/gws"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/crypto/ssh"
@@ -79,10 +77,6 @@ func TestConnectionManager_StateTransitions(t *testing.T) {
cm.handleStateChange(SSHConnected)
assert.Equal(t, SSHConnected, cm.State, "State should change to SSHConnected")
// Prevent handleStateChange from spawning its async reconnect goroutine:
// this test only checks the synchronous state machine, and the goroutine
// would otherwise race with the direct field writes below.
cm.setConnecting(true)
cm.handleStateChange(Disconnected)
assert.Equal(t, Disconnected, cm.State, "State should change to Disconnected")
@@ -96,12 +90,12 @@ func TestConnectionManager_StateTransitions(t *testing.T) {
func TestConnectionManager_EventHandling(t *testing.T) {
agent := createTestAgent(t)
cm := agent.connectionManager
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "true")
cm.wsClient = &WebSocketClient{
hubURL: &url.URL{
Host: "localhost:8080",
},
}
testCases := []struct {
name string
initialState ConnectionState
@@ -114,24 +108,12 @@ func TestConnectionManager_EventHandling(t *testing.T) {
event: WebSocketConnect,
expectedState: WebSocketConnected,
},
{
name: "WebSocket connect from SSH connected",
initialState: SSHConnected,
event: WebSocketConnect,
expectedState: WebSocketConnected,
},
{
name: "SSH connect from disconnected",
initialState: Disconnected,
event: SSHConnect,
expectedState: SSHConnected,
},
{
name: "SSH connect from WebSocket connected (no change)",
initialState: WebSocketConnected,
event: SSHConnect,
expectedState: WebSocketConnected,
},
{
name: "WebSocket disconnect from connected",
initialState: WebSocketConnected,
@@ -160,25 +142,6 @@ func TestConnectionManager_EventHandling(t *testing.T) {
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
// Prevent handleStateChange from spawning its async reconnect
// goroutine: this test only checks the synchronous state machine,
// and the goroutine would otherwise race with the direct field
// writes here and in later subtests.
cm.setConnecting(true)
cm.mu.Lock()
cm.sshConnections = 0
if tc.event == SSHConnect {
cm.sshConnections = 1
}
cm.mu.Unlock()
cm.wsClient.connMu.Lock()
cm.wsClient.Conn = nil
cm.wsClient.hubVerified = false
if tc.event == WebSocketConnect {
cm.wsClient.Conn = &gws.Conn{}
cm.wsClient.hubVerified = true
}
cm.wsClient.connMu.Unlock()
cm.State = tc.initialState
cm.handleEvent(tc.event)
assert.Equal(t, tc.expectedState, cm.State, "State should match expected after event")
@@ -252,74 +215,12 @@ func TestConnectionManager_ReconnectionLogic(t *testing.T) {
// Test that isConnecting flag prevents duplicate reconnection attempts
// Start from connected state, then simulate disconnect
cm.State = WebSocketConnected
cm.setConnecting(false)
cm.isConnecting = false
// First disconnect should trigger reconnection logic
cm.handleStateChange(Disconnected)
assert.Equal(t, Disconnected, cm.State, "Should change to disconnected")
assert.True(t, cm.isConnectingNow(), "Should set isConnecting flag")
}
func TestWebSocketDisconnectStartsSSH(t *testing.T) {
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
agent := createTestAgent(t)
cm := agent.connectionManager
cm.serverOptions = createTestServerOptions(t)
cm.State = WebSocketConnected
cm.setConnecting(true) // keep this test focused on the synchronous fallback
defer cm.stopWsTicker()
cm.handleEvent(WebSocketDisconnect)
require.Equal(t, Disconnected, cm.getState())
agent.serverMu.Lock()
listener := agent.serverListener
agent.serverMu.Unlock()
require.NotNil(t, listener, "SSH should be ready as soon as an established WS closes")
require.NoError(t, agent.StopServer())
}
// TestConnectionManager_TickerSurvivesStaleDisconnect reproduces the freeze from
// https://github.com/henrygd/beszel/issues/2326: a reconnect attempt's handshake
// can fail asynchronously (after connect() already returned with a nil error)
// while the manager is still in the Disconnected state. Previously the ticker
// was only re-armed from connect()'s synchronous error branch, so once that
// window was missed, the agent stopped retrying forever. The ticker must keep
// running any time the manager transitions into Disconnected, regardless of
// what happens to the in-flight handshake afterwards.
func TestConnectionManager_TickerSurvivesStaleDisconnect(t *testing.T) {
agent := createTestAgent(t)
cm := agent.connectionManager
cm.eventChan = make(chan ConnectionEvent, 1)
// Run on synctest's fake clock so the ticker fires without waiting a real
// wsTickerInterval. The ticker must be created inside the bubble.
synctest.Test(t, func(t *testing.T) {
// Simulate a healthy WebSocket connection, then a disconnect - mirroring
// handleStateChange's own Disconnected branch, but without launching the
// real async connect() goroutine so the ticker state can be asserted
// deterministically.
cm.State = WebSocketConnected
cm.stopWsTicker()
cm.setConnecting(true)
cm.handleStateChange(Disconnected)
require.NotNil(t, cm.wsTicker, "ticker must be armed as soon as the manager becomes Disconnected")
defer cm.stopWsTicker()
// Now simulate connect()'s in-flight handshake dying asynchronously with the
// manager still Disconnected (e.g. a late OnClose on an unauthenticated
// connection). This event is dropped by handleEvent since State is not
// WebSocketConnected, but the ticker armed above must still be running so
// the manager keeps retrying.
cm.setConnecting(false)
cm.handleEvent(WebSocketDisconnect)
assert.Equal(t, Disconnected, cm.State)
select {
case <-cm.wsTicker.C:
case <-time.After(wsTickerInterval + 2*time.Second):
t.Fatal("ticker did not fire after a stale disconnect event - agent would freeze forever")
}
})
assert.True(t, cm.isConnecting, "Should set isConnecting flag")
}
// TestConnectionManager_ConnectWithRateLimit tests connection rate limiting
@@ -364,19 +265,6 @@ func TestConnectionManager_StartWithInvalidConfig(t *testing.T) {
assert.Error(t, err, "Should error when starting already started connection manager")
}
func TestConnectionManager_StartRejectsInvalidCACertFile(t *testing.T) {
agent := createTestAgent(t)
cm := agent.connectionManager
t.Setenv("BESZEL_AGENT_HUB_URL", "https://hub.example.com")
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
t.Setenv("BESZEL_AGENT_CA_CERT_FILE", t.TempDir())
err := cm.Start(ServerOptions{})
require.Error(t, err)
assert.Contains(t, err.Error(), "read CA_CERT_FILE")
assert.Nil(t, cm.eventChan)
}
// TestConnectionManager_CloseWebSocket tests WebSocket closing
func TestConnectionManager_CloseWebSocket(t *testing.T) {
agent := createTestAgent(t)

View File

@@ -30,20 +30,11 @@ type CpuMetrics struct {
Iowait float64
Steal float64
Idle float64
// fromCgroup is set when Total comes from cgroup accounting (LXC) rather
// than /proc/stat, so per-core /proc/stat usage would not match it.
fromCgroup bool
}
// getCpuMetrics calculates detailed CPU usage metrics using cached previous measurements.
// It returns percentages for total, user, system, iowait, and steal time.
func getCpuMetrics(cacheTimeMs uint16) (CpuMetrics, error) {
// Inside LXC, lxcfs serves /proc/stat with the host cores' counters, not
// the guest's own usage. Prefer the cgroup's CPU accounting there. (#2332)
if metrics, ok := containerCpuMetrics(cacheTimeMs); ok {
metrics.fromCgroup = true
return metrics, nil
}
times, err := cpu.Times(false)
if err != nil || len(times) == 0 {
return CpuMetrics{}, err
@@ -128,8 +119,7 @@ func calculateBusy(t1, t2 cpu.TimesStat) float64 {
// On Linux, it excludes guest and guest_nice time from the total to match kernel behavior.
// Returns total CPU time and busy CPU time (total minus idle and I/O wait time).
func getAllBusy(t cpu.TimesStat) (float64, float64) {
tot := t.User + t.System + t.Idle + t.Nice + t.Iowait + t.Irq +
t.Softirq + t.Steal + t.Guest + t.GuestNice
tot := t.Total()
if runtime.GOOS == "linux" {
tot -= t.Guest // Linux 2.6.24+
tot -= t.GuestNice // Linux 3.2.0+

View File

@@ -1,372 +0,0 @@
//go:build linux
package agent
import (
"os"
"path/filepath"
"runtime"
"strconv"
"strings"
"sync"
"time"
"github.com/henrygd/beszel/agent/utils"
)
// LXC-aware CPU accounting (issue #2332).
//
// Inside an LXC guest, lxcfs serves /proc/stat with the raw counters of the
// host cores in the guest's cpuset, not the guest's own usage. An idle guest
// sharing a host core with a busy neighbor then reports near-100% CPU while
// doing nothing. The cgroup's own accounting (cpu.stat / cpuacct.usage)
// reflects only the guest's processes, so inside LXC we derive CPU% from that
// instead.
//
// Other runtimes (Docker, Podman, k8s) are deliberately left alone: the agent
// is normally deployed there to monitor the host, and the host's /proc/stat is
// exactly what it should report.
// File paths and hooks are variables so tests can point them at fixtures.
var (
cpuCgroupRoot = "/sys/fs/cgroup" // default cgroup v2 mount point
cpuCgroupMountinfo = "/proc/self/mountinfo"
cpuProcSelfCgroup = "/proc/self/cgroup"
cpuSystemdContPath = "/run/systemd/container"
cpuNumCPU = runtime.NumCPU
cpuNow = time.Now
)
// cpuUserHZ is the USER_HZ jiffies-per-second rate cpuacct.stat reports in.
const cpuUserHZ = 100
// inLxc reports whether the agent itself runs inside an LXC guest.
// The result is cached because it cannot change during the process lifetime.
var inLxc = sync.OnceValue(detectLxc)
// detectLxc looks for LXC guest markers that are readable without root
// (the agent usually runs as an unprivileged user, so /proc/1/environ is not).
func detectLxc() bool {
// lxcfs mounted over /proc/stat is the direct cause of the host-core
// counters. Only match that mount point: an LXC host also has lxcfs
// mounted, but at /var/lib/lxcfs.
if data, err := os.ReadFile(cpuCgroupMountinfo); err == nil && procStatFromLxcfs(data) {
return true
}
// set by liblxc for the container init and inherited on non-systemd guests
if os.Getenv("container") == "lxc" {
return true
}
// written by systemd on systemd-based guests
if data, err := os.ReadFile(cpuSystemdContPath); err == nil &&
strings.TrimSpace(string(data)) == "lxc" {
return true
}
return false
}
// procStatFromLxcfs reports whether mountinfo shows lxcfs mounted on /proc/stat.
func procStatFromLxcfs(mountinfo []byte) bool {
for line := range strings.SplitSeq(string(mountinfo), "\n") {
left, right, found := strings.Cut(line, " - ")
if !found {
continue
}
fields, post := strings.Fields(left), strings.Fields(right)
if len(fields) >= 5 && len(post) > 0 && fields[4] == "/proc/stat" && post[0] == "fuse.lxcfs" {
return true
}
}
return false
}
// cgroupCpuSample is one read of the container's cumulative CPU accounting.
type cgroupCpuSample struct {
usageUsec uint64
userUsec uint64
systemUsec uint64
cores float64 // usable CPU cores: affinity ∩ cpuset ∩ quota
at time.Time
}
var lastCgroupCpuSamples = make(map[uint16]cgroupCpuSample)
// init seeds the LXC CPU baseline so the first reported value is a real
// delta since startup rather than zero.
func init() {
if !inLxc() {
return
}
if s, ok := readContainerCpuSample(); ok {
s.at = cpuNow()
lastCgroupCpuSamples[60000] = s
}
}
// containerCpuMetrics derives CPU metrics from the guest's own cgroup
// accounting when running inside LXC. It returns ok=false everywhere else and
// whenever cgroup accounting is unreadable, so callers keep the /proc/stat
// path.
func containerCpuMetrics(cacheTimeMs uint16) (CpuMetrics, bool) {
if !inLxc() {
return CpuMetrics{}, false
}
cur, ok := readContainerCpuSample()
if !ok {
return CpuMetrics{}, false
}
cur.at = cpuNow()
prev, ok := lastCgroupCpuSamples[cacheTimeMs]
if !ok {
prev = lastCgroupCpuSamples[60000]
}
lastCgroupCpuSamples[cacheTimeMs] = cur
// No baseline yet, a backwards counter (cgroup recreated), or a
// non-positive clock delta: report zero this tick instead of guessing.
elapsedUsec := cur.at.Sub(prev.at).Microseconds()
if prev.at.IsZero() || elapsedUsec <= 0 || cur.usageUsec < prev.usageUsec {
return CpuMetrics{}, true
}
cores := cur.cores
if cores <= 0 {
cores = 1
}
window := float64(elapsedUsec) * cores
metrics := CpuMetrics{
Total: clampPercent(float64(cur.usageUsec-prev.usageUsec) / window * 100),
User: clampPercent(float64(cur.userUsec-prev.userUsec) / window * 100),
System: clampPercent(float64(cur.systemUsec-prev.systemUsec) / window * 100),
}
// cgroup accounting has no iowait/steal; everything not busy is idle.
metrics.Idle = clampPercent(100 - metrics.Total)
return metrics, true
}
// readContainerCpuSample reads the container's cumulative CPU usage, preferring
// the cgroup v2 unified hierarchy and falling back to the v1 cpuacct
// controller.
func readContainerCpuSample() (cgroupCpuSample, bool) {
if s, ok := readCgroupV2CpuSample(); ok {
return s, true
}
return readCgroupV1CpuSample()
}
// readCgroupV2CpuSample reads usage from the unified hierarchy's cpu.stat.
//
// The mount root is always the cgroup to read: an LXC guest has a private
// cgroup namespace, so /sys/fs/cgroup already is the guest's root cgroup, and
// its cpu.stat accounts for every process in the guest. The agent's own path
// in /proc/self/cgroup (its service cgroup, or the ".lxc" leaf when started
// from an attached shell) only covers a subset and must not be descended into.
func readCgroupV2CpuSample() (cgroupCpuSample, bool) {
if !inCgroupV2() {
return cgroupCpuSample{}, false // no v2 membership; try v1
}
dir := cpuCgroupRoot
if mount := cgroupMountPoint("cgroup2", ""); mount != "" {
dir = mount
}
stat := filepath.Join(dir, "cpu.stat")
usage, ok := cgroupStatValue(stat, "usage_usec")
if !ok {
return cgroupCpuSample{}, false
}
s := cgroupCpuSample{usageUsec: usage, cores: cpuCgroupCores(dir)}
s.userUsec, _ = cgroupStatValue(stat, "user_usec")
s.systemUsec, _ = cgroupStatValue(stat, "system_usec")
return s, true
}
// readCgroupV1CpuSample reads usage from the legacy cpuacct controller.
// As with v2, the hierarchy mount root is the guest's own cgroup and its
// accounting includes every child cgroup, so it is read directly rather than
// the agent's own sub-cgroup.
func readCgroupV1CpuSample() (cgroupCpuSample, bool) {
dir := cgroupMountPoint("cgroup", "cpuacct")
if dir == "" {
return cgroupCpuSample{}, false
}
usageNs, ok := utils.ReadUintFile(filepath.Join(dir, "cpuacct.usage"))
if !ok {
return cgroupCpuSample{}, false
}
s := cgroupCpuSample{usageUsec: usageNs / 1000, cores: cpuCgroupCores(dir)}
// cpuacct.stat reports user/system in USER_HZ jiffies.
if v, ok := cgroupStatValue(filepath.Join(dir, "cpuacct.stat"), "user"); ok {
s.userUsec = v * 1e6 / cpuUserHZ
}
if v, ok := cgroupStatValue(filepath.Join(dir, "cpuacct.stat"), "system"); ok {
s.systemUsec = v * 1e6 / cpuUserHZ
}
return s, true
}
// inCgroupV2 reports whether /proc/self/cgroup lists the v2 unified hierarchy
// (a "0::<path>" entry).
func inCgroupV2() bool {
data, err := os.ReadFile(cpuProcSelfCgroup)
if err != nil {
return false
}
for line := range strings.SplitSeq(string(data), "\n") {
if strings.HasPrefix(line, "0::") {
return true
}
}
return false
}
// cgroupMountPoint returns the mount point of a cgroup hierarchy from
// /proc/self/mountinfo: the cgroup2 mount for v2, or the cgroup mount whose
// super options list the wanted v1 controller.
func cgroupMountPoint(fstype, v1ctrl string) string {
data, err := os.ReadFile(cpuCgroupMountinfo)
if err != nil {
return ""
}
for _, line := range strings.Split(string(data), "\n") {
left, right, found := strings.Cut(line, " - ")
if !found {
continue
}
post := strings.Fields(right)
if len(post) == 0 || post[0] != fstype {
continue
}
if v1ctrl != "" && !mountOptHas(post, v1ctrl) {
continue
}
fields := strings.Fields(left)
if len(fields) >= 5 {
return unescapeMountPoint(fields[4])
}
}
return ""
}
// mountOptHas reports whether the comma-separated super options (field 3 after
// the " - " separator) contain opt.
func mountOptHas(post []string, opt string) bool {
if len(post) < 3 {
return false
}
for _, o := range strings.Split(post[2], ",") {
if o == opt {
return true
}
}
return false
}
// unescapeMountPoint decodes octal escapes (e.g. \040 for space) used in
// mountinfo paths.
func unescapeMountPoint(s string) string {
return strings.NewReplacer(`\040`, " ", `\011`, "\t", `\012`, "\n", `\134`, `\`).Replace(s)
}
// cpuCgroupCores returns how many CPU cores the cgroup at dir may use: the
// smallest of the process affinity mask, the cgroup cpuset, and the CPU quota.
func cpuCgroupCores(dir string) float64 {
cores := float64(cpuNumCPU())
if n := cpusetCount(dir); n > 0 && n < cores {
cores = n
}
if q, ok := cpuQuotaCores(dir); ok && q < cores {
cores = q
}
if cores <= 0 {
cores = 1
}
return cores
}
// cpusetCount returns the number of CPUs in the cgroup's cpuset, e.g. "0-3" or
// "2,5-7". An empty or missing file means unconstrained.
func cpusetCount(dir string) float64 {
for _, name := range []string{"cpuset.cpus.effective", "cpuset.cpus"} {
raw, err := os.ReadFile(filepath.Join(dir, name))
if err != nil {
continue
}
if n := countCpuList(strings.TrimSpace(string(raw))); n > 0 {
return float64(n)
}
}
return 0
}
// countCpuList counts the CPUs in a Linux CPU list like "0-3,5,8-9".
func countCpuList(list string) int {
total := 0
for part := range strings.SplitSeq(list, ",") {
lo, hi, ranged := strings.Cut(part, "-")
a, err := strconv.Atoi(lo)
if err != nil {
continue
}
b := a
if ranged {
if v, err := strconv.Atoi(hi); err == nil {
b = v
}
}
if b >= a {
total += b - a + 1
}
}
return total
}
// cpuQuotaCores returns the cgroup's CPU quota in cores. v2 uses cpu.max
// ("<quota|max> <period>"), v1 uses cpu.cfs_quota_us / cpu.cfs_period_us.
func cpuQuotaCores(dir string) (float64, bool) {
if raw, err := os.ReadFile(filepath.Join(dir, "cpu.max")); err == nil {
fields := strings.Fields(string(raw))
if len(fields) == 2 && fields[0] != "max" {
if quota, err := strconv.ParseFloat(fields[0], 64); err == nil && quota > 0 {
if period, err := strconv.ParseFloat(fields[1], 64); err == nil && period > 0 {
return quota / period, true
}
}
}
}
if quota, ok := readCgroupInt(filepath.Join(dir, "cpu.cfs_quota_us")); ok && quota > 0 {
if period, ok := readCgroupInt(filepath.Join(dir, "cpu.cfs_period_us")); ok && period > 0 {
return float64(quota) / float64(period), true
}
}
return 0, false
}
// cgroupStatValue returns the value of key in a cgroup "key value" stat file.
func cgroupStatValue(path, key string) (uint64, bool) {
data, err := os.ReadFile(path)
if err != nil {
return 0, false
}
for line := range strings.SplitSeq(string(data), "\n") {
name, value, found := strings.Cut(line, " ")
if !found || name != key {
continue
}
v, err := strconv.ParseUint(strings.TrimSpace(value), 10, 64)
return v, err == nil
}
return 0, false
}
// readCgroupInt reads a file containing a single signed integer
// (cpu.cfs_quota_us is -1 when no quota is set).
func readCgroupInt(path string) (int64, bool) {
data, err := os.ReadFile(path)
if err != nil {
return 0, false
}
v, err := strconv.ParseInt(strings.TrimSpace(string(data)), 10, 64)
return v, err == nil
}

View File

@@ -1,350 +0,0 @@
//go:build testing && linux
package agent
import (
"os"
"path/filepath"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// swapCpuContainerSeams points every container-detection and cgroup path at
// empty fixtures under a temp dir, then restores them on cleanup.
func swapCpuContainerSeams(t *testing.T) {
t.Helper()
backup := struct {
root, mountinfo, selfCgroup, systemdCont string
numCPU func() int
now func() time.Time
}{
cpuCgroupRoot, cpuCgroupMountinfo, cpuProcSelfCgroup, cpuSystemdContPath, cpuNumCPU, cpuNow,
}
origInLxc := inLxc
samples := lastCgroupCpuSamples
env, hadEnv := os.LookupEnv("container")
t.Cleanup(func() {
cpuCgroupRoot, cpuCgroupMountinfo, cpuProcSelfCgroup, cpuSystemdContPath = backup.root, backup.mountinfo, backup.selfCgroup, backup.systemdCont
cpuNumCPU, cpuNow = backup.numCPU, backup.now
inLxc = origInLxc
lastCgroupCpuSamples = samples
if hadEnv {
os.Setenv("container", env)
}
})
inLxc = sync.OnceValue(detectLxc)
lastCgroupCpuSamples = make(map[uint16]cgroupCpuSample)
os.Unsetenv("container")
tmp := t.TempDir()
cpuCgroupRoot = filepath.Join(tmp, "cgroup")
cpuCgroupMountinfo = filepath.Join(tmp, "mountinfo")
cpuProcSelfCgroup = filepath.Join(tmp, "self-cgroup")
cpuSystemdContPath = filepath.Join(tmp, "systemd-container")
}
func writeCpuFixture(t *testing.T, path, contents string) {
t.Helper()
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o755))
require.NoError(t, os.WriteFile(path, []byte(contents), 0o644))
}
// markLxc makes detection see a systemd-based LXC guest.
func markLxc(t *testing.T) {
t.Helper()
writeCpuFixture(t, cpuSystemdContPath, "lxc\n")
}
// fakeNow installs a controllable clock and returns a function to advance it.
func fakeNow(t *testing.T) func(time.Duration) {
t.Helper()
cur := time.Unix(1_700_000_000, 0)
cpuNow = func() time.Time { return cur }
return func(d time.Duration) { cur = cur.Add(d) }
}
func TestDetectLxc(t *testing.T) {
tests := []struct {
name string
setup func(t *testing.T)
want bool
}{
{"plain host", func(t *testing.T) {}, false},
{"container env lxc", func(t *testing.T) { t.Setenv("container", "lxc") }, true},
{"container env podman", func(t *testing.T) { t.Setenv("container", "podman") }, false},
{"systemd container lxc", func(t *testing.T) { writeCpuFixture(t, cpuSystemdContPath, "lxc\n") }, true},
{"systemd container nspawn", func(t *testing.T) { writeCpuFixture(t, cpuSystemdContPath, "systemd-nspawn\n") }, false},
{"lxcfs serving /proc/stat", func(t *testing.T) {
writeCpuFixture(t, cpuCgroupMountinfo,
"31 25 0:28 / /proc/stat rw,nosuid,nodev,relatime - fuse.lxcfs lxcfs rw,user_id=0,group_id=0\n")
}, true},
// an LXC host (e.g. Proxmox) mounts lxcfs too, but not over its own /proc
{"lxcfs mounted on host", func(t *testing.T) {
writeCpuFixture(t, cpuCgroupMountinfo,
"45 25 0:40 / /var/lib/lxcfs rw,nosuid,nodev,relatime - fuse.lxcfs lxcfs rw,user_id=0,group_id=0\n")
}, false},
{"cgroup-only mountinfo", func(t *testing.T) {
writeCpuFixture(t, cpuCgroupMountinfo,
"36 25 0:32 / /sys/fs/cgroup rw - cgroup2 cgroup2 rw,nsdelegate\n")
}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
swapCpuContainerSeams(t)
tt.setup(t)
assert.Equal(t, tt.want, detectLxc())
})
}
}
func TestReadCgroupV2CpuSample(t *testing.T) {
swapCpuContainerSeams(t)
writeCpuFixture(t, cpuProcSelfCgroup, "0::/\n")
writeCpuFixture(t, cpuCgroupMountinfo, "")
require.NoError(t, os.MkdirAll(cpuCgroupRoot, 0o755))
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"),
"usage_usec 3000000\nuser_usec 2000000\nsystem_usec 1000000\nnr_throttled 7\n")
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpuset.cpus.effective"), "2,5-7\n")
cpuNumCPU = func() int { return 8 }
s, ok := readCgroupV2CpuSample()
require.True(t, ok)
assert.EqualValues(t, 3000000, s.usageUsec)
assert.EqualValues(t, 2000000, s.userUsec)
assert.EqualValues(t, 1000000, s.systemUsec)
assert.InDelta(t, 4, s.cores, 0.001) // cpuset 2,5-7 = 4 cores
}
// The agent may sit in a sub-cgroup of the guest (a systemd service, or the
// ".lxc" leaf when started from an attached shell); the mount root still
// accounts for the whole guest and must win.
func TestReadCgroupV2PrefersContainerRoot(t *testing.T) {
for _, rel := range []string{"system.slice/beszel-agent.service", ".lxc"} {
t.Run(rel, func(t *testing.T) {
swapCpuContainerSeams(t)
writeCpuFixture(t, cpuProcSelfCgroup, "0::/"+rel+"\n")
writeCpuFixture(t, cpuCgroupMountinfo, "")
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"), "usage_usec 9000\n")
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, rel, "cpu.stat"), "usage_usec 5\n")
s, ok := readCgroupV2CpuSample()
require.True(t, ok)
assert.EqualValues(t, 9000, s.usageUsec)
})
}
}
// Same for v1: the cpuacct mount root covers the agent's sibling services.
func TestReadCgroupV1PrefersContainerRoot(t *testing.T) {
swapCpuContainerSeams(t)
writeCpuFixture(t, cpuProcSelfCgroup, "3:cpu,cpuacct:/system.slice/beszel-agent.service\n")
v1 := filepath.Join(t.TempDir(), "cpu,cpuacct")
writeCpuFixture(t, cpuCgroupMountinfo,
"30 25 0:26 / "+v1+" rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,cpu,cpuacct\n")
writeCpuFixture(t, filepath.Join(v1, "cpuacct.usage"), "9000000\n")
writeCpuFixture(t, filepath.Join(v1, "system.slice/beszel-agent.service/cpuacct.usage"), "5000\n")
s, ok := readContainerCpuSample()
require.True(t, ok)
assert.EqualValues(t, 9000, s.usageUsec)
}
func TestReadCgroupV1CpuSample(t *testing.T) {
swapCpuContainerSeams(t)
writeCpuFixture(t, cpuProcSelfCgroup, "3:cpuacct:/\n2:memory:/\n")
v1 := filepath.Join(t.TempDir(), "cpuacct")
writeCpuFixture(t, cpuCgroupMountinfo,
"30 25 0:26 / "+v1+" rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,cpuacct\n")
writeCpuFixture(t, filepath.Join(v1, "cpuacct.usage"), "2000000000\n")
writeCpuFixture(t, filepath.Join(v1, "cpuacct.stat"), "user 100\nsystem 50\n")
cpuNumCPU = func() int { return 4 }
s, ok := readContainerCpuSample() // no 0:: line -> falls through to v1
require.True(t, ok)
assert.EqualValues(t, 2000000, s.usageUsec) // ns -> usec
assert.EqualValues(t, 1000000, s.userUsec) // 100 jiffies * 1e6/100
assert.EqualValues(t, 500000, s.systemUsec) // 50 jiffies
assert.InDelta(t, 4, s.cores, 0.001)
}
func TestContainerCpuMetricsMath(t *testing.T) {
swapCpuContainerSeams(t)
markLxc(t)
writeCpuFixture(t, cpuProcSelfCgroup, "0::/\n")
writeCpuFixture(t, cpuCgroupMountinfo, "")
require.NoError(t, os.MkdirAll(cpuCgroupRoot, 0o755))
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpuset.cpus.effective"), "0-3\n")
cpuNumCPU = func() int { return 8 }
advance := fakeNow(t)
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"),
"usage_usec 1000000\nuser_usec 600000\nsystem_usec 400000\n")
m, ok := containerCpuMetrics(60000)
require.True(t, ok)
assert.Zero(t, m.Total) // first call only seeds the baseline
// 1s elapsed, container burned 2 core-seconds on 4 usable cores
advance(time.Second)
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"),
"usage_usec 3000000\nuser_usec 1600000\nsystem_usec 900000\n")
m, ok = containerCpuMetrics(60000)
require.True(t, ok)
assert.InDelta(t, 50, m.Total, 0.01)
assert.InDelta(t, 25, m.User, 0.01)
assert.InDelta(t, 12.5, m.System, 0.01)
assert.Zero(t, m.Iowait)
assert.Zero(t, m.Steal)
assert.InDelta(t, 50, m.Idle, 0.01)
}
func TestContainerCpuMetricsHonorsQuota(t *testing.T) {
swapCpuContainerSeams(t)
markLxc(t)
writeCpuFixture(t, cpuProcSelfCgroup, "0::/\n")
writeCpuFixture(t, cpuCgroupMountinfo, "")
require.NoError(t, os.MkdirAll(cpuCgroupRoot, 0o755))
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.max"), "200000 100000\n") // 2 cores
cpuNumCPU = func() int { return 8 }
advance := fakeNow(t)
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"), "usage_usec 1000000\n")
containerCpuMetrics(60000)
advance(time.Second)
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"), "usage_usec 2000000\n")
m, ok := containerCpuMetrics(60000)
require.True(t, ok)
assert.InDelta(t, 50, m.Total, 0.01) // 1 core-second against a 2-core quota
}
func TestContainerCpuMetricsZeroAndBackwardDelta(t *testing.T) {
swapCpuContainerSeams(t)
markLxc(t)
writeCpuFixture(t, cpuProcSelfCgroup, "0::/\n")
writeCpuFixture(t, cpuCgroupMountinfo, "")
require.NoError(t, os.MkdirAll(cpuCgroupRoot, 0o755))
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"), "usage_usec 5000000\n")
cpuNumCPU = func() int { return 4 }
advance := fakeNow(t)
// seed the baseline, then do not advance the clock: elapsed <= 0
containerCpuMetrics(60000)
m, ok := containerCpuMetrics(60000)
require.True(t, ok)
assert.Zero(t, m.Total)
// counter goes backwards (cgroup recreated): report zero and re-baseline
advance(time.Second)
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"), "usage_usec 100000\n")
m, ok = containerCpuMetrics(60000)
require.True(t, ok)
assert.Zero(t, m.Total)
// next tick measures from the new baseline, not the stale one
advance(time.Second)
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"), "usage_usec 1100000\n")
m, ok = containerCpuMetrics(60000)
require.True(t, ok)
assert.InDelta(t, 25, m.Total, 0.01) // 1e6 usec / (1s * 4 cores)
}
func TestContainerCpuMetricsFallbacks(t *testing.T) {
t.Run("not in lxc", func(t *testing.T) {
swapCpuContainerSeams(t)
_, ok := containerCpuMetrics(60000)
assert.False(t, ok)
})
// Docker agents monitor the host, so cgroup accounting must not kick in
// even when it is readable and no LXC marker is present.
t.Run("docker container", func(t *testing.T) {
swapCpuContainerSeams(t)
writeCpuFixture(t, cpuProcSelfCgroup, "0::/\n")
writeCpuFixture(t, cpuCgroupMountinfo, "")
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"), "usage_usec 5000000\n")
_, ok := containerCpuMetrics(60000)
assert.False(t, ok)
})
t.Run("in lxc without cgroup accounting", func(t *testing.T) {
swapCpuContainerSeams(t)
markLxc(t)
writeCpuFixture(t, cpuProcSelfCgroup, "0::/\n")
writeCpuFixture(t, cpuCgroupMountinfo, "")
// cpuCgroupRoot has no cpu.stat
_, ok := containerCpuMetrics(60000)
assert.False(t, ok)
})
}
// The host path must keep reporting through gopsutil untouched.
func TestGetCpuMetricsHostFallback(t *testing.T) {
swapCpuContainerSeams(t)
m, err := getCpuMetrics(60000)
require.NoError(t, err)
assert.False(t, m.fromCgroup)
assert.GreaterOrEqual(t, m.Total, 0.0)
assert.LessOrEqual(t, m.Total, 100.0)
}
// Inside LXC getCpuMetrics must report the cgroup-derived value, not
// the host core counters from /proc/stat.
func TestGetCpuMetricsPrefersCgroup(t *testing.T) {
swapCpuContainerSeams(t)
markLxc(t)
writeCpuFixture(t, cpuProcSelfCgroup, "0::/\n")
writeCpuFixture(t, cpuCgroupMountinfo, "")
require.NoError(t, os.MkdirAll(cpuCgroupRoot, 0o755))
cpuNumCPU = func() int { return 4 }
advance := fakeNow(t)
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"), "usage_usec 0\n")
m, err := getCpuMetrics(60000)
require.NoError(t, err)
assert.Zero(t, m.Total)
advance(time.Second)
writeCpuFixture(t, filepath.Join(cpuCgroupRoot, "cpu.stat"), "usage_usec 2000000\n")
m, err = getCpuMetrics(60000)
require.NoError(t, err)
assert.InDelta(t, 50, m.Total, 0.01)
assert.True(t, m.fromCgroup) // per-core usage is skipped for this source
}
func TestCountCpuList(t *testing.T) {
assert.Equal(t, 4, countCpuList("0-3"))
assert.Equal(t, 4, countCpuList("2,5-7"))
assert.Equal(t, 1, countCpuList("2"))
assert.Equal(t, 0, countCpuList(""))
assert.Equal(t, 0, countCpuList("max"))
assert.Equal(t, 6, countCpuList("0-3,8-9"))
}
func TestCpuQuotaCores(t *testing.T) {
dir := t.TempDir()
_, ok := cpuQuotaCores(dir)
assert.False(t, ok) // no quota files
writeCpuFixture(t, filepath.Join(dir, "cpu.max"), "max 100000\n")
_, ok = cpuQuotaCores(dir)
assert.False(t, ok) // unlimited
writeCpuFixture(t, filepath.Join(dir, "cpu.max"), "150000 100000\n")
q, ok := cpuQuotaCores(dir)
require.True(t, ok)
assert.InDelta(t, 1.5, q, 0.001)
// v1 files
v1 := t.TempDir()
writeCpuFixture(t, filepath.Join(v1, "cpu.cfs_quota_us"), "-1\n")
writeCpuFixture(t, filepath.Join(v1, "cpu.cfs_period_us"), "100000\n")
_, ok = cpuQuotaCores(v1)
assert.False(t, ok)
writeCpuFixture(t, filepath.Join(v1, "cpu.cfs_quota_us"), "50000\n")
q, ok = cpuQuotaCores(v1)
require.True(t, ok)
assert.InDelta(t, 0.5, q, 0.001)
}

View File

@@ -1,28 +0,0 @@
//go:build testing
package agent
import (
"runtime"
"testing"
"github.com/shirou/gopsutil/v4/cpu"
"github.com/stretchr/testify/assert"
)
func TestGetAllBusy(t *testing.T) {
times := cpu.TimesStat{
User: 1, System: 2, Idle: 3, Nice: 4, Iowait: 5,
Irq: 6, Softirq: 7, Steal: 8, Guest: 9, GuestNice: 10,
}
wantTotal, wantBusy := 55.0, 47.0
if runtime.GOOS == "linux" {
wantTotal, wantBusy = 36, 28
}
total, busy := getAllBusy(times)
assert.Equal(t, wantTotal, total)
assert.Equal(t, wantBusy, busy)
assert.InDelta(t, wantBusy/wantTotal*100, calculateBusy(cpu.TimesStat{}, times), 1e-10)
assert.Zero(t, calculateBusy(times, times))
assert.Zero(t, calculateBusy(times, cpu.TimesStat{}))
}

View File

@@ -1,9 +0,0 @@
//go:build !linux
package agent
// containerCpuMetrics is Linux-only (cgroup accounting). Other platforms keep
// the gopsutil /proc path.
func containerCpuMetrics(uint16) (CpuMetrics, bool) {
return CpuMetrics{}, false
}

View File

@@ -12,14 +12,6 @@ import (
"github.com/stretchr/testify/require"
)
func invalidDataDir(t *testing.T) string {
t.Helper()
filePath := filepath.Join(t.TempDir(), "file")
require.NoError(t, os.WriteFile(filePath, nil, 0644))
return filepath.Join(filePath, "data")
}
func TestGetDataDir(t *testing.T) {
// Test with explicit dataDir parameter
t.Run("explicit data dir", func(t *testing.T) {
@@ -56,7 +48,7 @@ func TestGetDataDir(t *testing.T) {
// Test with invalid explicit dataDir
t.Run("invalid explicit data dir", func(t *testing.T) {
invalidPath := invalidDataDir(t)
invalidPath := "/invalid/path/that/cannot/be/created"
_, err := GetDataDir(invalidPath)
assert.Error(t, err)
})
@@ -86,7 +78,7 @@ func TestTestDataDirs(t *testing.T) {
// Test with multiple directories, first one valid
t.Run("multiple dirs - first valid", func(t *testing.T) {
tempDir := t.TempDir()
invalidDir := invalidDataDir(t)
invalidDir := "/invalid/path"
result, err := testDataDirs([]string{tempDir, invalidDir})
require.NoError(t, err)
assert.Equal(t, tempDir, result)
@@ -95,7 +87,7 @@ func TestTestDataDirs(t *testing.T) {
// Test with multiple directories, second one valid
t.Run("multiple dirs - second valid", func(t *testing.T) {
tempDir := t.TempDir()
invalidDir := invalidDataDir(t)
invalidDir := "/invalid/path"
result, err := testDataDirs([]string{invalidDir, tempDir})
require.NoError(t, err)
assert.Equal(t, tempDir, result)
@@ -117,7 +109,7 @@ func TestTestDataDirs(t *testing.T) {
// Test with no valid directories
t.Run("no valid directories", func(t *testing.T) {
invalidPaths := []string{invalidDataDir(t), invalidDataDir(t)}
invalidPaths := []string{"/invalid/path1", "/invalid/path2"}
_, err := testDataDirs(invalidPaths)
assert.Error(t, err)
assert.Contains(t, err.Error(), "data directory not found")

View File

@@ -3,7 +3,6 @@ package agent
import (
"context"
"log/slog"
"math"
"os"
"path/filepath"
"runtime"
@@ -19,8 +18,7 @@ import (
// fsRegistrationContext holds the shared lookup state needed to resolve a
// filesystem into the tracked fsStats key and metadata.
type fsRegistrationContext struct {
filesystem string // device part of optional FILESYSTEM env var
filesystemName string // optional custom name from FILESYSTEM=device__name
filesystem string // value of optional FILESYSTEM env var
isWindows bool
efPath string // path to extra filesystems (default "/extra-filesystems")
diskIoCounters map[string]disk.IOCountersStat
@@ -97,41 +95,21 @@ func isDockerSpecialMountpoint(mountpoint string) bool {
return false
}
// evalSymlinks resolves device symlinks; it is a seam so tests can fake the
// /dev topology (e.g. /dev/vg/lv -> /dev/dm-N created by udev for LVM).
var evalSymlinks = filepath.EvalSymlinks
// registerFilesystemStats resolves the tracked key and stats payload for a
// filesystem before it is inserted into fsStats.
func registerFilesystemStats(existing map[string]*system.FsStats, device, mountpoint string, root bool, customName string, ctx fsRegistrationContext) (string, *system.FsStats, bool) {
key := device
resolvedKey := ""
if !ctx.isWindows {
key = filepath.Base(device)
// Device-mapper mounts appear as symlinked paths like /dev/vg/lv whose
// base name matches neither the diskstats name (dm-N) nor the dm label
// (vg-lv); the resolved target's base is one of those existing names.
// Bare names (folder devices, ZFS datasets) are skipped because they
// would resolve relative to the agent's working directory.
if filepath.IsAbs(device) {
if resolved, err := evalSymlinks(device); err == nil {
if base := filepath.Base(resolved); base != key {
resolvedKey = base
}
}
}
}
if root {
// Try to map root device to a diskIoCounters entry. First checks for an
// exact key match, then uses findIoDevice for normalized / prefix-based
// matching (e.g. nda0p2 -> nda0) and the symlink-resolved device name,
// and finally falls back to FILESYSTEM.
// matching (e.g. nda0p2 -> nda0), and finally falls back to FILESYSTEM.
if _, ioMatch := ctx.diskIoCounters[key]; !ioMatch {
if matchedKey, match := findIoDevice(key, ctx.diskIoCounters); match {
key = matchedKey
} else if matchedKey, match := findIoDevice(resolvedKey, ctx.diskIoCounters); match {
key = matchedKey
} else if ctx.filesystem != "" {
if matchedKey, match := findIoDevice(ctx.filesystem, ctx.diskIoCounters); match {
key = matchedKey
@@ -157,8 +135,6 @@ func registerFilesystemStats(existing map[string]*system.FsStats, device, mountp
if _, ioMatch = ctx.diskIoCounters[key]; !ioMatch {
if matchedKey, match := findIoDevice(key, ctx.diskIoCounters); match {
key = matchedKey
} else if matchedKey, match := findIoDevice(resolvedKey, ctx.diskIoCounters); match {
key = matchedKey
}
}
}
@@ -176,12 +152,12 @@ func registerFilesystemStats(existing map[string]*system.FsStats, device, mountp
}
// addFsStat inserts a discovered filesystem if it resolves to a new tracking
// key and reports whether it was added. The key selection itself lives in
// registerFilesystemStats so that logic can stay directly unit-tested.
func (d *diskDiscovery) addFsStat(device, mountpoint string, root bool, customName string) bool {
// key. The key selection itself lives in buildFsStatRegistration so that logic
// can stay directly unit-tested.
func (d *diskDiscovery) addFsStat(device, mountpoint string, root bool, customName string) {
key, fsStats, ok := registerFilesystemStats(d.agent.fsStats, device, mountpoint, root, customName, d.ctx)
if !ok {
return false
return
}
d.agent.fsStats[key] = fsStats
name := key
@@ -189,7 +165,6 @@ func (d *diskDiscovery) addFsStat(device, mountpoint string, root bool, customNa
name = customName
}
slog.Info("Detected disk", "name", name, "device", device, "mount", mountpoint, "io", key, "root", root)
return true
}
// addConfiguredRootFs resolves FILESYSTEM against partitions first, then falls
@@ -202,7 +177,7 @@ func (d *diskDiscovery) addConfiguredRootFs() bool {
for _, p := range d.partitions {
if filesystemMatchesPartitionSetting(d.ctx.filesystem, p) {
d.addFsStat(p.Device, p.Mountpoint, true, d.ctx.filesystemName)
d.addFsStat(p.Device, p.Mountpoint, true, "")
return true
}
}
@@ -210,7 +185,7 @@ func (d *diskDiscovery) addConfiguredRootFs() bool {
// FILESYSTEM may name a physical disk absent from partitions (e.g. ZFS lists
// dataset paths like zroot/ROOT/default, not block devices).
if ioKey, match := findIoDevice(d.ctx.filesystem, d.ctx.diskIoCounters); match {
d.agent.fsStats[ioKey] = &system.FsStats{Root: true, Mountpoint: d.rootMountPoint, Name: d.ctx.filesystemName}
d.agent.fsStats[ioKey] = &system.FsStats{Root: true, Mountpoint: d.rootMountPoint}
return true
}
@@ -227,24 +202,14 @@ func isRootFallbackPartition(p disk.PartitionStat, rootMountPoint string) bool {
// partition looks like the active root mount but still needs translating to an
// I/O device key.
func (d *diskDiscovery) addPartitionRootFs(device, mountpoint string) bool {
// device is passed through as-is: findIoDevice normalizes it, and
// filepath.Base would turn a Windows volume name such as "C:" into "\"
// on the way in (#2417).
fs, match := findIoDevice(device, d.ctx.diskIoCounters)
fs, match := findIoDevice(filepath.Base(device), d.ctx.diskIoCounters)
if !match {
return false
}
// The root device is already resolved, so if it was registered earlier as an
// extra filesystem (e.g. root drive listed in EXTRA_FILESYSTEMS), promote that
// entry rather than letting addLastResortRootFs guess a different device.
if stats, exists := d.agent.fsStats[fs]; exists {
stats.Root = true
stats.Mountpoint = mountpoint
return true
}
// Use the resolved I/O device directly to avoid a second fallback search
// inside registerFilesystemStats.
return d.addFsStat(fs, mountpoint, true, "")
// The resolved I/O device is already known here, so use it directly to avoid
// a second fallback search inside buildFsStatRegistration.
d.addFsStat(fs, mountpoint, true, "")
return true
}
// addLastResortRootFs is only used when neither FILESYSTEM nor partition-based
@@ -335,8 +300,7 @@ func (d *diskDiscovery) addExtraFilesystemFolders(folderNames []string) {
// Sets up the filesystems to monitor for disk usage and I/O.
func (a *Agent) initializeDiskInfo() {
filesystemRaw, _ := utils.GetEnv("FILESYSTEM")
filesystem, filesystemName := parseFilesystemEntry(filesystemRaw)
filesystem, _ := utils.GetEnv("FILESYSTEM")
hasRoot := false
isWindows := runtime.GOOS == "windows"
@@ -360,7 +324,6 @@ func (a *Agent) initializeDiskInfo() {
slog.Debug("Disk I/O", "diskstats", diskIoCounters)
ctx := fsRegistrationContext{
filesystem: filesystem,
filesystemName: filesystemName,
isWindows: isWindows,
diskIoCounters: diskIoCounters,
efPath: "/extra-filesystems",
@@ -560,57 +523,18 @@ func filesystemMatchesPartitionSetting(filesystem string, p disk.PartitionStat)
// normalizeDeviceName canonicalizes device strings for comparisons.
func normalizeDeviceName(value string) string {
name := strings.TrimSpace(value)
if volume, ok := windowsVolumeName(name); ok {
return volume
}
name = filepath.Base(name)
name := filepath.Base(strings.TrimSpace(value))
if name == "." {
return ""
}
return name
}
// windowsVolumeName returns the canonical form of a bare Windows volume
// specifier, so that "C:", "c:", `C:\` and "C:/" all name the same drive.
// Drive letters are case-insensitive on Windows, so the letter is uppercased.
//
// filepath.Base cannot do this. On Windows it treats "C:" as a volume name
// with no path element to take the base of and returns "\", so every drive
// letter normalizes to the same key. findIoDevice then returns whichever
// counter the map happened to yield first, which registers the root
// filesystem under a random drive (#2417).
func windowsVolumeName(value string) (string, bool) {
if len(value) < 2 || value[1] != ':' {
return "", false
}
if c := value[0]; !('a' <= c && c <= 'z' || 'A' <= c && c <= 'Z') {
return "", false
}
// Only separators may follow the specifier. "C:data" is a drive-relative
// path, not a volume.
for i := 2; i < len(value); i++ {
if value[i] != '\\' && value[i] != '/' {
return "", false
}
}
return strings.ToUpper(value[:2]), true
}
// Sets start values for disk I/O stats.
func (a *Agent) initializeDiskIoStats(diskIoCounters map[string]disk.IOCountersStat) {
a.fsNames = a.fsNames[:0]
now := time.Now()
// ZFS datasets have no /proc/diskstats entry, so they are excluded from
// I/O tracking instead of warning about a missing device (#1541).
var zfsMountpoints map[string]bool
if a.storagePoolManager != nil {
zfsMountpoints = a.storagePoolManager.ZfsMountpoints()
}
for device, stats := range a.fsStats {
if zfsMountpoints[stats.Mountpoint] {
continue
}
// skip if not in diskIoCounters
d, exists := diskIoCounters[device]
if !exists {
@@ -618,9 +542,9 @@ func (a *Agent) initializeDiskIoStats(diskIoCounters map[string]disk.IOCountersS
continue
}
// populate initial values
stats.Time = now
stats.TotalRead = d.ReadBytes
stats.TotalWrite = d.WriteBytes
a.setDiskBaseline(device, prevDiskFromCounter(d, now))
// add to list of valid io device names
a.fsNames = append(a.fsNames, device)
}
@@ -635,31 +559,20 @@ func (a *Agent) updateDiskUsage(systemStats *system.Stats) {
!a.lastDiskUsageUpdate.IsZero() &&
time.Since(a.lastDiskUsageUpdate) < a.diskUsageCacheDuration
// ZFS dataset mountpoints use `zfs list` values because statfs(2) reports
// dataset-level usage that excludes child datasets (#1541).
var zfsUsage map[string]zfsDatasetUsage
if a.storagePoolManager != nil {
zfsUsage = a.storagePoolManager.DatasetUsage()
}
// disk usage
for _, stats := range a.fsStats {
// Skip non-root filesystems if caching is active
if cacheExtraFs && !stats.Root {
continue
}
var total, used uint64
var usedPct float64
if u, ok := zfsUsage[stats.Mountpoint]; ok {
total = u.used + u.avail
used = u.used
if total > 0 {
usedPct = float64(used) / float64(total) * 100
if d, err := disk.Usage(stats.Mountpoint); err == nil {
stats.DiskTotal = utils.BytesToGigabytes(d.Total)
stats.DiskUsed = utils.BytesToGigabytes(d.Used)
if stats.Root {
systemStats.DiskTotal = utils.BytesToGigabytes(d.Total)
systemStats.DiskUsed = utils.BytesToGigabytes(d.Used)
systemStats.DiskPct = utils.TwoDecimals(d.UsedPercent)
}
} else if d, err := disk.Usage(stats.Mountpoint); err == nil {
total = d.Total
used = d.Used
usedPct = d.UsedPercent
} else {
// reset stats if error (likely unmounted)
slog.Error("Error getting disk stats", "name", stats.Mountpoint, "err", err)
@@ -667,14 +580,6 @@ func (a *Agent) updateDiskUsage(systemStats *system.Stats) {
stats.DiskUsed = 0
stats.TotalRead = 0
stats.TotalWrite = 0
continue
}
stats.DiskTotal = utils.BytesToGigabytes(total)
stats.DiskUsed = utils.BytesToGigabytes(used)
if stats.Root {
systemStats.DiskTotal = stats.DiskTotal
systemStats.DiskUsed = stats.DiskUsed
systemStats.DiskPct = utils.TwoDecimals(usedPct)
}
}
@@ -701,11 +606,21 @@ func (a *Agent) updateDiskIo(cacheTimeMs uint16, systemStats *system.Stats) {
}
// Previous snapshot for this interval and device
prev, ok := a.diskPrev[cacheTimeMs][name]
firstSample := !ok
if firstSample {
// Seed from the latest counters of any interval, else seed from current
if prev, ok = a.diskBaseline[name]; !ok {
prev, hasPrev := a.diskPrev[cacheTimeMs][name]
if !hasPrev {
// Seed from agent-level fsStats if present, else seed from current
prev = prevDisk{
readBytes: stats.TotalRead,
writeBytes: stats.TotalWrite,
readTime: d.ReadTime,
writeTime: d.WriteTime,
ioTime: d.IoTime,
weightedIO: d.WeightedIO,
readCount: d.ReadCount,
writeCount: d.WriteCount,
at: stats.Time,
}
if prev.at.IsZero() {
prev = prevDiskFromCounter(d, now)
}
}
@@ -719,12 +634,6 @@ func (a *Agent) updateDiskIo(cacheTimeMs uint16, systemStats *system.Stats) {
if msElapsed < 100 {
continue
}
// The first sample of an interval must span at least half the interval.
// Right after agent start the baseline is only a second or so old, and a
// burst of startup I/O would be recorded as the rate for the whole interval.
if firstSample && msElapsed < uint64(cacheTimeMs)/2 {
continue
}
diskIORead := (d.ReadBytes - prev.readBytes) * 1000 / msElapsed
diskIOWrite := (d.WriteBytes - prev.writeBytes) * 1000 / msElapsed
@@ -746,31 +655,29 @@ func (a *Agent) updateDiskIo(cacheTimeMs uint16, systemStats *system.Stats) {
// This is the total number of milliseconds spent by all reads (as
// measured from __make_request() to end_that_request_last()).
// https://www.kernel.org/doc/Documentation/iostats.txt (fields 4, 8)
deltaReadTime := ioTimeDelta(d.ReadTime, prev.readTime)
deltaWriteTime := ioTimeDelta(d.WriteTime, prev.writeTime)
diskReadTime := utils.TwoDecimals(float64(deltaReadTime) / float64(msElapsed) * 100)
diskWriteTime := utils.TwoDecimals(float64(deltaWriteTime) / float64(msElapsed) * 100)
diskReadTime := utils.TwoDecimals(float64(d.ReadTime-prev.readTime) / float64(msElapsed) * 100)
diskWriteTime := utils.TwoDecimals(float64(d.WriteTime-prev.writeTime) / float64(msElapsed) * 100)
// I/O utilization %: fraction of wall time the device had any I/O in progress (0-100).
diskIoUtilPct := utils.TwoDecimals(float64(ioTimeDelta(d.IoTime, prev.ioTime)) / float64(msElapsed) * 100)
diskIoUtilPct := utils.TwoDecimals(float64(d.IoTime-prev.ioTime) / float64(msElapsed) * 100)
// Weighted I/O: queue-depth weighted I/O time, normalized to interval (can exceed 100%).
// Linux kernel field 11: incremented by iops_in_progress × ms_since_last_update.
// Used to display queue depth. Multipled by 100 to increase accuracy of digit truncation (divided by 100 in UI).
diskWeightedIO := utils.TwoDecimals(float64(ioTimeDelta(d.WeightedIO, prev.weightedIO)) / float64(msElapsed) * 100)
diskWeightedIO := utils.TwoDecimals(float64(d.WeightedIO-prev.weightedIO) / float64(msElapsed) * 100)
// r_await / w_await: average time per read/write operation in milliseconds.
// Equivalent to r_await and w_await in iostat.
var rAwait, wAwait float64
if deltaReadCount := d.ReadCount - prev.readCount; deltaReadCount > 0 {
rAwait = utils.TwoDecimals(float64(deltaReadTime) / float64(deltaReadCount))
rAwait = utils.TwoDecimals(float64(d.ReadTime-prev.readTime) / float64(deltaReadCount))
}
if deltaWriteCount := d.WriteCount - prev.writeCount; deltaWriteCount > 0 {
wAwait = utils.TwoDecimals(float64(deltaWriteTime) / float64(deltaWriteCount))
wAwait = utils.TwoDecimals(float64(d.WriteTime-prev.writeTime) / float64(deltaWriteCount))
}
// Update the baseline that seeds new intervals
a.setDiskBaseline(name, prevDiskFromCounter(d, now))
// Update global fsStats baseline for cross-interval correctness
stats.Time = now
stats.TotalRead = d.ReadBytes
stats.TotalWrite = d.WriteBytes
stats.DiskReadPs = readMbPerSecond
@@ -789,8 +696,6 @@ func (a *Agent) updateDiskIo(cacheTimeMs uint16, systemStats *system.Stats) {
systemStats.DiskWritePs = stats.DiskWritePs
systemStats.DiskIO[0] = diskIORead
systemStats.DiskIO[1] = diskIOWrite
systemStats.DiskIOTotal[0] = d.ReadBytes
systemStats.DiskIOTotal[1] = d.WriteBytes
systemStats.DiskIoStats[0] = diskReadTime
systemStats.DiskIoStats[1] = diskWriteTime
systemStats.DiskIoStats[2] = diskIoUtilPct
@@ -802,30 +707,6 @@ func (a *Agent) updateDiskIo(cacheTimeMs uint16, systemStats *system.Stats) {
}
}
// setDiskBaseline stores the latest counters of a device. A cache interval
// without its own snapshot measures its first sample from them.
func (a *Agent) setDiskBaseline(name string, d prevDisk) {
if a.diskBaseline == nil {
a.diskBaseline = make(map[string]prevDisk)
}
a.diskBaseline[name] = d
}
// ioTimeDelta returns the increase of a cumulative millisecond counter from
// the disk I/O stats. Linux prints these fields of /proc/diskstats as 32-bit
// unsigned ints, so they wrap to zero at 2^32. A busy disk reaches that in
// days for the weighted I/O time. Other platforms report 64-bit counters,
// so a lower value there is a reset.
func ioTimeDelta(current, previous uint64) uint64 {
if current >= previous {
return current - previous
}
if runtime.GOOS == "linux" && previous <= math.MaxUint32 {
return current + (math.MaxUint32 + 1 - previous)
}
return 0
}
// getRootMountPoint returns the appropriate root mount point for the system.
// On Windows it returns the system drive (e.g. "C:").
// For immutable systems like Fedora Silverblue, it returns /sysroot instead of /

View File

@@ -1,174 +0,0 @@
//go:build linux
package agent
import (
"fmt"
"os"
"path/filepath"
"testing"
"time"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/shirou/gopsutil/v4/disk"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// Linux prints four millisecond fields of /proc/diskstats as 32-bit unsigned ints:
// read time, write time, io time and weighted io time. They wrap to zero at 2^32.
func TestUpdateDiskIoTimeCounterWrap(t *testing.T) {
const wrap = uint64(1) << 32
tests := []struct {
name string
base uint64 // added to every previous time counter
}{
{"no wrap", 0},
{"32-bit wrap", wrap - 1000},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// Deltas over 60s: read 300ms / 10 ops, write 400ms / 20 ops,
// io time 1200ms, weighted io 3000ms.
prev := prevDisk{
readBytes: 20000 * 512,
writeBytes: 10000 * 512,
readTime: tt.base + 900,
writeTime: tt.base + 700,
ioTime: tt.base + 400,
weightedIO: tt.base,
readCount: 1000,
writeCount: 500,
at: time.Now().Add(-60 * time.Second),
}
cur := func(v uint64) uint64 { return v % wrap }
line := fmt.Sprintf(" 8 0 sda %d 0 %d %d %d 0 %d %d 0 %d %d\n",
1010, 21200, cur(prev.readTime+300),
520, 10400, cur(prev.writeTime+400),
cur(prev.ioTime+1200), cur(prev.weightedIO+3000))
dir := t.TempDir()
require.NoError(t, os.WriteFile(filepath.Join(dir, "diskstats"), []byte(line), 0o644))
t.Setenv("HOST_PROC", dir)
t.Setenv("HOST_SYS", dir)
t.Setenv("HOST_DEV", dir)
t.Setenv("HOST_RUN", dir)
fs := &system.FsStats{Root: true}
a := &Agent{
fsNames: []string{"sda"},
fsStats: map[string]*system.FsStats{"sda": fs},
diskPrev: map[uint16]map[string]prevDisk{60000: {"sda": prev}},
}
var stats system.Stats
a.updateDiskIo(60000, &stats)
// Same order as DiskIoStats in system.FsStats.
want := [6]float64{0.5, 0.67, 2, 30, 20, 5}
for i := range want {
assert.InDelta(t, want[i], fs.DiskIoStats[i], 0.01, "DiskIoStats[%d]", i)
assert.InDelta(t, want[i], stats.DiskIoStats[i], 0.01, "system DiskIoStats[%d]", i)
}
})
}
}
// backdateDiskBaseline moves the baseline of a device into the past so the
// next seeded sample of an interval spans d.
func backdateDiskBaseline(a *Agent, name string, d time.Duration) {
b := a.diskBaseline[name]
b.at = time.Now().Add(-d)
a.diskBaseline[name] = b
}
// setupDiskstats points gopsutil at a temp dir and returns a writer for its diskstats file.
func setupDiskstats(t *testing.T) func(line string) {
dir := t.TempDir()
t.Setenv("HOST_PROC", dir)
t.Setenv("HOST_SYS", dir)
t.Setenv("HOST_DEV", dir)
t.Setenv("HOST_RUN", dir)
return func(line string) {
require.NoError(t, os.WriteFile(filepath.Join(dir, "diskstats"), []byte(line), 0o644))
}
}
// The first sample of a cache interval has no snapshot of its own. It must
// measure the time counters from the same baseline as the byte counters.
func TestUpdateDiskIoFirstSampleOfInterval(t *testing.T) {
writeDiskstats := setupDiskstats(t)
writeDiskstats(" 8 0 sda 1000 0 20000 900 500 0 10000 700 0 400 0\n")
counters, err := disk.IOCounters("sda")
require.NoError(t, err)
fs := &system.FsStats{Root: true}
a := &Agent{
fsStats: map[string]*system.FsStats{"sda": fs},
diskPrev: map[uint16]map[string]prevDisk{},
}
a.initializeDiskIoStats(counters)
backdateDiskBaseline(a, "sda", 60*time.Second)
// Deltas: read 300ms / 10 ops, write 400ms / 20 ops, io time 1200ms, weighted io 3000ms.
writeDiskstats(" 8 0 sda 1010 0 21200 1200 520 0 10400 1100 0 1600 3000\n")
var stats system.Stats
a.updateDiskIo(60000, &stats)
require.NotZero(t, fs.DiskReadBytes, "bytes are measured from the baseline")
for i := range 3 {
assert.NotZero(t, fs.DiskIoStats[i], "DiskIoStats[%d]", i)
}
assert.InDelta(t, 30, fs.DiskIoStats[3], 0.01, "r_await")
assert.InDelta(t, 20, fs.DiskIoStats[4], 0.01, "w_await")
assert.NotZero(t, fs.DiskIoStats[5], "weighted io")
// A second interval starts from the latest counters, not from the ones at start.
backdateDiskBaseline(a, "sda", time.Second)
// Deltas: read 100ms / 10 ops, write 100ms / 20 ops.
writeDiskstats(" 8 0 sda 1020 0 22400 1300 540 0 10800 1200 0 1800 3500\n")
a.updateDiskIo(1000, &stats)
assert.InDelta(t, 10, fs.DiskIoStats[3], 0.01, "r_await")
assert.InDelta(t, 5, fs.DiskIoStats[4], 0.01, "w_await")
}
// Right after agent start the baseline is too recent to stand for a whole
// interval. The first sample only stores a snapshot, and the next one is
// measured from it.
func TestUpdateDiskIoSkipsShortSeededSample(t *testing.T) {
writeDiskstats := setupDiskstats(t)
writeDiskstats(" 8 0 sda 1000 0 20000 900 500 0 10000 700 0 400 0\n")
counters, err := disk.IOCounters("sda")
require.NoError(t, err)
fs := &system.FsStats{Root: true}
a := &Agent{
fsStats: map[string]*system.FsStats{"sda": fs},
diskPrev: map[uint16]map[string]prevDisk{},
}
a.initializeDiskIoStats(counters)
backdateDiskBaseline(a, "sda", 2*time.Second)
// 1000 MB read in the 2s after start.
writeDiskstats(" 8 0 sda 2000 0 2068000 900 500 0 10000 700 0 400 0\n")
var stats system.Stats
a.updateDiskIo(60000, &stats)
assert.Zero(t, fs.DiskReadBytes)
assert.Zero(t, stats.DiskIO[0])
require.Contains(t, a.diskPrev[60000], "sda", "snapshot is stored")
// Next sample is measured from the stored snapshot: 60 MB over 60s.
snap := a.diskPrev[60000]["sda"]
snap.at = time.Now().Add(-60 * time.Second)
a.diskPrev[60000]["sda"] = snap
writeDiskstats(" 8 0 sda 3000 0 2190880 900 500 0 10000 700 0 400 0\n")
stats = system.Stats{}
a.updateDiskIo(60000, &stats)
assert.InDelta(t, 1_048_576, float64(stats.DiskIO[0]), 20_000)
}

View File

@@ -3,17 +3,14 @@
package agent
import (
"math"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"time"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/shirou/gopsutil/v4/disk"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestParseFilesystemEntry(t *testing.T) {
@@ -81,7 +78,14 @@ func TestParseFilesystemEntry(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
fs, customName := parseFilesystemEntry(tt.input)
fsEntry := strings.TrimSpace(tt.input)
var fs, customName string
if parts := strings.SplitN(fsEntry, "__", 2); len(parts) == 2 {
fs = strings.TrimSpace(parts[0])
customName = strings.TrimSpace(parts[1])
} else {
fs = fsEntry
}
assert.Equal(t, tt.expectedFs, fs)
assert.Equal(t, tt.expectedName, customName)
@@ -273,175 +277,6 @@ func TestBuildFsStatRegistration(t *testing.T) {
assert.Empty(t, key)
assert.Nil(t, stats)
})
t.Run("maps lvm symlinked device to io device through resolved name", func(t *testing.T) {
setEvalSymlinks(t, func(path string) (string, error) {
if path == "/dev/vg1/volume_1" {
return "/dev/dm-1", nil
}
return path, nil
})
key, stats, ok := registerFilesystemStats(
map[string]*system.FsStats{},
"/dev/vg1/volume_1",
"/volume1",
false,
"",
fsRegistrationContext{
isWindows: false,
efPath: "/extra-filesystems",
diskIoCounters: map[string]disk.IOCountersStat{
"dm-0": {Name: "dm-0", Label: "vg1-syno_vg_reserved_area"},
"dm-1": {Name: "dm-1", Label: "vg1-volume_1"},
},
},
)
assert.True(t, ok)
assert.Equal(t, "dm-1", key)
assert.Equal(t, "/volume1", stats.Mountpoint)
})
t.Run("maps lvm device through resolved mapper label", func(t *testing.T) {
setEvalSymlinks(t, func(path string) (string, error) {
if path == "/dev/vg1/volume_1" {
return "/dev/mapper/vg1-volume_1", nil
}
return path, nil
})
key, _, ok := registerFilesystemStats(
map[string]*system.FsStats{},
"/dev/vg1/volume_1",
"/volume1",
false,
"",
fsRegistrationContext{
isWindows: false,
efPath: "/extra-filesystems",
diskIoCounters: map[string]disk.IOCountersStat{
"dm-1": {Name: "dm-1", Label: "vg1-volume_1"},
},
},
)
assert.True(t, ok)
assert.Equal(t, "dm-1", key)
})
t.Run("maps lvm root device through resolved name", func(t *testing.T) {
setEvalSymlinks(t, func(path string) (string, error) {
if path == "/dev/vg1/volume_1" {
return "/dev/dm-1", nil
}
return path, nil
})
key, stats, ok := registerFilesystemStats(
map[string]*system.FsStats{},
"/dev/vg1/volume_1",
"/",
true,
"",
fsRegistrationContext{
isWindows: false,
efPath: "/extra-filesystems",
diskIoCounters: map[string]disk.IOCountersStat{
"dm-1": {Name: "dm-1", Label: "vg1-volume_1"},
},
},
)
assert.True(t, ok)
assert.Equal(t, "dm-1", key)
assert.True(t, stats.Root)
})
t.Run("resolved device wins over filesystem fallback", func(t *testing.T) {
setEvalSymlinks(t, func(path string) (string, error) {
if path == "/dev/vg1/volume_1" {
return "/dev/dm-1", nil
}
return path, nil
})
key, _, ok := registerFilesystemStats(
map[string]*system.FsStats{},
"/dev/vg1/volume_1",
"/",
true,
"",
fsRegistrationContext{
filesystem: "sda",
isWindows: false,
efPath: "/extra-filesystems",
diskIoCounters: map[string]disk.IOCountersStat{
"dm-1": {Name: "dm-1", Label: "vg1-volume_1"},
"sda": {Name: "sda"},
},
},
)
assert.True(t, ok)
assert.Equal(t, "dm-1", key)
})
t.Run("keeps base name when symlink resolution fails", func(t *testing.T) {
setEvalSymlinks(t, func(path string) (string, error) {
return "", os.ErrNotExist
})
key, _, ok := registerFilesystemStats(
map[string]*system.FsStats{},
"/dev/vg1/volume_1",
"/volume1",
false,
"",
fsRegistrationContext{
isWindows: false,
efPath: "/extra-filesystems",
diskIoCounters: map[string]disk.IOCountersStat{
"dm-1": {Name: "dm-1", Label: "vg1-volume_1"},
},
},
)
assert.True(t, ok)
assert.Equal(t, "volume_1", key)
})
t.Run("does not resolve symlinks for relative device names", func(t *testing.T) {
setEvalSymlinks(t, func(path string) (string, error) {
return "/dev/sdb1", nil
})
key, _, ok := registerFilesystemStats(
map[string]*system.FsStats{},
"data",
"/mnt/data",
false,
"",
fsRegistrationContext{
isWindows: false,
efPath: "/extra-filesystems",
diskIoCounters: map[string]disk.IOCountersStat{
"sdb1": {Name: "sdb1"},
},
},
)
assert.True(t, ok)
assert.Equal(t, "data", key)
})
}
// setEvalSymlinks swaps the device-symlink resolver for the duration of a test.
func setEvalSymlinks(t *testing.T, fn func(string) (string, error)) {
t.Helper()
old := evalSymlinks
evalSymlinks = fn
t.Cleanup(func() { evalSymlinks = old })
}
func TestAddConfiguredRootFs(t *testing.T) {
@@ -452,9 +287,8 @@ func TestAddConfiguredRootFs(t *testing.T) {
rootMountPoint: "/",
partitions: []disk.PartitionStat{{Device: "/dev/ada0p2", Mountpoint: "/"}},
ctx: fsRegistrationContext{
filesystem: "/dev/ada0p2",
filesystemName: "root disk",
isWindows: false,
filesystem: "/dev/ada0p2",
isWindows: false,
diskIoCounters: map[string]disk.IOCountersStat{
"ada0": {Name: "ada0", ReadBytes: 1000, WriteBytes: 1000},
},
@@ -468,7 +302,6 @@ func TestAddConfiguredRootFs(t *testing.T) {
assert.True(t, exists)
assert.True(t, stats.Root)
assert.Equal(t, "/", stats.Mountpoint)
assert.Equal(t, "root disk", stats.Name)
})
t.Run("adds root from io device when partition is missing", func(t *testing.T) {
@@ -477,9 +310,8 @@ func TestAddConfiguredRootFs(t *testing.T) {
agent: agent,
rootMountPoint: "/sysroot",
ctx: fsRegistrationContext{
filesystem: "zroot",
filesystemName: "root pool",
isWindows: false,
filesystem: "zroot",
isWindows: false,
diskIoCounters: map[string]disk.IOCountersStat{
"nda0": {Name: "nda0", Label: "zroot", ReadBytes: 1000, WriteBytes: 1000},
},
@@ -493,7 +325,6 @@ func TestAddConfiguredRootFs(t *testing.T) {
assert.True(t, exists)
assert.True(t, stats.Root)
assert.Equal(t, "/sysroot", stats.Mountpoint)
assert.Equal(t, "root pool", stats.Name)
})
t.Run("returns false when filesystem cannot be resolved", func(t *testing.T) {
@@ -676,13 +507,13 @@ func TestAddConfiguredExtraFilesystems(t *testing.T) {
func TestAddExtraFilesystemFolders(t *testing.T) {
t.Run("adds missing folders and skips existing mountpoints", func(t *testing.T) {
agent := &Agent{fsStats: map[string]*system.FsStats{
"existing": {Mountpoint: filepath.FromSlash("/extra-filesystems/existing")},
"existing": {Mountpoint: "/extra-filesystems/existing"},
}}
discovery := diskDiscovery{
agent: agent,
ctx: fsRegistrationContext{
isWindows: false,
efPath: filepath.FromSlash("/extra-filesystems"),
efPath: "/extra-filesystems",
diskIoCounters: map[string]disk.IOCountersStat{
"newdisk": {Name: "newdisk"},
},
@@ -691,10 +522,10 @@ func TestAddExtraFilesystemFolders(t *testing.T) {
discovery.addExtraFilesystemFolders([]string{"existing", "newdisk__Archive"})
require.Len(t, agent.fsStats, 2)
assert.Len(t, agent.fsStats, 2)
stats, exists := agent.fsStats["newdisk"]
require.True(t, exists)
assert.Equal(t, filepath.FromSlash("/extra-filesystems/newdisk__Archive"), stats.Mountpoint)
assert.True(t, exists)
assert.Equal(t, "/extra-filesystems/newdisk__Archive", stats.Mountpoint)
assert.Equal(t, "Archive", stats.Name)
})
}
@@ -705,7 +536,7 @@ func TestAddPartitionExtraFs(t *testing.T) {
agent: agent,
ctx: fsRegistrationContext{
isWindows: false,
efPath: filepath.FromSlash("/extra-filesystems"),
efPath: "/extra-filesystems",
diskIoCounters: map[string]disk.IOCountersStat{
"nvme0n1p1": {Name: "nvme0n1p1"},
"nvme1n1": {Name: "nvme1n1"},
@@ -720,12 +551,12 @@ func TestAddPartitionExtraFs(t *testing.T) {
d.addPartitionExtraFs(disk.PartitionStat{
Device: "/dev/nvme0n1p1",
Mountpoint: filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root"),
Mountpoint: "/extra-filesystems/nvme0n1p1__caddy1-root",
})
stats, exists := agent.fsStats["nvme0n1p1"]
require.True(t, exists)
assert.Equal(t, filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root"), stats.Mountpoint)
assert.True(t, exists)
assert.Equal(t, "/extra-filesystems/nvme0n1p1__caddy1-root", stats.Mountpoint)
assert.Equal(t, "caddy1-root", stats.Name)
})
@@ -736,10 +567,10 @@ func TestAddPartitionExtraFs(t *testing.T) {
// These simulate the virtual mounts that appear when host / is bind-mounted
// with disk.Partitions(all=true) — e.g. /proc, /sys, /dev visible under the mount.
for _, nested := range []string{
filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root/proc"),
filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root/sys"),
filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root/dev"),
filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root/run"),
"/extra-filesystems/nvme0n1p1__caddy1-root/proc",
"/extra-filesystems/nvme0n1p1__caddy1-root/sys",
"/extra-filesystems/nvme0n1p1__caddy1-root/dev",
"/extra-filesystems/nvme0n1p1__caddy1-root/run",
} {
d.addPartitionExtraFs(disk.PartitionStat{Device: "tmpfs", Mountpoint: nested})
}
@@ -752,20 +583,18 @@ func TestAddPartitionExtraFs(t *testing.T) {
d := makeDiscovery(agent)
partitions := []disk.PartitionStat{
{Device: "/dev/nvme0n1p1", Mountpoint: filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root")},
{Device: "/dev/nvme1n1", Mountpoint: filepath.FromSlash("/extra-filesystems/nvme1n1__caddy1-docker")},
{Device: "proc", Mountpoint: filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root/proc")},
{Device: "sysfs", Mountpoint: filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root/sys")},
{Device: "overlay", Mountpoint: filepath.FromSlash("/extra-filesystems/nvme0n1p1__caddy1-root/var/lib/docker")},
{Device: "/dev/nvme0n1p1", Mountpoint: "/extra-filesystems/nvme0n1p1__caddy1-root"},
{Device: "/dev/nvme1n1", Mountpoint: "/extra-filesystems/nvme1n1__caddy1-docker"},
{Device: "proc", Mountpoint: "/extra-filesystems/nvme0n1p1__caddy1-root/proc"},
{Device: "sysfs", Mountpoint: "/extra-filesystems/nvme0n1p1__caddy1-root/sys"},
{Device: "overlay", Mountpoint: "/extra-filesystems/nvme0n1p1__caddy1-root/var/lib/docker"},
}
for _, p := range partitions {
d.addPartitionExtraFs(p)
}
require.Len(t, agent.fsStats, 2)
require.Contains(t, agent.fsStats, "nvme0n1p1")
assert.Len(t, agent.fsStats, 2)
assert.Equal(t, "caddy1-root", agent.fsStats["nvme0n1p1"].Name)
require.Contains(t, agent.fsStats, "nvme1n1")
assert.Equal(t, "caddy1-docker", agent.fsStats["nvme1n1"].Name)
})
@@ -938,6 +767,82 @@ func TestIsDockerSpecialMountpoint(t *testing.T) {
}
}
func TestInitializeDiskInfoWithCustomNames(t *testing.T) {
// Test with custom names
t.Setenv("EXTRA_FILESYSTEMS", "sda1__my-storage,/dev/sdb1__backup-drive,nvme0n1p2")
// Mock disk partitions (we'll just test the parsing logic)
// Since the actual disk operations are system-dependent, we'll focus on the parsing
testCases := []struct {
envValue string
expectedFs []string
expectedNames map[string]string
}{
{
envValue: "sda1__my-storage,sdb1__backup-drive",
expectedFs: []string{"sda1", "sdb1"},
expectedNames: map[string]string{
"sda1": "my-storage",
"sdb1": "backup-drive",
},
},
{
envValue: "sda1,nvme0n1p2__fast-ssd",
expectedFs: []string{"sda1", "nvme0n1p2"},
expectedNames: map[string]string{
"nvme0n1p2": "fast-ssd",
},
},
}
for _, tc := range testCases {
t.Run("env_"+tc.envValue, func(t *testing.T) {
t.Setenv("EXTRA_FILESYSTEMS", tc.envValue)
// Create mock partitions that would match our test cases
partitions := []disk.PartitionStat{}
for _, fs := range tc.expectedFs {
if strings.HasPrefix(fs, "/dev/") {
partitions = append(partitions, disk.PartitionStat{
Device: fs,
Mountpoint: fs,
})
} else {
partitions = append(partitions, disk.PartitionStat{
Device: "/dev/" + fs,
Mountpoint: "/" + fs,
})
}
}
// Test the parsing logic by calling the relevant part
// We'll create a simplified version to test just the parsing
extraFilesystems := tc.envValue
for fsEntry := range strings.SplitSeq(extraFilesystems, ",") {
// Parse the entry
fsEntry = strings.TrimSpace(fsEntry)
var fs, customName string
if parts := strings.SplitN(fsEntry, "__", 2); len(parts) == 2 {
fs = strings.TrimSpace(parts[0])
customName = strings.TrimSpace(parts[1])
} else {
fs = fsEntry
}
// Verify the device is in our expected list
assert.Contains(t, tc.expectedFs, fs, "parsed device should be in expected list")
// Check if custom name should exist
if expectedName, exists := tc.expectedNames[fs]; exists {
assert.Equal(t, expectedName, customName, "custom name should match expected")
} else {
assert.Empty(t, customName, "custom name should be empty when not expected")
}
}
})
}
}
func TestFsStatsWithCustomNames(t *testing.T) {
// Test that FsStats properly stores custom names
fsStats := &system.FsStats{
@@ -1128,10 +1033,8 @@ func TestInitializeDiskIoStatsResetsTrackedDevices(t *testing.T) {
assert.Len(t, agent.fsNames, 2)
assert.Equal(t, uint64(10), agent.fsStats["sda"].TotalRead)
assert.Equal(t, uint64(20), agent.fsStats["sda"].TotalWrite)
assert.Equal(t, uint64(10), agent.diskBaseline["sda"].readBytes)
assert.Equal(t, uint64(40), agent.diskBaseline["sdb"].writeBytes)
assert.False(t, agent.diskBaseline["sda"].at.IsZero())
assert.False(t, agent.diskBaseline["sdb"].at.IsZero())
assert.False(t, agent.fsStats["sda"].Time.IsZero())
assert.False(t, agent.fsStats["sdb"].Time.IsZero())
agent.initializeDiskIoStats(map[string]disk.IOCountersStat{
"sdb": {Name: "sdb", ReadBytes: 50, WriteBytes: 60},
@@ -1141,114 +1044,3 @@ func TestInitializeDiskIoStatsResetsTrackedDevices(t *testing.T) {
assert.Equal(t, uint64(50), agent.fsStats["sdb"].TotalRead)
assert.Equal(t, uint64(60), agent.fsStats["sdb"].TotalWrite)
}
func TestIoTimeDelta(t *testing.T) {
assert.Equal(t, uint64(300), ioTimeDelta(1200, 900))
// A lower value is a 32-bit wrap only on Linux. Other platforms
// report 64-bit counters, so there it is a reset.
var want uint64
if runtime.GOOS == "linux" {
want = 1200
}
assert.Equal(t, want, ioTimeDelta(200, math.MaxUint32+1-1000))
assert.Equal(t, uint64(0), ioTimeDelta(200, math.MaxUint32+1000))
}
func TestNormalizeDeviceName(t *testing.T) {
// A Windows volume name is not a path element, so every spelling of the
// same drive has to normalize to the same key. filepath.Base cannot do
// this: on Windows it strips the "C:" specifier and returns "\", which
// collapses every drive letter onto one key (#2417).
for _, spelling := range []string{"C:", `C:\`, "C:/", `C:\\`} {
assert.Equal(t, "C:", normalizeDeviceName(spelling), "spelling %q", spelling)
}
// Drive letters are case-insensitive, so the letter is uppercased.
assert.Equal(t, "D:", normalizeDeviceName("d:"))
assert.Equal(t, "C:", normalizeDeviceName(" c: "))
assert.Equal(t, "C:", normalizeDeviceName(`c:\`))
// Non-volume inputs keep using filepath.Base.
assert.Equal(t, "sda1", normalizeDeviceName("/dev/sda1"))
assert.Equal(t, "sda1", normalizeDeviceName("/dev/sda1/"))
assert.Equal(t, "nvme0n1p2", normalizeDeviceName(" /dev/nvme0n1p2 "))
assert.Equal(t, "", normalizeDeviceName("."))
assert.Equal(t, "", normalizeDeviceName(" "))
// A drive-relative path is a path, not a volume.
assert.Equal(t, filepath.Base(`C:data`), normalizeDeviceName(`C:data`))
}
func TestFindIoDeviceWindowsVolumeNames(t *testing.T) {
// Every drive normalizes to a distinct key, so the root drive resolves
// exactly instead of to whichever counter the map yielded first (#2417).
ioCounters := map[string]disk.IOCountersStat{
"C:": {Name: "C:", ReadBytes: 10, WriteBytes: 10},
"D:": {Name: "D:", ReadBytes: 20, WriteBytes: 20},
"P:": {Name: "P:", ReadBytes: 30, WriteBytes: 30},
}
for i := 0; i < 32; i++ {
device, ok := findIoDevice("C:", ioCounters)
assert.True(t, ok)
assert.Equal(t, "C:", device)
}
// The drive may arrive with a trailing separator, as a mount point does.
device, ok := findIoDevice(`C:\`, ioCounters)
assert.True(t, ok)
assert.Equal(t, "C:", device)
}
func TestAddPartitionRootFsWindowsDrive(t *testing.T) {
agent := &Agent{fsStats: make(map[string]*system.FsStats)}
discovery := diskDiscovery{
agent: agent,
ctx: fsRegistrationContext{
isWindows: true,
diskIoCounters: map[string]disk.IOCountersStat{
"C:": {Name: "C:"},
"D:": {Name: "D:"},
"P:": {Name: "P:"},
},
},
}
ok := discovery.addPartitionRootFs("C:", `C:\`)
assert.True(t, ok)
assert.Len(t, agent.fsStats, 1)
stats, exists := agent.fsStats["C:"]
assert.True(t, exists)
assert.True(t, stats.Root)
}
func TestAddPartitionRootFsKeyAlreadyRegistered(t *testing.T) {
// The root drive is also listed in EXTRA_FILESYSTEMS, so its key is taken
// before the root fallback runs. The existing entry must be promoted to root
// rather than falling back to the most active device, which here is D:.
agent := &Agent{fsStats: map[string]*system.FsStats{
"C:": {Mountpoint: `C:\`, Name: "System"},
"D:": {Mountpoint: `D:\`},
}}
discovery := diskDiscovery{
agent: agent,
rootMountPoint: `C:\`,
ctx: fsRegistrationContext{
isWindows: true,
diskIoCounters: map[string]disk.IOCountersStat{
"C:": {Name: "C:", ReadBytes: 10},
"D:": {Name: "D:", ReadBytes: 100},
},
},
}
ok := discovery.addPartitionRootFs("C:", `C:\`)
assert.True(t, ok)
assert.Len(t, agent.fsStats, 2)
assert.True(t, agent.fsStats["C:"].Root)
assert.Equal(t, `C:\`, agent.fsStats["C:"].Mountpoint)
assert.Equal(t, "System", agent.fsStats["C:"].Name)
assert.False(t, agent.fsStats["D:"].Root)
}

View File

@@ -1,110 +0,0 @@
//go:build testing
package agent
import (
"testing"
"time"
"github.com/henrygd/beszel/agent/zfs"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/shirou/gopsutil/v4/disk"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// TestUpdateDiskUsageZfsMountpoint verifies that a filesystem whose mountpoint
// is a ZFS dataset reports `zfs list` usage (which includes child datasets)
// instead of the dataset-scoped statfs values (#1541).
func TestUpdateDiskUsageZfsMountpoint(t *testing.T) {
zm := &StoragePoolManager{detailInterval: time.Hour, backends: []*poolBackend{{name: "zfs"}}}
zm.backends[0].datasetsFn = func() ([]zfs.Dataset, error) {
return []zfs.Dataset{
{Name: "tank", Used: 12000000000000, Avail: 11999000000000, Mountpoint: "/tank"},
}, nil
}
agent := &Agent{
fsStats: map[string]*system.FsStats{
"tank": {Root: false, Mountpoint: "/tank"},
},
storagePoolManager: zm,
}
var stats system.Stats
agent.updateDiskUsage(&stats)
fs := agent.fsStats["tank"]
require.NotNil(t, fs)
assert.Equal(t, 22350.81, fs.DiskTotal) // (used + avail) in GiB
assert.Equal(t, 11175.87, fs.DiskUsed)
// Non-root filesystems do not populate system-level stats.
assert.Equal(t, float64(0), stats.DiskTotal)
}
// TestUpdateDiskUsageZfsRootPopulatesSystemStats verifies the root disk values
// are derived from ZFS usage when the root mountpoint is a ZFS dataset.
func TestUpdateDiskUsageZfsRootPopulatesSystemStats(t *testing.T) {
zm := &StoragePoolManager{detailInterval: time.Hour, backends: []*poolBackend{{name: "zfs"}}}
zm.backends[0].datasetsFn = func() ([]zfs.Dataset, error) {
return []zfs.Dataset{
{Name: "rpool/ROOT/pve-1", Used: 900000000000, Avail: 300000000000, Mountpoint: "/"},
}, nil
}
agent := &Agent{
fsStats: map[string]*system.FsStats{
"rpool/ROOT/pve-1": {Root: true, Mountpoint: "/"},
},
storagePoolManager: zm,
}
var stats system.Stats
agent.updateDiskUsage(&stats)
assert.Equal(t, 1117.59, agent.fsStats["rpool/ROOT/pve-1"].DiskTotal)
assert.Equal(t, 838.19, agent.fsStats["rpool/ROOT/pve-1"].DiskUsed)
assert.Equal(t, 75.0, stats.DiskPct)
assert.Equal(t, 1117.59, stats.DiskTotal)
assert.Equal(t, 838.19, stats.DiskUsed)
}
// TestUpdateDiskUsageWithoutZfsManager falls back to statfs when no manager is
// present (e.g. tests constructing bare Agent values).
func TestUpdateDiskUsageWithoutZfsManager(t *testing.T) {
agent := &Agent{
fsStats: map[string]*system.FsStats{
"root": {Root: true, Mountpoint: "/"},
},
}
var stats system.Stats
agent.updateDiskUsage(&stats)
assert.True(t, agent.fsStats["root"].DiskTotal > 0, "root usage should come from statfs")
assert.True(t, stats.DiskTotal > 0)
}
// TestInitializeDiskIoStatsSkipsZfsMountpoints verifies ZFS filesystems are
// excluded from diskstats I/O tracking instead of warning about a missing device.
func TestInitializeDiskIoStatsSkipsZfsMountpoints(t *testing.T) {
zm := &StoragePoolManager{detailInterval: time.Hour, backends: []*poolBackend{{name: "zfs"}}}
zm.backends[0].datasetsFn = func() ([]zfs.Dataset, error) {
return []zfs.Dataset{{Name: "tank", Mountpoint: "/tank"}}, nil
}
agent := &Agent{
fsStats: map[string]*system.FsStats{
"tank": {Root: false, Mountpoint: "/tank"},
"sda1": {Root: false, Mountpoint: "/mnt/data"},
},
storagePoolManager: zm,
diskPrev: make(map[uint16]map[string]prevDisk),
}
agent.initializeDiskIoStats(map[string]disk.IOCountersStat{
"sda1": {Name: "sda1", ReadBytes: 100, WriteBytes: 100},
})
assert.Equal(t, []string{"sda1"}, agent.fsNames)
assert.Equal(t, uint64(100), agent.fsStats["sda1"].TotalRead)
// ZFS entry is present but untouched by diskstats initialization.
assert.Equal(t, uint64(0), agent.fsStats["tank"].TotalRead)
}

View File

@@ -65,15 +65,11 @@ type dockerManager struct {
dockerVersionChecked bool // Whether a version probe has completed successfully
isWindows bool // Whether the Docker Engine API is running on Windows
buf *bytes.Buffer // Buffer to store and read response bodies
decoder *json.Decoder // Reusable JSON decoder that reads from buf
apiStats *container.ApiStats // Reusable API stats object
excludeContainers []string // Patterns to exclude containers by name
usingPodman bool // Whether the Docker Engine API is running on Podman
registryClient *http.Client // Client for registry requests; nil uses a client with a 10-second timeout
imageUpdatesDisabled bool // Whether image update checks are disabled by configuration
imageUpdatesMutex sync.RWMutex // Protects imageUpdates, its entries, and imageUpdatesRunning
imageUpdates map[string]*imageUpdateStatus // Shared update status keyed by normalized image reference
imageUpdatesRunning bool // Whether a background image-update batch is in progress
// Cache-time-aware tracking for CPU stats (similar to cpu.go)
// Maps cache time intervals to container-specific CPU usage tracking
lastCpuContainer map[uint16]map[string]uint64 // cacheTimeMs -> containerId -> last cpu container usage
@@ -166,9 +162,6 @@ func (dm *dockerManager) getDockerStats(cacheTimeMs uint16) ([]*container.Stats,
clear(dm.validIds)
}
// Only schedule auxiliary work here; metrics never wait for image discovery.
dm.refreshImageUpdates(dm.apiContainerList, time.Now())
var failedContainers []*container.ApiInfo
for _, ctr := range dm.apiContainerList {
@@ -381,26 +374,16 @@ func convertContainerPortsToString(ctr *container.ApiInfo) string {
return ""
}
sort.Slice(ctr.Ports, func(i, j int) bool {
if ctr.Ports[i].PublicPort != ctr.Ports[j].PublicPort {
return ctr.Ports[i].PublicPort < ctr.Ports[j].PublicPort
}
return ctr.Ports[i].IP < ctr.Ports[j].IP
return ctr.Ports[i].PublicPort < ctr.Ports[j].PublicPort
})
var builder strings.Builder
seen := make(map[string]struct{})
seenPorts := make(map[uint16]struct{})
for _, p := range ctr.Ports {
if p.PublicPort == 0 {
_, ok := seenPorts[p.PublicPort]
if p.PublicPort == 0 || ok {
continue
}
keyIP := p.IP
if keyIP == "0.0.0.0" || keyIP == "::" {
keyIP = ""
}
key := keyIP + ":" + strconv.Itoa(int(p.PublicPort))
if _, ok := seen[key]; ok {
continue
}
seen[key] = struct{}{}
seenPorts[p.PublicPort] = struct{}{}
if builder.Len() > 0 {
builder.WriteString(", ")
}
@@ -514,17 +497,6 @@ func (dm *dockerManager) updateContainerStats(ctr *container.ApiInfo, cacheTimeM
}
}
// Read and decode the response before locking shared stats to avoid blocking
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("container stats request failed: %s", resp.Status)
}
res := &container.ApiStats{}
if err := json.NewDecoder(resp.Body).Decode(res); err != nil {
return err
}
updateAvailable := dm.cachedImageUpdate(ctr.Image)
dm.containerStatsMutex.Lock()
defer dm.containerStatsMutex.Unlock()
@@ -539,9 +511,6 @@ func (dm *dockerManager) updateContainerStats(ctr *container.ApiInfo, cacheTimeM
stats.Status = statusText
stats.Health = health
stats.Image = ctr.Image
stats.UpdateAvailable = updateAvailable
if len(ctr.Ports) > 0 {
stats.Ports = convertContainerPortsToString(ctr)
}
@@ -554,24 +523,23 @@ func (dm *dockerManager) updateContainerStats(ctr *container.ApiInfo, cacheTimeM
stats.NetworkSent = 0
stats.NetworkRecv = 0
res := dm.apiStats
res.Networks = nil
if err := dm.decode(resp, res); err != nil {
return err
}
// Initialize CPU tracking for this cache time interval
dm.initializeCpuTracking(cacheTimeMs)
// Get previous CPU values
prevCpuContainer, prevCpuSystem := dm.getCpuPreviousValues(cacheTimeMs, ctr.IdShort)
// Calculate CPU percentage based on platform.
// Podman reports system_cpu_usage from cgroup cpu.stat (not /proc/stat), so it reflects
// only cgroup-tracked activity rather than total host capacity. Use a time-based method
// instead so the result is comparable to host CPU utilization. See:
// https://github.com/henrygd/beszel/issues/2049
// Calculate CPU percentage based on platform
var cpuPct float64
if dm.isWindows {
prevRead := dm.lastCpuReadTime[cacheTimeMs][ctr.IdShort]
cpuPct = res.CalculateCpuPercentWindows(prevCpuContainer, prevRead)
} else if dm.usingPodman && res.CPUStats.OnlineCPUs > 0 {
prevRead := dm.lastCpuReadTime[cacheTimeMs][ctr.IdShort]
cpuPct = res.CalculateCpuPercentPodman(prevCpuContainer, prevRead)
} else {
cpuPct = res.CalculateCpuPercentLinux(prevCpuContainer, prevCpuSystem)
}
@@ -689,8 +657,6 @@ func newDockerManager(agent *Agent) *dockerManager {
userAgent: "Docker-Client/",
}
dockerImageCheck, _ := utils.GetEnv("DOCKER_IMAGE_CHECK")
// Read container exclusion patterns from environment variable
var excludeContainers []string
if excludeStr, set := utils.GetEnv("EXCLUDE_CONTAINERS"); set && excludeStr != "" {
@@ -710,11 +676,11 @@ func newDockerManager(agent *Agent) *dockerManager {
Timeout: timeout,
Transport: userAgentTransport,
},
containerStatsMap: make(map[string]*container.Stats),
sem: make(chan struct{}, 5),
apiContainerList: []*container.ApiInfo{},
excludeContainers: excludeContainers,
imageUpdatesDisabled: dockerImageCheck == "false",
containerStatsMap: make(map[string]*container.Stats),
sem: make(chan struct{}, 5),
apiContainerList: []*container.ApiInfo{},
apiStats: &container.ApiStats{},
excludeContainers: excludeContainers,
// Initialize cache-time-aware tracking structures
lastCpuContainer: make(map[uint16]map[string]uint64),
@@ -781,18 +747,20 @@ func (dm *dockerManager) applyDockerVersionInfo(serverHeader string, versionInfo
}
}
// Decodes a Docker API JSON response using a reusable buffer. Not thread safe.
// Decodes Docker API JSON response using a reusable buffer and decoder. Not thread safe.
func (dm *dockerManager) decode(resp *http.Response, d any) error {
if dm.buf == nil {
// initialize buffer with 256kb starting size
dm.buf = bytes.NewBuffer(make([]byte, 0, 1024*256))
dm.decoder = json.NewDecoder(dm.buf)
}
defer resp.Body.Close()
defer dm.buf.Reset()
if _, err := dm.buf.ReadFrom(resp.Body); err != nil {
_, err := dm.buf.ReadFrom(resp.Body)
if err != nil {
return err
}
return json.Unmarshal(dm.buf.Bytes(), d)
return dm.decoder.Decode(d)
}
// Test docker / podman sockets and return if one exists
@@ -866,10 +834,9 @@ func (dm *dockerManager) getContainerInfo(ctx context.Context, containerID strin
// getLogs fetches the logs for a container
func (dm *dockerManager) getLogs(ctx context.Context, containerID string) (string, error) {
query := url.Values{
"timestamps": []string{"1"},
"stdout": []string{"1"},
"stderr": []string{"1"},
"tail": []string{fmt.Sprintf("%d", dockerLogsTail)},
"stdout": []string{"1"},
"stderr": []string{"1"},
"tail": []string{fmt.Sprintf("%d", dockerLogsTail)},
}
endpoint, err := buildDockerContainerEndpoint(containerID, "logs", query)
if err != nil {

View File

@@ -1,108 +0,0 @@
package agent
import (
"log/slog"
"sync"
"time"
"github.com/distribution/reference"
"github.com/henrygd/beszel/internal/entities/container"
)
const imageUpdateInterval = time.Hour
type imageUpdateStatus struct {
available bool
checkedAt time.Time
}
func normalizedImageReference(image string) string {
named, err := reference.ParseNormalizedNamed(image)
if err != nil {
return ""
}
// Digest-pinned references cannot move to a new version.
if _, pinned := named.(reference.Digested); pinned {
return ""
}
return reference.TagNameOnly(named).String()
}
// refreshImageUpdates starts at most one background batch. Neither its network
// work nor its completion is part of the container metrics wait group.
func (dm *dockerManager) refreshImageUpdates(containers []*container.ApiInfo, now time.Time) {
if dm.imageUpdatesDisabled {
return
}
dm.imageUpdatesMutex.Lock()
defer dm.imageUpdatesMutex.Unlock()
if dm.imageUpdatesRunning {
return
}
if dm.imageUpdates == nil {
dm.imageUpdates = make(map[string]*imageUpdateStatus)
}
active := make(map[string]struct{}, len(containers))
pending := make(map[string]*imageUpdateStatus)
for _, ctr := range containers {
if len(ctr.Names) > 0 && dm.shouldExcludeContainer(ctr.Names[0][1:]) {
continue
}
key := normalizedImageReference(ctr.Image)
if key == "" {
continue
}
active[key] = struct{}{}
entry := dm.imageUpdates[key]
if entry == nil {
entry = &imageUpdateStatus{}
dm.imageUpdates[key] = entry
}
if entry.checkedAt.IsZero() || now.Sub(entry.checkedAt) >= imageUpdateInterval {
pending[key] = entry
}
}
for key := range dm.imageUpdates {
if _, ok := active[key]; !ok {
delete(dm.imageUpdates, key)
}
}
if len(pending) == 0 {
return
}
dm.imageUpdatesRunning = true
go func() {
// Limit auxiliary requests even on hosts running many different images.
sem := make(chan struct{}, 2)
var wg sync.WaitGroup
for key, entry := range pending {
sem <- struct{}{}
wg.Add(1)
go func() {
defer wg.Done()
defer func() { <-sem }()
available, err := dm.checkImageUpdate(key)
if err != nil {
available = false
slog.Debug("Image update check failed", "image", key, "err", err)
}
dm.imageUpdatesMutex.Lock()
entry.available = available
entry.checkedAt = time.Now()
dm.imageUpdatesMutex.Unlock()
}()
}
wg.Wait()
dm.imageUpdatesMutex.Lock()
dm.imageUpdatesRunning = false
dm.imageUpdatesMutex.Unlock()
}()
}
func (dm *dockerManager) cachedImageUpdate(image string) bool {
key := normalizedImageReference(image)
dm.imageUpdatesMutex.RLock()
defer dm.imageUpdatesMutex.RUnlock()
entry := dm.imageUpdates[key]
return entry != nil && entry.available
}

View File

@@ -1,248 +0,0 @@
//go:build testing
package agent
import (
"encoding/json"
"fmt"
"github.com/fxamacker/cbor/v2"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"github.com/henrygd/beszel/internal/entities/container"
"github.com/stretchr/testify/require"
)
func waitForImageUpdates(t *testing.T, dm *dockerManager) {
t.Helper()
require.Eventually(t, func() bool {
dm.imageUpdatesMutex.RLock()
defer dm.imageUpdatesMutex.RUnlock()
return !dm.imageUpdatesRunning
}, time.Second*3, time.Millisecond)
}
func TestDisableDockerImageUpdateCheck(t *testing.T) {
t.Setenv("BESZEL_AGENT_DOCKER_IMAGE_CHECK", "false")
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/version" {
fmt.Fprint(w, `{"Version":"25.0.0"}`)
return
}
http.NotFound(w, r)
}))
defer server.Close()
t.Setenv("BESZEL_AGENT_DOCKER_HOST", server.URL)
dm := newDockerManager(nil)
require.True(t, dm.imageUpdatesDisabled)
dm.registryClient = &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
t.Fatal("disabled image update check made a registry request")
return nil, nil
})}
dm.refreshImageUpdates([]*container.ApiInfo{{Image: "nginx", Names: []string{"/nginx"}}}, time.Now())
require.False(t, dm.imageUpdatesRunning)
require.Nil(t, dm.imageUpdates)
}
func TestImageUpdateCacheAndStats(t *testing.T) {
local := "sha256:" + strings.Repeat("a", 64)
remote := "sha256:" + strings.Repeat("b", 64)
var inspections, lookups atomic.Int32
var fail atomic.Bool
var upToDate atomic.Bool
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch {
case strings.HasPrefix(r.URL.Path, "/images/"):
inspections.Add(1)
fmt.Fprintf(w, `{"RepoDigests":["docker.io/library/nginx@%s"]}`, local)
case r.URL.Path == "/containers/json":
fmt.Fprint(w, `[{"Id":"aaaaaaaaaaaa","Names":["/one"],"Image":"nginx","Status":"Up 2 hours"},{"Id":"bbbbbbbbbbbb","Names":["/two"],"Image":"docker.io/library/nginx:latest","Status":"Up 2 hours"}]`)
case strings.Contains(r.URL.Path, "/stats"):
fmt.Fprint(w, `{"memory_stats":{"usage":1048576},"cpu_stats":{},"networks":{}}`)
default:
http.NotFound(w, r)
}
}))
defer server.Close()
dm := newDockerManagerForVersionTest(server)
dm.dockerVersionChecked = true
dm.registryClient = &http.Client{Timeout: time.Second, Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) {
if fail.Load() {
return nil, fmt.Errorf("registry unavailable")
}
response := &http.Response{StatusCode: 200, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(`{"token":"test"}`))}
if r.Method == http.MethodHead {
lookups.Add(1)
digest := remote
if upToDate.Load() {
digest = local
}
response.Header.Set("Docker-Content-Digest", digest)
}
return response, nil
})}
stats, err := dm.getDockerStats(defaultCacheTimeMs)
require.NoError(t, err)
require.Len(t, stats, 2)
waitForImageUpdates(t, dm)
require.EqualValues(t, 1, lookups.Load())
require.EqualValues(t, 1, inspections.Load())
stats, err = dm.getDockerStats(defaultCacheTimeMs)
require.NoError(t, err)
for _, stat := range stats {
require.True(t, stat.UpdateAvailable)
if stat.Id == "aaaaaaaaaaaa" {
require.Equal(t, "nginx", stat.Image)
} else {
require.Equal(t, "docker.io/library/nginx:latest", stat.Image)
}
}
require.EqualValues(t, 1, lookups.Load())
expire := func() {
dm.imageUpdatesMutex.Lock()
dm.imageUpdates["docker.io/library/nginx:latest"].checkedAt = time.Now().Add(-imageUpdateInterval)
dm.imageUpdatesMutex.Unlock()
}
upToDate.Store(true)
expire()
_, err = dm.getDockerStats(defaultCacheTimeMs)
require.NoError(t, err)
waitForImageUpdates(t, dm)
require.EqualValues(t, 2, lookups.Load())
require.False(t, dm.cachedImageUpdate("nginx:latest"))
// An expired positive result is cleared on failure, and the failure itself
// is cached so realtime stats do not retry a broken registry every second.
dm.imageUpdatesMutex.Lock()
dm.imageUpdates["docker.io/library/nginx:latest"].available = true
dm.imageUpdatesMutex.Unlock()
fail.Store(true)
expire()
_, err = dm.getDockerStats(defaultCacheTimeMs)
require.NoError(t, err)
waitForImageUpdates(t, dm)
failedInspections := inspections.Load()
stats, err = dm.getDockerStats(defaultCacheTimeMs)
require.NoError(t, err)
require.Len(t, stats, 2)
require.Equal(t, failedInspections, inspections.Load())
for _, stat := range stats {
require.False(t, stat.UpdateAvailable)
require.Equal(t, 1.0, stat.Mem)
}
}
func TestImageDiscoveryDoesNotBlockStats(t *testing.T) {
started := make(chan struct{}, 1)
release := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasPrefix(r.URL.Path, "/images/") {
fmt.Fprintf(w, `{"RepoDigests":["example.com/app@sha256:%s"]}`, strings.Repeat("a", 64))
} else {
fmt.Fprint(w, `{"memory_stats":{"usage":1048576}}`)
}
}))
defer server.Close()
dm := newDockerManagerForVersionTest(server)
defer func() { close(release); waitForImageUpdates(t, dm) }()
dm.registryClient = &http.Client{Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) {
started <- struct{}{}
<-release
return nil, fmt.Errorf("timeout")
})}
ctr := &container.ApiInfo{IdShort: "aaaaaaaaaaaa", Image: "example.com/app", Names: []string{"/one"}}
dm.refreshImageUpdates([]*container.ApiInfo{ctr}, time.Now())
select {
case <-started:
case <-time.After(3 * time.Second):
t.Fatal("check did not start")
}
done := make(chan error, 1)
go func() { done <- dm.updateContainerStats(ctr, defaultCacheTimeMs) }()
select {
case err := <-done:
require.NoError(t, err)
case <-time.After(time.Second):
t.Fatal("registry blocked stats")
}
dm.imageUpdatesMutex.RLock()
require.True(t, dm.imageUpdatesRunning)
dm.imageUpdatesMutex.RUnlock()
}
func TestNormalizeImageUpdateReferences(t *testing.T) {
require.Equal(t, normalizedImageReference("nginx"), normalizedImageReference("docker.io/library/nginx:latest"))
require.Empty(t, normalizedImageReference("bad reference"))
require.Empty(t, normalizedImageReference("nginx@sha256:"+strings.Repeat("a", 64)))
}
// A stats request can return headers promptly and then stall while reading its
// body. The stats-map mutex must remain available during that read.
func TestStatsResponseBodyDoesNotHoldStatsLock(t *testing.T) {
started := make(chan struct{})
release := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
w.(http.Flusher).Flush()
close(started)
<-release
fmt.Fprint(w, `{"memory_stats":{"usage":1048576}}`)
}))
defer server.Close()
dm := newDockerManagerForVersionTest(server)
done := make(chan error, 1)
go func() {
done <- dm.updateContainerStats(&container.ApiInfo{IdShort: "aaaaaaaaaaaa", Names: []string{"/one"}, Image: "nginx"}, defaultCacheTimeMs)
}()
<-started
locked := make(chan struct{})
go func() { dm.containerStatsMutex.Lock(); dm.containerStatsMutex.Unlock(); close(locked) }()
select {
case <-locked:
case <-time.After(time.Second):
close(release)
<-done
t.Fatal("Docker response body held the stats mutex")
}
close(release)
require.NoError(t, <-done)
}
func TestImageUpdateStatsEncoding(t *testing.T) {
original := container.Stats{Image: "nginx:latest", UpdateAvailable: true}
encoded, err := cbor.Marshal(original)
require.NoError(t, err)
var fields map[int]any
require.NoError(t, cbor.Unmarshal(encoded, &fields))
require.Equal(t, true, fields[11])
require.Equal(t, "nginx:latest", fields[8])
var decoded container.Stats
require.NoError(t, cbor.Unmarshal(encoded, &decoded))
require.True(t, decoded.UpdateAvailable)
require.Equal(t, original.Image, decoded.Image)
encoded, err = json.Marshal(original)
require.NoError(t, err)
require.Contains(t, string(encoded), `"u":true`)
}
func TestImageUpdateCacheExpiryBoundaryAndPruning(t *testing.T) {
now := time.Now()
key := normalizedImageReference("nginx")
dm := &dockerManager{imageUpdates: map[string]*imageUpdateStatus{
key: {available: true, checkedAt: now},
"unused.example/image:latest": {checkedAt: now},
}}
dm.refreshImageUpdates([]*container.ApiInfo{{Image: "nginx"}}, now.Add(imageUpdateInterval-time.Nanosecond))
require.False(t, dm.imageUpdatesRunning)
require.Len(t, dm.imageUpdates, 1)
require.True(t, dm.cachedImageUpdate("nginx:latest"))
dm.refreshImageUpdates(nil, now)
require.Empty(t, dm.imageUpdates)
}

View File

@@ -1,224 +0,0 @@
package agent
import (
_ "crypto/sha256"
"encoding/json"
"fmt"
"net/http"
"net/url"
"slices"
"strings"
"time"
"github.com/distribution/reference"
"github.com/opencontainers/go-digest"
)
const imageRegistryTimeout = 10 * time.Second
const imageManifestAccept = "application/vnd.docker.distribution.manifest.list.v2+json, " +
"application/vnd.docker.distribution.manifest.v2+json, " +
"application/vnd.oci.image.manifest.v1+json, " +
"application/vnd.oci.image.index.v1+json"
// checkImageUpdate compares the digest recorded by Docker for image with the
// digest currently advertised by its registry. A digest-pinned reference is
// immutable and therefore never has an update available.
func (dm *dockerManager) checkImageUpdate(image string) (bool, error) {
named, err := reference.ParseNormalizedNamed(image)
if err != nil {
return false, fmt.Errorf("parse image reference %q: %w", image, err)
}
if _, pinned := named.(reference.Digested); pinned {
return false, nil
}
named = reference.TagNameOnly(named)
registry := reference.Domain(named)
repository := reference.Path(named)
tag := named.(reference.Tagged).Tag()
localDigests, err := dm.inspectImageDigests(image, registry, repository)
if err != nil {
return false, err
}
remoteDigest, err := dm.registryImageDigest(registry, repository, tag)
if err != nil {
return false, err
}
return !slices.Contains(localDigests, remoteDigest), nil
}
// inspectImageDigests reads Docker's image metadata without using dm.decode.
// The checker runs in the image-discovery goroutine, so it must not hold any
// of the container statistics locks while waiting on the Docker API.
func (dm *dockerManager) inspectImageDigests(image, registry, repository string) ([]string, error) {
if dm.client == nil {
return nil, fmt.Errorf("inspect image %q: Docker client is unavailable", image)
}
endpoint := "http://localhost/images/" + url.PathEscape(image) + "/json"
resp, err := dm.client.Get(endpoint)
if err != nil {
return nil, fmt.Errorf("inspect image %q: %w", image, err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("inspect image %q failed: %s", image, responseStatus(resp))
}
var inspect struct {
RepoDigests []string `json:"RepoDigests"`
}
if err := json.NewDecoder(resp.Body).Decode(&inspect); err != nil {
return nil, fmt.Errorf("decode image inspect %q: %w", image, err)
}
if len(inspect.RepoDigests) == 0 {
return nil, fmt.Errorf("inspect image %q returned no repository digests", image)
}
localDigests := matchingRepositoryDigests(inspect.RepoDigests, registry, repository)
if len(localDigests) == 0 {
return nil, fmt.Errorf("inspect image %q returned no valid digest for %s/%s", image, registry, repository)
}
return localDigests, nil
}
// matchingRepositoryDigests returns all valid digests belonging to the requested
// repository. Container engines can return both index and platform manifest digests for one
// local image, in either order.
func matchingRepositoryDigests(repoDigests []string, registry, repository string) []string {
var digests []string
for _, repoDigest := range repoDigests {
repoDigest = strings.TrimSpace(repoDigest)
at := strings.LastIndexByte(repoDigest, '@')
if at <= 0 || at == len(repoDigest)-1 || strings.Contains(repoDigest[:at], "@") {
continue
}
repoRef, err := reference.ParseNormalizedNamed(repoDigest[:at])
if err != nil || reference.Path(repoRef) != repository || !sameRegistry(reference.Domain(repoRef), registry) {
continue
}
if _, hasTag := repoRef.(reference.Tagged); hasTag {
continue
}
d, err := digest.Parse(repoDigest[at+1:])
if err != nil {
continue
}
digests = append(digests, d.String())
}
return digests
}
func sameRegistry(left, right string) bool {
left = canonicalRegistry(left)
right = canonicalRegistry(right)
return left == right ||
(left == "ghcr.io" && right == "lscr.io") ||
(left == "lscr.io" && right == "ghcr.io")
}
func canonicalRegistry(registry string) string {
if registry == "index.docker.io" {
return "docker.io"
}
return registry
}
func (dm *dockerManager) registryImageDigest(registry, repository, tag string) (string, error) {
client := dm.registryClient
if client == nil {
client = &http.Client{Timeout: imageRegistryTimeout}
}
token, err := dm.registryToken(client, registry, repository)
if err != nil {
return "", err
}
host := registry
if registry == "docker.io" {
host = "registry-1.docker.io"
}
manifestURL := "https://" + host + "/v2/" + repository + "/manifests/" + url.PathEscape(tag)
req, err := http.NewRequest(http.MethodHead, manifestURL, nil)
if err != nil {
return "", fmt.Errorf("create manifest request: %w", err)
}
req.Header.Set("Accept", imageManifestAccept)
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
resp, err := client.Do(req)
if err != nil {
return "", fmt.Errorf("fetch manifest %s:%s: %w", registry, repository, err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("manifest request for %s:%s failed: %s", repository, tag, responseStatus(resp))
}
remote := strings.TrimSpace(resp.Header.Get("Docker-Content-Digest"))
d, err := digest.Parse(remote)
if err != nil {
return "", fmt.Errorf("manifest request for %s:%s returned invalid digest: %w", repository, tag, err)
}
return d.String(), nil
}
func (dm *dockerManager) registryToken(client *http.Client, registry, repository string) (string, error) {
var authURL string
switch registry {
case "docker.io":
authURL = "https://auth.docker.io/token?service=registry.docker.io&scope=" + url.QueryEscape("repository:"+repository+":pull")
case "ghcr.io", "lscr.io":
// lscr.io is the LinuxServer alias for its GHCR-backed images.
authURL = "https://ghcr.io/token?service=ghcr.io&scope=" + url.QueryEscape("repository:"+repository+":pull")
default:
// Anonymous registries remain supported, as they were before the
// authenticated Docker Hub and GHCR paths were added.
return "", nil
}
req, err := http.NewRequest(http.MethodGet, authURL, nil)
if err != nil {
return "", fmt.Errorf("create registry auth request: %w", err)
}
resp, err := client.Do(req)
if err != nil {
return "", fmt.Errorf("fetch registry auth token for %s: %w", repository, err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("registry auth request for %s failed: %s", repository, responseStatus(resp))
}
var tokenResponse struct {
Token string `json:"token"`
AccessToken string `json:"access_token"`
}
if err := json.NewDecoder(resp.Body).Decode(&tokenResponse); err != nil {
return "", fmt.Errorf("decode registry auth response for %s: %w", repository, err)
}
token := strings.TrimSpace(tokenResponse.Token)
if token == "" {
token = strings.TrimSpace(tokenResponse.AccessToken)
}
if token == "" {
return "", fmt.Errorf("registry auth response for %s contained no token", repository)
}
return token, nil
}
func responseStatus(resp *http.Response) string {
if resp.Status != "" {
return resp.Status
}
return http.StatusText(resp.StatusCode)
}

View File

@@ -1,243 +0,0 @@
//go:build testing
package agent
import (
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"github.com/stretchr/testify/require"
)
type registryTransportFunc func(*http.Request) (*http.Response, error)
func (fn registryTransportFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return fn(req)
}
func registryResponse(status int, body string) *http.Response {
return &http.Response{
StatusCode: status,
Status: fmt.Sprintf("%d %s", status, http.StatusText(status)),
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(body)),
}
}
func registryDigest(fill byte) string {
return "sha256:" + strings.Repeat(string(fill), 64)
}
func newRegistryChecker(t *testing.T, inspectBody string, transport http.RoundTripper) *dockerManager {
t.Helper()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasPrefix(r.URL.Path, "/images/") {
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, inspectBody)
return
}
http.NotFound(w, r)
}))
t.Cleanup(server.Close)
return &dockerManager{
client: newDockerManagerForVersionTest(server).client,
registryClient: &http.Client{Transport: transport},
}
}
func TestCheckImageUpdateUsesInspectAndManifestDigests(t *testing.T) {
local := registryDigest('a')
remote := registryDigest('b')
var authCalls, manifestCalls atomic.Int32
dm := newRegistryChecker(t, fmt.Sprintf(`{"RepoDigests":["docker.io/library/alpine@%s"]}`, local), registryTransportFunc(func(req *http.Request) (*http.Response, error) {
switch {
case req.Method == http.MethodGet && req.URL.Host == "auth.docker.io":
authCalls.Add(1)
require.Equal(t, "/token", req.URL.Path)
return registryResponse(http.StatusOK, `{"token":"test-token"}`), nil
case req.Method == http.MethodHead && req.URL.Host == "registry-1.docker.io":
manifestCalls.Add(1)
require.Equal(t, "/v2/library/alpine/manifests/latest", req.URL.Path)
require.Equal(t, "Bearer test-token", req.Header.Get("Authorization"))
resp := registryResponse(http.StatusOK, "")
resp.Header.Set("Docker-Content-Digest", remote)
return resp, nil
default:
return registryResponse(http.StatusNotFound, ""), nil
}
}))
available, err := dm.checkImageUpdate("alpine")
require.NoError(t, err)
require.True(t, available)
require.EqualValues(t, 1, authCalls.Load())
require.EqualValues(t, 1, manifestCalls.Load())
}
func TestCheckImageUpdateMatchesAnyRepositoryDigest(t *testing.T) {
platform := registryDigest('a')
index := registryDigest('b')
other := registryDigest('c')
for _, test := range []struct {
name string
digests []string
remote string
available bool
}{
{name: "platform then index, remote index", digests: []string{platform, index}, remote: index},
{name: "index then platform, remote index", digests: []string{index, platform}, remote: index},
{name: "platform then index, remote platform", digests: []string{platform, index}, remote: platform},
{name: "index then platform, remote platform", digests: []string{index, platform}, remote: platform},
{name: "neither matches", digests: []string{platform, index}, remote: other, available: true},
} {
t.Run(test.name, func(t *testing.T) {
inspect := fmt.Sprintf(`{"RepoDigests":["docker.io/library/busybox@%s","docker.io/library/alpine@%s","docker.io/library/alpine@sha256:invalid","docker.io/library/alpine@%s"]}`, test.remote, test.digests[0], test.digests[1])
var manifestCalls atomic.Int32
dm := newRegistryChecker(t, inspect, registryTransportFunc(func(req *http.Request) (*http.Response, error) {
if req.Method == http.MethodGet {
return registryResponse(http.StatusOK, `{"token":"test"}`), nil
}
manifestCalls.Add(1)
require.Equal(t, http.MethodHead, req.Method)
resp := registryResponse(http.StatusOK, "")
resp.Header.Set("Docker-Content-Digest", test.remote)
return resp, nil
}))
available, err := dm.checkImageUpdate("alpine")
require.NoError(t, err)
require.Equal(t, test.available, available)
require.EqualValues(t, 1, manifestCalls.Load())
})
}
}
func TestCheckImageUpdateReportsUnknownInspectState(t *testing.T) {
for _, test := range []struct {
name string
body string
}{
{name: "missing field", body: `{}`},
{name: "empty field", body: `{"RepoDigests":[]}`},
{name: "malformed reference", body: `{"RepoDigests":["not-a-repo-digest"]}`},
{name: "wrong repository", body: `{"RepoDigests":["docker.io/library/busybox@` + registryDigest('a') + `"]}`},
{name: "malformed digest", body: `{"RepoDigests":["docker.io/library/alpine@sha256:not-a-digest"]}`},
} {
t.Run(test.name, func(t *testing.T) {
var registryCalls atomic.Int32
dm := newRegistryChecker(t, test.body, registryTransportFunc(func(req *http.Request) (*http.Response, error) {
registryCalls.Add(1)
return registryResponse(http.StatusOK, `{"token":"unexpected"}`), nil
}))
available, err := dm.checkImageUpdate("alpine")
require.Error(t, err)
require.False(t, available)
require.EqualValues(t, 0, registryCalls.Load(), "invalid local state must not query a registry")
})
}
}
func TestCheckImageUpdateChecksInspectAuthAndManifestStatuses(t *testing.T) {
local := registryDigest('a')
validInspect := fmt.Sprintf(`{"RepoDigests":["docker.io/library/alpine@%s"]}`, local)
tests := []struct {
name string
inspectCode int
authCode int
manifestCode int
remote string
want string
}{
{name: "inspect status", inspectCode: http.StatusNotFound, want: "inspect image"},
{name: "auth status", inspectCode: http.StatusOK, authCode: http.StatusUnauthorized, want: "registry auth"},
{name: "manifest status", inspectCode: http.StatusOK, authCode: http.StatusOK, manifestCode: http.StatusNotFound, remote: local, want: "manifest request"},
{name: "missing digest", inspectCode: http.StatusOK, authCode: http.StatusOK, manifestCode: http.StatusOK, want: "invalid digest"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if test.inspectCode != http.StatusOK && strings.HasPrefix(r.URL.Path, "/images/") {
w.WriteHeader(test.inspectCode)
return
}
_, _ = io.WriteString(w, validInspect)
}))
t.Cleanup(server.Close)
calls := 0
dm := &dockerManager{client: newDockerManagerForVersionTest(server).client, registryClient: &http.Client{Transport: registryTransportFunc(func(req *http.Request) (*http.Response, error) {
calls++
if req.Method == http.MethodGet {
return registryResponse(test.authCode, `{"token":"test"}`), nil
}
response := registryResponse(test.manifestCode, "")
response.Header.Set("Docker-Content-Digest", test.remote)
return response, nil
})}}
_, err := dm.checkImageUpdate("alpine")
require.Error(t, err)
require.Contains(t, err.Error(), test.want)
if test.inspectCode != http.StatusOK {
require.Zero(t, calls)
}
})
}
}
func TestCheckImageUpdateSupportsAnonymousAndLSCRRegistries(t *testing.T) {
t.Run("anonymous registry", func(t *testing.T) {
local := registryDigest('a')
var calls atomic.Int32
dm := newRegistryChecker(t, fmt.Sprintf(`{"RepoDigests":["example.com/app@%s"]}`, local), registryTransportFunc(func(req *http.Request) (*http.Response, error) {
calls.Add(1)
require.Equal(t, http.MethodHead, req.Method)
require.Equal(t, "example.com", req.URL.Host)
resp := registryResponse(http.StatusOK, "")
resp.Header.Set("Docker-Content-Digest", local)
return resp, nil
}))
available, err := dm.checkImageUpdate("example.com/app")
require.NoError(t, err)
require.False(t, available)
require.EqualValues(t, 1, calls.Load())
})
t.Run("lscr ghcr alias", func(t *testing.T) {
local := registryDigest('a')
var authCalls, manifestCalls atomic.Int32
dm := newRegistryChecker(t, fmt.Sprintf(`{"RepoDigests":["ghcr.io/linuxserver/app@%s"]}`, local), registryTransportFunc(func(req *http.Request) (*http.Response, error) {
if req.Method == http.MethodGet {
authCalls.Add(1)
return registryResponse(http.StatusOK, `{"token":"test"}`), nil
}
manifestCalls.Add(1)
require.Equal(t, "lscr.io", req.URL.Host)
resp := registryResponse(http.StatusOK, "")
resp.Header.Set("Docker-Content-Digest", local)
return resp, nil
}))
available, err := dm.checkImageUpdate("lscr.io/linuxserver/app")
require.NoError(t, err)
require.False(t, available)
require.EqualValues(t, 1, authCalls.Load())
require.EqualValues(t, 1, manifestCalls.Load())
})
}
func TestCheckImageUpdateSkipsPinnedDigest(t *testing.T) {
image := "docker.io/library/alpine@" + registryDigest('a')
dm := &dockerManager{}
available, err := dm.checkImageUpdate(image)
require.NoError(t, err)
require.False(t, available)
}

View File

@@ -729,7 +729,6 @@ func TestGetDockerStatsChecksDockerVersionAfterContainerList(t *testing.T) {
stats, err := dm.getDockerStats(defaultCacheTimeMs)
require.NoError(t, err)
require.NotNil(t, stats, "A successful empty snapshot must remain distinguishable from a collection failure")
assert.Empty(t, stats)
assert.True(t, dm.dockerVersionChecked)
assert.Equal(t, tt.expectedGood, dm.goodDockerVersion)
@@ -743,7 +742,6 @@ func TestGetDockerStatsChecksDockerVersionAfterContainerList(t *testing.T) {
stats, err = dm.getDockerStats(defaultCacheTimeMs)
require.NoError(t, err)
require.NotNil(t, stats, "A successful empty snapshot must remain distinguishable from a collection failure")
assert.Empty(t, stats)
assert.Equal(t, tt.expectedGood, dm.goodDockerVersion)
assert.Equal(t, tt.expectedPodman, dm.usingPodman)
@@ -806,24 +804,6 @@ func TestGetDockerStatsRetriesVersionCheckUntilSuccess(t *testing.T) {
assert.Equal(t, 2, requestCounts["/version"])
}
// A failed decode must not break later decodes. Previously the reused json.Decoder
// stayed desynced after one truncated response, breaking decode until restart.
func TestDecodeRecoversFromError(t *testing.T) {
dm := &dockerManager{}
// truncated JSON: body reads fine, decode fails
var bad []container.ApiInfo
err := dm.decode(&http.Response{Body: io.NopCloser(strings.NewReader(`[{"Id":"abc`))}, &bad)
require.Error(t, err)
// the next decode must still succeed
var good []container.ApiInfo
err = dm.decode(&http.Response{Body: io.NopCloser(strings.NewReader(`[{"Id":"abcdef012345","Names":["/ok"]}]`))}, &good)
require.NoError(t, err)
require.Len(t, good, 1)
assert.Equal(t, "abcdef012345", good[0].Id)
}
func TestCycleCpuDeltas(t *testing.T) {
dm := &dockerManager{
lastCpuContainer: map[uint16]map[string]uint64{
@@ -1023,199 +1003,6 @@ func TestCpuPercentageCalculationWithRealData(t *testing.T) {
assert.InDelta(t, expectedPct, actualPct, 0.01)
}
func TestCpuPercentageHandlesCounterRollback(t *testing.T) {
// If a stats response is processed after a newer one for the same container,
// or an accounting counter resets, the current total can be lower than the
// stored previous value. Unsigned subtraction wraps to ~2^64 instead of
// going negative, so the percentage explodes, validateCpuPercentage rejects
// the sample, and the whole collection is discarded - network stats too.
stats := &container.ApiStats{
CPUStats: container.CPUStats{
CPUUsage: container.CPUUsage{TotalUsage: 1_000_000},
SystemUsage: 20_000_000,
},
}
// Container counter went backwards.
assert.Equal(t, 0.0, stats.CalculateCpuPercentLinux(2_000_000, 10_000_000))
// System counter went backwards.
assert.Equal(t, 0.0, stats.CalculateCpuPercentLinux(500_000, 30_000_000))
// A normal forward sample is unaffected: 500000 / 10000000 * 100 = 5%.
assert.InDelta(t, 5.0, stats.CalculateCpuPercentLinux(500_000, 10_000_000), 0.001)
}
func TestCpuPercentageWindowsHandlesCounterRollback(t *testing.T) {
now := time.Now()
stats := &container.ApiStats{
Read: now,
NumProcs: 4,
CPUStats: container.CPUStats{
CPUUsage: container.CPUUsage{TotalUsage: 1_000_000},
},
}
prevRead := now.Add(-time.Second)
// Container counter went backwards.
assert.Equal(t, 0.0, stats.CalculateCpuPercentWindows(2_000_000, prevRead))
// A normal forward sample is unaffected.
assert.Greater(t, stats.CalculateCpuPercentWindows(500_000, prevRead), 0.0)
}
func TestCalculateCpuPercentPodman(t *testing.T) {
baseTime := time.Date(2026, 3, 15, 12, 0, 0, 0, time.UTC)
tests := []struct {
name string
prevCpuContainer uint64
prevRead time.Time
currentUsage uint64
currentRead time.Time
onlineCPUs uint32
expectedPct float64
}{
{
name: "normal calculation",
// container used 2ms of CPU over 1s with 2 CPUs → 0.1%
prevCpuContainer: 1_000_000_000,
prevRead: baseTime,
currentUsage: 1_002_000_000, // +2ms CPU time
currentRead: baseTime.Add(time.Second),
onlineCPUs: 2,
expectedPct: 0.1, // 2e6 / (1e9 * 2) * 100
},
{
name: "first run returns zero",
prevCpuContainer: 0,
prevRead: baseTime,
currentUsage: 5_000_000,
currentRead: baseTime.Add(time.Second),
onlineCPUs: 4,
expectedPct: 0.0,
},
{
name: "zero online cpus returns zero",
prevCpuContainer: 1_000_000_000,
prevRead: baseTime,
currentUsage: 1_010_000_000,
currentRead: baseTime.Add(time.Second),
onlineCPUs: 0,
expectedPct: 0.0,
},
{
name: "same read time returns zero",
prevCpuContainer: 1_000_000_000,
prevRead: baseTime,
currentUsage: 1_010_000_000,
currentRead: baseTime, // no elapsed time
onlineCPUs: 2,
expectedPct: 0.0,
},
{
name: "counter rollback returns zero",
prevCpuContainer: 2_000_000_000,
prevRead: baseTime,
currentUsage: 1_000_000_000,
currentRead: baseTime.Add(time.Second),
onlineCPUs: 2,
expectedPct: 0.0,
},
{
name: "100% single cpu",
// container consumed a full CPU-second over 1s on a 1-CPU host → 100%
prevCpuContainer: 1_000_000_000,
prevRead: baseTime,
currentUsage: 2_000_000_000, // +1s CPU time
currentRead: baseTime.Add(time.Second),
onlineCPUs: 1,
expectedPct: 100.0, // 1e9 / (1e9 * 1) * 100
},
{
name: "high utilization on multi-cpu host",
// container used 800ms on a 4-CPU host over 1s → 20%
prevCpuContainer: 10_000_000_000,
prevRead: baseTime,
currentUsage: 10_800_000_000,
currentRead: baseTime.Add(time.Second),
onlineCPUs: 4,
expectedPct: 20.0, // 800e6 / (1e9 * 4) * 100
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
s := &container.ApiStats{
Read: tt.currentRead,
CPUStats: container.CPUStats{
CPUUsage: container.CPUUsage{TotalUsage: tt.currentUsage},
OnlineCPUs: tt.onlineCPUs,
},
}
got := s.CalculateCpuPercentPodman(tt.prevCpuContainer, tt.prevRead)
assert.InDelta(t, tt.expectedPct, got, 0.001, "test %q", tt.name)
})
}
}
func TestUpdateContainerStatsPodmanCpuCalculation(t *testing.T) {
// Verify that Podman containers use the time-based CPU calculation
// when online_cpus is provided in the stats response.
// container used 20ms CPU over 1s with 2 CPUs → 1%
prevReadTime := time.Date(2026, 3, 15, 21, 26, 58, 0, time.UTC) // 1 second before stats read
const prevCpuUsage = uint64(5_000_000_000)
dm := &dockerManager{
client: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
switch req.URL.EscapedPath() {
case "/containers/0123456789ab/stats":
return &http.Response{
StatusCode: http.StatusOK,
Status: "200 OK",
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{
"read":"2026-03-15T21:26:59Z",
"cpu_stats":{"cpu_usage":{"total_usage":5020000000},"system_cpu_usage":9999999,"online_cpus":2},
"memory_stats":{"usage":1048576,"stats":{"inactive_file":262144}},
"networks":{"eth0":{"rx_bytes":0,"tx_bytes":0}}
}`)),
Request: req,
}, nil
default:
return nil, fmt.Errorf("unexpected path: %s", req.URL.EscapedPath())
}
})},
containerStatsMap: make(map[string]*container.Stats),
usingPodman: true,
lastCpuContainer: map[uint16]map[string]uint64{
defaultCacheTimeMs: {"0123456789ab": prevCpuUsage},
},
lastCpuSystem: map[uint16]map[string]uint64{
defaultCacheTimeMs: {"0123456789ab": 1}, // intentionally tiny — should NOT be used
},
lastCpuReadTime: map[uint16]map[string]time.Time{
defaultCacheTimeMs: {"0123456789ab": prevReadTime},
},
networkSentTrackers: make(map[uint16]*deltatracker.DeltaTracker[string, uint64]),
networkRecvTrackers: make(map[uint16]*deltatracker.DeltaTracker[string, uint64]),
lastNetworkReadTime: make(map[uint16]map[string]time.Time),
}
ctr := &container.ApiInfo{
IdShort: "0123456789ab",
Names: []string{"/myapp"},
Status: "Up 5 minutes",
Image: "myapp:latest",
}
err := dm.updateContainerStats(ctr, defaultCacheTimeMs)
require.NoError(t, err)
// cpu delta = 5020000000 - 5000000000 = 20000000 ns (20ms)
// elapsed = 1s = 1000000000 ns, online_cpus = 2
// expected = 20000000 / (1000000000 * 2) * 100 = 1.0%
expectedCpu := 1.0
assert.InDelta(t, expectedCpu, dm.containerStatsMap[ctr.IdShort].Cpu, 0.01)
}
func TestNetworkStatsCalculationWithRealData(t *testing.T) {
// Create synthetic test data to avoid timing issues
apiStats1 := &container.ApiStats{
@@ -1675,6 +1462,7 @@ func TestUpdateContainerStatsUsesPodmanInspectHealthFallback(t *testing.T) {
}
})},
containerStatsMap: make(map[string]*container.Stats),
apiStats: &container.ApiStats{},
usingPodman: true,
lastCpuContainer: make(map[uint16]map[string]uint64),
lastCpuSystem: make(map[uint16]map[string]uint64),
@@ -2076,14 +1864,6 @@ func TestConvertContainerPortsToString(t *testing.T) {
},
expected: "80, 443",
},
{
name: "ipv4 and ipv6 wildcard bindings are deduplicated",
ports: []port{
{PublicPort: 80, IP: "0.0.0.0"},
{PublicPort: 80, IP: "::"},
},
expected: "80",
},
{
name: "multiple ports with different IPs",
ports: []port{
@@ -2092,22 +1872,6 @@ func TestConvertContainerPortsToString(t *testing.T) {
},
expected: "80, 1.2.3.4:443",
},
{
name: "same port bound to multiple IPs shows all entries",
ports: []port{
{PublicPort: 65533, IP: "172.16.151.72"},
{PublicPort: 65533, IP: "172.16.156.25"},
},
expected: "172.16.151.72:65533, 172.16.156.25:65533",
},
{
name: "same port bound to IPv4 and IPv6",
ports: []port{
{PublicPort: 65534, IP: "172.16.151.72"},
{PublicPort: 65534, IP: "fd04:38e2:98c6:3fd::72"},
},
expected: "172.16.151.72:65534, fd04:38e2:98c6:3fd::72:65534",
},
{
name: "ports slice is nilled after call",
ports: []port{

View File

@@ -1,133 +0,0 @@
package agent
import (
"log/slog"
"os"
"path/filepath"
"strings"
"sync"
"github.com/henrygd/beszel/agent/utils"
"github.com/henrygd/beszel/internal/entities/system"
)
type fanSensor struct {
key, path, chip string
}
var getFanSensors = newFanSensorCache(hwmonRoot)
func newFanSensorCache(root string) func() ([]fanSensor, error) {
return sync.OnceValues(func() ([]fanSensor, error) {
return discoverHwmonFans(root)
})
}
// updateFans populates systemStats.Fans from the host's hwmon sysfs tree.
// No-op on platforms where hwmon isn't available (see fans_other.go).
func (a *Agent) updateFans(systemStats *system.Stats) {
if hwmonRoot == "" {
return
}
sensors, err := getFanSensors()
if err != nil {
slog.Debug("Error reading fans", "err", err)
return
}
// Filter before reading fan*_input: each read can wake an idle GPU.
if a.sensorConfig != nil && a.sensorConfig.skipGPU {
sensors = filterGpuFans(sensors)
}
fans := readFanSensors(sensors)
if len(fans) == 0 {
return
}
systemStats.Fans = fans
// Note: Commented out because we don't currently use this value in the UI.
// Compute the single "dashboard" value used by the FanSpeed alert.
// Per-sensor RPMs live in Stats.Fans and drive the multi-line FanChart
// in the UI; the alert path only needs one number to compare against
// the user's threshold, so we use the highest RPM across all fans
// a.systemInfo.DashboardFan = 0
// for _, rpm := range fans {
// if rpm > a.systemInfo.DashboardFan {
// a.systemInfo.DashboardFan = rpm
// }
// }
}
// readHwmonFans walks the given hwmon root (typically /sys/class/hwmon) and
// returns a map of "<chip>_<label-or-fan-idx>" → RPM for every fan*_input
// file it finds. Zero RPM is retained because it can represent a real fan that
// has stopped; negative and malformed readings are ignored.
func readHwmonFans(root string) (map[string]uint16, error) {
sensors, err := discoverHwmonFans(root)
if err != nil {
return nil, err
}
return readFanSensors(sensors), nil
}
func discoverHwmonFans(root string) ([]fanSensor, error) {
entries, err := os.ReadDir(root)
if err != nil {
return nil, err
}
var sensors []fanSensor
for _, entry := range entries {
chipDir := filepath.Join(root, entry.Name())
sensorDir := chipDir
inputs, _ := filepath.Glob(filepath.Join(sensorDir, "fan*_input"))
// Some legacy hwmon drivers (notably applesmc) register a hwmon class
// device but create fan attributes on the parent platform device. In
// sysfs that parent is exposed through hwmonN/device.
if len(inputs) == 0 {
deviceDir := filepath.Join(chipDir, "device")
if deviceInputs, _ := filepath.Glob(filepath.Join(deviceDir, "fan*_input")); len(deviceInputs) > 0 {
sensorDir = deviceDir
inputs = deviceInputs
}
}
chipName := utils.ReadStringFile(filepath.Join(sensorDir, "name"))
if chipName == "" {
chipName = utils.ReadStringFile(filepath.Join(chipDir, "name"))
}
if chipName == "" {
chipName = entry.Name()
}
for _, inputPath := range inputs {
base := strings.TrimSuffix(filepath.Base(inputPath), "_input")
label := utils.ReadStringFile(filepath.Join(sensorDir, base+"_label"))
key := chipName + "_" + base
if label != "" {
key = chipName + "_" + label
}
sensors = append(sensors, fanSensor{key, inputPath, chipName})
}
}
return sensors, nil
}
func readFanSensors(sensors []fanSensor) map[string]uint16 {
fans := make(map[string]uint16, len(sensors))
for _, sensor := range sensors {
if rpm, ok := utils.ReadUintFile(sensor.path); ok {
fans[sensor.key] = uint16(rpm)
}
}
return fans
}
// filterGpuFans drops GPU chips without touching the shared cache backing array.
func filterGpuFans(sensors []fanSensor) []fanSensor {
kept := make([]fanSensor, 0, len(sensors))
for _, sensor := range sensors {
if isGpuChipName(sensor.chip) {
continue
}
kept = append(kept, sensor)
}
return kept
}

View File

@@ -1,8 +0,0 @@
//go:build linux
package agent
// hwmonRoot is the sysfs entry point for hardware monitor chips. Each
// subdirectory (hwmon0, hwmon1, …) is one chip; fan*_input files inside it
// expose RPM readings.
const hwmonRoot = "/sys/class/hwmon"

View File

@@ -1,7 +0,0 @@
//go:build !linux
package agent
// hwmonRoot is empty on non-Linux platforms — fan RPM reporting via sysfs
// hwmon is Linux-specific. updateFans() short-circuits when this is empty.
const hwmonRoot = ""

View File

@@ -1,122 +0,0 @@
//go:build testing
package agent
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// writeFile creates path with parents and writes contents.
func writeFile(t *testing.T, path, contents string) {
t.Helper()
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o755))
require.NoError(t, os.WriteFile(path, []byte(contents), 0o644))
}
// TestReadHwmonFans verifies the /sys/class/hwmon walker:
// - picks up fan*_input from every chip,
// - keys entries by chip name + sensor label (or fan idx if no label),
// - retains 0 RPM for stopped fans,
// - tolerates chips with no fan files at all.
func TestReadHwmonFans(t *testing.T) {
root := t.TempDir()
// hwmon0: Raspberry Pi 5 active cooler — one fan, no label.
writeFile(t, filepath.Join(root, "hwmon0", "name"), "pwmfan\n")
writeFile(t, filepath.Join(root, "hwmon0", "fan1_input"), "6500\n")
// hwmon1: a thermal-only chip, no fan files. Must not error.
writeFile(t, filepath.Join(root, "hwmon1", "name"), "cpu_thermal\n")
writeFile(t, filepath.Join(root, "hwmon1", "temp1_input"), "55000\n")
// hwmon2: two fans — one stopped (0 RPM) and one labeled "chassis".
writeFile(t, filepath.Join(root, "hwmon2", "name"), "nct6798\n")
writeFile(t, filepath.Join(root, "hwmon2", "fan1_input"), "0\n")
writeFile(t, filepath.Join(root, "hwmon2", "fan2_input"), "1200\n")
writeFile(t, filepath.Join(root, "hwmon2", "fan2_label"), "chassis\n")
fans, err := readHwmonFans(root)
require.NoError(t, err)
assert.Equal(t, map[string]uint16{
"pwmfan_fan1": 6500,
"nct6798_fan1": 0,
"nct6798_chassis": 1200,
}, fans)
}
// TestReadHwmonFansLegacyParent verifies legacy hwmon layouts such as applesmc,
// where the hwmon class node exists but fan attributes live on hwmonN/device.
func TestReadHwmonFansLegacyParent(t *testing.T) {
root := t.TempDir()
deviceDir := filepath.Join(root, "devices", "applesmc.768")
writeFile(t, filepath.Join(deviceDir, "name"), "applesmc\n")
writeFile(t, filepath.Join(deviceDir, "fan1_input"), "1202\n")
writeFile(t, filepath.Join(deviceDir, "fan1_label"), "Exhaust\n")
chipDir := filepath.Join(root, "hwmon1")
require.NoError(t, os.MkdirAll(chipDir, 0o755))
require.NoError(t, os.Symlink(deviceDir, filepath.Join(chipDir, "device")))
fans, err := readHwmonFans(root)
require.NoError(t, err)
assert.Equal(t, map[string]uint16{"applesmc_Exhaust": 1202}, fans)
}
// TestReadHwmonFansMissingRoot returns an error rather than panicking when the
// hwmon root doesn't exist (e.g. running on a kernel without hwmon support).
func TestReadHwmonFansMissingRoot(t *testing.T) {
_, err := readHwmonFans(filepath.Join(t.TempDir(), "does-not-exist"))
assert.Error(t, err)
}
// TestReadHwmonFansEmpty returns an empty map (not nil error) when the root
// exists but contains no chips at all.
func TestReadHwmonFansEmpty(t *testing.T) {
root := t.TempDir()
fans, err := readHwmonFans(root)
require.NoError(t, err)
assert.Empty(t, fans)
}
func TestFanDiscoveryCache(t *testing.T) {
root := t.TempDir()
input := filepath.Join(root, "hwmon0", "fan1_input")
writeFile(t, filepath.Join(root, "hwmon0", "name"), "chip\n")
writeFile(t, input, "1000\n")
getSensors := newFanSensorCache(root)
sensors, err := getSensors()
require.NoError(t, err)
fans := readFanSensors(sensors)
assert.Equal(t, uint16(1000), fans["chip_fan1"])
writeFile(t, input, "1200\n")
writeFile(t, filepath.Join(root, "hwmon0", "fan1_label"), "case\n")
sensors, err = getSensors()
require.NoError(t, err)
fans = readFanSensors(sensors)
assert.Equal(t, map[string]uint16{"chip_fan1": 1200}, fans)
}
func TestFilterGpuFans(t *testing.T) {
root := t.TempDir()
writeFile(t, filepath.Join(root, "hwmon0", "name"), "xe\n")
writeFile(t, filepath.Join(root, "hwmon0", "fan1_input"), "1200\n")
writeFile(t, filepath.Join(root, "hwmon1", "name"), "nct6798\n")
writeFile(t, filepath.Join(root, "hwmon1", "fan1_input"), "800\n")
discovered, err := discoverHwmonFans(root)
require.NoError(t, err)
require.Len(t, discovered, 2)
filtered := filterGpuFans(discovered)
require.Len(t, filtered, 1)
assert.Equal(t, "nct6798_fan1", filtered[0].key)
assert.Len(t, discovered, 2)
}

View File

@@ -50,9 +50,6 @@ func generateFingerprint(hostname, cpuModel string) string {
if info, err := cpu.Info(); err == nil && len(info) > 0 {
cpuModel = info[0].ModelName
}
if cpuModel == "" {
cpuModel = getCpuModelFromCpuinfo()
}
}
fingerprint = hostname + cpuModel
}

View File

@@ -48,8 +48,6 @@ type GPUManager struct {
// Per-cache-key tracking for delta calculations
// cacheKey -> gpuId -> snapshot of last count/usage/power values
lastSnapshots map[uint16]map[string]*gpuSnapshot
// Per-card energy snapshots for Intel sysfs power calculation.
intelSysfsEnergySnapshots map[string]intelSysfsEnergySnapshot
}
// gpuSnapshot stores the last observed incremental values for delta tracking
@@ -92,7 +90,6 @@ const (
collectorSourceNVML collectorSource = "nvml"
collectorSourceNvidiaSMI collectorSource = collectorSource(nvidiaSmiCmd)
collectorSourceIntelGpuTop collectorSource = collectorSource(intelGpuStatsCmd)
collectorSourceIntelSysfs collectorSource = "intel_sysfs"
collectorSourceAmdSysfs collectorSource = "amd_sysfs"
collectorSourceRocmSMI collectorSource = collectorSource(rocmSmiCmd)
collectorSourceMacmon collectorSource = collectorSource(macmonCmd)
@@ -109,7 +106,6 @@ func isValidCollectorSource(source collectorSource) bool {
collectorSourceNVML,
collectorSourceNvidiaSMI,
collectorSourceIntelGpuTop,
collectorSourceIntelSysfs,
collectorSourceAmdSysfs,
collectorSourceRocmSMI,
collectorSourceMacmon,
@@ -126,8 +122,6 @@ type gpuCapabilities struct {
hasAmdSysfs bool
hasTegrastats bool
hasIntelGpuTop bool
hasXe bool
hasIntelSysfs bool
hasNvtop bool
hasMacmon bool
hasPowermetrics bool
@@ -361,16 +355,12 @@ func (gm *GPUManager) calculateGPUAverage(id string, gpu *system.GPUData, cacheK
// If no new data arrived
if deltaCount == 0 {
// Only discrete GPUs report temp/memory, so treat all-zero as suspended (return zeros).
// Engine-based (Intel) GPUs don't, so carry the last average forward across sample gaps.
if gpu.Engines == nil && gpu.Temperature == 0 && gpu.MemoryUsed == 0 {
// If GPU appears suspended (instantaneous values are 0), return zero values
// Otherwise return last known average for temporary collection gaps
if gpu.Temperature == 0 && gpu.MemoryUsed == 0 {
return system.GPUData{Name: gpu.Name}
}
lastAvg := gm.lastAvgData[id] // zero value if not found
if lastAvg.Name == "" {
lastAvg.Name = gpu.Name
}
return lastAvg
return gm.lastAvgData[id] // zero value if not found
}
// Calculate new average
@@ -379,13 +369,12 @@ func (gm *GPUManager) calculateGPUAverage(id string, gpu *system.GPUData, cacheK
gpuAvg.Power = utils.TwoDecimals(deltaPower / float64(deltaCount))
gpuAvg.PowerPkg = utils.TwoDecimals(deltaPowerPkg / float64(deltaCount))
if gpu.Engines != nil {
// make fresh map for averaged engine metrics to avoid mutating
// the accumulator map stored in gm.GpuDataMap
gpuAvg.Engines = make(map[string]float64, len(gpu.Engines))
gpuAvg.Usage = gm.calculateIntelGPUUsage(&gpuAvg, gpu, lastSnapshot, deltaCount)
gpuAvg.PowerPkg = utils.TwoDecimals(deltaPowerPkg / float64(deltaCount))
} else {
gpuAvg.Usage = utils.TwoDecimals(deltaUsage / float64(deltaCount))
}
@@ -454,9 +443,7 @@ func (gm *GPUManager) storeSnapshot(id string, gpu *system.GPUData, cacheKey uin
// It only reports capability presence and does not apply policy decisions.
func (gm *GPUManager) discoverGpuCapabilities() gpuCapabilities {
caps := gpuCapabilities{
hasAmdSysfs: gm.hasAmdSysfs(),
hasXe: gm.hasXe(),
hasIntelSysfs: gm.hasIntelSysfs(),
hasAmdSysfs: gm.hasAmdSysfs(),
}
if _, err := exec.LookPath(nvidiaSmiCmd); err == nil {
caps.hasNvidiaSmi = true
@@ -485,7 +472,7 @@ func (gm *GPUManager) discoverGpuCapabilities() gpuCapabilities {
}
func hasAnyGpuCollector(caps gpuCapabilities) bool {
return caps.hasNvidiaSmi || caps.hasRocmSmi || caps.hasAmdSysfs || caps.hasTegrastats || caps.hasIntelGpuTop || caps.hasIntelSysfs || caps.hasNvtop || caps.hasMacmon || caps.hasPowermetrics
return caps.hasNvidiaSmi || caps.hasRocmSmi || caps.hasAmdSysfs || caps.hasTegrastats || caps.hasIntelGpuTop || caps.hasNvtop || caps.hasMacmon || caps.hasPowermetrics
}
func (gm *GPUManager) startIntelCollector() {
@@ -576,13 +563,6 @@ func (gm *GPUManager) collectorDefinitions(caps gpuCapabilities) map[collectorSo
return true
},
},
collectorSourceIntelSysfs: {
group: collectorGroupIntel,
available: caps.hasIntelSysfs,
start: func(_ func()) bool {
return gm.startIntelSysfsCollector()
},
},
collectorSourceAmdSysfs: {
group: collectorGroupAmd,
available: caps.hasAmdSysfs,
@@ -725,12 +705,9 @@ func (gm *GPUManager) resolveLegacyCollectorPriority(caps gpuCapabilities) []col
priorities = append(priorities, collectorSourceAmdSysfs)
}
if caps.hasIntelGpuTop && !caps.hasXe {
if caps.hasIntelGpuTop {
priorities = append(priorities, collectorSourceIntelGpuTop)
}
if caps.hasIntelSysfs {
priorities = append(priorities, collectorSourceIntelSysfs)
}
// Apple collectors are currently opt-in only for testing.
// Enable them with GPU_COLLECTOR=macmon or GPU_COLLECTOR=powermetrics.
@@ -750,36 +727,9 @@ func (gm *GPUManager) resolveLegacyCollectorPriority(caps gpuCapabilities) []col
return priorities
}
// gpuHwmonChips are hwmon chip names belonging to GPUs. Sensor reads on some
// of these drivers (notably Intel Xe, where each read is a runtime PM resume)
// wake the card, so SKIP_GPU must avoid touching them, not just hide them.
var gpuHwmonChips = []string{"xe", "i915", "amdgpu", "radeon", "nvidia", "nouveau"}
func isGpuChipName(name string) bool {
name = strings.ToLower(strings.TrimSpace(name))
for _, chip := range gpuHwmonChips {
if name == chip {
return true
}
}
return false
}
// SensorKeys are "<chip>" or "<chip>_<label>".
func isGpuSensorKey(key string) bool {
key = strings.ToLower(strings.TrimSpace(key))
for _, chip := range gpuHwmonChips {
if key == chip || strings.HasPrefix(key, chip+"_") {
return true
}
}
return false
}
// NewGPUManager creates and initializes a new GPUManager
func NewGPUManager() (*GPUManager, error) {
if skipGPU, _ := utils.GetEnv("SKIP_GPU"); skipGPU == "true" {
slog.Info("SKIP_GPU enabled, skipping GPU monitoring (collectors, temperatures, and fans)")
return nil, nil
}
var gm GPUManager

View File

@@ -1,85 +0,0 @@
//go:build testing
package agent
import (
"fmt"
"io"
"os"
"os/exec"
"path/filepath"
"runtime"
"strings"
"testing"
"github.com/stretchr/testify/require"
)
// Run a copy of the test binary as a GPU command so fixtures do not need a shell.
func TestMain(m *testing.M) {
executable, err := os.Executable()
if err != nil {
panic(err)
}
switch strings.TrimSuffix(filepath.Base(executable), ".exe") {
case nvidiaSmiCmd, rocmSmiCmd, tegraStatsCmd, nvtopCmd, intelGpuStatsCmd:
output, err := os.ReadFile(executable + ".stdout")
if err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
// Only the parent creates files; a late collector must not undo cleanup.
args, err := os.OpenFile(executable+".args", os.O_WRONLY|os.O_TRUNC, 0)
if err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
_, err = io.WriteString(args, strings.Join(os.Args[1:], " "))
closeErr := args.Close()
if err == nil {
err = closeErr
}
if err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
fmt.Print(string(output))
os.Exit(0)
}
os.Exit(m.Run())
}
func gpuCommandFixture(t *testing.T, dir, name, output string) string {
t.Helper()
executable, err := os.Executable()
require.NoError(t, err)
if runtime.GOOS == "windows" {
name += ".exe"
}
path := filepath.Join(dir, name)
if err := os.Link(executable, path); err != nil {
src, err := os.Open(executable)
require.NoError(t, err)
defer src.Close()
dst, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_EXCL, 0755)
require.NoError(t, err)
_, err = io.Copy(dst, src)
closeErr := dst.Close()
require.NoError(t, err)
require.NoError(t, closeErr)
}
require.NoError(t, os.WriteFile(path+".stdout", []byte(output), 0600))
require.NoError(t, os.WriteFile(path+".args", nil, 0600))
return path + ".args"
}
func TestGPUFixtureDoesNotRecreateRemovedArgs(t *testing.T) {
argsFile := gpuCommandFixture(t, t.TempDir(), nvidiaSmiCmd, "fixture output\n")
require.NoError(t, os.WriteFile(argsFile, nil, 0600))
require.NoError(t, os.Remove(argsFile))
cmd := exec.Command(strings.TrimSuffix(argsFile, ".args"))
err := cmd.Run()
require.NoFileExists(t, argsFile, "a late fixture process must not recreate files removed by cleanup")
require.Error(t, err, "the fixture must report a missing argument-capture file")
}

View File

@@ -2,9 +2,7 @@ package agent
import (
"bufio"
"encoding/json"
"io"
"log/slog"
"os/exec"
"strconv"
"strings"
@@ -51,10 +49,10 @@ func (gm *GPUManager) updateIntelFromStats(sample *intelGpuStats) bool {
return true
}
// collectIntelStats executes intel_gpu_top in JSON mode (-J) and parses the output.
// collectIntelStats executes intel_gpu_top in text mode (-l) and parses the output
func (gm *GPUManager) collectIntelStats() (err error) {
// Build command arguments, optionally selecting a device via -d
args := []string{"-s", intelGpuStatsInterval, "-J"}
args := []string{"-s", intelGpuStatsInterval, "-l"}
if dev, ok := utils.GetEnv("INTEL_GPU_DEVICE"); ok && dev != "" {
args = append(args, "-d", dev)
}
@@ -82,64 +80,48 @@ func (gm *GPUManager) collectIntelStats() (err error) {
}
}()
if err := gm.parseIntelJSONStream(stdout); err != nil {
return err
}
// The closing "]" is printed as the process exits, so read to EOF to let
// it finish instead of killing it.
_, _ = io.Copy(io.Discard, stdout)
return nil
}
// parseIntelJSONStream decodes samples from intel_gpu_top -J output and
// aggregates them. Since v1.28 the samples are wrapped in an array ("[", then
// comma separated objects, and "]" only when the process exits). Older
// versions print the same comma separated objects without the opening "[", so
// it is added here to let both formats decode as an array.
func (gm *GPUManager) parseIntelJSONStream(r io.Reader) error {
er := &eofReader{r: r}
br := bufio.NewReader(er)
first, err := peekNonSpace(br)
if err != nil {
if err == io.EOF {
return errNoValidData
}
return err
}
var src io.Reader = br
if first != '[' {
src = io.MultiReader(strings.NewReader("["), br)
}
dec := json.NewDecoder(src)
if _, err := dec.Token(); err != nil { // opening "["
return err
}
scanner := bufio.NewScanner(stdout)
var header1 string
var engineNames []string
var friendlyNames []string
var preEngineCols int
var powerIndex int
var hadDataRow bool
// skip first data row because it sometimes has erroneous data
var skippedFirstDataRow bool
// Decode reads one object and skips the commas between them. The array is
// usually never closed, so output ending mid-array or mid-sample (the
// process was killed) is the normal end of the stream rather than an error.
for dec.More() {
var sample intelGpuJSONSample
if err := dec.Decode(&sample); err != nil {
if er.eof {
break
}
return err
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" {
continue
}
// first header line
if strings.HasPrefix(line, "Freq") {
header1 = line
continue
}
// second header line
if strings.HasPrefix(line, "req") {
engineNames, friendlyNames, powerIndex, preEngineCols = gm.parseIntelHeaders(header1, line)
continue
}
// Data row
if !skippedFirstDataRow {
skippedFirstDataRow = true
continue
}
stats := parseIntelJSONSample(sample)
if !validIntelPower(stats.PowerGPU) || !validIntelPower(stats.PowerPkg) {
slog.Debug("Skipping intel_gpu_top sample with invalid power", "gpu", stats.PowerGPU, "pkg", stats.PowerPkg)
continue
sample, err := gm.parseIntelData(line, engineNames, friendlyNames, powerIndex, preEngineCols)
if err != nil {
return err
}
hadDataRow = true
gm.updateIntelFromStats(&stats)
gm.updateIntelFromStats(&sample)
}
if scanErr := scanner.Err(); scanErr != nil {
return scanErr
}
if !hadDataRow {
return errNoValidData
@@ -147,82 +129,80 @@ func (gm *GPUManager) parseIntelJSONStream(r io.Reader) error {
return nil
}
// eofReader records whether the underlying reader has returned io.EOF. The
// json decoder reports a stream ending mid-value as a syntax error, so this
// is how a truncated final sample is told apart from invalid output.
type eofReader struct {
r io.Reader
eof bool
}
func (e *eofReader) Read(p []byte) (int, error) {
n, err := e.r.Read(p)
if err == io.EOF {
e.eof = true
}
return n, err
}
// peekNonSpace discards leading JSON whitespace and returns the next byte without consuming it.
func peekNonSpace(br *bufio.Reader) (byte, error) {
for {
b, err := br.Peek(1)
if err != nil {
return 0, err
}
switch b[0] {
case ' ', '\t', '\n', '\r':
_, _ = br.ReadByte()
func (gm *GPUManager) parseIntelHeaders(header1 string, header2 string) (engineNames []string, friendlyNames []string, powerIndex int, preEngineCols int) {
// Build indexes
h1 := strings.Fields(header1)
h2 := strings.Fields(header2)
powerIndex = -1 // Initialize to -1, will be set to actual index if found
// Collect engine names from header1
for _, col := range h1 {
key := strings.TrimRightFunc(col, func(r rune) bool {
return (r >= '0' && r <= '9') || r == '/'
})
var friendly string
switch key {
case "RCS":
friendly = "Render/3D"
case "BCS":
friendly = "Blitter"
case "VCS":
friendly = "Video"
case "VECS":
friendly = "VideoEnhance"
case "CCS":
friendly = "Compute"
default:
return b[0], nil
continue
}
engineNames = append(engineNames, key)
friendlyNames = append(friendlyNames, friendly)
}
// find power gpu index among pre-engine columns
if n := len(engineNames); n > 0 {
preEngineCols = max(len(h2)-3*n, 0)
limit := min(len(h2), preEngineCols)
for i := range limit {
if strings.EqualFold(h2[i], "gpu") {
powerIndex = i
break
}
}
}
return engineNames, friendlyNames, powerIndex, preEngineCols
}
// intelGpuJSONSample is a single sample from intel_gpu_top -J output. Only the
// needed fields are mapped.
type intelGpuJSONSample struct {
Power *struct {
GPU float64 `json:"GPU"`
Package float64 `json:"Package"`
} `json:"power"`
Engines map[string]struct {
Busy float64 `json:"busy"`
} `json:"engines"`
}
// validIntelPower reports whether a power reading from intel_gpu_top is plausible.
func validIntelPower(watts float64) bool {
// 5000 is well above any real GPU or package draw. intel_gpu_top
// computes power from unsigned energy counter deltas, so a counter that reads
// lower than the previous sample produces an enormous value for that period.
return watts >= 0 && watts <= 5000
}
// parseIntelJSONSample converts one intel_gpu_top JSON sample into intelGpuStats.
func parseIntelJSONSample(sample intelGpuJSONSample) (stats intelGpuStats) {
if sample.Power != nil {
stats.PowerGPU = sample.Power.GPU
stats.PowerPkg = sample.Power.Package
func (gm *GPUManager) parseIntelData(line string, engineNames []string, friendlyNames []string, powerIndex int, preEngineCols int) (sample intelGpuStats, err error) {
fields := strings.Fields(line)
if len(fields) == 0 {
return sample, errNoValidData
}
if len(sample.Engines) > 0 {
stats.Engines = make(map[string]float64, len(sample.Engines))
for key, engine := range sample.Engines {
stats.Engines[intelEngineClass(key)] += engine.Busy
// Make sure row has enough columns for engines
if need := preEngineCols + 3*len(engineNames); len(fields) < need {
return sample, errNoValidData
}
if powerIndex >= 0 && powerIndex < len(fields) {
if v, perr := strconv.ParseFloat(fields[powerIndex], 64); perr == nil {
sample.PowerGPU = v
}
if v, perr := strconv.ParseFloat(fields[powerIndex+1], 64); perr == nil {
sample.PowerPkg = v
}
}
return stats
}
// intelEngineClass returns the engine class name for an engine key. Keys are
// class names ("Render/3D", "Video") in class view, which JSON output uses by
// default since v1.28, and instance names ("Render/3D/0", "Video/1") in
// physical view, which older versions use.
func intelEngineClass(key string) string {
if i := strings.LastIndexByte(key, '/'); i >= 0 {
if _, err := strconv.ParseUint(key[i+1:], 10, 32); err == nil {
return key[:i]
if len(engineNames) > 0 {
sample.Engines = make(map[string]float64, len(engineNames))
for k := range engineNames {
base := preEngineCols + 3*k
if base < len(fields) {
busy := 0.0
if v, e := strconv.ParseFloat(fields[base], 64); e == nil {
busy = v
}
cur := sample.Engines[friendlyNames[k]]
sample.Engines[friendlyNames[k]] = cur + busy
} else {
sample.Engines[friendlyNames[k]] = 0
}
}
}
return key
return sample, nil
}

View File

@@ -1,280 +0,0 @@
//go:build linux
package agent
import (
"fmt"
"log/slog"
"os"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/henrygd/beszel/agent/utils"
"github.com/henrygd/beszel/internal/entities/system"
)
var (
drmSysfsRoot = "/sys/class/drm"
intelSysfsNow = time.Now
)
type intelSysfsEnergySnapshot struct {
microjoules uint64
timestamp time.Time
}
type intelSysfsCard struct {
cardPath string
hwmonDir string
}
// hasIntelSysfs returns true if any Intel DRM card exposes an hwmon energy counter.
func (gm *GPUManager) hasIntelSysfs() bool {
cards, err := discoverIntelSysfsCards()
return err == nil && len(cards) > 0
}
// startIntelSysfsCollector starts Intel GPU collection via sysfs.
func (gm *GPUManager) startIntelSysfsCollector() bool {
go func() {
if err := gm.collectIntelSysfsStats(); err != nil {
slog.Warn("Error collecting Intel GPU data via sysfs", "err", err)
}
}()
return true
}
// collectIntelSysfsStats collects Intel GPU metrics directly from DRM sysfs / hwmon.
func (gm *GPUManager) collectIntelSysfsStats() error {
sysfsPollInterval := 3000 * time.Millisecond
cards, err := discoverIntelSysfsCards()
if err != nil {
return err
}
if len(cards) == 0 {
return errNoValidData
}
slog.Debug("Using sysfs for Intel GPU data collection", "cards", len(cards))
for _, card := range cards {
slog.Debug("Intel sysfs card detected", "card", filepath.Base(card.cardPath), "hwmon", card.hwmonDir)
}
failures := 0
for {
hasData := false
for _, card := range cards {
if gm.updateIntelSysfsGpuData(card.cardPath, card.hwmonDir) {
hasData = true
}
}
if !hasData {
failures++
if failures > maxFailureRetries {
return errNoValidData
}
slog.Warn("No Intel GPU data from sysfs", "failures", failures)
time.Sleep(retryWaitTime)
continue
}
failures = 0
time.Sleep(sysfsPollInterval)
}
}
func discoverIntelSysfsCards() ([]intelSysfsCard, error) {
paths, err := filepath.Glob(filepath.Join(drmSysfsRoot, "card*"))
if err != nil {
return nil, err
}
var cards []intelSysfsCard
for _, cardPath := range paths {
if strings.Contains(filepath.Base(cardPath), "-") || !isIntelGpu(cardPath) {
continue
}
hwmonDir := findIntelEnergyHwmon(filepath.Join(cardPath, "device"))
if hwmonDir == "" {
continue
}
cards = append(cards, intelSysfsCard{cardPath: cardPath, hwmonDir: hwmonDir})
}
return cards, nil
}
func isIntelGpu(cardPath string) bool {
vendor, err := utils.ReadStringFileLimited(filepath.Join(cardPath, "device/vendor"), 64)
if err != nil {
return false
}
return strings.EqualFold(strings.TrimSpace(vendor), "0x8086")
}
func findIntelEnergyHwmon(devicePath string) string {
hwmons, _ := filepath.Glob(filepath.Join(devicePath, "hwmon/hwmon*"))
var fallback string
for _, hwmonDir := range hwmons {
if !sysfsFileExists(filepath.Join(hwmonDir, "energy1_input")) {
continue
}
if name, err := utils.ReadStringFileLimited(filepath.Join(hwmonDir, "name"), 64); err == nil && strings.EqualFold(strings.TrimSpace(name), "xe") {
return hwmonDir
}
if fallback == "" {
fallback = hwmonDir
}
}
return fallback
}
func sysfsFileExists(path string) bool {
_, err := utils.ReadStringFileLimited(path, 1)
return err == nil
}
// updateIntelSysfsGpuData reads GPU metrics from sysfs and updates the GPU data map.
// Returns true if the required energy counter was read successfully.
func (gm *GPUManager) updateIntelSysfsGpuData(cardPath, hwmonDir string) bool {
devicePath := filepath.Join(cardPath, "device")
id := filepath.Base(cardPath)
energy, err := readSysfsUint(filepath.Join(hwmonDir, "energy1_input"))
if err != nil {
return false
}
now := intelSysfsNow()
power, hasPower := gm.calculateIntelSysfsPower(id, energy, now)
powerPkg, hasPowerPkg := gm.readIntelSysfsPowerPkg(id, hwmonDir, now)
temp := readIntelSysfsTemperature(hwmonDir)
usage, usageErr := readOptionalSysfsFloat(filepath.Join(devicePath, "gpu_busy_percent"))
memUsed, memUsedErr := readFirstOptionalSysfsFloat(
filepath.Join(devicePath, "mem_info_vram_used"),
filepath.Join(devicePath, "mem_info_lmem_used"),
filepath.Join(devicePath, "mem_info_local_mem_used"),
)
memTotal, memTotalErr := readFirstOptionalSysfsFloat(
filepath.Join(devicePath, "mem_info_vram_total"),
filepath.Join(devicePath, "mem_info_lmem_total"),
filepath.Join(devicePath, "mem_info_local_mem_total"),
)
gm.Lock()
defer gm.Unlock()
gpu, ok := gm.GpuDataMap[id]
if !ok {
gpu = &system.GPUData{Name: getIntelSysfsGpuName(cardPath)}
gm.GpuDataMap[id] = gpu
}
if usageErr == nil {
gpu.Usage += usage
}
if memUsedErr == nil {
gpu.MemoryUsed = utils.BytesToMegabytes(memUsed)
}
if memTotalErr == nil {
gpu.MemoryTotal = utils.BytesToMegabytes(memTotal)
}
if temp > 0 {
gpu.Temperature = temp
}
if hasPower {
gpu.Power += power
slog.Debug("Computed Intel sysfs GPU power", "card", id, "watts", power)
}
if hasPowerPkg {
gpu.PowerPkg += powerPkg
}
gpu.Count++
return true
}
func (gm *GPUManager) calculateIntelSysfsPower(cardID string, microjoules uint64, timestamp time.Time) (float64, bool) {
if gm.intelSysfsEnergySnapshots == nil {
gm.intelSysfsEnergySnapshots = make(map[string]intelSysfsEnergySnapshot)
}
last, ok := gm.intelSysfsEnergySnapshots[cardID]
gm.intelSysfsEnergySnapshots[cardID] = intelSysfsEnergySnapshot{microjoules: microjoules, timestamp: timestamp}
if !ok {
return 0, false
}
if microjoules < last.microjoules {
slog.Debug("Intel sysfs energy counter reset", "card", cardID)
return 0, false
}
elapsed := timestamp.Sub(last.timestamp).Seconds()
if elapsed <= 0 {
return 0, false
}
delta := microjoules - last.microjoules
return float64(delta) / 1_000_000.0 / elapsed, true
}
func (gm *GPUManager) readIntelSysfsPowerPkg(cardID, hwmonDir string, timestamp time.Time) (float64, bool) {
energyPaths, _ := filepath.Glob(filepath.Join(hwmonDir, "energy*_input"))
for _, path := range energyPaths {
if filepath.Base(path) == "energy1_input" {
continue
}
energy, err := readSysfsUint(path)
if err != nil {
continue
}
return gm.calculateIntelSysfsPower(cardID+":"+filepath.Base(path), energy, timestamp)
}
return 0, false
}
func readIntelSysfsTemperature(hwmonDir string) float64 {
tempPaths, _ := filepath.Glob(filepath.Join(hwmonDir, "temp*_input"))
for _, path := range tempPaths {
temp, err := readSysfsFloat(path)
if err == nil && temp > 0 {
return temp / 1000.0
}
}
return 0
}
func readSysfsUint(path string) (uint64, error) {
val, err := utils.ReadStringFileLimited(path, 64)
if err != nil {
slog.Debug("Failed to read sysfs value", "path", path, "error", err)
return 0, err
}
return strconv.ParseUint(strings.TrimSpace(val), 10, 64)
}
func readOptionalSysfsFloat(path string) (float64, error) {
val, err := os.ReadFile(path)
if err != nil {
return 0, err
}
return strconv.ParseFloat(strings.TrimSpace(string(val)), 64)
}
func readFirstOptionalSysfsFloat(paths ...string) (float64, error) {
for _, path := range paths {
val, err := readOptionalSysfsFloat(path)
if err == nil {
return val, nil
}
}
return 0, fmt.Errorf("no sysfs values found")
}
func getIntelSysfsGpuName(cardPath string) string {
devicePath := filepath.Join(cardPath, "device")
if product, err := utils.ReadStringFileLimited(filepath.Join(devicePath, "product_name"), 128); err == nil && strings.TrimSpace(product) != "" {
return strings.TrimSpace(product)
}
if name, err := utils.ReadStringFileLimited(filepath.Join(devicePath, "name"), 128); err == nil && strings.TrimSpace(name) != "" {
return strings.TrimSpace(name)
}
return fmt.Sprintf("Intel GPU %s", filepath.Base(cardPath))
}

View File

@@ -1,217 +0,0 @@
//go:build linux
package agent
import (
"os"
"path/filepath"
"testing"
"time"
"github.com/henrygd/beszel/agent/utils"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func setupIntelSysfsTest(t *testing.T) (root, cardPath, hwmonPath string) {
t.Helper()
root = t.TempDir()
oldRoot := drmSysfsRoot
drmSysfsRoot = root
t.Cleanup(func() {
drmSysfsRoot = oldRoot
})
cardPath = filepath.Join(root, "card0")
devicePath := filepath.Join(cardPath, "device")
hwmonPath = filepath.Join(devicePath, "hwmon", "hwmon0")
require.NoError(t, os.MkdirAll(hwmonPath, 0o755))
return root, cardPath, hwmonPath
}
func writeIntelSysfsFile(t *testing.T, basePath, name, content string) {
t.Helper()
require.NoError(t, os.WriteFile(filepath.Join(basePath, name), []byte(content), 0o644))
}
func setIntelSysfsTime(t *testing.T, now time.Time) {
t.Helper()
oldNow := intelSysfsNow
intelSysfsNow = func() time.Time { return now }
t.Cleanup(func() {
intelSysfsNow = oldNow
})
}
func TestIntelSysfsDetectsIntelCardWithEnergy(t *testing.T) {
_, cardPath, hwmonPath := setupIntelSysfsTest(t)
devicePath := filepath.Join(cardPath, "device")
writeIntelSysfsFile(t, devicePath, "vendor", "0x8086\n")
writeIntelSysfsFile(t, hwmonPath, "name", "xe\n")
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "1000000\n")
gm := &GPUManager{}
assert.True(t, gm.hasIntelSysfs())
cards, err := discoverIntelSysfsCards()
require.NoError(t, err)
require.Len(t, cards, 1)
assert.Equal(t, cardPath, cards[0].cardPath)
assert.Equal(t, hwmonPath, cards[0].hwmonDir)
}
func TestIntelSysfsRejectsNonIntelCard(t *testing.T) {
_, cardPath, hwmonPath := setupIntelSysfsTest(t)
devicePath := filepath.Join(cardPath, "device")
writeIntelSysfsFile(t, devicePath, "vendor", "0x1002\n")
writeIntelSysfsFile(t, hwmonPath, "name", "xe\n")
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "1000000\n")
gm := &GPUManager{}
assert.False(t, gm.hasIntelSysfs())
}
func TestIntelSysfsRequiresEnergyInput(t *testing.T) {
_, cardPath, hwmonPath := setupIntelSysfsTest(t)
devicePath := filepath.Join(cardPath, "device")
writeIntelSysfsFile(t, devicePath, "vendor", "0x8086\n")
writeIntelSysfsFile(t, hwmonPath, "name", "xe\n")
gm := &GPUManager{}
assert.False(t, gm.hasIntelSysfs())
}
func TestIntelSysfsFirstSampleInitializesWithoutBogusPower(t *testing.T) {
_, cardPath, hwmonPath := setupIntelSysfsTest(t)
devicePath := filepath.Join(cardPath, "device")
writeIntelSysfsFile(t, devicePath, "vendor", "0x8086\n")
writeIntelSysfsFile(t, hwmonPath, "name", "xe\n")
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "1000000\n")
setIntelSysfsTime(t, time.Unix(100, 0))
gm := &GPUManager{GpuDataMap: make(map[string]*system.GPUData)}
ok := gm.updateIntelSysfsGpuData(cardPath, hwmonPath)
require.True(t, ok)
gpu := gm.GpuDataMap["card0"]
require.NotNil(t, gpu)
assert.Equal(t, "Intel GPU card0", gpu.Name)
assert.Equal(t, 0.0, gpu.Power)
assert.Equal(t, 1.0, gpu.Count)
}
func TestIntelSysfsSecondSampleComputesWatts(t *testing.T) {
_, cardPath, hwmonPath := setupIntelSysfsTest(t)
devicePath := filepath.Join(cardPath, "device")
writeIntelSysfsFile(t, devicePath, "vendor", "0x8086\n")
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "1000000\n")
gm := &GPUManager{GpuDataMap: make(map[string]*system.GPUData)}
oldNow := intelSysfsNow
intelSysfsNow = func() time.Time { return time.Unix(100, 0) }
t.Cleanup(func() { intelSysfsNow = oldNow })
require.True(t, gm.updateIntelSysfsGpuData(cardPath, hwmonPath))
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "6000000\n")
intelSysfsNow = func() time.Time { return time.Unix(102, 0) }
require.True(t, gm.updateIntelSysfsGpuData(cardPath, hwmonPath))
gpu := gm.GpuDataMap["card0"]
require.NotNil(t, gpu)
assert.Equal(t, 2.5, gpu.Power)
assert.Equal(t, 2.0, gpu.Count)
}
func TestIntelSysfsSecondEnergyCounterMapsToPowerPkg(t *testing.T) {
_, cardPath, hwmonPath := setupIntelSysfsTest(t)
devicePath := filepath.Join(cardPath, "device")
writeIntelSysfsFile(t, devicePath, "vendor", "0x8086\n")
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "1000000\n")
writeIntelSysfsFile(t, hwmonPath, "energy2_input", "2000000\n")
oldNow := intelSysfsNow
t.Cleanup(func() { intelSysfsNow = oldNow })
gm := &GPUManager{GpuDataMap: make(map[string]*system.GPUData)}
intelSysfsNow = func() time.Time { return time.Unix(100, 0) }
require.True(t, gm.updateIntelSysfsGpuData(cardPath, hwmonPath))
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "2000000\n")
writeIntelSysfsFile(t, hwmonPath, "energy2_input", "8000000\n")
intelSysfsNow = func() time.Time { return time.Unix(102, 0) }
require.True(t, gm.updateIntelSysfsGpuData(cardPath, hwmonPath))
gpu := gm.GpuDataMap["card0"]
require.NotNil(t, gpu)
assert.Equal(t, 0.5, gpu.Power)
assert.Equal(t, 3.0, gpu.PowerPkg)
}
func TestIntelSysfsCounterResetSkipsOneSample(t *testing.T) {
gm := &GPUManager{}
power, ok := gm.calculateIntelSysfsPower("card0", 5000000, time.Unix(100, 0))
assert.False(t, ok)
assert.Equal(t, 0.0, power)
power, ok = gm.calculateIntelSysfsPower("card0", 1000000, time.Unix(101, 0))
assert.False(t, ok)
assert.Equal(t, 0.0, power)
power, ok = gm.calculateIntelSysfsPower("card0", 3000000, time.Unix(103, 0))
assert.True(t, ok)
assert.Equal(t, 1.0, power)
}
func TestIntelSysfsTempInputMapsToCelsius(t *testing.T) {
_, cardPath, hwmonPath := setupIntelSysfsTest(t)
devicePath := filepath.Join(cardPath, "device")
writeIntelSysfsFile(t, devicePath, "vendor", "0x8086\n")
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "1000000\n")
writeIntelSysfsFile(t, hwmonPath, "temp1_input", "43500\n")
setIntelSysfsTime(t, time.Unix(100, 0))
gm := &GPUManager{GpuDataMap: make(map[string]*system.GPUData)}
require.True(t, gm.updateIntelSysfsGpuData(cardPath, hwmonPath))
gpu := gm.GpuDataMap["card0"]
require.NotNil(t, gpu)
assert.Equal(t, 43.5, gpu.Temperature)
}
func TestIntelSysfsMissingOptionalFilesDoNotFail(t *testing.T) {
_, cardPath, hwmonPath := setupIntelSysfsTest(t)
devicePath := filepath.Join(cardPath, "device")
writeIntelSysfsFile(t, devicePath, "vendor", "0x8086\n")
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "1000000\n")
setIntelSysfsTime(t, time.Unix(100, 0))
gm := &GPUManager{GpuDataMap: make(map[string]*system.GPUData)}
require.True(t, gm.updateIntelSysfsGpuData(cardPath, hwmonPath))
gpu := gm.GpuDataMap["card0"]
require.NotNil(t, gpu)
assert.Equal(t, 0.0, gpu.Usage)
assert.Equal(t, 0.0, gpu.MemoryUsed)
assert.Equal(t, 0.0, gpu.MemoryTotal)
assert.Equal(t, 0.0, gpu.Temperature)
}
func TestIntelSysfsMapsOpportunisticMemoryAndUsage(t *testing.T) {
_, cardPath, hwmonPath := setupIntelSysfsTest(t)
devicePath := filepath.Join(cardPath, "device")
writeIntelSysfsFile(t, devicePath, "vendor", "0x8086\n")
writeIntelSysfsFile(t, devicePath, "gpu_busy_percent", "37\n")
writeIntelSysfsFile(t, devicePath, "mem_info_lmem_used", "1073741824\n")
writeIntelSysfsFile(t, devicePath, "mem_info_lmem_total", "2147483648\n")
writeIntelSysfsFile(t, hwmonPath, "energy1_input", "1000000\n")
setIntelSysfsTime(t, time.Unix(100, 0))
gm := &GPUManager{GpuDataMap: make(map[string]*system.GPUData)}
require.True(t, gm.updateIntelSysfsGpuData(cardPath, hwmonPath))
gpu := gm.GpuDataMap["card0"]
require.NotNil(t, gpu)
assert.Equal(t, 37.0, gpu.Usage)
assert.Equal(t, utils.BytesToMegabytes(1073741824), gpu.MemoryUsed)
assert.Equal(t, utils.BytesToMegabytes(2147483648), gpu.MemoryTotal)
}

View File

@@ -1,13 +0,0 @@
//go:build !linux
package agent
type intelSysfsEnergySnapshot struct{}
func (gm *GPUManager) hasIntelSysfs() bool {
return false
}
func (gm *GPUManager) startIntelSysfsCollector() bool {
return false
}

View File

@@ -1,30 +0,0 @@
//go:build testing && !(amd64 && (windows || (linux && glibc)))
package agent
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// This fallback requires NVML initialisation to fail, as guaranteed by the unsupported implementation.
func TestNewGPUManagerPriorityNvmlFallbackToNvidiaSmi(t *testing.T) {
dir := t.TempDir()
t.Setenv("PATH", dir)
t.Setenv("BESZEL_AGENT_GPU_COLLECTOR", "nvml,nvidia-smi")
gpuCommandFixture(t, dir, "nvidia-smi", `0, NVIDIA Fallback GPU, 41, 256, 1024, 8, 14`+"\n")
gm, err := NewGPUManager()
require.NoError(t, err)
require.NotNil(t, gm)
waitGPUs(t, gm, "0")
gm.Lock()
defer gm.Unlock()
gpu, ok := gm.GpuDataMap["0"]
require.True(t, ok)
assert.Equal(t, "Fallback GPU", gpu.Name)
}

View File

@@ -5,7 +5,6 @@ import (
"io"
"log/slog"
"os/exec"
"path/filepath"
"strconv"
"strings"
"time"
@@ -49,14 +48,9 @@ func (gm *GPUManager) updateNvtopSnapshots(snapshots []nvtopSnapshot) bool {
valid := false
usedIDs := make(map[string]struct{}, len(snapshots))
var xeName string
for i, sample := range snapshots {
// nvtop leaves device_name unset on xe devices.
if sample.DeviceName == "" {
if xeName == "" {
xeName = xeGpuName()
}
sample.DeviceName = xeName
continue
}
indexID := "n" + strconv.Itoa(i)
id := indexID
@@ -164,38 +158,3 @@ func (gm *GPUManager) startNvtopCollector(interval string, onFailure func()) {
}
}()
}
// xeDevicePath returns the sysfs device path of the first xe GPU, or "".
func xeDevicePath() string {
cards, err := filepath.Glob("/sys/class/drm/card*")
if err != nil {
return ""
}
for _, card := range cards {
if strings.Contains(filepath.Base(card), "-") {
continue
}
if uevent, err := utils.ReadStringFileLimited(filepath.Join(card, "device", "uevent"), 4096); err == nil && strings.Contains(uevent, "DRIVER=xe") {
return filepath.Join(card, "device")
}
}
return ""
}
func (gm *GPUManager) hasXe() bool {
return xeDevicePath() != ""
}
// xeGpuName names an xe GPU from its PCI device id; nvtop leaves device_name unset on xe.
func xeGpuName() string {
devicePath := xeDevicePath()
if devicePath == "" {
return "GPU"
}
id, err := utils.ReadStringFileLimited(filepath.Join(devicePath, "device"), 64)
if err != nil {
return "GPU"
}
id = strings.ToLower(strings.TrimSpace(strings.TrimPrefix(id, "0x")))
return "Intel GPU (" + id + ")"
}

View File

@@ -3,10 +3,9 @@
package agent
import (
"encoding/json"
"fmt"
"os"
"slices"
"path/filepath"
"strings"
"testing"
"time"
@@ -333,12 +332,11 @@ func TestUpdateNvtopSnapshotsKeepsDeviceAssociationWhenOrderChanges(t *testing.T
}
func TestParseCollectorPriority(t *testing.T) {
got := parseCollectorPriority(" nvml, nvidia-smi, intel_gpu_top, intel_sysfs, amd_sysfs, nvtop, rocm-smi, bad ")
got := parseCollectorPriority(" nvml, nvidia-smi, intel_gpu_top, amd_sysfs, nvtop, rocm-smi, bad ")
want := []collectorSource{
collectorSourceNVML,
collectorSourceNvidiaSMI,
collectorSourceIntelGpuTop,
collectorSourceIntelSysfs,
collectorSourceAmdSysfs,
collectorSourceNVTop,
collectorSourceRocmSMI,
@@ -567,42 +565,6 @@ func TestGetCurrentData(t *testing.T) {
assert.EqualValues(t, 2, gm.GpuDataMap["0"].Count, "Count should still be 2")
})
t.Run("carries Intel GPU average forward between samples", func(t *testing.T) {
// Intel GPUs report no temp/memory, so between-sample gaps (delta 0) must
// reuse the last average instead of returning zeros and blanking the chart.
gm := &GPUManager{
GpuDataMap: map[string]*system.GPUData{
"0": {
Name: "GPU",
Usage: 0, // derived from engines for Intel
Power: 200, // averages to 100 over 2 counts
PowerPkg: 60, // averages to 30 over 2 counts
Count: 2,
Engines: map[string]float64{
"Render/3D": 80, // averages to 40
"Video": 20, // averages to 10
},
},
},
}
cacheKey := uint16(1000) // realtime cache key
// First collection - computes and stores averages
result1 := gm.GetCurrentData(cacheKey)
assert.InDelta(t, 100.0, result1["0"].Power, 0.01)
assert.InDelta(t, 30.0, result1["0"].PowerPkg, 0.01)
assert.InDelta(t, 40.0, result1["0"].Engines["Render/3D"], 0.01)
// Second collection with no new sample (count unchanged, temp/mem still 0).
// Must carry the last average forward rather than blanking to zero.
result2 := gm.GetCurrentData(cacheKey)
assert.Equal(t, "GPU", result2["0"].Name, "Name should be preserved")
assert.InDelta(t, 100.0, result2["0"].Power, 0.01, "Should reuse last average power, not 0")
assert.InDelta(t, 30.0, result2["0"].PowerPkg, 0.01, "Should reuse last average package power, not 0")
assert.InDelta(t, 40.0, result2["0"].Engines["Render/3D"], 0.01, "Should reuse last average engine usage")
})
t.Run("tracks separate averages per cache key", func(t *testing.T) {
gm := &GPUManager{
GpuDataMap: map[string]*system.GPUData{
@@ -1120,11 +1082,12 @@ func TestCalculateGPUAverage(t *testing.T) {
}
func TestGPUCapabilitiesAndLegacyPriority(t *testing.T) {
// Save original PATH
hasAmdSysfs := (&GPUManager{}).hasAmdSysfs()
tests := []struct {
name string
setupCommands func(*testing.T, string)
setupCommands func(string) error
wantNvidiaSmi bool
wantRocmSmi bool
wantTegrastats bool
@@ -1132,8 +1095,10 @@ func TestGPUCapabilitiesAndLegacyPriority(t *testing.T) {
wantErr bool
}{
{
name: "nvidia-smi not available",
setupCommands: func(*testing.T, string) {},
name: "nvidia-smi not available",
setupCommands: func(_ string) error {
return nil
},
wantNvidiaSmi: false,
wantRocmSmi: false,
wantTegrastats: false,
@@ -1142,8 +1107,14 @@ func TestGPUCapabilitiesAndLegacyPriority(t *testing.T) {
},
{
name: "nvidia-smi available",
setupCommands: func(t *testing.T, tempDir string) {
gpuCommandFixture(t, tempDir, "nvidia-smi", "test\n")
setupCommands: func(tempDir string) error {
path := filepath.Join(tempDir, "nvidia-smi")
script := `#!/bin/sh
echo "test"`
if err := os.WriteFile(path, []byte(script), 0755); err != nil {
return err
}
return nil
},
wantNvidiaSmi: true,
wantTegrastats: false,
@@ -1153,8 +1124,14 @@ func TestGPUCapabilitiesAndLegacyPriority(t *testing.T) {
},
{
name: "rocm-smi available",
setupCommands: func(t *testing.T, tempDir string) {
gpuCommandFixture(t, tempDir, "rocm-smi", "test\n")
setupCommands: func(tempDir string) error {
path := filepath.Join(tempDir, "rocm-smi")
script := `#!/bin/sh
echo "test"`
if err := os.WriteFile(path, []byte(script), 0755); err != nil {
return err
}
return nil
},
wantNvidiaSmi: false,
wantRocmSmi: true,
@@ -1164,8 +1141,14 @@ func TestGPUCapabilitiesAndLegacyPriority(t *testing.T) {
},
{
name: "tegrastats available",
setupCommands: func(t *testing.T, tempDir string) {
gpuCommandFixture(t, tempDir, "tegrastats", "test\n")
setupCommands: func(tempDir string) error {
path := filepath.Join(tempDir, "tegrastats")
script := `#!/bin/sh
echo "test"`
if err := os.WriteFile(path, []byte(script), 0755); err != nil {
return err
}
return nil
},
wantNvidiaSmi: false,
wantRocmSmi: false,
@@ -1175,8 +1158,14 @@ func TestGPUCapabilitiesAndLegacyPriority(t *testing.T) {
},
{
name: "nvtop available",
setupCommands: func(t *testing.T, tempDir string) {
gpuCommandFixture(t, tempDir, "nvtop", "test\n")
setupCommands: func(tempDir string) error {
path := filepath.Join(tempDir, "nvtop")
script := `#!/bin/sh
echo "[]"`
if err := os.WriteFile(path, []byte(script), 0755); err != nil {
return err
}
return nil
},
wantNvidiaSmi: false,
wantRocmSmi: false,
@@ -1185,9 +1174,12 @@ func TestGPUCapabilitiesAndLegacyPriority(t *testing.T) {
wantErr: false,
},
{
name: "no gpu tools available",
setupCommands: func(*testing.T, string) {},
wantErr: true,
name: "no gpu tools available",
setupCommands: func(_ string) error {
t.Setenv("PATH", "")
return nil
},
wantErr: true,
},
}
@@ -1195,7 +1187,9 @@ func TestGPUCapabilitiesAndLegacyPriority(t *testing.T) {
t.Run(tt.name, func(t *testing.T) {
tempDir := t.TempDir()
t.Setenv("PATH", tempDir)
tt.setupCommands(t, tempDir)
if err := tt.setupCommands(tempDir); err != nil {
t.Fatal(err)
}
gm := &GPUManager{}
caps := gm.discoverGpuCapabilities()
@@ -1237,21 +1231,6 @@ func TestGPUCapabilitiesAndLegacyPriority(t *testing.T) {
}
}
func waitGPUs(t *testing.T, gm *GPUManager, ids ...string) {
t.Helper()
require.Eventually(t, func() bool {
gm.Lock()
defer gm.Unlock()
for _, id := range ids {
gpu := gm.GpuDataMap[id]
if gpu == nil || gpu.Count == 0 {
return false
}
}
return true
}, 5*time.Second, 10*time.Millisecond, "GPU collectors did not produce data for %v", ids)
}
func TestCollectorStartHelpers(t *testing.T) {
// Set up temp dir with the commands
dir := t.TempDir()
@@ -1260,17 +1239,21 @@ func TestCollectorStartHelpers(t *testing.T) {
tests := []struct {
name string
command string
gpuID string
setup func(t *testing.T)
setup func(t *testing.T) error
validate func(t *testing.T, gm *GPUManager)
gm *GPUManager
}{
{
name: "nvidia-smi collector",
command: "nvidia-smi",
gpuID: "0",
setup: func(t *testing.T) {
gpuCommandFixture(t, dir, "nvidia-smi", `0, NVIDIA Test GPU, 50, 1024, 4096, 25, 100`+"\n")
setup: func(t *testing.T) error {
path := filepath.Join(dir, "nvidia-smi")
script := `#!/bin/sh
echo "0, NVIDIA Test GPU, 50, 1024, 4096, 25, 100"`
if err := os.WriteFile(path, []byte(script), 0755); err != nil {
return err
}
return nil
},
validate: func(t *testing.T, gm *GPUManager) {
gpu, exists := gm.GpuDataMap["0"]
@@ -1285,9 +1268,14 @@ func TestCollectorStartHelpers(t *testing.T) {
{
name: "rocm-smi collector",
command: "rocm-smi",
gpuID: "34756",
setup: func(t *testing.T) {
gpuCommandFixture(t, dir, "rocm-smi", `{"card0": {"Temperature (Sensor edge) (C)": "49.0", "Current Socket Graphics Package Power (W)": "28.159", "GPU use (%)": "0", "VRAM Total Memory (B)": "536870912", "VRAM Total Used Memory (B)": "445550592", "Card Series": "Rembrandt [Radeon 680M]", "Card Model": "0x1681", "Card Vendor": "Advanced Micro Devices, Inc. [AMD/ATI]", "Card SKU": "REMBRANDT", "Subsystem ID": "0x8a22", "Device Rev": "0xc8", "Node ID": "1", "GUID": "34756", "GFX Version": "gfx1035"}}`+"\n")
setup: func(t *testing.T) error {
path := filepath.Join(dir, "rocm-smi")
script := `#!/bin/sh
echo '{"card0": {"Temperature (Sensor edge) (C)": "49.0", "Current Socket Graphics Package Power (W)": "28.159", "GPU use (%)": "0", "VRAM Total Memory (B)": "536870912", "VRAM Total Used Memory (B)": "445550592", "Card Series": "Rembrandt [Radeon 680M]", "Card Model": "0x1681", "Card Vendor": "Advanced Micro Devices, Inc. [AMD/ATI]", "Card SKU": "REMBRANDT", "Subsystem ID": "0x8a22", "Device Rev": "0xc8", "Node ID": "1", "GUID": "34756", "GFX Version": "gfx1035"}}'`
if err := os.WriteFile(path, []byte(script), 0755); err != nil {
return err
}
return nil
},
validate: func(t *testing.T, gm *GPUManager) {
gpu, exists := gm.GpuDataMap["34756"]
@@ -1302,9 +1290,14 @@ func TestCollectorStartHelpers(t *testing.T) {
{
name: "tegrastats collector",
command: "tegrastats",
gpuID: "0",
setup: func(t *testing.T) {
gpuCommandFixture(t, dir, "tegrastats", `11-14-2024 22:54:33 RAM 1024/4096MB GR3D_FREQ 80% tj@70C VDD_GPU_SOC 1000mW`+"\n")
setup: func(t *testing.T) error {
path := filepath.Join(dir, "tegrastats")
script := `#!/bin/sh
echo "11-14-2024 22:54:33 RAM 1024/4096MB GR3D_FREQ 80% tj@70C VDD_GPU_SOC 1000mW"`
if err := os.WriteFile(path, []byte(script), 0755); err != nil {
return err
}
return nil
},
validate: func(t *testing.T, gm *GPUManager) {
gpu, exists := gm.GpuDataMap["0"]
@@ -1322,9 +1315,14 @@ func TestCollectorStartHelpers(t *testing.T) {
{
name: "nvtop collector",
command: "nvtop",
gpuID: "n0",
setup: func(t *testing.T) {
gpuCommandFixture(t, dir, "nvtop", `[{"device_name":"NVIDIA Test GPU","temp":"52C","power_draw":"31W","gpu_util":"37%","mem_total":"4294967296","mem_used":"536870912","processes":[]}]`+"\n")
setup: func(t *testing.T) error {
path := filepath.Join(dir, "nvtop")
script := `#!/bin/sh
echo '[{"device_name":"NVIDIA Test GPU","temp":"52C","power_draw":"31W","gpu_util":"37%","mem_total":"4294967296","mem_used":"536870912","processes":[]}]'`
if err := os.WriteFile(path, []byte(script), 0755); err != nil {
return err
}
return nil
},
validate: func(t *testing.T, gm *GPUManager) {
gpu, exists := gm.GpuDataMap["n0"]
@@ -1339,7 +1337,9 @@ func TestCollectorStartHelpers(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
tt.setup(t)
if err := tt.setup(t); err != nil {
t.Fatal(err)
}
if tt.gm == nil {
tt.gm = &GPUManager{
GpuDataMap: make(map[string]*system.GPUData),
@@ -1357,9 +1357,7 @@ func TestCollectorStartHelpers(t *testing.T) {
default:
t.Fatalf("unknown test command %q", tt.command)
}
waitGPUs(t, tt.gm, tt.gpuID)
tt.gm.Lock()
defer tt.gm.Unlock()
time.Sleep(50 * time.Millisecond) // Give collector time to run
tt.validate(t, tt.gm)
})
}
@@ -1370,17 +1368,21 @@ func TestNewGPUManagerPriorityNvtopFallback(t *testing.T) {
t.Setenv("PATH", dir)
t.Setenv("BESZEL_AGENT_GPU_COLLECTOR", "nvtop,nvidia-smi")
gpuCommandFixture(t, dir, "nvtop", `not-json`+"\n")
nvtopPath := filepath.Join(dir, "nvtop")
nvtopScript := `#!/bin/sh
echo 'not-json'`
require.NoError(t, os.WriteFile(nvtopPath, []byte(nvtopScript), 0755))
gpuCommandFixture(t, dir, "nvidia-smi", `0, NVIDIA Priority GPU, 45, 512, 2048, 12, 25`+"\n")
nvidiaPath := filepath.Join(dir, "nvidia-smi")
nvidiaScript := `#!/bin/sh
echo "0, NVIDIA Priority GPU, 45, 512, 2048, 12, 25"`
require.NoError(t, os.WriteFile(nvidiaPath, []byte(nvidiaScript), 0755))
gm, err := NewGPUManager()
require.NoError(t, err)
require.NotNil(t, gm)
waitGPUs(t, gm, "0")
gm.Lock()
defer gm.Unlock()
time.Sleep(150 * time.Millisecond)
gpu, ok := gm.GpuDataMap["0"]
require.True(t, ok)
assert.Equal(t, "Priority GPU", gpu.Name)
@@ -1392,27 +1394,52 @@ func TestNewGPUManagerPriorityMixedCollectors(t *testing.T) {
t.Setenv("PATH", dir)
t.Setenv("BESZEL_AGENT_GPU_COLLECTOR", "intel_gpu_top,rocm-smi")
intelOutput := intelJSONStream(true,
intelJSONSample(2, 2.69, map[string]float64{"Render/3D": 0, "Video": 0}),
intelJSONSample(1.8, 2.45, map[string]float64{"Render/3D": 8.5, "Video": 15}),
)
gpuCommandFixture(t, dir, intelGpuStatsCmd, intelOutput+"\n")
intelPath := filepath.Join(dir, "intel_gpu_top")
intelScript := `#!/bin/sh
echo "Freq MHz IRQ RC6 Power W IMC MiB/s RCS VCS"
echo " req act /s % gpu pkg rd wr % se wa % se wa"
echo "226 223 338 58 2.00 2.69 1820 965 0.00 0 0 0.00 0 0"
echo "189 187 412 67 1.80 2.45 1950 823 8.50 2 1 15.00 1 0"
`
require.NoError(t, os.WriteFile(intelPath, []byte(intelScript), 0755))
gpuCommandFixture(t, dir, "rocm-smi", `{"card0": {"Temperature (Sensor edge) (C)": "49.0", "Current Socket Graphics Package Power (W)": "28.159", "GPU use (%)": "0", "VRAM Total Memory (B)": "536870912", "VRAM Total Used Memory (B)": "445550592", "Card Series": "Rembrandt [Radeon 680M]", "GUID": "34756"}}`+"\n")
rocmPath := filepath.Join(dir, "rocm-smi")
rocmScript := `#!/bin/sh
echo '{"card0": {"Temperature (Sensor edge) (C)": "49.0", "Current Socket Graphics Package Power (W)": "28.159", "GPU use (%)": "0", "VRAM Total Memory (B)": "536870912", "VRAM Total Used Memory (B)": "445550592", "Card Series": "Rembrandt [Radeon 680M]", "GUID": "34756"}}'
`
require.NoError(t, os.WriteFile(rocmPath, []byte(rocmScript), 0755))
gm, err := NewGPUManager()
require.NoError(t, err)
require.NotNil(t, gm)
waitGPUs(t, gm, "i0", "34756")
gm.Lock()
defer gm.Unlock()
time.Sleep(150 * time.Millisecond)
_, intelOk := gm.GpuDataMap["i0"]
_, amdOk := gm.GpuDataMap["34756"]
assert.True(t, intelOk)
assert.True(t, amdOk)
}
func TestNewGPUManagerPriorityNvmlFallbackToNvidiaSmi(t *testing.T) {
dir := t.TempDir()
t.Setenv("PATH", dir)
t.Setenv("BESZEL_AGENT_GPU_COLLECTOR", "nvml,nvidia-smi")
nvidiaPath := filepath.Join(dir, "nvidia-smi")
nvidiaScript := `#!/bin/sh
echo "0, NVIDIA Fallback GPU, 41, 256, 1024, 8, 14"`
require.NoError(t, os.WriteFile(nvidiaPath, []byte(nvidiaScript), 0755))
gm, err := NewGPUManager()
require.NoError(t, err)
require.NotNil(t, gm)
time.Sleep(150 * time.Millisecond)
gpu, ok := gm.GpuDataMap["0"]
require.True(t, ok)
assert.Equal(t, "Fallback GPU", gpu.Name)
}
func TestNewGPUManagerConfiguredCollectorsMustStart(t *testing.T) {
dir := t.TempDir()
t.Setenv("PATH", dir)
@@ -1447,12 +1474,8 @@ func TestNewGPUManagerConfiguredNvmlBypassesCapabilityGate(t *testing.T) {
t.Setenv("BESZEL_AGENT_GPU_COLLECTOR", "nvml")
gm, err := NewGPUManager()
if err == nil {
// Native NVML can be available even with no tools on PATH.
require.NotNil(t, gm)
return
}
require.Nil(t, gm)
require.Error(t, err)
assert.Contains(t, err.Error(), "no configured GPU collectors are available")
assert.NotContains(t, err.Error(), noGPUFoundMsg)
}
@@ -1462,15 +1485,16 @@ func TestNewGPUManagerJetsonIgnoresCollectorConfig(t *testing.T) {
t.Setenv("PATH", dir)
t.Setenv("BESZEL_AGENT_GPU_COLLECTOR", "nvidia-smi")
gpuCommandFixture(t, dir, "tegrastats", `11-14-2024 22:54:33 RAM 1024/4096MB GR3D_FREQ 80% tj@70C VDD_GPU_SOC 1000mW`+"\n")
tegraPath := filepath.Join(dir, "tegrastats")
tegraScript := `#!/bin/sh
echo "11-14-2024 22:54:33 RAM 1024/4096MB GR3D_FREQ 80% tj@70C VDD_GPU_SOC 1000mW"`
require.NoError(t, os.WriteFile(tegraPath, []byte(tegraScript), 0755))
gm, err := NewGPUManager()
require.NoError(t, err)
require.NotNil(t, gm)
waitGPUs(t, gm, "0")
gm.Lock()
defer gm.Unlock()
time.Sleep(100 * time.Millisecond)
gpu, ok := gm.GpuDataMap["0"]
require.True(t, ok)
assert.Equal(t, "GPU", gpu.Name)
@@ -1692,60 +1716,22 @@ func TestIntelUpdateFromStats(t *testing.T) {
assert.Equal(t, float64(2), gpu.Count)
}
// intelJSONSample returns one sample object formatted like intel_gpu_top -J output
func intelJSONSample(powerGPU, powerPkg float64, engines map[string]float64) string {
var sb strings.Builder
sb.WriteString("{\n\t\"period\": {\n\t\t\"duration\": 3300.123456,\n\t\t\"unit\": \"ms\"\n\t},\n")
sb.WriteString("\t\"frequency\": {\n\t\t\"requested\": 373.000000,\n\t\t\"actual\": 373.000000,\n\t\t\"unit\": \"MHz\"\n\t},\n")
fmt.Fprintf(&sb, "\t\"power\": {\n\t\t\"GPU\": %f,\n\t\t\"Package\": %f,\n\t\t\"unit\": \"W\"\n\t},\n", powerGPU, powerPkg)
sb.WriteString("\t\"engines\": {")
names := make([]string, 0, len(engines))
for name := range engines {
names = append(names, name)
}
slices.Sort(names)
for i, name := range names {
if i > 0 {
sb.WriteString(",")
}
fmt.Fprintf(&sb, "\n\t\t%q: {\n\t\t\t\"busy\": %f,\n\t\t\t\"sema\": 0.000000,\n\t\t\t\"wait\": 0.000000,\n\t\t\t\"unit\": \"%%\"\n\t\t}", name, engines[name])
}
sb.WriteString("\n\t}\n}")
return sb.String()
}
// intelJSONStream joins samples as intel_gpu_top -J prints them. Since v1.28
// the output starts with "[" (withArray); older versions omit it.
func intelJSONStream(withArray bool, samples ...string) string {
var sb strings.Builder
if withArray {
sb.WriteString("[\n")
}
for i, s := range samples {
if i > 0 {
sb.WriteString(",\n")
}
sb.WriteString(s)
}
return sb.String()
}
func TestIntelCollectorStreaming(t *testing.T) {
dir := t.TempDir()
t.Setenv("PATH", dir)
engines := func(render, blitter, video float64) map[string]float64 {
return map[string]float64{"Render/3D": render, "Blitter": blitter, "Video": video}
// Create a fake intel_gpu_top that prints -l format with four samples (first will be skipped) and exits
scriptPath := filepath.Join(dir, "intel_gpu_top")
script := `#!/bin/sh
echo "Freq MHz IRQ RC6 Power W IMC MiB/s RCS BCS VCS"
echo " req act /s % gpu pkg rd wr % se wa % se wa % se wa"
echo "373 373 224 45 1.50 4.13 2554 714 12.34 0 0 0.00 0 0 5.00 0 0"
echo "226 223 338 58 2.00 2.69 1820 965 0.00 0 0 0.00 0 0 0.00 0 0"
echo "189 187 412 67 1.80 2.45 1950 823 8.50 2 1 15.00 1 0 22.00 0 1"
echo "298 295 278 51 2.20 3.12 1675 942 5.75 1 2 9.50 3 1 12.00 1 0"`
if err := os.WriteFile(scriptPath, []byte(script), 0755); err != nil {
t.Fatal(err)
}
output := intelJSONStream(true,
intelJSONSample(1.5, 4.13, engines(12.34, 0, 5)),
intelJSONSample(2.0, 2.69, engines(0, 0, 0)),
intelJSONSample(1.8, 2.45, engines(8.5, 15, 22)),
intelJSONSample(2.2, 3.12, engines(5.75, 9.5, 12)),
) + "\n]"
// Create a fake intel_gpu_top that prints -J output with four samples (first will be skipped) and exits
gpuCommandFixture(t, dir, intelGpuStatsCmd, output+"\n")
gm := &GPUManager{
GpuDataMap: make(map[string]*system.GPUData),
@@ -1759,168 +1745,229 @@ func TestIntelCollectorStreaming(t *testing.T) {
gpu := gm.GpuDataMap["i0"]
require.NotNil(t, gpu)
// Power should be sum of samples 2-4 (first is skipped): 2.0 + 1.8 + 2.2 = 6.0
assert.InDelta(t, 6.0, gpu.Power, 0.001)
assert.EqualValues(t, 6.0, gpu.Power)
assert.InDelta(t, 8.26, gpu.PowerPkg, 0.01) // Allow small floating point differences
// Engines aggregated from samples 2-4
assert.InDelta(t, 14.25, gpu.Engines["Render/3D"], 0.001) // 0.00 + 8.50 + 5.75
assert.InDelta(t, 34.0, gpu.Engines["Video"], 0.001) // 0.00 + 22.00 + 12.00
assert.InDelta(t, 24.5, gpu.Engines["Blitter"], 0.001) // 0.00 + 15.00 + 9.50
assert.EqualValues(t, 14.25, gpu.Engines["Render/3D"]) // 0.00 + 8.50 + 5.75
assert.EqualValues(t, 34.0, gpu.Engines["Video"]) // 0.00 + 22.00 + 12.00
assert.EqualValues(t, 24.5, gpu.Engines["Blitter"]) // 0.00 + 15.00 + 9.50
// Count should be 3 samples (first is skipped)
assert.Equal(t, float64(3), gpu.Count)
}
func TestParseIntelJSONStream(t *testing.T) {
first := intelJSONSample(9, 9, map[string]float64{"Render/3D": 99, "Compute": 99})
classView := []string{
intelJSONSample(2, 3, map[string]float64{"Render/3D": 10, "Blitter": 1, "Video": 5, "VideoEnhance": 0, "Compute": 40}),
intelJSONSample(1, 2, map[string]float64{"Render/3D": 20, "Blitter": 0, "Video": 5, "VideoEnhance": 3, "Compute": 60}),
}
classViewWant := map[string]float64{"Render/3D": 30, "Blitter": 1, "Video": 10, "VideoEnhance": 3, "Compute": 100}
func TestParseIntelHeaders(t *testing.T) {
tests := []struct {
name string
input string
wantErr error
wantAnyErr bool
wantCount float64
wantPower float64
wantPkg float64
wantEngines map[string]float64
name string
header1 string
header2 string
wantEngineNames []string
wantFriendlyNames []string
wantPowerIndex int
wantPreEngineCols int
}{
{
name: "array still open while process runs",
input: intelJSONStream(true, first, classView[0], classView[1]),
wantCount: 2,
wantPower: 3,
wantPkg: 5,
wantEngines: classViewWant,
name: "basic headers with RCS BCS VCS",
header1: "Freq MHz IRQ RC6 Power W IMC MiB/s RCS BCS VCS",
header2: " req act /s % gpu pkg rd wr % se wa % se wa % se wa",
wantEngineNames: []string{"RCS", "BCS", "VCS"},
wantFriendlyNames: []string{"Render/3D", "Blitter", "Video"},
wantPowerIndex: 4, // "gpu" is at index 4
wantPreEngineCols: 8, // 17 total cols - 3*3 = 8
},
{
name: "closed array",
input: intelJSONStream(true, first, classView[0], classView[1]) + "\n]\n",
wantCount: 2,
wantPower: 3,
wantPkg: 5,
wantEngines: classViewWant,
name: "basic headers with RCS BCS VCS using index in name",
header1: "Freq MHz IRQ RC6 Power W IMC MiB/s RCS/0 BCS/1 VCS/2",
header2: " req act /s % gpu pkg rd wr % se wa % se wa % se wa",
wantEngineNames: []string{"RCS", "BCS", "VCS"},
wantFriendlyNames: []string{"Render/3D", "Blitter", "Video"},
wantPowerIndex: 4, // "gpu" is at index 4
wantPreEngineCols: 8, // 17 total cols - 3*3 = 8
},
{
name: "truncated final sample",
input: intelJSONStream(true, first, classView[0], classView[1], `{"period": {"duration": 33`),
wantCount: 2,
wantPower: 3,
wantPkg: 5,
wantEngines: classViewWant,
name: "headers with only RCS",
header1: "Freq MHz IRQ RC6 Power W IMC MiB/s RCS",
header2: " req act /s % gpu pkg rd wr % se wa",
wantEngineNames: []string{"RCS"},
wantFriendlyNames: []string{"Render/3D"},
wantPowerIndex: 4,
wantPreEngineCols: 8, // 11 total - 3*1 = 8
},
{
// intel_gpu_top < 1.28 omits the opening "[" and uses physical engine names
name: "legacy output without array and with engine instances",
input: intelJSONStream(false,
intelJSONSample(9, 9, map[string]float64{"Render/3D/0": 99}),
intelJSONSample(1.5, 2.5, map[string]float64{"Render/3D/0": 12, "Blitter/0": 1, "Video/0": 4, "Video/1": 6, "VideoEnhance/0": 2}),
),
wantCount: 1,
wantPower: 1.5,
wantPkg: 2.5,
wantEngines: map[string]float64{"Render/3D": 12, "Blitter": 1, "Video": 10, "VideoEnhance": 2},
name: "headers with VECS and CCS",
header1: "Freq MHz IRQ RC6 Power W IMC MiB/s VECS CCS",
header2: " req act /s % gpu pkg rd wr % se wa % se wa",
wantEngineNames: []string{"VECS", "CCS"},
wantFriendlyNames: []string{"VideoEnhance", "Compute"},
wantPowerIndex: 4,
wantPreEngineCols: 8, // 14 total - 3*2 = 8
},
{
// energy counter read lower than the previous sample in intel_gpu_top
name: "sample with invalid power is skipped",
input: intelJSONStream(true, first, classView[0],
intelJSONSample(86_000_000, 3, map[string]float64{"Render/3D": 50}),
intelJSONSample(2, 90_000_000, map[string]float64{"Render/3D": 50}),
classView[1],
),
wantCount: 2,
wantPower: 3,
wantPkg: 5,
wantEngines: classViewWant,
name: "no engines",
header1: "Freq MHz IRQ RC6 Power W IMC MiB/s",
header2: " req act /s % gpu pkg rd wr",
wantEngineNames: nil, // no engines found, slices remain nil
wantFriendlyNames: nil,
wantPowerIndex: -1, // no engines, so no search
wantPreEngineCols: 0,
},
{
name: "only samples with invalid power",
input: intelJSONStream(true, first, intelJSONSample(86_000_000, 3, map[string]float64{"Render/3D": 50})),
wantErr: errNoValidData,
name: "power index not found",
header1: "Freq MHz IRQ RC6 Power W IMC MiB/s RCS",
header2: " req act /s % pkg cpu rd wr % se wa", // no "gpu"
wantEngineNames: []string{"RCS"},
wantFriendlyNames: []string{"Render/3D"},
wantPowerIndex: -1, // "gpu" not found
wantPreEngineCols: 8, // 11 total - 3*1 = 8
},
{
name: "empty output",
input: "",
wantErr: errNoValidData,
},
{
name: "only first sample, which is skipped",
input: intelJSONStream(true, first),
wantErr: errNoValidData,
},
{
name: "invalid output",
input: "intel_gpu_top: command failed",
wantAnyErr: true,
name: "empty headers",
header1: "",
header2: "",
wantEngineNames: nil, // empty input, slices remain nil
wantFriendlyNames: nil,
wantPowerIndex: -1,
wantPreEngineCols: 0,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
gm := &GPUManager{GpuDataMap: make(map[string]*system.GPUData)}
err := gm.parseIntelJSONStream(strings.NewReader(tt.input))
if tt.wantAnyErr {
assert.Error(t, err)
return
}
if tt.wantErr != nil {
assert.ErrorIs(t, err, tt.wantErr)
assert.Empty(t, gm.GpuDataMap)
return
}
require.NoError(t, err)
gm := &GPUManager{}
engineNames, friendlyNames, powerIndex, preEngineCols := gm.parseIntelHeaders(tt.header1, tt.header2)
gpu := gm.GpuDataMap["i0"]
require.NotNil(t, gpu)
assert.Equal(t, tt.wantCount, gpu.Count)
assert.InDelta(t, tt.wantPower, gpu.Power, 0.001)
assert.InDelta(t, tt.wantPkg, gpu.PowerPkg, 0.001)
assert.Len(t, gpu.Engines, len(tt.wantEngines))
for name, want := range tt.wantEngines {
assert.InDelta(t, want, gpu.Engines[name], 0.001, name)
}
assert.Equal(t, tt.wantEngineNames, engineNames)
assert.Equal(t, tt.wantFriendlyNames, friendlyNames)
assert.Equal(t, tt.wantPowerIndex, powerIndex)
assert.Equal(t, tt.wantPreEngineCols, preEngineCols)
})
}
}
func TestParseIntelJSONSample(t *testing.T) {
t.Run("without power", func(t *testing.T) {
var sample intelGpuJSONSample
require.NoError(t, json.Unmarshal([]byte(`{"engines": {"Render/3D": {"busy": 7.5, "unit": "%"}}}`), &sample))
stats := parseIntelJSONSample(sample)
assert.Zero(t, stats.PowerGPU)
assert.Zero(t, stats.PowerPkg)
assert.Equal(t, map[string]float64{"Render/3D": 7.5}, stats.Engines)
})
t.Run("without engines", func(t *testing.T) {
var sample intelGpuJSONSample
require.NoError(t, json.Unmarshal([]byte(`{"power": {"GPU": 1.25, "Package": 4.5, "unit": "W"}}`), &sample))
stats := parseIntelJSONSample(sample)
assert.Equal(t, 1.25, stats.PowerGPU)
assert.Equal(t, 4.5, stats.PowerPkg)
assert.Nil(t, stats.Engines)
})
}
func TestIntelEngineClass(t *testing.T) {
tests := map[string]string{
"Render/3D": "Render/3D",
"Render/3D/0": "Render/3D",
"Blitter": "Blitter",
"Blitter/0": "Blitter",
"Video/1": "Video",
"VideoEnhance/0": "VideoEnhance",
"Compute/3": "Compute",
"[unknown]": "[unknown]",
"[unknown]/0": "[unknown]",
"Video/": "Video/",
func TestParseIntelData(t *testing.T) {
tests := []struct {
name string
line string
engineNames []string
friendlyNames []string
powerIndex int
preEngineCols int
wantPowerGPU float64
wantEngines map[string]float64
wantErr error
}{
{
name: "basic data with power and engines",
line: "373 373 224 45 1.50 4.13 2554 714 12.34 0 0 0.00 0 0 5.00 0 0",
engineNames: []string{"RCS", "BCS", "VCS"},
friendlyNames: []string{"Render/3D", "Blitter", "Video"},
powerIndex: 4,
preEngineCols: 8,
wantPowerGPU: 1.50,
wantEngines: map[string]float64{
"Render/3D": 12.34,
"Blitter": 0.00,
"Video": 5.00,
},
},
{
name: "data with zero power",
line: "226 223 338 58 0.00 2.69 1820 965 0.00 0 0 0.00 0 0 0.00 0 0",
engineNames: []string{"RCS", "BCS", "VCS"},
friendlyNames: []string{"Render/3D", "Blitter", "Video"},
powerIndex: 4,
preEngineCols: 8,
wantPowerGPU: 0.00,
wantEngines: map[string]float64{
"Render/3D": 0.00,
"Blitter": 0.00,
"Video": 0.00,
},
},
{
name: "data with no power index",
line: "373 373 224 45 1.50 4.13 2554 714 12.34 0 0 0.00 0 0 5.00 0 0",
engineNames: []string{"RCS", "BCS", "VCS"},
friendlyNames: []string{"Render/3D", "Blitter", "Video"},
powerIndex: -1,
preEngineCols: 8,
wantPowerGPU: 0.0, // no power parsed
wantEngines: map[string]float64{
"Render/3D": 12.34,
"Blitter": 0.00,
"Video": 5.00,
},
},
{
name: "data with insufficient columns",
line: "373 373 224 45 1.50", // too few columns
engineNames: []string{"RCS", "BCS", "VCS"},
friendlyNames: []string{"Render/3D", "Blitter", "Video"},
powerIndex: 4,
preEngineCols: 8,
wantPowerGPU: 0.0,
wantEngines: nil, // empty sample returned
wantErr: errNoValidData,
},
{
name: "empty line",
line: "",
engineNames: []string{"RCS"},
friendlyNames: []string{"Render/3D"},
powerIndex: 4,
preEngineCols: 8,
wantPowerGPU: 0.0,
wantEngines: nil,
wantErr: errNoValidData,
},
{
name: "data with invalid power value",
line: "373 373 224 45 N/A 4.13 2554 714 12.34 0 0 0.00 0 0 5.00 0 0",
engineNames: []string{"RCS", "BCS", "VCS"},
friendlyNames: []string{"Render/3D", "Blitter", "Video"},
powerIndex: 4,
preEngineCols: 8,
wantPowerGPU: 0.0, // N/A can't be parsed
wantEngines: map[string]float64{
"Render/3D": 12.34,
"Blitter": 0.00,
"Video": 5.00,
},
},
{
name: "data with invalid engine value",
line: "373 373 224 45 1.50 4.13 2554 714 N/A 0 0 0.00 0 0 5.00 0 0",
engineNames: []string{"RCS", "BCS", "VCS"},
friendlyNames: []string{"Render/3D", "Blitter", "Video"},
powerIndex: 4,
preEngineCols: 8,
wantPowerGPU: 1.50,
wantEngines: map[string]float64{
"Render/3D": 0.0, // N/A becomes 0
"Blitter": 0.00,
"Video": 5.00,
},
},
{
name: "data with no engines",
line: "373 373 224 45 1.50 4.13 2554 714",
engineNames: []string{},
friendlyNames: []string{},
powerIndex: 4,
preEngineCols: 8,
wantPowerGPU: 1.50,
wantEngines: nil,
},
}
for key, want := range tests {
assert.Equal(t, want, intelEngineClass(key), key)
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
gm := &GPUManager{}
sample, err := gm.parseIntelData(tt.line, tt.engineNames, tt.friendlyNames, tt.powerIndex, tt.preEngineCols)
assert.Equal(t, tt.wantErr, err)
assert.Equal(t, tt.wantPowerGPU, sample.PowerGPU)
assert.Equal(t, tt.wantEngines, sample.Engines)
})
}
}
@@ -1928,12 +1975,21 @@ func TestIntelCollectorDeviceEnv(t *testing.T) {
dir := t.TempDir()
t.Setenv("PATH", dir)
// Prepare a file to capture args
argsFile := filepath.Join(dir, "args.txt")
// Create a fake intel_gpu_top that records its arguments and prints minimal valid output
output := intelJSONStream(true,
intelJSONSample(2, 2.69, map[string]float64{"Render/3D": 0, "Video": 0}),
intelJSONSample(1.8, 2.45, map[string]float64{"Render/3D": 8.5, "Video": 15}),
)
argsFile := gpuCommandFixture(t, dir, intelGpuStatsCmd, output)
scriptPath := filepath.Join(dir, "intel_gpu_top")
script := fmt.Sprintf(`#!/bin/sh
echo "$@" > %s
echo "Freq MHz IRQ RC6 Power W IMC MiB/s RCS VCS"
echo " req act /s %% gpu pkg rd wr %% se wa %% se wa"
echo "226 223 338 58 2.00 2.69 1820 965 0.00 0 0 0.00 0 0"
echo "189 187 412 67 1.80 2.45 1950 823 8.50 2 1 15.00 1 0"
`, argsFile)
if err := os.WriteFile(scriptPath, []byte(script), 0755); err != nil {
t.Fatal(err)
}
// Set device selector via prefixed env var
t.Setenv("BESZEL_AGENT_INTEL_GPU_DEVICE", "sriov")
@@ -1951,5 +2007,5 @@ func TestIntelCollectorDeviceEnv(t *testing.T) {
argsStr := strings.TrimSpace(string(data))
require.Contains(t, argsStr, "-d sriov")
require.Contains(t, argsStr, "-s ")
require.Contains(t, argsStr, "-J")
require.Contains(t, argsStr, "-l")
}

View File

@@ -7,23 +7,19 @@ import (
"github.com/fxamacker/cbor/v2"
"github.com/henrygd/beszel/internal/common"
"github.com/henrygd/beszel/internal/entities/monitor"
"github.com/henrygd/beszel/internal/entities/probe"
"github.com/henrygd/beszel/internal/entities/smart"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/lxzan/gws"
"log/slog"
)
// HandlerContext provides context for request handlers
type HandlerContext struct {
Client *WebSocketClient
Conn *gws.Conn // WebSocket that carried this request, if any
Agent *Agent
Request *common.HubRequest[cbor.RawMessage]
RequestID *uint32
HubVerified bool
ConnectionType system.ConnectionType // Transport that carried this request
Client *WebSocketClient
Agent *Agent
Request *common.HubRequest[cbor.RawMessage]
RequestID *uint32
HubVerified bool
// SendResponse abstracts how a handler sends responses (WS or SSH)
SendResponse func(data any, requestID *uint32) error
}
@@ -56,10 +52,7 @@ func NewHandlerRegistry() *HandlerRegistry {
registry.Register(common.GetContainerInfo, &GetContainerInfoHandler{})
registry.Register(common.GetSmartData, &GetSmartDataHandler{})
registry.Register(common.GetSystemdInfo, &GetSystemdInfoHandler{})
registry.Register(common.GetSystemdLogs, &GetSystemdLogsHandler{})
registry.Register(common.SyncNetworkMonitors, &SyncNetworkMonitorsHandler{})
registry.Register(common.GetZfsData, &GetZfsDataHandler{})
registry.Register(common.GetPackageUpdates, &GetPackageUpdatesHandler{})
registry.Register(common.SyncNetworkProbes, &SyncNetworkProbesHandler{})
return registry
}
@@ -104,11 +97,7 @@ func (h *GetDataHandler) Handle(hctx *HandlerContext) error {
_ = cbor.Unmarshal(hctx.Request.Data, &options)
sysStats := hctx.Agent.gatherStats(options)
// Cached stats may be shared by concurrent SSH and WebSocket requests.
// Set the transport on the response copy, not on the cached data.
response := *sysStats
response.Info.ConnectionType = hctx.ConnectionType
return hctx.SendResponse(&response, hctx.RequestID)
return hctx.SendResponse(sysStats, hctx.RequestID)
}
////////////////////////////////////////////////////////////////////////////
@@ -118,7 +107,7 @@ func (h *GetDataHandler) Handle(hctx *HandlerContext) error {
type CheckFingerprintHandler struct{}
func (h *CheckFingerprintHandler) Handle(hctx *HandlerContext) error {
return hctx.Client.handleAuthChallenge(hctx.Request, hctx.RequestID, hctx.Conn)
return hctx.Client.handleAuthChallenge(hctx.Request, hctx.RequestID)
}
////////////////////////////////////////////////////////////////////////////
@@ -179,47 +168,14 @@ type GetSmartDataHandler struct{}
func (h *GetSmartDataHandler) Handle(hctx *HandlerContext) error {
if hctx.Agent.smartManager == nil {
return hctx.SendResponse(smart.SmartDataResponse{Data: map[string]smart.SmartData{}}, hctx.RequestID)
// return empty map to indicate no data
return hctx.SendResponse(map[string]smart.SmartData{}, hctx.RequestID)
}
complete, err := hctx.Agent.smartManager.Refresh(false)
if err != nil {
if err := hctx.Agent.smartManager.Refresh(false); err != nil {
slog.Debug("smart refresh failed", "err", err)
}
return hctx.SendResponse(smart.SmartDataResponse{
Data: hctx.Agent.smartManager.GetCurrentData(),
Complete: complete,
}, hctx.RequestID)
}
////////////////////////////////////////////////////////////////////////////
////////////////////////////////////////////////////////////////////////////
// GetZfsDataHandler handles ZFS detail data requests
type GetZfsDataHandler struct{}
func (h *GetZfsDataHandler) Handle(hctx *HandlerContext) error {
if hctx.Agent.storagePoolManager == nil {
return hctx.SendResponse(nil, hctx.RequestID)
}
var req common.ZfsDataRequest
if err := cbor.Unmarshal(hctx.Request.Data, &req); err != nil {
return err
}
return hctx.SendResponse(hctx.Agent.storagePoolManager.GetDetail(req.Force), hctx.RequestID)
}
////////////////////////////////////////////////////////////////////////////
////////////////////////////////////////////////////////////////////////////
// GetPackageUpdatesHandler returns the pending package updates found by the
// last background check. It never runs a check itself.
type GetPackageUpdatesHandler struct{}
func (h *GetPackageUpdatesHandler) Handle(hctx *HandlerContext) error {
if hctx.Agent.packageUpdates == nil {
return hctx.SendResponse(system.PackageUpdates{}, hctx.RequestID)
}
return hctx.SendResponse(hctx.Agent.packageUpdates.list(), hctx.RequestID)
data := hctx.Agent.smartManager.GetCurrentData()
return hctx.SendResponse(data, hctx.RequestID)
}
////////////////////////////////////////////////////////////////////////////
@@ -253,42 +209,15 @@ func (h *GetSystemdInfoHandler) Handle(hctx *HandlerContext) error {
////////////////////////////////////////////////////////////////////////////
////////////////////////////////////////////////////////////////////////////
// GetSystemdLogsHandler handles recent systemd service log requests.
type GetSystemdLogsHandler struct{}
// SyncNetworkProbesHandler handles probe configuration sync from hub
type SyncNetworkProbesHandler struct{}
func (h *GetSystemdLogsHandler) Handle(hctx *HandlerContext) error {
if hctx.Agent.systemdManager == nil {
return errors.ErrUnsupported
}
var req common.SystemdLogsRequest
func (h *SyncNetworkProbesHandler) Handle(hctx *HandlerContext) error {
var req probe.SyncRequest
if err := cbor.Unmarshal(hctx.Request.Data, &req); err != nil {
return err
}
if req.ServiceName == "" {
return errors.New("service name is required")
}
logs, err := hctx.Agent.systemdManager.getServiceLogs(req.ServiceName)
if err != nil {
return err
}
return hctx.SendResponse(logs, hctx.RequestID)
}
////////////////////////////////////////////////////////////////////////////
////////////////////////////////////////////////////////////////////////////
// SyncNetworkMonitorsHandler handles monitor configuration sync from hub
type SyncNetworkMonitorsHandler struct{}
func (h *SyncNetworkMonitorsHandler) Handle(hctx *HandlerContext) error {
var req monitor.SyncRequest
if err := cbor.Unmarshal(hctx.Request.Data, &req); err != nil {
return err
}
resp, err := hctx.Agent.monitorManager.HandleSyncRequest(req)
resp, err := hctx.Agent.probeManager.HandleSyncRequest(req)
if err != nil {
return err
}

View File

@@ -4,13 +4,9 @@ package agent
import (
"testing"
"time"
"github.com/fxamacker/cbor/v2"
"github.com/henrygd/beszel/agent/zfs"
"github.com/henrygd/beszel/internal/common"
"github.com/henrygd/beszel/internal/entities/smart"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/stretchr/testify/assert"
)
@@ -21,68 +17,6 @@ type MockHandler struct {
handleFunc func(ctx *HandlerContext) error
}
func TestNewAgentResponseSmartData(t *testing.T) {
response := newAgentResponse(smart.SmartDataResponse{
Data: map[string]smart.SmartData{
"AAA": {SerialNumber: "AAA"},
},
Complete: true,
}, nil)
assert.Equal(t, "AAA", response.SmartData["AAA"].SerialNumber)
assert.True(t, response.SmartComplete)
}
func TestGetDataHandlerReportsRequestTransport(t *testing.T) {
cache := NewSystemDataCache()
cached := &system.CombinedData{}
cache.Set(cached, defaultDataCacheTimeMs)
agent := &Agent{cache: cache}
options, err := cbor.Marshal(common.DataRequestOptions{CacheTimeMs: defaultDataCacheTimeMs})
assert.NoError(t, err)
request := &common.HubRequest[cbor.RawMessage]{Action: common.GetData, Data: options}
for _, transport := range []system.ConnectionType{system.ConnectionTypeSSH, system.ConnectionTypeWebSocket} {
ctx := &HandlerContext{
Agent: agent,
Request: request,
ConnectionType: transport,
SendResponse: func(data any, _ *uint32) error {
response := data.(*system.CombinedData)
assert.Equal(t, transport, response.Info.ConnectionType)
return nil
},
}
assert.NoError(t, (&GetDataHandler{}).Handle(ctx))
assert.Equal(t, system.ConnectionTypeNone, cached.Info.ConnectionType)
}
}
func TestGetZfsDataHandlerForceRefresh(t *testing.T) {
poolCalls := 0
zm := &StoragePoolManager{detailInterval: time.Hour, backends: []*poolBackend{{name: "zfs"}}}
zm.backends[0].poolStatsFn = func() ([]zfs.PoolStat, error) {
poolCalls++
return []zfs.PoolStat{{Name: "tank", Alloc: uint64(poolCalls)}}, nil
}
zm.backends[0].poolStatusesFn = func() ([]zfs.PoolStatus, error) { return nil, nil }
zm.backends[0].datasetsFn = func() ([]zfs.Dataset, error) { return nil, nil }
zm.GetDetail(false)
requestData, err := cbor.Marshal(common.ZfsDataRequest{Force: true})
assert.NoError(t, err)
ctx := &HandlerContext{
Agent: &Agent{storagePoolManager: zm},
Request: &common.HubRequest[cbor.RawMessage]{
Action: common.GetZfsData,
Data: requestData,
},
SendResponse: func(any, *uint32) error { return nil },
}
assert.NoError(t, (&GetZfsDataHandler{}).Handle(ctx))
assert.Equal(t, 2, poolCalls)
}
func (m *MockHandler) Handle(ctx *HandlerContext) error {
if m.handleFunc != nil {
return m.handleFunc(ctx)

View File

@@ -17,17 +17,15 @@ import (
var mdraidSysfsRoot = "/sys"
type mdraidHealth struct {
level string
arrayState string
degraded uint64
faultyDisks uint64
populatedDisks uint64
raidDisks uint64
syncAction string
syncCompleted string
syncSpeed string
mismatchCnt uint64
capacity uint64
level string
arrayState string
degraded uint64
raidDisks uint64
syncAction string
syncCompleted string
syncSpeed string
mismatchCnt uint64
capacity uint64
}
// scanMdraidDevices discovers Linux md arrays exposed in sysfs.
@@ -94,9 +92,6 @@ func (sm *SmartManager) collectMdraidHealth(deviceInfo *DeviceInfo) (bool, error
if health.degraded > 0 {
attrs = append(attrs, &smart.SmartAttribute{Name: "Degraded", RawValue: health.degraded})
}
if health.faultyDisks > 0 {
attrs = append(attrs, &smart.SmartAttribute{Name: "FaultyDisks", RawValue: health.faultyDisks})
}
if health.syncAction != "" {
attrs = append(attrs, &smart.SmartAttribute{Name: "SyncAction", RawString: health.syncAction})
}
@@ -157,7 +152,6 @@ func readMdraidHealth(blockName string) (mdraidHealth, bool) {
if val, ok := utils.ReadUintFile(filepath.Join(mdDir, "degraded")); ok {
out.degraded = val
}
out.faultyDisks, out.populatedDisks = countMdraidMemberStates(blockName, mdraidSysfsRoot)
if val, ok := utils.ReadUintFile(filepath.Join(mdDir, "mismatch_cnt")); ok {
out.mismatchCnt = val
}
@@ -183,27 +177,13 @@ func mdraidSmartStatus(health mdraidHealth) string {
case "resync", "recover", "reshape":
return "WARNING"
}
// Use actual faulty member count rather than the degraded counter, which
// equals raid_disks minus active_disks. On QNAP systems raid_disks may be
// set to a large value (e.g. 32) while only a few slots are ever used,
// making degraded misleadingly large despite zero failed disks.
if health.faultyDisks > 0 {
return "FAILED"
}
if health.degraded > 0 {
if isSparseSlotDegraded(health) {
// A sysfs snapshot cannot distinguish reserved slots from a removed
// member on sparse arrays, so report the ambiguity as a warning.
return "WARNING"
}
return "FAILED"
}
if health.mismatchCnt > 0 {
switch syncAction {
case "check", "repair":
return "WARNING"
}
// "check" and "repair" are requested consistency scans, not evidence of
// array failure. With no health issues above, keep scrubbing green while
// reporting the sync action and progress attributes.
switch state {
case "clean", "active", "active-idle", "write-pending", "read-auto", "readonly":
return "PASSED"
@@ -211,43 +191,6 @@ func mdraidSmartStatus(health mdraidHealth) string {
return "UNKNOWN"
}
// countMdraidMemberStates reads member device directories under
// block/<name>/md and returns how many are explicitly marked "faulty", plus
// how many are populated at all (regardless of state). populatedDisks lets
// callers distinguish RAID slots that were never used (QNAP reserves far
// more raid_disks than it ever populates) from members that went missing.
func countMdraidMemberStates(blockName, root string) (faultyDisks, populatedDisks uint64) {
devDir := filepath.Join(root, "block", blockName, "md")
entries, err := os.ReadDir(devDir)
if err != nil {
return 0, 0
}
for _, ent := range entries {
if !strings.HasPrefix(ent.Name(), "dev-") {
continue
}
populatedDisks++
statePath := filepath.Join(devDir, ent.Name(), "state")
state := utils.ReadStringFile(statePath)
if strings.Contains(state, "faulty") {
faultyDisks++
}
}
return faultyDisks, populatedDisks
}
// isSparseSlotDegraded reports whether a non-zero "degraded" count may be
// explained by RAID slots that were never populated. QNAP configures system
// arrays with raid_disks set to a large fixed maximum (e.g. 32) far beyond the
// handful of slots it ever populates, so sparse slots outnumber populated ones.
func isSparseSlotDegraded(health mdraidHealth) bool {
if health.populatedDisks == 0 || health.raidDisks <= health.populatedDisks {
return false
}
sparseSlots := health.raidDisks - health.populatedDisks
return sparseSlots > health.populatedDisks
}
// isMdraidBlockName matches /dev/mdN-style block device names.
func isMdraidBlockName(name string) bool {
if !strings.HasPrefix(name, "md") {

View File

@@ -40,15 +40,6 @@ func TestMdraidMockSysfsScanAndCollect(t *testing.T) {
write(filepath.Join(mdDir, "sync_completed"), "10%\n")
write(filepath.Join(mdDir, "sync_speed"), "100M\n")
write(filepath.Join(mdDir, "mismatch_cnt"), "0\n")
// Simulate two healthy member devices (no faulty state).
for _, dev := range []string{"dev-sda", "dev-sdb"} {
devPath := filepath.Join(mdDir, dev)
if err := os.MkdirAll(devPath, 0o755); err != nil {
t.Fatal(err)
}
write(filepath.Join(devPath, "state"), "in_sync\n")
}
write(filepath.Join(queueDir, "logical_block_size"), "512\n")
write(filepath.Join(tmp, "block", "md0", "size"), "2048\n")
@@ -90,110 +81,19 @@ func TestMdraidMockSysfsScanAndCollect(t *testing.T) {
}
}
func TestCountMdraidMemberStates(t *testing.T) {
tmp := t.TempDir()
write := func(path, content string) {
t.Helper()
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
mdDir := filepath.Join(tmp, "block", "md0", "md")
// No dev-* entries: zero faulty, zero populated.
if faulty, populated := countMdraidMemberStates("md0", tmp); faulty != 0 || populated != 0 {
t.Fatalf("no members: got (faulty=%d populated=%d), want (0,0)", faulty, populated)
}
// Two healthy members.
write(filepath.Join(mdDir, "dev-sda", "state"), "in_sync\n")
write(filepath.Join(mdDir, "dev-sdb", "state"), "in_sync\n")
if faulty, populated := countMdraidMemberStates("md0", tmp); faulty != 0 || populated != 2 {
t.Fatalf("all in_sync: got (faulty=%d populated=%d), want (0,2)", faulty, populated)
}
// One faulty member.
write(filepath.Join(mdDir, "dev-sdb", "state"), "faulty\n")
if faulty, populated := countMdraidMemberStates("md0", tmp); faulty != 1 || populated != 2 {
t.Fatalf("one faulty: got (faulty=%d populated=%d), want (1,2)", faulty, populated)
}
// QNAP-style: 28 degraded slots but no dev-* entries for them, 4 in_sync.
write(filepath.Join(mdDir, "dev-sdb", "state"), "in_sync\n")
write(filepath.Join(mdDir, "dev-sdc", "state"), "in_sync\n")
write(filepath.Join(mdDir, "dev-sdd", "state"), "in_sync\n")
if faulty, populated := countMdraidMemberStates("md0", tmp); faulty != 0 || populated != 4 {
t.Fatalf("qnap sparse: got (faulty=%d populated=%d), want (0,4)", faulty, populated)
}
}
func TestMdraidSmartStatus(t *testing.T) {
if got := mdraidSmartStatus(mdraidHealth{arrayState: "inactive"}); got != "FAILED" {
t.Fatalf("mdraidSmartStatus(inactive) = %q, want FAILED", got)
}
if got := mdraidSmartStatus(mdraidHealth{arrayState: "active", degraded: 1, faultyDisks: 1, syncAction: "recover"}); got != "WARNING" {
if got := mdraidSmartStatus(mdraidHealth{arrayState: "active", degraded: 1, syncAction: "recover"}); got != "WARNING" {
t.Fatalf("mdraidSmartStatus(degraded+recover) = %q, want WARNING", got)
}
if got := mdraidSmartStatus(mdraidHealth{arrayState: "active", degraded: 1, faultyDisks: 1}); got != "FAILED" {
t.Fatalf("mdraidSmartStatus(degraded+faulty) = %q, want FAILED", got)
}
// QNAP-style: raid_disks=32 but only 4 populated; degraded=28 but no faulty devices.
if got := mdraidSmartStatus(mdraidHealth{arrayState: "clean", degraded: 28, faultyDisks: 0, raidDisks: 32, populatedDisks: 4}); got != "WARNING" {
t.Fatalf("mdraidSmartStatus(qnap sparse) = %q, want WARNING", got)
}
// A member disappearing from the same sparse array is indistinguishable
// from another reserved slot, so it must not be reported as healthy.
if got := mdraidSmartStatus(mdraidHealth{arrayState: "clean", degraded: 29, faultyDisks: 0, raidDisks: 32, populatedDisks: 3}); got != "WARNING" {
t.Fatalf("mdraidSmartStatus(qnap sparse missing member) = %q, want WARNING", got)
}
// A genuinely missing member (removed dev-* entry, not just an unpopulated
// QNAP reserve slot) must still fail: raid_disks=4, only 3 populated, all
// of them in_sync, so faultyDisks==0 but degraded==1.
if got := mdraidSmartStatus(mdraidHealth{arrayState: "clean", degraded: 1, faultyDisks: 0, raidDisks: 4, populatedDisks: 3}); got != "FAILED" {
t.Fatalf("mdraidSmartStatus(missing member) = %q, want FAILED", got)
}
// Degraded with no member-state info at all (e.g. sysfs read failed) must
// still fail rather than being silently treated as a sparse QNAP array.
if got := mdraidSmartStatus(mdraidHealth{arrayState: "clean", degraded: 1, faultyDisks: 0, raidDisks: 4, populatedDisks: 0}); got != "FAILED" {
t.Fatalf("mdraidSmartStatus(degraded, no member info) = %q, want FAILED", got)
if got := mdraidSmartStatus(mdraidHealth{arrayState: "active", degraded: 1}); got != "FAILED" {
t.Fatalf("mdraidSmartStatus(degraded) = %q, want FAILED", got)
}
if got := mdraidSmartStatus(mdraidHealth{arrayState: "active", syncAction: "recover"}); got != "WARNING" {
t.Fatalf("mdraidSmartStatus(recover) = %q, want WARNING", got)
}
if got := mdraidSmartStatus(mdraidHealth{arrayState: "clean", syncAction: "check"}); got != "PASSED" {
t.Fatalf("mdraidSmartStatus(clean+check) = %q, want PASSED", got)
}
if got := mdraidSmartStatus(mdraidHealth{arrayState: "clean", syncAction: "check", mismatchCnt: 1}); got != "WARNING" {
t.Fatalf("mdraidSmartStatus(clean+check+mismatch) = %q, want WARNING", got)
}
if got := mdraidSmartStatus(mdraidHealth{arrayState: "clean", mismatchCnt: 1}); got != "WARNING" {
t.Fatalf("mdraidSmartStatus(clean+mismatch) = %q, want WARNING", got)
}
for _, tc := range []struct {
name string
health mdraidHealth
want string
}{
{"clean", mdraidHealth{arrayState: "clean"}, "PASSED"},
{"active", mdraidHealth{arrayState: "active"}, "PASSED"},
{"mismatch", mdraidHealth{arrayState: "active", mismatchCnt: 1}, "WARNING"},
{"degraded", mdraidHealth{arrayState: "active", degraded: 1}, "FAILED"},
{"faulty member", mdraidHealth{arrayState: "active", faultyDisks: 1}, "FAILED"},
{"inactive", mdraidHealth{arrayState: "inactive"}, "FAILED"},
{"unknown", mdraidHealth{arrayState: "unknown"}, "UNKNOWN"},
} {
t.Run("repair/"+tc.name, func(t *testing.T) {
tc.health.syncAction = "repair"
if got := mdraidSmartStatus(tc.health); got != tc.want {
t.Fatalf("mdraidSmartStatus(%+v) = %q, want %s", tc.health, got, tc.want)
}
})
}
if got := mdraidSmartStatus(mdraidHealth{arrayState: "clean"}); got != "PASSED" {
t.Fatalf("mdraidSmartStatus(clean) = %q, want PASSED", got)
}

View File

@@ -92,7 +92,7 @@ func (a *Agent) updateNetworkStats(cacheTimeMs uint16, systemStats *system.Stats
func (a *Agent) initializeNetIoStats() {
// reset valid network interfaces
a.netInterfaces = make(map[string]bool, 0)
a.netInterfaces = make(map[string]struct{}, 0)
// parse NICS env var for whitelist / blacklist
nicsEnvVal, nicsEnvExists := utils.GetEnv("NICS")
@@ -107,14 +107,9 @@ func (a *Agent) initializeNetIoStats() {
if skipNetworkInterface(v, nicCfg) {
continue
}
// driver is checked only here so updates don't pay for it on non-Jetson systems
useMacCounters := isNvidiaEthernet(v.Name)
if useMacCounters {
correctNvethernetCounters(&v)
}
slog.Info("Detected network interface", "name", v.Name, "sent", v.BytesSent, "recv", v.BytesRecv)
// store as a valid network interface
a.netInterfaces[v.Name] = useMacCounters
a.netInterfaces[v.Name] = struct{}{}
}
}
@@ -164,13 +159,9 @@ func (a *Agent) sumAndTrackPerNicDeltas(cacheTimeMs uint16, msElapsed uint64, ne
tracker.Cycle()
for _, v := range netIO {
useMacCounters, exists := a.netInterfaces[v.Name]
if !exists {
if _, exists := a.netInterfaces[v.Name]; !exists {
continue
}
if useMacCounters {
correctNvethernetCounters(&v)
}
totalBytesSent += v.BytesSent
totalBytesRecv += v.BytesRecv

View File

@@ -1,71 +0,0 @@
//go:build linux
package agent
import (
"log/slog"
"math"
"os"
"path/filepath"
"strings"
"github.com/safchain/ethtool"
psutilNet "github.com/shirou/gopsutil/v4/net"
)
// correctNvethernetCounters replaces the inflated sysfs byte counters of an
// nvethernet NIC with its MAC octet counters.
func correctNvethernetCounters(v *psutilNet.IOCountersStat) {
tx, rx, ok := readEthtoolMACOctets(v.Name)
if !ok {
return
}
v.BytesSent = tx
v.BytesRecv = rx
}
func isNvidiaEthernet(name string) bool {
if name == "" || strings.Contains(name, "/") {
return false
}
driverPath := filepath.Join("/sys/class/net", name, "device/driver")
target, err := os.Readlink(driverPath)
if err != nil {
return false
}
return filepath.Base(target) == "nvethernet"
}
func readEthtoolMACOctets(name string) (tx, rx uint64, ok bool) {
stats, err := ethtool.Stats(name)
if err != nil {
slog.Debug("Failed to read ethtool network counters", "interface", name, "err", err)
return 0, 0, false
}
tx, okTx := ethtoolCounter(stats, "mmc_tx_octetcount_gb", "mmc_tx_octetcount_gb_h")
rx, okRx := ethtoolCounter(stats, "mmc_rx_octetcount_gb", "mmc_rx_octetcount_gb_h")
if !okTx || !okRx {
return 0, 0, false
}
return tx, rx, true
}
// ethtoolCounter combines nvethernet's split MMC counters. The driver accumulates
// the low and high registers into independent 64-bit fields, so the low value can
// exceed 32 bits and must be added rather than OR'd into the shifted high word.
func ethtoolCounter(stats map[string]uint64, lowKey, highKey string) (uint64, bool) {
low, ok := stats[lowKey]
if !ok {
return 0, false
}
high, hasHigh := stats[highKey]
if !hasHigh {
return low, true
}
if high > (math.MaxUint64-low)>>32 {
return 0, false
}
return high<<32 + low, true
}

View File

@@ -1,53 +0,0 @@
//go:build linux && testing
package agent
import (
"math"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestEthtoolCounterCombinesHighWord(t *testing.T) {
stats := map[string]uint64{
"counter": 123,
"counter_h": 2,
}
value, ok := ethtoolCounter(stats, "counter", "counter_h")
require.True(t, ok)
assert.Equal(t, uint64(2<<32+123), value)
}
func TestEthtoolCounterAddsLowWordAbove32Bits(t *testing.T) {
// nvethernet accumulates each register in 64-bit software fields, so the low
// word can carry past 32 bits; OR would drop the overlapping bit.
stats := map[string]uint64{
"counter": 1<<32 + 5,
"counter_h": 1,
}
value, ok := ethtoolCounter(stats, "counter", "counter_h")
require.True(t, ok)
assert.Equal(t, uint64(2<<32+5), value)
}
func TestEthtoolCounterRejectsOverflow(t *testing.T) {
stats := map[string]uint64{
"counter": math.MaxUint64,
"counter_h": 1,
}
_, ok := ethtoolCounter(stats, "counter", "counter_h")
assert.False(t, ok)
}
func TestEthtoolCounterFallsBackToLowWord(t *testing.T) {
stats := map[string]uint64{"counter": 456}
value, ok := ethtoolCounter(stats, "counter", "counter_h")
require.True(t, ok)
assert.Equal(t, uint64(456), value)
}

View File

@@ -1,11 +0,0 @@
//go:build !linux
package agent
import psutilNet "github.com/shirou/gopsutil/v4/net"
func isNvidiaEthernet(name string) bool {
return false
}
func correctNvethernetCounters(v *psutilNet.IOCountersStat) {}

View File

@@ -1,195 +0,0 @@
package agent
import (
"errors"
"fmt"
"net/http"
"sync"
"time"
"github.com/henrygd/beszel/internal/entities/monitor"
)
// MonitorManager manages network monitor configurations and task lifetimes.
type MonitorManager struct {
mu sync.RWMutex
monitors map[string]*monitorTask // keyed by monitor ID
probe monitorProbe
certCheck certChecker
resumeGuard monitorResumeGuard
}
func newMonitorManager() *MonitorManager {
return newMonitorManagerWithProbe(networkMonitorProbe(&http.Client{Timeout: monitor.MaxProbeTimeout}))
}
func newMonitorManagerWithProbe(probe monitorProbe) *MonitorManager {
return &MonitorManager{monitors: make(map[string]*monitorTask), probe: probe, certCheck: checkCert}
}
// SyncMonitors replaces all monitor tasks with the given configs.
func (pm *MonitorManager) SyncMonitors(configs []monitor.Config) {
pm.mu.Lock()
defer pm.mu.Unlock()
// Build set of new keys
newKeys := make(map[string]monitor.Config, len(configs))
for _, cfg := range configs {
if cfg.ID == "" {
continue
}
newKeys[cfg.ID] = cfg
}
// Stop removed monitors
for key, task := range pm.monitors {
if _, exists := newKeys[key]; !exists {
task.cancel()
delete(pm.monitors, key)
}
}
// Start new monitors and restart tasks whose config changed.
for key, cfg := range newKeys {
task, exists := pm.monitors[key]
if exists && task.config == cfg {
continue
}
if exists {
task.cancel()
}
task = newMonitorTaskFromExisting(cfg, task)
task.resumeGuard = &pm.resumeGuard
pm.resumeGuard.start()
pm.monitors[key] = task
pm.startMonitor(task)
}
if len(pm.monitors) == 0 {
pm.resumeGuard.shutdown()
}
}
// HandleSyncRequest applies a full or incremental monitor sync request.
func (pm *MonitorManager) HandleSyncRequest(req monitor.SyncRequest) (monitor.SyncResponse, error) {
switch req.Action {
case monitor.SyncActionReplace:
pm.SyncMonitors(req.Configs)
return monitor.SyncResponse{}, nil
case monitor.SyncActionUpsert:
result, err := pm.UpsertMonitor(req.Config, req.RunNow)
if err != nil {
return monitor.SyncResponse{}, err
}
if result == nil {
return monitor.SyncResponse{}, nil
}
return monitor.SyncResponse{Result: *result}, nil
case monitor.SyncActionDelete:
if req.Config.ID == "" {
return monitor.SyncResponse{}, errors.New("missing monitor ID for delete")
}
pm.DeleteMonitor(req.Config.ID)
return monitor.SyncResponse{}, nil
default:
return monitor.SyncResponse{}, fmt.Errorf("unknown monitor sync action: %d", req.Action)
}
}
// UpsertMonitor creates or replaces a single monitor task.
func (pm *MonitorManager) UpsertMonitor(config monitor.Config, runNow bool) (*monitor.Result, error) {
if config.ID == "" {
return nil, errors.New("missing monitor ID")
}
pm.mu.Lock()
task, exists := pm.monitors[config.ID]
if exists && task.config == config {
pm.mu.Unlock()
if !runNow {
return nil, nil
}
return pm.runNow(task), nil
}
if exists {
task.cancel()
}
task = newMonitorTaskFromExisting(config, task)
task.resumeGuard = &pm.resumeGuard
pm.resumeGuard.start()
pm.monitors[config.ID] = task
pm.mu.Unlock()
if runNow {
result := pm.runNow(task)
pm.startMonitor(task)
return result, nil
}
pm.startMonitor(task)
return nil, nil
}
// runNow runs a probe and any due certificate check concurrently, so the
// response fits within the hub's single probe timeout budget.
func (pm *MonitorManager) runNow(task *monitorTask) *monitor.Result {
var wg sync.WaitGroup
wg.Go(func() { task.refreshCert(pm.certCheck) })
result := task.runProbe(pm.probe)
wg.Wait()
if result != nil {
result.Cert = task.certInfo()
}
return result
}
// DeleteMonitor stops and removes a single monitor task.
func (pm *MonitorManager) DeleteMonitor(id string) {
if id == "" {
return
}
pm.mu.Lock()
defer pm.mu.Unlock()
if task, exists := pm.monitors[id]; exists {
task.cancel()
delete(pm.monitors, id)
}
if len(pm.monitors) == 0 {
pm.resumeGuard.shutdown()
}
}
// GetResults returns aggregated results for all monitors over the last supplied duration in ms.
func (pm *MonitorManager) GetResults(durationMs uint16) map[string]monitor.Result {
pm.mu.RLock()
defer pm.mu.RUnlock()
results := make(map[string]monitor.Result, len(pm.monitors))
now := time.Now()
duration := time.Duration(durationMs) * time.Millisecond
for _, task := range pm.monitors {
result, ok := task.history.result(duration, now)
if !ok {
continue
}
// Only the default interval updates monitor records on the hub, so
// realtime requests must not consume the unsent certificate.
if durationMs == defaultDataCacheTimeMs {
result.Cert = task.takeUnsentCert()
}
results[task.config.ID] = result
}
return results
}
// Stop stops all monitor tasks.
func (pm *MonitorManager) Stop() {
pm.mu.Lock()
defer pm.mu.Unlock()
for key, task := range pm.monitors {
task.cancel()
delete(pm.monitors, key)
}
pm.resumeGuard.shutdown()
}

View File

@@ -1,74 +0,0 @@
package agent
import (
"context"
"crypto/tls"
"errors"
"fmt"
"net"
"net/url"
"strings"
"time"
"github.com/henrygd/beszel/internal/entities/monitor"
)
const (
certCheckInterval = 24 * time.Hour
certCheckRetryInterval = time.Hour
)
// certChecker fetches the leaf certificate for an HTTPS target.
type certChecker func(context.Context, string) (monitor.CertInfo, error)
// certCheckEnabled reports whether a monitor's certificate is checked, which is
// the case for every HTTP monitor with an https target.
func certCheckEnabled(config monitor.Config) bool {
return config.Protocol == "http" && len(config.Target) > 8 && strings.EqualFold(config.Target[:8], "https://")
}
// checkCert reads the leaf certificate presented by an HTTPS target. The chain is
// not verified, so expired or self-signed certificates are still reported.
func checkCert(ctx context.Context, target string) (monitor.CertInfo, error) {
address, host, err := certAddress(target)
if err != nil {
return monitor.CertInfo{}, err
}
ctx, cancel := context.WithTimeout(ctx, monitor.MaxProbeTimeout)
defer cancel()
dialer := tls.Dialer{Config: &tls.Config{ServerName: host, InsecureSkipVerify: true}}
conn, err := dialer.DialContext(ctx, "tcp", address)
if err != nil {
return monitor.CertInfo{}, err
}
defer conn.Close()
certs := conn.(*tls.Conn).ConnectionState().PeerCertificates
if len(certs) == 0 {
return monitor.CertInfo{}, errors.New("no peer certificates")
}
leaf := certs[0]
return monitor.CertInfo{
Expires: leaf.NotAfter.UnixMilli(),
Issuer: leaf.Issuer.CommonName,
}, nil
}
// certAddress returns the dial address and server name for an HTTPS URL.
func certAddress(target string) (address, host string, err error) {
u, err := url.Parse(target)
if err != nil {
return "", "", err
}
if !strings.EqualFold(u.Scheme, "https") {
return "", "", fmt.Errorf("certificate check requires an https target: %s", target)
}
host = u.Hostname()
if host == "" {
return "", "", fmt.Errorf("missing host in target: %s", target)
}
port := u.Port()
if port == "" {
port = "443"
}
return net.JoinHostPort(host, port), host, nil
}

View File

@@ -1,184 +0,0 @@
//go:build testing
package agent
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"testing"
"testing/synctest"
"time"
"github.com/henrygd/beszel/internal/entities/monitor"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestCheckCertReadsUnverifiedLeaf(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {}))
defer server.Close()
// httptest uses a self-signed certificate, which must still be reported.
info, err := checkCert(context.Background(), server.URL)
require.NoError(t, err)
leaf := server.Certificate()
assert.Equal(t, leaf.NotAfter.UnixMilli(), info.Expires)
assert.Equal(t, leaf.Issuer.CommonName, info.Issuer)
}
func TestCertAddress(t *testing.T) {
tests := []struct {
target, address, host string
wantErr bool
}{
{target: "https://example.com", address: "example.com:443", host: "example.com"},
{target: "https://example.com:8443/path?q=1", address: "example.com:8443", host: "example.com"},
{target: "HTTPS://[::1]:9443", address: "[::1]:9443", host: "::1"},
{target: "http://example.com", wantErr: true},
{target: "https://", wantErr: true},
}
for _, tt := range tests {
address, host, err := certAddress(tt.target)
if tt.wantErr {
assert.Error(t, err, tt.target)
continue
}
require.NoError(t, err, tt.target)
assert.Equal(t, tt.address, address)
assert.Equal(t, tt.host, host)
}
}
func TestRefreshCertCadence(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
task := newMonitorTask(monitor.Config{ID: "test", Target: "https://example.test", Protocol: "http"})
defer task.cancel()
var calls int
var fail error
// Far enough out that the regular interval applies for the whole test.
expires := time.Now().Add(365 * 24 * time.Hour).UnixMilli()
check := func(context.Context, string) (monitor.CertInfo, error) {
calls++
if fail != nil {
return monitor.CertInfo{}, fail
}
return monitor.CertInfo{Expires: expires + int64(calls)}, nil
}
task.refreshCert(check)
require.NotNil(t, task.certInfo())
assert.Equal(t, expires+1, task.certInfo().Expires)
// Not due again until the check interval passes.
time.Sleep(certCheckInterval - time.Second)
task.refreshCert(check)
assert.Equal(t, 1, calls)
time.Sleep(time.Second)
task.refreshCert(check)
assert.Equal(t, 2, calls)
// Failures keep the last known certificate and retry sooner.
fail = errors.New("connection refused")
time.Sleep(certCheckInterval)
task.refreshCert(check)
assert.Equal(t, 3, calls)
assert.Equal(t, expires+2, task.certInfo().Expires)
time.Sleep(certCheckRetryInterval)
fail = nil
task.refreshCert(check)
assert.Equal(t, 4, calls)
assert.Equal(t, expires+4, task.certInfo().Expires)
})
}
func TestRefreshCertRetriesSoonerNearExpiry(t *testing.T) {
for _, tc := range []struct {
name string
expires time.Duration // relative to the check
interval time.Duration
}{
{"expired", -time.Hour, certCheckRetryInterval},
{"expires before next regular check", certCheckInterval - time.Minute, certCheckRetryInterval},
{"expires after next regular check", certCheckInterval + time.Minute, certCheckInterval},
} {
t.Run(tc.name, func(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
task := newMonitorTask(monitor.Config{ID: "test", Target: "https://example.test", Protocol: "http"})
defer task.cancel()
var calls int
check := func(context.Context, string) (monitor.CertInfo, error) {
calls++
return monitor.CertInfo{Expires: time.Now().Add(tc.expires).UnixMilli()}, nil
}
task.refreshCert(check)
time.Sleep(tc.interval - time.Second)
task.refreshCert(check)
assert.Equal(t, 1, calls)
time.Sleep(time.Second)
task.refreshCert(check)
assert.Equal(t, 2, calls)
})
})
}
}
func TestCertCheckEnabled(t *testing.T) {
tests := []struct {
protocol, target string
want bool
}{
{"http", "https://example.com", true},
{"http", "HTTPS://example.com:8443/path", true},
{"http", "http://example.com", false},
{"http", "https://", false},
{"tcp", "https://example.com", false},
{"icmp", "example.com", false},
}
for _, tt := range tests {
assert.Equal(t, tt.want, certCheckEnabled(monitor.Config{Protocol: tt.protocol, Target: tt.target}), tt.protocol+" "+tt.target)
}
}
func TestRefreshCertSkipsNonHTTPS(t *testing.T) {
task := newMonitorTask(monitor.Config{ID: "test", Target: "http://example.test", Protocol: "http"})
defer task.cancel()
task.refreshCert(func(context.Context, string) (monitor.CertInfo, error) {
t.Fatal("certificate check must not run for non-https targets")
return monitor.CertInfo{}, nil
})
assert.Nil(t, task.certInfo())
}
func TestUpsertMonitorRunNowIncludesCert(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {}))
defer server.Close()
pm := newMonitorManagerWithProbe(func(context.Context, monitor.Config) (int64, error) { return 100, nil })
defer pm.Stop()
config := monitor.Config{ID: "cert", Target: server.URL, Protocol: "http", Interval: 60}
result, err := pm.UpsertMonitor(config, true)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Cert)
assert.Equal(t, server.Certificate().NotAfter.UnixMilli(), result.Cert.Expires)
// Realtime results never carry the certificate, and the default interval
// sends it only once per check.
assert.Nil(t, pm.GetResults(1000)["cert"].Cert)
results := pm.GetResults(defaultDataCacheTimeMs)
require.NotNil(t, results["cert"].Cert)
assert.Equal(t, result.Cert.Expires, results["cert"].Cert.Expires)
assert.Nil(t, pm.GetResults(defaultDataCacheTimeMs)["cert"].Cert)
// Changing the interval keeps the known certificate without resending it.
config.Interval = 30
_, err = pm.UpsertMonitor(config, false)
require.NoError(t, err)
pm.mu.RLock()
task := pm.monitors["cert"]
pm.mu.RUnlock()
assert.NotNil(t, task.certInfo())
assert.Nil(t, pm.GetResults(defaultDataCacheTimeMs)["cert"].Cert)
}

View File

@@ -1,274 +0,0 @@
package agent
import (
"math"
"sync"
"time"
"github.com/henrygd/beszel/internal/entities/monitor"
)
// Monitors run at user-defined intervals (e.g., every 10s).
// To keep memory usage low and constant, data is stored in two layers:
// 1. Raw samples: The most recent individual results (kept for monitorRawRetention).
// 2. Minute buckets: A ring buffer of 61 buckets, each representing one
// wall-clock minute. Samples collected within the same minute are aggregated
// (sum, min, max, count) into a single bucket.
//
// Short-term requests (<= 61s) use raw samples.
// Long-term requests (up to 1h) use the minute buckets to avoid storing thousands
// of individual data points.
const (
// monitorRawRetention is the duration to keep individual samples
monitorRawRetention = 61 * time.Second
// monitorMinuteBucketLen is the number of 1-minute buckets to keep (1 hour + 1 for partials)
monitorMinuteBucketLen int32 = 61
)
// monitorHistory owns retention and aggregation, independently of probe execution.
type monitorHistory struct {
mu sync.Mutex
sampleCount int64
samples []monitorSample
buckets [monitorMinuteBucketLen]monitorBucket
}
func newMonitorHistory() *monitorHistory {
// Start small for typical intervals; append grows the buffer for faster probes.
return &monitorHistory{samples: make([]monitorSample, 0, 4)}
}
func (h *monitorHistory) clone() *monitorHistory {
h.mu.Lock()
defer h.mu.Unlock()
cloned := newMonitorHistory()
cloned.samples = append(cloned.samples, h.samples...)
cloned.buckets = h.buckets
cloned.sampleCount = h.sampleCount
return cloned
}
func (h *monitorHistory) result(duration time.Duration, now time.Time) (monitor.Result, bool) {
h.mu.Lock()
defer h.mu.Unlock()
return h.resultLocked(duration, now)
}
func (h *monitorHistory) record(sample monitorSample) monitor.Result {
h.mu.Lock()
defer h.mu.Unlock()
h.addSampleLocked(sample)
result, _ := h.resultLocked(time.Minute, sample.timestamp)
return result
}
// monitorSample stores one monitor attempt and its collection time.
type monitorSample struct {
responseUs int64 // -1 means loss
timestamp time.Time
}
// monitorBucket stores one minute of aggregated monitor data.
type monitorBucket struct {
minute int32
filled bool
stats monitorAggregate
}
// monitorAggregate accumulates successful response stats and total sample counts.
type monitorAggregate struct {
sumUs int64
minUs int64
maxUs int64
totalCount int64
successCount int64
}
// newMonitorAggregate initializes an aggregate with an unset minimum value.
func newMonitorAggregate() monitorAggregate {
return monitorAggregate{minUs: math.MaxInt64}
}
// addResponse folds a single monitor sample into the aggregate.
func (agg *monitorAggregate) addResponse(responseUs int64) {
agg.totalCount++
if responseUs < 0 {
return
}
agg.successCount++
agg.sumUs += responseUs
if responseUs < agg.minUs {
agg.minUs = responseUs
}
if responseUs > agg.maxUs {
agg.maxUs = responseUs
}
}
// addAggregate merges another aggregate into this one.
func (agg *monitorAggregate) addAggregate(other monitorAggregate) {
if other.totalCount == 0 {
return
}
agg.totalCount += other.totalCount
agg.successCount += other.successCount
agg.sumUs += other.sumUs
if other.successCount == 0 {
return
}
if agg.minUs == math.MaxInt64 || other.minUs < agg.minUs {
agg.minUs = other.minUs
}
if other.maxUs > agg.maxUs {
agg.maxUs = other.maxUs
}
}
// hasData reports whether the aggregate contains any samples.
func (agg monitorAggregate) hasData() bool {
return agg.totalCount > 0
}
// result converts the aggregate into the monitor result format.
func (agg monitorAggregate) result() monitor.Result {
avg := agg.avgResponse()
result := monitor.Result{
AvgResponse: avg,
MinResponse: agg.minUs,
MaxResponse: agg.maxUs,
PacketLoss: agg.lossPercentage(),
TotalCount: agg.totalCount,
SuccessCount: agg.successCount,
ResponseSum: agg.sumUs,
}
if agg.successCount == 0 {
result.MinResponse, result.MaxResponse = 0, 0
}
return result
}
// avgResponse returns the rounded average of successful samples.
func (agg monitorAggregate) avgResponse() int64 {
if agg.successCount == 0 {
return 0
}
return agg.sumUs / agg.successCount
}
// lossPercentage returns the rounded failure rate for the aggregate.
func (agg monitorAggregate) lossPercentage() float64 {
if agg.totalCount == 0 {
return 0
}
return math.Round(float64(agg.totalCount-agg.successCount)/float64(agg.totalCount)*10000) / 100
}
// resultLocked returns the aggregated monitor result for the requested duration along with a bool indicating whether any data was available.
func (h *monitorHistory) resultLocked(duration time.Duration, now time.Time) (monitor.Result, bool) {
agg := h.aggregateLocked(duration, now)
if !agg.hasData() {
// short realtime windows (e.g. the 1s window used for 1m/realtime charts) often fall
// between monitor samples since monitors run at longer, user-defined intervals; fall back to
// the most recent sample so realtime requests still report current status.
agg = h.latestSampleAggregateLocked()
}
hourAgg := h.aggregateLocked(time.Hour, now)
if !agg.hasData() {
return monitor.Result{}, false
}
result := agg.result()
if len(h.samples) > 0 {
result.LastProbeAt = h.samples[len(h.samples)-1].timestamp.UnixMilli()
}
result.AvgResponse1h = hourAgg.avgResponse()
result.MinResponse1h = hourAgg.minUs
result.MaxResponse1h = hourAgg.maxUs
result.PacketLoss1h = hourAgg.lossPercentage()
result.SampleCount = h.sampleCount
if hourAgg.successCount == 0 {
result.MinResponse1h, result.MaxResponse1h = 0, 0
}
return result, true
}
// latestSampleAggregateLocked returns an aggregate containing only the most recent sample, if any.
func (h *monitorHistory) latestSampleAggregateLocked() monitorAggregate {
agg := newMonitorAggregate()
if len(h.samples) == 0 {
return agg
}
agg.addResponse(h.samples[len(h.samples)-1].responseUs)
return agg
}
// aggregateLocked collects monitor data for the requested time window.
func (h *monitorHistory) aggregateLocked(duration time.Duration, now time.Time) monitorAggregate {
cutoff := now.Add(-duration)
// Keep short windows exact; longer windows read from minute buckets to avoid raw-sample retention.
if duration <= monitorRawRetention {
return aggregateSamplesSince(h.samples, cutoff)
}
return aggregateBucketsSince(h.buckets[:], cutoff, now)
}
// aggregateSamplesSince aggregates raw samples newer than the cutoff.
func aggregateSamplesSince(samples []monitorSample, cutoff time.Time) monitorAggregate {
agg := newMonitorAggregate()
for _, sample := range samples {
if sample.timestamp.Before(cutoff) {
continue
}
agg.addResponse(sample.responseUs)
}
return agg
}
// aggregateBucketsSince aggregates minute buckets overlapping the requested window.
func aggregateBucketsSince(buckets []monitorBucket, cutoff, now time.Time) monitorAggregate {
agg := newMonitorAggregate()
startMinute := int32(cutoff.Unix() / 60)
endMinute := int32(now.Unix() / 60)
for _, bucket := range buckets {
if !bucket.filled || bucket.minute < startMinute || bucket.minute > endMinute {
continue
}
agg.addAggregate(bucket.stats)
}
return agg
}
// addSampleLocked stores a fresh sample in both raw and per-minute retention buffers.
func (h *monitorHistory) addSampleLocked(sample monitorSample) {
h.sampleCount++
cutoff := sample.timestamp.Add(-monitorRawRetention)
start := 0
for i := range h.samples {
if !h.samples[i].timestamp.Before(cutoff) {
start = i
break
}
if i == len(h.samples)-1 {
start = len(h.samples)
}
}
if start > 0 {
size := copy(h.samples, h.samples[start:])
h.samples = h.samples[:size]
}
h.samples = append(h.samples, sample)
minute := int32(sample.timestamp.Unix() / 60)
// Each slot stores one wall-clock minute, so the ring stays fixed-size at ~1h per monitor.
bucket := &h.buckets[minute%monitorMinuteBucketLen]
if !bucket.filled || bucket.minute != minute {
bucket.minute = minute
bucket.filled = true
bucket.stats = newMonitorAggregate()
}
bucket.stats.addResponse(sample.responseUs)
}

View File

@@ -1,154 +0,0 @@
package agent
import (
"testing"
"time"
"github.com/fxamacker/cbor/v2"
"github.com/henrygd/beszel/internal/entities/monitor"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestMonitorHistoryWindowCounts(t *testing.T) {
history := newMonitorHistory()
now := time.Now()
// This older success counts toward lifetime warm-up, but not this window.
history.record(monitorSample{responseUs: 1000, timestamp: now.Add(-2 * time.Minute)})
history.record(monitorSample{responseUs: 10, timestamp: now.Add(-30 * time.Second)})
history.record(monitorSample{responseUs: 21, timestamp: now.Add(-20 * time.Second)})
history.record(monitorSample{responseUs: -1, timestamp: now.Add(-10 * time.Second)})
result, ok := history.result(time.Minute, now)
require.True(t, ok)
assert.EqualValues(t, 4, result.SampleCount)
assert.EqualValues(t, 3, result.TotalCount)
assert.EqualValues(t, 2, result.SuccessCount)
assert.EqualValues(t, 31, result.ResponseSum, "preserve the sum before average rounding")
assert.EqualValues(t, 15, result.AvgResponse)
assert.Equal(t, 33.33, result.PacketLoss)
encoded, err := cbor.Marshal(result)
require.NoError(t, err)
var decoded monitor.Result
require.NoError(t, cbor.Unmarshal(encoded, &decoded))
assert.Equal(t, result, decoded)
stats := monitor.Stats{}.FromResult(decoded)
assert.Equal(t, result.TotalCount, stats.TotalCount)
assert.Equal(t, result.SuccessCount, stats.SuccessCount)
assert.Equal(t, result.ResponseSum, stats.ResponseSum)
// Reads do not consume samples. A short window's latest-sample fallback
// carries the count for that single failure, not the minute or lifetime count.
repeated, _ := history.result(time.Minute, now)
assert.Equal(t, result, repeated)
fallback, ok := history.result(time.Second, now)
require.True(t, ok)
assert.EqualValues(t, 1, fallback.TotalCount)
assert.Zero(t, fallback.SuccessCount)
assert.Zero(t, fallback.ResponseSum)
assert.Equal(t, 100.0, fallback.PacketLoss)
assert.EqualValues(t, 4, fallback.SampleCount)
}
func TestMonitorHistoryAggregateLockedUsesRawSamplesForShortWindows(t *testing.T) {
now := time.Date(2026, time.April, 21, 12, 0, 0, 0, time.UTC)
history := newMonitorHistory()
history.addSampleLocked(monitorSample{responseUs: 10, timestamp: now.Add(-90 * time.Second)})
history.addSampleLocked(monitorSample{responseUs: 20, timestamp: now.Add(-30 * time.Second)})
history.addSampleLocked(monitorSample{responseUs: -1, timestamp: now.Add(-10 * time.Second)})
agg := history.aggregateLocked(time.Minute, now)
require.True(t, agg.hasData())
assert.Equal(t, int64(2), agg.totalCount)
assert.Equal(t, int64(1), agg.successCount)
result := agg.result()
assert.Equal(t, int64(20), result.AvgResponse)
assert.Equal(t, int64(20), result.MinResponse)
assert.Equal(t, int64(20), result.MaxResponse)
assert.Equal(t, 50.0, result.PacketLoss)
}
func TestMonitorHistoryAggregateLockedUsesMinuteBucketsForLongWindows(t *testing.T) {
now := time.Date(2026, time.April, 21, 12, 0, 30, 0, time.UTC)
history := newMonitorHistory()
history.addSampleLocked(monitorSample{responseUs: 10, timestamp: now.Add(-11 * time.Minute)})
history.addSampleLocked(monitorSample{responseUs: 20, timestamp: now.Add(-9 * time.Minute)})
history.addSampleLocked(monitorSample{responseUs: 40, timestamp: now.Add(-5 * time.Minute)})
history.addSampleLocked(monitorSample{responseUs: -1, timestamp: now.Add(-90 * time.Second)})
history.addSampleLocked(monitorSample{responseUs: 30, timestamp: now.Add(-30 * time.Second)})
agg := history.aggregateLocked(10*time.Minute, now)
require.True(t, agg.hasData())
assert.Equal(t, int64(4), agg.totalCount)
assert.Equal(t, int64(3), agg.successCount)
result := agg.result()
assert.Equal(t, int64(30), result.AvgResponse)
assert.Equal(t, int64(20), result.MinResponse)
assert.Equal(t, int64(40), result.MaxResponse)
assert.Equal(t, 25.0, result.PacketLoss)
}
func TestMonitorHistoryAddSampleLockedTrimsRawSamplesButKeepsBucketHistory(t *testing.T) {
now := time.Date(2026, time.April, 21, 12, 0, 0, 0, time.UTC)
history := newMonitorHistory()
history.addSampleLocked(monitorSample{responseUs: 10, timestamp: now.Add(-10 * time.Minute)})
history.addSampleLocked(monitorSample{responseUs: 20, timestamp: now})
require.Len(t, history.samples, 1)
assert.Equal(t, int64(20), history.samples[0].responseUs)
agg := history.aggregateLocked(10*time.Minute, now)
require.True(t, agg.hasData())
assert.Equal(t, int64(2), agg.totalCount)
assert.Equal(t, int64(2), agg.successCount)
result := agg.result()
assert.Equal(t, int64(15), result.AvgResponse)
assert.Equal(t, int64(10), result.MinResponse)
assert.Equal(t, int64(20), result.MaxResponse)
assert.Equal(t, 0.0, result.PacketLoss)
}
func TestMonitorHistoryProbeTimestamp(t *testing.T) {
history := newMonitorHistory()
start := time.Date(2026, time.September, 14, 12, 0, 0, 0, time.UTC)
_, ok := history.result(time.Minute, start)
require.False(t, ok)
first := history.record(monitorSample{responseUs: 20, timestamp: start})
assert.Equal(t, start.UnixMilli(), first.LastProbeAt)
for minute := 0; minute < 5; minute++ {
now := start.Add(time.Duration(minute)*time.Minute + time.Second)
// Realtime reads must not consume freshness for the persistence request.
for _, window := range []time.Duration{time.Second, time.Minute} {
result, ok := history.result(window, now)
require.True(t, ok)
assert.Equal(t, first.LastProbeAt, result.LastProbeAt)
assert.Equal(t, int64(20), result.AvgResponse)
}
}
next := start.Add(5 * time.Minute)
failed := history.record(monitorSample{responseUs: -1, timestamp: next})
assert.Equal(t, next.UnixMilli(), failed.LastProbeAt)
assert.Equal(t, float64(100), failed.PacketLoss)
repeated, ok := history.result(time.Minute, next.Add(2*time.Minute))
require.True(t, ok)
assert.Equal(t, failed.LastProbeAt, repeated.LastProbeAt)
assert.Equal(t, float64(100), repeated.PacketLoss)
}
func TestMonitorHistorySampleCount(t *testing.T) {
history := newMonitorHistory()
now := time.Now()
// Both failed and successful probes count, including older samples so
// monitors with hourly intervals can finish warming up.
history.record(monitorSample{responseUs: -1, timestamp: now.Add(-2 * time.Hour)})
for i, response := range []int64{10, -1, 20} {
result := history.record(monitorSample{responseUs: response, timestamp: now.Add(time.Duration(i) * time.Second)})
assert.EqualValues(t, i+2, result.SampleCount)
}
result, ok := history.clone().result(time.Minute, now.Add(3*time.Second))
require.True(t, ok)
assert.EqualValues(t, 4, result.SampleCount)
}

View File

@@ -1,312 +0,0 @@
package agent
import (
"bytes"
"context"
"crypto/rand"
"errors"
"fmt"
"math"
"net"
"os"
"os/exec"
"regexp"
"runtime"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
"log/slog"
)
// Match the numeric RTT independently of the localized label used by Windows.
var pingTimeRegex = regexp.MustCompile(`(?i)[=<]\s*([0-9]+(?:[.,][0-9]+)?)\s*ms\b`)
var icmpSequence atomic.Uint32
type icmpPacketConn interface {
Close() error
}
// icmpMethod tracks which ICMP approach to use. Once a method succeeds or
// all native methods fail, the choice is cached so subsequent monitors skip
// the trial-and-error overhead.
type icmpMethod uint8
const (
icmpUntried icmpMethod = iota // haven't tried yet
icmpRaw // privileged raw socket
icmpDatagram // unprivileged datagram socket
icmpExecFallback // shell out to system ping command
)
// icmpFamily holds the network parameters and cached detection result for one address family.
type icmpFamily struct {
rawNetwork string // e.g. "ip4:icmp" or "ip6:ipv6-icmp"
dgramNetwork string // e.g. "udp4" or "udp6"
listenAddr string // "0.0.0.0" or "::"
echoType icmp.Type // outgoing echo request type
replyType icmp.Type // expected echo reply type
proto int // IANA protocol number for parsing replies
isIPv6 bool
mode icmpMethod // cached detection result (guarded by icmpModeMu)
}
var (
icmpV4 = icmpFamily{
rawNetwork: "ip4:icmp",
dgramNetwork: "udp4",
listenAddr: "0.0.0.0",
echoType: ipv4.ICMPTypeEcho,
replyType: ipv4.ICMPTypeEchoReply,
proto: 1,
}
icmpV6 = icmpFamily{
rawNetwork: "ip6:ipv6-icmp",
dgramNetwork: "udp6",
listenAddr: "::",
echoType: ipv6.ICMPTypeEchoRequest,
replyType: ipv6.ICMPTypeEchoReply,
proto: 58,
isIPv6: true,
}
icmpModeMu sync.Mutex
icmpListen = func(network, listenAddr string) (icmpPacketConn, error) {
return icmp.ListenPacket(network, listenAddr)
}
)
// monitorICMP sends an ICMP echo request and measures round-trip response.
// Supports both IPv4 and IPv6 targets. The ICMP method (raw socket,
// unprivileged datagram, or exec fallback) is detected once per address
// family and cached for subsequent monitors.
// Returns response in microseconds, or -1 and an error on failure.
func monitorICMP(ctx context.Context, target string) (int64, error) {
ctx, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
family, ip, err := resolveICMPTarget(ctx, target)
if err != nil {
return -1, err
}
icmpModeMu.Lock()
if family.mode == icmpUntried {
family.mode = detectICMPMode(family, icmpListen)
}
mode := family.mode
icmpModeMu.Unlock()
switch mode {
case icmpRaw:
return monitorICMPNative(ctx, family.rawNetwork, family, &net.IPAddr{IP: ip})
case icmpDatagram:
return monitorICMPNative(ctx, family.dgramNetwork, family, &net.UDPAddr{IP: ip})
case icmpExecFallback:
return monitorICMPExec(ctx, ip.String(), family.isIPv6)
default:
return -1, errors.New("unsupported ICMP mode")
}
}
// resolveICMPTarget resolves a target hostname or IP to determine the address
// family and concrete IP address. Prefers IPv4 for dual-stack hostnames.
func resolveICMPTarget(ctx context.Context, target string) (*icmpFamily, net.IP, error) {
if ip := net.ParseIP(target); ip != nil {
if ip.To4() != nil {
return &icmpV4, ip.To4(), nil
}
return &icmpV6, ip, nil
}
ips, err := net.DefaultResolver.LookupIP(ctx, "ip", target)
if err != nil || len(ips) == 0 {
return nil, nil, err
}
for _, ip := range ips {
if v4 := ip.To4(); v4 != nil {
return &icmpV4, v4, nil
}
}
return &icmpV6, ips[0], nil
}
func detectICMPMode(family *icmpFamily, listen func(network, listenAddr string) (icmpPacketConn, error)) icmpMethod {
label := "IPv4"
if family.isIPv6 {
label = "IPv6"
}
conn, err := listen(family.rawNetwork, family.listenAddr)
slog.Debug("ICMP raw socket test", "family", label, "err", err)
if err == nil {
conn.Close()
return icmpRaw
}
conn, err = listen(family.dgramNetwork, family.listenAddr)
slog.Debug("ICMP datagram socket test", "family", label, "err", err)
if err == nil {
conn.Close()
return icmpDatagram
}
return icmpExecFallback
}
// monitorICMPNative sends an ICMP echo request using Go's x/net/icmp package.
func monitorICMPNative(ctx context.Context, network string, family *icmpFamily, dst net.Addr) (int64, error) {
conn, err := icmp.ListenPacket(network, family.listenAddr)
if err != nil {
return -1, err
}
defer conn.Close()
return monitorICMPPacket(ctx, conn, family, dst)
}
func monitorICMPPacket(ctx context.Context, conn net.PacketConn, family *icmpFamily, dst net.Addr) (int64, error) {
if err := ctx.Err(); err != nil {
return -1, err
}
// Closing the socket interrupts both reads and writes on cancellation.
stop := context.AfterFunc(ctx, func() { _ = conn.Close() })
defer stop()
// Prepare correlation data before starting the round-trip timer. The token
// also distinguishes delayed replies after the 16-bit sequence wraps.
token := make([]byte, 16)
if _, err := rand.Read(token); err != nil {
return -1, err
}
echo := &icmp.Echo{
ID: os.Getpid() & 0xffff,
Seq: int(icmpSequence.Add(1) & 0xffff),
Data: token,
}
// Linux ping sockets replace the Echo ID with their bound port. Darwin
// datagram sockets and raw sockets preserve the supplied ID.
if local, ok := conn.LocalAddr().(*net.UDPAddr); ok && runtime.GOOS == "linux" {
echo.ID = local.Port
}
targetIP := icmpAddrIP(dst)
msg := &icmp.Message{
Type: family.echoType,
Code: 0,
Body: echo,
}
msgBytes, err := msg.Marshal(nil)
if err != nil {
return -1, err
}
// Set deadline before sending
if err := conn.SetDeadline(time.Now().Add(3 * time.Second)); err != nil {
return -1, err
}
buf := make([]byte, 1500)
start := time.Now()
if _, err := conn.WriteTo(msgBytes, dst); err != nil {
return -1, err
}
// Read reply
for {
n, peer, err := conn.ReadFrom(buf)
received := time.Now()
if err != nil {
return -1, err
}
if !targetIP.Equal(icmpAddrIP(peer)) {
continue
}
reply, err := icmp.ParseMessage(family.proto, buf[:n])
if err != nil || reply.Type != family.replyType || reply.Code != 0 {
continue
}
body, ok := reply.Body.(*icmp.Echo)
if ok && body.ID == echo.ID && body.Seq == echo.Seq && bytes.Equal(body.Data, echo.Data) {
return received.Sub(start).Microseconds(), nil
}
// Keep waiting for our reply without extending the original deadline.
}
}
func icmpAddrIP(addr net.Addr) net.IP {
switch addr := addr.(type) {
case *net.IPAddr:
return addr.IP
case *net.UDPAddr:
return addr.IP
default:
return nil
}
}
// pingCommand selects the executable and arguments for the supported agent platforms.
// The context deadline enforces the timeout: -W has incompatible meanings across
// Linux, BSD IPv4 ping, and macOS ping6.
func pingCommand(goos, target string, isIPv6 bool) (string, []string, error) {
family := "-4"
if isIPv6 {
family = "-6"
}
switch goos {
case "windows":
return "ping", []string{family, "-n", "1", "-w", "3000", target}, nil
case "linux":
return "ping", []string{family, "-n", "-c", "1", target}, nil
case "darwin", "freebsd", "openbsd":
command := "ping"
if isIPv6 {
command = "ping6"
}
return command, []string{"-n", "-c", "1", target}, nil
default:
return "", nil, fmt.Errorf("ping fallback is unsupported on %s", goos)
}
}
// monitorICMPExec falls back to the system ping command. Returns -1 and an error on failure.
func monitorICMPExec(ctx context.Context, target string, isIPv6 bool) (int64, error) {
ctx, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
name, args, err := pingCommand(runtime.GOOS, target, isIPv6)
if err != nil {
return -1, err
}
cmd := exec.CommandContext(ctx, name, args...)
// Keep Unix output and decimal formatting stable. Windows ignores LC_ALL.
cmd.Env = append(os.Environ(), "LC_ALL=C")
output, err := cmd.Output()
if ctx.Err() != nil {
return -1, ctx.Err()
}
if err != nil {
return -1, fmt.Errorf("%s failed: %w", name, err)
}
return parsePingResponse(output)
}
// parsePingResponse returns the reported RTT, never subprocess execution time.
// For a bounded value such as Windows' time<1ms, retain the reported upper bound.
func parsePingResponse(output []byte) (int64, error) {
matches := pingTimeRegex.FindSubmatch(output)
if len(matches) < 2 {
return -1, errors.New("ping output contains no round-trip time")
}
ms, err := strconv.ParseFloat(strings.ReplaceAll(string(matches[1]), ",", "."), 64)
if err != nil || math.IsInf(ms, 0) || ms >= float64(math.MaxInt64)/1000 {
return -1, errors.New("invalid round-trip time in ping output")
}
return int64(math.Round(ms * 1000)), nil
}

View File

@@ -1,433 +0,0 @@
//go:build testing
package agent
import (
"context"
"errors"
"fmt"
"net"
"os"
"path/filepath"
"runtime"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/net/icmp"
)
type testICMPPacketConn struct{}
func (testICMPPacketConn) Close() error { return nil }
type blockingICMPConn struct {
net.PacketConn
reading chan struct{}
}
func (c *blockingICMPConn) WriteTo(p []byte, addr net.Addr) (int, error) {
return len(p), nil
}
func (c *blockingICMPConn) ReadFrom(p []byte) (int, net.Addr, error) {
close(c.reading)
return c.PacketConn.ReadFrom(p)
}
func TestMonitorICMPPacketCancellation(t *testing.T) {
conn, err := net.ListenPacket("udp4", "127.0.0.1:0")
require.NoError(t, err)
defer conn.Close()
blocking := &blockingICMPConn{PacketConn: conn, reading: make(chan struct{})}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
done := make(chan error, 1)
go func() {
_, err := monitorICMPPacket(ctx, blocking, &icmpV4, conn.LocalAddr())
done <- err
}()
select {
case <-blocking.reading:
case <-time.After(time.Second):
t.Fatal("probe did not begin reading")
}
cancel()
select {
case err := <-done:
require.Error(t, err)
case <-time.After(time.Second):
t.Fatal("cancellation did not interrupt the socket read")
}
}
func TestMonitorICMPExecCancellation(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("test uses a POSIX shell stub for ping")
}
dir := t.TempDir()
require.NoError(t, os.WriteFile(filepath.Join(dir, "ping"), []byte("#!/bin/sh\nexec sleep 30\n"), 0o755))
t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH"))
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
defer cancel()
done := make(chan error, 1)
go func() {
_, err := monitorICMPExec(ctx, "127.0.0.1", false)
done <- err
}()
select {
case err := <-done:
require.ErrorIs(t, err, context.DeadlineExceeded)
case <-time.After(time.Second):
t.Fatal("cancellation did not terminate ping")
}
}
func TestPingCommand(t *testing.T) {
for _, goos := range []string{"linux", "windows", "darwin", "freebsd", "openbsd"} {
for _, ipv6 := range []bool{false, true} {
t.Run(fmt.Sprintf("%s/ipv6=%t", goos, ipv6), func(t *testing.T) {
target, family := "192.0.2.1", "-4"
if ipv6 {
target, family = "2001:db8::1", "-6"
}
name, args, err := pingCommand(goos, target, ipv6)
require.NoError(t, err)
wantName := "ping"
wantArgs := []string{"-n", "-c", "1", target}
switch goos {
case "windows":
wantArgs = []string{family, "-n", "1", "-w", "3000", target}
case "linux":
wantArgs = append([]string{family}, wantArgs...)
default:
if ipv6 {
wantName = "ping6"
}
}
assert.Equal(t, wantName, name)
assert.Equal(t, wantArgs, args)
})
}
}
_, _, err := pingCommand("unsupported", "192.0.2.1", false)
require.Error(t, err)
}
func TestParsePingResponse(t *testing.T) {
for _, tc := range []struct {
name string
output string
wantUs int64
}{
{"linux", "64 bytes from 192.0.2.1: icmp_seq=1 ttl=64 time=12.345 ms", 12345},
{"bsd", "64 bytes from 192.0.2.1: icmp_seq=0 ttl=64 time=0.023 ms", 23},
{"ipv6", "64 bytes from 2001:db8::1: icmp_seq=0 hlim=64 time=1.234 ms", 1234},
{"windows", "Reply from 192.0.2.1: bytes=32 time=12ms TTL=128", 12000},
{"windows submillisecond", "Reply from ::1: time<1ms", 1000},
{"localized windows", "Antwort von 192.0.2.1: Bytes=32 Zeit=12ms TTL=128", 12000},
{"decimal comma", "64 bytes from 192.0.2.1: time=1,234 ms", 1234},
{"rounding", "time=0.1236 ms", 124},
{"empty", "", -1},
{"timeout", "Request timed out.", -1},
{"unreachable", "Reply from 192.0.2.1: Destination host unreachable.", -1},
{"malformed", "time=oops ms", -1},
{"negative", "time=-1 ms", -1},
{"overflow", "time=999999999999999999999 ms", -1},
} {
t.Run(tc.name, func(t *testing.T) {
responseUs, err := parsePingResponse([]byte(tc.output))
if tc.wantUs < 0 {
require.Error(t, err)
} else {
require.NoError(t, err)
}
assert.Equal(t, tc.wantUs, responseUs)
})
}
}
func TestMonitorICMPExecOutput(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("test uses a POSIX shell stub for ping")
}
for _, tc := range []struct {
name string
output string
exit int
wantUs int64
}{
{"success", "time=1.234 ms", 0, 1234},
{"missing RTT", "unrecognized output", 0, -1},
{"failed command with RTT", "time=1.234 ms", 1, -1},
} {
t.Run(tc.name, func(t *testing.T) {
dir := t.TempDir()
// Also verify an inherited locale cannot override the C locale.
script := fmt.Sprintf("#!/bin/sh\n[ \"$LC_ALL\" = C ] || exit 2\nprintf '%%s\\n' '%s'\nexit %d\n", tc.output, tc.exit)
require.NoError(t, os.WriteFile(filepath.Join(dir, "ping"), []byte(script), 0o755))
t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH"))
t.Setenv("LC_ALL", "de_DE.UTF-8")
responseUs, err := monitorICMPExec(t.Context(), "127.0.0.1", false)
if tc.wantUs < 0 {
require.Error(t, err)
} else {
require.NoError(t, err)
}
assert.Equal(t, tc.wantUs, responseUs)
})
}
}
type icmpTestReply struct {
data []byte
peer net.Addr
}
type scriptedICMPConn struct {
net.PacketConn
local net.Addr
onWrite func([]byte, net.Addr)
replies []icmpTestReply
reads int
deadlineSets int
}
func (c *scriptedICMPConn) LocalAddr() net.Addr { return c.local }
func (c *scriptedICMPConn) SetDeadline(deadline time.Time) error {
c.deadlineSets++
return nil
}
func (c *scriptedICMPConn) WriteTo(data []byte, dst net.Addr) (int, error) {
c.onWrite(data, dst)
return len(data), nil
}
func (c *scriptedICMPConn) ReadFrom(buf []byte) (int, net.Addr, error) {
c.reads++
if len(c.replies) == 0 {
return 0, nil, os.ErrDeadlineExceeded
}
reply := c.replies[0]
c.replies = c.replies[1:]
return copy(buf, reply.data), reply.peer, nil
}
func TestMonitorICMPReplyCorrelation(t *testing.T) {
for _, family := range []*icmpFamily{&icmpV4, &icmpV6} {
for _, datagram := range []bool{false, true} {
network := family.rawNetwork
ip, other := net.ParseIP("192.0.2.1"), net.ParseIP("192.0.2.2")
if family.isIPv6 {
ip, other = net.ParseIP("2001:db8::1"), net.ParseIP("2001:db8::2")
}
var dst net.Addr = &net.IPAddr{IP: ip}
var wrongPeer net.Addr = &net.IPAddr{IP: other}
if datagram {
network = family.dgramNetwork
dst = &net.UDPAddr{IP: ip}
wrongPeer = &net.UDPAddr{IP: other}
}
for _, mismatch := range []string{"source", "id", "sequence", "payload", "type", "code", "malformed"} {
for _, eventuallyMatches := range []bool{false, true} {
ending := "timeout"
if eventuallyMatches {
ending = "success"
}
t.Run(network+"/"+mismatch+"/"+ending, func(t *testing.T) {
conn := &scriptedICMPConn{local: &net.IPAddr{IP: net.IPv4zero}}
if datagram {
conn.local = &net.UDPAddr{Port: 12345}
if runtime.GOOS == "linux" {
// Deliberately differ from the process ID.
conn.local = &net.UDPAddr{Port: (os.Getpid() % 65534) + 1}
}
}
conn.onWrite = func(data []byte, target net.Addr) {
require.Equal(t, dst, target)
request, err := icmp.ParseMessage(family.proto, data)
require.NoError(t, err)
echo := request.Body.(*icmp.Echo)
expectedID := os.Getpid() & 0xffff
if datagram && runtime.GOOS == "linux" {
expectedID = conn.local.(*net.UDPAddr).Port
}
require.Equal(t, expectedID, echo.ID)
reply := &icmp.Message{Type: family.replyType, Body: echo}
valid, err := reply.Marshal(nil)
require.NoError(t, err)
peer := dst
switch mismatch {
case "source":
peer = wrongPeer
case "id":
echo.ID ^= 1
case "sequence":
echo.Seq ^= 1
case "payload":
echo.Data[0] ^= 1
case "type":
reply.Type = family.echoType
case "code":
reply.Code = 1
}
invalid, err := reply.Marshal(nil)
require.NoError(t, err)
if mismatch == "malformed" {
invalid = invalid[:2]
}
conn.replies = []icmpTestReply{{invalid, peer}}
if eventuallyMatches {
conn.replies = append(conn.replies, icmpTestReply{valid, dst})
}
}
elapsed, err := monitorICMPPacket(context.Background(), conn, family, dst)
if eventuallyMatches {
require.NoError(t, err)
assert.GreaterOrEqual(t, elapsed, int64(0))
} else {
require.ErrorIs(t, err, os.ErrDeadlineExceeded)
assert.Equal(t, int64(-1), elapsed)
}
assert.Equal(t, 2, conn.reads)
assert.Equal(t, 1, conn.deadlineSets)
})
}
}
}
}
}
func TestMonitorICMPLoopback(t *testing.T) {
for _, family := range []*icmpFamily{&icmpV4, &icmpV6} {
for _, network := range []string{family.rawNetwork, family.dgramNetwork} {
t.Run(network, func(t *testing.T) {
conn, err := icmp.ListenPacket(network, family.listenAddr)
if err != nil {
t.Skipf("ICMP socket unavailable: %v", err)
}
defer conn.Close()
ip := net.ParseIP("127.0.0.1")
if family.isIPv6 {
ip = net.ParseIP("::1")
}
var dst net.Addr = &net.IPAddr{IP: ip}
if network == family.dgramNetwork {
dst = &net.UDPAddr{IP: ip}
}
elapsed, err := monitorICMPPacket(context.Background(), conn, family, dst)
require.NoError(t, err)
assert.GreaterOrEqual(t, elapsed, int64(0))
})
}
}
}
func TestDetectICMPMode(t *testing.T) {
tests := []struct {
name string
family *icmpFamily
rawErr error
udpErr error
want icmpMethod
wantNetworks []string
}{
{
name: "IPv4 prefers raw socket when available",
family: &icmpV4,
want: icmpRaw,
wantNetworks: []string{"ip4:icmp"},
},
{
name: "IPv4 uses datagram when raw unavailable",
family: &icmpV4,
rawErr: errors.New("operation not permitted"),
want: icmpDatagram,
wantNetworks: []string{"ip4:icmp", "udp4"},
},
{
name: "IPv4 falls back to exec when both unavailable",
family: &icmpV4,
rawErr: errors.New("operation not permitted"),
udpErr: errors.New("protocol not supported"),
want: icmpExecFallback,
wantNetworks: []string{"ip4:icmp", "udp4"},
},
{
name: "IPv6 prefers raw socket when available",
family: &icmpV6,
want: icmpRaw,
wantNetworks: []string{"ip6:ipv6-icmp"},
},
{
name: "IPv6 uses datagram when raw unavailable",
family: &icmpV6,
rawErr: errors.New("operation not permitted"),
want: icmpDatagram,
wantNetworks: []string{"ip6:ipv6-icmp", "udp6"},
},
{
name: "IPv6 falls back to exec when both unavailable",
family: &icmpV6,
rawErr: errors.New("operation not permitted"),
udpErr: errors.New("protocol not supported"),
want: icmpExecFallback,
wantNetworks: []string{"ip6:ipv6-icmp", "udp6"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
calls := make([]string, 0, 2)
listen := func(network, listenAddr string) (icmpPacketConn, error) {
require.Equal(t, tt.family.listenAddr, listenAddr)
calls = append(calls, network)
switch network {
case tt.family.rawNetwork:
if tt.rawErr != nil {
return nil, tt.rawErr
}
case tt.family.dgramNetwork:
if tt.udpErr != nil {
return nil, tt.udpErr
}
default:
t.Fatalf("unexpected network %q", network)
}
return testICMPPacketConn{}, nil
}
assert.Equal(t, tt.want, detectICMPMode(tt.family, listen))
assert.Equal(t, tt.wantNetworks, calls)
})
}
}
func TestResolveICMPTarget(t *testing.T) {
t.Run("IPv4 literal", func(t *testing.T) {
family, ip, err := resolveICMPTarget(context.Background(), "127.0.0.1")
require.NoError(t, err)
require.NotNil(t, family)
assert.False(t, family.isIPv6)
assert.Equal(t, "127.0.0.1", ip.String())
})
t.Run("IPv6 literal", func(t *testing.T) {
family, ip, err := resolveICMPTarget(context.Background(), "::1")
require.NoError(t, err)
require.NotNil(t, family)
assert.True(t, family.isIPv6)
assert.Equal(t, "::1", ip.String())
})
t.Run("IPv4-mapped IPv6 resolves as IPv4", func(t *testing.T) {
family, ip, err := resolveICMPTarget(context.Background(), "::ffff:127.0.0.1")
require.NoError(t, err)
require.NotNil(t, family)
assert.False(t, family.isIPv6)
assert.Equal(t, "127.0.0.1", ip.String())
})
}

View File

@@ -1,133 +0,0 @@
package agent
import (
"context"
"errors"
"fmt"
"net"
"net/http"
"time"
"github.com/henrygd/beszel"
"github.com/henrygd/beszel/internal/entities/monitor"
)
const networkMonitorUserAgent = "Beszel-Agent/" + beszel.Version + " (+https://beszel.dev)"
// monitorProbe performs one check. Errors are recorded as loss by the task runner.
// Implementations must honor cancellation and bound their execution time.
type monitorProbe func(context.Context, monitor.Config) (int64, error)
func networkMonitorProbe(client *http.Client) monitorProbe {
return func(ctx context.Context, config monitor.Config) (int64, error) {
switch config.Protocol {
case "icmp":
return monitorICMP(ctx, config.Target)
case "tcp":
return monitorTCP(ctx, config.Target, config.Port)
case "http":
return monitorHTTP(ctx, client, config.Target)
case "dns":
return monitorDNS(ctx, config.Target, config.Server)
default:
return -1, fmt.Errorf("unknown monitor protocol: %s", config.Protocol)
}
}
}
// monitorTCP measures connection establishment time, including address fallback
// but excluding DNS resolution.
// Returns -1 and an error on failure.
func monitorTCP(ctx context.Context, target string, port uint16) (int64, error) {
ctx, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
// Resolve DNS first, outside the timing window but within the probe deadline.
ips, err := net.DefaultResolver.LookupHost(ctx, target)
if err != nil {
return -1, err
}
if len(ips) == 0 {
return -1, errors.New("no addresses resolved for TCP monitor")
}
portString := fmt.Sprintf("%d", port)
deadline, _ := ctx.Deadline()
// Share the remaining probe budget across addresses so an unresponsive
// first address cannot consume all the time available for alternatives.
start := time.Now()
for i, ip := range ips {
if err := ctx.Err(); err != nil {
return -1, err
}
dialer := net.Dialer{Timeout: time.Until(deadline) / time.Duration(len(ips)-i)}
var conn net.Conn
conn, err = dialer.DialContext(ctx, "tcp", net.JoinHostPort(ip, portString))
if err != nil {
continue
}
responseUs := time.Since(start).Microseconds()
conn.Close()
return responseUs, nil
}
return -1, err
}
// monitorDNS measures DNS resolution response time in microseconds. If server is
// non-empty, the lookup is sent to that DNS server (host or host:port, default
// port 53) instead of the system resolver. Returns -1 and an error on failure.
func monitorDNS(ctx context.Context, target, server string) (int64, error) {
ctx, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
resolver := net.DefaultResolver
if server != "" {
resolver = dnsResolverForServer(server)
}
start := time.Now()
ips, err := resolver.LookupHost(ctx, target)
if err != nil || len(ips) == 0 {
return -1, err
}
return time.Since(start).Microseconds(), nil
}
// dnsResolverForServer builds a resolver that sends lookups to the given DNS
// server address instead of the system resolver. server may be a bare host or
// host:port; when no port is given, the standard DNS port 53 is used.
func dnsResolverForServer(server string) *net.Resolver {
address := server
if _, _, err := net.SplitHostPort(server); err != nil {
address = net.JoinHostPort(server, "53")
}
return &net.Resolver{
PreferGo: true,
Dial: func(ctx context.Context, network, _ string) (net.Conn, error) {
var dialer net.Dialer
return dialer.DialContext(ctx, network, address)
},
}
}
// monitorHTTP measures HTTP GET request response in microseconds. Returns -1 and an error on failure.
func monitorHTTP(ctx context.Context, client *http.Client, url string) (int64, error) {
if client == nil {
client = http.DefaultClient
}
start := time.Now()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return -1, err
}
req.Header.Set("User-Agent", networkMonitorUserAgent)
resp, err := client.Do(req)
if err != nil {
return -1, err
}
resp.Body.Close()
if resp.StatusCode >= 400 {
return -1, fmt.Errorf("HTTP error: %s", resp.Status)
}
return time.Since(start).Microseconds(), nil
}

View File

@@ -1,88 +0,0 @@
package agent
import (
"sync"
"time"
)
const (
monitorResumeHeartbeat = 10 * time.Second
// Allow scheduling jitter without mistaking an ordinary tick for resume.
monitorResumeGap = 2 * monitorResumeHeartbeat
monitorResumePause = 10 * time.Second
)
// monitorResumeGuard detects likely suspend/resume using wall time. A long
// process stall or forward clock adjustment can also trigger the bounded pause.
// One heartbeat is shared by all configured monitors.
type monitorResumeGuard struct {
mu sync.Mutex
stop chan struct{}
lastTick time.Time
pauseUntil time.Time
generation uint32
}
func (g *monitorResumeGuard) start() {
g.mu.Lock()
defer g.mu.Unlock()
if g.stop != nil {
return
}
stop := make(chan struct{})
g.stop = stop
g.lastTick = time.Now().Round(0)
g.pauseUntil = time.Time{}
go func() {
ticker := time.NewTicker(monitorResumeHeartbeat)
defer ticker.Stop()
for {
select {
case <-stop:
return
case <-ticker.C:
g.mu.Lock()
if g.stop == stop {
g.observe(time.Now())
}
g.mu.Unlock()
}
}
}()
}
func (g *monitorResumeGuard) shutdown() {
g.mu.Lock()
defer g.mu.Unlock()
if g.stop != nil {
close(g.stop)
g.stop = nil
g.generation++
}
}
// observe requires mu. Strip the monotonic component because it can stop during
// suspend. Read the current time rather than the ticker's queued timestamp.
func (g *monitorResumeGuard) observe(now time.Time) {
now = now.Round(0)
if now.Sub(g.lastTick) > monitorResumeGap {
g.pauseUntil = now.Add(monitorResumePause)
g.generation++
}
g.lastTick = now
}
// snapshot also observes time so a probe waking before the heartbeat detects
// resume itself. A changed generation invalidates probes spanning suspend.
func (g *monitorResumeGuard) snapshot() (generation uint32, allowed bool) {
if g == nil {
return 0, true
}
g.mu.Lock()
defer g.mu.Unlock()
if g.stop == nil {
return g.generation, true
}
g.observe(time.Now())
return g.generation, !g.lastTick.Before(g.pauseUntil)
}

View File

@@ -1,121 +0,0 @@
//go:build testing
package agent
import (
"context"
"errors"
"sync/atomic"
"testing"
"testing/synctest"
"time"
"github.com/henrygd/beszel/internal/entities/monitor"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func simulateMonitorSleep(g *monitorResumeGuard) {
g.mu.Lock()
g.lastTick = time.Now().Add(-time.Hour).Round(0)
g.mu.Unlock()
}
func TestMonitorResumePause(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
var g monitorResumeGuard
g.start()
defer g.shutdown()
generation, allowed := g.snapshot()
require.True(t, allowed)
// Heartbeats alone must keep the guard current between infrequent probes.
time.Sleep(time.Minute)
synctest.Wait()
steadyGeneration, allowed := g.snapshot()
require.True(t, allowed)
require.Equal(t, generation, steadyGeneration)
// The probe, rather than the heartbeat, must detect this gap.
simulateMonitorSleep(&g)
next, allowed := g.snapshot()
assert.False(t, allowed)
assert.NotEqual(t, generation, next)
time.Sleep(9 * time.Second)
_, allowed = g.snapshot()
assert.False(t, allowed)
time.Sleep(time.Second)
end, allowed := g.snapshot()
assert.True(t, allowed)
assert.Equal(t, next, end)
})
}
func TestMonitorResumeGuardLifecycle(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
pm := newMonitorManagerWithProbe(func(context.Context, monitor.Config) (int64, error) { return 1, nil })
defer pm.Stop()
assert.Nil(t, pm.resumeGuard.stop)
pm.SyncMonitors([]monitor.Config{{ID: "a", Interval: 3600}, {ID: "b", Interval: 3600}})
stop := pm.resumeGuard.stop
require.NotNil(t, stop)
pm.DeleteMonitor("a")
assert.Equal(t, stop, pm.resumeGuard.stop)
pm.DeleteMonitor("b")
assert.Nil(t, pm.resumeGuard.stop)
select {
case <-stop:
default:
t.Fatal("heartbeat was not stopped")
}
time.Sleep(time.Hour)
_, err := pm.UpsertMonitor(monitor.Config{ID: "c", Interval: 3600}, false)
require.NoError(t, err)
_, allowed := pm.resumeGuard.snapshot()
assert.True(t, allowed, "idle time must not trigger a resume pause")
pm.SyncMonitors(nil)
assert.Nil(t, pm.resumeGuard.stop)
})
}
func TestMonitorResumeDiscardsInflightProbe(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
var g monitorResumeGuard
g.start()
defer g.shutdown()
task := newMonitorTask(monitor.Config{ID: "test"})
defer task.cancel()
task.resumeGuard = &g
result := task.runProbe(func(context.Context, monitor.Config) (int64, error) {
simulateMonitorSleep(&g)
return 0, errors.New("network not ready")
})
assert.Nil(t, result)
assert.Empty(t, task.history.samples)
// Explicit requests may still run during the pause and record real failures.
result = task.runProbe(func(context.Context, monitor.Config) (int64, error) {
return 0, errors.New("unreachable")
})
require.NotNil(t, result)
assert.Equal(t, 100.0, result.PacketLoss)
})
}
func TestMonitorResumeSkipsScheduledProbes(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
var calls atomic.Int32
pm := newMonitorManagerWithProbe(func(context.Context, monitor.Config) (int64, error) {
calls.Add(1)
return 1, nil
})
defer pm.Stop()
pm.SyncMonitors([]monitor.Config{{ID: "test", Interval: 1}})
simulateMonitorSleep(&pm.resumeGuard)
pm.resumeGuard.snapshot()
time.Sleep(9 * time.Second)
synctest.Wait()
assert.Zero(t, calls.Load())
assert.Empty(t, pm.GetResults(1000))
time.Sleep(2 * time.Second)
synctest.Wait()
assert.Positive(t, calls.Load())
})
}

View File

@@ -1,63 +0,0 @@
package agent
import (
"context"
"log/slog"
"math/rand"
"time"
)
func (pm *MonitorManager) startMonitor(task *monitorTask) {
interval := time.Duration(task.config.Interval) * time.Second
if interval < time.Second {
interval = 30 * time.Second
}
delay := getStagger(interval.Milliseconds())
slog.Debug("starting monitor task", "target", task.config.Target, "delay", delay, "interval", interval)
// Certificate checks piggyback on probe ticks, so they run at most once per
// probe interval after they become due.
go runMonitorSchedule(task.ctx, interval, delay, func() {
if _, allowed := task.resumeGuard.snapshot(); allowed {
task.runProbe(pm.probe)
task.refreshCert(pm.certCheck)
}
})
}
// runMonitorSchedule owns only timing. Checks run serially, and slow checks
// naturally drop missed ticks rather than building an execution backlog.
func runMonitorSchedule(ctx context.Context, interval, delay time.Duration, run func()) {
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-ctx.Done():
return
case <-timer.C:
}
if ctx.Err() != nil {
return
}
run()
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
if ctx.Err() != nil {
return
}
run()
}
}
}
// getStagger returns an initial delay between half an interval and one interval.
func getStagger(intervalMilli int64) time.Duration {
delay := rand.Intn(int(intervalMilli))
if delay < int(intervalMilli)/2 {
delay += int(intervalMilli) / 2
}
return time.Duration(delay) * time.Millisecond
}

View File

@@ -1,167 +0,0 @@
//go:build testing
package agent
import (
"context"
"sync/atomic"
"testing"
"testing/synctest"
"time"
"github.com/henrygd/beszel/internal/entities/monitor"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestMonitorScheduleTiming(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
ctx, cancel := context.WithCancel(t.Context())
defer cancel()
var calls atomic.Int32
go runMonitorSchedule(ctx, 10*time.Second, 5*time.Second, func() { calls.Add(1) })
synctest.Wait()
time.Sleep(4 * time.Second)
synctest.Wait()
assert.Equal(t, 0, int(calls.Load()))
time.Sleep(time.Second)
synctest.Wait()
assert.Equal(t, 1, int(calls.Load()))
time.Sleep(10 * time.Second)
synctest.Wait()
assert.Equal(t, 2, int(calls.Load()))
cancel()
synctest.Wait()
time.Sleep(time.Minute)
synctest.Wait()
assert.Equal(t, 2, int(calls.Load()))
})
}
func TestMonitorScheduleSlowProbe(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
ctx, cancel := context.WithCancel(t.Context())
defer cancel()
var calls atomic.Int32
release := make(chan struct{})
go runMonitorSchedule(ctx, time.Second, 0, func() {
calls.Add(1)
select {
case <-release:
case <-ctx.Done():
}
})
synctest.Wait()
assert.Equal(t, 1, int(calls.Load()))
time.Sleep(time.Minute)
synctest.Wait()
assert.Equal(t, 1, int(calls.Load()), "a slow probe must not spawn overlapping checks")
close(release)
synctest.Wait()
assert.Equal(t, 1, int(calls.Load()), "missed intervals must not accumulate a backlog")
time.Sleep(time.Second)
synctest.Wait()
assert.Equal(t, 2, int(calls.Load()))
cancel()
synctest.Wait()
})
}
func TestMonitorScheduledAndImmediateRequestsShareProbe(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
var calls atomic.Int32
release := make(chan struct{})
cfg := monitor.Config{ID: "test", Interval: 10}
pm := newMonitorManagerWithProbe(func(ctx context.Context, config monitor.Config) (int64, error) {
assert.Equal(t, cfg, config)
calls.Add(1)
<-release
return 42, nil
})
defer pm.Stop()
task := newMonitorTask(cfg)
pm.monitors[cfg.ID] = task
go runMonitorSchedule(task.ctx, 10*time.Second, 0, func() { task.runProbe(pm.probe) })
synctest.Wait()
results := make(chan *monitor.Result, 2)
for range 2 {
go func() {
result, _ := pm.UpsertMonitor(cfg, true)
results <- result
}()
}
synctest.Wait()
assert.Equal(t, 1, int(calls.Load()))
assert.Empty(t, pm.GetResults(1000), "reading history must not wait for network I/O")
close(release)
synctest.Wait()
first, second := <-results, <-results
require.NotNil(t, first)
require.NotNil(t, second)
assert.Equal(t, int64(42), first.AvgResponse)
assert.Equal(t, first, second)
assert.NotSame(t, first, second, "callers must not share mutable result pointers")
assert.Len(t, task.history.samples, 1)
// A later explicit request must still perform a fresh probe.
_, err := pm.UpsertMonitor(cfg, true)
require.NoError(t, err)
assert.Equal(t, 2, int(calls.Load()))
assert.Len(t, task.history.samples, 2)
})
}
func TestMonitorReplacementCancelsSharedProbe(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
cfg := monitor.Config{ID: "test", Interval: 10}
pm := newMonitorManagerWithProbe(func(ctx context.Context, config monitor.Config) (int64, error) {
if config.Interval == 10 {
<-ctx.Done()
return 0, ctx.Err()
}
return 30, nil
})
defer pm.Stop()
task := newMonitorTask(cfg)
task.history.record(monitorSample{responseUs: 10, timestamp: time.Now()})
pm.monitors[cfg.ID] = task
results := make(chan *monitor.Result, 2)
for range 2 {
go func() {
result, _ := pm.UpsertMonitor(cfg, true)
results <- result
}()
}
synctest.Wait()
updated := cfg
updated.Interval = 20
result, err := pm.UpsertMonitor(updated, true)
require.NoError(t, err)
require.NotNil(t, result)
assert.Equal(t, int64(20), result.AvgResponse)
assert.Zero(t, result.PacketLoss)
synctest.Wait()
assert.Nil(t, <-results)
assert.Nil(t, <-results)
assert.Len(t, task.history.samples, 1)
assert.Len(t, pm.monitors[cfg.ID].history.samples, 2)
})
}
func TestMonitorInjectedProbeTimeoutRecordsLoss(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
pm := newMonitorManagerWithProbe(func(ctx context.Context, _ monitor.Config) (int64, error) {
ctx, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
<-ctx.Done()
return 0, ctx.Err()
})
defer pm.Stop()
start := time.Now()
result, err := pm.UpsertMonitor(monitor.Config{ID: "test", Interval: 3600}, true)
require.NoError(t, err)
require.NotNil(t, result)
assert.Equal(t, 3*time.Second, time.Since(start))
assert.Equal(t, 100.0, result.PacketLoss)
assert.NoError(t, pm.monitors["test"].ctx.Err())
})
}

View File

@@ -1,191 +0,0 @@
package agent
import (
"context"
"log/slog"
"sync"
"time"
"github.com/henrygd/beszel/internal/entities/monitor"
)
const monitorFailureLogInterval = 5 * time.Minute
// monitorTask coordinates a probe and its history for one immutable configuration.
type monitorTask struct {
config monitor.Config
ctx context.Context
cancel context.CancelFunc
history *monitorHistory
resumeGuard *monitorResumeGuard
runMu sync.Mutex
inflight *monitorRun
lastFailureLog int64 // Unix nanoseconds
certMu sync.Mutex
cert *monitor.CertInfo
certUnsent bool // cert has not been included in a stats result yet
certChecking bool
nextCertCheck time.Time
}
type monitorRun struct {
done chan struct{}
result *monitor.Result // published by closing done; never mutated afterwards
}
func newMonitorTask(config monitor.Config) *monitorTask {
ctx, cancel := context.WithCancel(context.Background())
task := &monitorTask{config: config, ctx: ctx, history: newMonitorHistory()}
// Serialize cancellation with publication, so canceled probes cannot enter
// history copied into a replacement task.
task.cancel = func() {
task.runMu.Lock()
cancel()
task.runMu.Unlock()
}
return task
}
func newMonitorTaskFromExisting(config monitor.Config, existing *monitorTask) *monitorTask {
task := newMonitorTask(config)
if existing != nil {
task.history = existing.history.clone()
// Keep the last known certificate, but check again soon for the new config.
// The hub already stores it, so it is not marked unsent.
if config.Target == existing.config.Target {
task.cert = existing.certInfo()
}
}
return task
}
// runProbe shares an in-flight check between scheduled and immediate requests.
// Every completed check contributes exactly one sample, regardless of how many
// callers were waiting for it. No task or history lock is held during network I/O.
func (task *monitorTask) runProbe(probe monitorProbe) *monitor.Result {
task.runMu.Lock()
if task.ctx.Err() != nil {
task.runMu.Unlock()
return nil
}
if run := task.inflight; run != nil {
task.runMu.Unlock()
select {
case <-task.ctx.Done():
return nil
case <-run.done:
if task.ctx.Err() != nil {
return nil
}
return copyMonitorResult(run.result)
}
}
run := &monitorRun{done: make(chan struct{})}
task.inflight = run
task.runMu.Unlock()
generation, _ := task.resumeGuard.snapshot()
responseUs, err := probe(task.ctx, task.config)
var logFailure bool
task.runMu.Lock()
currentGeneration, _ := task.resumeGuard.snapshot()
if task.ctx.Err() == nil && generation == currentGeneration {
now := time.Now()
if err != nil {
responseUs = -1
logAt := now.UnixNano()
if task.lastFailureLog == 0 || logAt < task.lastFailureLog || logAt-task.lastFailureLog >= int64(monitorFailureLogInterval) {
logFailure = true
task.lastFailureLog = logAt
}
} else {
task.lastFailureLog = 0
}
result := task.history.record(monitorSample{responseUs: responseUs, timestamp: now})
run.result = &result
}
task.inflight = nil
close(run.done)
task.runMu.Unlock()
if logFailure {
slog.Warn("monitor failed", "err", err, "target", task.config.Target, "protocol", task.config.Protocol)
}
if task.ctx.Err() != nil {
return nil
}
return copyMonitorResult(run.result)
}
// refreshCert checks the certificate of an HTTPS target when due. A failed
// check keeps the last known certificate and retries sooner, as does a
// certificate that expires before the next regular check, so renewals show up
// quickly. Concurrent callers skip rather than wait, and no lock is held during
// network I/O.
func (task *monitorTask) refreshCert(check certChecker) {
if check == nil || !certCheckEnabled(task.config) {
return
}
task.certMu.Lock()
if task.certChecking || time.Now().Before(task.nextCertCheck) {
task.certMu.Unlock()
return
}
task.certChecking = true
task.certMu.Unlock()
info, err := check(task.ctx, task.config.Target)
task.certMu.Lock()
defer task.certMu.Unlock()
task.certChecking = false
if task.ctx.Err() != nil {
return
}
if err != nil {
task.nextCertCheck = time.Now().Add(certCheckRetryInterval)
slog.Warn("certificate check failed", "err", err, "target", task.config.Target)
return
}
task.cert = &info
task.certUnsent = true
now := time.Now()
interval := certCheckInterval
if time.UnixMilli(info.Expires).Before(now.Add(certCheckInterval)) {
interval = certCheckRetryInterval
}
task.nextCertCheck = now.Add(interval)
}
// certInfo returns a copy of the latest certificate info, or nil if unknown.
func (task *monitorTask) certInfo() *monitor.CertInfo {
task.certMu.Lock()
defer task.certMu.Unlock()
if task.cert == nil {
return nil
}
cert := *task.cert
return &cert
}
// takeUnsentCert returns the latest certificate info once after each successful
// check, so unchanged info is not resent with every stats result.
func (task *monitorTask) takeUnsentCert() *monitor.CertInfo {
task.certMu.Lock()
defer task.certMu.Unlock()
if !task.certUnsent {
return nil
}
task.certUnsent = false
cert := *task.cert
return &cert
}
func copyMonitorResult(result *monitor.Result) *monitor.Result {
if result == nil {
return nil
}
copy := *result
return &copy
}

View File

@@ -1,79 +0,0 @@
//go:build testing
package agent
import (
"bytes"
"context"
"errors"
"log/slog"
"testing"
"testing/synctest"
"time"
"github.com/henrygd/beszel/internal/entities/monitor"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestMonitorFailureLogCooldown(t *testing.T) {
var logs bytes.Buffer
previous := slog.Default()
slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil)))
t.Cleanup(func() { slog.SetDefault(previous) })
synctest.Test(t, func(t *testing.T) {
task := newMonitorTask(monitor.Config{ID: "test", Target: "example.test", Protocol: "tcp"})
defer task.cancel()
failure := errors.New("connection refused")
probe := func(context.Context, monitor.Config) (int64, error) { return 42, failure }
var samples int64
check := func(wantLog bool) {
t.Helper()
logs.Reset()
result := task.runProbe(probe)
require.NotNil(t, result)
samples++
assert.Equal(t, samples, result.SampleCount, "suppressed warnings must still record samples")
if !wantLog {
assert.Empty(t, logs.String())
} else {
assert.Contains(t, logs.String(), `msg="monitor failed"`)
assert.Equal(t, 1, bytes.Count(logs.Bytes(), []byte("\n")))
}
}
check(true)
check(false)
time.Sleep(5*time.Minute - time.Nanosecond)
check(false)
time.Sleep(time.Nanosecond)
check(true)
check(false)
time.Sleep(5 * time.Minute)
check(true)
check(false)
// Recovery clears the cooldown.
failure = nil
check(false)
failure = errors.New("connection refused again")
check(true)
// Another monitor has its own cooldown.
other := newMonitorTask(task.config)
defer other.cancel()
logs.Reset()
require.NotNil(t, other.runProbe(probe))
assert.Contains(t, logs.String(), `msg="monitor failed"`)
// A canceled probe must not publish a failure or emit a warning.
logs.Reset()
result := other.runProbe(func(context.Context, monitor.Config) (int64, error) {
other.cancel()
return -1, context.Canceled
})
assert.Nil(t, result)
assert.Empty(t, logs.String())
})
}

View File

@@ -1,590 +0,0 @@
package agent
import (
"context"
"encoding/binary"
"io"
"net"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/henrygd/beszel"
"github.com/henrygd/beszel/internal/entities/monitor"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/net/dns/dnsmessage"
)
func TestMonitorManagerGetResultsIncludesHourResponseRange(t *testing.T) {
now := time.Now().UTC()
task := newMonitorTask(monitor.Config{ID: "monitor-1"})
task.history.addSampleLocked(monitorSample{responseUs: 10, timestamp: now.Add(-30 * time.Minute)})
task.history.addSampleLocked(monitorSample{responseUs: 20, timestamp: now.Add(-9 * time.Minute)})
task.history.addSampleLocked(monitorSample{responseUs: 40, timestamp: now.Add(-5 * time.Minute)})
task.history.addSampleLocked(monitorSample{responseUs: 30, timestamp: now.Add(-50 * time.Second)})
task.history.addSampleLocked(monitorSample{responseUs: -1, timestamp: now.Add(-30 * time.Second)})
pm := newMonitorManager()
pm.monitors = map[string]*monitorTask{"icmp:example.com": task}
results := pm.GetResults(uint16(time.Minute / time.Millisecond))
result, ok := results["monitor-1"]
require.True(t, ok)
assert.Equal(t, int64(30), result.AvgResponse)
assert.Equal(t, int64(25), result.AvgResponse1h)
assert.Equal(t, int64(30), result.MinResponse)
assert.Equal(t, int64(10), result.MinResponse1h)
assert.Equal(t, int64(30), result.MaxResponse)
assert.Equal(t, int64(40), result.MaxResponse1h)
assert.Equal(t, 50.0, result.PacketLoss)
assert.Equal(t, 20.0, result.PacketLoss1h)
}
func TestMonitorManagerGetResultsIncludesLossOnlyHourData(t *testing.T) {
now := time.Now().UTC()
task := newMonitorTask(monitor.Config{ID: "monitor-1"})
task.history.addSampleLocked(monitorSample{responseUs: -1, timestamp: now.Add(-30 * time.Second)})
task.history.addSampleLocked(monitorSample{responseUs: -1, timestamp: now.Add(-10 * time.Second)})
pm := newMonitorManager()
pm.monitors = map[string]*monitorTask{"icmp:example.com": task}
results := pm.GetResults(uint16(time.Minute / time.Millisecond))
result, ok := results["monitor-1"]
require.True(t, ok)
assert.Equal(t, int64(0), result.AvgResponse)
assert.Equal(t, int64(0), result.AvgResponse1h)
assert.Equal(t, int64(0), result.MinResponse)
assert.Equal(t, int64(0), result.MinResponse1h)
assert.Equal(t, int64(0), result.MaxResponse)
assert.Equal(t, int64(0), result.MaxResponse1h)
assert.Equal(t, 100.0, result.PacketLoss)
assert.Equal(t, 100.0, result.PacketLoss1h)
}
func TestMonitorConfigResultKeyUsesSyncedID(t *testing.T) {
cfg := monitor.Config{ID: "monitor-1", Target: "1.1.1.1", Protocol: "icmp", Interval: 10}
assert.Equal(t, "monitor-1", cfg.ID)
}
func TestMonitorManagerSyncMonitorsSkipsConfigsWithoutStableID(t *testing.T) {
validCfg := monitor.Config{ID: "monitor-1", Target: "ignored", Protocol: "noop", Interval: 10}
invalidCfg := monitor.Config{Target: "ignored", Protocol: "noop", Interval: 10}
pm := newMonitorManager()
pm.SyncMonitors([]monitor.Config{validCfg, invalidCfg})
defer pm.Stop()
_, validExists := pm.monitors[validCfg.ID]
_, invalidExists := pm.monitors[invalidCfg.ID]
assert.True(t, validExists)
assert.False(t, invalidExists)
}
func TestMonitorManagerSyncMonitorsStopsRemovedTasksButKeepsExisting(t *testing.T) {
keepCfg := monitor.Config{ID: "monitor-1", Target: "ignored", Protocol: "noop", Interval: 10}
removeCfg := monitor.Config{ID: "monitor-2", Target: "ignored", Protocol: "noop", Interval: 10}
keptTask := newMonitorTask(keepCfg)
removedTask := newMonitorTask(removeCfg)
pm := newMonitorManager()
pm.monitors = map[string]*monitorTask{
keepCfg.ID: keptTask,
removeCfg.ID: removedTask,
}
pm.SyncMonitors([]monitor.Config{keepCfg})
assert.Same(t, keptTask, pm.monitors[keepCfg.ID])
_, exists := pm.monitors[removeCfg.ID]
assert.False(t, exists)
select {
case <-removedTask.ctx.Done():
default:
t.Fatal("expected removed monitor task to be cancelled")
}
select {
case <-keptTask.ctx.Done():
t.Fatal("expected existing monitor task to remain active")
default:
}
}
func TestMonitorManagerSyncMonitorsRestartsChangedConfig(t *testing.T) {
originalCfg := monitor.Config{ID: "monitor-1", Target: "ignored-a", Protocol: "noop", Interval: 10}
updatedCfg := monitor.Config{ID: "monitor-1", Target: "ignored-b", Protocol: "noop", Interval: 10}
originalTask := newMonitorTask(originalCfg)
pm := newMonitorManager()
pm.monitors = map[string]*monitorTask{
originalCfg.ID: originalTask,
}
pm.SyncMonitors([]monitor.Config{updatedCfg})
defer pm.Stop()
restartedTask := pm.monitors[updatedCfg.ID]
assert.NotSame(t, originalTask, restartedTask)
assert.Equal(t, updatedCfg, restartedTask.config)
select {
case <-originalTask.ctx.Done():
default:
t.Fatal("expected changed monitor task to be cancelled")
}
}
func TestMonitorManagerApplySyncUpsertRunsImmediatelyAndReturnsResult(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
defer server.Close()
pm := &MonitorManager{
monitors: make(map[string]*monitorTask),
probe: networkMonitorProbe(server.Client()),
}
resp, err := pm.HandleSyncRequest(monitor.SyncRequest{
Action: monitor.SyncActionUpsert,
Config: monitor.Config{ID: "monitor-1", Target: server.URL, Protocol: "http", Interval: 10},
RunNow: true,
})
defer pm.Stop()
require.NoError(t, err)
assert.GreaterOrEqual(t, resp.Result.AvgResponse, int64(0))
assert.Equal(t, 0.0, resp.Result.PacketLoss)
assert.Equal(t, 0.0, resp.Result.PacketLoss1h)
task := pm.monitors["monitor-1"]
require.NotNil(t, task)
task.history.mu.Lock()
defer task.history.mu.Unlock()
require.Len(t, task.history.samples, 1)
}
func TestMonitorManagerUpsertMonitorKeepsHistoryWhenOnlyIntervalChanges(t *testing.T) {
originalCfg := monitor.Config{ID: "monitor-1", Target: "1.1.1.1", Protocol: "icmp", Interval: 10}
updatedCfg := monitor.Config{ID: "monitor-1", Target: "1.1.1.1", Protocol: "icmp", Interval: 30}
now := time.Now().UTC()
existingTask := newMonitorTask(originalCfg)
existingTask.history.addSampleLocked(monitorSample{responseUs: 12, timestamp: now.Add(-50 * time.Minute)})
existingTask.history.addSampleLocked(monitorSample{responseUs: 24, timestamp: now.Add(-30 * time.Second)})
pm := newMonitorManager()
pm.monitors = map[string]*monitorTask{originalCfg.ID: existingTask}
result, err := pm.UpsertMonitor(updatedCfg, false)
defer pm.Stop()
require.NoError(t, err)
assert.Nil(t, result)
updatedTask := pm.monitors[updatedCfg.ID]
require.NotNil(t, updatedTask)
assert.NotSame(t, existingTask, updatedTask)
assert.Equal(t, updatedCfg, updatedTask.config)
updatedTask.history.mu.Lock()
defer updatedTask.history.mu.Unlock()
require.Len(t, updatedTask.history.samples, 1)
assert.Equal(t, int64(24), updatedTask.history.samples[0].responseUs)
agg := updatedTask.history.aggregateLocked(time.Hour, now)
require.True(t, agg.hasData())
assert.Equal(t, int64(2), agg.totalCount)
assert.Equal(t, int64(2), agg.successCount)
assert.Equal(t, int64(18), agg.avgResponse())
select {
case <-existingTask.ctx.Done():
default:
t.Fatal("expected original monitor task to be cancelled")
}
}
func TestMonitorManagerApplySyncDeleteRemovesTask(t *testing.T) {
config := monitor.Config{ID: "monitor-1", Target: "1.1.1.1", Protocol: "icmp", Interval: 10}
task := newMonitorTask(config)
pm := newMonitorManager()
pm.monitors = map[string]*monitorTask{config.ID: task}
_, err := pm.HandleSyncRequest(monitor.SyncRequest{
Action: monitor.SyncActionDelete,
Config: monitor.Config{ID: config.ID},
})
require.NoError(t, err)
_, exists := pm.monitors[config.ID]
assert.False(t, exists)
select {
case <-task.ctx.Done():
default:
t.Fatal("expected deleted monitor task to be cancelled")
}
}
func TestMonitorManagerGetRandomDelay(t *testing.T) {
for i := 1000; i < 360_000; i += 1000 {
delay := getStagger(int64(i))
assert.GreaterOrEqual(t, delay, time.Duration(i/2)*time.Millisecond)
assert.LessOrEqual(t, delay, time.Duration(i)*time.Millisecond)
}
}
func TestMonitorHTTP(t *testing.T) {
t.Run("success", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, "Beszel-Agent/"+beszel.Version+" (+https://beszel.dev)", r.Header.Get("User-Agent"))
w.WriteHeader(http.StatusNoContent)
}))
defer server.Close()
responseUs, err := monitorHTTP(context.Background(), server.Client(), server.URL)
require.NoError(t, err)
assert.GreaterOrEqual(t, responseUs, int64(0))
})
t.Run("server error", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "boom", http.StatusInternalServerError)
}))
defer server.Close()
responseUs, err := monitorHTTP(context.Background(), server.Client(), server.URL)
assert.Equal(t, int64(-1), responseUs)
require.Error(t, err)
})
}
func TestMonitorTCP(t *testing.T) {
t.Run("success", func(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
accepted := make(chan struct{})
go func() {
defer close(accepted)
conn, err := listener.Accept()
if err == nil {
_ = conn.Close()
}
}()
port := uint16(listener.Addr().(*net.TCPAddr).Port)
responseUs, err := monitorTCP(context.Background(), "127.0.0.1", port)
require.NoError(t, err)
assert.GreaterOrEqual(t, responseUs, int64(0))
<-accepted
})
t.Run("connection failure", func(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
port := uint16(listener.Addr().(*net.TCPAddr).Port)
require.NoError(t, listener.Close())
responseUs, err := monitorTCP(context.Background(), "127.0.0.1", port)
assert.Equal(t, int64(-1), responseUs)
require.Error(t, err)
})
}
func TestMonitorTCPAddressFallback(t *testing.T) {
for _, tc := range []struct {
name string
ips []string
loss bool
}{
{"first address fails", []string{"127.0.0.2", "127.0.0.1"}, false},
{"first address succeeds", []string{"127.0.0.1", "127.0.0.2"}, false},
{"all addresses fail", []string{"127.0.0.2", "127.0.0.3"}, true},
} {
t.Run(tc.name, func(t *testing.T) {
listener, err := net.Listen("tcp4", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
original := net.DefaultResolver
net.DefaultResolver = tcpMonitorTestResolver(tc.ips)
defer func() { net.DefaultResolver = original }()
// Verify the resolver preserves the intended order, so success cannot
// accidentally bypass the failed first address in the regression case.
ips, err := net.DefaultResolver.LookupHost(t.Context(), "tcp-monitor.invalid.")
require.NoError(t, err)
require.Equal(t, tc.ips, ips)
responseUs, err := monitorTCP(t.Context(), "tcp-monitor.invalid.", uint16(listener.Addr().(*net.TCPAddr).Port))
if tc.loss {
require.Error(t, err)
assert.Equal(t, int64(-1), responseUs)
} else {
require.NoError(t, err)
assert.GreaterOrEqual(t, responseUs, int64(0))
}
})
}
}
// tcpMonitorTestResolver supplies multiple A records without external DNS.
func tcpMonitorTestResolver(ips []string) *net.Resolver {
return &net.Resolver{PreferGo: true, Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
client, server := net.Pipe()
go func() {
defer server.Close()
// net.Resolver uses TCP framing when its connection is not a PacketConn.
var size uint16
if err := binary.Read(server, binary.BigEndian, &size); err != nil {
return
}
packet := make([]byte, size)
if _, err := io.ReadFull(server, packet); err != nil {
return
}
var msg dnsmessage.Message
if err := msg.Unpack(packet); err != nil {
return
}
msg.Header.Response = true
msg.Header.RecursionAvailable = true
for _, question := range msg.Questions {
if question.Type != dnsmessage.TypeA {
continue
}
for _, ip := range ips {
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{Name: question.Name, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET},
Body: &dnsmessage.AResource{A: [4]byte(net.ParseIP(ip).To4())},
})
}
}
packet, err := msg.Pack()
if err != nil {
return
}
response := binary.BigEndian.AppendUint16(nil, uint16(len(packet)))
_, _ = server.Write(append(response, packet...))
}()
return client, nil
}}
}
// udpDNSTestServer starts a UDP server on loopback that answers A queries with the
// given IPs, and returns its listen address (host:port).
func udpDNSTestServer(t *testing.T, ips []string) string {
t.Helper()
conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
require.NoError(t, err)
t.Cleanup(func() { conn.Close() })
go func() {
buf := make([]byte, 512)
for {
n, addr, err := conn.ReadFromUDP(buf)
if err != nil {
return
}
var msg dnsmessage.Message
if err := msg.Unpack(buf[:n]); err != nil {
continue
}
msg.Header.Response = true
msg.Header.RecursionAvailable = true
for _, question := range msg.Questions {
if question.Type != dnsmessage.TypeA {
continue
}
for _, ip := range ips {
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{Name: question.Name, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET},
Body: &dnsmessage.AResource{A: [4]byte(net.ParseIP(ip).To4())},
})
}
}
packet, err := msg.Pack()
if err != nil {
continue
}
_, _ = conn.WriteToUDP(packet, addr)
}
}()
return conn.LocalAddr().String()
}
func TestMonitorDNS(t *testing.T) {
t.Run("success", func(t *testing.T) {
responseUs, err := monitorDNS(context.Background(), "localhost", "")
require.NoError(t, err)
assert.GreaterOrEqual(t, responseUs, int64(0))
})
t.Run("lookup failure", func(t *testing.T) {
responseUs, err := monitorDNS(context.Background(), "", "")
assert.Equal(t, int64(-1), responseUs)
require.Error(t, err)
})
t.Run("custom server", func(t *testing.T) {
serverAddr := udpDNSTestServer(t, []string{"192.0.2.10"})
responseUs, err := monitorDNS(context.Background(), "example.test.", serverAddr)
require.NoError(t, err)
assert.GreaterOrEqual(t, responseUs, int64(0))
})
t.Run("custom server without port defaults to 53", func(t *testing.T) {
resolver := dnsResolverForServer("127.0.0.1")
conn, err := resolver.Dial(context.Background(), "udp", "")
require.NoError(t, err)
defer conn.Close()
assert.Equal(t, "127.0.0.1:53", conn.RemoteAddr().String())
})
t.Run("custom server unreachable", func(t *testing.T) {
responseUs, err := monitorDNS(context.Background(), "example.test.", "127.0.0.1:1")
assert.Equal(t, int64(-1), responseUs)
require.Error(t, err)
})
}
func TestMonitorManagerCancelsActiveProbe(t *testing.T) {
for _, action := range []string{"stop", "delete", "upsert", "sync replace", "sync remove"} {
t.Run(action, func(t *testing.T) {
started := make(chan struct{})
canceled := make(chan struct{})
release := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
close(started)
select {
case <-r.Context().Done():
close(canceled)
case <-release:
}
}))
defer server.Close()
defer close(release)
pm := newMonitorManager()
defer pm.Stop()
cfg := monitor.Config{ID: "test", Protocol: "http", Target: server.URL, Interval: 3600}
task := newMonitorTask(cfg)
// Seed history to ensure a canceled RunNow does not return an old result.
task.history.addSampleLocked(monitorSample{responseUs: 123, timestamp: time.Now()})
pm.monitors[cfg.ID] = task
done := make(chan *monitor.Result, 1)
go func() {
result, _ := pm.UpsertMonitor(cfg, true)
done <- result
}()
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("probe did not start")
}
updated := cfg
updated.Interval--
switch action {
case "stop":
pm.Stop()
case "delete":
pm.DeleteMonitor(cfg.ID)
case "upsert":
_, err := pm.UpsertMonitor(updated, false)
require.NoError(t, err)
case "sync replace":
pm.SyncMonitors([]monitor.Config{updated})
case "sync remove":
pm.SyncMonitors(nil)
}
select {
case <-canceled:
case <-time.After(time.Second):
t.Fatal("active HTTP request was not canceled")
}
select {
case result := <-done:
assert.Nil(t, result)
case <-time.After(time.Second):
t.Fatal("RunNow did not return after cancellation")
}
task.history.mu.Lock()
assert.Len(t, task.history.samples, 1, "cancellation must not record packet loss")
task.history.mu.Unlock()
})
}
}
func TestMonitorResolutionCancellation(t *testing.T) {
for _, protocol := range []string{"tcp", "dns", "icmp"} {
t.Run(protocol, func(t *testing.T) {
started := make(chan struct{}, 1)
original := net.DefaultResolver
net.DefaultResolver = &net.Resolver{PreferGo: true, Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
select {
case started <- struct{}{}:
default:
}
<-ctx.Done()
return nil, ctx.Err()
}}
defer func() { net.DefaultResolver = original }()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
done := make(chan error, 1)
go func() {
var err error
switch protocol {
case "tcp":
_, err = monitorTCP(ctx, "monitor-cancellation.invalid.", 80)
case "dns":
_, err = monitorDNS(ctx, "monitor-cancellation.invalid.", "")
case "icmp":
_, err = monitorICMP(ctx, "monitor-cancellation.invalid.")
}
done <- err
}()
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("lookup did not start")
}
cancel()
select {
case err := <-done:
require.Error(t, err)
case <-time.After(time.Second):
t.Fatal("lookup did not cancel")
}
})
}
}
func TestMonitorProbeTimeoutRecordsLoss(t *testing.T) {
release := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
select {
case <-r.Context().Done():
case <-release:
}
}))
defer server.Close()
defer close(release)
pm := newMonitorManager()
pm.probe = networkMonitorProbe(&http.Client{Timeout: 20 * time.Millisecond})
task := newMonitorTask(monitor.Config{ID: "timeout", Protocol: "http", Target: server.URL})
defer task.cancel()
result := task.runProbe(pm.probe)
require.NotNil(t, result)
assert.Equal(t, 100.0, result.PacketLoss)
assert.Equal(t, 100.0, result.PacketLoss1h)
require.Len(t, task.history.samples, 1)
assert.Equal(t, int64(-1), task.history.samples[0].responseUs)
assert.NoError(t, task.ctx.Err(), "a probe timeout must not cancel the task")
}

View File

@@ -344,7 +344,7 @@ func TestComputeBytesPerSecond(t *testing.T) {
func TestSumAndTrackPerNicDeltas(t *testing.T) {
a := &Agent{
netInterfaces: map[string]bool{"eth0": false, "wlan0": false},
netInterfaces: map[string]struct{}{"eth0": {}, "wlan0": {}},
netInterfaceDeltaTrackers: make(map[uint16]*deltatracker.DeltaTracker[string, uint64]),
}
@@ -373,7 +373,7 @@ func TestSumAndTrackPerNicDeltas(t *testing.T) {
func TestSumAndTrackPerNicDeltasHandlesCounterReset(t *testing.T) {
a := &Agent{
netInterfaces: map[string]bool{"eth0": false},
netInterfaces: map[string]struct{}{"eth0": {}},
netInterfaceDeltaTrackers: make(map[uint16]*deltatracker.DeltaTracker[string, uint64]),
}
@@ -469,7 +469,7 @@ func TestApplyNetworkTotals(t *testing.T) {
t.Run(tt.name, func(t *testing.T) {
// Setup agent with initialized maps
a := &Agent{
netInterfaces: make(map[string]bool),
netInterfaces: make(map[string]struct{}),
netIoStats: make(map[uint16]system.NetIoStats),
netInterfaceDeltaTrackers: make(map[uint16]*deltatracker.DeltaTracker[string, uint64]),
}
@@ -511,22 +511,3 @@ func TestApplyNetworkTotals(t *testing.T) {
})
}
}
func TestSumAndTrackPerNicDeltasKeepsCountersWhenMacReadFails(t *testing.T) {
a := &Agent{
// missing0 is flagged for MAC counters but has no ethtool stats to read
netInterfaces: map[string]bool{"eth0": false, "missing0": true},
netInterfaceDeltaTrackers: make(map[uint16]*deltatracker.DeltaTracker[string, uint64]),
}
netIO := []psutilNet.IOCountersStat{
{Name: "eth0", BytesSent: 100, BytesRecv: 200},
{Name: "missing0", BytesSent: 300, BytesRecv: 400},
}
stats := &system.Stats{}
a.ensureNetworkInterfacesMap(stats)
tx, rx := a.sumAndTrackPerNicDeltas(1, 0, netIO, stats)
assert.Equal(t, uint64(400), tx)
assert.Equal(t, uint64(600), rx)
assert.Equal(t, [4]uint64{0, 0, 300, 400}, stats.NetworkInterfaces["missing0"])
}

View File

@@ -1,508 +0,0 @@
package agent
import (
"bufio"
"context"
"errors"
"log/slog"
"os"
"os/exec"
"path/filepath"
"runtime"
"slices"
"strings"
"sync"
"time"
"github.com/henrygd/beszel/agent/utils"
"github.com/henrygd/beszel/internal/entities/system"
)
const (
defaultPackageUpdatesInterval = time.Hour
packageUpdatesTimeout = 5 * time.Minute
// pacmanSyncInterval limits how often checkupdates downloads fresh sync
// databases. Checks in between reuse the last synced copy.
pacmanSyncInterval = 12 * time.Hour
)
// packageUpdatesResult is the outcome of one package manager check.
type packageUpdatesResult struct {
// counts is [total] or [total, security] pending package updates.
counts []uint16
packages []system.PackageUpdate
// securityKnown is true if packages carry per-package security flags.
securityKnown bool
}
type packageUpdatesCheck func(ctx context.Context) (packageUpdatesResult, error)
// packageUpdatesManager periodically checks the host package manager for pending
// updates in the background and caches the result, so checks never delay metrics.
type packageUpdatesManager struct {
sync.Mutex
name string
check packageUpdatesCheck
interval time.Duration
result packageUpdatesResult
checkedAt time.Time
running bool
}
// newPackageUpdatesManager returns nil if disabled or no supported package manager
// is found. Agents running in a container are skipped because the container's
// package database is not the host's. dataDir holds pacman's private sync databases.
func newPackageUpdatesManager(dataDir string) *packageUpdatesManager {
if runtime.GOOS != "linux" || runningInContainer() {
return nil
}
interval, enabled := packageUpdatesInterval()
if !enabled {
return nil
}
name, check := detectPackageManager(dataDir)
if check == nil {
return nil
}
slog.Debug("Package updates", "manager", name, "interval", interval)
return &packageUpdatesManager{name: name, check: check, interval: interval}
}
// packageUpdatesInterval reads PACKAGE_UPDATES_INTERVAL as a Go duration such as
// "30m" or "6h". "0" disables checks. Invalid or negative values keep the default.
func packageUpdatesInterval() (interval time.Duration, enabled bool) {
env, exists := utils.GetEnv("PACKAGE_UPDATES_INTERVAL")
if !exists {
return defaultPackageUpdatesInterval, true
}
duration, err := time.ParseDuration(env)
switch {
case err == nil && duration == 0:
slog.Info("PACKAGE_UPDATES_INTERVAL", "duration", "disabled")
return 0, false
case err == nil && duration > 0:
slog.Info("PACKAGE_UPDATES_INTERVAL", "duration", duration)
return duration, true
default:
slog.Warn("Invalid PACKAGE_UPDATES_INTERVAL", "value", env)
return defaultPackageUpdatesInterval, true
}
}
// get returns the last cached counts and starts a background check if they are stale.
func (pm *packageUpdatesManager) get(now time.Time) []uint16 {
pm.Lock()
defer pm.Unlock()
if !pm.running && (pm.checkedAt.IsZero() || now.Sub(pm.checkedAt) >= pm.interval) {
pm.running = true
go pm.refresh()
}
return pm.result.counts
}
// list returns the per-package details of the last check. It never starts a check.
func (pm *packageUpdatesManager) list() system.PackageUpdates {
pm.Lock()
defer pm.Unlock()
data := system.PackageUpdates{
Manager: pm.name,
SecurityKnown: pm.result.securityKnown,
Packages: pm.result.packages,
}
if !pm.checkedAt.IsZero() {
data.CheckedAt = pm.checkedAt.Unix()
}
return data
}
func (pm *packageUpdatesManager) refresh() {
ctx, cancel := context.WithTimeout(context.Background(), packageUpdatesTimeout)
defer cancel()
result, err := pm.check(ctx)
if err != nil {
slog.Debug("Package updates check failed", "err", err)
result = packageUpdatesResult{}
}
pm.Lock()
pm.result = result
pm.checkedAt = time.Now()
pm.running = false
pm.Unlock()
}
func runningInContainer() bool {
for _, path := range []string{"/.dockerenv", "/run/.containerenv"} {
if _, err := os.Stat(path); err == nil {
return true
}
}
return false
}
func detectPackageManager(dataDir string) (string, packageUpdatesCheck) {
switch {
case commandExists("apt-get"):
return "apt", checkApt
case commandExists("dnf"):
return "dnf", checkDnf
case commandExists("zypper"):
return "zypper", checkZypper
case commandExists("checkupdates"):
return "pacman", newPacmanCheck(dataDir)
case commandExists("apk"):
return "apk", checkApk
}
return "", nil
}
func commandExists(name string) bool {
_, err := exec.LookPath(name)
return err == nil
}
// runPackageCommand runs a read-only package manager command and returns stdout.
// okCodes lists non-zero exit codes that still mean success.
func runPackageCommand(ctx context.Context, okCodes []int, name string, args ...string) (string, error) {
return runPackageCommandEnv(ctx, nil, okCodes, name, args...)
}
// runPackageCommandEnv is runPackageCommand with extra environment variables.
func runPackageCommandEnv(ctx context.Context, env []string, okCodes []int, name string, args ...string) (string, error) {
cmd := exec.CommandContext(ctx, name, args...)
cmd.Env = append(os.Environ(), "LC_ALL=C")
cmd.Env = append(cmd.Env, env...)
// checkupdates is a shell script, so a timeout kills only the script and its
// children can keep stdout open. WaitDelay stops Output from waiting on them.
cmd.WaitDelay = 10 * time.Second
out, err := cmd.Output()
if exitErr, ok := errors.AsType[*exec.ExitError](err); ok && slices.Contains(okCodes, exitErr.ExitCode()) {
return string(out), nil
}
return string(out), err
}
// countSecurity returns the number of packages flagged as security updates.
func countSecurity(packages []system.PackageUpdate) (count uint16) {
for _, pkg := range packages {
if pkg.Security {
count++
}
}
return count
}
// checkApt simulates a full upgrade against the current package lists.
// It never refreshes the lists; apt-daily or the user does that.
func checkApt(ctx context.Context) (packageUpdatesResult, error) {
out, err := runPackageCommand(ctx, nil, "apt-get", "-s", "dist-upgrade")
if err != nil {
return packageUpdatesResult{}, err
}
packages := parseAptSimulate(out)
return packageUpdatesResult{
counts: []uint16{uint16(len(packages)), countSecurity(packages)},
packages: packages,
securityKnown: true,
}, nil
}
// checkDnf uses the system metadata cache only (-C), so it never downloads metadata.
// check-update lists only available versions, so installed versions come from rpm.
func checkDnf(ctx context.Context) (packageUpdatesResult, error) {
out, err := runPackageCommand(ctx, []int{100}, "dnf", "-q", "-C", "check-update")
if err != nil {
return packageUpdatesResult{}, err
}
packages := parseDnfCheckUpdate(out)
result := packageUpdatesResult{packages: packages}
if len(packages) > 0 {
args := []string{"-q", "--qf", rpmInstalledQueryFormat}
for _, pkg := range packages {
args = append(args, pkg.Name)
}
// rpm exits non-zero if any package is not installed; keep what it printed
out, _ = runPackageCommand(ctx, nil, "rpm", args...)
installed := parseRpmInstalled(out)
for i := range packages {
packages[i].Current = installed[packages[i].Name]
}
}
out, err = runPackageCommand(ctx, []int{100}, "dnf", "-q", "-C", "check-update", "--security")
if err == nil {
// --security lists the lowest version that fixes an advisory, which may be
// older than the version check-update offers, so match on name.arch only
security := make(map[string]struct{})
for _, pkg := range parseDnfCheckUpdate(out) {
security[pkg.Name] = struct{}{}
}
for i := range packages {
_, packages[i].Security = security[packages[i].Name]
}
result.securityKnown = true
}
for i := range packages {
packages[i].Name = trimRpmArch(packages[i].Name)
}
result.counts = []uint16{uint16(len(packages))}
if result.securityKnown {
result.counts = append(result.counts, countSecurity(packages))
}
return result, nil
}
// checkZypper lists package updates. Security updates come from patches, which
// zypper does not map to packages here, so only the security count is known.
func checkZypper(ctx context.Context) (packageUpdatesResult, error) {
out, err := runPackageCommand(ctx, nil, "zypper", "--no-refresh", "-q", "list-updates")
if err != nil {
return packageUpdatesResult{}, err
}
packages := parseZypperListUpdates(out)
result := packageUpdatesResult{packages: packages, counts: []uint16{uint16(len(packages))}}
out, err = runPackageCommand(ctx, nil, "zypper", "--no-refresh", "-q", "list-patches", "--category", "security")
if err == nil {
result.counts = append(result.counts, parseZypperTable(out))
}
return result, nil
}
// newPacmanCheck uses checkupdates (pacman-contrib), which syncs a private copy of
// the databases and never touches pacman's own. The copy lives in dataDir because
// the systemd unit's ProtectSystem=strict makes the default /tmp location read-only.
// It syncs every pacmanSyncInterval and uses the existing copy (-n) in between.
// Local upgrades show up right away since checkupdates links the live local DB.
// Exit code 2 means no updates.
func newPacmanCheck(dataDir string) packageUpdatesCheck {
var env []string
var syncDir string
if dataDir != "" {
dbPath := filepath.Join(dataDir, "checkup-db")
env = []string{"CHECKUPDATES_DB=" + dbPath}
syncDir = filepath.Join(dbPath, "sync")
}
// checks never overlap (packageUpdatesManager.running), so no lock is needed
var lastSync time.Time
return func(ctx context.Context) (packageUpdatesResult, error) {
// -n with a missing database reports no updates rather than failing,
// so always sync first and whenever the private copy is missing
sync := lastSync.IsZero() || time.Since(lastSync) >= pacmanSyncInterval
if !sync && syncDir != "" {
if _, err := os.Stat(syncDir); err != nil {
sync = true
}
}
var args []string
if !sync {
args = append(args, "-n")
}
out, err := runPackageCommandEnv(ctx, env, []int{2}, "checkupdates", args...)
if err != nil {
return packageUpdatesResult{}, err
}
if sync {
lastSync = time.Now()
}
packages := parsePacmanCheckUpdates(out)
return packageUpdatesResult{counts: []uint16{uint16(len(packages))}, packages: packages}, nil
}
}
func checkApk(ctx context.Context) (packageUpdatesResult, error) {
out, err := runPackageCommand(ctx, nil, "apk", "--no-network", "-u", "list")
if err != nil {
return packageUpdatesResult{}, err
}
packages := parseApkUpgradable(out)
return packageUpdatesResult{counts: []uint16{uint16(len(packages))}, packages: packages}, nil
}
// parseAptSimulate parses upgrades in `apt-get -s` output. Upgrade lines look like
// "Inst libc6 [2.35-0ubuntu3.4] (2.35-0ubuntu3.15 Ubuntu:22.04/jammy-updates, Ubuntu:22.04/jammy-security [arm64])".
// New dependencies have no "[old version]" and are skipped.
func parseAptSimulate(out string) (packages []system.PackageUpdate) {
scanner := bufio.NewScanner(strings.NewReader(out))
for scanner.Scan() {
line := scanner.Text()
fields := strings.Fields(line)
if len(fields) < 4 || fields[0] != "Inst" || !strings.HasPrefix(fields[2], "[") || !strings.HasPrefix(fields[3], "(") {
continue
}
pkg := system.PackageUpdate{
Name: fields[1],
Current: strings.Trim(fields[2], "[]"),
Available: strings.TrimPrefix(fields[3], "("),
}
start := strings.IndexByte(line, '(')
end := strings.IndexByte(line, ')')
pkg.Security = start >= 0 && end > start && strings.Contains(line[start:end], "-security")
packages = append(packages, pkg)
}
return packages
}
// parseDnfCheckUpdate parses "name.arch version repo" lines, stopping at the
// obsoletes section so obsoleted packages are not listed twice. Names keep the
// arch so they can be matched with rpm output. dnf4 wraps a long name.arch onto
// its own line, with the version and repo on the next line.
func parseDnfCheckUpdate(out string) (packages []system.PackageUpdate) {
var wrappedName string
scanner := bufio.NewScanner(strings.NewReader(out))
for scanner.Scan() {
line := scanner.Text()
if strings.HasPrefix(line, "Obsoleting") {
break
}
fields := strings.Fields(line)
if wrappedName != "" && len(fields) == 2 {
fields = []string{wrappedName, fields[0], fields[1]}
}
wrappedName = ""
switch {
case len(fields) == 3 && strings.Contains(fields[0], "."):
packages = append(packages, system.PackageUpdate{Name: fields[0], Available: fields[1]})
case len(fields) == 1 && strings.Contains(fields[0], ".") && !strings.HasPrefix(line, " "):
wrappedName = fields[0]
}
}
return packages
}
// rpmInstalledQueryFormat prints "name.arch [epoch:]version-release", matching
// the version format of dnf check-update.
const rpmInstalledQueryFormat = `%{NAME}.%{ARCH} %|EPOCH?{%{EPOCH}:}:{}|%{VERSION}-%{RELEASE}\n`
// parseRpmInstalled maps name.arch to its installed version. For packages with
// several installed versions, such as kernels, the last one listed wins.
func parseRpmInstalled(out string) map[string]string {
installed := make(map[string]string)
scanner := bufio.NewScanner(strings.NewReader(out))
for scanner.Scan() {
// "package foo.x86_64 is not installed" has more than two fields
if fields := strings.Fields(scanner.Text()); len(fields) == 2 {
installed[fields[0]] = fields[1]
}
}
return installed
}
// trimRpmArch removes the ".arch" suffix from a dnf package name.
func trimRpmArch(name string) string {
if i := strings.LastIndexByte(name, '.'); i > 0 {
return name[:i]
}
return name
}
// parseZypperTable counts the data rows of a zypper table (the lines after the
// "---+---" separator).
func parseZypperTable(out string) (count uint16) {
inTable := false
scanner := bufio.NewScanner(strings.NewReader(out))
for scanner.Scan() {
line := scanner.Text()
switch {
case !inTable:
inTable = strings.HasPrefix(line, "--") && strings.Contains(line, "-+-")
case strings.Contains(line, "|"):
count++
default:
return count
}
}
return count
}
// parseZypperListUpdates parses the `zypper list-updates` table, locating the
// columns by their header names.
func parseZypperListUpdates(out string) (packages []system.PackageUpdate) {
nameCol, currentCol, availableCol := -1, -1, -1
var header []string
inTable := false
scanner := bufio.NewScanner(strings.NewReader(out))
for scanner.Scan() {
line := scanner.Text()
switch {
case !inTable && strings.HasPrefix(line, "--") && strings.Contains(line, "-+-"):
for i, col := range header {
switch strings.TrimSpace(col) {
case "Name":
nameCol = i
case "Current Version":
currentCol = i
case "Available Version":
availableCol = i
}
}
if nameCol < 0 || availableCol < 0 {
return nil
}
inTable = true
case !inTable:
header = strings.Split(line, "|")
case strings.Contains(line, "|"):
cols := strings.Split(line, "|")
if len(cols) != len(header) {
continue
}
pkg := system.PackageUpdate{
Name: strings.TrimSpace(cols[nameCol]),
Available: strings.TrimSpace(cols[availableCol]),
}
if currentCol >= 0 {
pkg.Current = strings.TrimSpace(cols[currentCol])
}
packages = append(packages, pkg)
default:
return packages
}
}
return packages
}
// parsePacmanCheckUpdates parses "name old -> new" lines.
func parsePacmanCheckUpdates(out string) (packages []system.PackageUpdate) {
scanner := bufio.NewScanner(strings.NewReader(out))
for scanner.Scan() {
fields := strings.Fields(scanner.Text())
if len(fields) >= 4 && fields[2] == "->" {
packages = append(packages, system.PackageUpdate{Name: fields[0], Current: fields[1], Available: fields[3]})
}
}
return packages
}
// parseApkUpgradable parses lines of `apk -u list`, which look like
// "musl-1.2.5-r3 aarch64 {musl} (MIT) [upgradable from: musl-1.2.5-r0]".
func parseApkUpgradable(out string) (packages []system.PackageUpdate) {
const marker = "[upgradable from:"
scanner := bufio.NewScanner(strings.NewReader(out))
for scanner.Scan() {
line := scanner.Text()
i := strings.Index(line, marker)
fields := strings.Fields(line)
if i < 0 || len(fields) == 0 {
continue
}
name, available := splitApkNameVersion(fields[0])
_, current := splitApkNameVersion(strings.TrimSuffix(strings.TrimSpace(line[i+len(marker):]), "]"))
packages = append(packages, system.PackageUpdate{Name: name, Current: current, Available: available})
}
return packages
}
// splitApkNameVersion splits "name-version-rN" into name and "version-rN".
// Names may contain dashes, but versions do not.
func splitApkNameVersion(s string) (name, version string) {
rel := strings.LastIndexByte(s, '-')
if rel <= 0 || !strings.HasPrefix(s[rel+1:], "r") {
return s, ""
}
ver := strings.LastIndexByte(s[:rel], '-')
if ver <= 0 {
return s, ""
}
return s[:ver], s[ver+1:]
}

View File

@@ -1,445 +0,0 @@
//go:build testing
package agent
import (
"context"
"errors"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"time"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func readPackageUpdatesTestData(t *testing.T, name string) string {
t.Helper()
data, err := os.ReadFile(filepath.Join("test-data", "package_updates", name))
require.NoError(t, err)
return string(data)
}
// Test data files are real command outputs captured in containers.
// findPackage returns the named package from a parsed list.
func findPackage(t *testing.T, packages []system.PackageUpdate, name string) system.PackageUpdate {
t.Helper()
for _, pkg := range packages {
if pkg.Name == name {
return pkg
}
}
t.Fatalf("package %q not found", name)
return system.PackageUpdate{}
}
// fakeCommands puts shell scripts named after package manager commands first on PATH.
func fakeCommands(t *testing.T, scripts map[string]string) {
t.Helper()
if runtime.GOOS == "windows" {
t.Skip("requires shell scripts on PATH")
}
binDir := t.TempDir()
for name, script := range scripts {
require.NoError(t, os.WriteFile(filepath.Join(binDir, name), []byte("#!/bin/sh\n"+script), 0o755))
}
t.Setenv("PATH", binDir+string(os.PathListSeparator)+os.Getenv("PATH"))
}
func testDataPath(t *testing.T, name string) string {
t.Helper()
path, err := filepath.Abs(filepath.Join("test-data", "package_updates", name))
require.NoError(t, err)
return path
}
func TestPackageUpdatesInterval(t *testing.T) {
tests := []struct {
name string
value *string
interval time.Duration
enabled bool
}{
{"unset", nil, time.Hour, true},
{"duration", new("30m"), 30 * time.Minute, true},
{"compound duration", new("1h30m"), 90 * time.Minute, true},
{"zero disables", new("0"), 0, false},
{"zero with unit disables", new("0s"), 0, false},
{"negative keeps default", new("-5m"), time.Hour, true},
{"no unit keeps default", new("60"), time.Hour, true},
{"invalid keeps default", new("hourly"), time.Hour, true},
{"empty keeps default", new(""), time.Hour, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Setenv("BESZEL_AGENT_PACKAGE_UPDATES_INTERVAL", "")
require.NoError(t, os.Unsetenv("BESZEL_AGENT_PACKAGE_UPDATES_INTERVAL"))
t.Setenv("PACKAGE_UPDATES_INTERVAL", "")
require.NoError(t, os.Unsetenv("PACKAGE_UPDATES_INTERVAL"))
if tt.value != nil {
t.Setenv("PACKAGE_UPDATES_INTERVAL", *tt.value)
}
interval, enabled := packageUpdatesInterval()
assert.Equal(t, tt.interval, interval)
assert.Equal(t, tt.enabled, enabled)
})
}
t.Run("prefixed variable takes precedence", func(t *testing.T) {
t.Setenv("PACKAGE_UPDATES_INTERVAL", "0")
t.Setenv("BESZEL_AGENT_PACKAGE_UPDATES_INTERVAL", "6h")
interval, enabled := packageUpdatesInterval()
assert.Equal(t, 6*time.Hour, interval)
assert.True(t, enabled)
})
}
func TestParseAptSimulate(t *testing.T) {
tests := []struct {
file string
total, security int
}{
{"apt_debian12.txt", 44, 5},
{"apt_ubuntu2204.txt", 58, 45},
}
for _, tt := range tests {
t.Run(tt.file, func(t *testing.T) {
packages := parseAptSimulate(readPackageUpdatesTestData(t, tt.file))
assert.Len(t, packages, tt.total)
assert.EqualValues(t, tt.security, countSecurity(packages))
})
}
t.Run("versions", func(t *testing.T) {
packages := parseAptSimulate(readPackageUpdatesTestData(t, "apt_ubuntu2204.txt"))
assert.Equal(t, system.PackageUpdate{Name: "libc6", Current: "2.35-0ubuntu3.4", Available: "2.35-0ubuntu3.15", Security: true}, findPackage(t, packages, "libc6"))
assert.Equal(t, system.PackageUpdate{Name: "base-files", Current: "12ubuntu4.4", Available: "12ubuntu4.7"}, findPackage(t, packages, "base-files"))
packages = parseAptSimulate(readPackageUpdatesTestData(t, "apt_debian12.txt"))
assert.Equal(t, system.PackageUpdate{Name: "tzdata", Current: "2023c-5+deb12u1", Available: "2026b-0+deb12u1"}, findPackage(t, packages, "tzdata"))
})
t.Run("new dependencies and trailing brackets", func(t *testing.T) {
out := `Inst linux-image-6.8.0-50-generic (6.8.0-50.51 Ubuntu:24.04/noble-updates, Ubuntu:24.04/noble-security [amd64])
Inst linux-image-generic [6.8.0-49.49] (6.8.0-50.50 Ubuntu:24.04/noble-updates, Ubuntu:24.04/noble-security [amd64])
Inst gcc-12-base [12.3.0-1ubuntu1~22.04] (12.3.0-1ubuntu1~22.04.3 Ubuntu:22.04/jammy-updates [arm64]) [libstdc++6:arm64 libgcc-s1:arm64 ]
Conf linux-image-generic (6.8.0-50.50 Ubuntu:24.04/noble-updates, Ubuntu:24.04/noble-security [amd64])
Remv oldpkg [1.0]`
assert.Equal(t, []system.PackageUpdate{
{Name: "linux-image-generic", Current: "6.8.0-49.49", Available: "6.8.0-50.50", Security: true},
{Name: "gcc-12-base", Current: "12.3.0-1ubuntu1~22.04", Available: "12.3.0-1ubuntu1~22.04.3"},
}, parseAptSimulate(out))
})
t.Run("no updates", func(t *testing.T) {
assert.Empty(t, parseAptSimulate("Reading package lists...\n0 upgraded, 0 newly installed, 0 to remove and 0 not upgraded.\n"))
})
}
func TestParseDnfCheckUpdate(t *testing.T) {
tests := []struct {
file string
count int
}{
{"dnf4_rocky9_check_update.txt", 110},
{"dnf4_rocky9_check_update_security.txt", 53},
{"dnf5_fedora42_check_update.txt", 20},
{"dnf5_fedora42_check_update_security.txt", 5},
}
for _, tt := range tests {
t.Run(tt.file, func(t *testing.T) {
assert.Len(t, parseDnfCheckUpdate(readPackageUpdatesTestData(t, tt.file)), tt.count)
})
}
t.Run("versions keep epoch and arch", func(t *testing.T) {
packages := parseDnfCheckUpdate(readPackageUpdatesTestData(t, "dnf5_fedora42_check_update.txt"))
assert.Equal(t, system.PackageUpdate{Name: "openssl-libs.aarch64", Available: "1:3.2.6-4.fc42"}, findPackage(t, packages, "openssl-libs.aarch64"))
})
t.Run("obsoletes section, notices and wrapped names", func(t *testing.T) {
out := `
kernel.x86_64 5.14.0-503.el9 baseos
Security: kernel-core-5.14.0-427.el9.x86_64 is an installed security update
python3-some-very-long-package-name-that-wraps.noarch
1.2.3-4.el9 appstream
Obsoleting Packages
grub2-tools.x86_64 1:2.06-80.el9 baseos
grub2-tools.x86_64 1:2.06-77.el9 @baseos
`
assert.Equal(t, []system.PackageUpdate{
{Name: "kernel.x86_64", Available: "5.14.0-503.el9"},
{Name: "python3-some-very-long-package-name-that-wraps.noarch", Available: "1.2.3-4.el9"},
}, parseDnfCheckUpdate(out))
})
}
func TestParseRpmInstalled(t *testing.T) {
installed := parseRpmInstalled(readPackageUpdatesTestData(t, "dnf4_rocky9_rpm_installed.txt"))
assert.Len(t, installed, 110)
assert.Equal(t, "2.34-83.el9.7", installed["glibc.aarch64"])
assert.Equal(t, "1:3.0.7-24.el9", installed["openssl-libs.aarch64"])
assert.NotContains(t, installed, "package")
// several installed kernels: the last one wins
installed = parseRpmInstalled("kernel.x86_64 5.14.0-427.el9\nkernel.x86_64 5.14.0-503.el9\n")
assert.Equal(t, "5.14.0-503.el9", installed["kernel.x86_64"])
}
func TestCheckDnf(t *testing.T) {
tests := []struct {
name, updates, security, installed string
total, securityCount int
pkg system.PackageUpdate
}{
{
name: "dnf4",
updates: "dnf4_rocky9_check_update.txt", security: "dnf4_rocky9_check_update_security.txt", installed: "dnf4_rocky9_rpm_installed.txt",
total: 110, securityCount: 53,
pkg: system.PackageUpdate{Name: "vim-minimal", Current: "2:8.2.2637-20.el9_1", Available: "2:8.2.2637-26.el9_8.21", Security: true},
},
{
name: "dnf5",
updates: "dnf5_fedora42_check_update.txt", security: "dnf5_fedora42_check_update_security.txt", installed: "dnf5_fedora42_rpm_installed.txt",
total: 20, securityCount: 5,
pkg: system.PackageUpdate{Name: "openssl-libs", Current: "1:3.2.6-3.fc42", Available: "1:3.2.6-4.fc42", Security: true},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
fakeCommands(t, map[string]string{
"dnf": `case "$*" in *--security*) cat "` + testDataPath(t, tt.security) + `" ;; *) cat "` + testDataPath(t, tt.updates) + `" ;; esac
exit 100`,
"rpm": `cat "` + testDataPath(t, tt.installed) + `"
exit 1`,
})
result, err := checkDnf(context.Background())
require.NoError(t, err)
assert.Equal(t, []uint16{uint16(tt.total), uint16(tt.securityCount)}, result.counts)
assert.True(t, result.securityKnown)
assert.Len(t, result.packages, tt.total)
assert.Equal(t, tt.pkg, findPackage(t, result.packages, tt.pkg.Name))
for _, pkg := range result.packages {
assert.NotEmpty(t, pkg.Current, pkg.Name)
}
})
}
t.Run("security query fails", func(t *testing.T) {
fakeCommands(t, map[string]string{
"dnf": `case "$*" in *--security*) exit 1 ;; esac
echo "bash.x86_64 5.1.8-9.el9 baseos"
exit 100`,
"rpm": `echo "bash.x86_64 5.1.8-6.el9_1"`,
})
result, err := checkDnf(context.Background())
require.NoError(t, err)
assert.Equal(t, []uint16{1}, result.counts)
assert.False(t, result.securityKnown)
assert.Equal(t, []system.PackageUpdate{{Name: "bash", Current: "5.1.8-6.el9_1", Available: "5.1.8-9.el9"}}, result.packages)
})
}
func TestParseZypperTable(t *testing.T) {
tests := []struct {
file string
count uint16
}{
{"zypper_leap155_list_updates.txt", 22},
{"zypper_leap155_list_patches_security.txt", 4},
{"zypper_leap156_list_updates_none.txt", 0},
}
for _, tt := range tests {
t.Run(tt.file, func(t *testing.T) {
assert.Equal(t, tt.count, parseZypperTable(readPackageUpdatesTestData(t, tt.file)))
})
}
}
func TestParseZypperListUpdates(t *testing.T) {
packages := parseZypperListUpdates(readPackageUpdatesTestData(t, "zypper_leap155_list_updates.txt"))
assert.Len(t, packages, 22)
assert.Equal(t, system.PackageUpdate{Name: "zypper", Current: "1.14.76-150500.6.6.15", Available: "1.14.78-150500.6.14.1"}, findPackage(t, packages, "zypper"))
assert.Equal(t, system.PackageUpdate{Name: "aaa_base", Current: "84.87+git20180409.04c9dae-150300.10.20.1", Available: "84.87+git20180409.04c9dae-150300.10.23.1"}, findPackage(t, packages, "aaa_base"))
assert.Empty(t, parseZypperListUpdates(readPackageUpdatesTestData(t, "zypper_leap156_list_updates_none.txt")))
// patch tables have no version columns
assert.Empty(t, parseZypperListUpdates(readPackageUpdatesTestData(t, "zypper_leap155_list_patches_security.txt")))
}
func TestCheckZypper(t *testing.T) {
fakeCommands(t, map[string]string{
"zypper": `case "$*" in *list-patches*) cat "` + testDataPath(t, "zypper_leap155_list_patches_security.txt") + `" ;; *) cat "` + testDataPath(t, "zypper_leap155_list_updates.txt") + `" ;; esac`,
})
result, err := checkZypper(context.Background())
require.NoError(t, err)
// security patches don't map to packages, so only the count is known
assert.Equal(t, []uint16{22, 4}, result.counts)
assert.False(t, result.securityKnown)
assert.Len(t, result.packages, 22)
}
func TestParsePacmanCheckUpdates(t *testing.T) {
assert.Equal(t, []system.PackageUpdate{
{Name: "libpcap", Current: "1.10.7-1", Available: "1.11.0-1"},
{Name: "libsecret", Current: "0.21.7-1", Available: "0.21.8.2-1"},
{Name: "libtirpc", Current: "1.3.7-1", Available: "1.3.8-1"},
{Name: "tzdata", Current: "2026c-1", Available: "2026d-1"},
}, parsePacmanCheckUpdates(readPackageUpdatesTestData(t, "pacman_checkupdates.txt")))
assert.Empty(t, parsePacmanCheckUpdates(""))
}
func TestParseApkUpgradable(t *testing.T) {
packages := parseApkUpgradable(readPackageUpdatesTestData(t, "apk_alpine320_list_upgradable.txt"))
assert.Len(t, packages, 10)
assert.Equal(t, system.PackageUpdate{Name: "musl", Current: "1.2.5-r0", Available: "1.2.5-r3"}, packages[6])
// names with dashes and digits
assert.Equal(t, system.PackageUpdate{Name: "busybox-binsh", Current: "1.36.1-r28", Available: "1.36.1-r31"}, packages[2])
assert.Equal(t, system.PackageUpdate{Name: "ca-certificates-bundle", Current: "20240226-r0", Available: "20260413-r0"}, packages[3])
assert.Equal(t, system.PackageUpdate{Name: "libcrypto3", Current: "3.3.0-r2", Available: "3.3.7-r0"}, packages[4])
assert.Empty(t, parseApkUpgradable(""))
}
func TestSplitApkNameVersion(t *testing.T) {
tests := []struct{ in, name, version string }{
{"musl-1.2.5-r3", "musl", "1.2.5-r3"},
{"py3-foo-bar-2.0_rc1-r0", "py3-foo-bar", "2.0_rc1-r0"},
{"apk-tools-2.14.4-r1", "apk-tools", "2.14.4-r1"},
// unexpected formats keep the whole string as the name
{"noversion", "noversion", ""},
{"name-1.0", "name-1.0", ""},
{"-1.0-r0", "-1.0-r0", ""},
}
for _, tt := range tests {
name, version := splitApkNameVersion(tt.in)
assert.Equal(t, tt.name, name, tt.in)
assert.Equal(t, tt.version, version, tt.in)
}
}
func TestPackageUpdatesManagerCaching(t *testing.T) {
calls := make(chan struct{}, 10)
packages := []system.PackageUpdate{{Name: "libc6", Current: "1", Available: "2", Security: true}}
result := packageUpdatesResult{counts: []uint16{3, 1}, packages: packages, securityKnown: true}
var resultErr error
pm := &packageUpdatesManager{
name: "apt",
interval: time.Hour,
check: func(context.Context) (packageUpdatesResult, error) {
calls <- struct{}{}
return result, resultErr
},
}
waitIdle := func() {
require.Eventually(t, func() bool {
pm.Lock()
defer pm.Unlock()
return !pm.running
}, time.Second, time.Millisecond)
}
// no check has finished yet
assert.Equal(t, system.PackageUpdates{Manager: "apt"}, pm.list())
now := time.Now()
// first call starts a background check and returns nothing yet
assert.Nil(t, pm.get(now))
waitIdle()
assert.Len(t, calls, 1)
// cached result within interval, no new check
assert.Equal(t, []uint16{3, 1}, pm.get(now.Add(time.Minute)))
assert.Len(t, calls, 1)
list := pm.list()
assert.Equal(t, "apt", list.Manager)
assert.True(t, list.SecurityKnown)
assert.Equal(t, packages, list.Packages)
assert.NotZero(t, list.CheckedAt)
// list never starts a check
assert.Len(t, calls, 1)
// stale after interval: returns cached value and refreshes in background
result, resultErr = packageUpdatesResult{}, errors.New("boom")
assert.Equal(t, []uint16{3, 1}, pm.get(now.Add(2*time.Hour)))
waitIdle()
assert.Len(t, calls, 2)
// failed check clears the counts and the list
assert.Nil(t, pm.get(time.Now()))
assert.Nil(t, pm.list().Packages)
assert.False(t, pm.list().SecurityKnown)
}
func TestGetPackageUpdatesHandler(t *testing.T) {
var sent any
hctx := &HandlerContext{
Agent: &Agent{},
SendResponse: func(data any, _ *uint32) error {
sent = data
return nil
},
}
handler := &GetPackageUpdatesHandler{}
// no supported package manager
require.NoError(t, handler.Handle(hctx))
assert.Equal(t, system.PackageUpdates{}, sent)
packages := []system.PackageUpdate{{Name: "musl", Current: "1.2.5-r0", Available: "1.2.5-r3"}}
hctx.Agent.packageUpdates = &packageUpdatesManager{
name: "apk",
result: packageUpdatesResult{counts: []uint16{1}, packages: packages},
checkedAt: time.Unix(1700000000, 0),
}
require.NoError(t, handler.Handle(hctx))
assert.Equal(t, system.PackageUpdates{Manager: "apk", CheckedAt: 1700000000, Packages: packages}, sent)
}
func TestPacmanCheckSync(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("requires a shell script on PATH")
}
binDir := t.TempDir()
dataDir := t.TempDir()
logFile := filepath.Join(binDir, "calls.log")
// fake checkupdates logs its args and db path, and creates the sync dir when syncing
script := `#!/bin/sh
echo "args=[$*] db=$CHECKUPDATES_DB" >> ` + logFile + `
[ "$1" = "-n" ] || mkdir -p "$CHECKUPDATES_DB/sync"
echo "linux 6.1-1 -> 6.2-1"
`
require.NoError(t, os.WriteFile(filepath.Join(binDir, "checkupdates"), []byte(script), 0o755))
t.Setenv("PATH", binDir+string(os.PathListSeparator)+os.Getenv("PATH"))
check := newPacmanCheck(dataDir)
dbPath := filepath.Join(dataDir, "checkup-db")
readCalls := func() []string {
data, err := os.ReadFile(logFile)
require.NoError(t, err)
return strings.Split(strings.TrimSpace(string(data)), "\n")
}
// first check syncs
result, err := check(context.Background())
require.NoError(t, err)
assert.Equal(t, []uint16{1}, result.counts)
assert.Equal(t, []system.PackageUpdate{{Name: "linux", Current: "6.1-1", Available: "6.2-1"}}, result.packages)
// later checks reuse the synced copy
_, err = check(context.Background())
require.NoError(t, err)
// a missing private copy forces a sync
require.NoError(t, os.RemoveAll(dbPath))
_, err = check(context.Background())
require.NoError(t, err)
assert.Equal(t, []string{
"args=[] db=" + dbPath,
"args=[-n] db=" + dbPath,
"args=[] db=" + dbPath,
}, readCalls())
}

538
agent/probe.go Normal file
View File

@@ -0,0 +1,538 @@
package agent
import (
"errors"
"fmt"
"math"
"math/rand"
"net"
"net/http"
// "strconv"
"sync"
"time"
"log/slog"
"github.com/henrygd/beszel/internal/entities/probe"
)
// Probes run at user-defined intervals (e.g., every 10s).
// To keep memory usage low and constant, data is stored in two layers:
// 1. Raw samples: The most recent individual results (kept for probeRawRetention).
// 2. Minute buckets: A ring buffer of 61 buckets, each representing one
// wall-clock minute. Samples collected within the same minute are aggregated
// (sum, min, max, count) into a single bucket.
//
// Short-term requests (<= 70s) use raw samples.
// Long-term requests (up to 1h) use the minute buckets to avoid storing thousands
// of individual data points.
const (
// probeRawRetention is the duration to keep individual samples
probeRawRetention = 61 * time.Second
// probeMinuteBucketLen is the number of 1-minute buckets to keep (1 hour + 1 for partials)
probeMinuteBucketLen int32 = 61
)
// ProbeManager manages network probe tasks.
type ProbeManager struct {
mu sync.RWMutex
probes map[string]*probeTask // key = probe.Config.Key()
httpClient *http.Client
}
// probeTask owns retention buffers and cancellation for a single probe config.
type probeTask struct {
config probe.Config
cancel chan struct{}
mu sync.Mutex
samples []probeSample
buckets [probeMinuteBucketLen]probeBucket
}
// probeSample stores one probe attempt and its collection time.
type probeSample struct {
responseUs int64 // -1 means loss
timestamp time.Time
}
// probeBucket stores one minute of aggregated probe data.
type probeBucket struct {
minute int32
filled bool
stats probeAggregate
}
// probeAggregate accumulates successful response stats and total sample counts.
type probeAggregate struct {
sumUs int64
minUs int64
maxUs int64
totalCount int64
successCount int64
}
func newProbeManager() *ProbeManager {
return &ProbeManager{
probes: make(map[string]*probeTask),
httpClient: &http.Client{Timeout: 10 * time.Second},
}
}
func newProbeTask(config probe.Config) *probeTask {
return &probeTask{
config: config,
cancel: make(chan struct{}),
samples: make([]probeSample, 0, 64),
}
}
func newProbeTaskFromExisting(config probe.Config, existing *probeTask) *probeTask {
task := newProbeTask(config)
if existing == nil {
return task
}
existing.mu.Lock()
defer existing.mu.Unlock()
task.samples = append(task.samples, existing.samples...)
task.buckets = existing.buckets
return task
}
// newProbeAggregate initializes an aggregate with an unset minimum value.
func newProbeAggregate() probeAggregate {
return probeAggregate{minUs: math.MaxInt64}
}
// addResponse folds a single probe sample into the aggregate.
func (agg *probeAggregate) addResponse(responseUs int64) {
agg.totalCount++
if responseUs < 0 {
return
}
agg.successCount++
agg.sumUs += responseUs
if responseUs < agg.minUs {
agg.minUs = responseUs
}
if responseUs > agg.maxUs {
agg.maxUs = responseUs
}
}
// addAggregate merges another aggregate into this one.
func (agg *probeAggregate) addAggregate(other probeAggregate) {
if other.totalCount == 0 {
return
}
agg.totalCount += other.totalCount
agg.successCount += other.successCount
agg.sumUs += other.sumUs
if other.successCount == 0 {
return
}
if agg.minUs == math.MaxInt64 || other.minUs < agg.minUs {
agg.minUs = other.minUs
}
if other.maxUs > agg.maxUs {
agg.maxUs = other.maxUs
}
}
// hasData reports whether the aggregate contains any samples.
func (agg probeAggregate) hasData() bool {
return agg.totalCount > 0
}
// result converts the aggregate into the probe result format.
func (agg probeAggregate) result() probe.Result {
avg := agg.avgResponse()
result := probe.Result{
AvgResponse: avg,
MinResponse: agg.minUs,
MaxResponse: agg.maxUs,
PacketLoss: agg.lossPercentage(),
}
if agg.successCount == 0 {
result.MinResponse, result.MaxResponse = 0, 0
}
return result
}
// avgResponse returns the rounded average of successful samples.
func (agg probeAggregate) avgResponse() int64 {
if agg.successCount == 0 {
return 0
}
return agg.sumUs / agg.successCount
}
// lossPercentage returns the rounded failure rate for the aggregate.
func (agg probeAggregate) lossPercentage() float64 {
if agg.totalCount == 0 {
return 0
}
return math.Round(float64(agg.totalCount-agg.successCount)/float64(agg.totalCount)*10000) / 100
}
// SyncProbes replaces all probe tasks with the given configs.
func (pm *ProbeManager) SyncProbes(configs []probe.Config) {
pm.mu.Lock()
defer pm.mu.Unlock()
// Build set of new keys
newKeys := make(map[string]probe.Config, len(configs))
for _, cfg := range configs {
if cfg.ID == "" {
continue
}
newKeys[cfg.ID] = cfg
}
// Stop removed probes
for key, task := range pm.probes {
if _, exists := newKeys[key]; !exists {
close(task.cancel)
delete(pm.probes, key)
}
}
// Start new probes and restart tasks whose config changed.
for key, cfg := range newKeys {
task, exists := pm.probes[key]
if exists && task.config == cfg {
continue
}
if exists {
close(task.cancel)
}
task = newProbeTaskFromExisting(cfg, task)
pm.probes[key] = task
go pm.runProbe(task, false)
}
}
// HandleSyncRequest applies a full or incremental probe sync request.
func (pm *ProbeManager) HandleSyncRequest(req probe.SyncRequest) (probe.SyncResponse, error) {
switch req.Action {
case probe.SyncActionReplace:
pm.SyncProbes(req.Configs)
return probe.SyncResponse{}, nil
case probe.SyncActionUpsert:
result, err := pm.UpsertProbe(req.Config, req.RunNow)
if err != nil {
return probe.SyncResponse{}, err
}
if result == nil {
return probe.SyncResponse{}, nil
}
return probe.SyncResponse{Result: *result}, nil
case probe.SyncActionDelete:
if req.Config.ID == "" {
return probe.SyncResponse{}, errors.New("missing probe ID for delete")
}
pm.DeleteProbe(req.Config.ID)
return probe.SyncResponse{}, nil
default:
return probe.SyncResponse{}, fmt.Errorf("unknown probe sync action: %d", req.Action)
}
}
// UpsertProbe creates or replaces a single probe task.
func (pm *ProbeManager) UpsertProbe(config probe.Config, runNow bool) (*probe.Result, error) {
if config.ID == "" {
return nil, errors.New("missing probe ID")
}
pm.mu.Lock()
task, exists := pm.probes[config.ID]
startTask := false
if exists && task.config == config {
pm.mu.Unlock()
if !runNow {
return nil, nil
}
return pm.runProbeNow(task), nil
}
if exists {
close(task.cancel)
}
task = newProbeTaskFromExisting(config, task)
pm.probes[config.ID] = task
startTask = true
pm.mu.Unlock()
if runNow {
result := pm.runProbeNow(task)
if startTask {
go pm.runProbe(task, false)
}
return result, nil
}
if startTask {
go pm.runProbe(task, false)
}
return nil, nil
}
// DeleteProbe stops and removes a single probe task.
func (pm *ProbeManager) DeleteProbe(id string) {
if id == "" {
return
}
pm.mu.Lock()
defer pm.mu.Unlock()
if task, exists := pm.probes[id]; exists {
close(task.cancel)
delete(pm.probes, id)
}
}
// GetResults returns aggregated results for all probes over the last supplied duration in ms.
func (pm *ProbeManager) GetResults(durationMs uint16) map[string]probe.Result {
pm.mu.RLock()
defer pm.mu.RUnlock()
results := make(map[string]probe.Result, len(pm.probes))
now := time.Now()
duration := time.Duration(durationMs) * time.Millisecond
for _, task := range pm.probes {
task.mu.Lock()
result, ok := task.resultLocked(duration, now)
task.mu.Unlock()
if !ok {
continue
}
results[task.config.ID] = result
}
return results
}
// Stop stops all probe tasks.
func (pm *ProbeManager) Stop() {
pm.mu.Lock()
defer pm.mu.Unlock()
for key, task := range pm.probes {
close(task.cancel)
delete(pm.probes, key)
}
}
// runProbe executes a single probe task in a loop.
func (pm *ProbeManager) runProbe(task *probeTask, runNow bool) {
interval := time.Duration(task.config.Interval) * time.Second
if interval < time.Second {
interval = 30 * time.Second
}
stagger := getStagger(interval.Milliseconds())
slog.Debug("starting probe task", "target", task.config.Target, "delay", stagger.String(), "interval", interval.String())
if runNow {
pm.executeProbe(task)
}
select {
case <-task.cancel:
// slog.Info("removed probe", "target", task.config.Target)
return
case <-time.After(stagger):
pm.executeProbe(task)
}
ticker := time.Tick(interval)
for {
select {
case <-task.cancel:
// slog.Info("removed probe", "target", task.config.Target)
return
case <-ticker:
pm.executeProbe(task)
}
}
}
// getStagger returns a random duration between intervalSeconds/2 and intervalSeconds to stagger initial probe executions
func getStagger(intervalMilli int64) time.Duration {
intervalMilliInt := int(intervalMilli)
randomDelayInt := rand.Intn(intervalMilliInt)
if randomDelayInt < intervalMilliInt/2 {
randomDelayInt += intervalMilliInt / 2
}
return time.Duration(randomDelayInt) * time.Millisecond
}
func (pm *ProbeManager) runProbeNow(task *probeTask) *probe.Result {
pm.executeProbe(task)
task.mu.Lock()
defer task.mu.Unlock()
result, ok := task.resultLocked(time.Minute, time.Now())
if !ok {
return nil
}
return &result
}
// resultLocked returns the aggregated probe result for the requested duration along with a bool indicating whether any data was available.
func (task *probeTask) resultLocked(duration time.Duration, now time.Time) (probe.Result, bool) {
agg := task.aggregateLocked(duration, now)
hourAgg := task.aggregateLocked(time.Hour, now)
if !agg.hasData() {
return probe.Result{}, false
}
result := agg.result()
result.AvgResponse1h = hourAgg.avgResponse()
result.MinResponse1h = hourAgg.minUs
result.MaxResponse1h = hourAgg.maxUs
result.PacketLoss1h = hourAgg.lossPercentage()
if hourAgg.successCount == 0 {
result.MinResponse1h, result.MaxResponse1h = 0, 0
}
return result, true
}
// aggregateLocked collects probe data for the requested time window.
func (task *probeTask) aggregateLocked(duration time.Duration, now time.Time) probeAggregate {
cutoff := now.Add(-duration)
// Keep short windows exact; longer windows read from minute buckets to avoid raw-sample retention.
if duration <= probeRawRetention {
return aggregateSamplesSince(task.samples, cutoff)
}
return aggregateBucketsSince(task.buckets[:], cutoff, now)
}
// aggregateSamplesSince aggregates raw samples newer than the cutoff.
func aggregateSamplesSince(samples []probeSample, cutoff time.Time) probeAggregate {
agg := newProbeAggregate()
for _, sample := range samples {
if sample.timestamp.Before(cutoff) {
continue
}
agg.addResponse(sample.responseUs)
}
return agg
}
// aggregateBucketsSince aggregates minute buckets overlapping the requested window.
func aggregateBucketsSince(buckets []probeBucket, cutoff, now time.Time) probeAggregate {
agg := newProbeAggregate()
startMinute := int32(cutoff.Unix() / 60)
endMinute := int32(now.Unix() / 60)
for _, bucket := range buckets {
if !bucket.filled || bucket.minute < startMinute || bucket.minute > endMinute {
continue
}
agg.addAggregate(bucket.stats)
}
return agg
}
// addSampleLocked stores a fresh sample in both raw and per-minute retention buffers.
func (task *probeTask) addSampleLocked(sample probeSample) {
cutoff := sample.timestamp.Add(-probeRawRetention)
start := 0
for i := range task.samples {
if !task.samples[i].timestamp.Before(cutoff) {
start = i
break
}
if i == len(task.samples)-1 {
start = len(task.samples)
}
}
if start > 0 {
size := copy(task.samples, task.samples[start:])
task.samples = task.samples[:size]
}
task.samples = append(task.samples, sample)
minute := int32(sample.timestamp.Unix() / 60)
// Each slot stores one wall-clock minute, so the ring stays fixed-size at ~1h per probe.
bucket := &task.buckets[minute%probeMinuteBucketLen]
if !bucket.filled || bucket.minute != minute {
bucket.minute = minute
bucket.filled = true
bucket.stats = newProbeAggregate()
}
bucket.stats.addResponse(sample.responseUs)
}
// executeProbe runs the configured probe and records the sample.
func (pm *ProbeManager) executeProbe(task *probeTask) {
// slog.Info("running probe", "id", task.config.ID, "interval", task.config.Interval)
var responseUs int64
var err error
switch task.config.Protocol {
case "icmp":
responseUs, err = probeICMP(task.config.Target)
case "tcp":
responseUs, err = probeTCP(task.config.Target, task.config.Port)
case "http":
responseUs, err = probeHTTP(pm.httpClient, task.config.Target)
default:
slog.Warn("unknown probe protocol", "protocol", task.config.Protocol)
return
}
if err != nil {
slog.Warn("probe failed", "err", err, "target", task.config.Target, "protocol", task.config.Protocol)
}
sample := probeSample{
responseUs: responseUs,
timestamp: time.Now(),
}
task.mu.Lock()
task.addSampleLocked(sample)
task.mu.Unlock()
}
// probeTCP measures pure TCP handshake response (excluding DNS resolution).
// Returns -1 and an error on failure.
func probeTCP(target string, port uint16) (int64, error) {
// Resolve DNS first, outside the timing window
ips, err := net.LookupHost(target)
if err != nil || len(ips) == 0 {
return -1, err
}
addr := net.JoinHostPort(ips[0], fmt.Sprintf("%d", port))
// Measure only the TCP handshake
start := time.Now()
conn, err := net.DialTimeout("tcp", addr, 3*time.Second)
if err != nil {
return -1, err
}
conn.Close()
return time.Since(start).Microseconds(), nil
}
// probeHTTP measures HTTP GET request response in microseconds. Returns -1 and an error on failure.
func probeHTTP(client *http.Client, url string) (int64, error) {
if client == nil {
client = http.DefaultClient
}
start := time.Now()
resp, err := client.Get(url)
if err != nil {
return -1, err
}
resp.Body.Close()
if resp.StatusCode >= 400 {
return -1, fmt.Errorf("HTTP error: %s", resp.Status)
}
return time.Since(start).Microseconds(), nil
}

241
agent/probe_ping.go Normal file
View File

@@ -0,0 +1,241 @@
package agent
import (
"errors"
"math"
"net"
"os"
"os/exec"
"regexp"
"runtime"
"strconv"
"sync"
"time"
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
"log/slog"
)
var pingTimeRegex = regexp.MustCompile(`time[=<]([\d.]+)\s*ms`)
type icmpPacketConn interface {
Close() error
}
// icmpMethod tracks which ICMP approach to use. Once a method succeeds or
// all native methods fail, the choice is cached so subsequent probes skip
// the trial-and-error overhead.
type icmpMethod uint8
const (
icmpUntried icmpMethod = iota // haven't tried yet
icmpRaw // privileged raw socket
icmpDatagram // unprivileged datagram socket
icmpExecFallback // shell out to system ping command
)
// icmpFamily holds the network parameters and cached detection result for one address family.
type icmpFamily struct {
rawNetwork string // e.g. "ip4:icmp" or "ip6:ipv6-icmp"
dgramNetwork string // e.g. "udp4" or "udp6"
listenAddr string // "0.0.0.0" or "::"
echoType icmp.Type // outgoing echo request type
replyType icmp.Type // expected echo reply type
proto int // IANA protocol number for parsing replies
isIPv6 bool
mode icmpMethod // cached detection result (guarded by icmpModeMu)
}
var (
icmpV4 = icmpFamily{
rawNetwork: "ip4:icmp",
dgramNetwork: "udp4",
listenAddr: "0.0.0.0",
echoType: ipv4.ICMPTypeEcho,
replyType: ipv4.ICMPTypeEchoReply,
proto: 1,
}
icmpV6 = icmpFamily{
rawNetwork: "ip6:ipv6-icmp",
dgramNetwork: "udp6",
listenAddr: "::",
echoType: ipv6.ICMPTypeEchoRequest,
replyType: ipv6.ICMPTypeEchoReply,
proto: 58,
isIPv6: true,
}
icmpModeMu sync.Mutex
icmpListen = func(network, listenAddr string) (icmpPacketConn, error) {
return icmp.ListenPacket(network, listenAddr)
}
)
// probeICMP sends an ICMP echo request and measures round-trip response.
// Supports both IPv4 and IPv6 targets. The ICMP method (raw socket,
// unprivileged datagram, or exec fallback) is detected once per address
// family and cached for subsequent probes.
// Returns response in microseconds, or -1 and an error on failure.
func probeICMP(target string) (int64, error) {
family, ip, err := resolveICMPTarget(target)
if err != nil {
return -1, err
}
icmpModeMu.Lock()
if family.mode == icmpUntried {
family.mode = detectICMPMode(family, icmpListen)
}
mode := family.mode
icmpModeMu.Unlock()
switch mode {
case icmpRaw:
return probeICMPNative(family.rawNetwork, family, &net.IPAddr{IP: ip})
case icmpDatagram:
return probeICMPNative(family.dgramNetwork, family, &net.UDPAddr{IP: ip})
case icmpExecFallback:
return probeICMPExec(target, family.isIPv6)
default:
return -1, errors.New("unsupported ICMP mode")
}
}
// resolveICMPTarget resolves a target hostname or IP to determine the address
// family and concrete IP address. Prefers IPv4 for dual-stack hostnames.
func resolveICMPTarget(target string) (*icmpFamily, net.IP, error) {
if ip := net.ParseIP(target); ip != nil {
if ip.To4() != nil {
return &icmpV4, ip.To4(), nil
}
return &icmpV6, ip, nil
}
ips, err := net.LookupIP(target)
if err != nil || len(ips) == 0 {
return nil, nil, err
}
for _, ip := range ips {
if v4 := ip.To4(); v4 != nil {
return &icmpV4, v4, nil
}
}
return &icmpV6, ips[0], nil
}
func detectICMPMode(family *icmpFamily, listen func(network, listenAddr string) (icmpPacketConn, error)) icmpMethod {
label := "IPv4"
if family.isIPv6 {
label = "IPv6"
}
conn, err := listen(family.rawNetwork, family.listenAddr)
slog.Debug("ICMP raw socket test", "family", label, "err", err)
if err == nil {
conn.Close()
return icmpRaw
}
conn, err = listen(family.dgramNetwork, family.listenAddr)
slog.Debug("ICMP datagram socket test", "family", label, "err", err)
if err == nil {
conn.Close()
return icmpDatagram
}
return icmpExecFallback
}
// probeICMPNative sends an ICMP echo request using Go's x/net/icmp package.
func probeICMPNative(network string, family *icmpFamily, dst net.Addr) (int64, error) {
conn, err := icmp.ListenPacket(network, family.listenAddr)
if err != nil {
return -1, err
}
defer conn.Close()
// Build ICMP echo request
msg := &icmp.Message{
Type: family.echoType,
Code: 0,
Body: &icmp.Echo{
ID: os.Getpid() & 0xffff,
Seq: 1,
Data: []byte("beszel-probe"),
},
}
msgBytes, err := msg.Marshal(nil)
if err != nil {
return -1, err
}
// Set deadline before sending
conn.SetDeadline(time.Now().Add(3 * time.Second))
start := time.Now()
if _, err := conn.WriteTo(msgBytes, dst); err != nil {
return -1, err
}
// Read reply
buf := make([]byte, 1500)
for {
n, _, err := conn.ReadFrom(buf)
if err != nil {
return -1, err
}
reply, err := icmp.ParseMessage(family.proto, buf[:n])
if err != nil {
return -1, err
}
if reply.Type == family.replyType {
return time.Since(start).Microseconds(), nil
}
// Ignore non-echo-reply messages (e.g. destination unreachable) and keep reading
}
}
// probeICMPExec falls back to the system ping command. Returns -1 and an error on failure.
func probeICMPExec(target string, isIPv6 bool) (int64, error) {
var cmd *exec.Cmd
switch runtime.GOOS {
case "windows":
if isIPv6 {
cmd = exec.Command("ping", "-6", "-n", "1", "-w", "3000", target)
} else {
cmd = exec.Command("ping", "-n", "1", "-w", "3000", target)
}
default:
if isIPv6 {
cmd = exec.Command("ping", "-6", "-c", "1", "-W", "3", target)
} else {
cmd = exec.Command("ping", "-c", "1", "-W", "3", target)
}
}
start := time.Now()
output, err := cmd.Output()
if err != nil {
// If ping fails but we got output, still try to parse
if len(output) == 0 {
return -1, err
}
}
matches := pingTimeRegex.FindSubmatch(output)
if len(matches) >= 2 {
if ms, err := strconv.ParseFloat(string(matches[1]), 64); err == nil {
return int64(math.Round(ms * 1000)), nil
}
}
// Fallback: use wall clock time if ping succeeded but parsing failed
if err == nil {
return time.Since(start).Microseconds(), nil
}
return -1, err
}

121
agent/probe_ping_test.go Normal file
View File

@@ -0,0 +1,121 @@
//go:build testing
package agent
import (
"errors"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type testICMPPacketConn struct{}
func (testICMPPacketConn) Close() error { return nil }
func TestDetectICMPMode(t *testing.T) {
tests := []struct {
name string
family *icmpFamily
rawErr error
udpErr error
want icmpMethod
wantNetworks []string
}{
{
name: "IPv4 prefers raw socket when available",
family: &icmpV4,
want: icmpRaw,
wantNetworks: []string{"ip4:icmp"},
},
{
name: "IPv4 uses datagram when raw unavailable",
family: &icmpV4,
rawErr: errors.New("operation not permitted"),
want: icmpDatagram,
wantNetworks: []string{"ip4:icmp", "udp4"},
},
{
name: "IPv4 falls back to exec when both unavailable",
family: &icmpV4,
rawErr: errors.New("operation not permitted"),
udpErr: errors.New("protocol not supported"),
want: icmpExecFallback,
wantNetworks: []string{"ip4:icmp", "udp4"},
},
{
name: "IPv6 prefers raw socket when available",
family: &icmpV6,
want: icmpRaw,
wantNetworks: []string{"ip6:ipv6-icmp"},
},
{
name: "IPv6 uses datagram when raw unavailable",
family: &icmpV6,
rawErr: errors.New("operation not permitted"),
want: icmpDatagram,
wantNetworks: []string{"ip6:ipv6-icmp", "udp6"},
},
{
name: "IPv6 falls back to exec when both unavailable",
family: &icmpV6,
rawErr: errors.New("operation not permitted"),
udpErr: errors.New("protocol not supported"),
want: icmpExecFallback,
wantNetworks: []string{"ip6:ipv6-icmp", "udp6"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
calls := make([]string, 0, 2)
listen := func(network, listenAddr string) (icmpPacketConn, error) {
require.Equal(t, tt.family.listenAddr, listenAddr)
calls = append(calls, network)
switch network {
case tt.family.rawNetwork:
if tt.rawErr != nil {
return nil, tt.rawErr
}
case tt.family.dgramNetwork:
if tt.udpErr != nil {
return nil, tt.udpErr
}
default:
t.Fatalf("unexpected network %q", network)
}
return testICMPPacketConn{}, nil
}
assert.Equal(t, tt.want, detectICMPMode(tt.family, listen))
assert.Equal(t, tt.wantNetworks, calls)
})
}
}
func TestResolveICMPTarget(t *testing.T) {
t.Run("IPv4 literal", func(t *testing.T) {
family, ip, err := resolveICMPTarget("127.0.0.1")
require.NoError(t, err)
require.NotNil(t, family)
assert.False(t, family.isIPv6)
assert.Equal(t, "127.0.0.1", ip.String())
})
t.Run("IPv6 literal", func(t *testing.T) {
family, ip, err := resolveICMPTarget("::1")
require.NoError(t, err)
require.NotNil(t, family)
assert.True(t, family.isIPv6)
assert.Equal(t, "::1", ip.String())
})
t.Run("IPv4-mapped IPv6 resolves as IPv4", func(t *testing.T) {
family, ip, err := resolveICMPTarget("::ffff:127.0.0.1")
require.NoError(t, err)
require.NotNil(t, family)
assert.False(t, family.isIPv6)
assert.Equal(t, "127.0.0.1", ip.String())
})
}

356
agent/probe_test.go Normal file
View File

@@ -0,0 +1,356 @@
package agent
import (
"net"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/henrygd/beszel/internal/entities/probe"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestProbeTaskAggregateLockedUsesRawSamplesForShortWindows(t *testing.T) {
now := time.Date(2026, time.April, 21, 12, 0, 0, 0, time.UTC)
task := &probeTask{}
task.addSampleLocked(probeSample{responseUs: 10, timestamp: now.Add(-90 * time.Second)})
task.addSampleLocked(probeSample{responseUs: 20, timestamp: now.Add(-30 * time.Second)})
task.addSampleLocked(probeSample{responseUs: -1, timestamp: now.Add(-10 * time.Second)})
agg := task.aggregateLocked(time.Minute, now)
require.True(t, agg.hasData())
assert.Equal(t, int64(2), agg.totalCount)
assert.Equal(t, int64(1), agg.successCount)
result := agg.result()
assert.Equal(t, int64(20), result.AvgResponse)
assert.Equal(t, int64(20), result.MinResponse)
assert.Equal(t, int64(20), result.MaxResponse)
assert.Equal(t, 50.0, result.PacketLoss)
}
func TestProbeTaskAggregateLockedUsesMinuteBucketsForLongWindows(t *testing.T) {
now := time.Date(2026, time.April, 21, 12, 0, 30, 0, time.UTC)
task := &probeTask{}
task.addSampleLocked(probeSample{responseUs: 10, timestamp: now.Add(-11 * time.Minute)})
task.addSampleLocked(probeSample{responseUs: 20, timestamp: now.Add(-9 * time.Minute)})
task.addSampleLocked(probeSample{responseUs: 40, timestamp: now.Add(-5 * time.Minute)})
task.addSampleLocked(probeSample{responseUs: -1, timestamp: now.Add(-90 * time.Second)})
task.addSampleLocked(probeSample{responseUs: 30, timestamp: now.Add(-30 * time.Second)})
agg := task.aggregateLocked(10*time.Minute, now)
require.True(t, agg.hasData())
assert.Equal(t, int64(4), agg.totalCount)
assert.Equal(t, int64(3), agg.successCount)
result := agg.result()
assert.Equal(t, int64(30), result.AvgResponse)
assert.Equal(t, int64(20), result.MinResponse)
assert.Equal(t, int64(40), result.MaxResponse)
assert.Equal(t, 25.0, result.PacketLoss)
}
func TestProbeTaskAddSampleLockedTrimsRawSamplesButKeepsBucketHistory(t *testing.T) {
now := time.Date(2026, time.April, 21, 12, 0, 0, 0, time.UTC)
task := &probeTask{}
task.addSampleLocked(probeSample{responseUs: 10, timestamp: now.Add(-10 * time.Minute)})
task.addSampleLocked(probeSample{responseUs: 20, timestamp: now})
require.Len(t, task.samples, 1)
assert.Equal(t, int64(20), task.samples[0].responseUs)
agg := task.aggregateLocked(10*time.Minute, now)
require.True(t, agg.hasData())
assert.Equal(t, int64(2), agg.totalCount)
assert.Equal(t, int64(2), agg.successCount)
result := agg.result()
assert.Equal(t, int64(15), result.AvgResponse)
assert.Equal(t, int64(10), result.MinResponse)
assert.Equal(t, int64(20), result.MaxResponse)
assert.Equal(t, 0.0, result.PacketLoss)
}
func TestProbeManagerGetResultsIncludesHourResponseRange(t *testing.T) {
now := time.Now().UTC()
task := &probeTask{config: probe.Config{ID: "probe-1"}}
task.addSampleLocked(probeSample{responseUs: 10, timestamp: now.Add(-30 * time.Minute)})
task.addSampleLocked(probeSample{responseUs: 20, timestamp: now.Add(-9 * time.Minute)})
task.addSampleLocked(probeSample{responseUs: 40, timestamp: now.Add(-5 * time.Minute)})
task.addSampleLocked(probeSample{responseUs: 30, timestamp: now.Add(-50 * time.Second)})
task.addSampleLocked(probeSample{responseUs: -1, timestamp: now.Add(-30 * time.Second)})
pm := &ProbeManager{probes: map[string]*probeTask{"icmp:example.com": task}}
results := pm.GetResults(uint16(time.Minute / time.Millisecond))
result, ok := results["probe-1"]
require.True(t, ok)
assert.Equal(t, int64(30), result.AvgResponse)
assert.Equal(t, int64(25), result.AvgResponse1h)
assert.Equal(t, int64(30), result.MinResponse)
assert.Equal(t, int64(10), result.MinResponse1h)
assert.Equal(t, int64(30), result.MaxResponse)
assert.Equal(t, int64(40), result.MaxResponse1h)
assert.Equal(t, 50.0, result.PacketLoss)
assert.Equal(t, 20.0, result.PacketLoss1h)
}
func TestProbeManagerGetResultsIncludesLossOnlyHourData(t *testing.T) {
now := time.Now().UTC()
task := &probeTask{config: probe.Config{ID: "probe-1"}}
task.addSampleLocked(probeSample{responseUs: -1, timestamp: now.Add(-30 * time.Second)})
task.addSampleLocked(probeSample{responseUs: -1, timestamp: now.Add(-10 * time.Second)})
pm := &ProbeManager{probes: map[string]*probeTask{"icmp:example.com": task}}
results := pm.GetResults(uint16(time.Minute / time.Millisecond))
result, ok := results["probe-1"]
require.True(t, ok)
assert.Equal(t, int64(0), result.AvgResponse)
assert.Equal(t, int64(0), result.AvgResponse1h)
assert.Equal(t, int64(0), result.MinResponse)
assert.Equal(t, int64(0), result.MinResponse1h)
assert.Equal(t, int64(0), result.MaxResponse)
assert.Equal(t, int64(0), result.MaxResponse1h)
assert.Equal(t, 100.0, result.PacketLoss)
assert.Equal(t, 100.0, result.PacketLoss1h)
}
func TestProbeConfigResultKeyUsesSyncedID(t *testing.T) {
cfg := probe.Config{ID: "probe-1", Target: "1.1.1.1", Protocol: "icmp", Interval: 10}
assert.Equal(t, "probe-1", cfg.ID)
}
func TestProbeManagerSyncProbesSkipsConfigsWithoutStableID(t *testing.T) {
validCfg := probe.Config{ID: "probe-1", Target: "ignored", Protocol: "noop", Interval: 10}
invalidCfg := probe.Config{Target: "ignored", Protocol: "noop", Interval: 10}
pm := newProbeManager()
pm.SyncProbes([]probe.Config{validCfg, invalidCfg})
defer pm.Stop()
_, validExists := pm.probes[validCfg.ID]
_, invalidExists := pm.probes[invalidCfg.ID]
assert.True(t, validExists)
assert.False(t, invalidExists)
}
func TestProbeManagerSyncProbesStopsRemovedTasksButKeepsExisting(t *testing.T) {
keepCfg := probe.Config{ID: "probe-1", Target: "ignored", Protocol: "noop", Interval: 10}
removeCfg := probe.Config{ID: "probe-2", Target: "ignored", Protocol: "noop", Interval: 10}
keptTask := &probeTask{config: keepCfg, cancel: make(chan struct{})}
removedTask := &probeTask{config: removeCfg, cancel: make(chan struct{})}
pm := &ProbeManager{
probes: map[string]*probeTask{
keepCfg.ID: keptTask,
removeCfg.ID: removedTask,
},
}
pm.SyncProbes([]probe.Config{keepCfg})
assert.Same(t, keptTask, pm.probes[keepCfg.ID])
_, exists := pm.probes[removeCfg.ID]
assert.False(t, exists)
select {
case <-removedTask.cancel:
default:
t.Fatal("expected removed probe task to be cancelled")
}
select {
case <-keptTask.cancel:
t.Fatal("expected existing probe task to remain active")
default:
}
}
func TestProbeManagerSyncProbesRestartsChangedConfig(t *testing.T) {
originalCfg := probe.Config{ID: "probe-1", Target: "ignored-a", Protocol: "noop", Interval: 10}
updatedCfg := probe.Config{ID: "probe-1", Target: "ignored-b", Protocol: "noop", Interval: 10}
originalTask := &probeTask{config: originalCfg, cancel: make(chan struct{})}
pm := &ProbeManager{
probes: map[string]*probeTask{
originalCfg.ID: originalTask,
},
}
pm.SyncProbes([]probe.Config{updatedCfg})
defer pm.Stop()
restartedTask := pm.probes[updatedCfg.ID]
assert.NotSame(t, originalTask, restartedTask)
assert.Equal(t, updatedCfg, restartedTask.config)
select {
case <-originalTask.cancel:
default:
t.Fatal("expected changed probe task to be cancelled")
}
}
func TestProbeManagerApplySyncUpsertRunsImmediatelyAndReturnsResult(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
defer server.Close()
pm := &ProbeManager{
probes: make(map[string]*probeTask),
httpClient: server.Client(),
}
resp, err := pm.HandleSyncRequest(probe.SyncRequest{
Action: probe.SyncActionUpsert,
Config: probe.Config{ID: "probe-1", Target: server.URL, Protocol: "http", Interval: 10},
RunNow: true,
})
defer pm.Stop()
require.NoError(t, err)
assert.GreaterOrEqual(t, resp.Result.AvgResponse, int64(0))
assert.Equal(t, 0.0, resp.Result.PacketLoss)
assert.Equal(t, 0.0, resp.Result.PacketLoss1h)
task := pm.probes["probe-1"]
require.NotNil(t, task)
task.mu.Lock()
defer task.mu.Unlock()
require.Len(t, task.samples, 1)
}
func TestProbeManagerUpsertProbeKeepsHistoryWhenOnlyIntervalChanges(t *testing.T) {
originalCfg := probe.Config{ID: "probe-1", Target: "1.1.1.1", Protocol: "icmp", Interval: 10}
updatedCfg := probe.Config{ID: "probe-1", Target: "1.1.1.1", Protocol: "icmp", Interval: 30}
now := time.Now().UTC()
existingTask := &probeTask{config: originalCfg, cancel: make(chan struct{})}
existingTask.addSampleLocked(probeSample{responseUs: 12, timestamp: now.Add(-50 * time.Minute)})
existingTask.addSampleLocked(probeSample{responseUs: 24, timestamp: now.Add(-30 * time.Second)})
pm := &ProbeManager{
probes: map[string]*probeTask{originalCfg.ID: existingTask},
}
result, err := pm.UpsertProbe(updatedCfg, false)
defer pm.Stop()
require.NoError(t, err)
assert.Nil(t, result)
updatedTask := pm.probes[updatedCfg.ID]
require.NotNil(t, updatedTask)
assert.NotSame(t, existingTask, updatedTask)
assert.Equal(t, updatedCfg, updatedTask.config)
updatedTask.mu.Lock()
defer updatedTask.mu.Unlock()
require.Len(t, updatedTask.samples, 1)
assert.Equal(t, int64(24), updatedTask.samples[0].responseUs)
agg := updatedTask.aggregateLocked(time.Hour, now)
require.True(t, agg.hasData())
assert.Equal(t, int64(2), agg.totalCount)
assert.Equal(t, int64(2), agg.successCount)
assert.Equal(t, int64(18), agg.avgResponse())
select {
case <-existingTask.cancel:
default:
t.Fatal("expected original probe task to be cancelled")
}
}
func TestProbeManagerApplySyncDeleteRemovesTask(t *testing.T) {
config := probe.Config{ID: "probe-1", Target: "1.1.1.1", Protocol: "icmp", Interval: 10}
task := &probeTask{config: config, cancel: make(chan struct{})}
pm := &ProbeManager{
probes: map[string]*probeTask{config.ID: task},
}
_, err := pm.HandleSyncRequest(probe.SyncRequest{
Action: probe.SyncActionDelete,
Config: probe.Config{ID: config.ID},
})
require.NoError(t, err)
_, exists := pm.probes[config.ID]
assert.False(t, exists)
select {
case <-task.cancel:
default:
t.Fatal("expected deleted probe task to be cancelled")
}
}
func TestProbeManagerGetRandomDelay(t *testing.T) {
for i := 1000; i < 360_000; i += 1000 {
delay := getStagger(int64(i))
assert.GreaterOrEqual(t, delay, time.Duration(i/2)*time.Millisecond)
assert.LessOrEqual(t, delay, time.Duration(i)*time.Millisecond)
}
}
func TestProbeHTTP(t *testing.T) {
t.Run("success", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
defer server.Close()
responseUs, err := probeHTTP(server.Client(), server.URL)
require.NoError(t, err)
assert.GreaterOrEqual(t, responseUs, int64(0))
})
t.Run("server error", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "boom", http.StatusInternalServerError)
}))
defer server.Close()
responseUs, err := probeHTTP(server.Client(), server.URL)
assert.Equal(t, int64(-1), responseUs)
require.Error(t, err)
})
}
func TestProbeTCP(t *testing.T) {
t.Run("success", func(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer listener.Close()
accepted := make(chan struct{})
go func() {
defer close(accepted)
conn, err := listener.Accept()
if err == nil {
_ = conn.Close()
}
}()
port := uint16(listener.Addr().(*net.TCPAddr).Port)
responseUs, err := probeTCP("127.0.0.1", port)
require.NoError(t, err)
assert.GreaterOrEqual(t, responseUs, int64(0))
<-accepted
})
t.Run("connection failure", func(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
port := uint16(listener.Addr().(*net.TCPAddr).Port)
require.NoError(t, listener.Close())
responseUs, err := probeTCP("127.0.0.1", port)
assert.Equal(t, int64(-1), responseUs)
require.Error(t, err)
})
}

View File

@@ -21,9 +21,6 @@ func newAgentResponse(data any, requestID *uint32) common.AgentResponse {
response.String = &v
case map[string]smart.SmartData:
response.SmartData = v
case smart.SmartDataResponse:
response.SmartData = v.Data
response.SmartComplete = v.Complete
case systemd.ServiceDetails:
response.ServiceInfo = v
default:

View File

@@ -5,9 +5,7 @@ import (
"errors"
"fmt"
"log/slog"
"os"
"path"
"path/filepath"
"runtime"
"strconv"
"strings"
@@ -34,8 +32,6 @@ type SensorConfig struct {
isBlacklist bool
hasWildcards bool
skipCollection bool
skipGPU bool
sensorShadow string
firstRun bool
}
@@ -45,14 +41,13 @@ func (a *Agent) newSensorConfig() *SensorConfig {
sensorsEnvVal, sensorsSet := utils.GetEnv("SENSORS")
skipCollection := sensorsSet && sensorsEnvVal == ""
sensorsTimeout, _ := utils.GetEnv("SENSORS_TIMEOUT")
skipGPU, _ := utils.GetEnv("SKIP_GPU")
return a.newSensorConfigWithEnv(primarySensor, sysSensors, sensorsEnvVal, sensorsTimeout, skipCollection, skipGPU == "true")
return a.newSensorConfigWithEnv(primarySensor, sysSensors, sensorsEnvVal, sensorsTimeout, skipCollection)
}
// newSensorConfigWithEnv creates a SensorConfig with the provided environment variables
// sensorsSet indicates if the SENSORS environment variable was explicitly set (even to empty string)
func (a *Agent) newSensorConfigWithEnv(primarySensor, sysSensors, sensorsEnvVal, sensorsTimeout string, skipCollection, skipGPU bool) *SensorConfig {
func (a *Agent) newSensorConfigWithEnv(primarySensor, sysSensors, sensorsEnvVal, sensorsTimeout string, skipCollection bool) *SensorConfig {
timeout := 2 * time.Second
if sensorsTimeout != "" {
if d, err := time.ParseDuration(sensorsTimeout); err == nil {
@@ -67,7 +62,6 @@ func (a *Agent) newSensorConfigWithEnv(primarySensor, sysSensors, sensorsEnvVal,
primarySensor: primarySensor,
timeout: timeout,
skipCollection: skipCollection,
skipGPU: skipGPU,
firstRun: true,
sensors: make(map[string]struct{}),
}
@@ -79,19 +73,6 @@ func (a *Agent) newSensorConfigWithEnv(primarySensor, sysSensors, sensorsEnvVal,
common.EnvKey, common.EnvMap{common.HostSysEnvKey: sysSensors},
)
}
if skipGPU && runtime.GOOS == "linux" {
// gopsutil reads every temp*_input before results can be filtered, so
// point it at a shadow tree built from the effective sysfs root instead.
if shadow, err := buildNonGpuSysShadow(effectiveSysRoot(config.context)); err == nil {
slog.Info("SKIP_GPU enabled, using non-GPU sensor sysfs shadow", "path", shadow)
config.sensorShadow = shadow
config.context = context.WithValue(config.context,
common.EnvKey, common.EnvMap{common.HostSysEnvKey: shadow},
)
} else {
slog.Warn("SKIP_GPU sensor shadow unavailable, falling back to post-read filtering", "err", err)
}
}
// handle blacklist
if strings.HasPrefix(sensorsEnvVal, "-") {
@@ -168,9 +149,6 @@ func (a *Agent) updateTemperatures(systemStats *system.Stats) {
if !isValidSensor(sensorName, a.sensorConfig) {
continue
}
if a.sensorConfig.skipGPU && isGpuSensorKey(sensorName) {
continue
}
// set dashboard temperature
switch a.sensorConfig.primarySensor {
case "":
@@ -267,102 +245,3 @@ func scaleTemperature(temp float64) float64 {
}
return scaled100
}
// effectiveSysRoot mirrors gopsutil's HostSys lookup, which lives in its
// internal package: context override, then HOST_SYS env, then /sys.
func effectiveSysRoot(ctx context.Context) string {
if envMap, ok := ctx.Value(common.EnvKey).(common.EnvMap); ok {
if v := envMap[common.HostSysEnvKey]; v != "" {
return v
}
}
if v := os.Getenv("HOST_SYS"); v != "" {
return v
}
return "/sys"
}
func (config *SensorConfig) cleanupSensorShadow() {
if config.sensorShadow == "" {
return
}
if err := os.RemoveAll(config.sensorShadow); err != nil {
slog.Warn("Error removing sensor sysfs shadow", "path", config.sensorShadow, "err", err)
return
}
config.sensorShadow = ""
}
func (a *Agent) cleanupSensorShadow() {
if a.sensorConfig != nil {
a.sensorConfig.cleanupSensorShadow()
}
}
func isGpuThermalZone(zoneType string) bool {
zoneType = strings.ToLower(strings.TrimSpace(zoneType))
return isGpuChipName(zoneType) || strings.Contains(zoneType, "gpu")
}
// buildNonGpuSysShadow links non-GPU sensor directories into a temp dir. Only
// static chip names and thermal-zone types are read; no sensor values are touched.
func buildNonGpuSysShadow(sysRoot string) (string, error) {
shadow, err := os.MkdirTemp("", "beszel-sensors-*")
if err != nil {
return "", err
}
shadowHwmon := filepath.Join(shadow, "class", "hwmon")
if err := os.MkdirAll(shadowHwmon, 0o755); err != nil {
os.RemoveAll(shadow)
return "", err
}
entries, err := os.ReadDir(filepath.Join(sysRoot, "class", "hwmon"))
if err != nil && !os.IsNotExist(err) {
os.RemoveAll(shadow)
return "", err
}
for _, entry := range entries {
chipDir := filepath.Join(sysRoot, "class", "hwmon", entry.Name())
// Some hwmon devices expose name under device/ (gopsutil's CentOS fallback).
name, ok := utils.ReadStringFileOK(filepath.Join(chipDir, "name"))
if !ok {
name, ok = utils.ReadStringFileOK(filepath.Join(chipDir, "device", "name"))
}
if !ok || isGpuChipName(name) {
continue
}
if err := os.Symlink(chipDir, filepath.Join(shadowHwmon, entry.Name())); err != nil {
os.RemoveAll(shadow)
return "", err
}
}
thermalEntries, err := os.ReadDir(filepath.Join(sysRoot, "class", "thermal"))
if err != nil {
if os.IsNotExist(err) {
return shadow, nil
}
os.RemoveAll(shadow)
return "", err
}
shadowThermal := filepath.Join(shadow, "class", "thermal")
if err := os.MkdirAll(shadowThermal, 0o755); err != nil {
os.RemoveAll(shadow)
return "", err
}
for _, entry := range thermalEntries {
if !strings.HasPrefix(entry.Name(), "thermal_zone") {
continue
}
zoneDir := filepath.Join(sysRoot, "class", "thermal", entry.Name())
zoneType, ok := utils.ReadStringFileOK(filepath.Join(zoneDir, "type"))
if !ok || isGpuThermalZone(zoneType) {
continue
}
if err := os.Symlink(zoneDir, filepath.Join(shadowThermal, entry.Name())); err != nil {
os.RemoveAll(shadow)
return "", err
}
}
return shadow, nil
}

View File

@@ -1,4 +1,4 @@
//go:build !windows && !freebsd
//go:build !windows
package agent

View File

@@ -1,14 +0,0 @@
//go:build freebsd
package agent
import (
"context"
"github.com/shirou/gopsutil/v4/sensors"
"golang.org/x/sys/unix"
)
var getSensorTemps = func(ctx context.Context) ([]sensors.TemperatureStat, error) {
return getFreeBSDSensorTemps(ctx, unix.SysctlUint32)
}

View File

@@ -1,81 +0,0 @@
//go:build freebsd || testing
package agent
import (
"context"
"fmt"
"github.com/shirou/gopsutil/v4/sensors"
)
const (
freebsdZeroCelsiusDeciKelvin = 2731
freebsdAcpiThermalZoneCount = 16
)
type freebsdSysctlUintReader func(name string) (uint32, error)
func getFreeBSDSensorTemps(ctx context.Context, readSysctl freebsdSysctlUintReader) ([]sensors.TemperatureStat, error) {
cpuCount, err := readSysctl("hw.ncpu")
if err != nil {
return nil, err
}
temps := make([]sensors.TemperatureStat, 0, int(cpuCount)+freebsdAcpiThermalZoneCount)
for cpu := range cpuCount {
select {
case <-ctx.Done():
return temps, ctx.Err()
default:
}
sysctlName := fmt.Sprintf("dev.cpu.%d.temperature", cpu)
value, err := readSysctl(sysctlName)
if err != nil {
continue
}
temp, ok := freebsdDeciKelvinToCelsius(value)
if !ok {
continue
}
temps = append(temps, sensors.TemperatureStat{
SensorKey: fmt.Sprintf("cpu.%d", cpu),
Temperature: temp,
})
}
for zone := 0; zone < freebsdAcpiThermalZoneCount; zone++ {
select {
case <-ctx.Done():
return temps, ctx.Err()
default:
}
sysctlName := fmt.Sprintf("hw.acpi.thermal.tz%d.temperature", zone)
value, err := readSysctl(sysctlName)
if err != nil {
continue
}
temp, ok := freebsdDeciKelvinToCelsius(value)
if !ok {
continue
}
temps = append(temps, sensors.TemperatureStat{
SensorKey: fmt.Sprintf("acpi.thermal.tz%d", zone),
Temperature: temp,
})
}
return temps, nil
}
func freebsdDeciKelvinToCelsius(value uint32) (float64, bool) {
if value <= freebsdZeroCelsiusDeciKelvin {
return 0, false
}
temp := float64(int64(value)-freebsdZeroCelsiusDeciKelvin) / 10
if temp <= 0 || temp >= 200 {
return 0, false
}
return temp, true
}

View File

@@ -1,167 +0,0 @@
//go:build testing
package agent
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
var errFakeFreeBSDSysctlNotFound = errors.New("sysctl not found")
type fakeFreeBSDSysctls struct {
values map[string]uint32
errs map[string]error
}
func (f fakeFreeBSDSysctls) read(name string) (uint32, error) {
if err, ok := f.errs[name]; ok {
return 0, err
}
if value, ok := f.values[name]; ok {
return value, nil
}
return 0, errFakeFreeBSDSysctlNotFound
}
func TestFreeBSDDeciKelvinToCelsius(t *testing.T) {
tests := []struct {
name string
value uint32
expected float64
ok bool
}{
{
name: "45 Celsius",
value: 3181,
expected: 45,
ok: true,
},
{
name: "fractional Celsius",
value: 3186,
expected: 45.5,
ok: true,
},
{
name: "zero deci-Kelvin",
value: 0,
ok: false,
},
{
name: "zero Celsius",
value: freebsdZeroCelsiusDeciKelvin,
ok: false,
},
{
name: "below zero Celsius",
value: freebsdZeroCelsiusDeciKelvin - 1,
ok: false,
},
{
name: "invalid signed integer",
value: 1<<32 - 1,
ok: false,
},
{
name: "unreasonably high Celsius",
value: freebsdZeroCelsiusDeciKelvin + 2000,
ok: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result, ok := freebsdDeciKelvinToCelsius(tt.value)
assert.Equal(t, tt.ok, ok)
assert.InDelta(t, tt.expected, result, 0.001)
})
}
}
func TestGetFreeBSDSensorTemps(t *testing.T) {
reader := fakeFreeBSDSysctls{
values: map[string]uint32{
"hw.ncpu": 4,
"dev.cpu.0.temperature": 3231,
"dev.cpu.1.temperature": 3242,
"dev.cpu.3.temperature": freebsdZeroCelsiusDeciKelvin,
"hw.acpi.thermal.tz0.temperature": 3101,
"hw.acpi.thermal.tz2.temperature": 3116,
"hw.acpi.thermal.tz3.temperature": freebsdZeroCelsiusDeciKelvin,
"unrelated.sensor.value": 9999,
"dev.cpu.99.temperature": 9999,
"dev.amdtemp.0.core0.foo": 9999,
},
}
temps, err := getFreeBSDSensorTemps(context.Background(), reader.read)
require.NoError(t, err)
require.Len(t, temps, 4)
assert.Equal(t, "cpu.0", temps[0].SensorKey)
assert.InDelta(t, 50.0, temps[0].Temperature, 0.001)
assert.Equal(t, "cpu.1", temps[1].SensorKey)
assert.InDelta(t, 51.1, temps[1].Temperature, 0.001)
assert.Equal(t, "acpi.thermal.tz0", temps[2].SensorKey)
assert.InDelta(t, 37.0, temps[2].Temperature, 0.001)
assert.Equal(t, "acpi.thermal.tz2", temps[3].SensorKey)
assert.InDelta(t, 38.5, temps[3].Temperature, 0.001)
}
func TestGetFreeBSDSensorTempsCpuCountError(t *testing.T) {
reader := fakeFreeBSDSysctls{
errs: map[string]error{
"hw.ncpu": errors.New("permission denied"),
},
}
temps, err := getFreeBSDSensorTemps(context.Background(), reader.read)
assert.Nil(t, temps)
assert.EqualError(t, err, "permission denied")
}
func TestGetFreeBSDSensorTempsNoTemperatureSysctls(t *testing.T) {
reader := fakeFreeBSDSysctls{
values: map[string]uint32{"hw.ncpu": 2},
}
temps, err := getFreeBSDSensorTemps(context.Background(), reader.read)
require.NoError(t, err)
assert.Empty(t, temps)
}
func TestGetFreeBSDSensorTempsAcpiOnly(t *testing.T) {
reader := fakeFreeBSDSysctls{
values: map[string]uint32{
"hw.ncpu": 0,
"hw.acpi.thermal.tz0.temperature": 3081,
},
}
temps, err := getFreeBSDSensorTemps(context.Background(), reader.read)
require.NoError(t, err)
require.Len(t, temps, 1)
assert.Equal(t, "acpi.thermal.tz0", temps[0].SensorKey)
assert.InDelta(t, 35.0, temps[0].Temperature, 0.001)
}
func TestGetFreeBSDSensorTempsContextCancelled(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
reader := fakeFreeBSDSysctls{
values: map[string]uint32{"hw.ncpu": 2},
}
temps, err := getFreeBSDSensorTemps(ctx, reader.read)
assert.Empty(t, temps)
assert.ErrorIs(t, err, context.Canceled)
}

View File

@@ -5,9 +5,6 @@ package agent
import (
"context"
"fmt"
"os"
"path/filepath"
"runtime"
"testing"
"time"
@@ -331,7 +328,7 @@ func TestNewSensorConfigWithEnv(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := agent.newSensorConfigWithEnv(tt.primarySensor, tt.sysSensors, tt.sensors, tt.sensorsTimeout, tt.skipCollection, false)
result := agent.newSensorConfigWithEnv(tt.primarySensor, tt.sysSensors, tt.sensors, tt.sensorsTimeout, tt.skipCollection)
// Check primary sensor
assert.Equal(t, tt.expectedConfig.primarySensor, result.primarySensor)
@@ -605,9 +602,8 @@ func TestUpdateTemperaturesSkipsOnTimeout(t *testing.T) {
},
}
originalGetSensorTemps := getSensorTemps
t.Cleanup(func() {
getSensorTemps = originalGetSensorTemps
getSensorTemps = sensors.TemperaturesWithContext
})
getSensorTemps = func(ctx context.Context) ([]sensors.TemperatureStat, error) {
time.Sleep(50 * time.Millisecond)
@@ -623,153 +619,3 @@ func TestUpdateTemperaturesSkipsOnTimeout(t *testing.T) {
assert.Equal(t, 0.0, agent.systemInfo.DashboardTemp)
assert.Equal(t, map[string]float64{}, stats.Temperatures)
}
func TestIsGpuSensorKey(t *testing.T) {
for _, key := range []string{"xe", "XE_temp1", "amdgpu_edge", "NVIDIA"} {
assert.True(t, isGpuSensorKey(key), key)
}
for _, key := range []string{"coretemp_core_0", "acpitz", "xen_temp", "myxe", ""} {
assert.False(t, isGpuSensorKey(key), key)
}
}
func TestSkipGpuSensorShadow(t *testing.T) {
sysRoot := t.TempDir()
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon0", "name"), "coretemp\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon0", "temp1_input"), "55000\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon1", "name"), "xe\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon1", "temp1_input"), "48000\n")
writeFile(t, filepath.Join(sysRoot, "class", "thermal", "thermal_zone0", "type"), "cpu-thermal\n")
writeFile(t, filepath.Join(sysRoot, "class", "thermal", "thermal_zone0", "temp"), "55000\n")
writeFile(t, filepath.Join(sysRoot, "class", "thermal", "thermal_zone1", "type"), "gpu\n")
writeFile(t, filepath.Join(sysRoot, "class", "thermal", "thermal_zone1", "temp"), "48000\n")
shadow, err := buildNonGpuSysShadow(sysRoot)
require.NoError(t, err)
t.Cleanup(func() { os.RemoveAll(shadow) })
assert.FileExists(t, filepath.Join(shadow, "class", "hwmon", "hwmon0", "temp1_input"))
assert.NoFileExists(t, filepath.Join(shadow, "class", "hwmon", "hwmon1"))
assert.FileExists(t, filepath.Join(shadow, "class", "thermal", "thermal_zone0", "temp"))
assert.NoFileExists(t, filepath.Join(shadow, "class", "thermal", "thermal_zone1"))
}
func TestSkipGpuSensorShadowDeviceName(t *testing.T) {
sysRoot := t.TempDir()
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon0", "device", "name"), "coretemp\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon0", "device", "temp1_input"), "55000\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon1", "device", "name"), "xe\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon1", "device", "temp1_input"), "48000\n")
shadow, err := buildNonGpuSysShadow(sysRoot)
require.NoError(t, err)
t.Cleanup(func() { os.RemoveAll(shadow) })
assert.FileExists(t, filepath.Join(shadow, "class", "hwmon", "hwmon0", "device", "temp1_input"))
assert.NoFileExists(t, filepath.Join(shadow, "class", "hwmon", "hwmon1"))
}
func TestSkipGpuSensorShadowKeepsThermalZonesWithoutNonGpuHwmon(t *testing.T) {
sysRoot := t.TempDir()
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon0", "name"), "xe\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon0", "temp1_input"), "48000\n")
writeFile(t, filepath.Join(sysRoot, "class", "thermal", "thermal_zone0", "type"), "cpu-thermal\n")
writeFile(t, filepath.Join(sysRoot, "class", "thermal", "thermal_zone0", "temp"), "55000\n")
writeFile(t, filepath.Join(sysRoot, "class", "thermal", "thermal_zone1", "type"), "gpu\n")
writeFile(t, filepath.Join(sysRoot, "class", "thermal", "thermal_zone1", "temp"), "48000\n")
shadow, err := buildNonGpuSysShadow(sysRoot)
require.NoError(t, err)
t.Cleanup(func() { os.RemoveAll(shadow) })
hwmonTemps, err := filepath.Glob(filepath.Join(shadow, "class", "hwmon", "hwmon*", "temp*_input"))
require.NoError(t, err)
assert.Empty(t, hwmonTemps)
assert.FileExists(t, filepath.Join(shadow, "class", "thermal", "thermal_zone0", "temp"))
assert.NoFileExists(t, filepath.Join(shadow, "class", "thermal", "thermal_zone1"))
}
func TestNewSensorConfigSkipGpuWiresShadow(t *testing.T) {
t.Setenv("SKIP_GPU", "true")
agent := &Agent{}
config := agent.newSensorConfig()
assert.True(t, config.skipGPU)
if runtime.GOOS != "linux" {
assert.Empty(t, config.sensorShadow)
assert.Nil(t, config.context.Value(common.EnvKey))
return
}
envMap, ok := config.context.Value(common.EnvKey).(common.EnvMap)
require.True(t, ok, "SKIP_GPU should point the sensor context at a sysfs shadow")
shadow, ok := envMap[common.HostSysEnvKey]
require.True(t, ok)
assert.DirExists(t, filepath.Join(shadow, "class", "hwmon"))
assert.Equal(t, shadow, config.sensorShadow)
config.cleanupSensorShadow()
assert.NoDirExists(t, shadow)
assert.Empty(t, config.sensorShadow)
}
func TestSkipGpuShadowUsesSysSensorsRoot(t *testing.T) {
sysRoot := t.TempDir()
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon0", "name"), "coretemp\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon0", "temp1_input"), "55000\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon1", "name"), "xe\n")
writeFile(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon1", "temp1_input"), "48000\n")
agent := &Agent{}
config := agent.newSensorConfigWithEnv("", sysRoot, "", "", false, true)
t.Cleanup(config.cleanupSensorShadow)
envMap, ok := config.context.Value(common.EnvKey).(common.EnvMap)
require.True(t, ok, "SKIP_GPU should point the sensor context at a sysfs shadow")
shadow, ok := envMap[common.HostSysEnvKey]
require.True(t, ok)
if runtime.GOOS != "linux" {
assert.Equal(t, sysRoot, shadow)
assert.Empty(t, config.sensorShadow)
return
}
require.NotEqual(t, sysRoot, shadow, "shadow must not be the SYS_SENSORS tree itself")
target, err := os.Readlink(filepath.Join(shadow, "class", "hwmon", "hwmon0"))
require.NoError(t, err)
assert.Equal(t, filepath.Join(sysRoot, "class", "hwmon", "hwmon0"), target)
assert.NoFileExists(t, filepath.Join(shadow, "class", "hwmon", "hwmon1"))
}
func TestUpdateTemperaturesSkipGpu(t *testing.T) {
originalGetSensorTemps := getSensorTemps
t.Cleanup(func() {
getSensorTemps = originalGetSensorTemps
})
getSensorTemps = func(ctx context.Context) ([]sensors.TemperatureStat, error) {
return []sensors.TemperatureStat{
{SensorKey: "coretemp_core_0", Temperature: 55},
{SensorKey: "XE", Temperature: 48},
}, nil
}
newAgent := func(skipGPU bool) *Agent {
agent := &Agent{
systemInfo: system.Info{},
sensorConfig: &SensorConfig{
context: context.Background(),
timeout: 2 * time.Second,
sensors: map[string]struct{}{},
skipGPU: skipGPU,
},
}
return agent
}
stats := &system.Stats{}
newAgent(true).updateTemperatures(stats)
assert.Equal(t, map[string]float64{"coretemp_core_0": 55}, stats.Temperatures)
stats = &system.Stats{}
newAgent(false).updateTemperatures(stats)
assert.Len(t, stats.Temperatures, 2)
}

View File

@@ -214,12 +214,9 @@ func (lhm *lhmProcess) getTemps(ctx context.Context) (temps []sensors.Temperatur
return temps, nil
}
// getSensorTemps is a variable so tests can replace the platform sensor collector.
var getSensorTemps = getWindowsSensorTemps
// getWindowsSensorTemps attempts to pull sensor temperatures from the embedded LHM process.
// getSensorTemps attempts to pull sensor temperatures from the embedded LHM process.
// NB: LibreHardwareMonitorLib requires admin privileges to access all available sensors.
func getWindowsSensorTemps(ctx context.Context) (temps []sensors.TemperatureStat, err error) {
func getSensorTemps(ctx context.Context) (temps []sensors.TemperatureStat, err error) {
defer func() {
if err != nil {
slog.Debug("Error reading sensors", "err", err)

View File

@@ -29,31 +29,19 @@ type ServerOptions struct {
Keys []gossh.PublicKey // SSH public keys for authentication
}
// hubVersions caches hub versions by session ID to avoid repeated parsing.
var hubVersions map[string]semver.Version
// StartServer starts the SSH server with the provided options.
// It configures the server with secure defaults, sets up authentication,
// and begins listening for connections. Returns an error if the server
// is already running or if there's an issue starting the server.
func (a *Agent) StartServer(opts ServerOptions) error {
server, listener, err := a.prepareSSHServer(opts)
if err != nil {
return err
}
return a.serveSSHServer(server, listener)
}
var errSSHServerRunning = errors.New("server already started")
// prepareSSHServer binds the listener before Serve starts so a concurrent stop
// can always close it, including when the WebSocket wins the connection race.
func (a *Agent) prepareSSHServer(opts ServerOptions) (*ssh.Server, net.Listener, error) {
a.serverMu.Lock()
defer a.serverMu.Unlock()
if disableSSH, _ := utils.GetEnv("DISABLE_SSH"); disableSSH == "true" {
return nil, nil, errors.New("SSH disabled")
return errors.New("SSH disabled")
}
if a.server != nil {
return nil, nil, errSSHServerRunning
return errors.New("server already started")
}
slog.Info("Starting SSH server", "addr", opts.Addr, "network", opts.Network)
@@ -61,19 +49,32 @@ func (a *Agent) prepareSSHServer(opts ServerOptions) (*ssh.Server, net.Listener,
if opts.Network == "unix" {
// remove existing socket file if it exists
if err := os.Remove(opts.Addr); err != nil && !os.IsNotExist(err) {
return nil, nil, err
return err
}
}
// start listening on the address
ln, err := net.Listen(opts.Network, opts.Addr)
if err != nil {
return nil, nil, err
return err
}
defer ln.Close()
server := &ssh.Server{
Handler: a.handleSession,
ServerConfigCallback: newSSHServerConfig,
// base config (limit to allowed algorithms)
config := &gossh.ServerConfig{
ServerVersion: fmt.Sprintf("SSH-2.0-%s_%s", beszel.AppName, beszel.Version),
}
config.KeyExchanges = common.DefaultKeyExchanges
config.MACs = common.DefaultMACs
config.Ciphers = common.DefaultCiphers
// set default handler
ssh.Handle(a.handleSession)
a.server = &ssh.Server{
ServerConfigCallback: func(ctx ssh.Context) *gossh.ServerConfig {
return config
},
// check public key(s)
PublicKeyHandler: func(ctx ssh.Context, key ssh.PublicKey) bool {
remoteAddr := ctx.RemoteAddr()
@@ -94,44 +95,28 @@ func (a *Agent) prepareSSHServer(opts ServerOptions) (*ssh.Server, net.Listener,
IdleTimeout: 70 * time.Second,
}
a.server = server
a.serverListener = ln
return server, ln, nil
// Start SSH server on the listener
return a.server.Serve(ln)
}
func (a *Agent) serveSSHServer(server *ssh.Server, listener net.Listener) error {
err := server.Serve(listener)
a.serverMu.Lock()
if a.server == server {
a.server = nil
a.serverListener = nil
// getHubVersion retrieves and caches the hub version for a given session.
// It extracts the version from the SSH client version string and caches
// it to avoid repeated parsing. Returns a zero version if parsing fails.
func (a *Agent) getHubVersion(sessionId string, sessionCtx ssh.Context) semver.Version {
if hubVersions == nil {
hubVersions = make(map[string]semver.Version, 1)
}
a.serverMu.Unlock()
return err
}
// newSSHServerConfig returns a separate config for each connection because
// gliderlabs adds host keys and connection-specific callbacks to it.
func newSSHServerConfig(ssh.Context) *gossh.ServerConfig {
return &gossh.ServerConfig{
Config: gossh.Config{
KeyExchanges: common.DefaultKeyExchanges,
MACs: common.DefaultMACs,
Ciphers: common.DefaultCiphers,
},
ServerVersion: fmt.Sprintf("SSH-2.0-%s_%s", beszel.AppName, beszel.Version),
}
}
// getHubVersion extracts the hub version from the SSH client version string
// for a given session. Returns a zero version if parsing fails.
func (a *Agent) getHubVersion(sessionCtx ssh.Context) semver.Version {
clientVersion := sessionCtx.Value(ssh.ContextKeyClientVersion)
if versionStr, ok := clientVersion.(string); ok {
hubVersion, _ := extractHubVersion(versionStr)
hubVersion, ok := hubVersions[sessionId]
if ok {
return hubVersion
}
return semver.Version{}
// Extract hub version from SSH client version
clientVersion := sessionCtx.Value(ssh.ContextKeyClientVersion)
if versionStr, ok := clientVersion.(string); ok {
hubVersion, _ = extractHubVersion(versionStr)
}
hubVersions[sessionId] = hubVersion
return hubVersion
}
// handleSession handles an incoming SSH session by gathering system statistics
@@ -139,10 +124,12 @@ func (a *Agent) getHubVersion(sessionCtx ssh.Context) semver.Version {
// appropriate encoding format based on hub version, and exits with appropriate
// status codes.
func (a *Agent) handleSession(s ssh.Session) {
sessionCtx := s.Context()
a.connectionManager.sshConnectionOpened(sessionCtx)
a.connectionManager.eventChan <- SSHConnect
hubVersion := a.getHubVersion(sessionCtx)
sessionCtx := s.Context()
sessionID := sessionCtx.SessionID()
hubVersion := a.getHubVersion(sessionID, sessionCtx)
// Legacy one-shot behavior for older hubs
if hubVersion.LT(beszel.MinVersionAgentResponse) {
@@ -187,13 +174,12 @@ func (a *Agent) handleSSHRequest(w io.Writer, req *common.HubRequest[cbor.RawMes
}
ctx := &HandlerContext{
Client: nil,
Agent: a,
Request: req,
RequestID: nil,
HubVerified: true,
ConnectionType: system.ConnectionTypeSSH,
SendResponse: sshResponder,
Client: nil,
Agent: a,
Request: req,
RequestID: nil,
HubVerified: true,
SendResponse: sshResponder,
}
if handler, ok := a.handlerRegistry.GetHandler(req.Action); ok {
@@ -208,9 +194,7 @@ func (a *Agent) handleSSHRequest(w io.Writer, req *common.HubRequest[cbor.RawMes
// handleLegacyStats serves the legacy one-shot stats payload for older hubs
func (a *Agent) handleLegacyStats(w io.Writer, hubVersion semver.Version) error {
stats := a.gatherStats(common.DataRequestOptions{CacheTimeMs: defaultDataCacheTimeMs})
response := *stats
response.Info.ConnectionType = system.ConnectionTypeSSH
return a.writeToSession(w, &response, hubVersion)
return a.writeToSession(w, stats, hubVersion)
}
// writeToSession encodes and writes system statistics to the session.
@@ -287,21 +271,13 @@ func GetNetwork(addr string) string {
// StopServer stops the SSH server if it's running.
// It returns an error if the server is not running or if there's an error stopping it.
func (a *Agent) StopServer() error {
a.serverMu.Lock()
if a.server == nil {
a.serverMu.Unlock()
return errors.New("SSH server not running")
}
server := a.server
listener := a.serverListener
a.server = nil
a.serverListener = nil
a.serverMu.Unlock()
slog.Info("Stopping SSH server")
if listener != nil {
_ = listener.Close()
}
_ = server.Close()
_ = a.server.Close()
a.server = nil
a.connectionManager.eventChan <- SSHDisconnect
return nil
}

View File

@@ -1,111 +0,0 @@
//go:build testing
package agent
import (
"crypto/ed25519"
"fmt"
"net"
"sync"
"testing"
"time"
"github.com/henrygd/beszel"
"github.com/henrygd/beszel/internal/common"
"github.com/gliderlabs/ssh"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
gossh "golang.org/x/crypto/ssh"
)
func TestSSHServerConfigConcurrentConnections(t *testing.T) {
_, key, err := ed25519.GenerateKey(nil)
require.NoError(t, err)
signer, err := gossh.NewSignerFromKey(key)
require.NoError(t, err)
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
const connections = 2
configs := make(chan *gossh.ServerConfig, connections)
release := make(chan struct{})
var releaseOnce sync.Once
unblock := func() { releaseOnce.Do(func() { close(release) }) }
server := &ssh.Server{
HostSigners: []ssh.Signer{signer},
ServerConfigCallback: func(ctx ssh.Context) *gossh.ServerConfig {
config := newSSHServerConfig(ctx)
configs <- config
// Both connections must obtain their configuration before either
// lets gliderlabs add host keys and connection-specific callbacks.
<-release
return config
},
PublicKeyHandler: func(_ ssh.Context, key ssh.PublicKey) bool {
return ssh.KeysEqual(key, signer.PublicKey())
},
Handler: func(session ssh.Session) { _ = session.Exit(0) },
}
served := make(chan error, 1)
go func() { served <- server.Serve(listener) }()
t.Cleanup(func() {
unblock()
_ = listener.Close()
_ = server.Close()
select {
case <-served:
case <-time.After(5 * time.Second):
t.Error("SSH test server did not stop")
}
})
results := make(chan error, connections)
for range connections {
go func() {
conn, err := net.DialTimeout("tcp", listener.Addr().String(), 5*time.Second)
if err != nil {
results <- err
return
}
defer conn.Close()
_ = conn.SetDeadline(time.Now().Add(5 * time.Second))
client, _, _, err := gossh.NewClientConn(conn, listener.Addr().String(), &gossh.ClientConfig{
User: "test",
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
HostKeyCallback: gossh.FixedHostKey(signer.PublicKey()),
})
if err == nil {
err = client.Close()
}
results <- err
}()
}
var first *gossh.ServerConfig
for range connections {
select {
case config := <-configs:
assert.Equal(t, fmt.Sprintf("SSH-2.0-%s_%s", beszel.AppName, beszel.Version), config.ServerVersion)
assert.Equal(t, common.DefaultKeyExchanges, config.KeyExchanges)
assert.Equal(t, common.DefaultMACs, config.MACs)
assert.Equal(t, common.DefaultCiphers, config.Ciphers)
if first == nil {
first = config
} else {
assert.NotSame(t, first, config, "SSH connections must not share mutable configuration")
}
case <-time.After(5 * time.Second):
t.Fatal("SSH connections did not reach their config callbacks")
}
}
unblock()
for range connections {
select {
case err := <-results:
require.NoError(t, err)
case <-time.After(5 * time.Second):
t.Fatal("SSH handshake did not finish")
}
}
}

View File

@@ -7,11 +7,7 @@ import (
"crypto/ed25519"
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"net/http/httptest"
"net/url"
"os"
"path/filepath"
"strings"
@@ -25,7 +21,6 @@ import (
"github.com/blang/semver"
"github.com/fxamacker/cbor/v2"
"github.com/gliderlabs/ssh"
"github.com/lxzan/gws"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
gossh "golang.org/x/crypto/ssh"
@@ -203,370 +198,6 @@ func TestStartServerDisableSSH(t *testing.T) {
assert.Contains(t, err.Error(), "SSH disabled")
}
func TestStopServerDoesNotBlockWhenEventQueueFull(t *testing.T) {
agent := createTestAgent(t)
agent.server = &ssh.Server{}
agent.connectionManager.eventChan = make(chan ConnectionEvent, 1)
agent.connectionManager.eventChan <- WebSocketConnect
done := make(chan error, 1)
go func() {
done <- agent.StopServer()
}()
select {
case err := <-done:
require.NoError(t, err)
case <-time.After(time.Second):
t.Fatal("StopServer blocked on the connection event queue")
}
assert.Nil(t, agent.server)
assert.Equal(t, WebSocketConnect, <-agent.connectionManager.eventChan)
}
func TestSSHConnectionFallbackLifecycle(t *testing.T) {
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
agent := createTestAgent(t)
cm := agent.connectionManager
cm.eventChan = make(chan ConnectionEvent, 4)
_, privateKey, err := ed25519.GenerateKey(nil)
require.NoError(t, err)
signer, err := gossh.NewSignerFromKey(privateKey)
require.NoError(t, err)
cm.serverOptions = ServerOptions{
Network: "tcp",
Addr: "127.0.0.1:0",
Keys: []gossh.PublicKey{signer.PublicKey()},
}
// A WebSocket that closed after its upgrade returned nil must start SSH.
cm.handleEvent(WebSocketDisconnect)
agent.serverMu.Lock()
require.NotNil(t, agent.serverListener)
addr := agent.serverListener.Addr().String()
agent.serverMu.Unlock()
defer func() { _ = agent.StopServer() }()
clientConfig := &gossh.ClientConfig{
User: "hub",
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
}
client, err := gossh.Dial("tcp", addr, clientConfig)
require.NoError(t, err)
defer client.Close()
// A connection is counted when it starts its first session.
startSession := func(c *gossh.Client) *gossh.Session {
session, err := c.NewSession()
require.NoError(t, err)
require.NoError(t, session.Shell())
return session
}
session := startSession(client)
select {
case <-cm.sshChanged:
cm.handleSSHChange()
case <-time.After(5 * time.Second):
t.Fatal("SSH connection did not notify the manager")
}
require.Equal(t, SSHConnected, cm.getState())
wsAttempt := make(chan struct{}, 1)
releaseWS := make(chan struct{})
var releaseOnce sync.Once
release := func() { releaseOnce.Do(func() { close(releaseWS) }) }
hub := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
wsAttempt <- struct{}{}
<-releaseWS
w.WriteHeader(http.StatusServiceUnavailable)
}))
defer hub.Close()
defer release()
t.Setenv("BESZEL_AGENT_HUB_URL", hub.URL)
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
cm.wsClient, err = newWebSocketClient(agent)
require.NoError(t, err)
// A normal short-lived session must not be mistaken for a lost connection,
// and further sessions must not count the same connection again.
_ = session.Close()
_ = startSession(client).Close()
select {
case <-cm.sshChanged:
t.Fatal("session close unexpectedly changed SSH connection state")
case <-time.After(100 * time.Millisecond):
}
cm.mu.Lock()
assert.Equal(t, 1, cm.sshConnections, "sessions should not be counted as connections")
cm.mu.Unlock()
secondClient, err := gossh.Dial("tcp", addr, clientConfig)
require.NoError(t, err)
defer secondClient.Close()
defer startSession(secondClient).Close()
require.Eventually(t, func() bool {
cm.mu.Lock()
defer cm.mu.Unlock()
return cm.sshConnections == 2
}, 5*time.Second, 10*time.Millisecond, "second SSH connection was not counted")
require.NoError(t, client.Close())
select {
case <-cm.sshChanged:
t.Fatal("closing one of two SSH connections changed SSH connection state")
case <-time.After(100 * time.Millisecond):
}
require.Equal(t, SSHConnected, cm.getState())
require.NoError(t, secondClient.Close())
select {
case <-cm.sshChanged:
cm.handleSSHChange()
case <-time.After(5 * time.Second):
t.Fatal("SSH TCP close did not notify the manager")
}
require.Equal(t, Disconnected, cm.getState())
require.NotNil(t, cm.wsTicker)
select {
case <-wsAttempt:
case <-time.After(5 * time.Second):
t.Fatal("agent did not retry WebSocket after SSH disconnected")
}
// The hub may redial straight away, so the listener must stay open while
// the WebSocket attempt is pending and after it fails.
requireSameListener := func(msg string) {
agent.serverMu.Lock()
defer agent.serverMu.Unlock()
require.NotNil(t, agent.serverListener, msg)
assert.Equal(t, addr, agent.serverListener.Addr().String(), msg)
}
requireSameListener("SSH listener should stay open during the WebSocket attempt")
thirdClient, err := gossh.Dial("tcp", addr, clientConfig)
require.NoError(t, err, "SSH should accept a redial during the WebSocket attempt")
require.NoError(t, thirdClient.Close())
release()
require.Eventually(t, func() bool {
return !cm.isConnectingNow()
}, 5*time.Second, 10*time.Millisecond, "reconnect attempt did not finish")
requireSameListener("SSH listener should stay open after the WebSocket attempt fails")
cm.stopWsTicker()
}
// offeredKeySigner offers an authorized public key without proving possession
// of its private key: Sign blocks until released, then signs with another key.
type offeredKeySigner struct {
gossh.Signer
publicKey gossh.PublicKey
signing chan struct{}
release chan struct{}
}
func (s *offeredKeySigner) PublicKey() gossh.PublicKey { return s.publicKey }
func (s *offeredKeySigner) Sign(rand io.Reader, data []byte) (*gossh.Signature, error) {
close(s.signing)
<-s.release
return s.Signer.Sign(rand, data)
}
// The public key handler runs when a key is offered, before the client signs
// anything, so it must not be what marks an SSH connection as established.
func TestSSHPublicKeyOfferIsNotAConnection(t *testing.T) {
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
agent := createTestAgent(t)
cm := agent.connectionManager
cm.eventChan = make(chan ConnectionEvent, 4)
newSigner := func() gossh.Signer {
_, privateKey, err := ed25519.GenerateKey(nil)
require.NoError(t, err)
signer, err := gossh.NewSignerFromKey(privateKey)
require.NoError(t, err)
return signer
}
hubKey := newSigner().PublicKey()
cm.serverOptions = ServerOptions{
Network: "tcp",
Addr: "127.0.0.1:0",
Keys: []gossh.PublicKey{hubKey},
}
cm.handleEvent(WebSocketDisconnect)
agent.serverMu.Lock()
require.NotNil(t, agent.serverListener)
addr := agent.serverListener.Addr().String()
agent.serverMu.Unlock()
defer func() { _ = agent.StopServer() }()
signer := &offeredKeySigner{
Signer: newSigner(),
publicKey: hubKey,
signing: make(chan struct{}),
release: make(chan struct{}),
}
dialErr := make(chan error, 1)
go func() {
client, err := gossh.Dial("tcp", addr, &gossh.ClientConfig{
User: "hub",
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
})
if client != nil {
client.Close()
}
dialErr <- err
}()
// The server has accepted the offered key and is waiting for a signature.
select {
case <-signer.signing:
case <-time.After(5 * time.Second):
t.Fatal("server did not accept the offered public key")
}
select {
case <-cm.sshChanged:
t.Fatal("offering a public key changed SSH connection state")
case <-time.After(100 * time.Millisecond):
}
assert.False(t, cm.hasSSHConnection())
assert.Equal(t, Disconnected, cm.getState())
close(signer.release)
require.Error(t, <-dialErr, "a signature from another key must be rejected")
assert.False(t, cm.hasSSHConnection())
}
// startSSHFallbackServer starts the fallback SSH server for a disconnected
// agent and returns its address and a client config that can authenticate.
func startSSHFallbackServer(t *testing.T) (*Agent, string, *gossh.ClientConfig) {
t.Helper()
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
agent := createTestAgent(t)
cm := agent.connectionManager
cm.eventChan = make(chan ConnectionEvent, 4)
_, privateKey, err := ed25519.GenerateKey(nil)
require.NoError(t, err)
signer, err := gossh.NewSignerFromKey(privateKey)
require.NoError(t, err)
cm.serverOptions = ServerOptions{
Network: "tcp",
Addr: "127.0.0.1:0",
Keys: []gossh.PublicKey{signer.PublicKey()},
}
cm.startSSHServer()
agent.serverMu.Lock()
require.NotNil(t, agent.serverListener)
addr := agent.serverListener.Addr().String()
agent.serverMu.Unlock()
t.Cleanup(func() { _ = agent.StopServer() })
return agent, addr, &gossh.ClientConfig{
User: "hub",
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
}
}
// connectSSHSession dials the agent and starts a session, which is what marks
// the connection as established.
func connectSSHSession(t *testing.T, addr string, config *gossh.ClientConfig) *gossh.Client {
t.Helper()
client, err := gossh.Dial("tcp", addr, config)
require.NoError(t, err)
t.Cleanup(func() { _ = client.Close() })
session, err := client.NewSession()
require.NoError(t, err)
require.NoError(t, session.Shell())
return client
}
// handleNextSSHChange applies the next SSH connection notification, as the
// connection manager's event loop would.
func handleNextSSHChange(t *testing.T, cm *ConnectionManager) {
t.Helper()
select {
case <-cm.sshChanged:
cm.handleSSHChange()
case <-time.After(5 * time.Second):
t.Fatal("SSH connection change did not notify the manager")
}
}
// An agent without a WebSocket client only has SSH, so losing the hub's SSH
// connection must leave the listener in place for it to reconnect.
func TestSSHDisconnectKeepsListenerWithoutWebSocket(t *testing.T) {
agent, addr, clientConfig := startSSHFallbackServer(t)
cm := agent.connectionManager
require.Nil(t, cm.wsClient)
defer cm.stopWsTicker()
client := connectSSHSession(t, addr, clientConfig)
handleNextSSHChange(t, cm)
require.Equal(t, SSHConnected, cm.getState())
require.NoError(t, client.Close())
handleNextSSHChange(t, cm)
require.Equal(t, Disconnected, cm.getState())
require.Eventually(t, func() bool {
return !cm.isConnectingNow()
}, 5*time.Second, 10*time.Millisecond, "reconnect attempt did not finish")
agent.serverMu.Lock()
require.NotNil(t, agent.serverListener, "SSH listener should stay open")
assert.Equal(t, addr, agent.serverListener.Addr().String(), "SSH listener should not be restarted")
agent.serverMu.Unlock()
connectSSHSession(t, addr, clientConfig)
handleNextSSHChange(t, cm)
assert.Equal(t, SSHConnected, cm.getState())
}
// A WebSocket attempt that was already in flight can authenticate after SSH has
// connected. WebSocket is preferred, so it takes over and SSH is shut down.
func TestWebSocketTakesOverFromSSH(t *testing.T) {
agent, addr, clientConfig := startSSHFallbackServer(t)
cm := agent.connectionManager
client := connectSSHSession(t, addr, clientConfig)
handleNextSSHChange(t, cm)
require.Equal(t, SSHConnected, cm.getState())
cm.wsClient = &WebSocketClient{
agent: agent,
hubURL: &url.URL{Host: "localhost:8080"},
Conn: &gws.Conn{},
hubVerified: true,
}
cm.handleEvent(WebSocketConnect)
require.Equal(t, WebSocketConnected, cm.getState())
agent.serverMu.Lock()
assert.Nil(t, agent.serverListener, "SSH listener should close once WebSocket takes over")
agent.serverMu.Unlock()
closed := make(chan struct{})
go func() {
_ = client.Wait()
close(closed)
}()
select {
case <-closed:
case <-time.After(5 * time.Second):
t.Fatal("SSH connection was not closed when WebSocket took over")
}
// The SSH connection closing must not disturb the WebSocket state.
handleNextSSHChange(t, cm)
assert.False(t, cm.hasSSHConnection())
assert.Equal(t, WebSocketConnected, cm.getState())
}
/////////////////////////////////////////////////////////////////
//////////////////// ParseKeys Tests ////////////////////////////
/////////////////////////////////////////////////////////////////
@@ -773,23 +404,27 @@ func TestGetHubVersion(t *testing.T) {
clientVersion: "SSH-2.0-beszel_0.12.0",
}
// Test first call - should extract version
version := agent.getHubVersion(mockCtx)
// Test first call - should extract and cache version
version := agent.getHubVersion("test-session-123", mockCtx)
assert.Equal(t, "0.12.0", version.String())
// Test that version reflects the current client version (no stale caching)
mockCtx.clientVersion = "SSH-2.0-beszel_0.11.0"
version = agent.getHubVersion(mockCtx)
// Test second call - should return cached version
mockCtx.clientVersion = "SSH-2.0-beszel_0.11.0" // Change version but should still return cached
version = agent.getHubVersion("test-session-123", mockCtx)
assert.Equal(t, "0.12.0", version.String()) // Should still be cached version
// Test different session - should extract new version
version = agent.getHubVersion("different-session", mockCtx)
assert.Equal(t, "0.11.0", version.String())
// Test with invalid version string (non-beszel client)
mockCtx.clientVersion = "SSH-2.0-OpenSSH_8.0"
version = agent.getHubVersion(mockCtx)
version = agent.getHubVersion("invalid-session", mockCtx)
assert.Equal(t, "0.0.0", version.String()) // Should be empty version for non-beszel clients
// Test with no client version
mockCtx.clientVersion = ""
version = agent.getHubVersion(mockCtx)
version = agent.getHubVersion("no-version-session", mockCtx)
assert.True(t, version.EQ(semver.Version{})) // Should be empty version
}
@@ -866,6 +501,9 @@ func TestWriteToSessionEncoding(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// Reset the global hubVersions map to ensure clean state for each test
hubVersions = nil
agent, err := NewAgent("")
require.NoError(t, err)
@@ -947,28 +585,39 @@ func createTestCombinedData() *system.CombinedData {
}
}
// TestGetHubVersionConcurrent guards against a regression of the
// "concurrent map writes" panic previously caused by a shared, unsynchronized
// hubVersions cache (see https://github.com/henrygd/beszel/issues/2128).
// getHubVersion no longer shares mutable state between sessions, so calling
// it concurrently from many goroutines must be safe under `go test -race`.
func TestGetHubVersionConcurrent(t *testing.T) {
func TestHubVersionCaching(t *testing.T) {
// Reset the global hubVersions map to ensure clean state
hubVersions = nil
agent, err := NewAgent("")
require.NoError(t, err)
const goroutines = 50
var wg sync.WaitGroup
wg.Add(goroutines)
for i := 0; i < goroutines; i++ {
go func(i int) {
defer wg.Done()
ctx := &mockSSHContext{
sessionID: fmt.Sprintf("session-%d", i),
clientVersion: "SSH-2.0-beszel_0.12.0",
}
version := agent.getHubVersion(ctx)
assert.Equal(t, "0.12.0", version.String())
}(i)
ctx1 := &mockSSHContext{
sessionID: "session1",
clientVersion: "SSH-2.0-beszel_0.12.0",
}
wg.Wait()
ctx2 := &mockSSHContext{
sessionID: "session2",
clientVersion: "SSH-2.0-beszel_0.11.0",
}
// First calls should cache the versions
v1 := agent.getHubVersion("session1", ctx1)
v2 := agent.getHubVersion("session2", ctx2)
assert.Equal(t, "0.12.0", v1.String())
assert.Equal(t, "0.11.0", v2.String())
// Verify caching by changing context but keeping same session ID
ctx1.clientVersion = "SSH-2.0-beszel_0.10.0"
v1Cached := agent.getHubVersion("session1", ctx1)
assert.Equal(t, "0.12.0", v1Cached.String()) // Should still be cached version
// New session should get new version
ctx3 := &mockSSHContext{
sessionID: "session3",
clientVersion: "SSH-2.0-beszel_0.13.0",
}
v3 := agent.getHubVersion("session3", ctx3)
assert.Equal(t, "0.13.0", v3.String())
}

View File

@@ -55,11 +55,6 @@ type DeviceInfo struct {
typeVerified bool
// parserType holds the parser type (nvme, sat, scsi) that last succeeded.
parserType string
// explicitType reports whether Type came from an explicit ":type" hint in
// SMART_DEVICES. Such a type is a deliberate user override and must always be
// passed to smartctl via -d, even for scsi/ata where a scan-detected type is
// otherwise left off (see smartctlArgs and issue #1345).
explicitType bool
}
// deviceKey is a composite key for a device, used to identify a device uniquely.
@@ -70,9 +65,8 @@ type deviceKey struct {
var errNoValidSmartData = fmt.Errorf("no valid SMART data found") // Error for missing data
// Refresh updates SMART data for all known devices and reports whether every
// discovered device was collected successfully.
func (sm *SmartManager) Refresh(forceScan bool) (bool, error) {
// Refresh updates SMART data for all known devices
func (sm *SmartManager) Refresh(forceScan bool) error {
sm.refreshMutex.Lock()
defer sm.refreshMutex.Unlock()
@@ -93,7 +87,7 @@ func (sm *SmartManager) Refresh(forceScan bool) (bool, error) {
}
}
return scanErr == nil && collectErr == nil, sm.resolveRefreshError(scanErr, collectErr)
return sm.resolveRefreshError(scanErr, collectErr)
}
// devicesSnapshot returns a copy of the current device slice to avoid iterating
@@ -257,9 +251,8 @@ func (sm *SmartManager) parseConfiguredDevices(config string) ([]*DeviceInfo, er
}
devices = append(devices, &DeviceInfo{
Name: name,
Type: devType,
explicitType: devType != "",
Name: name,
Type: devType,
})
}
@@ -311,13 +304,11 @@ func (sm *SmartManager) filterExcludedDevices(devices []*DeviceInfo) []*DeviceIn
return filtered
}
// detectSmartOutputType inspects protocol-specific sections and the reported
// device type to choose a parser, including when the NVMe health log is missing.
// detectSmartOutputType inspects sections that are unique to each smartctl
// JSON schema (NVMe, ATA/SATA, SCSI) to determine which parser should be used
// when the reported device type is ambiguous or missing.
func detectSmartOutputType(output []byte) string {
var hints struct {
Device struct {
Type string `json:"type"`
} `json:"device"`
AtaSmartAttributes json.RawMessage `json:"ata_smart_attributes"`
NVMeSmartHealthInformationLog json.RawMessage `json:"nvme_smart_health_information_log"`
ScsiErrorCounterLog json.RawMessage `json:"scsi_error_counter_log"`
@@ -328,7 +319,7 @@ func detectSmartOutputType(output []byte) string {
}
switch {
case hasJSONValue(hints.NVMeSmartHealthInformationLog), normalizeParserType(hints.Device.Type) == "nvme":
case hasJSONValue(hints.NVMeSmartHealthInformationLog):
return "nvme"
case hasJSONValue(hints.AtaSmartAttributes):
return "sat"
@@ -377,15 +368,9 @@ func (sm *SmartManager) parseSmartOutput(deviceInfo *DeviceInfo, output []byte)
Type string
Parse func([]byte) (bool, int)
}{
{Type: "nvme", Parse: func(output []byte) (bool, int) {
return sm.parseSmartForNvme(output, deviceInfo.Type)
}},
{Type: "sat", Parse: func(output []byte) (bool, int) {
return sm.parseSmartForSata(output, deviceInfo.Type)
}},
{Type: "scsi", Parse: func(output []byte) (bool, int) {
return sm.parseSmartForScsi(output, deviceInfo.Type)
}},
{Type: "nvme", Parse: sm.parseSmartForNvme},
{Type: "sat", Parse: sm.parseSmartForSata},
{Type: "scsi", Parse: sm.parseSmartForScsi},
}
deviceType := normalizeParserType(deviceInfo.parserType)
@@ -399,11 +384,11 @@ func (sm *SmartManager) parseSmartOutput(deviceInfo *DeviceInfo, output []byte)
}
}
// Inspect every response so a failed NVMe query cannot reach other parsers.
structureType := detectSmartOutputType(output)
// Update the stored parser only when it is not yet verified.
// Only run the type detection when we do not yet know which parser works
// or the previous attempt failed.
needsDetection := deviceType == "" || !deviceInfo.typeVerified
if needsDetection {
structureType := detectSmartOutputType(output)
if deviceType != structureType {
deviceType = structureType
deviceInfo.parserType = structureType
@@ -444,11 +429,6 @@ func (sm *SmartManager) parseSmartOutput(deviceInfo *DeviceInfo, output []byte)
// Try the selected parsers in order until we find one that succeeds.
for _, parser := range selectedParsers {
// A failed NVMe response may still contain a serial number, which is
// enough for the SATA and SCSI parsers to accept incorrect zero values.
if structureType == "nvme" && parser.Type != "nvme" {
continue
}
hasData, _ := parser.Parse(output)
if hasData {
deviceInfo.parserType = parser.Type
@@ -499,11 +479,10 @@ func (sm *SmartManager) CollectSmart(deviceInfo *DeviceInfo) error {
return errNoValidSmartData
}
// slog.Info("collecting SMART data", "device", deviceInfo.Name, "type", deviceInfo.Type, "has_existing_data", sm.hasDataForDevice(deviceInfo))
// slog.Info("collecting SMART data", "device", deviceInfo.Name, "type", deviceInfo.Type, "has_existing_data", sm.hasDataForDevice(deviceInfo.Name))
// Check if we have existing data for this exact device identity. Multiple
// bridge slots can share a path, so a name-only match is not sufficient.
hasExistingData := sm.hasDataForDevice(deviceInfo)
// Check if we have any existing data for this device
hasExistingData := sm.hasDataForDevice(deviceInfo.Name)
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
@@ -579,9 +558,7 @@ func (sm *SmartManager) smartctlArgs(deviceInfo *DeviceInfo, includeStandby bool
deviceType = strings.ToLower(deviceInfo.Type)
parserType = strings.ToLower(deviceInfo.parserType)
// types sometimes misidentified in scan; see github.com/henrygd/beszel/issues/1345
// An explicit SMART_DEVICES ":type" hint is a deliberate override, so always
// pass it through; otherwise scsi/ata are left off so smartctl can auto-detect.
if deviceType != "" && (deviceInfo.explicitType || (deviceType != "scsi" && deviceType != "ata")) {
if deviceType != "" && deviceType != "scsi" && deviceType != "ata" {
args = append(args, "-d", deviceInfo.Type)
}
}
@@ -606,18 +583,14 @@ func (sm *SmartManager) smartctlArgs(deviceInfo *DeviceInfo, includeStandby bool
return args
}
// hasDataForDevice checks if we have cached SMART data for a specific device identity.
func (sm *SmartManager) hasDataForDevice(deviceInfo *DeviceInfo) bool {
if deviceInfo == nil {
return false
}
// hasDataForDevice checks if we have cached SMART data for a specific device
func (sm *SmartManager) hasDataForDevice(deviceName string) bool {
sm.Lock()
defer sm.Unlock()
deviceKey := makeDeviceKey(deviceInfo.Name, deviceInfo.Type)
// Check if any cached data has this device name
for _, data := range sm.SmartDataMap {
if data != nil && makeDeviceKey(data.DiskName, data.DiskType) == deviceKey {
if data != nil && data.DiskName == deviceName {
return true
}
}
@@ -690,9 +663,6 @@ func mergeDeviceLists(existing, scanned, configured []*DeviceInfo) []*DeviceInfo
target.Type = prev.Type
target.typeVerified = true
target.parserType = prev.parserType
if prev.explicitType {
target.explicitType = true
}
}
// applyConfiguredMetadata updates a matched device with any configured
@@ -706,9 +676,6 @@ func mergeDeviceLists(existing, scanned, configured []*DeviceInfo) []*DeviceInfo
existingDev.typeVerified = false
existingDev.parserType = normalizeParserType(newType)
}
if configuredDev.explicitType {
existingDev.explicitType = true
}
if configuredDev.InfoName != "" {
existingDev.InfoName = configuredDev.InfoName
}
@@ -765,14 +732,7 @@ func mergeDeviceLists(existing, scanned, configured []*DeviceInfo) []*DeviceInfo
continue
}
if existingDev := deviceIndexByName[configuredDevice.Name]; existingDev != nil {
oldKey := makeDeviceKey(existingDev.Name, existingDev.Type)
if prev := existingIndex[key]; prev != nil {
preserveVerifiedType(existingDev, prev)
}
applyConfiguredMetadata(existingDev, configuredDevice)
delete(deviceIndex, oldKey)
deviceIndex[makeDeviceKey(existingDev.Name, existingDev.Type)] = existingDev
delete(deviceIndexByName, configuredDevice.Name)
continue
}
@@ -876,11 +836,9 @@ func (sm *SmartManager) isVirtualDeviceFromStrings(fields ...string) bool {
return false
}
// parseSmartForSata parses the output of smartctl --all -j for SATA/ATA devices and updates the SmartDataMap.
// deviceType is the exact type used to identify and query the device; when set,
// it takes precedence over the generic type reported by smartctl.
// parseSmartForSata parses the output of smartctl --all -j for SATA/ATA devices and updates the SmartDataMap
// Returns hasValidData and exitStatus
func (sm *SmartManager) parseSmartForSata(output []byte, deviceType string) (bool, int) {
func (sm *SmartManager) parseSmartForSata(output []byte) (bool, int) {
var data smart.SmartInfoForSata
if err := json.Unmarshal(output, &data); err != nil {
@@ -919,9 +877,6 @@ func (sm *SmartManager) parseSmartForSata(output []byte, deviceType string) (boo
smartData.SmartStatus = getSmartStatus(smartData.Temperature, data.SmartStatus.Passed)
smartData.DiskName = data.Device.Name
smartData.DiskType = data.Device.Type
if deviceType != "" {
smartData.DiskType = deviceType
}
// get values from ata_device_statistics if necessary
var ataDeviceStats smart.AtaDeviceStatistics
@@ -995,7 +950,7 @@ func findAtaDeviceStatisticsValue(data *smart.SmartInfoForSata, ataDeviceStats *
return nil
}
func (sm *SmartManager) parseSmartForScsi(output []byte, deviceType string) (bool, int) {
func (sm *SmartManager) parseSmartForScsi(output []byte) (bool, int) {
var data smart.SmartInfoForScsi
if err := json.Unmarshal(output, &data); err != nil {
@@ -1030,9 +985,6 @@ func (sm *SmartManager) parseSmartForScsi(output []byte, deviceType string) (boo
smartData.SmartStatus = getSmartStatus(smartData.Temperature, data.SmartStatus.Passed)
smartData.DiskName = data.Device.Name
smartData.DiskType = data.Device.Type
if deviceType != "" {
smartData.DiskType = deviceType
}
attributes := make([]*smart.SmartAttribute, 0, 10)
attributes = append(attributes, &smart.SmartAttribute{Name: "PowerOnHours", RawValue: data.PowerOnTime.Hours})
@@ -1130,11 +1082,9 @@ func (sm *SmartManager) lookupDarwinNvmeCapacity(serial string) uint64 {
return sm.darwinNvmeCapacity[serial]
}
// parseSmartForNvme parses the output of smartctl --all -j /dev/nvmeX and updates the SmartDataMap.
// deviceType is the exact type used to identify and query the device; when set,
// it takes precedence over the generic type reported by smartctl.
// parseSmartForNvme parses the output of smartctl --all -j /dev/nvmeX and updates the SmartDataMap
// Returns hasValidData and exitStatus
func (sm *SmartManager) parseSmartForNvme(output []byte, deviceType string) (bool, int) {
func (sm *SmartManager) parseSmartForNvme(output []byte) (bool, int) {
data := &smart.SmartInfoForNvme{}
if err := json.Unmarshal(output, &data); err != nil {
@@ -1152,16 +1102,6 @@ func (sm *SmartManager) parseSmartForNvme(output []byte, deviceType string) (boo
return false, data.Smartctl.ExitStatus
}
// smartctl may return device identity fields before failing to read the NVMe
// health log (for example, on an unsupported controller path or with insufficient
// permissions). Do not accept that partial response as valid SMART data: doing
// so stores incorrect zero values and prevents the namespace-path fallback.
log := data.NVMeSmartHealthInformationLog
if log == nil {
slog.Debug("no NVMe SMART health information", "device", data.Device.Name)
return false, data.Smartctl.ExitStatus
}
sm.Lock()
defer sm.Unlock()
@@ -1184,16 +1124,14 @@ func (sm *SmartManager) parseSmartForNvme(output []byte, deviceType string) (boo
if smartData.Capacity == 0 && (runtime.GOOS == "darwin" || sm.darwinNvmeProvider != nil) {
smartData.Capacity = sm.lookupDarwinNvmeCapacity(data.SerialNumber)
}
smartData.Temperature = log.Temperature
smartData.Temperature = data.NVMeSmartHealthInformationLog.Temperature
smartData.SmartStatus = getSmartStatus(smartData.Temperature, data.SmartStatus.Passed)
smartData.DiskName = data.Device.Name
smartData.DiskType = data.Device.Type
if deviceType != "" {
smartData.DiskType = deviceType
}
// nvme attributes does not follow the same format as ata attributes,
// so we manually map each field to SmartAttributes
log := data.NVMeSmartHealthInformationLog
smartData.Attributes = []*smart.SmartAttribute{
{Name: "CriticalWarning", RawValue: uint64(log.CriticalWarning)},
{Name: "Temperature", RawValue: uint64(log.Temperature)},

Some files were not shown because too many files have changed in this diff Show More