diff --git a/TODO.md b/TODO.md index f96341a..3f0510f 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, counted by `go-chi/httprate`, 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: one container image (issue #52): the root `Dockerfile` builds the only image, and `Dockerfile.backend` is gone. nginx serves the frontend on port 8080 and proxies `/api/` and `/.well-known/healthcheck` to the backend, diff --git a/backend/README.md b/backend/README.md index fa4af1b..73138fd 100644 --- a/backend/README.md +++ b/backend/README.md @@ -75,13 +75,16 @@ Internal packages in `internal/` follow standard Go project layout: ### Configuration -| Variable | Default | Description | -| ----------------- | -------------------- | -------------------------------------------------------------------------------------------------------- | -| `BIND_ADDRESS` | empty | IP address to listen on; empty listens on every interface | -| `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 | +| ---------------------- | -------------------- | -------------------------------------------------------------------------------------------------------- | +| `BIND_ADDRESS` | empty | IP address to listen on; empty listens on every interface | +| `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 @@ -103,6 +106,38 @@ 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; past that it gets 429 with + `Retry-After: 60`. The minute slides: reports from the minute before still + count, fading out over the current one, so an address is sure never to be + refused only while it sends at most half of `REPORTS_PER_MINUTE` in any 60 + seconds. The page sends one report a minute from each open tab, so the default + of 60 refuses nothing from up to 30 tabs behind one address, such as a + household or an office sharing it, however their reports bunch up. Report + responses also carry `X-RateLimit-Limit`, `X-RateLimit-Remaining` and + `X-RateLimit-Reset` headers. +- **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..a2462e8 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -5,6 +5,7 @@ go 1.25.5 require ( github.com/go-chi/chi/v5 v5.2.5 github.com/go-chi/cors v1.2.2 + github.com/go-chi/httprate v0.16.0 github.com/joho/godotenv v1.5.1 github.com/klauspost/compress v1.18.4 github.com/spf13/viper v1.21.0 @@ -14,6 +15,7 @@ require ( require ( github.com/fsnotify/fsnotify v1.9.0 // indirect github.com/go-viper/mapstructure/v2 v2.4.0 // indirect + github.com/klauspost/cpuid/v2 v2.2.10 // indirect github.com/pelletier/go-toml/v2 v2.2.4 // indirect github.com/sagikazarmark/locafero v0.11.0 // indirect github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 // indirect @@ -21,10 +23,11 @@ require ( github.com/spf13/cast v1.10.0 // indirect github.com/spf13/pflag v1.0.10 // indirect github.com/subosito/gotenv v1.6.0 // indirect + github.com/zeebo/xxh3 v1.0.2 // indirect go.uber.org/dig v1.19.0 // indirect go.uber.org/multierr v1.10.0 // indirect go.uber.org/zap v1.26.0 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect - golang.org/x/sys v0.29.0 // indirect + golang.org/x/sys v0.30.0 // indirect golang.org/x/text v0.28.0 // indirect ) diff --git a/backend/go.sum b/backend/go.sum index de6ddf2..cac5e22 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -8,6 +8,8 @@ github.com/go-chi/chi/v5 v5.2.5 h1:Eg4myHZBjyvJmAFjFvWgrqDTXFyOzjj7YIm3L3mu6Ug= github.com/go-chi/chi/v5 v5.2.5/go.mod h1:X7Gx4mteadT3eDOMTsXzmI4/rwUpOwBHLpAfupzFJP0= github.com/go-chi/cors v1.2.2 h1:Jmey33TE+b+rB7fT8MUy1u0I4L+NARQlK6LhzKPSyQE= github.com/go-chi/cors v1.2.2/go.mod h1:sSbTewc+6wYHBBCW7ytsFSn836hqM7JxpglAy2Vzc58= +github.com/go-chi/httprate v0.16.0 h1:8V5DH9j6pSK6UQoBsTpvMyFxycqaKEIToyPKzHJjUa8= +github.com/go-chi/httprate v0.16.0/go.mod h1:A8lo+qRhk+s9LiuP5saS7XCGDXRXMcrueq0NfIuCa/I= github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs= github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= @@ -16,6 +18,8 @@ github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0= github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4= github.com/klauspost/compress v1.18.4 h1:RPhnKRAQ4Fh8zU2FY/6ZFDwTVTxgJ/EMydqSTzE9a2c= github.com/klauspost/compress v1.18.4/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4= +github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE= +github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0= github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= @@ -42,6 +46,10 @@ github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8= github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU= +github.com/zeebo/assert v1.3.0 h1:g7C04CbJuIDKNPFHmsk4hwZDO5O+kntRxzaUoNXj+IQ= +github.com/zeebo/assert v1.3.0/go.mod h1:Pq9JiuJQpG8JLJdtkwrJESF0Foym2/D9XMU5ciN/wJ0= +github.com/zeebo/xxh3 v1.0.2 h1:xZmwmqxHZA8AI603jOQ0tMqmBr9lPeFwGg6d+xy9DC0= +github.com/zeebo/xxh3 v1.0.2/go.mod h1:5NWz9Sef7zIDm2JHfFlcQvNekmcEl9ekUZQQKCYaDcA= go.uber.org/dig v1.19.0 h1:BACLhebsYdpQ7IROQ1AGPjrXcP5dF80U3gKoFzbaq/4= go.uber.org/dig v1.19.0/go.mod h1:Us0rSJiThwCv2GteUN0Q7OKvU7n5J4dxZ9JKUXozFdE= go.uber.org/fx v1.24.0 h1:wE8mruvpg2kiiL1Vqd0CC+tr0/24XIB10Iwp2lLWzkg= @@ -54,8 +62,8 @@ go.uber.org/zap v1.26.0 h1:sI7k6L95XOKS281NhVKOFCUNIvv9e0w4BF8N3u+tCRo= go.uber.org/zap v1.26.0/go.mod h1:dtElttAiwGvoJ/vj4IwHBS/gXsEu/pZ50mUIRWuG0so= go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= -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/sys v0.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc= +golang.org/x/sys v0.30.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= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 4f60de4..1974e40 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,16 +43,19 @@ type Params struct { // Config holds the resolved application configuration. type Config struct { - BindAddress string - DataDir string - Debug bool - MetricsPassword string - MetricsUsername string - Port int - SentryDSN string - TrustedProxies []string - log *slog.Logger - params *Params + BindAddress string + 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 @@ -61,11 +74,15 @@ 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") // An empty BIND_ADDRESS listens on every interface. viper.SetDefault("BIND_ADDRESS", "") viper.SetDefault("PORT", "8080") + viper.SetDefault("REPORTS_PER_MINUTE", defaultReportsPerMinute) viper.SetDefault("SENTRY_DSN", "") viper.SetDefault("METRICS_USERNAME", "") viper.SetDefault("METRICS_PASSWORD", "") @@ -81,16 +98,31 @@ func New( } s := &Config{ - BindAddress: viper.GetString("BIND_ADDRESS"), - 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, + BindAddress: viper.GetString("BIND_ADDRESS"), + 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..2e685fc 100644 --- a/backend/internal/middleware/middleware.go +++ b/backend/internal/middleware/middleware.go @@ -20,6 +20,7 @@ import ( "github.com/go-chi/chi/v5/middleware" "github.com/go-chi/cors" + "github.com/go-chi/httprate" "go.uber.org/fx" ) @@ -320,21 +321,43 @@ 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 and answers the rest with 429, the +// Retry-After header httprate sets, and the usual error body. 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 { + return httprate.LimitBy(perMinute, time.Minute, + func(r *http.Request) (string, error) { + return clientIP(r.RemoteAddr, r.Header, s.trustedProxies), nil + }, + httprate.WithLimitHandler( + func(w http.ResponseWriter, _ *http.Request) { + writeJSONError(w, http.StatusTooManyRequests) + }, + ), + ) +} diff --git a/backend/internal/middleware/middleware_test.go b/backend/internal/middleware/middleware_test.go index 8016256..126469a 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,170 @@ 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) + }) +} + +// TestRateLimitRefusesPastAllowanceThenResets checks one client +// address: it may use its whole allowance at once, the next request +// is refused with 429, and later it may send again. +func TestRateLimitRefusesPastAllowanceThenResets(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") + } + + if got := rec.Header().Get("Retry-After"); got != "60" { + t.Fatalf("Retry-After = %q, want %q", got, "60") + } + + // httprate also counts the previous minute's requests, fading + // them out over the current one, so two minutes on the whole + // allowance is back. + time.Sleep(2 * time.Minute) + + for i := range perMinute { + if code := post().Code; code != http.StatusOK { + t.Fatalf("two minutes later, request %d: status = %d, want %d", + i+1, 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 71f5117..515ce5d 100644 --- a/backend/internal/server/routes_test.go +++ b/backend/internal/server/routes_test.go @@ -49,6 +49,38 @@ func newServer(t *testing.T) *server.Server { 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) + srv.SetupRoutes() + + 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.