fix(backend): rate-limit and cap report ingest, drop wildcard CORS (closes #20)
check / check (push) Successful in 1m1s
check / check (push) Successful in 1m1s
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) 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:
@@ -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,170 @@ func TestRecovererRepanicsOnAbortHandler(t *testing.T) {
|
||||
t.Errorf("abort was logged: %q", logbuf.String())
|
||||
}
|
||||
}
|
||||
|
||||
// okHandler stands in for the route a middleware guards.
|
||||
func okHandler() http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
}
|
||||
|
||||
// TestRateLimitRefusesPastAllowanceThenResets checks one client
|
||||
// address: it may use its whole allowance at once, the next request
|
||||
// is refused with 429, and later it may send again.
|
||||
func TestRateLimitRefusesPastAllowanceThenResets(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// synctest runs this on a fake clock: time.Sleep returns at once,
|
||||
// with the clock moved on.
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
const perMinute = 2
|
||||
|
||||
handler := (&middleware.Middleware{}).RateLimit(perMinute)(okHandler())
|
||||
|
||||
post := func() *httptest.ResponseRecorder {
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequestWithContext(t.Context(),
|
||||
http.MethodPost, "/api/v1/reports", http.NoBody)
|
||||
handler.ServeHTTP(rec, req)
|
||||
|
||||
return rec
|
||||
}
|
||||
|
||||
for i := range perMinute {
|
||||
if code := post().Code; code != http.StatusOK {
|
||||
t.Fatalf("request %d: status = %d, want %d",
|
||||
i+1, code, http.StatusOK)
|
||||
}
|
||||
}
|
||||
|
||||
rec := post()
|
||||
if rec.Code != http.StatusTooManyRequests {
|
||||
t.Fatalf("request past the allowance: status = %d, want %d",
|
||||
rec.Code, http.StatusTooManyRequests)
|
||||
}
|
||||
|
||||
if got := rec.Body.String(); got != "{\"status\":\"error\"}\n" {
|
||||
t.Errorf("body = %q, want %q", got, "{\"status\":\"error\"}\n")
|
||||
}
|
||||
|
||||
if got := rec.Header().Get("Retry-After"); got != "60" {
|
||||
t.Fatalf("Retry-After = %q, want %q", got, "60")
|
||||
}
|
||||
|
||||
// httprate also counts the previous minute's requests, fading
|
||||
// them out over the current one, so two minutes on the whole
|
||||
// allowance is back.
|
||||
time.Sleep(2 * time.Minute)
|
||||
|
||||
for i := range perMinute {
|
||||
if code := post().Code; code != http.StatusOK {
|
||||
t.Fatalf("two minutes later, request %d: status = %d, want %d",
|
||||
i+1, code, http.StatusOK)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestRateLimitIsPerForwardedClient checks that clients behind a
|
||||
// trusted proxy each get their own allowance: the limit is keyed on
|
||||
// the client address clientIP resolves, not on the proxy's.
|
||||
func TestRateLimitIsPerForwardedClient(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const otherClient = "203.0.113.8"
|
||||
|
||||
mw := middleware.NewWithTrustedProxies(mustPrefixes(t, "127.0.0.1/32"))
|
||||
handler := mw.RateLimit(1)(okHandler())
|
||||
|
||||
post := func(client string) int {
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequestWithContext(t.Context(),
|
||||
http.MethodPost, "/api/v1/reports", http.NoBody)
|
||||
req.RemoteAddr = loopbackPeer
|
||||
req.Header.Set("X-Forwarded-For", client)
|
||||
handler.ServeHTTP(rec, req)
|
||||
|
||||
return rec.Code
|
||||
}
|
||||
|
||||
if code := post(forwardedIP); code != http.StatusOK {
|
||||
t.Fatalf("first request: status = %d, want %d", code, http.StatusOK)
|
||||
}
|
||||
|
||||
if code := post(forwardedIP); code != http.StatusTooManyRequests {
|
||||
t.Fatalf("same client again: status = %d, want %d",
|
||||
code, http.StatusTooManyRequests)
|
||||
}
|
||||
|
||||
if code := post(otherClient); code != http.StatusOK {
|
||||
t.Fatalf("other client behind the same proxy: status = %d, want %d",
|
||||
code, http.StatusOK)
|
||||
}
|
||||
}
|
||||
|
||||
// preflight sends cors the preflight request a browser makes before
|
||||
// it POSTs JSON from origin.
|
||||
func preflight(
|
||||
t *testing.T,
|
||||
cors func(http.Handler) http.Handler,
|
||||
origin string,
|
||||
) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequestWithContext(t.Context(),
|
||||
http.MethodOptions, "/api/v1/reports", http.NoBody)
|
||||
req.Header.Set("Origin", origin)
|
||||
req.Header.Set("Access-Control-Request-Method", http.MethodPost)
|
||||
req.Header.Set("Access-Control-Request-Headers", "content-type")
|
||||
cors(okHandler()).ServeHTTP(rec, req)
|
||||
|
||||
return rec
|
||||
}
|
||||
|
||||
// TestCORSWithoutOriginsAddsNoHeaders checks the default: with no
|
||||
// origins configured, no origin is given any CORS header.
|
||||
func TestCORSWithoutOriginsAddsNoHeaders(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
rec := preflight(t,
|
||||
(&middleware.Middleware{}).CORS(nil), "https://elsewhere.example")
|
||||
|
||||
for name := range rec.Header() {
|
||||
if strings.HasPrefix(name, "Access-Control-") {
|
||||
t.Errorf("CORS header %s set with no origins configured", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCORSAllowsOnlyListedOrigins(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const listed = "https://netwatch.example"
|
||||
|
||||
cors := (&middleware.Middleware{}).CORS([]string{listed})
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
origin string
|
||||
want string
|
||||
}{
|
||||
{name: "listed origin allowed", origin: listed, want: listed},
|
||||
{name: "other origin refused", origin: "https://elsewhere.example"},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
rec := preflight(t, cors, tc.origin)
|
||||
|
||||
got := rec.Header().Get("Access-Control-Allow-Origin")
|
||||
if got != tc.want {
|
||||
t.Errorf("Access-Control-Allow-Origin = %q, want %q",
|
||||
got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user