From ca44509020392549398f8ff7d9fc3d9957f5a3e2 Mon Sep 17 00:00:00 2001 From: n30nex Date: Sat, 12 Sep 2026 14:39:16 -0400 Subject: [PATCH] feat(ws): limit connection attempts and signal capacity shedding --- cmd/beacon/main.go | 3 +- config.yaml.example | 6 + docs/docs.go | 2 +- docs/swagger.json | 2 +- docs/swagger.yaml | 1 + internal/api/router/admin_config_test.go | 2 +- internal/api/router/auth_test.go | 2 +- internal/api/router/rate_limit_test.go | 12 +- internal/api/router/router.go | 4 +- internal/api/router/router_test.go | 30 ++-- internal/api/router/ws_attempt_test.go | 53 +++++++ internal/config/config.go | 15 +- internal/config/ws_connect_test.go | 37 +++++ internal/ws/connect_limit_test.go | 174 +++++++++++++++++++++++ internal/ws/handler.go | 21 ++- 15 files changed, 335 insertions(+), 29 deletions(-) create mode 100644 internal/api/router/ws_attempt_test.go create mode 100644 internal/config/ws_connect_test.go create mode 100644 internal/ws/connect_limit_test.go diff --git a/cmd/beacon/main.go b/cmd/beacon/main.go index 380f8ac..219b139 100644 --- a/cmd/beacon/main.go +++ b/cmd/beacon/main.go @@ -43,6 +43,7 @@ var version = "dev" // @version 1.6.0 // @description MeshCore network observation backend. Ingests LoRa packets from MQTT brokers, stores in PostgreSQL, and streams live events via WebSocket. // @description REST requests share a configurable per-client rate limit (default 300/minute and a 300-request one-second burst cap). Exceeded limits return HTTP 429 with error.code=rate_limited and a Retry-After header in seconds. CORS preflights and WebSocket upgrades do not consume this API budget. +// @description WebSocket upgrade attempts at /ws have a configurable per-client limit (default 10/minute, including failed handshakes). Rate exhaustion returns HTTP 429 with Retry-After before upgrade. An accepted socket exceeding the concurrent cap closes with code 1013 before hello; established connections remain open. // @termsOfService https://github.com/MeshCore-Beacon/beacon-server // @contact.name MeshCore Beacon @@ -292,7 +293,7 @@ func main() { go scheduler.Start(ctx) // ── HTTP server ────────────────────────────────────────────────────────── - r := router.New(h, reader, []*ingest.Worker{broker1, broker2}, resolved.MaxConnsPerIP, cfg.CORS, cfg.Server, cfg.Auth, resolved.RateLimit) + r := router.New(h, reader, []*ingest.Worker{broker1, broker2}, resolved.MaxConnsPerIP, resolved.MaxConnectsPerMinute, cfg.CORS, cfg.Server, cfg.Auth, resolved.RateLimit) srv := &http.Server{ Addr: addr, diff --git a/config.yaml.example b/config.yaml.example index a805495..9686dcd 100644 --- a/config.yaml.example +++ b/config.yaml.example @@ -136,6 +136,12 @@ routes: websocket: max_connections_per_ip: 5 # default: 5 + # Upgrade attempts, including failed handshakes; zero/omitted defaults to 10. + # IPv6 clients share a /64 attempt budget. Negative values prevent startup. + # Rate exhaustion rejects the handshake with 429 and Retry-After: 60. + # 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 # Node staleness, deletion, and clock-drift thresholds. #nodes: diff --git a/docs/docs.go b/docs/docs.go index 7e443c3..efc4bd1 100644 --- a/docs/docs.go +++ b/docs/docs.go @@ -4383,7 +4383,7 @@ var SwaggerInfo = &swag.Spec{ BasePath: "/api/v1", Schemes: []string{"http", "https"}, Title: "MeshCore Beacon API", - Description: "MeshCore network observation backend. Ingests LoRa packets from MQTT brokers, stores in PostgreSQL, and streams live events via WebSocket.\nREST requests share a configurable per-client rate limit (default 300/minute and a 300-request one-second burst cap). Exceeded limits return HTTP 429 with error.code=rate_limited and a Retry-After header in seconds. CORS preflights and WebSocket upgrades do not consume this API budget.", + Description: "MeshCore network observation backend. Ingests LoRa packets from MQTT brokers, stores in PostgreSQL, and streams live events via WebSocket.\nREST requests share a configurable per-client rate limit (default 300/minute and a 300-request one-second burst cap). Exceeded limits return HTTP 429 with error.code=rate_limited and a Retry-After header in seconds. CORS preflights and WebSocket upgrades do not consume this API budget.\nWebSocket upgrade attempts at /ws have a configurable per-client limit (default 10/minute, including failed handshakes). Rate exhaustion returns HTTP 429 with Retry-After before upgrade. An accepted socket exceeding the concurrent cap closes with code 1013 before hello; established connections remain open.", InfoInstanceName: "swagger", SwaggerTemplate: docTemplate, LeftDelim: "{{", diff --git a/docs/swagger.json b/docs/swagger.json index aa2eba9..018641c 100644 --- a/docs/swagger.json +++ b/docs/swagger.json @@ -5,7 +5,7 @@ ], "swagger": "2.0", "info": { - "description": "MeshCore network observation backend. Ingests LoRa packets from MQTT brokers, stores in PostgreSQL, and streams live events via WebSocket.\nREST requests share a configurable per-client rate limit (default 300/minute and a 300-request one-second burst cap). Exceeded limits return HTTP 429 with error.code=rate_limited and a Retry-After header in seconds. CORS preflights and WebSocket upgrades do not consume this API budget.", + "description": "MeshCore network observation backend. Ingests LoRa packets from MQTT brokers, stores in PostgreSQL, and streams live events via WebSocket.\nREST requests share a configurable per-client rate limit (default 300/minute and a 300-request one-second burst cap). Exceeded limits return HTTP 429 with error.code=rate_limited and a Retry-After header in seconds. CORS preflights and WebSocket upgrades do not consume this API budget.\nWebSocket upgrade attempts at /ws have a configurable per-client limit (default 10/minute, including failed handshakes). Rate exhaustion returns HTTP 429 with Retry-After before upgrade. An accepted socket exceeding the concurrent cap closes with code 1013 before hello; established connections remain open.", "title": "MeshCore Beacon API", "termsOfService": "https://github.com/MeshCore-Beacon/beacon-server", "contact": { diff --git a/docs/swagger.yaml b/docs/swagger.yaml index d10fb48..6e7923e 100644 --- a/docs/swagger.yaml +++ b/docs/swagger.yaml @@ -1247,6 +1247,7 @@ info: description: |- MeshCore network observation backend. Ingests LoRa packets from MQTT brokers, stores in PostgreSQL, and streams live events via WebSocket. REST requests share a configurable per-client rate limit (default 300/minute and a 300-request one-second burst cap). Exceeded limits return HTTP 429 with error.code=rate_limited and a Retry-After header in seconds. CORS preflights and WebSocket upgrades do not consume this API budget. + WebSocket upgrade attempts at /ws have a configurable per-client limit (default 10/minute, including failed handshakes). Rate exhaustion returns HTTP 429 with Retry-After before upgrade. An accepted socket exceeding the concurrent cap closes with code 1013 before hello; established connections remain open. license: name: AGPL-3-or-later termsOfService: https://github.com/MeshCore-Beacon/beacon-server diff --git a/internal/api/router/admin_config_test.go b/internal/api/router/admin_config_test.go index 544c8a6..fd24286 100644 --- a/internal/api/router/admin_config_test.go +++ b/internal/api/router/admin_config_test.go @@ -32,7 +32,7 @@ func TestAdminConfig(t *testing.T) { want = map[string]any{"allowed_origins": []any{"https://example.test"}, "allowed_methods": []any{"GET"}, "allowed_headers": []any{"Authorization", "X-Operator"}, "allow_credentials": true, "max_age": float64(600)} } - handler := New(nil, nil, []*ingest.Worker{{}, {}}, 5, cfg, config.ServerConfig{}, config.AuthConfig{APIKey: key}, config.ResolvedRateLimitConfig{}) + handler := New(nil, nil, []*ingest.Worker{{}, {}}, 5, 1000, cfg, config.ServerConfig{}, config.AuthConfig{APIKey: key}, config.ResolvedRateLimitConfig{}) if custom { cfg.AllowedOrigins[0] = "https://changed.test" cfg.AllowedMethods[0] = "DELETE" diff --git a/internal/api/router/auth_test.go b/internal/api/router/auth_test.go index 57bc5ac..6c7b196 100644 --- a/internal/api/router/auth_test.go +++ b/internal/api/router/auth_test.go @@ -13,7 +13,7 @@ import ( func TestAdminAuthBoundary(t *testing.T) { for _, key := range []string{"", "synthetic-test-key"} { - handler := New(nil, nil, nil, 5, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{APIKey: key}, config.ResolvedRateLimitConfig{}) + handler := New(nil, nil, nil, 5, 1000, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{APIKey: key}, config.ResolvedRateLimitConfig{}) for _, path := range []string{"/api/v1/admin", "/api/v1/admin/", "/api/v1/admin/config", "/api/v1/admin//config"} { for _, method := range []string{"GET", "POST", "PUT", "DELETE", "OPTIONS"} { for _, token := range []string{"", "wrong", "synthetic-test-key"} { diff --git a/internal/api/router/rate_limit_test.go b/internal/api/router/rate_limit_test.go index ac7f7ae..95ab42f 100644 --- a/internal/api/router/rate_limit_test.go +++ b/internal/api/router/rate_limit_test.go @@ -22,7 +22,7 @@ import ( ) func TestDefaultAPIRateLimit(t *testing.T) { - handler := New(nil, nil, nil, 5, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{}, config.Resolve(&config.Config{}).RateLimit) + handler := New(nil, nil, nil, 5, 1000, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{}, config.Resolve(&config.Config{}).RateLimit) for i := 0; i <= 300; i++ { request := httptest.NewRequest(http.MethodGet, "/api/v1/brokers", nil) request.RemoteAddr = "198.51.100.1:1234" @@ -56,7 +56,7 @@ func TestAPIRateLimitContractAndLogging(t *testing.T) { slog.SetDefault(slog.New(slog.NewJSONHandler(&logs, nil))) t.Cleanup(func() { slog.SetDefault(previous); log.SetOutput(writer); log.SetFlags(flags) }) proxy := config.ServerConfig{TrustedProxies: []netip.Prefix{netip.MustParsePrefix("127.0.0.1/32")}} - handler := New(nil, nil, nil, 5, config.CORSConfig{}, proxy, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 2, Burst: 10}) + handler := New(nil, nil, nil, 5, 1000, config.CORSConfig{}, proxy, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 2, Burst: 10}) for _, path := range []string{"/api/v1/brokers", "/api/v1/packets?limit=0"} { response := rateRequest(handler, path, "127.0.0.1:1234", "198.51.100.25") if response.Code != http.StatusOK && response.Code != http.StatusBadRequest { @@ -111,7 +111,7 @@ func TestAPIRateLimitClientIdentity(t *testing.T) { if tc.trusted { proxy.TrustedProxies = []netip.Prefix{netip.MustParsePrefix("127.0.0.1/32")} } - handler := New(nil, nil, nil, 5, config.CORSConfig{}, proxy, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 1, Burst: 1}) + handler := New(nil, nil, nil, 5, 1000, config.CORSConfig{}, proxy, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 1, Burst: 1}) if response := rateRequest(handler, "/api/v1/brokers", tc.firstPeer, tc.firstHeader); response.Code != http.StatusOK { t.Fatalf("first client: %d", response.Code) } @@ -128,7 +128,7 @@ func TestAPIRateLimitClientIdentity(t *testing.T) { func TestAPIRateLimitWindowsAndExclusions(t *testing.T) { synctest.Test(t, func(t *testing.T) { - handler := New(nil, nil, nil, 5, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 4, Burst: 2}) + handler := New(nil, nil, nil, 5, 1000, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 4, Burst: 2}) request := func() *httptest.ResponseRecorder { return rateRequest(handler, "/api/v1/brokers", "198.51.100.1:1", "") } @@ -164,7 +164,7 @@ func TestAPIRateLimitWindowsAndExclusions(t *testing.T) { t.Fatal("minute budget did not recover") } }) - handler := New(nil, nil, nil, 5, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: false, RequestsPerMinute: 1, Burst: 1}) + handler := New(nil, nil, nil, 5, 1000, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: false, RequestsPerMinute: 1, Burst: 1}) for range 5 { if response := rateRequest(handler, "/api/v1/brokers", "198.51.100.1:1", ""); response.Code != http.StatusOK || response.Header().Get("Retry-After") != "" { t.Fatal("disabled limiter still applied") @@ -173,7 +173,7 @@ func TestAPIRateLimitWindowsAndExclusions(t *testing.T) { } func TestAPIRateLimitConcurrentRequests(t *testing.T) { - handler := New(nil, nil, nil, 5, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 5, Burst: 5}) + handler := New(nil, nil, nil, 5, 1000, config.CORSConfig{}, config.ServerConfig{}, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 5, Burst: 5}) var accepted, rejected atomic.Int32 var group sync.WaitGroup for range 20 { diff --git a/internal/api/router/router.go b/internal/api/router/router.go index 2e91ab8..34db94a 100644 --- a/internal/api/router/router.go +++ b/internal/api/router/router.go @@ -39,7 +39,7 @@ import ( // /stats → stats subrouter // // Admin endpoints require a configured bearer key. -func New(h *hub.Hub, reader api.Reader, workers []*ingest.Worker, maxConnsPerIP int, corsCfg config.CORSConfig, serverCfg config.ServerConfig, authCfg config.AuthConfig, rateLimitCfg config.ResolvedRateLimitConfig) http.Handler { +func New(h *hub.Hub, reader api.Reader, workers []*ingest.Worker, maxConnsPerIP, maxConnectsPerMinute int, corsCfg config.CORSConfig, serverCfg config.ServerConfig, authCfg config.AuthConfig, rateLimitCfg config.ResolvedRateLimitConfig) http.Handler { r := chi.NewRouter() // ── CORS ───────────────────────────────────────────────────────────────── @@ -99,7 +99,7 @@ func New(h *hub.Hub, reader api.Reader, workers []*ingest.Worker, maxConnsPerIP )) // ── WebSocket ──────────────────────────────────────────────────────────── - r.Get("/ws", ws.Handler(h, reader, maxConnsPerIP)) + r.Get("/ws", ws.Handler(h, reader, maxConnsPerIP, maxConnectsPerMinute)) // ── Public REST API (v1) ───────────────────────────────────────────────── r.Route("/api/v1", func(r chi.Router) { diff --git a/internal/api/router/router_test.go b/internal/api/router/router_test.go index d3565bc..88927e2 100644 --- a/internal/api/router/router_test.go +++ b/internal/api/router/router_test.go @@ -19,7 +19,7 @@ import ( ) func TestWebSocketLimitIgnoresUntrustedHeaders(t *testing.T) { - checkWebSocketLimit(t, config.ServerConfig{}, "198.51.100.2", http.StatusTooManyRequests) + checkWebSocketLimit(t, config.ServerConfig{}, "198.51.100.2", true) } func TestWebSocketLimitUsesTrustedClientIP(t *testing.T) { @@ -27,16 +27,16 @@ func TestWebSocketLimitUsesTrustedClientIP(t *testing.T) { netip.MustParsePrefix("127.0.0.1/8"), netip.MustParsePrefix("::1/128"), }} t.Run("different clients", func(t *testing.T) { - checkWebSocketLimit(t, cfg, "198.51.100.2", http.StatusSwitchingProtocols) + checkWebSocketLimit(t, cfg, "198.51.100.2", false) }) t.Run("same client", func(t *testing.T) { - checkWebSocketLimit(t, cfg, "198.51.100.1", http.StatusTooManyRequests) + checkWebSocketLimit(t, cfg, "198.51.100.1", true) }) } -func checkWebSocketLimit(t *testing.T, cfg config.ServerConfig, secondIP string, wantStatus int) { +func checkWebSocketLimit(t *testing.T, cfg config.ServerConfig, secondIP string, wantShed bool) { t.Helper() - server := httptest.NewServer(New(hub.New(), nil, nil, 1, config.CORSConfig{}, cfg, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 1, Burst: 1})) + server := httptest.NewServer(New(hub.New(), nil, nil, 1, 10, config.CORSConfig{}, cfg, config.AuthConfig{}, config.ResolvedRateLimitConfig{Enabled: true, RequestsPerMinute: 1, Burst: 1})) defer server.Close() // Exhaust this client's REST budget before checking its independent WS cap. for _, want := range []int{http.StatusOK, http.StatusTooManyRequests} { @@ -71,10 +71,22 @@ func checkWebSocketLimit(t *testing.T, cfg config.ServerConfig, secondIP string, t.Fatalf("read hello: %v", err) } second, response, err := dial(secondIP) - if second != nil { - second.CloseNow() + if err != nil || response == nil || response.StatusCode != http.StatusSwitchingProtocols { + t.Fatalf("second handshake: response=%v err=%v", response, err) } - if response == nil || response.StatusCode != wantStatus || (err == nil) != (wantStatus == http.StatusSwitchingProtocols) { - t.Fatalf("second connection: response=%v err=%v, want status %d", response, err, wantStatus) + defer second.CloseNow() + _, _, err = second.Read(ctx) + if wantShed && websocket.CloseStatus(err) != websocket.StatusTryAgainLater { + t.Fatalf("expected close 1013, got %v", err) + } + if !wantShed && err != nil { + t.Fatalf("independent client did not receive hello: %v", err) + } + if err := first.Write(ctx, websocket.MessageText, []byte(`{"v":1,"type":"ping","id":"still-active"}`)); err != nil { + t.Fatal(err) + } + _, pong, err := first.Read(ctx) + if err != nil || !strings.Contains(string(pong), `"type":"pong"`) { + t.Fatalf("established client stopped responding: %s %v", pong, err) } } diff --git a/internal/api/router/ws_attempt_test.go b/internal/api/router/ws_attempt_test.go new file mode 100644 index 0000000..b0f7398 --- /dev/null +++ b/internal/api/router/ws_attempt_test.go @@ -0,0 +1,53 @@ +// Copyright 2026 Beacon Contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +package router + +import ( + "net/http" + "net/http/httptest" + "net/netip" + "testing" + + "github.com/MeshCore-Beacon/beacon-server/internal/config" +) + +func TestWebSocketAttemptClientIdentity(t *testing.T) { + for _, tc := range []struct { + name, firstPeer, secondPeer, firstHeader, secondHeader string + trusted, shared bool + }{ + {"untrusted rotation", "198.51.100.1:1", "198.51.100.1:2", "192.0.2.1", "192.0.2.2", false, true}, + {"independent peers", "198.51.100.1:1", "198.51.100.2:1", "192.0.2.1", "192.0.2.1", false, false}, + {"trusted independent clients", "127.0.0.1:1", "127.0.0.1:2", "198.51.100.1", "198.51.100.2", true, false}, + {"trusted shared client", "127.0.0.1:1", "127.0.0.1:2", "198.51.100.1", "198.51.100.1", true, true}, + {"IPv6 shared prefix", "[2001:db8:1::1]:1", "[2001:db8:1::2]:2", "", "", false, true}, + {"IPv6 independent prefixes", "[2001:db8:1::1]:1", "[2001:db8:2::1]:2", "", "", false, false}, + } { + t.Run(tc.name, func(t *testing.T) { + cfg := config.ServerConfig{} + if tc.trusted { + cfg.TrustedProxies = []netip.Prefix{netip.MustParsePrefix("127.0.0.1/32")} + } + handler := New(nil, nil, nil, 5, 1, config.CORSConfig{}, cfg, config.AuthConfig{}, config.ResolvedRateLimitConfig{}) + attempt := func(peer, header string) int { + request := httptest.NewRequest(http.MethodGet, "/ws", nil) + request.RemoteAddr = peer + request.Header.Set("X-Real-IP", header) + request.Header.Set("X-Forwarded-For", header) + request.Header.Set("True-Client-IP", header) + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + return response.Code + } + first := attempt(tc.firstPeer, tc.firstHeader) + if first < 400 || first == http.StatusTooManyRequests { + t.Fatalf("first failed handshake: %d", first) + } + second := attempt(tc.secondPeer, tc.secondHeader) + if (second == http.StatusTooManyRequests) != tc.shared || second < 400 { + t.Fatalf("second handshake status %d; shared budget=%v", second, tc.shared) + } + }) + } +} diff --git a/internal/config/config.go b/internal/config/config.go index 9455a2e..9147b87 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -85,6 +85,7 @@ type ResolvedConfig struct { RouteGrace time.Duration RouteMinObservations int MaxConnsPerIP int + MaxConnectsPerMinute int ViewRefreshInterval time.Duration ReconfirmInterval time.Duration CleanupInterval time.Duration @@ -231,6 +232,9 @@ type WebSocketConfig struct { // MaxConnectionsPerIP is the maximum number of concurrent WebSocket // connections allowed from a single IP address. Defaults to 5 if not set. MaxConnectionsPerIP int `yaml:"max_connections_per_ip"` + // 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"` } // PacketsConfig controls packet retention behaviour. @@ -382,6 +386,9 @@ func Load(path string) (*Config, error) { if cfg.RateLimit.RequestsPerMinute < 0 || cfg.RateLimit.Burst < 0 { return nil, fmt.Errorf("ratelimit.requests_per_minute and ratelimit.burst must be positive or zero for defaults") } + if cfg.WebSocket.MaxConnectsPerMinute < 0 { + return nil, fmt.Errorf("websocket.max_connects_per_minute must be positive or zero for the default") + } configDir := filepath.Dir(path) for iata, details := range cfg.IATAs { if details.BorderFile != "" && !filepath.IsAbs(details.BorderFile) { @@ -407,6 +414,7 @@ func Resolve(cfg *Config) ResolvedConfig { RouteGrace: cfg.Routes.Grace.Duration, RouteMinObservations: cfg.Routes.MinObservations, MaxConnsPerIP: cfg.WebSocket.MaxConnectionsPerIP, + MaxConnectsPerMinute: cfg.WebSocket.MaxConnectsPerMinute, ViewRefreshInterval: cfg.Background.ViewRefresh.Duration, ReconfirmInterval: cfg.Background.Reconfirm.Duration, CleanupInterval: cfg.Background.Cleanup.Duration, @@ -446,6 +454,9 @@ func Resolve(cfg *Config) ResolvedConfig { if r.MaxConnsPerIP == 0 { r.MaxConnsPerIP = 5 } + if r.MaxConnectsPerMinute == 0 { + r.MaxConnectsPerMinute = 10 + } if r.ViewRefreshInterval == 0 { r.ViewRefreshInterval = time.Hour } @@ -477,9 +488,9 @@ func Resolve(cfg *Config) ResolvedConfig { func (r ResolvedConfig) String() string { return fmt.Sprintf( - "telemetryResolution=%s telemetryRetention=%s packetRetention=%s routeRetention=%s routeGrace=%s routeMinObs=%d maxConnsPerIP=%d viewRefresh=%s reconfirm=%s cleanup=%s presenceFlush=%s presencePacketTTL=%s clockDriftThreshold=%s nodeStaleThreshold=%s nodeDeleteAfter=%s observerDeleteAfter=%s rateLimitEnabled=%t requestsPerMinute=%d burst=%d", + "telemetryResolution=%s telemetryRetention=%s packetRetention=%s routeRetention=%s routeGrace=%s routeMinObs=%d maxConnsPerIP=%d maxConnectsPerMinute=%d viewRefresh=%s reconfirm=%s cleanup=%s presenceFlush=%s presencePacketTTL=%s clockDriftThreshold=%s nodeStaleThreshold=%s nodeDeleteAfter=%s observerDeleteAfter=%s rateLimitEnabled=%t requestsPerMinute=%d burst=%d", r.TelemetryResolution, r.TelemetryRetention, r.PacketRetention, r.RouteRetention, r.RouteGrace, r.RouteMinObservations, - r.MaxConnsPerIP, r.ViewRefreshInterval, r.ReconfirmInterval, r.CleanupInterval, + r.MaxConnsPerIP, r.MaxConnectsPerMinute, r.ViewRefreshInterval, r.ReconfirmInterval, r.CleanupInterval, r.PresenceFlushInterval, r.PresencePacketTTL, r.ClockDriftThreshold, r.NodeStaleThreshold, r.NodeDeleteAfter, r.ObserverDeleteAfter, r.RateLimit.Enabled, r.RateLimit.RequestsPerMinute, r.RateLimit.Burst, diff --git a/internal/config/ws_connect_test.go b/internal/config/ws_connect_test.go new file mode 100644 index 0000000..5914707 --- /dev/null +++ b/internal/config/ws_connect_test.go @@ -0,0 +1,37 @@ +// Copyright 2026 Beacon Contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +package config + +import ( + "os" + "path/filepath" + "testing" +) + +func TestWebSocketConnectConfig(t *testing.T) { + for _, tc := range []struct { + name, yaml string + want int + wantError bool + }{ + {"omitted", "{}", 10, false}, + {"zero default", "websocket: {max_connects_per_minute: 0}", 10, false}, + {"custom", "websocket: {max_connects_per_minute: 30}", 30, false}, + {"negative", "websocket: {max_connects_per_minute: -1}", 0, true}, + } { + t.Run(tc.name, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "config.yaml") + if err := os.WriteFile(path, []byte(tc.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 && (Resolve(cfg).MaxConnectsPerMinute != tc.want || Resolve(cfg).MaxConnsPerIP != 5) { + t.Fatalf("unexpected resolved WebSocket limits: %+v", Resolve(cfg)) + } + }) + } +} diff --git a/internal/ws/connect_limit_test.go b/internal/ws/connect_limit_test.go new file mode 100644 index 0000000..2e55a8b --- /dev/null +++ b/internal/ws/connect_limit_test.go @@ -0,0 +1,174 @@ +// Copyright 2026 Beacon Contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +package ws + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/MeshCore-Beacon/beacon-server/internal/hub" + "github.com/coder/websocket" +) + +func connectTestServer(t *testing.T, maxConns, maxConnects int) (*httptest.Server, <-chan struct{}) { + t.Helper() + handler := Handler(hub.New(), nil, maxConns, maxConnects) + done := make(chan struct{}, 64) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + handler(w, r) + done <- struct{}{} + })) + t.Cleanup(server.Close) + return server, done +} + +func TestFailedHandshakeDoesNotUseConnectionSlot(t *testing.T) { + server, done := connectTestServer(t, 1, 10) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + response, err := server.Client().Get(server.URL) + if err != nil { + t.Fatal(err) + } + response.Body.Close() + if response.StatusCode < 400 || response.StatusCode == http.StatusTooManyRequests { + t.Fatalf("expected a failed handshake: %d", response.StatusCode) + } + waitConnectionEnd(t, ctx, done) + conn, _, err := websocket.Dial(ctx, "ws"+strings.TrimPrefix(server.URL, "http"), nil) + if err != nil { + t.Fatal(err) + } + defer conn.CloseNow() + if _, _, err := conn.Read(ctx); err != nil { + t.Fatalf("failed handshake consumed a connection slot: %v", err) + } + conn.CloseNow() + waitConnectionEnd(t, ctx, done) +} + +func TestConnectAttemptBudgetAndRecovery(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + handler := Handler(nil, nil, 1, 2) + attempt := func() *httptest.ResponseRecorder { + request := httptest.NewRequest(http.MethodGet, "/ws", nil) + request.RemoteAddr = "198.51.100.1:1234" + response := httptest.NewRecorder() + handler(response, request) + return response + } + for range 2 { + if response := attempt(); response.Code < 400 || response.Code == http.StatusTooManyRequests { + t.Fatalf("expected failed handshake within budget: %d", response.Code) + } + } + if response := attempt(); response.Code != http.StatusTooManyRequests || response.Header().Get("Retry-After") != "60" { + t.Fatalf("failed handshakes bypassed attempt budget: %d %v", response.Code, response.Header()) + } + time.Sleep(2 * time.Minute) + if response := attempt(); response.Code < 400 || response.Code == http.StatusTooManyRequests { + t.Fatalf("attempt budget did not recover: %d", response.Code) + } + }) +} + +func TestConcurrentConnectsRespectBothLimits(t *testing.T) { + for _, tc := range []struct { + name string + conns, attempts int + wantOpen, wantShed int32 + wantRejected int32 + }{ + {"concurrent cap", 2, 20, 2, 6, 0}, + {"attempt cap", 20, 2, 2, 0, 6}, + } { + t.Run(tc.name, func(t *testing.T) { + server, done := connectTestServer(t, tc.conns, tc.attempts) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + var opened, shed, rejected atomic.Int32 + connections := make(chan *websocket.Conn, 8) + defer func() { + close(connections) + for conn := range connections { + conn.CloseNow() + } + for range 8 { + waitConnectionEnd(t, ctx, done) + } + }() + var group sync.WaitGroup + for range 8 { + group.Go(func() { + conn, response, err := websocket.Dial(ctx, "ws"+strings.TrimPrefix(server.URL, "http"), nil) + if err != nil { + if response != nil && response.StatusCode == http.StatusTooManyRequests { + rejected.Add(1) + } else { + t.Errorf("unexpected handshake failure: %v", err) + } + return + } + connections <- conn // Keep accepted connections open until every attempt completes. + if _, _, err := conn.Read(ctx); err == nil { + opened.Add(1) + } else if websocket.CloseStatus(err) == websocket.StatusTryAgainLater { + shed.Add(1) + } else { + t.Errorf("unexpected WebSocket read error: %v", err) + } + }) + } + group.Wait() + if opened.Load() != tc.wantOpen || shed.Load() != tc.wantShed || rejected.Load() != tc.wantRejected { + t.Fatalf("opened=%d shed=%d rejected=%d", opened.Load(), shed.Load(), rejected.Load()) + } + }) + } +} + +func waitConnectionEnd(t *testing.T, ctx context.Context, done <-chan struct{}) { + t.Helper() + select { + case <-done: + case <-ctx.Done(): + t.Fatal("connection handler did not release its slot") + } +} + +func TestConnectChurnLimited(t *testing.T) { + server, done := connectTestServer(t, 1, 10) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + url := "ws" + strings.TrimPrefix(server.URL, "http") + for range 10 { + conn, _, err := websocket.Dial(ctx, url, nil) + if err != nil { + t.Fatal(err) + } + if _, _, err := conn.Read(ctx); err != nil { + conn.CloseNow() + t.Fatalf("read hello: %v", err) + } + conn.CloseNow() + waitConnectionEnd(t, ctx, done) + } + conn, response, err := websocket.Dial(ctx, url, nil) + if conn != nil { + conn.CloseNow() + } + if err == nil || response == nil || response.StatusCode != http.StatusTooManyRequests { + t.Fatal("eleventh connection attempt was not rejected with HTTP 429") + } + if response.Header.Get("Retry-After") != "60" { + t.Fatal("missing upgrade retry guidance") + } +} diff --git a/internal/ws/handler.go b/internal/ws/handler.go index 1e4772e..cb632ea 100644 --- a/internal/ws/handler.go +++ b/internal/ws/handler.go @@ -39,6 +39,7 @@ import ( "time" "github.com/coder/websocket" + "github.com/go-chi/httprate" "github.com/google/uuid" "github.com/MeshCore-Beacon/beacon-server/internal/api" @@ -52,24 +53,34 @@ 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 int) http.HandlerFunc { +func Handler(h *hub.Hub, reader api.Reader, maxConnsPerIP, maxConnectsPerMinute int) http.HandlerFunc { limiter := newIPLimiter(maxConnsPerIP) + attempts := httprate.NewRateLimiter(maxConnectsPerMinute, time.Minute, + httprate.WithResponseHeaders(httprate.ResponseHeaders{RetryAfter: "Retry-After"})) return func(w http.ResponseWriter, r *http.Request) { ip := r.RemoteAddr if host, _, err := net.SplitHostPort(ip); err == nil { ip = host } - if !limiter.acquire(ip) { - slog.Warn(fmt.Sprintf("ws: connection limit reached for IP %s", ip), "component", "ws") - http.Error(w, "too many connections from this IP", http.StatusTooManyRequests) + // Count attempts before Accept, including failed handshakes. RemoteAddr + // was already resolved by TrustedProxyIP; never read forwarding headers here. + if attempts.RespondOnLimit(w, r, httprate.CanonicalizeIP(ip)) { return } - defer limiter.release(ip) conn, err := websocket.Accept(w, r, nil) if err != nil { slog.Warn("ws: failed to accept connection", "component", "ws", "error", err) return } + defer conn.CloseNow() + if !limiter.acquire(ip) { + // This attempt passed the rate budget. Give the accepted socket a + // browser-visible backoff signal before allocating a hub client. + slog.Warn("ws: connection limit reached", "component", "ws", "client_ip", ip) + _ = conn.Close(websocket.StatusTryAgainLater, "connection limit reached") + return + } + defer limiter.release(ip) connID := uuid.NewString() client := h.NewClient()