fix(backend): rate-limit and cap report ingest, drop wildcard CORS (closes #20) #63

Merged
clawbot merged 1 commits from fix/bound-report-endpoint into next 2026-09-29 04:22:20 +02:00
16 changed files with 908 additions and 54 deletions
Showing only changes of commit f426513c21 - Show all commits
+10
View File
@@ -23,6 +23,16 @@ latest run passes.
# Completed Steps # 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, and an entry that is not a plain
`scheme://host[:port]` origin, `*` included, stops the server from starting.
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 - 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 only image, and `Dockerfile.backend` is gone. nginx serves the frontend on
port 8080 and proxies `/api/` and `/.well-known/healthcheck` to the backend, port 8080 and proxies `/api/` and `/.well-known/healthcheck` to the backend,
+39 -1
View File
@@ -76,12 +76,15 @@ Internal packages in `internal/` follow standard Go project layout:
### Configuration ### Configuration
| Variable | Default | Description | | Variable | Default | Description |
| ----------------- | -------------------- | -------------------------------------------------------------------------------------------------------- | | ---------------------- | -------------------- | -------------------------------------------------------------------------------------------------------- |
| `BIND_ADDRESS` | empty | IP address to listen on; empty listens on every interface | | `BIND_ADDRESS` | empty | IP address to listen on; empty listens on every interface |
| `PORT` | `8080` | HTTP listen port | | `PORT` | `8080` | HTTP listen port |
| `DATA_DIR` | `./data/reports` | Directory for compressed reports | | `DATA_DIR` | `./data/reports` | Directory for compressed reports |
| `DATA_DIR_MAX_BYTES` | `1073741824` (1 GiB) | Largest total size of the report files in `DATA_DIR`; see [Report limits](#report-limits) |
| `DEBUG` | `false` | Enable debug logging | | `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 | | `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`. `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 The loopback entries cover the reverse proxy that shares the container; the
@@ -103,6 +106,41 @@ Reports are written as `reports-<timestamp>.jsonl.zst` files in `DATA_DIR`.
Each file contains one JSON object per line, compressed with zstd. Files are Each file contains one JSON object per line, compressed with zstd. Files are
created with `O_EXCL` to prevent overwrites. 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. Each entry must be a plain
origin, `scheme://host` with an optional `:port`, as browsers send it: no path,
not even a trailing `/`, and no `*`. Any other entry stops the server from
starting, with an error naming `CORS_ALLOWED_ORIGINS`.
## TODO ## TODO
- Add integration test that POSTs a report and verifies the compressed output - Add integration test that POSTs a report and verifies the compressed output
+4 -1
View File
@@ -5,6 +5,7 @@ go 1.25.5
require ( require (
github.com/go-chi/chi/v5 v5.2.5 github.com/go-chi/chi/v5 v5.2.5
github.com/go-chi/cors v1.2.2 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/joho/godotenv v1.5.1
github.com/klauspost/compress v1.18.4 github.com/klauspost/compress v1.18.4
github.com/spf13/viper v1.21.0 github.com/spf13/viper v1.21.0
@@ -14,6 +15,7 @@ require (
require ( require (
github.com/fsnotify/fsnotify v1.9.0 // indirect github.com/fsnotify/fsnotify v1.9.0 // indirect
github.com/go-viper/mapstructure/v2 v2.4.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/pelletier/go-toml/v2 v2.2.4 // indirect
github.com/sagikazarmark/locafero v0.11.0 // indirect github.com/sagikazarmark/locafero v0.11.0 // indirect
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 // 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/cast v1.10.0 // indirect
github.com/spf13/pflag v1.0.10 // indirect github.com/spf13/pflag v1.0.10 // indirect
github.com/subosito/gotenv v1.6.0 // 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/dig v1.19.0 // indirect
go.uber.org/multierr v1.10.0 // indirect go.uber.org/multierr v1.10.0 // indirect
go.uber.org/zap v1.26.0 // indirect go.uber.org/zap v1.26.0 // indirect
go.yaml.in/yaml/v3 v3.0.4 // 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 golang.org/x/text v0.28.0 // indirect
) )
+10 -2
View File
@@ -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/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 h1:Jmey33TE+b+rB7fT8MUy1u0I4L+NARQlK6LhzKPSyQE=
github.com/go-chi/cors v1.2.2/go.mod h1:sSbTewc+6wYHBBCW7ytsFSn836hqM7JxpglAy2Vzc58= 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 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs=
github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= 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/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 h1:RPhnKRAQ4Fh8zU2FY/6ZFDwTVTxgJ/EMydqSTzE9a2c=
github.com/klauspost/compress v1.18.4/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4= 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 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= 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/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 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8=
github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU= 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 h1:BACLhebsYdpQ7IROQ1AGPjrXcP5dF80U3gKoFzbaq/4=
go.uber.org/dig v1.19.0/go.mod h1:Us0rSJiThwCv2GteUN0Q7OKvU7n5J4dxZ9JKUXozFdE= go.uber.org/dig v1.19.0/go.mod h1:Us0rSJiThwCv2GteUN0Q7OKvU7n5J4dxZ9JKUXozFdE=
go.uber.org/fx v1.24.0 h1:wE8mruvpg2kiiL1Vqd0CC+tr0/24XIB10Iwp2lLWzkg= 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.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 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= 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.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc=
golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= 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 h1:rhazDwis8INMIwQ4tpjLDzUhx6RlXqZNPEM0huQojng=
golang.org/x/text v0.28.0/go.mod h1:U8nCwOR8jO/marOQ0QbDiOngZVEBB7MAiitBuMjXiNU= 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= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
+62
View File
@@ -4,7 +4,9 @@ package config
import ( import (
"errors" "errors"
"fmt"
"log/slog" "log/slog"
"net/url"
"strings" "strings"
"sneak.berlin/go/netwatch/internal/globals" "sneak.berlin/go/netwatch/internal/globals"
@@ -23,6 +25,20 @@ import (
const defaultTrustedProxies = "127.0.0.1/32,::1/128," + const defaultTrustedProxies = "127.0.0.1/32,::1/128," +
"10.0.0.0/8,172.16.0.0/12,192.168.0.0/16" "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")
errNotOrigin = errors.New(
"must be an origin, scheme://host with an optional port",
)
)
// Params defines the dependencies for Config. // Params defines the dependencies for Config.
type Params struct { type Params struct {
fx.In fx.In
@@ -34,11 +50,14 @@ type Params struct {
// Config holds the resolved application configuration. // Config holds the resolved application configuration.
type Config struct { type Config struct {
BindAddress string BindAddress string
CORSAllowedOrigins []string
DataDir string DataDir string
DataDirMaxBytes int64
Debug bool Debug bool
MetricsPassword string MetricsPassword string
MetricsUsername string MetricsUsername string
Port int Port int
ReportsPerMinute int
SentryDSN string SentryDSN string
TrustedProxies []string TrustedProxies []string
log *slog.Logger log *slog.Logger
@@ -61,11 +80,15 @@ func New(
viper.AutomaticEnv() 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", "./data/reports")
viper.SetDefault("DATA_DIR_MAX_BYTES", defaultDataDirMaxBytes)
viper.SetDefault("DEBUG", "false") viper.SetDefault("DEBUG", "false")
// An empty BIND_ADDRESS listens on every interface. // An empty BIND_ADDRESS listens on every interface.
viper.SetDefault("BIND_ADDRESS", "") viper.SetDefault("BIND_ADDRESS", "")
viper.SetDefault("PORT", "8080") viper.SetDefault("PORT", "8080")
viper.SetDefault("REPORTS_PER_MINUTE", defaultReportsPerMinute)
viper.SetDefault("SENTRY_DSN", "") viper.SetDefault("SENTRY_DSN", "")
viper.SetDefault("METRICS_USERNAME", "") viper.SetDefault("METRICS_USERNAME", "")
viper.SetDefault("METRICS_PASSWORD", "") viper.SetDefault("METRICS_PASSWORD", "")
@@ -82,17 +105,37 @@ func New(
s := &Config{ s := &Config{
BindAddress: viper.GetString("BIND_ADDRESS"), BindAddress: viper.GetString("BIND_ADDRESS"),
CORSAllowedOrigins: splitList(viper.GetString("CORS_ALLOWED_ORIGINS")),
DataDir: viper.GetString("DATA_DIR"), DataDir: viper.GetString("DATA_DIR"),
DataDirMaxBytes: viper.GetInt64("DATA_DIR_MAX_BYTES"),
Debug: viper.GetBool("DEBUG"), Debug: viper.GetBool("DEBUG"),
MetricsPassword: viper.GetString("METRICS_PASSWORD"), MetricsPassword: viper.GetString("METRICS_PASSWORD"),
MetricsUsername: viper.GetString("METRICS_USERNAME"), MetricsUsername: viper.GetString("METRICS_USERNAME"),
Port: viper.GetInt("PORT"), Port: viper.GetInt("PORT"),
ReportsPerMinute: viper.GetInt("REPORTS_PER_MINUTE"),
SentryDSN: viper.GetString("SENTRY_DSN"), SentryDSN: viper.GetString("SENTRY_DSN"),
TrustedProxies: splitList(viper.GetString("TRUSTED_PROXIES")), TrustedProxies: splitList(viper.GetString("TRUSTED_PROXIES")),
log: log, log: log,
params: &params, params: &params,
} }
// 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)
}
err = checkOrigins(s.CORSAllowedOrigins)
if err != nil {
return nil, err
}
if s.Debug { if s.Debug {
params.Logger.EnableDebugLogging() params.Logger.EnableDebugLogging()
s.log = params.Logger.Get() s.log = params.Logger.Get()
@@ -101,6 +144,25 @@ func New(
return s, nil return s, nil
} }
// checkOrigins fails on the first CORS_ALLOWED_ORIGINS entry that is
// not a plain origin, scheme://host with an optional port, as browsers
// send it; anything more, such as a trailing "/", would match no page.
// go-chi/cors reads a "*" anywhere in an entry as a wildcard, so no
// entry may contain one.
func checkOrigins(origins []string) error {
for _, origin := range origins {
u, err := url.Parse(origin)
if err != nil || u.Scheme == "" || u.Host == "" ||
strings.Contains(origin, "*") ||
origin != u.Scheme+"://"+u.Host {
return fmt.Errorf("CORS_ALLOWED_ORIGINS %q: %w",
origin, errNotOrigin)
}
}
return nil
}
// splitList turns a comma-separated setting into a trimmed // splitList turns a comma-separated setting into a trimmed
// slice, dropping empty entries. // slice, dropping empty entries.
func splitList(raw string) []string { func splitList(raw string) []string {
+64
View File
@@ -0,0 +1,64 @@
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")
}
// TestCORSAllowedOriginsMustBeOrigins: "*" would let every origin in,
// and an entry that is not a plain origin would match no page.
func TestCORSAllowedOriginsMustBeOrigins(t *testing.T) {
for _, entry := range []string{
"*",
"https://*.netwatch.example",
"netwatch.example",
"https://netwatch.example/",
} {
t.Run(entry, func(t *testing.T) {
t.Setenv("CORS_ALLOWED_ORIGINS", entry)
requireConfigError(t, "CORS_ALLOWED_ORIGINS")
})
}
}
+18 -2
View File
@@ -4,6 +4,8 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"net/http" "net/http"
"sneak.berlin/go/netwatch/internal/reportbuf"
) )
// maxLoggedFieldBytes bounds untrusted text (string fields, // maxLoggedFieldBytes bounds untrusted text (string fields,
@@ -55,10 +57,9 @@ func (s *Handlers) HandleReport() http.HandlerFunc {
err = s.buf.Append(rpt) err = s.buf.Append(rpt)
if err != nil { if err != nil {
s.log.Error("failed to buffer report", "error", err)
s.respondJSON(w, r, s.respondJSON(w, r,
&response{Status: "error"}, &response{Status: "error"},
http.StatusInternalServerError, s.appendErrorStatus(err),
) )
return return
@@ -88,6 +89,21 @@ func (s *Handlers) decodeErrorStatus(err error) int {
return http.StatusBadRequest 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 // logReportReceived logs an accepted report. Untrusted fields are
// bounded (client_id, timestamp) or reduced to a length // bounded (client_id, timestamp) or reduced to a length
// (geo_bytes) so the raw attacker-controlled body never reaches // (geo_bytes) so the raw attacker-controlled body never reaches
+27
View File
@@ -13,6 +13,7 @@ import (
"sneak.berlin/go/netwatch/internal/handlers" "sneak.berlin/go/netwatch/internal/handlers"
"sneak.berlin/go/netwatch/internal/middleware" "sneak.berlin/go/netwatch/internal/middleware"
"sneak.berlin/go/netwatch/internal/reportbuf"
) )
var errStorageFailed = errors.New("storage failed") 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) { func TestHandleReportMalformedJSONIs400(t *testing.T) {
t.Parallel() t.Parallel()
@@ -15,6 +15,13 @@ func NewWithLogger(log *slog.Logger) *Middleware {
return &Middleware{log: log} 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( func ClientIP(
remoteAddr string, remoteAddr string,
header http.Header, header http.Header,
+36 -13
View File
@@ -20,6 +20,7 @@ import (
"github.com/go-chi/chi/v5/middleware" "github.com/go-chi/chi/v5/middleware"
"github.com/go-chi/cors" "github.com/go-chi/cors"
"github.com/go-chi/httprate"
"go.uber.org/fx" "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. // CORS returns middleware that lets pages served from the given
func (s *Middleware) CORS() func(http.Handler) http.Handler { // 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{ return cors.Handler(cors.Options{
AllowedOrigins: []string{"*"}, AllowedOrigins: origins,
AllowedMethods: []string{ AllowedMethods: []string{http.MethodGet, http.MethodPost},
"GET", "POST", "PUT", "DELETE", "OPTIONS", AllowedHeaders: []string{"Content-Type"},
},
AllowedHeaders: []string{
"Accept",
"Authorization",
"Content-Type",
"X-CSRF-Token",
},
ExposedHeaders: []string{"Link"},
AllowCredentials: false, AllowCredentials: false,
MaxAge: corsMaxAgeSec, 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)
},
),
)
}
@@ -10,6 +10,8 @@ import (
"net/netip" "net/netip"
"strings" "strings"
"testing" "testing"
"testing/synctest"
"time"
"sneak.berlin/go/netwatch/internal/middleware" "sneak.berlin/go/netwatch/internal/middleware"
) )
@@ -300,3 +302,204 @@ func TestRecovererRepanicsOnAbortHandler(t *testing.T) {
t.Errorf("abort was logged: %q", logbuf.String()) 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)
}
}
})
}
// postForwarded sends handler a report from peer that names client in
// X-Forwarded-For, and returns the status.
func postForwarded(
t *testing.T,
handler http.Handler,
peer, client string,
) int {
t.Helper()
rec := httptest.NewRecorder()
req := httptest.NewRequestWithContext(t.Context(),
http.MethodPost, "/api/v1/reports", http.NoBody)
req.RemoteAddr = peer
req.Header.Set("X-Forwarded-For", client)
handler.ServeHTTP(rec, req)
return rec.Code
}
// 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())
code := postForwarded(t, handler, loopbackPeer, forwardedIP)
if code != http.StatusOK {
t.Fatalf("first request: status = %d, want %d", code, http.StatusOK)
}
code = postForwarded(t, handler, loopbackPeer, forwardedIP)
if code != http.StatusTooManyRequests {
t.Fatalf("same client again: status = %d, want %d",
code, http.StatusTooManyRequests)
}
code = postForwarded(t, handler, loopbackPeer, otherClient)
if code != http.StatusOK {
t.Fatalf("other client behind the same proxy: status = %d, want %d",
code, http.StatusOK)
}
}
// TestRateLimitIgnoresForwardedForFromUntrustedPeer checks that a
// peer that is not a trusted proxy cannot get a fresh allowance by
// naming a different client in X-Forwarded-For on each request.
func TestRateLimitIgnoresForwardedForFromUntrustedPeer(t *testing.T) {
t.Parallel()
const untrustedPeer = "198.51.100.4:5000"
mw := middleware.NewWithTrustedProxies(mustPrefixes(t, "127.0.0.1/32"))
handler := mw.RateLimit(1)(okHandler())
code := postForwarded(t, handler, untrustedPeer, "203.0.113.8")
if code != http.StatusOK {
t.Fatalf("first request: status = %d, want %d", code, http.StatusOK)
}
code = postForwarded(t, handler, untrustedPeer, "203.0.113.9")
if code != http.StatusTooManyRequests {
t.Fatalf("same peer naming another client: status = %d, want %d",
code, http.StatusTooManyRequests)
}
}
// 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)
}
})
}
}
@@ -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()
}
+78 -4
View File
@@ -6,11 +6,13 @@ import (
"bytes" "bytes"
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io/fs" "io/fs"
"log/slog" "log/slog"
"os" "os"
"path/filepath" "path/filepath"
"strings"
"sync" "sync"
"time" "time"
@@ -27,8 +29,16 @@ const (
defaultDataDir = "./data/reports" defaultDataDir = "./data/reports"
dirPerms fs.FileMode = 0o750 dirPerms fs.FileMode = 0o750
filePerms fs.FileMode = 0o640 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. // Params defines the dependencies for Buffer.
type Params struct { type Params struct {
fx.In fx.In
@@ -44,8 +54,13 @@ type Buffer struct {
dataDir string dataDir string
done chan struct{} done chan struct{}
log *slog.Logger log *slog.Logger
maxBytes int64
mu sync.Mutex mu sync.Mutex
stopOnce sync.Once 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 // New creates a Buffer and registers lifecycle hooks to
@@ -63,6 +78,7 @@ func New(
dataDir: dir, dataDir: dir,
done: make(chan struct{}), done: make(chan struct{}),
log: params.Logger.Get(), log: params.Logger.Get(),
maxBytes: params.Config.DataDirMaxBytes,
} }
lc.Append(fx.Hook{ lc.Append(fx.Hook{
@@ -72,6 +88,12 @@ func New(
return fmt.Errorf("create data dir: %w", err) 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() go b.flushLoop()
return nil return nil
@@ -97,15 +119,27 @@ func New(
} }
// Append marshals v as a single JSON line and appends it to // Append marshals v as a single JSON line and appends it to
// the buffer. If the buffer reaches the size threshold, it is // the buffer. It stores nothing and returns ErrFull if the line
// drained and written to disk asynchronously. // 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 { func (b *Buffer) Append(v any) error {
line, err := json.Marshal(v) line, err := json.Marshal(v)
if err != nil { if err != nil {
return fmt.Errorf("marshal report: %w", err) return fmt.Errorf("marshal report: %w", err)
} }
lineBytes := int64(len(line)) + 1 // with its newline
b.mu.Lock() b.mu.Lock()
if b.usedBytes+lineBytes > b.maxBytes {
b.mu.Unlock()
return ErrFull
}
b.usedBytes += lineBytes
b.buf.Write(line) b.buf.Write(line)
b.buf.WriteByte('\n') b.buf.WriteByte('\n')
@@ -178,8 +212,7 @@ func (b *Buffer) drainBuf() []byte {
// in the data directory. // in the data directory.
func (b *Buffer) writeFile(data []byte) error { func (b *Buffer) writeFile(data []byte) error {
ts := time.Now().UTC().Format("2006-01-02T15-04-05.000Z") ts := time.Now().UTC().Format("2006-01-02T15-04-05.000Z")
name := fmt.Sprintf("reports-%s.jsonl.zst", ts) path := filepath.Join(b.dataDir, filePrefix+ts+fileSuffix)
path := filepath.Join(b.dataDir, name)
// path is built from the operator-supplied dataDir plus a // path is built from the operator-supplied dataDir plus a
// generated timestamp, so it carries no external input. // 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) 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() err = f.Close()
if err != nil { if err != nil {
return fmt.Errorf("close report file: %w", err) 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 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
}
@@ -1,11 +1,17 @@
package reportbuf_test package reportbuf_test
import ( import (
"encoding/json"
"errors" "errors"
"io/fs" "io/fs"
"os" "os"
"path/filepath"
"strconv"
"strings" "strings"
"sync"
"sync/atomic"
"testing" "testing"
"time"
"sneak.berlin/go/netwatch/internal/config" "sneak.berlin/go/netwatch/internal/config"
"sneak.berlin/go/netwatch/internal/globals" "sneak.berlin/go/netwatch/internal/globals"
@@ -94,6 +100,253 @@ 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)
}
// The second report is written at shutdown, and must not land in
// the first file's millisecond (see TestWrittenReportsKeepCounting).
time.Sleep(time.Millisecond)
err = buf.Append(report)
if err != nil {
t.Fatalf("second report, after the first was written: %v", err)
}
}
// TestWrittenReportsKeepCounting writes one report file after another
// under a small cap: each report must be taken while the files on disk
// leave room for it, and refused once they do not.
func TestWrittenReportsKeepCounting(t *testing.T) {
const maxBytes = 200
report := map[string]string{"id": "written"}
size := int64(lineBytes(t, report))
dir := t.TempDir()
t.Setenv("DATA_DIR", dir)
t.Setenv("DATA_DIR_MAX_BYTES", strconv.Itoa(maxBytes))
buf := startBuffer(t)
// Every file takes at least a byte, so they fill the cap within
// maxBytes rounds.
for range maxBytes {
used := reportFilesBytes(t, dir)
err := buf.Append(report)
if used+size > maxBytes {
if !errors.Is(err, reportbuf.ErrFull) {
t.Fatalf("with %d bytes of report files: error = %v, "+
"want ErrFull", used, err)
}
return
}
if err != nil {
t.Fatalf("with %d bytes of report files: %v", used, err)
}
// Report files are named to the millisecond; two in the same
// one collide (https://git.eeqj.de/sneak/netwatch/issues/61).
time.Sleep(time.Millisecond)
err = buf.Flush()
if err != nil {
t.Fatalf("flush: %v", err)
}
}
t.Fatal("the report files never filled the cap")
}
// TestConcurrentAppendsStopAtCap appends from many goroutines at once
// with room for exactly roomFor reports: exactly that many must be
// taken, which holds only if Append checks and counts each report
// under one lock.
func TestConcurrentAppendsStopAtCap(t *testing.T) {
const (
roomFor = 5
senders = 50
)
// Large, so each Append takes long enough for the senders to
// overlap while the cap is reached.
report := map[string]string{"id": strings.Repeat("a", 1_000_000)}
t.Setenv("DATA_DIR", t.TempDir())
t.Setenv("DATA_DIR_MAX_BYTES",
strconv.Itoa(roomFor*lineBytes(t, report)))
buf := startBuffer(t)
var (
taken atomic.Int64
wg sync.WaitGroup
)
start := make(chan struct{})
for range senders {
wg.Go(func() {
<-start
err := buf.Append(report)
if err == nil {
taken.Add(1)
} else if !errors.Is(err, reportbuf.ErrFull) {
t.Errorf("append: %v", err)
}
})
}
close(start)
wg.Wait()
if got := taken.Load(); got != roomFor {
t.Fatalf("%d reports taken, want %d", got, roomFor)
}
}
// reportFilesBytes returns the total size of the report files in dir.
func reportFilesBytes(t *testing.T, dir string) int64 {
t.Helper()
paths, err := filepath.Glob(filepath.Join(dir, "reports-*.jsonl.zst"))
if err != nil {
t.Fatalf("list report files: %v", err)
}
var total int64
for _, path := range paths {
info, statErr := os.Stat(path)
if statErr != nil {
t.Fatalf("stat %s: %v", path, statErr)
}
total += info.Size()
}
return total
}
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 { func hasReportFile(t *testing.T, dir string) bool {
t.Helper() t.Helper()
+3 -2
View File
@@ -25,7 +25,7 @@ func (s *Server) SetupRoutes() {
s.router.Use(middleware.RequestID) s.router.Use(middleware.RequestID)
s.router.Use(s.mw.Logging()) s.router.Use(s.mw.Logging())
s.router.Use(s.mw.SecurityHeaders()) 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(s.mw.MaxBodyBytes(maxRequestBodyBytes))
s.router.Use(middleware.Timeout(requestTimeout)) s.router.Use(middleware.Timeout(requestTimeout))
@@ -35,6 +35,7 @@ func (s *Server) SetupRoutes() {
) )
s.router.Route("/api/v1", func(r chi.Router) { 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())
}) })
} }
+58
View File
@@ -49,6 +49,64 @@ func newServer(t *testing.T) *server.Server {
return srv 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)
}
}
// TestCORSAllowedOriginsReachTheRouter checks that an origin listed in
// CORS_ALLOWED_ORIGINS is allowed by the router, not only when handed
// to the CORS middleware directly.
func TestCORSAllowedOriginsReachTheRouter(t *testing.T) {
const origin = "https://netwatch.example:8443"
t.Setenv("CORS_ALLOWED_ORIGINS", origin)
srv := newServer(t)
srv.SetupRoutes()
// The preflight a browser sends before it POSTs JSON from origin.
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")
srv.ServeHTTP(rec, req)
got := rec.Header().Get("Access-Control-Allow-Origin")
if got != origin {
t.Fatalf("Access-Control-Allow-Origin = %q, want %q", got, origin)
}
}
// TestHealthCheckRejectsOversizeBody sends the health check, which // TestHealthCheckRejectsOversizeBody sends the health check, which
// never reads its body, a body one byte over the limit. Only the // never reads its body, a body one byte over the limit. Only the
// router-wide body limit can reject it. // router-wide body limit can reject it.