// Package middleware holds the HTTP middleware chain: request // identity, logging, metrics, panic recovery, timeouts, body caps, // security headers and CSRF. // // Order matters and is fixed in internal/server/routes.go, not here. package middleware import ( "context" "log/slog" "net/http" "strconv" "time" "github.com/go-chi/chi/v5" "github.com/google/uuid" "go.uber.org/fx" "sneak.berlin/go/simplexcalc/internal/config" "sneak.berlin/go/simplexcalc/internal/logger" "sneak.berlin/go/simplexcalc/internal/telemetry" ) // contextKey is this package's private context key type, so no other // package can collide with or read these values by accident. type contextKey string // requestIDKey carries the per-request id. const requestIDKey contextKey = "request-id" // RequestIDHeader is the response header the id is echoed in, so a // user reporting a failure can quote something that finds the log line. const RequestIDHeader = "X-Request-Id" // Params defines dependencies for Middleware. type Params struct { fx.In Config *config.Config Logger *logger.Logger Sentry *telemetry.Sentry Metrics *telemetry.Metrics } // Middleware is the set of handlers, built once and reused. type Middleware struct { params Params log *slog.Logger cfg *config.Config } // New creates the middleware set. // //nolint:revive // lc parameter is required by fx even if unused. func New(lc fx.Lifecycle, params Params) (*Middleware, error) { return &Middleware{ params: params, log: params.Logger.Get(), cfg: params.Config, }, nil } // RequestIDFrom returns the id assigned to r's context, or "" outside a // request that went through RequestID. func RequestIDFrom(ctx context.Context) string { id, _ := ctx.Value(requestIDKey).(string) return id } // RequestID assigns each request an id and echoes it. An id supplied by // the client is ignored: it is attacker-controlled, it would let a // caller collide two unrelated requests in the log, and there is no // trusted proxy contract here that would make it meaningful. func (m *Middleware) RequestID() func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { id := uuid.NewString() w.Header().Set(RequestIDHeader, id) next.ServeHTTP(w, r.WithContext( context.WithValue(r.Context(), requestIDKey, id), )) }) } } // RequestLogger logs one line per completed request. func (m *Middleware) RequestLogger() 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() rec := newResponseRecorder(w) next.ServeHTTP(rec, r) m.log.Info("request", "id", RequestIDFrom(r.Context()), "method", r.Method, "path", r.URL.Path, "route", routePattern(r), "status", rec.Status(), "bytes", rec.written, "duration_ms", time.Since(start).Milliseconds(), ) }) } } // Metrics records the Prometheus series for each request. func (m *Middleware) Metrics() 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() rec := newResponseRecorder(w) m.params.Metrics.InFlightAdd(1) defer m.params.Metrics.InFlightAdd(-1) next.ServeHTTP(rec, r) m.params.Metrics.Observe( r.Method, routePattern(r), strconv.Itoa(rec.Status()), time.Since(start), ) }) } } // Timeout bounds handler execution with the configured request timeout. // The handler sees a context with a deadline; a handler that ignores it // still runs to completion, so handlers must pass the context down to // everything that can block. func (m *Middleware) Timeout() func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { ctx, cancel := context.WithTimeout(r.Context(), m.cfg.RequestTimeout) defer cancel() next.ServeHTTP(w, r.WithContext(ctx)) }) } } // routePattern returns the chi route pattern for r, or "unmatched" when // no route matched (a 404). It is what the metrics and the log are // labelled by; see the comment on the requests counter for why the path // is not. func routePattern(r *http.Request) string { rctx := chi.RouteContext(r.Context()) if rctx == nil { return "unmatched" } pattern := rctx.RoutePattern() if pattern == "" { return "unmatched" } return pattern }