diff --git a/cmd/beacon/main.go b/cmd/beacon/main.go index 7c54b7c1..2f170439 100644 --- a/cmd/beacon/main.go +++ b/cmd/beacon/main.go @@ -309,6 +309,7 @@ func main() { r := router.New(h, reader, []*ingest.Worker{broker1, broker2}, router.Options{ MaxConnsPerIP: resolved.MaxConnsPerIP, MaxConnectsPerMinute: resolved.MaxConnectsPerMinute, + WSAllowedOrigins: cfg.WebSocket.AllowedOrigins, CORS: cfg.CORS, Server: cfg.Server, Auth: cfg.Auth, RateLimit: resolved.RateLimit, AdminRoutes: map[string]http.Handler{ "/accounts": handlers.AccountsRouter(store), diff --git a/config.yaml.example b/config.yaml.example index de12bef4..e50906ac 100644 --- a/config.yaml.example +++ b/config.yaml.example @@ -149,6 +149,9 @@ websocket: # Within this rate budget, a full concurrent cap accepts then closes with 1013 # before hello, allowing browsers to back off without evicting existing clients. max_connects_per_minute: 10 + # Other sites allowed to open /ws. Exact scheme://host[:port], no wildcards. + #allowed_origins: + # - https://example.com # Node staleness, deletion, and clock-drift thresholds. #nodes: diff --git a/internal/api/router/options.go b/internal/api/router/options.go index ec9e1fbf..757a6593 100644 --- a/internal/api/router/options.go +++ b/internal/api/router/options.go @@ -14,6 +14,7 @@ import ( type Options struct { MaxConnsPerIP int MaxConnectsPerMinute int + WSAllowedOrigins []string CORS config.CORSConfig Server config.ServerConfig Auth config.AuthConfig diff --git a/internal/api/router/router.go b/internal/api/router/router.go index 447d407e..58bbfac5 100644 --- a/internal/api/router/router.go +++ b/internal/api/router/router.go @@ -91,7 +91,7 @@ func New(h *hub.Hub, reader api.Reader, workers []*ingest.Worker, opts Options) )) // ── WebSocket ──────────────────────────────────────────────────────────── - r.Get("/ws", ws.Handler(h, reader, opts.MaxConnsPerIP, opts.MaxConnectsPerMinute)) + r.Get("/ws", ws.Handler(h, reader, opts.MaxConnsPerIP, opts.MaxConnectsPerMinute, opts.WSAllowedOrigins)) // ── Public REST API (v1) ───────────────────────────────────────────────── r.Route("/api/v1", func(r chi.Router) { diff --git a/internal/config/config.go b/internal/config/config.go index 5f532726..8473c771 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -8,8 +8,10 @@ package config import ( "fmt" "net/netip" + "net/url" "os" "path/filepath" + "strings" "time" "gopkg.in/yaml.v3" @@ -241,6 +243,8 @@ type WebSocketConfig struct { // MaxConnectsPerMinute limits upgrade attempts, including failed handshakes. // Zero/omitted defaults to 10; IPv6 addresses share a /64 attempt budget. MaxConnectsPerMinute int `yaml:"max_connects_per_minute"` + // AllowedOrigins are extra exact origins allowed to open /ws. Defaults to same-host only. + AllowedOrigins []string `yaml:"allowed_origins"` } // PacketsConfig controls packet retention behaviour. @@ -396,6 +400,11 @@ func Load(path string) (*Config, error) { if cfg.WebSocket.MaxConnectsPerMinute < 0 { return nil, fmt.Errorf("websocket.max_connects_per_minute must be positive or zero for the default") } + for i, origin := range cfg.WebSocket.AllowedOrigins { + if err := validateOrigin(origin); err != nil { + return nil, fmt.Errorf("websocket.allowed_origins[%d]: %w", i, err) + } + } configDir := filepath.Dir(path) for iata, details := range cfg.IATAs { if details.BorderFile != "" && !filepath.IsAbs(details.BorderFile) { @@ -406,6 +415,20 @@ func Load(path string) (*Config, error) { return cfg, nil } +// validateOrigin rejects wildcards, which the WebSocket library would treat as patterns. +func validateOrigin(origin string) error { + u, err := url.Parse(origin) + if err != nil { + return fmt.Errorf("invalid origin %q: %w", origin, err) + } + if (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" || u.User != nil || + u.Path != "" || u.RawQuery != "" || u.Fragment != "" || u.Opaque != "" || + strings.ContainsAny(origin, `*?[]\`) { + return fmt.Errorf("origin %q must be exactly scheme://host[:port] with an http or https scheme", origin) + } + return nil +} + // Resolve returns a ResolvedConfig with defaults applied for any zero values. func Resolve(cfg *Config) ResolvedConfig { r := ResolvedConfig{ diff --git a/internal/config/ws_connect_test.go b/internal/config/ws_connect_test.go index 5914707c..b285b82b 100644 --- a/internal/config/ws_connect_test.go +++ b/internal/config/ws_connect_test.go @@ -35,3 +35,35 @@ func TestWebSocketConnectConfig(t *testing.T) { }) } } + +func TestWebSocketAllowedOrigins(t *testing.T) { + for _, tc := range []struct { + name, origin string + wantError bool + }{ + {"https", "https://example.com", false}, + {"http with port", "http://localhost:5173", false}, + {"wildcard", "*", true}, + {"wildcard host", "https://*.example.com", true}, + {"no scheme", "example.com", true}, + {"other scheme", "ftp://example.com", true}, + {"path", "https://example.com/", true}, + {"query", "https://example.com?x=1", true}, + {"userinfo", "https://user@example.com", true}, + } { + t.Run(tc.name, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "config.yaml") + yaml := "websocket: {allowed_origins: [\"" + tc.origin + "\"]}" + if err := os.WriteFile(path, []byte(yaml), 0600); err != nil { + t.Fatal(err) + } + cfg, err := Load(path) + if (err != nil) != tc.wantError { + t.Fatalf("Load error = %v, want error = %v", err, tc.wantError) + } + if err == nil && (len(cfg.WebSocket.AllowedOrigins) != 1 || cfg.WebSocket.AllowedOrigins[0] != tc.origin) { + t.Fatalf("allowed_origins = %v", cfg.WebSocket.AllowedOrigins) + } + }) + } +} diff --git a/internal/hub/hub.go b/internal/hub/hub.go index 3671afeb..d4086a1d 100644 --- a/internal/hub/hub.go +++ b/internal/hub/hub.go @@ -20,6 +20,7 @@ import ( "fmt" "log/slog" "slices" + "sync/atomic" ) // EventType identifies the kind of server-push event. These match the @@ -40,10 +41,13 @@ const ( // fields for clients that opted into them via configure (currently just // resolvedPath on packetObservation events). Left nil for event types that // don't have an opt-in variant; the hub falls back to Payload in that case. +// The WithKey variants also carry observerPublicKey. type Event struct { - Type EventType - Payload json.RawMessage - PayloadResolved json.RawMessage + Type EventType + Payload json.RawMessage + PayloadResolved json.RawMessage + PayloadWithKey json.RawMessage + PayloadResolvedWithKey json.RawMessage // Routing metadata used by the hub to match subscriptions. // Populated by the ingest layer before calling Broadcast. @@ -73,11 +77,35 @@ type Client struct { subscriptions map[string]Scope // OR semantics: event matches if it matches any scope entry // ResolvePath is a connection-wide opt-in (not per-subscription), set via - // SetResolvePath ("configure" WS messages). Freely toggleable at any + // Configure ("configure" WS messages). Freely toggleable at any // point during the connection's lifetime. Only ever read/written inside // Run(), so it needs no locking despite Client being shared with the WS // goroutines. ResolvePath bool + // IncludeObserverKey works like ResolvePath, for observerPublicKey. + IncludeObserverKey bool +} + +// ClientOptions are the connection-wide settings a "configure" message sets. +type ClientOptions struct { + ResolvePath bool + IncludeObserverKey bool +} + +// payloadFor falls back to a narrower variant if the event lacks one. +func (e Event) payloadFor(c *Client) json.RawMessage { + if c.IncludeObserverKey { + if c.ResolvePath && e.PayloadResolvedWithKey != nil { + return e.PayloadResolvedWithKey + } + if !c.ResolvePath && e.PayloadWithKey != nil { + return e.PayloadWithKey + } + } + if c.ResolvePath && e.PayloadResolved != nil { + return e.PayloadResolved + } + return e.Payload } // matches returns true if the event satisfies at least one of the client's @@ -122,10 +150,13 @@ type Hub struct { unsubscribe chan unsubscribeMsg remove chan *Client broadcast chan Event + + // observerKeyClients counts registered clients with IncludeObserverKey set. + observerKeyClients atomic.Int64 } // subscribeMsg carries a client registration, a scope subscription, or a -// configure (resolvePath toggle) request — all three go through this single +// configure request — all three go through this single // channel, not separate ones, specifically so that Go's same-channel FIFO // guarantee orders them relative to NewClient's registration message. A // separate "configure" channel raced against registration: select() has no @@ -139,7 +170,7 @@ type subscribeMsg struct { subscriptionID string isConfigure bool - resolvePath bool + options ClientOptions } type unsubscribeMsg struct { @@ -147,14 +178,6 @@ type unsubscribeMsg struct { subscriptionID string } -// configureMsg carries a connection-wide setting change, decoupled from the -// subscribe/unsubscribe scope mechanics so it can be toggled independently -// and repeatedly over the life of a connection. -type configureMsg struct { - client *Client - resolvePath bool -} - // New creates a Hub. Call Run() in a goroutine before using it. func New() *Hub { return &Hub{ @@ -192,13 +215,15 @@ func (h *Hub) RemoveScope(c *Client, id string) { h.unsubscribe <- unsubscribeMsg{client: c, subscriptionID: id} } -// SetResolvePath toggles a client's opt-in to the resolvedPath variant of -// packetObservation events. Unlike scopes, this is a single connection-wide -// flag (not additive/OR'd) and can be flipped on or off at any point during -// the connection's lifetime — takes effect on the next broadcast after the -// hub processes it. -func (h *Hub) SetResolvePath(c *Client, enabled bool) { - h.subscribe <- subscribeMsg{client: c, isConfigure: true, resolvePath: enabled} +// Configure replaces a client's connection-wide options. Unlike scopes, it is +// not additive. +func (h *Hub) Configure(c *Client, opts ClientOptions) { + h.subscribe <- subscribeMsg{client: c, isConfigure: true, options: opts} +} + +// ObserverKeyWanted lets ingest skip the key variants when nobody wants them. +func (h *Hub) ObserverKeyWanted() bool { + return h.observerKeyClients.Load() > 0 } // Remove deregisters a client and closes its Send channel. @@ -236,9 +261,17 @@ func (h *Hub) Run() { // Registration with no scope yet (NewClient path). clients[msg.client] = struct{}{} case msg.isConfigure: - // SetResolvePath path — client must already be registered. + // Configure path — client must already be registered. if _, ok := clients[msg.client]; ok { - msg.client.ResolvePath = msg.resolvePath + if msg.client.IncludeObserverKey != msg.options.IncludeObserverKey { + if msg.options.IncludeObserverKey { + h.observerKeyClients.Add(1) + } else { + h.observerKeyClients.Add(-1) + } + } + msg.client.ResolvePath = msg.options.ResolvePath + msg.client.IncludeObserverKey = msg.options.IncludeObserverKey } default: // AddScope path — client must already be registered. @@ -255,6 +288,9 @@ func (h *Hub) Run() { case c := <-h.remove: if _, ok := clients[c]; ok { delete(clients, c) + if c.IncludeObserverKey { + h.observerKeyClients.Add(-1) + } close(c.Send) close(c.laggedCH) } @@ -265,9 +301,7 @@ func (h *Hub) Run() { continue } outEvt := evt - if c.ResolvePath && evt.PayloadResolved != nil { - outEvt.Payload = evt.PayloadResolved - } + outEvt.Payload = evt.payloadFor(c) select { case c.Send <- outEvt: default: diff --git a/internal/hub/hub_test.go b/internal/hub/hub_test.go index 5ced37be..3a992238 100644 --- a/internal/hub/hub_test.go +++ b/internal/hub/hub_test.go @@ -269,7 +269,7 @@ func TestHub_ResolvePath_OptedIn_GetsResolvedPayload(t *testing.T) { h := runHub(t) c := h.NewClient() h.AddScope(c, "sub1", Scope{Events: []EventType{EventPacketObservation}}) - h.SetResolvePath(c, true) + h.Configure(c, ClientOptions{ResolvePath: true}) time.Sleep(10 * time.Millisecond) @@ -294,7 +294,7 @@ func TestHub_ResolvePath_DefaultOff_GetsBasePayload(t *testing.T) { h := runHub(t) c := h.NewClient() h.AddScope(c, "sub1", Scope{Events: []EventType{EventPacketObservation}}) - // no SetResolvePath call — default is off + // no Configure call — default is off time.Sleep(10 * time.Millisecond) @@ -320,7 +320,7 @@ func TestHub_ResolvePath_OptedIn_NoResolvedVariant_FallsBackToBase(t *testing.T) c := h.NewClient() // e.g. nodeUpdate events never carry a PayloadResolved variant h.AddScope(c, "sub1", Scope{Events: []EventType{EventNodeUpdate}}) - h.SetResolvePath(c, true) + h.Configure(c, ClientOptions{ResolvePath: true}) time.Sleep(10 * time.Millisecond) @@ -367,15 +367,70 @@ func TestHub_ResolvePath_ToggleableLive(t *testing.T) { t.Errorf("expected base payload before opting in, got %s", got) } - h.SetResolvePath(c, true) + h.Configure(c, ClientOptions{ResolvePath: true}) time.Sleep(10 * time.Millisecond) if got := broadcastAndRead(); got != `{"resolvedPath":[{"confidence":"high"}]}` { t.Errorf("expected resolved payload after opting in, got %s", got) } - h.SetResolvePath(c, false) + h.Configure(c, ClientOptions{}) time.Sleep(10 * time.Millisecond) if got := broadcastAndRead(); got != `{"resolvedPath":null}` { t.Errorf("expected base payload after opting back out, got %s", got) } } + +func TestEvent_PayloadFor(t *testing.T) { + full := Event{ + Payload: json.RawMessage(`base`), + PayloadResolved: json.RawMessage(`resolved`), + PayloadWithKey: json.RawMessage(`key`), + PayloadResolvedWithKey: json.RawMessage(`resolved+key`), + } + baseOnly := Event{Payload: json.RawMessage(`base`), PayloadResolved: json.RawMessage(`resolved`)} + for _, tc := range []struct { + name string + evt Event + opts ClientOptions + want string + }{ + {"default", full, ClientOptions{}, "base"}, + {"resolve", full, ClientOptions{ResolvePath: true}, "resolved"}, + {"key", full, ClientOptions{IncludeObserverKey: true}, "key"}, + {"both", full, ClientOptions{ResolvePath: true, IncludeObserverKey: true}, "resolved+key"}, + {"key missing", baseOnly, ClientOptions{IncludeObserverKey: true}, "base"}, + {"both, key missing", baseOnly, ClientOptions{ResolvePath: true, IncludeObserverKey: true}, "resolved"}, + } { + t.Run(tc.name, func(t *testing.T) { + c := &Client{ResolvePath: tc.opts.ResolvePath, IncludeObserverKey: tc.opts.IncludeObserverKey} + if got := string(tc.evt.payloadFor(c)); got != tc.want { + t.Fatalf("payloadFor = %q, want %q", got, tc.want) + } + }) + } +} + +func TestHub_ObserverKeyWanted_TracksOptIns(t *testing.T) { + h := runHub(t) + waitFor := func(want bool) { + t.Helper() + deadline := time.Now().Add(time.Second) + for h.ObserverKeyWanted() != want { + if time.Now().After(deadline) { + t.Fatalf("ObserverKeyWanted = %t, want %t", !want, want) + } + time.Sleep(time.Millisecond) + } + } + a, b := h.NewClient(), h.NewClient() + if h.ObserverKeyWanted() { + t.Fatal("no client opted in yet") + } + h.Configure(a, ClientOptions{IncludeObserverKey: true}) + h.Configure(a, ClientOptions{IncludeObserverKey: true}) // repeat must not double count + h.Configure(b, ClientOptions{IncludeObserverKey: true}) + waitFor(true) + h.Configure(a, ClientOptions{ResolvePath: true}) + h.Remove(b) + waitFor(false) +} diff --git a/internal/ingest/ingest.go b/internal/ingest/ingest.go index b322569b..4961344a 100644 --- a/internal/ingest/ingest.go +++ b/internal/ingest/ingest.go @@ -378,32 +378,43 @@ func (w *Worker) broadcast(eventType hub.EventType, iata string, payloadType uin }) } -// broadcastPacketObservation marshals evt twice: once as-is (the default +// broadcastPacketObservation marshals evt once as-is (the default // payload every packetObservation subscriber gets) and once with // resolvedPath populated (delivered only to connections that opted in via // the "configure" WS message; see hub.Client.ResolvePath). resolvedPath is // passed in rather than computed here because the caller already has the // path-hash resolution results in hand from other per-packet work (known // route detection, capability detection) — this adds no extra DB calls. -func (w *Worker) broadcastPacketObservation(iata string, payloadType uint8, evt packetObservationEvent, resolvedPath []api.ResolvedHop) { +func (w *Worker) broadcastPacketObservation(iata string, payloadType uint8, evt packetObservationEvent, resolvedPath []api.ResolvedHop, observerKey string) { base, err := json.Marshal(evt) if err != nil { w.log.Error("failed to marshal packetObservation event", "error", err) return } + out := hub.Event{ + Type: hub.EventPacketObservation, + Payload: base, + IATA: iata, + PayloadType: payloadType, + } evt.Observation.ResolvedPath = resolvedPath - resolved, err := json.Marshal(evt) - if err != nil { + if out.PayloadResolved, err = json.Marshal(evt); err != nil { w.log.Error("failed to marshal packetObservation event (resolved variant)", "error", err) - resolved = nil // fall back to base-only; not fatal + out.PayloadResolved = nil } - w.hub.Broadcast(hub.Event{ - Type: hub.EventPacketObservation, - Payload: base, - PayloadResolved: resolved, - IATA: iata, - PayloadType: payloadType, - }) + if w.hub.ObserverKeyWanted() { + evt.Observation.ObserverPublicKey = observerKey + if out.PayloadResolvedWithKey, err = json.Marshal(evt); err != nil { + w.log.Error("failed to marshal packetObservation event (resolved key variant)", "error", err) + out.PayloadResolvedWithKey = nil + } + evt.Observation.ResolvedPath = nil + if out.PayloadWithKey, err = json.Marshal(evt); err != nil { + w.log.Error("failed to marshal packetObservation event (key variant)", "error", err) + out.PayloadWithKey = nil + } + } + w.hub.Broadcast(out) } // parseNumber handles RSSI and SNR fields that different observer types send as diff --git a/internal/ingest/observer_key_test.go b/internal/ingest/observer_key_test.go new file mode 100644 index 00000000..afd04b01 --- /dev/null +++ b/internal/ingest/observer_key_test.go @@ -0,0 +1,80 @@ +// Copyright 2026 Beacon Contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +package ingest + +import ( + "context" + "encoding/json" + "strings" + "testing" + "time" + + "github.com/MeshCore-Beacon/beacon-server/internal/api" + "github.com/MeshCore-Beacon/beacon-server/internal/hub" +) + +func TestBroadcastPacketObservation_ObserverKeyOptIn(t *testing.T) { + const key = "ab01020304050607080910111213141516171819202122232425262728293031" + w, _ := newTestWorker() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + go w.hub.Run() + client := w.hub.NewClient() + defer w.hub.Remove(client) + w.hub.AddScope(client, "all", hub.Scope{Events: []hub.EventType{hub.EventPacketObservation, hub.EventObserverStatus}}) + waitForSummarySubscriber(t, ctx, w.hub, client) + + broadcast := func() hub.Event { + t.Helper() + w.broadcastPacketObservation("YVR", 4, packetObservationEvent{}, []api.ResolvedHop{{}}, key) + for { + select { + case evt := <-client.Send: + if evt.Type == hub.EventPacketObservation { + return evt + } + case <-ctx.Done(): + t.Fatal("packetObservation not delivered") + } + } + } + observation := func(raw json.RawMessage) map[string]json.RawMessage { + t.Helper() + var evt struct { + Observation map[string]json.RawMessage `json:"observation"` + } + if err := json.Unmarshal(raw, &evt); err != nil { + t.Fatal(err) + } + return evt.Observation + } + + evt := broadcast() + if evt.PayloadWithKey != nil || evt.PayloadResolvedWithKey != nil { + t.Fatal("key variants built with no opted-in client") + } + if strings.Contains(string(evt.Payload)+string(evt.PayloadResolved), "observerPublicKey") { + t.Fatal("default payloads must not carry observerPublicKey") + } + + // client stays plain so its Payload is still the base variant. + keyed := w.hub.NewClient() + defer w.hub.Remove(keyed) + w.hub.Configure(keyed, hub.ClientOptions{IncludeObserverKey: true}) + for !w.hub.ObserverKeyWanted() { + time.Sleep(time.Millisecond) + } + evt = broadcast() + if _, ok := observation(evt.Payload)["observerPublicKey"]; ok { + t.Fatal("base payload must not carry observerPublicKey") + } + withKey := observation(evt.PayloadWithKey) + if string(withKey["observerPublicKey"]) != `"`+key+`"` || string(withKey["resolvedPath"]) != "null" { + t.Fatalf("key variant = %v", withKey) + } + both := observation(evt.PayloadResolvedWithKey) + if string(both["observerPublicKey"]) != `"`+key+`"` || string(both["resolvedPath"]) == "null" { + t.Fatalf("resolved key variant = %v", both) + } +} diff --git a/internal/ingest/packet.go b/internal/ingest/packet.go index 657da999..e34b55d1 100644 --- a/internal/ingest/packet.go +++ b/internal/ingest/packet.go @@ -82,15 +82,17 @@ type packetObservationEvent struct { Summary *string `json:"summary,omitempty"` // same advert name as REST list/backfill rows } `json:"packet"` Observation struct { - ObserverID string `json:"observerId"` - ObserverName string `json:"observerName"` - IATA string `json:"iata"` - HeardAt int64 `json:"heardAt"` - RSSI int16 `json:"rssi"` - SNR float32 `json:"snr"` - SourceBroker string `json:"sourceBroker"` - PathBytes string `json:"pathBytes"` - PathLength struct { + ObserverID string `json:"observerId"` + ObserverName string `json:"observerName"` + // Only set for clients that sent includeObserverKey. + ObserverPublicKey string `json:"observerPublicKey,omitempty"` + IATA string `json:"iata"` + HeardAt int64 `json:"heardAt"` + RSSI int16 `json:"rssi"` + SNR float32 `json:"snr"` + SourceBroker string `json:"sourceBroker"` + PathBytes string `json:"pathBytes"` + PathLength struct { Raw string `json:"raw"` HashSize uint8 `json:"hashSize"` HopCount uint8 `json:"hopCount"` @@ -912,7 +914,7 @@ func (w *Worker) handlePacket(ctx context.Context, iata, pubkeyHex string, raw [ resolvedPath := api.BuildResolvedPath(hashes, resolved) evt.Observation.ResolvedSource = resolvedSource evt.Observation.ResolvedDestination = resolvedDestination - w.broadcastPacketObservation(iata, packet.PayloadType(), evt, resolvedPath) + w.broadcastPacketObservation(iata, packet.PayloadType(), evt, resolvedPath, hex.EncodeToString(pubkeyBytes)) } } diff --git a/internal/ws/connect_limit_test.go b/internal/ws/connect_limit_test.go index 2e55a8b8..21d878be 100644 --- a/internal/ws/connect_limit_test.go +++ b/internal/ws/connect_limit_test.go @@ -20,7 +20,7 @@ import ( func connectTestServer(t *testing.T, maxConns, maxConnects int) (*httptest.Server, <-chan struct{}) { t.Helper() - handler := Handler(hub.New(), nil, maxConns, maxConnects) + handler := Handler(hub.New(), nil, maxConns, maxConnects, nil) done := make(chan struct{}, 64) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { handler(w, r) @@ -57,7 +57,7 @@ func TestFailedHandshakeDoesNotUseConnectionSlot(t *testing.T) { func TestConnectAttemptBudgetAndRecovery(t *testing.T) { synctest.Test(t, func(t *testing.T) { - handler := Handler(nil, nil, 1, 2) + handler := Handler(nil, nil, 1, 2, nil) attempt := func() *httptest.ResponseRecorder { request := httptest.NewRequest(http.MethodGet, "/ws", nil) request.RemoteAddr = "198.51.100.1:1234" diff --git a/internal/ws/handler.go b/internal/ws/handler.go index cb632eaa..c51a4b4a 100644 --- a/internal/ws/handler.go +++ b/internal/ws/handler.go @@ -10,14 +10,15 @@ // Client → Server: // subscribe { v, type, id, scope } → server replies subscribed { v, type, id, subscriptionId } // unsubscribe { v, type, id, subscriptionId } -// configure { v, type, id, resolvePath } → server replies configured { v, type, id, resolvePath } +// configure { v, type, id, resolvePath, includeObserverKey } +// → server replies configured { v, type, id, resolvePath, includeObserverKey } // ping { v, type, id } → server replies pong { v, type, id } // -// configure's resolvePath (bool) is a connection-wide setting, not -// per-subscription: enables/disables per-hop resolvedPath data on -// packetObservation events. Freely toggleable at any point during the -// connection; each configure call sets it to exactly the value sent -// (not additive across calls, unlike subscribe scopes). Default false. +// configure's flags are connection-wide settings, not per-subscription. +// resolvePath adds per-hop resolvedPath data to packetObservation events, +// includeObserverKey adds observation.observerPublicKey. Each configure +// sets both to exactly the values sent, so an omitted flag turns off. +// Both default false. // // Server → Client events (unsolicited): // packetObservation, observerStatus, nodeUpdate, channelMessage @@ -53,7 +54,11 @@ const ( // Handler returns an http.HandlerFunc that requires the hub to be injected. // Wire it via router.New(h) so the hub is available at startup. -func Handler(h *hub.Hub, reader api.Reader, maxConnsPerIP, maxConnectsPerMinute int) http.HandlerFunc { +func Handler(h *hub.Hub, reader api.Reader, maxConnsPerIP, maxConnectsPerMinute int, allowedOrigins []string) http.HandlerFunc { + var acceptOpts *websocket.AcceptOptions + if len(allowedOrigins) > 0 { + acceptOpts = &websocket.AcceptOptions{OriginPatterns: allowedOrigins} + } limiter := newIPLimiter(maxConnsPerIP) attempts := httprate.NewRateLimiter(maxConnectsPerMinute, time.Minute, httprate.WithResponseHeaders(httprate.ResponseHeaders{RetryAfter: "Retry-After"})) @@ -67,7 +72,7 @@ func Handler(h *hub.Hub, reader api.Reader, maxConnsPerIP, maxConnectsPerMinute if attempts.RespondOnLimit(w, r, httprate.CanonicalizeIP(ip)) { return } - conn, err := websocket.Accept(w, r, nil) + conn, err := websocket.Accept(w, r, acceptOpts) if err != nil { slog.Warn("ws: failed to accept connection", "component", "ws", "error", err) return @@ -172,10 +177,9 @@ type clientMessage struct { SubscriptionID string `json:"subscriptionId,omitempty"` Scope *subscribeScope `json:"scope,omitempty"` - // ResolvePath is only read for "configure" messages: enables/disables - // the resolvedPath variant of packetObservation events for the whole - // connection. See package doc. - ResolvePath bool `json:"resolvePath,omitempty"` + // Only read for "configure" messages. See package doc. + ResolvePath bool `json:"resolvePath,omitempty"` + IncludeObserverKey bool `json:"includeObserverKey,omitempty"` } // subscribeScope mirrors the scope object in the subscribe message. @@ -260,12 +264,13 @@ func handleClientMessage(ctx context.Context, client *hub.Client, reader api.Rea } case "configure": - h.SetResolvePath(client, msg.ResolvePath) + h.Configure(client, hub.ClientOptions{ResolvePath: msg.ResolvePath, IncludeObserverKey: msg.IncludeObserverKey}) reply, _ := json.Marshal(map[string]any{ - "v": 1, "type": "configured", "id": msg.ID, "resolvePath": msg.ResolvePath, + "v": 1, "type": "configured", "id": msg.ID, + "resolvePath": msg.ResolvePath, "includeObserverKey": msg.IncludeObserverKey, }) if slog.Default().Enabled(context.Background(), slog.LevelDebug) { - slog.Debug(fmt.Sprintf("ws[%s]: configured resolvePath=%t", connID, msg.ResolvePath), "component", "ws") + slog.Debug(fmt.Sprintf("ws[%s]: configured resolvePath=%t includeObserverKey=%t", connID, msg.ResolvePath, msg.IncludeObserverKey), "component", "ws") } if err := conn.Write(ctx, websocket.MessageText, reply); err != nil { slog.Warn(fmt.Sprintf("ws[%s]: failed to send configured reply", connID), "component", "ws", "error", err) diff --git a/internal/ws/handler_test.go b/internal/ws/handler_test.go new file mode 100644 index 00000000..8ac911f8 --- /dev/null +++ b/internal/ws/handler_test.go @@ -0,0 +1,122 @@ +// Copyright 2026 Beacon Contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +package ws + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/MeshCore-Beacon/beacon-server/internal/hub" + "github.com/coder/websocket" +) + +func TestAllowedOrigins(t *testing.T) { + for _, tc := range []struct { + name string + allowed []string + origin string + wantOK bool + }{ + {"no origin header", nil, "", true}, + {"foreign origin by default", nil, "https://example.com", false}, + {"listed origin", []string{"https://example.com"}, "https://example.com", true}, + {"listed origin, other case", []string{"https://example.com"}, "https://Example.com", true}, + {"scheme mismatch", []string{"https://example.com"}, "http://example.com", false}, + {"port mismatch", []string{"https://example.com"}, "https://example.com:8443", false}, + {"unlisted origin", []string{"https://example.com"}, "https://other.example", false}, + } { + t.Run(tc.name, func(t *testing.T) { + server := httptest.NewServer(Handler(hub.New(), nil, 5, 100, tc.allowed)) + t.Cleanup(server.Close) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + header := http.Header{} + if tc.origin != "" { + header.Set("Origin", tc.origin) + } + conn, resp, err := websocket.Dial(ctx, "ws"+strings.TrimPrefix(server.URL, "http"), &websocket.DialOptions{HTTPHeader: header}) + if tc.wantOK { + if err != nil { + t.Fatalf("dial: %v", err) + } + conn.CloseNow() + return + } + if err == nil { + conn.CloseNow() + t.Fatal("expected handshake to be refused") + } + if resp == nil || resp.StatusCode != http.StatusForbidden { + t.Fatalf("expected 403, got %v (%v)", resp, err) + } + }) + } +} + +func TestConfigureIncludeObserverKey(t *testing.T) { + h := hub.New() + go h.Run() + server := httptest.NewServer(Handler(h, nil, 5, 100, nil)) + t.Cleanup(server.Close) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + read := func(conn *websocket.Conn) map[string]json.RawMessage { + t.Helper() + _, data, err := conn.Read(ctx) + if err != nil { + t.Fatal(err) + } + var msg map[string]json.RawMessage + if err := json.Unmarshal(data, &msg); err != nil { + t.Fatal(err) + } + return msg + } + send := func(conn *websocket.Conn, msg string) { + t.Helper() + if err := conn.Write(ctx, websocket.MessageText, []byte(msg)); err != nil { + t.Fatal(err) + } + } + connect := func(configure string) *websocket.Conn { + t.Helper() + conn, _, err := websocket.Dial(ctx, "ws"+strings.TrimPrefix(server.URL, "http"), nil) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { conn.CloseNow() }) + read(conn) // hello + if configure != "" { + send(conn, configure) + reply := read(conn) + if string(reply["type"]) != `"configured"` || string(reply["includeObserverKey"]) != "true" || string(reply["resolvePath"]) != "false" { + t.Fatalf("unexpected configured reply: %v", reply) + } + } + send(conn, `{"v":1,"type":"subscribe","id":"s","scope":{}}`) + read(conn) // subscribed + return conn + } + + plain := connect("") + keyed := connect(`{"v":1,"type":"configure","id":"c","includeObserverKey":true}`) + time.Sleep(20 * time.Millisecond) // let the hub apply both subscriptions + h.Broadcast(hub.Event{ + Type: hub.EventPacketObservation, + Payload: json.RawMessage(`{"variant":"base"}`), + PayloadWithKey: json.RawMessage(`{"variant":"key"}`), + }) + if got := string(read(plain)["data"]); got != `{"variant":"base"}` { + t.Errorf("plain client got %s", got) + } + if got := string(read(keyed)["data"]); got != `{"variant":"key"}` { + t.Errorf("opted-in client got %s", got) + } +}