// Package middleware provides HTTP middleware for logging, // CORS, and other cross-cutting concerns. package middleware import ( "fmt" "log/slog" "net" "net/http" "net/netip" "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" "go.uber.org/fx" ) const corsMaxAgeSec = 300 // 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 CIDR strings 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 proxy %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) }, ) } } // CORS returns middleware that adds permissive CORS headers. func (s *Middleware) CORS() func(http.Handler) http.Handler { 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"}, AllowCredentials: false, MaxAge: corsMaxAgeSec, }) }