From 2bf52ba7ffc20da3b838f0d8d94700acf9a70faf Mon Sep 17 00:00:00 2001 From: clawbot <35+clawbot@noreply.example.org> Date: Tue, 29 Sep 2026 00:18:52 +0000 Subject: [PATCH] fix(backend): rate-limit and cap report ingest, drop wildcard CORS (closes #20) POST /api/v1/reports stays unauthenticated but is bounded. Each client address, as the trusted-proxy logic resolves it, may send REPORTS_PER_MINUTE reports a minute (default 60, all at once if it likes), using golang.org/x/time/rate; past that it gets 429 with Retry-After. Buckets that have refilled are dropped once a minute, so idle addresses do not pile up. reportbuf refuses a report that would take the report files past DATA_DIR_MAX_BYTES (default 1 GiB) with ErrFull, answered with 507; the count starts from the files already in DATA_DIR, and reports not yet written count at their uncompressed size. CORS adds nothing unless CORS_ALLOWED_ORIGINS lists origins. A limit that is not a positive number stops the server from starting. Model: opus-5-5 --- TODO.md | 9 + backend/README.md | 44 ++++- backend/go.mod | 1 + backend/go.sum | 2 + backend/internal/config/config.go | 68 +++++-- backend/internal/config/config_test.go | 47 +++++ backend/internal/handlers/report.go | 20 ++- backend/internal/handlers/report_test.go | 27 +++ backend/internal/middleware/export_test.go | 7 + backend/internal/middleware/middleware.go | 107 +++++++++-- .../middleware/middleware_internal_test.go | 46 +++++ .../internal/middleware/middleware_test.go | 166 ++++++++++++++++++ backend/internal/reportbuf/export_test.go | 7 + backend/internal/reportbuf/reportbuf.go | 88 +++++++++- backend/internal/reportbuf/reportbuf_test.go | 128 ++++++++++++++ backend/internal/server/routes.go | 5 +- backend/internal/server/routes_test.go | 53 +++++- 17 files changed, 770 insertions(+), 55 deletions(-) create mode 100644 backend/internal/config/config_test.go create mode 100644 backend/internal/middleware/middleware_internal_test.go create mode 100644 backend/internal/reportbuf/export_test.go diff --git a/TODO.md b/TODO.md index ecb3301..f529374 100644 --- a/TODO.md +++ b/TODO.md @@ -23,6 +23,15 @@ latest run passes. # Completed Steps +- 2026-09-29: bounded the report endpoint (issue #20): `POST /api/v1/reports` + still needs no credentials, but each client address, as resolved through + `TRUSTED_PROXIES`, may send `REPORTS_PER_MINUTE` (default 60) reports a minute + and past that gets 429 with `Retry-After`; the report files in `DATA_DIR`, + counted from start with those already there, may total at most + `DATA_DIR_MAX_BYTES` (default 1 GiB), past which reports get 507; and the + wildcard CORS is gone: no CORS headers unless `CORS_ALLOWED_ORIGINS` lists + origins. Deleting report files frees room only at the next start; pruning is + issue #54 - 2026-09-28: unified the gate (issue #16): the root `make check` covers the Go backend as well as the frontend, and the pre-commit hook with it; the backend moved onto scripts-to-rule-them-all (`backend/script/*`, `backend/Makefile` as diff --git a/backend/README.md b/backend/README.md index 1230c12..26147fb 100644 --- a/backend/README.md +++ b/backend/README.md @@ -74,12 +74,15 @@ Internal packages in `internal/` follow standard Go project layout: ### Configuration -| Variable | Default | Description | -| ----------------- | -------------------- | -------------------------------------------------------------------------------------------------------- | -| `PORT` | `8080` | HTTP listen port | -| `DATA_DIR` | `./data/reports` | Directory for compressed reports | -| `DEBUG` | `false` | Enable debug logging | -| `TRUSTED_PROXIES` | loopback + RFC1918 | Comma-separated CIDRs whose `X-Forwarded-For` / `X-Real-IP` headers are trusted for client IP resolution | +| Variable | Default | Description | +| ---------------------- | -------------------- | -------------------------------------------------------------------------------------------------------- | +| `PORT` | `8080` | HTTP listen port | +| `DATA_DIR` | `./data/reports` | Directory for compressed reports | +| `DATA_DIR_MAX_BYTES` | `1073741824` (1 GiB) | Most the report files in `DATA_DIR` may total; see [Report limits](#report-limits) | +| `DEBUG` | `false` | Enable debug logging | +| `TRUSTED_PROXIES` | loopback + RFC1918 | Comma-separated CIDRs whose `X-Forwarded-For` / `X-Real-IP` headers are trusted for client IP resolution | +| `REPORTS_PER_MINUTE` | `60` | Reports each client address may send a minute; see [Report limits](#report-limits) | +| `CORS_ALLOWED_ORIGINS` | empty | Comma-separated origins whose pages may call the API; see [CORS](#cors) | `TRUSTED_PROXIES` defaults to `127.0.0.1/32,::1/128,10.0.0.0/8,172.16.0.0/12,192.168.0.0/16`. The loopback entries cover the reverse proxy that shares the container; the @@ -92,6 +95,35 @@ Reports are written as `reports-.jsonl.zst` files in `DATA_DIR`. Each file contains one JSON object per line, compressed with zstd. Files are created with `O_EXCL` to prevent overwrites. +### Report limits + +`POST /api/v1/reports` takes reports from anyone who can reach it, without +credentials, so it is bounded instead. Both refusals below answer with the same +`{"status":"error"}` body as any other error. + +- **Rate limit.** Each client address, resolved through `TRUSTED_PROXIES`, may + send `REPORTS_PER_MINUTE` reports a minute, all at once if it likes. Past that + it gets 429 with a `Retry-After` header until its allowance refills, at one + report every 60 / `REPORTS_PER_MINUTE` seconds. The page sends one report a + minute from each open tab, so the default of 60 refuses nothing from up to 60 + tabs behind one address, such as a household or an office sharing it, even + when all their reports arrive together. +- **Size cap.** The report files in `DATA_DIR` may total at most + `DATA_DIR_MAX_BYTES`, counting the files already there at start. Reports + waiting in memory count at their uncompressed size until they are written, so + a report that would take the total past the cap is refused with 507, and + nothing of it is stored. Deleting report files frees room only at the next + start, when the files are counted again. The default of 1 GiB is small enough + for any host; set it to the space you can give `DATA_DIR`. + +### CORS + +The page calls the API from the origin it is served from, so by default the +server sends no CORS headers, and browsers let no other origin's pages call it. +To serve the page from elsewhere, list that origin in `CORS_ALLOWED_ORIGINS` +(for example `https://netwatch.example.com`); pages from a listed origin may +`GET` and `POST` with a `Content-Type` header. + ## TODO - Add integration test that POSTs a report and verifies the compressed output diff --git a/backend/go.mod b/backend/go.mod index 89e3ef4..ac4ad4f 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -9,6 +9,7 @@ require ( github.com/klauspost/compress v1.18.4 github.com/spf13/viper v1.21.0 go.uber.org/fx v1.24.0 + golang.org/x/time v0.15.0 ) require ( diff --git a/backend/go.sum b/backend/go.sum index de6ddf2..983942d 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -58,6 +58,8 @@ golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU= golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/text v0.28.0 h1:rhazDwis8INMIwQ4tpjLDzUhx6RlXqZNPEM0huQojng= golang.org/x/text v0.28.0/go.mod h1:U8nCwOR8jO/marOQ0QbDiOngZVEBB7MAiitBuMjXiNU= +golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= +golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15 h1:YR8cESwS4TdDjEe65xsg0ogRM/Nc3DYOhEAlW+xobZo= gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 02a6db9..499d117 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -4,6 +4,7 @@ package config import ( "errors" + "fmt" "log/slog" "strings" @@ -23,6 +24,15 @@ import ( const defaultTrustedProxies = "127.0.0.1/32,::1/128," + "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16" +// Default limits on stored reports; backend/README.md gives the +// reasons for these values. +const ( + defaultReportsPerMinute = 60 + defaultDataDirMaxBytes = 1 << 30 // 1 GiB +) + +var errNotPositive = errors.New("must be a positive whole number") + // Params defines the dependencies for Config. type Params struct { fx.In @@ -33,15 +43,18 @@ type Params struct { // Config holds the resolved application configuration. type Config struct { - DataDir string - Debug bool - MetricsPassword string - MetricsUsername string - Port int - SentryDSN string - TrustedProxies []string - log *slog.Logger - params *Params + CORSAllowedOrigins []string + DataDir string + DataDirMaxBytes int64 + Debug bool + MetricsPassword string + MetricsUsername string + Port int + ReportsPerMinute int + SentryDSN string + TrustedProxies []string + log *slog.Logger + params *Params } // New loads configuration from env, .env files, and config @@ -60,9 +73,13 @@ func New( viper.AutomaticEnv() + // An empty CORS_ALLOWED_ORIGINS allows no other origin. + viper.SetDefault("CORS_ALLOWED_ORIGINS", "") viper.SetDefault("DATA_DIR", "./data/reports") + viper.SetDefault("DATA_DIR_MAX_BYTES", defaultDataDirMaxBytes) viper.SetDefault("DEBUG", "false") viper.SetDefault("PORT", "8080") + viper.SetDefault("REPORTS_PER_MINUTE", defaultReportsPerMinute) viper.SetDefault("SENTRY_DSN", "") viper.SetDefault("METRICS_USERNAME", "") viper.SetDefault("METRICS_PASSWORD", "") @@ -78,15 +95,30 @@ func New( } s := &Config{ - DataDir: viper.GetString("DATA_DIR"), - Debug: viper.GetBool("DEBUG"), - MetricsPassword: viper.GetString("METRICS_PASSWORD"), - MetricsUsername: viper.GetString("METRICS_USERNAME"), - Port: viper.GetInt("PORT"), - SentryDSN: viper.GetString("SENTRY_DSN"), - TrustedProxies: splitList(viper.GetString("TRUSTED_PROXIES")), - log: log, - params: ¶ms, + CORSAllowedOrigins: splitList(viper.GetString("CORS_ALLOWED_ORIGINS")), + DataDir: viper.GetString("DATA_DIR"), + DataDirMaxBytes: viper.GetInt64("DATA_DIR_MAX_BYTES"), + Debug: viper.GetBool("DEBUG"), + MetricsPassword: viper.GetString("METRICS_PASSWORD"), + MetricsUsername: viper.GetString("METRICS_USERNAME"), + Port: viper.GetInt("PORT"), + ReportsPerMinute: viper.GetInt("REPORTS_PER_MINUTE"), + SentryDSN: viper.GetString("SENTRY_DSN"), + TrustedProxies: splitList(viper.GetString("TRUSTED_PROXIES")), + log: log, + params: ¶ms, + } + + // viper reads a value that is not a number as 0, so this also + // catches a mistyped setting. + if s.ReportsPerMinute <= 0 { + return nil, fmt.Errorf("REPORTS_PER_MINUTE %q: %w", + viper.GetString("REPORTS_PER_MINUTE"), errNotPositive) + } + + if s.DataDirMaxBytes <= 0 { + return nil, fmt.Errorf("DATA_DIR_MAX_BYTES %q: %w", + viper.GetString("DATA_DIR_MAX_BYTES"), errNotPositive) } if s.Debug { diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go new file mode 100644 index 0000000..38211e0 --- /dev/null +++ b/backend/internal/config/config_test.go @@ -0,0 +1,47 @@ +package config_test + +import ( + "strings" + "testing" + + "sneak.berlin/go/netwatch/internal/config" + "sneak.berlin/go/netwatch/internal/globals" + "sneak.berlin/go/netwatch/internal/logger" + + "go.uber.org/fx" +) + +// requireConfigError builds the config as main does and fails the +// test unless that fails with an error naming setting. It uses +// fx.New, because fxtest.New fails the test itself on an error. +func requireConfigError(t *testing.T, setting string) { + t.Helper() + + app := fx.New( + fx.NopLogger, + fx.Provide(globals.New, logger.New, config.New), + fx.Invoke(func(*config.Config) {}), + ) + + err := app.Err() + if err == nil || !strings.Contains(err.Error(), setting) { + t.Fatalf("config error = %v, want one naming %s", err, setting) + } +} + +// TestReportsPerMinuteMustBePositive: unchecked, zero would panic +// when the routes are built, and a negative rate would lift the +// limit. +func TestReportsPerMinuteMustBePositive(t *testing.T) { + t.Setenv("REPORTS_PER_MINUTE", "0") + + requireConfigError(t, "REPORTS_PER_MINUTE") +} + +// TestDataDirMaxBytesMustBeANumber: viper reads a value that is not +// a number, such as "1GB", as 0, which would refuse every report. +func TestDataDirMaxBytesMustBeANumber(t *testing.T) { + t.Setenv("DATA_DIR_MAX_BYTES", "1GB") + + requireConfigError(t, "DATA_DIR_MAX_BYTES") +} diff --git a/backend/internal/handlers/report.go b/backend/internal/handlers/report.go index 81954a2..859a3a7 100644 --- a/backend/internal/handlers/report.go +++ b/backend/internal/handlers/report.go @@ -4,6 +4,8 @@ import ( "encoding/json" "errors" "net/http" + + "sneak.berlin/go/netwatch/internal/reportbuf" ) // maxLoggedFieldBytes bounds untrusted text (string fields, @@ -55,10 +57,9 @@ func (s *Handlers) HandleReport() http.HandlerFunc { err = s.buf.Append(rpt) if err != nil { - s.log.Error("failed to buffer report", "error", err) s.respondJSON(w, r, &response{Status: "error"}, - http.StatusInternalServerError, + s.appendErrorStatus(err), ) return @@ -88,6 +89,21 @@ func (s *Handlers) decodeErrorStatus(err error) int { return http.StatusBadRequest } +// appendErrorStatus logs a failure to store a report and returns +// the status to send: 507 when the report files are at their size +// cap, otherwise 500. +func (s *Handlers) appendErrorStatus(err error) int { + if errors.Is(err, reportbuf.ErrFull) { + s.log.Warn("report refused: report files at their size cap") + + return http.StatusInsufficientStorage + } + + s.log.Error("failed to buffer report", "error", err) + + return http.StatusInternalServerError +} + // logReportReceived logs an accepted report. Untrusted fields are // bounded (client_id, timestamp) or reduced to a length // (geo_bytes) so the raw attacker-controlled body never reaches diff --git a/backend/internal/handlers/report_test.go b/backend/internal/handlers/report_test.go index 68e9c71..996ce5c 100644 --- a/backend/internal/handlers/report_test.go +++ b/backend/internal/handlers/report_test.go @@ -13,6 +13,7 @@ import ( "sneak.berlin/go/netwatch/internal/handlers" "sneak.berlin/go/netwatch/internal/middleware" + "sneak.berlin/go/netwatch/internal/reportbuf" ) var errStorageFailed = errors.New("storage failed") @@ -66,6 +67,32 @@ func TestHandleReportStorageFailureIsNon2xx(t *testing.T) { } } +// TestHandleReportFullIs507 checks the answer when the report files +// are at their size cap: 507 and the usual error body, which tells +// the client nothing more. +func TestHandleReportFullIs507(t *testing.T) { + t.Parallel() + + h := newTestHandlers(stubAppender{err: reportbuf.ErrFull}, io.Discard) + + rec := httptest.NewRecorder() + req := httptest.NewRequestWithContext(t.Context(), + http.MethodPost, "/api/v1/reports", + strings.NewReader(`{"clientId":"c1","hosts":[]}`), + ) + + h.HandleReport().ServeHTTP(rec, req) + + if rec.Code != http.StatusInsufficientStorage { + t.Fatalf("status = %d, want %d", + rec.Code, http.StatusInsufficientStorage) + } + + if got := rec.Body.String(); got != "{\"status\":\"error\"}\n" { + t.Errorf("body = %q, want %q", got, "{\"status\":\"error\"}\n") + } +} + func TestHandleReportMalformedJSONIs400(t *testing.T) { t.Parallel() diff --git a/backend/internal/middleware/export_test.go b/backend/internal/middleware/export_test.go index 03722f4..22e1e7e 100644 --- a/backend/internal/middleware/export_test.go +++ b/backend/internal/middleware/export_test.go @@ -15,6 +15,13 @@ func NewWithLogger(log *slog.Logger) *Middleware { return &Middleware{log: log} } +// NewWithTrustedProxies builds a Middleware that honours forwarded +// headers from the given networks, for tests of the client address +// paths without the fx graph. +func NewWithTrustedProxies(trusted []netip.Prefix) *Middleware { + return &Middleware{trustedProxies: trusted} +} + func ClientIP( remoteAddr string, header http.Header, diff --git a/backend/internal/middleware/middleware.go b/backend/internal/middleware/middleware.go index b74092f..15ee821 100644 --- a/backend/internal/middleware/middleware.go +++ b/backend/internal/middleware/middleware.go @@ -7,11 +7,14 @@ import ( "fmt" "io" "log/slog" + "math" "net" "net/http" "net/netip" "runtime/debug" + "strconv" "strings" + "sync" "time" "sneak.berlin/go/netwatch/internal/config" @@ -21,6 +24,7 @@ import ( "github.com/go-chi/chi/v5/middleware" "github.com/go-chi/cors" "go.uber.org/fx" + "golang.org/x/time/rate" ) const corsMaxAgeSec = 300 @@ -320,21 +324,98 @@ func (s *Middleware) Recoverer() func(http.Handler) http.Handler { } } -// CORS returns middleware that adds permissive CORS headers. -func (s *Middleware) CORS() func(http.Handler) http.Handler { +// CORS returns middleware that lets pages served from the given +// origins call the API. With no origins it adds no CORS headers at +// all, so only same-origin pages can use the API. That case must not +// reach cors.Handler, which treats an empty origin list as "allow +// every origin". +func (s *Middleware) CORS( + origins []string, +) func(http.Handler) http.Handler { + if len(origins) == 0 { + return func(next http.Handler) http.Handler { return next } + } + return cors.Handler(cors.Options{ - AllowedOrigins: []string{"*"}, - AllowedMethods: []string{ - "GET", "POST", "PUT", "DELETE", "OPTIONS", - }, - AllowedHeaders: []string{ - "Accept", - "Authorization", - "Content-Type", - "X-CSRF-Token", - }, - ExposedHeaders: []string{"Link"}, + AllowedOrigins: origins, + AllowedMethods: []string{http.MethodGet, http.MethodPost}, + AllowedHeaders: []string{"Content-Type"}, AllowCredentials: false, MaxAge: corsMaxAgeSec, }) } + +// RateLimit returns middleware that allows each client address +// perMinute requests a minute, all at once if it likes, and answers +// the rest with 429 and a Retry-After header. The address is the one +// clientIP resolves, so clients behind the reverse proxy are limited +// one by one, not together as the proxy. +func (s *Middleware) RateLimit( + perMinute int, +) func(http.Handler) http.Handler { + // One request's allowance comes back every interval, so a + // refused client can always retry after it. + interval := time.Minute / time.Duration(perMinute) + retryAfter := strconv.Itoa(int(math.Ceil(interval.Seconds()))) + + limiter := &addressLimiter{ + burst: perMinute, + byAddr: make(map[string]*rate.Limiter), + limit: rate.Every(interval), + } + + return func(next http.Handler) http.Handler { + return http.HandlerFunc( + func(w http.ResponseWriter, r *http.Request) { + addr := clientIP(r.RemoteAddr, r.Header, s.trustedProxies) + + if !limiter.allow(addr, time.Now()) { + w.Header().Set("Retry-After", retryAfter) + writeJSONError(w, http.StatusTooManyRequests) + + return + } + + next.ServeHTTP(w, r) + }, + ) + } +} + +// addressLimiter holds a token bucket for each client address that +// has made a request recently. +type addressLimiter struct { + burst int + byAddr map[string]*rate.Limiter + lastSweep time.Time + limit rate.Limit + mu sync.Mutex +} + +// allow takes a token from addr's bucket, which starts full, and +// reports whether there was one. At most once a minute it first drops +// every bucket that has filled up again: a full bucket behaves exactly +// like the new one that would replace it, so this changes no answer, +// and the map holds only the addresses heard from recently. +func (a *addressLimiter) allow(addr string, now time.Time) bool { + a.mu.Lock() + defer a.mu.Unlock() + + if now.Sub(a.lastSweep) >= time.Minute { + for key, bucket := range a.byAddr { + if bucket.TokensAt(now) >= float64(a.burst) { + delete(a.byAddr, key) + } + } + + a.lastSweep = now + } + + bucket, ok := a.byAddr[addr] + if !ok { + bucket = rate.NewLimiter(a.limit, a.burst) + a.byAddr[addr] = bucket + } + + return bucket.AllowN(now, 1) +} diff --git a/backend/internal/middleware/middleware_internal_test.go b/backend/internal/middleware/middleware_internal_test.go new file mode 100644 index 0000000..aceac80 --- /dev/null +++ b/backend/internal/middleware/middleware_internal_test.go @@ -0,0 +1,46 @@ +package middleware + +import ( + "testing" + "time" + + "golang.org/x/time/rate" +) + +// TestAddressLimiterDropsOnlyFullBuckets checks the sweep in allow: +// a minute after the last one, it drops an address whose bucket has +// filled up again, and keeps one still short of tokens, whose limit +// would otherwise start over. +func TestAddressLimiterDropsOnlyFullBuckets(t *testing.T) { + t.Parallel() + + const ( + refilled = "198.51.100.1" + drained = "198.51.100.2" + ) + + // Two a minute: one token back every 30 seconds. + limiter := &addressLimiter{ + burst: 2, + byAddr: make(map[string]*rate.Limiter), + limit: rate.Every(30 * time.Second), + } + + start := time.Now() + + // The first call sweeps the empty map and takes one of two tokens. + limiter.allow(refilled, start) + limiter.allow(drained, start.Add(59*time.Second)) + limiter.allow(drained, start.Add(59*time.Second)) + + // A minute after the first sweep, this call sweeps again. + limiter.allow("198.51.100.3", start.Add(time.Minute)) + + if _, ok := limiter.byAddr[refilled]; ok { + t.Error("address with a full bucket was kept") + } + + if _, ok := limiter.byAddr[drained]; !ok { + t.Error("address short of tokens was dropped") + } +} diff --git a/backend/internal/middleware/middleware_test.go b/backend/internal/middleware/middleware_test.go index 8016256..77eb3fb 100644 --- a/backend/internal/middleware/middleware_test.go +++ b/backend/internal/middleware/middleware_test.go @@ -10,6 +10,8 @@ import ( "net/netip" "strings" "testing" + "testing/synctest" + "time" "sneak.berlin/go/netwatch/internal/middleware" ) @@ -300,3 +302,167 @@ func TestRecovererRepanicsOnAbortHandler(t *testing.T) { t.Errorf("abort was logged: %q", logbuf.String()) } } + +// okHandler stands in for the route a middleware guards. +func okHandler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }) +} + +// TestRateLimitRefusesPastAllowanceUntilRetryAfter checks one client +// address: it may use its whole allowance at once, the next request +// is refused with 429, and once Retry-After has passed it may send +// again. +func TestRateLimitRefusesPastAllowanceUntilRetryAfter(t *testing.T) { + t.Parallel() + + // synctest runs this on a fake clock: time.Sleep returns at once, + // with the clock moved on. + synctest.Test(t, func(t *testing.T) { + const perMinute = 2 + + handler := (&middleware.Middleware{}).RateLimit(perMinute)(okHandler()) + + post := func() *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + req := httptest.NewRequestWithContext(t.Context(), + http.MethodPost, "/api/v1/reports", http.NoBody) + handler.ServeHTTP(rec, req) + + return rec + } + + for i := range perMinute { + if code := post().Code; code != http.StatusOK { + t.Fatalf("request %d: status = %d, want %d", + i+1, code, http.StatusOK) + } + } + + rec := post() + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("request past the allowance: status = %d, want %d", + rec.Code, http.StatusTooManyRequests) + } + + if got := rec.Body.String(); got != "{\"status\":\"error\"}\n" { + t.Errorf("body = %q, want %q", got, "{\"status\":\"error\"}\n") + } + + // Two a minute: one request's allowance comes back every 30s. + if got := rec.Header().Get("Retry-After"); got != "30" { + t.Fatalf("Retry-After = %q, want %q", got, "30") + } + + time.Sleep(30 * time.Second) + + if code := post().Code; code != http.StatusOK { + t.Fatalf("after Retry-After: status = %d, want %d", + code, http.StatusOK) + } + }) +} + +// TestRateLimitIsPerForwardedClient checks that clients behind a +// trusted proxy each get their own allowance: the limit is keyed on +// the client address clientIP resolves, not on the proxy's. +func TestRateLimitIsPerForwardedClient(t *testing.T) { + t.Parallel() + + const otherClient = "203.0.113.8" + + mw := middleware.NewWithTrustedProxies(mustPrefixes(t, "127.0.0.1/32")) + handler := mw.RateLimit(1)(okHandler()) + + post := func(client string) int { + rec := httptest.NewRecorder() + req := httptest.NewRequestWithContext(t.Context(), + http.MethodPost, "/api/v1/reports", http.NoBody) + req.RemoteAddr = loopbackPeer + req.Header.Set("X-Forwarded-For", client) + handler.ServeHTTP(rec, req) + + return rec.Code + } + + if code := post(forwardedIP); code != http.StatusOK { + t.Fatalf("first request: status = %d, want %d", code, http.StatusOK) + } + + if code := post(forwardedIP); code != http.StatusTooManyRequests { + t.Fatalf("same client again: status = %d, want %d", + code, http.StatusTooManyRequests) + } + + if code := post(otherClient); code != http.StatusOK { + t.Fatalf("other client behind the same proxy: status = %d, want %d", + code, http.StatusOK) + } +} + +// preflight sends cors the preflight request a browser makes before +// it POSTs JSON from origin. +func preflight( + t *testing.T, + cors func(http.Handler) http.Handler, + origin string, +) *httptest.ResponseRecorder { + t.Helper() + + rec := httptest.NewRecorder() + req := httptest.NewRequestWithContext(t.Context(), + http.MethodOptions, "/api/v1/reports", http.NoBody) + req.Header.Set("Origin", origin) + req.Header.Set("Access-Control-Request-Method", http.MethodPost) + req.Header.Set("Access-Control-Request-Headers", "content-type") + cors(okHandler()).ServeHTTP(rec, req) + + return rec +} + +// TestCORSWithoutOriginsAddsNoHeaders checks the default: with no +// origins configured, no origin is given any CORS header. +func TestCORSWithoutOriginsAddsNoHeaders(t *testing.T) { + t.Parallel() + + rec := preflight(t, + (&middleware.Middleware{}).CORS(nil), "https://elsewhere.example") + + for name := range rec.Header() { + if strings.HasPrefix(name, "Access-Control-") { + t.Errorf("CORS header %s set with no origins configured", name) + } + } +} + +func TestCORSAllowsOnlyListedOrigins(t *testing.T) { + t.Parallel() + + const listed = "https://netwatch.example" + + cors := (&middleware.Middleware{}).CORS([]string{listed}) + + cases := []struct { + name string + origin string + want string + }{ + {name: "listed origin allowed", origin: listed, want: listed}, + {name: "other origin refused", origin: "https://elsewhere.example"}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + rec := preflight(t, cors, tc.origin) + + got := rec.Header().Get("Access-Control-Allow-Origin") + if got != tc.want { + t.Errorf("Access-Control-Allow-Origin = %q, want %q", + got, tc.want) + } + }) + } +} diff --git a/backend/internal/reportbuf/export_test.go b/backend/internal/reportbuf/export_test.go new file mode 100644 index 0000000..6440af9 --- /dev/null +++ b/backend/internal/reportbuf/export_test.go @@ -0,0 +1,7 @@ +package reportbuf + +// Flush writes the buffered reports to a file now, as the periodic +// flush does, so tests need not wait a minute for it. +func (b *Buffer) Flush() error { + return b.flushLocked() +} diff --git a/backend/internal/reportbuf/reportbuf.go b/backend/internal/reportbuf/reportbuf.go index 4626770..318b77f 100644 --- a/backend/internal/reportbuf/reportbuf.go +++ b/backend/internal/reportbuf/reportbuf.go @@ -6,11 +6,13 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io/fs" "log/slog" "os" "path/filepath" + "strings" "sync" "time" @@ -27,8 +29,16 @@ const ( defaultDataDir = "./data/reports" dirPerms fs.FileMode = 0o750 filePerms fs.FileMode = 0o640 + + // Report files are named filePrefix + timestamp + fileSuffix. + filePrefix = "reports-" + fileSuffix = ".jsonl.zst" ) +// ErrFull is returned by Append when storing the report would +// take the report files past the configured maximum size. +var ErrFull = errors.New("report files at their size cap") + // Params defines the dependencies for Buffer. type Params struct { fx.In @@ -44,8 +54,13 @@ type Buffer struct { dataDir string done chan struct{} log *slog.Logger + maxBytes int64 mu sync.Mutex stopOnce sync.Once + // usedBytes is what Append checks against maxBytes: the size + // of the report files in dataDir, plus the reports not yet + // written to one at their uncompressed size. + usedBytes int64 } // New creates a Buffer and registers lifecycle hooks to @@ -60,9 +75,10 @@ func New( } b := &Buffer{ - dataDir: dir, - done: make(chan struct{}), - log: params.Logger.Get(), + dataDir: dir, + done: make(chan struct{}), + log: params.Logger.Get(), + maxBytes: params.Config.DataDirMaxBytes, } lc.Append(fx.Hook{ @@ -72,6 +88,12 @@ func New( return fmt.Errorf("create data dir: %w", err) } + // Report files left by earlier runs count too. + b.usedBytes, err = reportFilesSize(b.dataDir) + if err != nil { + return err + } + go b.flushLoop() return nil @@ -97,15 +119,27 @@ func New( } // Append marshals v as a single JSON line and appends it to -// the buffer. If the buffer reaches the size threshold, it is -// drained and written to disk asynchronously. +// the buffer. It stores nothing and returns ErrFull if the line +// would take usedBytes past maxBytes. If the buffer reaches the +// size threshold, it is drained and written to disk +// asynchronously. func (b *Buffer) Append(v any) error { line, err := json.Marshal(v) if err != nil { return fmt.Errorf("marshal report: %w", err) } + lineBytes := int64(len(line)) + 1 // with its newline + b.mu.Lock() + + if b.usedBytes+lineBytes > b.maxBytes { + b.mu.Unlock() + + return ErrFull + } + + b.usedBytes += lineBytes b.buf.Write(line) b.buf.WriteByte('\n') @@ -178,8 +212,7 @@ func (b *Buffer) drainBuf() []byte { // in the data directory. func (b *Buffer) writeFile(data []byte) error { ts := time.Now().UTC().Format("2006-01-02T15-04-05.000Z") - name := fmt.Sprintf("reports-%s.jsonl.zst", ts) - path := filepath.Join(b.dataDir, name) + path := filepath.Join(b.dataDir, filePrefix+ts+fileSuffix) // path is built from the operator-supplied dataDir plus a // generated timestamp, so it carries no external input. @@ -214,10 +247,51 @@ func (b *Buffer) writeFile(data []byte) error { return fmt.Errorf("close zstd encoder: %w", err) } + info, err := f.Stat() + if err != nil { + return fmt.Errorf("stat report file: %w", err) + } + err = f.Close() if err != nil { return fmt.Errorf("close report file: %w", err) } + // The reports counted at their uncompressed size while they + // waited; now they count as the file. After a failed write they + // stay counted as they were, which errs toward refusing reports + // early rather than letting the files pass the cap. + b.mu.Lock() + b.usedBytes += info.Size() - int64(len(data)) + b.mu.Unlock() + return nil } + +// reportFilesSize returns the total size of the report files in +// dir. +func reportFilesSize(dir string) (int64, error) { + entries, err := os.ReadDir(dir) + if err != nil { + return 0, fmt.Errorf("read data dir: %w", err) + } + + var total int64 + + for _, entry := range entries { + name := entry.Name() + if !strings.HasPrefix(name, filePrefix) || + !strings.HasSuffix(name, fileSuffix) { + continue + } + + info, err := entry.Info() + if err != nil { + return 0, fmt.Errorf("stat report file: %w", err) + } + + total += info.Size() + } + + return total, nil +} diff --git a/backend/internal/reportbuf/reportbuf_test.go b/backend/internal/reportbuf/reportbuf_test.go index ade9457..442c213 100644 --- a/backend/internal/reportbuf/reportbuf_test.go +++ b/backend/internal/reportbuf/reportbuf_test.go @@ -1,9 +1,12 @@ package reportbuf_test import ( + "encoding/json" "errors" "io/fs" "os" + "path/filepath" + "strconv" "strings" "testing" @@ -94,6 +97,131 @@ func TestFailedFinalFlushFailsStop(t *testing.T) { } } +// startBuffer starts a Buffer through fx, as main does, with the +// DATA_DIR and DATA_DIR_MAX_BYTES the calling test has set. +func startBuffer(t *testing.T) *reportbuf.Buffer { + t.Helper() + + var buf *reportbuf.Buffer + + app := fxtest.New(t, + fx.Provide( + globals.New, + logger.New, + config.New, + reportbuf.New, + ), + fx.Populate(&buf), + ) + + app.RequireStart() + t.Cleanup(app.RequireStop) + + return buf +} + +// lineBytes is what one report takes in the buffer: its JSON and a +// newline. +func lineBytes(t *testing.T, report any) int { + t.Helper() + + line, err := json.Marshal(report) + if err != nil { + t.Fatalf("marshal report: %v", err) + } + + return len(line) + 1 +} + +func TestAppendPastCapIsRefused(t *testing.T) { + report := map[string]string{"id": "cap"} + + t.Setenv("DATA_DIR", t.TempDir()) + t.Setenv("DATA_DIR_MAX_BYTES", strconv.Itoa(lineBytes(t, report))) + + buf := startBuffer(t) + + err := buf.Append(report) + if err != nil { + t.Fatalf("report that fills the cap exactly: %v", err) + } + + err = buf.Append(report) + if !errors.Is(err, reportbuf.ErrFull) { + t.Fatalf("report past the cap: error = %v, want ErrFull", err) + } +} + +// TestCapCountsReportFilesAlreadyInDataDir starts on a data +// directory holding a report file from an earlier run, and a file +// that is not a report, which must not count. +func TestCapCountsReportFilesAlreadyInDataDir(t *testing.T) { + const earlierBytes = 100 + + report := map[string]string{"id": "cap"} + dir := t.TempDir() + + writeBytes(t, filepath.Join(dir, "reports-2026-01-01T00-00-00.000Z.jsonl.zst"), + earlierBytes) + writeBytes(t, filepath.Join(dir, "notes.txt"), 10*earlierBytes) + + t.Setenv("DATA_DIR", dir) + t.Setenv("DATA_DIR_MAX_BYTES", + strconv.Itoa(earlierBytes+lineBytes(t, report))) + + buf := startBuffer(t) + + err := buf.Append(report) + if err != nil { + t.Fatalf("report that fills the cap exactly: %v", err) + } + + err = buf.Append(report) + if !errors.Is(err, reportbuf.ErrFull) { + t.Fatalf("report past the cap: error = %v, want ErrFull", err) + } +} + +// TestWrittenReportsCountAtFileSize checks that once reports are +// written, they count as their compressed file, not their +// uncompressed size, which frees room under the cap. +func TestWrittenReportsCountAtFileSize(t *testing.T) { + // Repetitive, so its file is far smaller than its JSON. + report := map[string]string{"id": strings.Repeat("a", 1000)} + size := lineBytes(t, report) + + t.Setenv("DATA_DIR", t.TempDir()) + // Room for the report twice over only if the first one counts + // at its file's size by the time the second arrives. + t.Setenv("DATA_DIR_MAX_BYTES", strconv.Itoa(2*size-1)) + + buf := startBuffer(t) + + err := buf.Append(report) + if err != nil { + t.Fatalf("first report: %v", err) + } + + err = buf.Flush() + if err != nil { + t.Fatalf("flush: %v", err) + } + + err = buf.Append(report) + if err != nil { + t.Fatalf("second report, after the first was written: %v", err) + } +} + +func writeBytes(t *testing.T, path string, n int) { + t.Helper() + + err := os.WriteFile(path, make([]byte, n), 0o600) + if err != nil { + t.Fatalf("write %s: %v", path, err) + } +} + func hasReportFile(t *testing.T, dir string) bool { t.Helper() diff --git a/backend/internal/server/routes.go b/backend/internal/server/routes.go index 9d4547d..4c87f59 100644 --- a/backend/internal/server/routes.go +++ b/backend/internal/server/routes.go @@ -25,7 +25,7 @@ func (s *Server) SetupRoutes() { s.router.Use(middleware.RequestID) s.router.Use(s.mw.Logging()) s.router.Use(s.mw.SecurityHeaders()) - s.router.Use(s.mw.CORS()) + s.router.Use(s.mw.CORS(s.params.Config.CORSAllowedOrigins)) s.router.Use(s.mw.MaxBodyBytes(maxRequestBodyBytes)) s.router.Use(middleware.Timeout(requestTimeout)) @@ -35,6 +35,7 @@ func (s *Server) SetupRoutes() { ) s.router.Route("/api/v1", func(r chi.Router) { - r.Post("/reports", s.h.HandleReport()) + r.With(s.mw.RateLimit(s.params.Config.ReportsPerMinute)). + Post("/reports", s.h.HandleReport()) }) } diff --git a/backend/internal/server/routes_test.go b/backend/internal/server/routes_test.go index 1a50a7f..5569b86 100644 --- a/backend/internal/server/routes_test.go +++ b/backend/internal/server/routes_test.go @@ -19,16 +19,13 @@ import ( "go.uber.org/fx/fxtest" ) -// TestHealthCheckRejectsOversizeBody sends the health check, which -// never reads its body, a body one byte over the limit. Only the -// router-wide body limit can reject it. -func TestHealthCheckRejectsOversizeBody(t *testing.T) { - t.Parallel() +// newServer builds the server from the same constructors as main, +// never started: SetupRoutes is called directly, so nothing listens. +func newServer(t *testing.T) *server.Server { + t.Helper() var srv *server.Server - // The same constructors as main, never started: SetupRoutes is - // called directly, so nothing listens. app := fxtest.New(t, fx.Provide( config.New, @@ -50,6 +47,48 @@ func TestHealthCheckRejectsOversizeBody(t *testing.T) { srv.SetupRoutes() + return srv +} + +// TestReportsAreRateLimited checks that POST /api/v1/reports is +// behind the per-address rate limit, set here to two a minute. +func TestReportsAreRateLimited(t *testing.T) { + t.Setenv("REPORTS_PER_MINUTE", "2") + + srv := newServer(t) + + post := func() int { + rec := httptest.NewRecorder() + req := httptest.NewRequestWithContext(t.Context(), + http.MethodPost, "/api/v1/reports", + strings.NewReader(`{"clientId":"c1","hosts":[]}`), + ) + srv.ServeHTTP(rec, req) + + return rec.Code + } + + for i := range 2 { + if code := post(); code != http.StatusOK { + t.Fatalf("report %d: status = %d, want %d", + i+1, code, http.StatusOK) + } + } + + if code := post(); code != http.StatusTooManyRequests { + t.Fatalf("third report in a minute: status = %d, want %d", + code, http.StatusTooManyRequests) + } +} + +// TestHealthCheckRejectsOversizeBody sends the health check, which +// never reads its body, a body one byte over the limit. Only the +// router-wide body limit can reject it. +func TestHealthCheckRejectsOversizeBody(t *testing.T) { + t.Parallel() + + srv := newServer(t) + rec := httptest.NewRecorder() req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/.well-known/healthcheck",