Compare commits

..

2 Commits

Author SHA1 Message Date
henrygd
dae8181f47 fix: improve agent installer error handling (#1971, #1972)
- Enable set -eu and handle expected failures explicitly
- Validate platform and service manager before making changes
- Fix argument parsing, version fallback, cron setup, and prompt EOF
- Verify downloads before stopping the existing agent
- Add atomic binary replacement, rollback, and cleanup
- Add regression tests for installation failure paths
2026-09-05 13:02:15 -04:00
Elan Ruusamäe
b8fb5d2367 Fix: Enable immediate exit on errors in install-agent.sh (#1972)
Add error handling to ensure script exits on errors.
2026-09-05 12:05:45 -04:00
27 changed files with 219 additions and 1070 deletions

View File

@@ -2,8 +2,6 @@
## Reporting a Vulnerability ## 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. If it's low severity (use best judgement) you may 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.

View File

@@ -25,9 +25,7 @@ import (
) )
const ( const (
// Keep the connection alive long enough for a slow collection cycle to wsDeadline = 70 * time.Second
// finish before the hub considers the agent disconnected.
wsDeadline = 120 * time.Second
) )
type caCertFileError struct { type caCertFileError struct {

View File

@@ -700,11 +700,3 @@ func TestGetToken(t *testing.T) {
assert.Equal(t, expectedToken, token, "Whitespace should be stripped from token file content") 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)
}
}

2
go.mod
View File

@@ -10,7 +10,7 @@ require (
github.com/gliderlabs/ssh v0.3.8 github.com/gliderlabs/ssh v0.3.8
github.com/google/uuid v1.6.0 github.com/google/uuid v1.6.0
github.com/lxzan/gws v1.10.1 github.com/lxzan/gws v1.10.1
github.com/nicholas-fedor/shoutrrr v0.20.0 github.com/nicholas-fedor/shoutrrr v0.19.0
github.com/pocketbase/dbx v1.12.0 github.com/pocketbase/dbx v1.12.0
github.com/pocketbase/pocketbase v0.40.2 github.com/pocketbase/pocketbase v0.40.2
github.com/shirou/gopsutil/v4 v4.26.8 github.com/shirou/gopsutil/v4 v4.26.8

8
go.sum
View File

@@ -54,8 +54,8 @@ github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArs
github.com/golang/protobuf v1.3.1/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= github.com/golang/protobuf v1.3.1/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/google/pprof v0.0.0-20260906184651-6331bc6350fe h1:QAinXoAFJdGQYztXn3VpFey7KCwpedbZ/EkzbplQ0cY= github.com/google/pprof v0.0.0-20260902005441-ca85771921e4 h1:/6mPXfWmhv8eKck12I0YNIcIjwHtxP3YRIMKiEgTjWg=
github.com/google/pprof v0.0.0-20260906184651-6331bc6350fe/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk= github.com/google/pprof v0.0.0-20260902005441-ca85771921e4/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
@@ -83,8 +83,8 @@ github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsRe
github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A= github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A=
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
github.com/nicholas-fedor/shoutrrr v0.20.0 h1:hMAxIYlfAeZ1FcTDgU0kUOvVXUsOirWo8IWlnzGLkac= github.com/nicholas-fedor/shoutrrr v0.19.0 h1:Rl6bpK3DXuR2Trtx2JV8t+wjUwkHdRHrc8nBKoEpHr0=
github.com/nicholas-fedor/shoutrrr v0.20.0/go.mod h1:hgde37yNWCXh8+N6WemyDRMNYLOFTf326GsBx8Z7CFA= github.com/nicholas-fedor/shoutrrr v0.19.0/go.mod h1:Glfdi8AGTbnEn2k2+hW62n8oL0i9vqRVFtXaUIthNks=
github.com/onsi/ginkgo/v2 v2.32.1 h1:6tlvcDm/3sE8lGJbZ4+d4mO3RLy24/tQWOFzVSQNIfw= github.com/onsi/ginkgo/v2 v2.32.1 h1:6tlvcDm/3sE8lGJbZ4+d4mO3RLy24/tQWOFzVSQNIfw=
github.com/onsi/ginkgo/v2 v2.32.1/go.mod h1:+aXOY+vzZ5mu2iI2HpTZUPmM//oQfsNFX6gU9kNcA44= github.com/onsi/ginkgo/v2 v2.32.1/go.mod h1:+aXOY+vzZ5mu2iI2HpTZUPmM//oQfsNFX6gU9kNcA44=
github.com/onsi/gomega v1.43.0 h1:VlG/1FxqNxhSO+lq/OHBNaaqwiBK/mO8JbVkX9Y+FeU= github.com/onsi/gomega v1.43.0 h1:VlG/1FxqNxhSO+lq/OHBNaaqwiBK/mO8JbVkX9Y+FeU=

View File

@@ -20,10 +20,10 @@ type hubLike interface {
} }
type AlertManager struct { type AlertManager struct {
hub hubLike hub hubLike
stopOnce sync.Once stopOnce sync.Once
pendingAlerts sync.Map pendingAlerts sync.Map
alertsCache *AlertsCache alertsCache *AlertsCache
} }
type AlertMessageData struct { type AlertMessageData struct {
@@ -231,20 +231,8 @@ func (am *AlertManager) SendAlert(data AlertMessageData) error {
am.hub.Logger().Error("Failed to unmarshal user settings", "err", err) am.hub.Logger().Error("Failed to unmarshal user settings", "err", err)
} }
// send alerts via webhooks // send alerts via webhooks
send := sendPublicNotification
if len(userAlertSettings.Webhooks) > 0 {
// Read the owner's current role at delivery time, including for URLs
// saved before an admin was demoted. Never fall back on lookup failure.
owner, err := am.hub.FindRecordById("users", data.UserID)
if err != nil {
return fmt.Errorf("load notification owner: %w", err)
}
if owner.GetString("role") == "admin" {
send = shoutrrr.Send
}
}
for _, webhook := range userAlertSettings.Webhooks { for _, webhook := range userAlertSettings.Webhooks {
if err := am.sendShoutrrrAlert(webhook, data.Title, data.Message, data.Link, data.LinkText, send); err != nil { if err := am.SendShoutrrrAlert(webhook, data.Title, data.Message, data.Link, data.LinkText); err != nil {
am.hub.Logger().Error("Failed to send shoutrrr alert", "err", err) am.hub.Logger().Error("Failed to send shoutrrr alert", "err", err)
} }
} }
@@ -275,10 +263,6 @@ func (am *AlertManager) SendAlert(data AlertMessageData) error {
// SendShoutrrrAlert sends an alert via a Shoutrrr URL // SendShoutrrrAlert sends an alert via a Shoutrrr URL
func (am *AlertManager) SendShoutrrrAlert(notificationUrl, title, message, link, linkText string) error { func (am *AlertManager) SendShoutrrrAlert(notificationUrl, title, message, link, linkText string) error {
return am.sendShoutrrrAlert(notificationUrl, title, message, link, linkText, shoutrrr.Send)
}
func (am *AlertManager) sendShoutrrrAlert(notificationUrl, title, message, link, linkText string, send func(string, string) error) error {
// Parse the URL // Parse the URL
parsedURL, err := url.Parse(notificationUrl) parsedURL, err := url.Parse(notificationUrl)
if err != nil { if err != nil {
@@ -321,7 +305,7 @@ func (am *AlertManager) sendShoutrrrAlert(notificationUrl, title, message, link,
parsedURL.RawQuery = queryParams.Encode() parsedURL.RawQuery = queryParams.Encode()
// log.Println("URL after modification:", parsedURL.String()) // log.Println("URL after modification:", parsedURL.String())
err = send(parsedURL.String(), message) err = shoutrrr.Send(parsedURL.String(), message)
if err == nil { if err == nil {
am.hub.Logger().Info("Sent shoutrrr alert", "title", title) am.hub.Logger().Info("Sent shoutrrr alert", "title", title)

View File

@@ -3,11 +3,13 @@ package alerts
import ( import (
"database/sql" "database/sql"
"errors" "errors"
"net"
"net/http" "net/http"
"net/url"
"slices" "slices"
"strings"
"github.com/henrygd/beszel/internal/hub/utils" "github.com/henrygd/beszel/internal/hub/utils"
"github.com/nicholas-fedor/shoutrrr"
"github.com/pocketbase/dbx" "github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/core" "github.com/pocketbase/pocketbase/core"
) )
@@ -145,16 +147,72 @@ func (am *AlertManager) SendTestNotification(e *core.RequestEvent) error {
if err != nil || data.URL == "" { if err != nil || data.URL == "" {
return e.BadRequestError("URL is required", err) return e.BadRequestError("URL is required", err)
} }
send := shoutrrr.Send // Only allow admins to send test notifications to internal URLs
if !e.Auth.IsSuperuser() && e.Auth.GetString("role") != "admin" { if !e.Auth.IsSuperuser() && e.Auth.GetString("role") != "admin" {
send = sendPublicNotification internalURL, err := isInternalURL(data.URL)
} if err != nil {
err = am.sendShoutrrrAlert(data.URL, "Test Alert", "This is a notification from Beszel.", am.hub.Settings().Meta.AppURL, "View Beszel", send) return e.BadRequestError(err.Error(), nil)
if errors.Is(err, errInternalDestination) || errors.Is(err, errUnrestrictedService) { }
return e.ForbiddenError(err.Error(), nil) if internalURL {
return e.ForbiddenError("Only admins can send to internal destinations", nil)
}
} }
err = am.SendShoutrrrAlert(data.URL, "Test Alert", "This is a notification from Beszel.", am.hub.Settings().Meta.AppURL, "View Beszel")
if err != nil { if err != nil {
return e.JSON(200, map[string]string{"err": err.Error()}) return e.JSON(200, map[string]string{"err": err.Error()})
} }
return e.JSON(200, map[string]bool{"err": false}) return e.JSON(200, map[string]bool{"err": false})
} }
// isInternalURL checks if the given shoutrrr URL points to an internal destination (localhost or private IP)
func isInternalURL(rawURL string) (bool, error) {
parsedURL, err := url.Parse(rawURL)
if err != nil {
return false, err
}
host := parsedURL.Hostname()
if host == "" {
return false, nil
}
if strings.EqualFold(host, "localhost") {
return true, nil
}
if ip := net.ParseIP(host); ip != nil {
return isInternalIP(ip), nil
}
// Some Shoutrrr URLs use the host position for service identifiers rather than a
// network hostname (for example, discord://token@webhookid). Restrict DNS lookups
// to names that look like actual hostnames so valid service URLs keep working.
if !strings.Contains(host, ".") {
return false, nil
}
ips, err := net.LookupIP(host)
if err != nil {
return false, nil
}
if slices.ContainsFunc(ips, isInternalIP) {
return true, nil
}
return false, nil
}
var cgnatNetwork = &net.IPNet{
IP: net.IPv4(100, 64, 0, 0),
Mask: net.CIDRMask(10, 32),
}
func isInternalIP(ip net.IP) bool {
return ip.IsPrivate() ||
ip.IsLoopback() ||
ip.IsUnspecified() ||
ip.IsLinkLocalUnicast() ||
ip.IsMulticast() ||
cgnatNetwork.Contains(ip)
}

View File

@@ -7,11 +7,10 @@ import (
"encoding/json" "encoding/json"
"io" "io"
"net/http" "net/http"
"net/http/httptest"
"strings" "strings"
"sync/atomic"
"testing" "testing"
"github.com/henrygd/beszel/internal/alerts"
beszelTests "github.com/henrygd/beszel/internal/tests" beszelTests "github.com/henrygd/beszel/internal/tests"
pbTests "github.com/pocketbase/pocketbase/tests" pbTests "github.com/pocketbase/pocketbase/tests"
@@ -30,6 +29,43 @@ func jsonReader(v any) io.Reader {
return bytes.NewReader(data) return bytes.NewReader(data)
} }
func TestIsInternalURL(t *testing.T) {
testCases := []struct {
name string
url string
internal bool
}{
{name: "loopback ipv4", url: "generic://127.0.0.1", internal: true},
{name: "private ipv4", url: "generic://10.0.0.1", internal: true},
{name: "localhost hostname", url: "generic://localhost", internal: true},
{name: "localhost with path", url: "generic+http://localhost/api/v1/postStuff", internal: true},
{name: "loopback with port and path", url: "generic+http://127.0.0.1:8080/api/v1/postStuff", internal: true},
{name: "public hostname", url: "generic+https://beszel.dev/api/v1/postStuff", internal: false},
{name: "cloud metadata ipv4", url: "generic://169.254.169.254", internal: true},
{name: "link-local ipv4", url: "generic://169.254.1.1", internal: true},
{name: "link-local ipv6", url: "generic://[fe80::1]", internal: true},
{name: "mapped link-local ipv4", url: "generic://[::ffff:169.254.169.254]", internal: true},
{name: "cgnat lower boundary", url: "generic://100.64.0.0", internal: true},
{name: "cgnat upper boundary", url: "generic://100.127.255.255", internal: true},
{name: "below cgnat", url: "generic://100.63.255.255", internal: false},
{name: "above cgnat", url: "generic://100.128.0.0", internal: false},
{name: "multicast ipv4", url: "generic://224.0.0.1", internal: true},
{name: "multicast ipv6", url: "generic://[ff02::1]", internal: true},
{name: "public ipv4", url: "generic://8.8.8.8", internal: false},
{name: "public ipv6", url: "generic://[2001:4860:4860::8888]", internal: false},
{name: "token style service url", url: "discord://abc123@123456789", internal: false},
{name: "single label service url", url: "slack://token@team/channel", internal: false},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
internal, err := alerts.IsInternalURL(testCase.url)
assert.NoError(t, err)
assert.Equal(t, testCase.internal, internal)
})
}
}
func TestUserAlertsApi(t *testing.T) { func TestUserAlertsApi(t *testing.T) {
hub, _ := beszelTests.NewTestHub(t.TempDir()) hub, _ := beszelTests.NewTestHub(t.TempDir())
defer hub.Cleanup() defer hub.Cleanup()
@@ -421,17 +457,6 @@ func TestSendTestNotification(t *testing.T) {
hub, user := beszelTests.GetHubWithUser(t) hub, user := beszelTests.GetHubWithUser(t)
defer hub.Cleanup() defer hub.Cleanup()
var delivered atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
delivered.Add(1)
}))
defer server.Close()
localURL := "generic+" + server.URL
readonlyUser, err := beszelTests.CreateUserWithRole(hub, "readonly@example.com", "password123", "readonly")
assert.NoError(t, err)
readonlyToken, err := readonlyUser.NewAuthToken()
assert.NoError(t, err)
userToken, err := user.NewAuthToken() userToken, err := user.NewAuthToken()
adminUser, err := beszelTests.CreateUserWithRole(hub, "admin@example.com", "password123", "admin") adminUser, err := beszelTests.CreateUserWithRole(hub, "admin@example.com", "password123", "admin")
@@ -456,11 +481,11 @@ func TestSendTestNotification(t *testing.T) {
ExpectedContent: []string{"requires valid"}, ExpectedContent: []string{"requires valid"},
TestAppFactory: testAppFactory, TestAppFactory: testAppFactory,
Body: jsonReader(map[string]any{ Body: jsonReader(map[string]any{
"url": localURL, "url": "generic://127.0.0.1",
}), }),
}, },
{ {
Name: "POST /test-notification - invalid service reports error", Name: "POST /test-notification - with external auth should succeed",
Method: http.MethodPost, Method: http.MethodPost,
URL: "/api/beszel/test-notification", URL: "/api/beszel/test-notification",
TestAppFactory: testAppFactory, TestAppFactory: testAppFactory,
@@ -468,7 +493,7 @@ func TestSendTestNotification(t *testing.T) {
"Authorization": userToken, "Authorization": userToken,
}, },
Body: jsonReader(map[string]any{ Body: jsonReader(map[string]any{
"url": "unknown://example.com", "url": "generic://8.8.8.8",
}), }),
ExpectedStatus: 200, ExpectedStatus: 200,
ExpectedContent: []string{"\"err\":"}, ExpectedContent: []string{"\"err\":"},
@@ -510,10 +535,10 @@ func TestSendTestNotification(t *testing.T) {
"Authorization": adminUserToken, "Authorization": adminUserToken,
}, },
Body: jsonReader(map[string]any{ Body: jsonReader(map[string]any{
"url": localURL, "url": "generic://127.0.0.1",
}), }),
ExpectedStatus: 200, ExpectedStatus: 200,
ExpectedContent: []string{"\"err\":false"}, ExpectedContent: []string{"\"err\":"},
}, },
{ {
Name: "POST /test-notification - internal url with superuser auth should succeed", Name: "POST /test-notification - internal url with superuser auth should succeed",
@@ -524,28 +549,14 @@ func TestSendTestNotification(t *testing.T) {
"Authorization": superuserToken, "Authorization": superuserToken,
}, },
Body: jsonReader(map[string]any{ Body: jsonReader(map[string]any{
"url": localURL, "url": "generic://127.0.0.1",
}), }),
ExpectedStatus: 200, ExpectedStatus: 200,
ExpectedContent: []string{"\"err\":"}, ExpectedContent: []string{"\"err\":"},
}, },
} }
for _, url := range []string{localURL, "smtp://user:pass@127.0.0.1/?fromAddress=sender@example.com&toAddresses=recipient@example.com", "mqtt://127.0.0.1/topic"} {
scenarios = append(scenarios, beszelTests.ApiScenario{
Name: "readonly cannot send to " + url,
Method: http.MethodPost,
URL: "/api/beszel/test-notification",
TestAppFactory: testAppFactory,
Headers: map[string]string{"Authorization": readonlyToken},
Body: jsonReader(map[string]any{"url": url}),
ExpectedStatus: 403,
ExpectedContent: []string{"Only admins"},
})
}
for _, scenario := range scenarios { for _, scenario := range scenarios {
scenario.Test(t) scenario.Test(t)
} }
assert.EqualValues(t, 2, delivered.Load(), "only admin and superuser requests should reach the server")
} }

View File

@@ -100,6 +100,10 @@ func (am *AlertManager) SetAlertTriggered(alert CachedAlertData, triggered bool)
return am.setAlertTriggered(alert, triggered) return am.setAlertTriggered(alert, triggered)
} }
func IsInternalURL(rawURL string) (bool, error) {
return isInternalURL(rawURL)
}
// BuildContainerLogExcerpt exposes buildContainerLogExcerpt for testing. // BuildContainerLogExcerpt exposes buildContainerLogExcerpt for testing.
func BuildContainerLogExcerpt(raw string) string { func BuildContainerLogExcerpt(raw string) string {
return buildContainerLogExcerpt(raw) return buildContainerLogExcerpt(raw)

View File

@@ -1,66 +0,0 @@
//go:build testing
package alerts_test
import (
"net/http"
"net/http/httptest"
"sync/atomic"
"testing"
"github.com/henrygd/beszel/internal/alerts"
beszelTests "github.com/henrygd/beszel/internal/tests"
"github.com/pocketbase/dbx"
"github.com/stretchr/testify/require"
)
func TestPersistedWebhooksUseCurrentOwnerRole(t *testing.T) {
hub, user := beszelTests.GetHubWithUser(t)
defer hub.Cleanup()
am := alerts.NewTestAlertManagerWithoutWorker(hub)
var delivered atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
delivered.Add(1)
}))
defer server.Close()
settings, err := hub.FindFirstRecordByFilter("user_settings", "user={:user}", dbx.Params{"user": user.Id})
require.NoError(t, err)
settings.Set("settings", alerts.UserNotificationSettings{Webhooks: []string{"generic+" + server.URL}})
require.NoError(t, hub.Save(settings))
message := alerts.AlertMessageData{UserID: user.Id, Title: "Test", Message: "Persisted webhook"}
// Keep the same URL and manager while changing roles, so cached privileges
// or treating previously saved URLs as trusted would fail this test.
for _, tc := range []struct {
name string
role string
want int32
}{
{"regular user", "user", 0},
{"readonly user", "readonly", 0},
{"promoted admin", "admin", 1},
{"demoted admin", "user", 1},
} {
t.Run(tc.name, func(t *testing.T) {
user.Set("role", tc.role)
require.NoError(t, hub.Save(user))
// Webhook errors are logged; SendAlert continues to email delivery.
require.NoError(t, am.SendAlert(message))
require.Equal(t, tc.want, delivered.Load())
})
}
t.Run("missing owner fails closed", func(t *testing.T) {
// Model an orphaned settings record without deleting it through the
// normal user deletion cascade.
const missingOwner = "missingowner123"
settings.Set("user", missingOwner)
require.NoError(t, hub.SaveNoValidate(settings))
message.UserID = missingOwner
err := am.SendAlert(message)
require.ErrorContains(t, err, "load notification owner")
require.EqualValues(t, 1, delivered.Load())
})
}

View File

@@ -1,150 +0,0 @@
package alerts
import (
"context"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/netip"
"sync/atomic"
"syscall"
"time"
"github.com/nicholas-fedor/shoutrrr/pkg/router"
"github.com/nicholas-fedor/shoutrrr/pkg/types"
)
var (
errInternalDestination = errors.New("Only admins can send to internal destinations")
errUnrestrictedService = errors.New("Only admins can use this notification service") // Restrict services w/o custom connection support
publicNotificationDialer = &net.Dialer{
Timeout: 10 * time.Second,
// Control checks each resolved address immediately before connecting.
Control: func(_, address string, _ syscall.RawConn) error { return checkNotificationAddress(address) },
}
publicNotificationClient = newPublicNotificationClient()
)
func newPublicNotificationClient() *http.Client {
return &http.Client{
Timeout: 15 * time.Second,
Transport: &http.Transport{
// Do not use proxies: they can resolve the target themselves and
// bypass the destination check on our socket.
DialContext: publicNotificationDialer.DialContext,
TLSHandshakeTimeout: 10 * time.Second,
IdleConnTimeout: 90 * time.Second,
},
}
}
func checkNotificationAddress(address string) error {
addr, err := netip.ParseAddrPort(address)
if err != nil || addr.Addr().Zone() != "" {
return errInternalDestination
}
ip := net.IP(addr.Addr().AsSlice())
if !ip.IsGlobalUnicast() || isInternalIP(ip) {
return errInternalDestination
}
return nil
}
func sendPublicNotification(rawURL, message string) error {
client := &notificationClient{Client: publicNotificationClient}
service, err := newPublicNotificationService(rawURL, types.SenderOptions{HTTPClient: client, DialContext: client.dialContext})
if err == nil {
if closer, ok := service.(io.Closer); ok {
defer closer.Close()
}
err = service.Send(message, &types.Params{})
}
// Some services format errors without preserving their error chain.
if client.blocked.Load() {
return errInternalDestination
}
return err
}
type notificationClient struct {
*http.Client
blocked atomic.Bool
}
func (c *notificationClient) Do(req *http.Request) (*http.Response, error) {
response, err := c.Client.Do(req)
if errors.Is(err, errInternalDestination) {
c.blocked.Store(true)
}
return response, err
}
func (c *notificationClient) dialContext(ctx context.Context, network, address string) (net.Conn, error) {
conn, err := publicNotificationDialer.DialContext(ctx, network, address)
if errors.Is(err, errInternalDestination) {
c.blocked.Store(true)
}
return conn, err
}
func newPublicNotificationService(rawURL string, opts types.SenderOptions) (types.Service, error) {
r := &router.ServiceRouter{}
scheme, serviceURL, err := r.ExtractServiceName(rawURL)
if err != nil {
return nil, err
}
service, err := r.NewService(scheme)
if err != nil {
return nil, err
}
httpSetter, httpOK := service.(types.HTTPClientSetter)
dialSetter, dialOK := service.(types.DialContextSetter)
if (!httpOK || opts.HTTPClient == nil) && (!dialOK || opts.DialContext == nil) {
return nil, errUnrestrictedService
}
if serviceURL.Scheme != scheme {
custom, ok := service.(types.CustomURLService)
if !ok {
return nil, fmt.Errorf("%w: %s", router.ErrCustomURLsNotSupported, scheme)
}
serviceURL, err = custom.GetServiceURLFromCustom(serviceURL)
if err != nil {
return nil, err
}
}
// Shoutrrr v0.20.0 CreateSenderWithOptions injects only AFTER Initialize.
// Matrix can log in during Initialize, so inject before it as well.
if httpOK {
httpSetter.SetHTTPClient(opts.HTTPClient)
}
if dialOK {
dialSetter.SetDialContext(opts.DialContext)
}
if err := service.Initialize(serviceURL, nil); err != nil {
return nil, err
}
// Some initializers replace their HTTP client with a default client.
if httpOK {
httpSetter.SetHTTPClient(opts.HTTPClient)
}
if dialOK {
dialSetter.SetDialContext(opts.DialContext)
}
return service, nil
}
var cgnatNetwork = &net.IPNet{
IP: net.IPv4(100, 64, 0, 0),
Mask: net.CIDRMask(10, 32),
}
func isInternalIP(ip net.IP) bool {
return ip.IsPrivate() ||
ip.IsLoopback() ||
ip.IsUnspecified() ||
ip.IsLinkLocalUnicast() ||
ip.IsMulticast() ||
cgnatNetwork.Contains(ip)
}

View File

@@ -1,217 +0,0 @@
package alerts
import (
"context"
"encoding/binary"
"errors"
"io"
"net"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"github.com/nicholas-fedor/shoutrrr/pkg/types"
"golang.org/x/net/dns/dnsmessage"
)
func TestCheckNotificationAddress(t *testing.T) {
for _, host := range []string{"127.0.0.1", "10.0.0.1", "172.16.0.1", "192.168.0.1", "169.254.169.254", "100.64.0.0", "100.127.255.255", "0.0.0.0", "224.0.0.1", "255.255.255.255", "::1", "::", "fc00::1", "fe80::1", "ff02::1", "::ffff:127.0.0.1", "::ffff:169.254.169.254", "fe80::1%lo", "localhost", "consul"} {
t.Run(host, func(t *testing.T) {
if err := checkNotificationAddress(net.JoinHostPort(host, "80")); !errors.Is(err, errInternalDestination) {
t.Fatalf("expected blocked address, got %v", err)
}
})
}
for _, host := range []string{"8.8.8.8", "100.63.255.255", "100.128.0.0", "2001:4860:4860::8888"} {
if err := checkNotificationAddress(net.JoinHostPort(host, "443")); err != nil {
t.Errorf("public address %s: %v", host, err)
}
}
}
func TestPublicNotificationBlocksInternalRequests(t *testing.T) {
var hits atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
hits.Add(1)
}))
defer server.Close()
host := strings.TrimPrefix(server.URL, "http://")
for _, rawURL := range []string{
"generic+http://" + host,
"generic+https://" + host,
"generic+http://localhost:" + strings.Split(host, ":")[1],
"matrix://user:password@" + host + "/room?disabletls=yes",
"mattermost://" + host + "/token?disabletls=yes",
} {
t.Run(rawURL, func(t *testing.T) {
if err := sendPublicNotification(rawURL, "test"); !errors.Is(err, errInternalDestination) {
t.Fatalf("expected internal destination error, got %v", err)
}
})
}
if hits.Load() != 0 {
t.Fatal("internal server received a request")
}
}
type notificationRoundTripper func(*http.Request) (*http.Response, error)
func (f notificationRoundTripper) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }
func TestPublicNotificationRedirect(t *testing.T) {
client := newPublicNotificationClient()
defer client.CloseIdleConnections()
transport := client.Transport
client.Transport = notificationRoundTripper(func(r *http.Request) (*http.Response, error) {
if r.URL.Host == "public.example" {
return &http.Response{StatusCode: 307, Header: http.Header{"Location": {"http://127.0.0.1/"}}, Body: io.NopCloser(strings.NewReader("")), Request: r}, nil
}
return transport.RoundTrip(r)
})
_, err := client.Get("http://public.example/")
if !errors.Is(err, errInternalDestination) {
t.Fatalf("expected redirect to be blocked, got %v", err)
}
}
func TestPublicNotificationServiceClient(t *testing.T) {
for _, rawURL := range []string{"generic+http://public.example/path", "discord://token@123456789", "slack://hook:AAAAAAAAA-BBBBBBBBB-123456789123456789123456@webhook"} {
t.Run(rawURL, func(t *testing.T) {
var hits int
client := &http.Client{Transport: notificationRoundTripper(func(r *http.Request) (*http.Response, error) {
hits++
body := `{"ok":true}`
if strings.HasPrefix(rawURL, "slack:") {
body = "ok"
}
return &http.Response{StatusCode: 200, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(body)), Request: r}, nil
})}
service, err := newPublicNotificationService(rawURL, types.SenderOptions{HTTPClient: client})
if err != nil {
t.Fatal(err)
}
if err := service.Send("test", &types.Params{}); err != nil {
t.Fatal(err)
}
if hits == 0 {
t.Fatal("injected client was not used")
}
})
}
}
func TestPublicNotificationDNS(t *testing.T) {
// Supply deterministic DNS responses over an in-memory TCP connection.
// The first lookup sees a public IP; subsequent lookups see loopback.
var rebound atomic.Bool
resolver := &net.Resolver{PreferGo: true, Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
client, server := net.Pipe()
go func() {
defer server.Close()
var size [2]byte
if _, err := io.ReadFull(server, size[:]); err != nil {
return
}
buf := make([]byte, binary.BigEndian.Uint16(size[:]))
if _, err := io.ReadFull(server, buf); err != nil {
return
}
var msg dnsmessage.Message
if err := msg.Unpack(buf); err != nil {
return
}
msg.Header.Response = true
msg.Header.RecursionAvailable = true
q := msg.Questions[0]
if q.Type == dnsmessage.TypeA {
ip := [4]byte{8, 8, 8, 8}
if rebound.Load() {
ip = [4]byte{127, 0, 0, 1}
}
msg.Answers = []dnsmessage.Resource{{Header: dnsmessage.ResourceHeader{Name: q.Name, Type: q.Type, Class: dnsmessage.ClassINET}, Body: &dnsmessage.AResource{A: ip}}}
}
buf, err := msg.Pack()
if err != nil {
return
}
binary.BigEndian.PutUint16(size[:], uint16(len(buf)))
server.Write(append(size[:], buf...))
}()
return client, nil
}}
// These tests do not run in parallel; restore the process resolver afterward.
previous := net.DefaultResolver
net.DefaultResolver = resolver
t.Cleanup(func() { net.DefaultResolver = previous })
ips, err := resolver.LookupIP(context.Background(), "ip4", "rebind.example")
if err != nil || len(ips) != 1 || !ips[0].Equal(net.IPv4(8, 8, 8, 8)) {
t.Fatalf("initial DNS lookup: %v, %v", ips, err)
}
rebound.Store(true)
client := newPublicNotificationClient()
defer client.CloseIdleConnections()
for _, host := range []string{"rebind.example", "consul"} {
guarded := &notificationClient{Client: client}
conn, dialErr := guarded.dialContext(context.Background(), "tcp", net.JoinHostPort(host, "25"))
if conn != nil {
conn.Close()
}
if !errors.Is(dialErr, errInternalDestination) || !guarded.blocked.Load() {
t.Errorf("expected TCP dial-time rejection for %s, got %v", host, dialErr)
}
_, err := client.Get("http://" + host + "/")
if !errors.Is(err, errInternalDestination) {
t.Errorf("expected dial-time rejection for %s, got %v", host, err)
}
}
}
func TestPublicNotificationTCP(t *testing.T) {
for _, rawURL := range []string{
"smtp://user:pass@HOST:25/?fromAddress=sender@example.com&toAddresses=recipient@example.com",
"smtp://user:pass@HOST:465/?fromAddress=sender@example.com&toAddresses=recipient@example.com",
"mqtt://HOST:1883/topic",
"mqtts://HOST:8883/topic",
} {
t.Run(rawURL, func(t *testing.T) {
t.Parallel()
t.Run("internal destination", func(t *testing.T) {
err := sendPublicNotification(strings.ReplaceAll(rawURL, "HOST", "127.0.0.1"), "test")
if !errors.Is(err, errInternalDestination) {
t.Fatalf("expected blocked destination, got %v", err)
}
})
t.Run("public destination uses injected dialer", func(t *testing.T) {
var calls atomic.Int32
stopped := errors.New("test dial stopped")
service, err := newPublicNotificationService(strings.ReplaceAll(rawURL, "HOST", "8.8.8.8"), types.SenderOptions{
DialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
calls.Add(1)
if network != "tcp" || !strings.HasPrefix(address, "8.8.8.8:") {
t.Errorf("unexpected dial: %s %s", network, address)
}
if err := checkNotificationAddress(address); err != nil {
t.Error(err)
}
return nil, stopped
},
})
if err != nil {
t.Fatal(err)
}
if closer, ok := service.(io.Closer); ok {
defer closer.Close()
}
if err := service.Send("test", &types.Params{}); err == nil {
t.Fatal("expected dial failure")
}
if calls.Load() == 0 {
t.Fatal("custom dialer was not used")
}
})
})
}
}

View File

@@ -43,11 +43,6 @@ COPY --from=builder /app/agent/test-data/amdgpu.ids /usr/share/libdrm/amdgpu.ids
# Copy smartmontools binaries and config files # Copy smartmontools binaries and config files
COPY --from=smartmontools-builder /usr/sbin/smartctl /usr/sbin/smartctl COPY --from=smartmontools-builder /usr/sbin/smartctl /usr/sbin/smartctl
# Install ZFS userspace utilities (zpool, zfs) for pool/dataset monitoring
RUN apt-get update && apt-get install -y --no-install-recommends \
zfsutils-linux \
&& apt-get clean && rm -rf /var/lib/apt/lists/*
# Ensure data persistence across container recreations # Ensure data persistence across container recreations
VOLUME ["/var/lib/beszel-agent"] VOLUME ["/var/lib/beszel-agent"]

View File

@@ -65,32 +65,6 @@ RUN set -eux; \
cp -v "$interp" "/out/rootfs$interp"; \ cp -v "$interp" "/out/rootfs$interp"; \
fi fi
# --------------------------
# ZFS utilities builder stage
# --------------------------
FROM --platform=$TARGETPLATFORM debian:bookworm-slim AS zfsutils-builder
RUN apt-get update && apt-get install -y --no-install-recommends \
zfsutils-linux \
&& rm -rf /var/lib/apt/lists/*
# Copy the zpool/zfs binaries and their required runtime libraries
RUN set -eux; \
mkdir -p /out/rootfs/lib /out/rootfs/lib64 /out/rootfs/usr/lib; \
for bin in /usr/sbin/zpool /usr/sbin/zfs; do \
mkdir -p "/out/rootfs$(dirname "$bin")"; \
cp -v "$bin" "/out/rootfs$bin"; \
ldd "$bin" \
| awk '{print $3}' \
| grep '^/' \
| xargs -r -I '{}' sh -c 'mkdir -p "/out/rootfs$(dirname "{}")"; cp -v "{}" "/out/rootfs{}"'; \
interp="$(ldd "$bin" | awk "/ld-linux/ {print \$1}")"; \
if [ -n "$interp" ] && [ -e "$interp" ]; then \
mkdir -p "/out/rootfs$(dirname "$interp")"; \
cp -v "$interp" "/out/rootfs$interp"; \
fi; \
done
# -------------------------- # --------------------------
# Final image: lightweight multi-arch NVIDIA agent (slim) # Final image: lightweight multi-arch NVIDIA agent (slim)
# -------------------------- # --------------------------
@@ -104,9 +78,6 @@ COPY --from=builder /app/agent/test-data/amdgpu.ids /usr/share/libdrm/amdgpu.ids
COPY --from=smartmontools-builder /usr/sbin/smartctl /usr/sbin/smartctl COPY --from=smartmontools-builder /usr/sbin/smartctl /usr/sbin/smartctl
COPY --from=smartmontools-builder /out/rootfs/ / COPY --from=smartmontools-builder /out/rootfs/ /
# Copy ZFS utilities (zpool, zfs) binaries and required runtime libraries
COPY --from=zfsutils-builder /out/rootfs/ /
# nvidia-smi is intentionally not bundled. # nvidia-smi is intentionally not bundled.
# Mount the host binary instead, for example: # Mount the host binary instead, for example:
# - /usr/bin/nvidia-smi:/usr/bin/nvidia-smi:ro # - /usr/bin/nvidia-smi:/usr/bin/nvidia-smi:ro

View File

@@ -99,8 +99,8 @@ func setCollectionAuthSettings(app core.App) error {
} }
if err := applyCollectionRules(app, []string{"fingerprints"}, collectionRules{ if err := applyCollectionRules(app, []string{"fingerprints"}, collectionRules{
list: &systemScopedWriteRule, list: &systemScopedReadRule,
view: &systemScopedWriteRule, view: &systemScopedReadRule,
create: &systemScopedWriteRule, create: &systemScopedWriteRule,
update: &systemScopedWriteRule, update: &systemScopedWriteRule,
delete: &systemScopedWriteRule, delete: &systemScopedWriteRule,

View File

@@ -88,8 +88,8 @@ func TestCollectionRulesDefault(t *testing.T) {
// fingerprints collection // fingerprints collection
fingerprintsCollection, err := hub.FindCollectionByNameOrId("fingerprints") fingerprintsCollection, err := hub.FindCollectionByNameOrId("fingerprints")
require.NoError(t, err, "Failed to find fingerprints collection") require.NoError(t, err, "Failed to find fingerprints collection")
assert.Equal(t, isUserInSystemUsersNotReadonly, *fingerprintsCollection.ListRule) assert.Equal(t, isUserInSystemUsers, *fingerprintsCollection.ListRule)
assert.Equal(t, isUserInSystemUsersNotReadonly, *fingerprintsCollection.ViewRule) assert.Equal(t, isUserInSystemUsers, *fingerprintsCollection.ViewRule)
assert.Equal(t, isUserInSystemUsersNotReadonly, *fingerprintsCollection.CreateRule) assert.Equal(t, isUserInSystemUsersNotReadonly, *fingerprintsCollection.CreateRule)
assert.Equal(t, isUserInSystemUsersNotReadonly, *fingerprintsCollection.UpdateRule) assert.Equal(t, isUserInSystemUsersNotReadonly, *fingerprintsCollection.UpdateRule)
assert.Equal(t, isUserInSystemUsersNotReadonly, *fingerprintsCollection.DeleteRule) assert.Equal(t, isUserInSystemUsersNotReadonly, *fingerprintsCollection.DeleteRule)
@@ -216,8 +216,8 @@ func TestCollectionRulesShareAllSystems(t *testing.T) {
// fingerprints collection // fingerprints collection
fingerprintsCollection, err := hub.FindCollectionByNameOrId("fingerprints") fingerprintsCollection, err := hub.FindCollectionByNameOrId("fingerprints")
require.NoError(t, err, "Failed to find fingerprints collection") require.NoError(t, err, "Failed to find fingerprints collection")
assert.Equal(t, isUserNotReadonly, *fingerprintsCollection.ListRule) assert.Equal(t, isUser, *fingerprintsCollection.ListRule)
assert.Equal(t, isUserNotReadonly, *fingerprintsCollection.ViewRule) assert.Equal(t, isUser, *fingerprintsCollection.ViewRule)
assert.Equal(t, isUserNotReadonly, *fingerprintsCollection.CreateRule) assert.Equal(t, isUserNotReadonly, *fingerprintsCollection.CreateRule)
assert.Equal(t, isUserNotReadonly, *fingerprintsCollection.UpdateRule) assert.Equal(t, isUserNotReadonly, *fingerprintsCollection.UpdateRule)
assert.Equal(t, isUserNotReadonly, *fingerprintsCollection.DeleteRule) assert.Equal(t, isUserNotReadonly, *fingerprintsCollection.DeleteRule)

View File

@@ -272,15 +272,7 @@ func (sys *System) createRecords(data *system.CombinedData) (*core.Record, error
// update system record (do this last because it triggers alerts and we need above records to be inserted first) // update system record (do this last because it triggers alerts and we need above records to be inserted first)
systemRecord.Set("status", up) systemRecord.Set("status", up)
// Distinguish an idle GPU from a system without GPU data (#2312) systemRecord.Set("info", data.Info)
info := struct {
system.Info
GpuPct *float64 `json:"g,omitempty"`
}{Info: data.Info}
if len(data.Stats.GPUData) > 0 {
info.GpuPct = &data.Info.GpuPct
}
systemRecord.Set("info", info)
if err := txApp.SaveNoValidate(systemRecord); err != nil { if err := txApp.SaveNoValidate(systemRecord); err != nil {
return err return err
} }
@@ -330,11 +322,6 @@ func createSystemdStatsRecords(app core.App, data []*systemd.Service, systemId s
valueStrings := make([]string, 0, len(data)) valueStrings := make([]string, 0, len(data))
for i, service := range data { for i, service := range data {
// Agent payloads can contain null entries. Reject the snapshot before
// executing any queries so existing service records remain intact.
if service == nil {
return fmt.Errorf("null systemd service at index %d", i)
}
suffix := fmt.Sprintf("%d", i) suffix := fmt.Sprintf("%d", i)
valueStrings = append(valueStrings, fmt.Sprintf("({:id%[1]s}, {:system}, {:name%[1]s}, {:state%[1]s}, {:sub%[1]s}, {:cpu%[1]s}, {:cpuPeak%[1]s}, {:memory%[1]s}, {:memPeak%[1]s}, {:updated})", suffix)) valueStrings = append(valueStrings, fmt.Sprintf("({:id%[1]s}, {:system}, {:name%[1]s}, {:state%[1]s}, {:sub%[1]s}, {:cpu%[1]s}, {:cpuPeak%[1]s}, {:memory%[1]s}, {:memPeak%[1]s}, {:updated})", suffix))
params["id"+suffix] = makeStableHashId(systemId, service.Name) params["id"+suffix] = makeStableHashId(systemId, service.Name)

View File

@@ -1,44 +0,0 @@
//go:build testing
package systems
import (
"testing"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestCreateRecordsGPUUtilization(t *testing.T) {
sys, app := newTestSystemWithHub(t)
for _, tc := range []struct {
name string
gpu bool
usage float64
}{
{"no GPU", false, 0},
{"active GPU", true, 42.5},
{"idle GPU", true, 0},
{"GPU removed", false, 0},
} {
t.Run(tc.name, func(t *testing.T) {
data := &system.CombinedData{Info: system.Info{GpuPct: tc.usage, Cpu: 12.5}}
if tc.gpu {
data.Stats.GPUData = map[string]system.GPUData{"0": {Name: "GPU", Usage: tc.usage}}
}
_, err := sys.createRecords(data)
require.NoError(t, err)
record, err := app.FindRecordById("systems", sys.Id)
require.NoError(t, err)
var info map[string]any
require.NoError(t, record.UnmarshalJSONField("info", &info))
assert.Equal(t, 12.5, info["cpu"])
if tc.gpu {
assert.Equal(t, tc.usage, info["g"])
} else {
assert.NotContains(t, info, "g")
}
})
}
}

View File

@@ -4,7 +4,6 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"sync"
"time" "time"
"github.com/henrygd/beszel/internal/hub/ws" "github.com/henrygd/beszel/internal/hub/ws"
@@ -43,17 +42,13 @@ var errSystemExists = errors.New("system exists")
// SystemManager manages a collection of monitored systems and their connections. // SystemManager manages a collection of monitored systems and their connections.
// It handles system lifecycle, status updates, and maintains both SSH and WebSocket connections. // It handles system lifecycle, status updates, and maintains both SSH and WebSocket connections.
type SystemManager struct { type SystemManager struct {
hub hubLike // Hub interface for database and alert operations hub hubLike // Hub interface for database and alert operations
systems *store.Store[string, *System] // Thread-safe store of active systems systems *store.Store[string, *System] // Thread-safe store of active systems
sshConfig *ssh.ClientConfig // SSH client configuration for system connections sshConfig *ssh.ClientConfig // SSH client configuration for system connections
smartFetchMap *expirymap.ExpiryMap[smartFetchState] // Stores last SMART fetch time/result; TTL is only for cleanup smartFetchMap *expirymap.ExpiryMap[smartFetchState] // Stores last SMART fetch time/result; TTL is only for cleanup
zfsFetchMap *expirymap.ExpiryMap[zfsFetchState] // Stores last ZFS fetch time/result; TTL is only for cleanup zfsFetchMap *expirymap.ExpiryMap[zfsFetchState] // Stores last ZFS fetch time/result; TTL is only for cleanup
realtimeMutex sync.Mutex // Protects all realtime worker and subscription state ctx context.Context // Cancelled when the app terminates
activeSubscriptions map[string]*subscriptionInfo // Realtime subscriptions keyed by system ID cancel context.CancelFunc // Cancels ctx and all child system contexts
realtimeWorkerStop chan struct{} // Stops the current realtime worker generation
realtimeWorkerRun bool // Whether a realtime worker has been started
ctx context.Context // Cancelled when the app terminates
cancel context.CancelFunc // Cancels ctx and all child system contexts
} }
// hubLike defines the interface requirements for the hub dependency. // hubLike defines the interface requirements for the hub dependency.
@@ -72,11 +67,10 @@ type hubLike interface {
// The hub must implement the hubLike interface to provide database and alert functionality. // The hub must implement the hubLike interface to provide database and alert functionality.
func NewSystemManager(hub hubLike) *SystemManager { func NewSystemManager(hub hubLike) *SystemManager {
sm := &SystemManager{ sm := &SystemManager{
systems: store.New(map[string]*System{}), systems: store.New(map[string]*System{}),
hub: hub, hub: hub,
smartFetchMap: expirymap.New[smartFetchState](time.Hour), smartFetchMap: expirymap.New[smartFetchState](time.Hour),
zfsFetchMap: expirymap.New[zfsFetchState](time.Hour), zfsFetchMap: expirymap.New[zfsFetchState](time.Hour),
activeSubscriptions: make(map[string]*subscriptionInfo),
} }
sm.ctx, sm.cancel = context.WithCancel(context.Background()) sm.ctx, sm.cancel = context.WithCancel(context.Background())
return sm return sm
@@ -144,7 +138,6 @@ func (sm *SystemManager) bindEventHooks() {
// onTerminate cancels SystemManager context on app shutdown // onTerminate cancels SystemManager context on app shutdown
func (sm *SystemManager) onTerminate(e *core.TerminateEvent) error { func (sm *SystemManager) onTerminate(e *core.TerminateEvent) error {
sm.cancel() sm.cancel()
sm.stopRealtimeWorker()
return e.Next() return e.Next()
} }

View File

@@ -3,27 +3,25 @@ package systems
import ( import (
"encoding/json" "encoding/json"
"strings" "strings"
"sync"
"time" "time"
"github.com/henrygd/beszel/internal/common" "github.com/henrygd/beszel/internal/common"
"github.com/henrygd/beszel/internal/hub/utils"
"github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/apis"
"github.com/pocketbase/pocketbase/core" "github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/tools/subscriptions" "github.com/pocketbase/pocketbase/tools/subscriptions"
) )
type subscriptionInfo struct { type subscriptionInfo struct {
subscription string subscription string
connectedClients int connectedClients uint8
fetching bool
} }
type realtimeFetch struct { var (
systemID string activeSubscriptions = make(map[string]*subscriptionInfo)
subscription string workerRunning bool
info *subscriptionInfo tickerStopChan chan struct{}
} realtimeMutex sync.Mutex
)
// onRealtimeConnectRequest handles client connection events for realtime subscriptions. // onRealtimeConnectRequest handles client connection events for realtime subscriptions.
// It cleans up existing subscriptions when a client connects. // It cleans up existing subscriptions when a client connects.
@@ -40,19 +38,6 @@ func (sm *SystemManager) onRealtimeConnectRequest(e *core.RealtimeConnectRequest
// onRealtimeSubscribeRequest handles client subscription events for realtime metrics. // onRealtimeSubscribeRequest handles client subscription events for realtime metrics.
// It tracks new subscriptions and unsubscriptions to manage the realtime worker lifecycle. // It tracks new subscriptions and unsubscriptions to manage the realtime worker lifecycle.
func (sm *SystemManager) onRealtimeSubscribeRequest(e *core.RealtimeSubscribeRequestEvent) error { func (sm *SystemManager) onRealtimeSubscribeRequest(e *core.RealtimeSubscribeRequestEvent) error {
// Parse with PocketBase's own subscription parser before changing the real
// client. Reject the entire request if any metrics target is inaccessible.
requested := subscriptions.NewDefaultClient()
requested.Subscribe(e.Subscriptions...)
for topic, options := range requested.Subscriptions() {
if !strings.HasPrefix(topic, "rt_metrics") {
continue
}
system, err := sm.GetSystem(options.Query["system"])
if err != nil || !system.HasUser(e.App, e.Auth) {
return e.NotFoundError("", nil)
}
}
oldSubs := e.Client.Subscriptions() oldSubs := e.Client.Subscriptions()
// after e.Next() is the result of the subscribe request // after e.Next() is the result of the subscribe request
err := e.Next() err := e.Next()
@@ -62,7 +47,14 @@ func (sm *SystemManager) onRealtimeSubscribeRequest(e *core.RealtimeSubscribeReq
for k, options := range newSubs { for k, options := range newSubs {
if _, ok := oldSubs[k]; !ok { if _, ok := oldSubs[k]; !ok {
if strings.HasPrefix(k, "rt_metrics") { if strings.HasPrefix(k, "rt_metrics") {
sm.addRealtimeSubscription(options.Query["system"], k) systemId := options.Query["system"]
if _, ok := activeSubscriptions[systemId]; !ok {
activeSubscriptions[systemId] = &subscriptionInfo{
subscription: k,
}
}
activeSubscriptions[systemId].connectedClients += 1
sm.onRealtimeSubscriptionAdded()
} }
} }
} }
@@ -76,76 +68,72 @@ func (sm *SystemManager) onRealtimeSubscribeRequest(e *core.RealtimeSubscribeReq
return err return err
} }
// addRealtimeSubscription tracks a subscriber and starts a worker if necessary. // onRealtimeSubscriptionAdded initializes or starts the realtime worker when the first subscription is added.
func (sm *SystemManager) addRealtimeSubscription(systemID, subscription string) { // It ensures only one worker runs at a time.
sm.realtimeMutex.Lock() func (sm *SystemManager) onRealtimeSubscriptionAdded() {
defer sm.realtimeMutex.Unlock() realtimeMutex.Lock()
defer realtimeMutex.Unlock()
if sm.activeSubscriptions == nil { // Start the worker if it's not already running
sm.activeSubscriptions = make(map[string]*subscriptionInfo) if !workerRunning {
} workerRunning = true
info, ok := sm.activeSubscriptions[systemID] // Create a new stop channel for this worker instance
if !ok { tickerStopChan = make(chan struct{})
info = &subscriptionInfo{subscription: subscription} go sm.startRealtimeWorker()
sm.activeSubscriptions[systemID] = info
}
info.connectedClients++
if !sm.realtimeWorkerRun {
sm.realtimeWorkerRun = true
stop := make(chan struct{})
sm.realtimeWorkerStop = stop
go sm.startRealtimeWorker(stop)
} }
} }
// stopRealtimeWorker stops the current worker generation, if any. // checkSubscriptions stops the realtime worker when there are no active subscriptions.
func (sm *SystemManager) stopRealtimeWorker() { // This prevents unnecessary resource usage when no clients are listening for realtime data.
sm.realtimeMutex.Lock() func (sm *SystemManager) checkSubscriptions() {
defer sm.realtimeMutex.Unlock() if !workerRunning || len(activeSubscriptions) > 0 {
sm.stopRealtimeWorkerLocked()
}
func (sm *SystemManager) stopRealtimeWorkerLocked() {
if !sm.realtimeWorkerRun {
return return
} }
close(sm.realtimeWorkerStop)
sm.realtimeWorkerStop = nil realtimeMutex.Lock()
sm.realtimeWorkerRun = false defer realtimeMutex.Unlock()
// Signal the worker to stop
if tickerStopChan != nil {
select {
case tickerStopChan <- struct{}{}:
default:
}
}
// Mark worker as stopped (will be reset when next subscription comes in)
workerRunning = false
} }
// removeRealtimeSubscription removes a realtime subscription and checks if the worker should be stopped. // removeRealtimeSubscription removes a realtime subscription and checks if the worker should be stopped.
// It only processes subscriptions with the "rt_metrics" prefix and triggers cleanup when subscriptions are removed. // It only processes subscriptions with the "rt_metrics" prefix and triggers cleanup when subscriptions are removed.
func (sm *SystemManager) removeRealtimeSubscription(subscription string, options subscriptions.SubscriptionOptions) { func (sm *SystemManager) removeRealtimeSubscription(subscription string, options subscriptions.SubscriptionOptions) {
if strings.HasPrefix(subscription, "rt_metrics") { if strings.HasPrefix(subscription, "rt_metrics") {
systemID := options.Query["system"] systemId := options.Query["system"]
sm.realtimeMutex.Lock() if info, ok := activeSubscriptions[systemId]; ok {
if info, ok := sm.activeSubscriptions[systemID]; ok { info.connectedClients -= 1
info.connectedClients--
if info.connectedClients <= 0 { if info.connectedClients <= 0 {
delete(sm.activeSubscriptions, systemID) delete(activeSubscriptions, systemId)
} }
} }
if len(sm.activeSubscriptions) == 0 { sm.checkSubscriptions()
sm.stopRealtimeWorkerLocked()
}
sm.realtimeMutex.Unlock()
} }
} }
// startRealtimeWorker runs the main loop for fetching realtime data from agents. // startRealtimeWorker runs the main loop for fetching realtime data from agents.
// It continuously fetches system data and broadcasts it to subscribed clients via WebSocket. // It continuously fetches system data and broadcasts it to subscribed clients via WebSocket.
func (sm *SystemManager) startRealtimeWorker(stop <-chan struct{}) { func (sm *SystemManager) startRealtimeWorker() {
sm.fetchRealtimeDataAndNotify() sm.fetchRealtimeDataAndNotify()
ticker := time.NewTicker(time.Second) tick := time.Tick(1 * time.Second)
defer ticker.Stop()
for { for {
select { select {
case <-stop: case <-tickerStopChan:
return return
case <-ticker.C: case <-tick:
if len(activeSubscriptions) == 0 {
return
}
sm.fetchRealtimeDataAndNotify() sm.fetchRealtimeDataAndNotify()
} }
} }
@@ -153,79 +141,27 @@ func (sm *SystemManager) startRealtimeWorker(stop <-chan struct{}) {
// fetchRealtimeDataAndNotify fetches realtime data for all active subscriptions and notifies the clients. // fetchRealtimeDataAndNotify fetches realtime data for all active subscriptions and notifies the clients.
func (sm *SystemManager) fetchRealtimeDataAndNotify() { func (sm *SystemManager) fetchRealtimeDataAndNotify() {
for _, fetch := range sm.claimRealtimeFetches() { for systemId, info := range activeSubscriptions {
system, err := sm.GetSystem(fetch.systemID) system, err := sm.GetSystem(systemId)
if err != nil { if err != nil {
sm.finishRealtimeFetch(fetch)
continue continue
} }
go func(fetch realtimeFetch) { go func() {
defer sm.finishRealtimeFetch(fetch)
data, err := system.fetchDataFromAgent(common.DataRequestOptions{CacheTimeMs: 1000}) data, err := system.fetchDataFromAgent(common.DataRequestOptions{CacheTimeMs: 1000})
if err != nil { if err != nil {
return return
} }
bytes, err := json.Marshal(data) bytes, err := json.Marshal(data)
if err == nil { if err == nil {
notify(sm.hub, system, fetch.subscription, bytes) notify(sm.hub, info.subscription, bytes)
} }
}(fetch) }()
}
}
// claimRealtimeFetches takes a stable snapshot and marks each selected system as
// in flight. Slow agents are skipped on later ticks until their fetch completes.
func (sm *SystemManager) claimRealtimeFetches() []realtimeFetch {
sm.realtimeMutex.Lock()
defer sm.realtimeMutex.Unlock()
fetches := make([]realtimeFetch, 0, len(sm.activeSubscriptions))
for systemID, info := range sm.activeSubscriptions {
if info.fetching {
continue
}
info.fetching = true
fetches = append(fetches, realtimeFetch{
systemID: systemID,
subscription: info.subscription,
info: info,
})
}
return fetches
}
func (sm *SystemManager) finishRealtimeFetch(fetch realtimeFetch) {
sm.realtimeMutex.Lock()
defer sm.realtimeMutex.Unlock()
// A subscription may have been removed and recreated while the old request
// was running. Only release the exact entry claimed by this request.
if info := sm.activeSubscriptions[fetch.systemID]; info == fetch.info {
info.fetching = false
} }
} }
// notify broadcasts realtime data to all clients subscribed to a specific subscription. // notify broadcasts realtime data to all clients subscribed to a specific subscription.
// Custom topics bypass collection rules, so check current access for every // It iterates through all connected clients and sends the data only to those with matching subscriptions.
// recipient, including clients whose authentication or membership was revoked. func notify(app core.App, subscription string, data []byte) error {
func notify(app core.App, system *System, subscription string, data []byte) error {
shareAll, _ := utils.GetEnv("SHARE_ALL_SYSTEMS")
members := make(map[string]struct{})
if shareAll != "true" {
// Refresh once per broadcast so membership changes take effect on the
// next update without querying the database for every recipient.
var recordData struct{ Users string }
if err := app.DB().NewQuery("SELECT users FROM systems WHERE id={:id}").
Bind(dbx.Params{"id": system.Id}).One(&recordData); err != nil {
return err
}
var userIDs []string
if err := json.Unmarshal([]byte(recordData.Users), &userIDs); err != nil {
return err
}
for _, id := range userIDs {
members[id] = struct{}{}
}
}
message := subscriptions.Message{ message := subscriptions.Message{
Name: subscription, Name: subscription,
Data: data, Data: data,
@@ -234,13 +170,6 @@ func notify(app core.App, system *System, subscription string, data []byte) erro
if !client.HasSubscription(subscription) { if !client.HasSubscription(subscription) {
continue continue
} }
auth, _ := client.Get(apis.RealtimeClientAuthKey).(*core.Record)
if auth == nil {
continue
}
if _, member := members[auth.Id]; shareAll != "true" && !member {
continue
}
client.Send(message) client.Send(message)
} }
return nil return nil

View File

@@ -1,229 +0,0 @@
package systems
import (
"testing"
"time"
"github.com/pocketbase/pocketbase/apis"
"github.com/pocketbase/pocketbase/core"
pbtests "github.com/pocketbase/pocketbase/tests"
"github.com/pocketbase/pocketbase/tools/hook"
"github.com/pocketbase/pocketbase/tools/store"
"github.com/pocketbase/pocketbase/tools/subscriptions"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestRealtimeAuthorization(t *testing.T) {
t.Setenv("SHARE_ALL_SYSTEMS", "false")
t.Setenv("BESZEL_HUB_SHARE_ALL_SYSTEMS", "")
app, err := pbtests.NewTestApp(t.TempDir())
require.NoError(t, err)
t.Cleanup(app.Cleanup)
_, err = app.DB().NewQuery(`CREATE TABLE IF NOT EXISTS systems (id TEXT PRIMARY KEY, users TEXT)`).Execute()
require.NoError(t, err)
_, err = app.DB().NewQuery(`INSERT INTO systems (id, users) VALUES ('target', '["member"]')`).Execute()
require.NoError(t, err)
member := core.NewRecord(core.NewAuthCollection("users"))
member.Id = "member"
outsider := core.NewRecord(member.Collection())
outsider.Id = "outsider"
system := &System{Id: "target"}
sm := newRealtimeTestManager()
sm.systems.Set(system.Id, system)
// Keep the lifecycle bookkeeping active without starting an agent worker.
sm.realtimeWorkerRun = true
sm.realtimeWorkerStop = make(chan struct{})
t.Cleanup(sm.stopRealtimeWorker)
topic := `rt_metrics?options={"query":{"system":"target"}}`
for _, tc := range []struct {
name string
auth *core.Record
topic string
share bool
allowed bool
}{
{"guest", nil, topic, false, false},
{"outsider", outsider, topic, false, false},
{"member", member, topic, false, true},
{"missing system", member, `rt_metrics`, false, false},
{"unknown system", member, `rt_metrics?options={"query":{"system":"missing"}}`, false, false},
{"malformed options", member, `rt_metrics?options=invalid`, false, false},
{"prefix variant", outsider, `rt_metrics_extra?options={"query":{"system":"target"}}`, false, false},
{"shared outsider", outsider, topic, true, true},
{"shared guest", nil, topic, true, false},
{"other topic", nil, "systems/*", false, true},
} {
t.Run(tc.name, func(t *testing.T) {
if tc.share {
t.Setenv("BESZEL_HUB_SHARE_ALL_SYSTEMS", "true")
}
client := subscriptions.NewDefaultClient()
client.Subscribe("existing")
e := &core.RealtimeSubscribeRequestEvent{
RequestEvent: &core.RequestEvent{App: app, Auth: tc.auth},
Client: client, Subscriptions: []string{tc.topic},
}
called := false
h := &hook.Hook[*core.RealtimeSubscribeRequestEvent]{}
h.BindFunc(sm.onRealtimeSubscribeRequest)
err := h.Trigger(e, func(e *core.RealtimeSubscribeRequestEvent) error {
called = true
client.Unsubscribe()
client.Subscribe(e.Subscriptions...)
return nil
})
if tc.allowed {
require.NoError(t, err)
assert.True(t, called)
} else {
require.Error(t, err)
assert.False(t, called)
assert.True(t, client.HasSubscription("existing"))
assert.False(t, client.HasSubscription(tc.topic))
}
})
}
t.Run("broadcast checks current access", func(t *testing.T) {
client := subscriptions.NewDefaultClient()
client.Subscribe(topic)
app.SubscriptionsBroker().Register(client)
defer app.SubscriptionsBroker().Unregister(client.Id())
secondClient := subscriptions.NewDefaultClient()
secondClient.Subscribe(topic)
app.SubscriptionsBroker().Register(secondClient)
defer app.SubscriptionsBroker().Unregister(secondClient.Id())
check := func(auth *core.Record, allowed bool) {
t.Helper()
client.Set(apis.RealtimeClientAuthKey, auth)
secondClient.Set(apis.RealtimeClientAuthKey, auth)
done := make(chan struct{})
go func() {
notify(app, system, topic, []byte(`{"cpu":42}`))
close(done)
}()
// Even on failure, drain pending sends and join the broadcaster before
// unregistering clients, which closes their channels.
defer func() {
for {
select {
case <-client.Channel():
case <-secondClient.Channel():
case <-done:
return
}
}
}()
var received [2]int
timer := time.NewTimer(time.Second)
defer timer.Stop()
for {
select {
case msg := <-client.Channel():
received[0]++
assert.Equal(t, topic, msg.Name)
case msg := <-secondClient.Channel():
received[1]++
assert.Equal(t, topic, msg.Name)
case <-done:
want := [2]int{}
if allowed {
want = [2]int{1, 1}
}
assert.Equal(t, want, received)
return
case <-timer.C:
t.Fatal("broadcast did not finish")
}
}
}
check(nil, false)
check(outsider, false)
check(member, true)
_, err := app.DB().NewQuery(`UPDATE systems SET users = '[]'`).Execute()
require.NoError(t, err)
check(member, false)
t.Setenv("BESZEL_HUB_SHARE_ALL_SYSTEMS", "true")
check(outsider, true)
check(nil, false)
})
}
func newRealtimeTestManager() *SystemManager {
return &SystemManager{
systems: store.New(map[string]*System{}),
activeSubscriptions: make(map[string]*subscriptionInfo),
}
}
func TestRealtimeFetchesDoNotOverlapPerSystem(t *testing.T) {
sm := newRealtimeTestManager()
sm.activeSubscriptions["one"] = &subscriptionInfo{subscription: "rt_metrics_one"}
sm.activeSubscriptions["two"] = &subscriptionInfo{subscription: "rt_metrics_two"}
first := sm.claimRealtimeFetches()
require.Len(t, first, 2)
assert.Empty(t, sm.claimRealtimeFetches())
sm.finishRealtimeFetch(first[0])
next := sm.claimRealtimeFetches()
require.Len(t, next, 1)
assert.Equal(t, first[0].systemID, next[0].systemID)
sm.finishRealtimeFetch(first[1])
sm.finishRealtimeFetch(next[0])
}
func TestFinishingOldRealtimeFetchDoesNotReleaseReplacement(t *testing.T) {
sm := newRealtimeTestManager()
oldInfo := &subscriptionInfo{subscription: "old"}
sm.activeSubscriptions["system"] = oldInfo
fetch := sm.claimRealtimeFetches()[0]
newInfo := &subscriptionInfo{subscription: "new", fetching: true}
sm.activeSubscriptions["system"] = newInfo
sm.finishRealtimeFetch(fetch)
assert.True(t, newInfo.fetching)
}
func TestRealtimeSubscriptionLifecycle(t *testing.T) {
sm := newRealtimeTestManager()
options := subscriptions.SubscriptionOptions{Query: map[string]string{"system": "system"}}
sm.addRealtimeSubscription("system", "rt_metrics")
sm.addRealtimeSubscription("system", "rt_metrics")
sm.realtimeMutex.Lock()
firstStop := sm.realtimeWorkerStop
assert.True(t, sm.realtimeWorkerRun)
assert.Equal(t, 2, sm.activeSubscriptions["system"].connectedClients)
sm.realtimeMutex.Unlock()
sm.removeRealtimeSubscription("rt_metrics", options)
sm.realtimeMutex.Lock()
assert.True(t, sm.realtimeWorkerRun)
assert.Equal(t, 1, sm.activeSubscriptions["system"].connectedClients)
sm.realtimeMutex.Unlock()
sm.removeRealtimeSubscription("rt_metrics", options)
sm.realtimeMutex.Lock()
assert.False(t, sm.realtimeWorkerRun)
assert.Empty(t, sm.activeSubscriptions)
sm.realtimeMutex.Unlock()
select {
case <-firstStop:
default:
t.Fatal("worker stop channel was not closed")
}
// A later subscription must get a new stop channel owned by its worker.
sm.addRealtimeSubscription("system", "rt_metrics")
sm.realtimeMutex.Lock()
secondStop := sm.realtimeWorkerStop
assert.NotEqual(t, firstStop, secondStop)
sm.realtimeMutex.Unlock()
sm.stopRealtimeWorker()
}

View File

@@ -3,11 +3,9 @@
package systems_test package systems_test
import ( import (
"encoding/json"
"testing" "testing"
"time" "time"
"github.com/fxamacker/cbor/v2"
"github.com/henrygd/beszel/internal/entities/system" "github.com/henrygd/beszel/internal/entities/system"
"github.com/henrygd/beszel/internal/entities/systemd" "github.com/henrygd/beszel/internal/entities/systemd"
"github.com/henrygd/beszel/internal/hub/systems" "github.com/henrygd/beszel/internal/hub/systems"
@@ -17,42 +15,6 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
func TestCreateRecordsRejectsNullSystemdService(t *testing.T) {
hub, user := tests.GetHubWithUser(t)
defer hub.Cleanup()
records, err := tests.CreateSystems(hub, 1, user.Id, "paused")
require.NoError(t, err)
sys, err := hub.GetSystemManager().GetSystem(records[0].Id)
require.NoError(t, err)
require.NoError(t, systems.CreateSystemdStatsRecords(hub, []*systemd.Service{
{Name: "existing.service", State: systemd.StatusFailed},
}, records[0].Id))
for _, services := range []string{`[null]`, `[{"name":"new.service"},null]`, `[null,{"name":"new.service"}]`} {
for _, encoding := range []string{"json", "cbor"} {
t.Run(encoding+"/"+services, func(t *testing.T) {
var data system.CombinedData
require.NoError(t, json.Unmarshal([]byte(`{"systemd":`+services+`}`), &data))
if encoding == "cbor" {
encoded, err := cbor.Marshal(data)
require.NoError(t, err)
data = system.CombinedData{}
require.NoError(t, cbor.Unmarshal(encoded, &data))
}
_, err := sys.CreateRecords(&data)
require.ErrorContains(t, err, "null systemd service")
var names []string
require.NoError(t, hub.DB().Select("name").From("systemd_services").
Where(dbx.HashExp{"system": records[0].Id}).Column(&names))
assert.Equal(t, []string{"existing.service"}, names)
count, err := hub.CountRecords("system_stats", dbx.HashExp{"system": records[0].Id})
require.NoError(t, err)
assert.Zero(t, count, "invalid snapshot must roll back system stats")
})
}
}
}
func TestCreateRecordsHandlesSystemdAlertLifecycle(t *testing.T) { func TestCreateRecordsHandlesSystemdAlertLifecycle(t *testing.T) {
hub, user := tests.GetHubWithUser(t) hub, user := tests.GetHubWithUser(t)
defer hub.Cleanup() defer hub.Cleanup()

View File

@@ -75,12 +75,7 @@ export const smartColumns: ColumnDef<SmartAttribute>[] = [
header: "Name", header: "Name",
}, },
{ {
accessorFn: (row) => { accessorFn: (row) => row.rs || row.rv?.toString(),
if (row.n === "DataUnitsWritten" || row.n === "DataUnitsRead") {
return formatDataUnits(Number(row.rv ?? 0))
}
return row.rs || row.rv?.toString()
},
header: "Value", header: "Value",
}, },
{ {
@@ -108,12 +103,6 @@ function formatCapacity(bytes: number): string {
return `${toFixedFloat(value, value >= 10 ? 1 : 2)} ${unit}` return `${toFixedFloat(value, value >= 10 ? 1 : 2)} ${unit}`
} }
// Function to format NVMe data units
// (1 unit = 1000 * 512 bytes) as a human-readable size
function formatDataUnits(units: number): string {
return formatCapacity(units * 1000 * 512)
}
const SMART_DEVICE_FIELDS = "id,system,name,model,state,capacity,temp,type,hours,cycles,updated" const SMART_DEVICE_FIELDS = "id,system,name,model,state,capacity,temp,type,hours,cycles,updated"
export const createColumns = ( export const createColumns = (

View File

@@ -193,7 +193,7 @@ export function SystemsTableColumns(viewMode: "table" | "grid"): ColumnDef<Syste
header: sortableHeader, header: sortableHeader,
}, },
{ {
accessorFn: ({ info }) => info.g, accessorFn: ({ info }) => info.g || undefined,
id: "gpu", id: "gpu",
name: () => "GPU", name: () => "GPU",
cell: (info) => { cell: (info) => {

View File

@@ -4,7 +4,7 @@ import { basePath } from "@/components/router"
import { toast } from "@/components/ui/use-toast" import { toast } from "@/components/ui/use-toast"
import type { ChartTimes, UserSettings } from "@/types" import type { ChartTimes, UserSettings } from "@/types"
import { $alerts, $allSystemsById, $allSystemsByName, $userSettings } from "./stores" import { $alerts, $allSystemsById, $allSystemsByName, $userSettings } from "./stores"
import { chartTimeData, debounce } from "./utils" import { chartTimeData } from "./utils"
/** PocketBase JS Client */ /** PocketBase JS Client */
export const pb = new PocketBase(basePath) export const pb = new PocketBase(basePath)
@@ -12,7 +12,7 @@ export const pb = new PocketBase(basePath)
export const isAdmin = () => pb.authStore.record?.role === "admin" export const isAdmin = () => pb.authStore.record?.role === "admin"
export const isReadOnlyUser = () => pb.authStore.record?.role === "readonly" export const isReadOnlyUser = () => pb.authStore.record?.role === "readonly"
const verifyAuth = () => { export const verifyAuth = () => {
pb.collection("users") pb.collection("users")
.authRefresh() .authRefresh()
.catch(() => { .catch(() => {
@@ -25,22 +25,6 @@ const verifyAuth = () => {
}) })
} }
const verifyAuthDebounced = debounce(verifyAuth, 100)
// verify the session whenever any API request returns a 4xx response (e.g. an
// expired JWT). The auth-refresh endpoint is excluded to avoid a loop, since
// it returns 401 itself when the token is no longer valid.
pb.afterSend = (response, data) => {
if (
(response.status === 401 || response.status === 403) &&
pb.authStore.token &&
!response.url.includes("auth-refresh")
) {
verifyAuthDebounced()
}
return data
}
/** Logs the user out by clearing the auth store and unsubscribing from realtime updates. */ /** Logs the user out by clearing the auth store and unsubscribing from realtime updates. */
export function logOut() { export function logOut() {
$allSystemsByName.set({}) $allSystemsByName.set({})

View File

@@ -1,6 +1,6 @@
/** biome-ignore-all lint/suspicious/noAssignInExpressions: it's fine :) */ /** biome-ignore-all lint/suspicious/noAssignInExpressions: it's fine :) */
import type { PreinitializedMapStore } from "nanostores" import type { PreinitializedMapStore } from "nanostores"
import { pb } from "@/lib/api" import { pb, verifyAuth } from "@/lib/api"
import { import {
$allSystemsById, $allSystemsById,
$allSystemsByName, $allSystemsByName,
@@ -167,6 +167,11 @@ export async function subscribe() {
export async function refresh() { export async function refresh() {
try { try {
const records = await fetchSystems() const records = await fetchSystems()
if (!records.length) {
// No systems found, verify authentication
verifyAuth()
return
}
for (const record of records) { for (const record of records) {
add(record) add(record)
} }

View File

@@ -812,11 +812,6 @@ elif is_freebsd; then
echo "Adding beszel to wheel group for self-updates" echo "Adding beszel to wheel group for self-updates"
pw group mod wheel -m beszel pw group mod wheel -m beszel
fi fi
# Add the user to the operator group for device access (SMART, /dev/xpt0, /dev/nvme*)
if pw group show operator >/dev/null 2>&1; then
echo "Adding beszel to operator group for device access"
pw group mod operator -m beszel
fi
fi fi
else else