From 6cb302fcf6eb489d299707c0ca1fd796a2cfe5de Mon Sep 17 00:00:00 2001 From: henrygd Date: Sun, 6 Sep 2026 13:24:53 -0400 Subject: [PATCH] fix(hub): bound realtime metric fetching and enforce access control by system - synchronize realtime subscription state - prevent overlapping fetches per system - add realtime worker lifecycle tests - reject subscriptions without system access - test access revocation and shared-system permissions --- internal/hub/systems/system_manager.go | 29 ++- internal/hub/systems/system_realtime.go | 199 ++++++++++------ internal/hub/systems/system_realtime_test.go | 229 +++++++++++++++++++ 3 files changed, 382 insertions(+), 75 deletions(-) create mode 100644 internal/hub/systems/system_realtime_test.go diff --git a/internal/hub/systems/system_manager.go b/internal/hub/systems/system_manager.go index 5c04ea00..c1a4cde4 100644 --- a/internal/hub/systems/system_manager.go +++ b/internal/hub/systems/system_manager.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "sync" "time" "github.com/henrygd/beszel/internal/hub/ws" @@ -42,13 +43,17 @@ var errSystemExists = errors.New("system exists") // SystemManager manages a collection of monitored systems and their connections. // It handles system lifecycle, status updates, and maintains both SSH and WebSocket connections. type SystemManager struct { - hub hubLike // Hub interface for database and alert operations - systems *store.Store[string, *System] // Thread-safe store of active systems - sshConfig *ssh.ClientConfig // SSH client configuration for system connections - 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 - ctx context.Context // Cancelled when the app terminates - cancel context.CancelFunc // Cancels ctx and all child system contexts + hub hubLike // Hub interface for database and alert operations + systems *store.Store[string, *System] // Thread-safe store of active systems + sshConfig *ssh.ClientConfig // SSH client configuration for system connections + 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 + realtimeMutex sync.Mutex // Protects all realtime worker and subscription state + activeSubscriptions map[string]*subscriptionInfo // Realtime subscriptions keyed by system ID + 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. @@ -67,10 +72,11 @@ type hubLike interface { // The hub must implement the hubLike interface to provide database and alert functionality. func NewSystemManager(hub hubLike) *SystemManager { sm := &SystemManager{ - systems: store.New(map[string]*System{}), - hub: hub, - smartFetchMap: expirymap.New[smartFetchState](time.Hour), - zfsFetchMap: expirymap.New[zfsFetchState](time.Hour), + systems: store.New(map[string]*System{}), + hub: hub, + smartFetchMap: expirymap.New[smartFetchState](time.Hour), + zfsFetchMap: expirymap.New[zfsFetchState](time.Hour), + activeSubscriptions: make(map[string]*subscriptionInfo), } sm.ctx, sm.cancel = context.WithCancel(context.Background()) return sm @@ -138,6 +144,7 @@ func (sm *SystemManager) bindEventHooks() { // onTerminate cancels SystemManager context on app shutdown func (sm *SystemManager) onTerminate(e *core.TerminateEvent) error { sm.cancel() + sm.stopRealtimeWorker() return e.Next() } diff --git a/internal/hub/systems/system_realtime.go b/internal/hub/systems/system_realtime.go index 6c27d0bd..0fb6b6f9 100644 --- a/internal/hub/systems/system_realtime.go +++ b/internal/hub/systems/system_realtime.go @@ -3,25 +3,27 @@ package systems import ( "encoding/json" "strings" - "sync" "time" "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/tools/subscriptions" ) type subscriptionInfo struct { subscription string - connectedClients uint8 + connectedClients int + fetching bool } -var ( - activeSubscriptions = make(map[string]*subscriptionInfo) - workerRunning bool - tickerStopChan chan struct{} - realtimeMutex sync.Mutex -) +type realtimeFetch struct { + systemID string + subscription string + info *subscriptionInfo +} // onRealtimeConnectRequest handles client connection events for realtime subscriptions. // It cleans up existing subscriptions when a client connects. @@ -38,6 +40,19 @@ func (sm *SystemManager) onRealtimeConnectRequest(e *core.RealtimeConnectRequest // onRealtimeSubscribeRequest handles client subscription events for realtime metrics. // It tracks new subscriptions and unsubscriptions to manage the realtime worker lifecycle. 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() // after e.Next() is the result of the subscribe request err := e.Next() @@ -47,14 +62,7 @@ func (sm *SystemManager) onRealtimeSubscribeRequest(e *core.RealtimeSubscribeReq for k, options := range newSubs { if _, ok := oldSubs[k]; !ok { if strings.HasPrefix(k, "rt_metrics") { - systemId := options.Query["system"] - if _, ok := activeSubscriptions[systemId]; !ok { - activeSubscriptions[systemId] = &subscriptionInfo{ - subscription: k, - } - } - activeSubscriptions[systemId].connectedClients += 1 - sm.onRealtimeSubscriptionAdded() + sm.addRealtimeSubscription(options.Query["system"], k) } } } @@ -68,72 +76,76 @@ func (sm *SystemManager) onRealtimeSubscribeRequest(e *core.RealtimeSubscribeReq return err } -// onRealtimeSubscriptionAdded initializes or starts the realtime worker when the first subscription is added. -// It ensures only one worker runs at a time. -func (sm *SystemManager) onRealtimeSubscriptionAdded() { - realtimeMutex.Lock() - defer realtimeMutex.Unlock() +// addRealtimeSubscription tracks a subscriber and starts a worker if necessary. +func (sm *SystemManager) addRealtimeSubscription(systemID, subscription string) { + sm.realtimeMutex.Lock() + defer sm.realtimeMutex.Unlock() - // Start the worker if it's not already running - if !workerRunning { - workerRunning = true - // Create a new stop channel for this worker instance - tickerStopChan = make(chan struct{}) - go sm.startRealtimeWorker() + if sm.activeSubscriptions == nil { + sm.activeSubscriptions = make(map[string]*subscriptionInfo) + } + info, ok := sm.activeSubscriptions[systemID] + if !ok { + info = &subscriptionInfo{subscription: subscription} + sm.activeSubscriptions[systemID] = info + } + info.connectedClients++ + + if !sm.realtimeWorkerRun { + sm.realtimeWorkerRun = true + stop := make(chan struct{}) + sm.realtimeWorkerStop = stop + go sm.startRealtimeWorker(stop) } } -// checkSubscriptions stops the realtime worker when there are no active subscriptions. -// This prevents unnecessary resource usage when no clients are listening for realtime data. -func (sm *SystemManager) checkSubscriptions() { - if !workerRunning || len(activeSubscriptions) > 0 { +// stopRealtimeWorker stops the current worker generation, if any. +func (sm *SystemManager) stopRealtimeWorker() { + sm.realtimeMutex.Lock() + defer sm.realtimeMutex.Unlock() + sm.stopRealtimeWorkerLocked() +} + +func (sm *SystemManager) stopRealtimeWorkerLocked() { + if !sm.realtimeWorkerRun { return } - - realtimeMutex.Lock() - 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 + close(sm.realtimeWorkerStop) + sm.realtimeWorkerStop = nil + sm.realtimeWorkerRun = false } // 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. func (sm *SystemManager) removeRealtimeSubscription(subscription string, options subscriptions.SubscriptionOptions) { if strings.HasPrefix(subscription, "rt_metrics") { - systemId := options.Query["system"] - if info, ok := activeSubscriptions[systemId]; ok { - info.connectedClients -= 1 + systemID := options.Query["system"] + sm.realtimeMutex.Lock() + if info, ok := sm.activeSubscriptions[systemID]; ok { + info.connectedClients-- if info.connectedClients <= 0 { - delete(activeSubscriptions, systemId) + delete(sm.activeSubscriptions, systemID) } } - sm.checkSubscriptions() + if len(sm.activeSubscriptions) == 0 { + sm.stopRealtimeWorkerLocked() + } + sm.realtimeMutex.Unlock() } } // startRealtimeWorker runs the main loop for fetching realtime data from agents. // It continuously fetches system data and broadcasts it to subscribed clients via WebSocket. -func (sm *SystemManager) startRealtimeWorker() { +func (sm *SystemManager) startRealtimeWorker(stop <-chan struct{}) { sm.fetchRealtimeDataAndNotify() - tick := time.Tick(1 * time.Second) + ticker := time.NewTicker(time.Second) + defer ticker.Stop() for { select { - case <-tickerStopChan: + case <-stop: return - case <-tick: - if len(activeSubscriptions) == 0 { - return - } + case <-ticker.C: sm.fetchRealtimeDataAndNotify() } } @@ -141,27 +153,79 @@ func (sm *SystemManager) startRealtimeWorker() { // fetchRealtimeDataAndNotify fetches realtime data for all active subscriptions and notifies the clients. func (sm *SystemManager) fetchRealtimeDataAndNotify() { - for systemId, info := range activeSubscriptions { - system, err := sm.GetSystem(systemId) + for _, fetch := range sm.claimRealtimeFetches() { + system, err := sm.GetSystem(fetch.systemID) if err != nil { + sm.finishRealtimeFetch(fetch) continue } - go func() { + go func(fetch realtimeFetch) { + defer sm.finishRealtimeFetch(fetch) data, err := system.fetchDataFromAgent(common.DataRequestOptions{CacheTimeMs: 1000}) if err != nil { return } bytes, err := json.Marshal(data) if err == nil { - notify(sm.hub, info.subscription, bytes) + notify(sm.hub, system, fetch.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. -// It iterates through all connected clients and sends the data only to those with matching subscriptions. -func notify(app core.App, subscription string, data []byte) error { +// Custom topics bypass collection rules, so check current access for every +// recipient, including clients whose authentication or membership was revoked. +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{ Name: subscription, Data: data, @@ -170,6 +234,13 @@ func notify(app core.App, subscription string, data []byte) error { if !client.HasSubscription(subscription) { continue } + auth, _ := client.Get(apis.RealtimeClientAuthKey).(*core.Record) + if auth == nil { + continue + } + if _, member := members[auth.Id]; shareAll != "true" && !member { + continue + } client.Send(message) } return nil diff --git a/internal/hub/systems/system_realtime_test.go b/internal/hub/systems/system_realtime_test.go new file mode 100644 index 00000000..ea72b50b --- /dev/null +++ b/internal/hub/systems/system_realtime_test.go @@ -0,0 +1,229 @@ +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() +}