// Package middleware provides HTTP middleware functions. package middleware import ( "context" "crypto/rand" "log/slog" "net/http" "net/netip" "regexp" "time" basicauth "github.com/99designs/basicauth-go" "github.com/go-chi/chi/v5/middleware" "github.com/go-chi/cors" "github.com/go-chi/httprate" metrics "github.com/slok/go-http-metrics/metrics/prometheus" ghmm "github.com/slok/go-http-metrics/middleware" "github.com/slok/go-http-metrics/middleware/std" "go.uber.org/fx" "sneak.berlin/go/pixa/internal/clientip" "sneak.berlin/go/pixa/internal/config" "sneak.berlin/go/pixa/internal/logger" ) // CORSMaxAgeSeconds is the max age for CORS preflight cache (24 hours). const CORSMaxAgeSeconds = 86400 // HSTSValue is the Strict-Transport-Security header value: one year with // includeSubDomains. Emitted unconditionally even though pixa listens plain // HTTP behind a TLS-terminating proxy; browsers ignore an HSTS header received // over plaintext (RFC 6797 section 8.1), so it never lies about the connection, // and emitting it here avoids trusting a forwarded-proto header. const HSTSValue = "max-age=31536000; includeSubDomains" // ContentSecurityPolicyValue is the Content-Security-Policy header value. // default-src 'self' is the baseline and frame-ancestors 'none' is the primary // clickjacking control. const ContentSecurityPolicyValue = "default-src 'self'; " + "script-src 'self'; " + "style-src 'self'; " + "object-src 'none'; " + "base-uri 'self'; " + "form-action 'self'; " + "frame-ancestors 'none'" // PermissionsPolicyValue is the Permissions-Policy header value. Every listed // feature is denied because pixa uses none of them. const PermissionsPolicyValue = "accelerometer=(), autoplay=(), camera=(), " + "display-capture=(), geolocation=(), gyroscope=(), magnetometer=(), " + "microphone=(), payment=(), usb=()" // Params defines dependencies for Middleware. type Params struct { fx.In Logger *logger.Logger Config *config.Config } // Middleware provides HTTP middleware functions. type Middleware struct { log *slog.Logger config *config.Config clientIP *clientip.Resolver } // New creates a new Middleware instance. func New(_ fx.Lifecycle, params Params) (*Middleware, error) { s := &Middleware{ log: params.Logger.Get(), config: params.Config, clientIP: clientip.NewResolver(params.Config.TrustedProxies), } return s, nil } // ClientIP returns a middleware that resolves the real client IP, // honoring X-Forwarded-For only from trusted proxies, and stores it in // the request context for the logging middleware and handlers to read. func (s *Middleware) ClientIP() func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { ip := s.clientIP.Resolve( r.RemoteAddr, r.Header.Values(clientip.ForwardedForHeader)) ctx := clientip.WithClientIP(r.Context(), ip) next.ServeHTTP(w, r.WithContext(ctx)) }) } } // RateLimit returns a middleware that limits each client to requestLimit // requests per window and refuses a request over the limit with 429 Too Many // Requests and a Retry-After header. Clients are told apart by the address // the ClientIP middleware stored in the request context, so ClientIP must // run first. An IPv6 client is counted by its /64, which one client usually // holds whole; an IPv4-mapped address (::ffff:a.b.c.d) is counted as the // IPv4 address it carries, since every such address falls in the same /64. // Counts are kept only for the current and the previous window, so memory // stays bounded. func (s *Middleware) RateLimit( requestLimit int, window time.Duration, ) func(http.Handler) http.Handler { return httprate.LimitBy(requestLimit, window, func(r *http.Request) (string, error) { ip := clientip.FromContext(r.Context()) addr, err := netip.ParseAddr(ip) if err == nil { ip = addr.Unmap().String() } return httprate.CanonicalizeIP(ip), nil }) } // requestIDPattern is what a request's own X-Request-Id must look like to be // kept as its ID: 1 to 64 letters, digits, '-', '_' or '.'. var requestIDPattern = regexp.MustCompile(`^[A-Za-z0-9._-]{1,64}$`) // RequestID returns a middleware that gives each request an ID and sends it as // the X-Request-Id response header, so a client can quote it when reporting a // problem. The ID is the request's own X-Request-Id when that matches // requestIDPattern, and otherwise a random one, which tells nothing about the // machine or the traffic. It is stored in the request context under chi's // RequestIDKey, where the logging middleware, the handlers and the upstream // fetch read it. func (s *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 := r.Header.Get(middleware.RequestIDHeader) if !requestIDPattern.MatchString(id) { id = rand.Text() } w.Header().Set(middleware.RequestIDHeader, id) ctx := context.WithValue(r.Context(), middleware.RequestIDKey, id) next.ServeHTTP(w, r.WithContext(ctx)) }) } } type loggingResponseWriter struct { http.ResponseWriter statusCode int bytesWritten int64 } func newLoggingResponseWriter(w http.ResponseWriter) *loggingResponseWriter { return &loggingResponseWriter{ResponseWriter: w, statusCode: http.StatusOK} } func (lrw *loggingResponseWriter) WriteHeader(code int) { lrw.statusCode = code lrw.ResponseWriter.WriteHeader(code) } func (lrw *loggingResponseWriter) Write(b []byte) (int, error) { n, err := lrw.ResponseWriter.Write(b) lrw.bytesWritten += int64(n) return n, err } // Logging returns a logging middleware. 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() lrw := newLoggingResponseWriter(w) ctx := r.Context() defer func() { latency := time.Since(start) reqID, _ := ctx.Value(middleware.RequestIDKey).(string) s.log.InfoContext(ctx, "request", "request_start", start, "method", r.Method, "url", r.URL.String(), "useragent", r.UserAgent(), "request_id", reqID, "referer", r.Referer(), "proto", r.Proto, "remoteIP", clientip.FromContext(ctx), "status", lrw.statusCode, "response_bytes", lrw.bytesWritten, "latency_ms", latency.Milliseconds(), ) }() next.ServeHTTP(lrw, r) }) } } // CORS returns a CORS middleware. func (s *Middleware) CORS() func(http.Handler) http.Handler { return cors.Handler(cors.Options{ AllowedOrigins: []string{s.config.AccessControlAllowOrigin}, AllowedMethods: []string{"GET", "HEAD", "OPTIONS"}, AllowedHeaders: []string{"Accept", "Authorization", "Content-Type"}, ExposedHeaders: []string{"Link"}, AllowCredentials: false, MaxAge: CORSMaxAgeSeconds, }) } // Metrics returns a Prometheus metrics middleware. func (s *Middleware) Metrics() func(http.Handler) http.Handler { mdlw := ghmm.New(ghmm.Config{ Recorder: metrics.NewRecorder(metrics.Config{}), }) return func(next http.Handler) http.Handler { return std.Handler("", mdlw, next) } } // MetricsAuth returns a basic auth middleware for the metrics endpoint. func (s *Middleware) MetricsAuth() func(http.Handler) http.Handler { return basicauth.New( "metrics", map[string][]string{ s.config.MetricsUsername: { s.config.MetricsPassword, }, }, ) } // SecurityHeaders returns a middleware that adds security headers to responses. // These headers help protect against common web vulnerabilities. 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) { // Prevent MIME type sniffing w.Header().Set("X-Content-Type-Options", "nosniff") // Prevent clickjacking w.Header().Set("X-Frame-Options", "DENY") // Control referrer information w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin") // Disable XSS filtering (modern browsers don't need it, can cause issues) w.Header().Set("X-XSS-Protection", "0") // Force HTTPS on future visits (ignored by browsers over plaintext) w.Header().Set("Strict-Transport-Security", HSTSValue) // Restrict content sources; frame-ancestors is the primary // clickjacking control, X-Frame-Options the legacy fallback w.Header().Set("Content-Security-Policy", ContentSecurityPolicyValue) // Deny browser features pixa does not use w.Header().Set("Permissions-Policy", PermissionsPolicyValue) next.ServeHTTP(w, r) }) } }