// Package middleware provides HTTP middleware for logging, // CORS, and other cross-cutting concerns. package middleware import ( "errors" "fmt" "io" "log/slog" "net" "net/http" "net/netip" "runtime/debug" "strings" "time" "sneak.berlin/go/netwatch/internal/config" "sneak.berlin/go/netwatch/internal/globals" "sneak.berlin/go/netwatch/internal/logger" "github.com/go-chi/chi/v5/middleware" "github.com/go-chi/cors" "github.com/go-chi/httprate" "go.uber.org/fx" ) const corsMaxAgeSec = 300 // jsonErrorBody is the body written for errors raised inside // middleware, matching the {"status":"error"} shape the handlers // return so clients see one error contract across the API. const ( jsonContentType = "application/json; charset=utf-8" jsonErrorBody = "{\"status\":\"error\"}\n" ) // Security header values. The backend is a JSON API with no // HTML surface, so the CSP forbids every resource type and // framing outright. const ( hstsValue = "max-age=31536000; includeSubDomains" cspValue = "default-src 'none'; frame-ancestors 'none'" permissionsPolicyValue = "camera=(), microphone=(), geolocation=()" ) // Params defines the dependencies for Middleware. type Params struct { fx.In Config *config.Config Globals *globals.Globals Logger *logger.Logger } // Middleware holds shared state for middleware factories. type Middleware struct { log *slog.Logger params *Params trustedProxies []netip.Prefix } // New creates a Middleware instance. func New( _ fx.Lifecycle, params Params, ) (*Middleware, error) { trusted, err := parseTrustedProxies(params.Config.TrustedProxies) if err != nil { return nil, err } s := new(Middleware) s.params = ¶ms s.log = params.Logger.Get() s.trustedProxies = trusted return s, nil } // parseTrustedProxies converts the TRUSTED_PROXIES entries into // prefixes, failing fast on any malformed entry. func parseTrustedProxies(cidrs []string) ([]netip.Prefix, error) { prefixes := make([]netip.Prefix, 0, len(cidrs)) for _, cidr := range cidrs { prefix, err := netip.ParsePrefix(cidr) if err != nil { return nil, fmt.Errorf( "TRUSTED_PROXIES %q: %w", cidr, err, ) } prefixes = append(prefixes, prefix.Masked()) } return prefixes, nil } type loggingResponseWriter struct { http.ResponseWriter statusCode int } func newLoggingResponseWriter( w http.ResponseWriter, ) *loggingResponseWriter { return &loggingResponseWriter{w, http.StatusOK} } func (lrw *loggingResponseWriter) WriteHeader(code int) { lrw.statusCode = code lrw.ResponseWriter.WriteHeader(code) } func ipFromHostPort(hostPort string) string { host, _, err := net.SplitHostPort(hostPort) if err != nil { return hostPort } return host } // clientIP resolves the caller's address. X-Forwarded-For and // X-Real-IP are honoured only when the direct peer is a // trusted proxy; otherwise the direct peer is returned so a // spoofed header cannot forge the logged address. func clientIP( remoteAddr string, header http.Header, trusted []netip.Prefix, ) string { peer := ipFromHostPort(remoteAddr) if !addrInAny(peer, trusted) { return peer } if xff := firstForwardedFor(header.Get("X-Forwarded-For")); xff != "" { return xff } if xr := strings.TrimSpace(header.Get("X-Real-IP")); validIP(xr) { return xr } return peer } // firstForwardedFor returns the left-most valid address in an // X-Forwarded-For list (the original client), or "" if none. func firstForwardedFor(value string) string { for part := range strings.SplitSeq(value, ",") { candidate := strings.TrimSpace(part) if validIP(candidate) { return candidate } } return "" } func validIP(s string) bool { _, err := netip.ParseAddr(s) return err == nil } // addrInAny reports whether s parses as an address contained // in any of the trusted prefixes. func addrInAny(s string, trusted []netip.Prefix) bool { addr, err := netip.ParseAddr(s) if err != nil { return false } addr = addr.Unmap() for _, prefix := range trusted { if prefix.Contains(addr) { return true } } return false } // Logging returns middleware that logs each request with // timing, status code, and client information. func (s *Middleware) Logging() func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc( func(w http.ResponseWriter, r *http.Request) { start := time.Now().UTC() lrw := newLoggingResponseWriter(w) ctx := r.Context() defer func() { latency := time.Since(start) s.log.InfoContext(ctx, "request", "request_start", start, "method", r.Method, "url", r.URL.String(), "useragent", r.UserAgent(), "request_id", ctx.Value( middleware.RequestIDKey, ), "referer", r.Referer(), "proto", r.Proto, "remote_ip", clientIP( r.RemoteAddr, r.Header, s.trustedProxies, ), "status", lrw.statusCode, "latency_ms", latency.Milliseconds(), ) }() next.ServeHTTP(lrw, r) }, ) } } // SecurityHeaders returns middleware that sets response // security headers. It runs before CORS so the headers are // present on preflight responses the CORS handler writes. func (s *Middleware) SecurityHeaders() func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc( func(w http.ResponseWriter, r *http.Request) { h := w.Header() h.Set("Strict-Transport-Security", hstsValue) h.Set("Content-Security-Policy", cspValue) h.Set("X-Frame-Options", "DENY") h.Set("X-Content-Type-Options", "nosniff") h.Set("Referrer-Policy", "no-referrer") h.Set("Permissions-Policy", permissionsPolicyValue) next.ServeHTTP(w, r) }, ) } } // writeJSONError writes the shared JSON error body with the // given status. Used where middleware must reject a request // before it reaches a handler. func writeJSONError(w http.ResponseWriter, status int) { w.Header().Set("Content-Type", jsonContentType) w.WriteHeader(status) _, _ = io.WriteString(w, jsonErrorBody) } // MaxBodyBytes returns middleware that caps the request body at // limit bytes. A declared Content-Length over the limit is // rejected immediately with 413. Bodies without a declared // length (or that understate it) are capped as they are read, so // a handler that reads the body sees a *http.MaxBytesError it can // map to 413. Mounted again on a route group, it can only lower // the limit: a cap applied earlier in the chain still holds. func (s *Middleware) MaxBodyBytes( limit int64, ) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc( func(w http.ResponseWriter, r *http.Request) { if r.ContentLength > limit { writeJSONError( w, http.StatusRequestEntityTooLarge, ) return } r.Body = http.MaxBytesReader(w, r.Body, limit) next.ServeHTTP(w, r) }, ) } } // Recoverer returns middleware that recovers from a panic in a // downstream handler, logs the panic and stack trace through // slog, and responds 500 with no body. http.ErrAbortHandler is // re-panicked so the server can abort the response as intended. func (s *Middleware) Recoverer() func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc( func(w http.ResponseWriter, r *http.Request) { defer func() { rec := recover() if rec == nil { return } err, ok := rec.(error) if ok && errors.Is(err, http.ErrAbortHandler) { panic(rec) } s.log.ErrorContext(r.Context(), "panic recovered", "panic", fmt.Sprintf("%v", rec), "stack", string(debug.Stack()), ) w.WriteHeader(http.StatusInternalServerError) }() next.ServeHTTP(w, r) }, ) } } // 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: 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) }, ), ) }