Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions cmd/beacon/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
3 changes: 3 additions & 0 deletions config.yaml.example
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
1 change: 1 addition & 0 deletions internal/api/router/options.go
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ import (
type Options struct {
MaxConnsPerIP int
MaxConnectsPerMinute int
WSAllowedOrigins []string
CORS config.CORSConfig
Server config.ServerConfig
Auth config.AuthConfig
Expand Down
2 changes: 1 addition & 1 deletion internal/api/router/router.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
23 changes: 23 additions & 0 deletions internal/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,8 +8,10 @@ package config
import (
"fmt"
"net/netip"
"net/url"
"os"
"path/filepath"
"strings"
"time"

"gopkg.in/yaml.v3"
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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) {
Expand All @@ -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{
Expand Down
32 changes: 32 additions & 0 deletions internal/config/ws_connect_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
})
}
}
86 changes: 60 additions & 26 deletions internal/hub/hub.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ import (
"fmt"
"log/slog"
"slices"
"sync/atomic"
)

// EventType identifies the kind of server-push event. These match the
Expand All @@ -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.
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -139,22 +170,14 @@ type subscribeMsg struct {
subscriptionID string

isConfigure bool
resolvePath bool
options ClientOptions
}

type unsubscribeMsg struct {
client *Client
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{
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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.
Expand All @@ -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)
}
Expand All @@ -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:
Expand Down
65 changes: 60 additions & 5 deletions internal/hub/hub_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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)

Expand All @@ -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)

Expand Down Expand Up @@ -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)
}
Loading
Loading