fix(backend): rate-limit and cap report ingest, drop wildcard CORS (closes #20)
check / check (push) Successful in 11s
check / check (push) Successful in 11s
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, counted by go-chi/httprate over a sliding minute); past that it gets 429 with Retry-After. reportbuf refuses a report that would take the report files past DATA_DIR_MAX_BYTES (default 1 GiB), counting the files already in DATA_DIR and unwritten reports at their uncompressed size; the handler answers 507. CORS adds nothing unless CORS_ALLOWED_ORIGINS lists origins. A limit that is not a positive number, or an origin that is not a plain scheme://host[:port], stops the server from starting. Model: opus-5-5
This commit was merged in pull request #63.
This commit is contained in:
@@ -4,7 +4,9 @@ package config
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"sneak.berlin/go/netwatch/internal/globals"
|
||||
@@ -23,6 +25,20 @@ 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")
|
||||
errNotOrigin = errors.New(
|
||||
"must be an origin, scheme://host with an optional port",
|
||||
)
|
||||
)
|
||||
|
||||
// Params defines the dependencies for Config.
|
||||
type Params struct {
|
||||
fx.In
|
||||
@@ -33,16 +49,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 +80,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 +104,36 @@ 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)
|
||||
}
|
||||
|
||||
err = checkOrigins(s.CORSAllowedOrigins)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if s.Debug {
|
||||
@@ -101,6 +144,25 @@ func New(
|
||||
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
|
||||
// slice, dropping empty entries.
|
||||
func splitList(raw string) []string {
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
},
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -10,6 +10,8 @@ import (
|
||||
"net/netip"
|
||||
"strings"
|
||||
"testing"
|
||||
"testing/synctest"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/netwatch/internal/middleware"
|
||||
)
|
||||
@@ -300,3 +302,204 @@ 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)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// 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()
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -1,11 +1,17 @@
|
||||
package reportbuf_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/netwatch/internal/config"
|
||||
"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 {
|
||||
t.Helper()
|
||||
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
}
|
||||
|
||||
@@ -49,6 +49,64 @@ 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)
|
||||
}
|
||||
}
|
||||
|
||||
// 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
|
||||
// never reads its body, a body one byte over the limit. Only the
|
||||
// router-wide body limit can reject it.
|
||||
|
||||
Reference in New Issue
Block a user