fix(backend): rate-limit and cap report ingest, drop wildcard CORS (closes #20)
check / check (push) Successful in 58s

POST /api/v1/reports stays unauthenticated but is bounded. Each client
address, as the trusted-proxy logic resolves it, may send
REPORTS_PER_MINUTE reports a minute (default 60, all at once if it
likes), using golang.org/x/time/rate; past that it gets 429 with
Retry-After. Buckets that have refilled are dropped once a minute, so
idle addresses do not pile up. reportbuf refuses a report that would
take the report files past DATA_DIR_MAX_BYTES (default 1 GiB) with
ErrFull, answered with 507; the count starts from the files already in
DATA_DIR, and reports not yet written count at their uncompressed size.
CORS adds nothing unless CORS_ALLOWED_ORIGINS lists origins. A limit
that is not a positive number stops the server from starting.

Model: opus-5-5
This commit is contained in:
2026-09-29 00:18:52 +00:00
parent de4e86c433
commit 2bf52ba7ff
17 changed files with 770 additions and 55 deletions
@@ -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,
+94 -13
View File
@@ -7,11 +7,14 @@ import (
"fmt"
"io"
"log/slog"
"math"
"net"
"net/http"
"net/netip"
"runtime/debug"
"strconv"
"strings"
"sync"
"time"
"sneak.berlin/go/netwatch/internal/config"
@@ -21,6 +24,7 @@ import (
"github.com/go-chi/chi/v5/middleware"
"github.com/go-chi/cors"
"go.uber.org/fx"
"golang.org/x/time/rate"
)
const corsMaxAgeSec = 300
@@ -320,21 +324,98 @@ func (s *Middleware) Recoverer() func(http.Handler) http.Handler {
}
}
// CORS returns middleware that adds permissive CORS headers.
func (s *Middleware) CORS() func(http.Handler) http.Handler {
// CORS returns middleware that lets pages served from the given
// origins call the API. With no origins it adds no CORS headers at
// all, so only same-origin pages can use the API. That case must not
// reach cors.Handler, which treats an empty origin list as "allow
// every origin".
func (s *Middleware) CORS(
origins []string,
) func(http.Handler) http.Handler {
if len(origins) == 0 {
return func(next http.Handler) http.Handler { return next }
}
return cors.Handler(cors.Options{
AllowedOrigins: []string{"*"},
AllowedMethods: []string{
"GET", "POST", "PUT", "DELETE", "OPTIONS",
},
AllowedHeaders: []string{
"Accept",
"Authorization",
"Content-Type",
"X-CSRF-Token",
},
ExposedHeaders: []string{"Link"},
AllowedOrigins: origins,
AllowedMethods: []string{http.MethodGet, http.MethodPost},
AllowedHeaders: []string{"Content-Type"},
AllowCredentials: false,
MaxAge: corsMaxAgeSec,
})
}
// RateLimit returns middleware that allows each client address
// perMinute requests a minute, all at once if it likes, and answers
// the rest with 429 and a Retry-After header. The address is the one
// clientIP resolves, so clients behind the reverse proxy are limited
// one by one, not together as the proxy.
func (s *Middleware) RateLimit(
perMinute int,
) func(http.Handler) http.Handler {
// One request's allowance comes back every interval, so a
// refused client can always retry after it.
interval := time.Minute / time.Duration(perMinute)
retryAfter := strconv.Itoa(int(math.Ceil(interval.Seconds())))
limiter := &addressLimiter{
burst: perMinute,
byAddr: make(map[string]*rate.Limiter),
limit: rate.Every(interval),
}
return func(next http.Handler) http.Handler {
return http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) {
addr := clientIP(r.RemoteAddr, r.Header, s.trustedProxies)
if !limiter.allow(addr, time.Now()) {
w.Header().Set("Retry-After", retryAfter)
writeJSONError(w, http.StatusTooManyRequests)
return
}
next.ServeHTTP(w, r)
},
)
}
}
// addressLimiter holds a token bucket for each client address that
// has made a request recently.
type addressLimiter struct {
burst int
byAddr map[string]*rate.Limiter
lastSweep time.Time
limit rate.Limit
mu sync.Mutex
}
// allow takes a token from addr's bucket, which starts full, and
// reports whether there was one. At most once a minute it first drops
// every bucket that has filled up again: a full bucket behaves exactly
// like the new one that would replace it, so this changes no answer,
// and the map holds only the addresses heard from recently.
func (a *addressLimiter) allow(addr string, now time.Time) bool {
a.mu.Lock()
defer a.mu.Unlock()
if now.Sub(a.lastSweep) >= time.Minute {
for key, bucket := range a.byAddr {
if bucket.TokensAt(now) >= float64(a.burst) {
delete(a.byAddr, key)
}
}
a.lastSweep = now
}
bucket, ok := a.byAddr[addr]
if !ok {
bucket = rate.NewLimiter(a.limit, a.burst)
a.byAddr[addr] = bucket
}
return bucket.AllowN(now, 1)
}
@@ -0,0 +1,46 @@
package middleware
import (
"testing"
"time"
"golang.org/x/time/rate"
)
// TestAddressLimiterDropsOnlyFullBuckets checks the sweep in allow:
// a minute after the last one, it drops an address whose bucket has
// filled up again, and keeps one still short of tokens, whose limit
// would otherwise start over.
func TestAddressLimiterDropsOnlyFullBuckets(t *testing.T) {
t.Parallel()
const (
refilled = "198.51.100.1"
drained = "198.51.100.2"
)
// Two a minute: one token back every 30 seconds.
limiter := &addressLimiter{
burst: 2,
byAddr: make(map[string]*rate.Limiter),
limit: rate.Every(30 * time.Second),
}
start := time.Now()
// The first call sweeps the empty map and takes one of two tokens.
limiter.allow(refilled, start)
limiter.allow(drained, start.Add(59*time.Second))
limiter.allow(drained, start.Add(59*time.Second))
// A minute after the first sweep, this call sweeps again.
limiter.allow("198.51.100.3", start.Add(time.Minute))
if _, ok := limiter.byAddr[refilled]; ok {
t.Error("address with a full bucket was kept")
}
if _, ok := limiter.byAddr[drained]; !ok {
t.Error("address short of tokens was dropped")
}
}
@@ -10,6 +10,8 @@ import (
"net/netip"
"strings"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/netwatch/internal/middleware"
)
@@ -300,3 +302,167 @@ func TestRecovererRepanicsOnAbortHandler(t *testing.T) {
t.Errorf("abort was logged: %q", logbuf.String())
}
}
// okHandler stands in for the route a middleware guards.
func okHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
})
}
// TestRateLimitRefusesPastAllowanceUntilRetryAfter checks one client
// address: it may use its whole allowance at once, the next request
// is refused with 429, and once Retry-After has passed it may send
// again.
func TestRateLimitRefusesPastAllowanceUntilRetryAfter(t *testing.T) {
t.Parallel()
// synctest runs this on a fake clock: time.Sleep returns at once,
// with the clock moved on.
synctest.Test(t, func(t *testing.T) {
const perMinute = 2
handler := (&middleware.Middleware{}).RateLimit(perMinute)(okHandler())
post := func() *httptest.ResponseRecorder {
rec := httptest.NewRecorder()
req := httptest.NewRequestWithContext(t.Context(),
http.MethodPost, "/api/v1/reports", http.NoBody)
handler.ServeHTTP(rec, req)
return rec
}
for i := range perMinute {
if code := post().Code; code != http.StatusOK {
t.Fatalf("request %d: status = %d, want %d",
i+1, code, http.StatusOK)
}
}
rec := post()
if rec.Code != http.StatusTooManyRequests {
t.Fatalf("request past the allowance: status = %d, want %d",
rec.Code, http.StatusTooManyRequests)
}
if got := rec.Body.String(); got != "{\"status\":\"error\"}\n" {
t.Errorf("body = %q, want %q", got, "{\"status\":\"error\"}\n")
}
// Two a minute: one request's allowance comes back every 30s.
if got := rec.Header().Get("Retry-After"); got != "30" {
t.Fatalf("Retry-After = %q, want %q", got, "30")
}
time.Sleep(30 * time.Second)
if code := post().Code; code != http.StatusOK {
t.Fatalf("after Retry-After: status = %d, want %d",
code, http.StatusOK)
}
})
}
// TestRateLimitIsPerForwardedClient checks that clients behind a
// trusted proxy each get their own allowance: the limit is keyed on
// the client address clientIP resolves, not on the proxy's.
func TestRateLimitIsPerForwardedClient(t *testing.T) {
t.Parallel()
const otherClient = "203.0.113.8"
mw := middleware.NewWithTrustedProxies(mustPrefixes(t, "127.0.0.1/32"))
handler := mw.RateLimit(1)(okHandler())
post := func(client string) int {
rec := httptest.NewRecorder()
req := httptest.NewRequestWithContext(t.Context(),
http.MethodPost, "/api/v1/reports", http.NoBody)
req.RemoteAddr = loopbackPeer
req.Header.Set("X-Forwarded-For", client)
handler.ServeHTTP(rec, req)
return rec.Code
}
if code := post(forwardedIP); code != http.StatusOK {
t.Fatalf("first request: status = %d, want %d", code, http.StatusOK)
}
if code := post(forwardedIP); code != http.StatusTooManyRequests {
t.Fatalf("same client again: status = %d, want %d",
code, http.StatusTooManyRequests)
}
if code := post(otherClient); code != http.StatusOK {
t.Fatalf("other client behind the same proxy: status = %d, want %d",
code, http.StatusOK)
}
}
// preflight sends cors the preflight request a browser makes before
// it POSTs JSON from origin.
func preflight(
t *testing.T,
cors func(http.Handler) http.Handler,
origin string,
) *httptest.ResponseRecorder {
t.Helper()
rec := httptest.NewRecorder()
req := httptest.NewRequestWithContext(t.Context(),
http.MethodOptions, "/api/v1/reports", http.NoBody)
req.Header.Set("Origin", origin)
req.Header.Set("Access-Control-Request-Method", http.MethodPost)
req.Header.Set("Access-Control-Request-Headers", "content-type")
cors(okHandler()).ServeHTTP(rec, req)
return rec
}
// TestCORSWithoutOriginsAddsNoHeaders checks the default: with no
// origins configured, no origin is given any CORS header.
func TestCORSWithoutOriginsAddsNoHeaders(t *testing.T) {
t.Parallel()
rec := preflight(t,
(&middleware.Middleware{}).CORS(nil), "https://elsewhere.example")
for name := range rec.Header() {
if strings.HasPrefix(name, "Access-Control-") {
t.Errorf("CORS header %s set with no origins configured", name)
}
}
}
func TestCORSAllowsOnlyListedOrigins(t *testing.T) {
t.Parallel()
const listed = "https://netwatch.example"
cors := (&middleware.Middleware{}).CORS([]string{listed})
cases := []struct {
name string
origin string
want string
}{
{name: "listed origin allowed", origin: listed, want: listed},
{name: "other origin refused", origin: "https://elsewhere.example"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
rec := preflight(t, cors, tc.origin)
got := rec.Header().Get("Access-Control-Allow-Origin")
if got != tc.want {
t.Errorf("Access-Control-Allow-Origin = %q, want %q",
got, tc.want)
}
})
}
}