// Package http provides reusable net/http middleware shared by every service // that exposes an HTTP surface. It holds no domain types — handlers belong in // the owning service, not here. package http import ( "context" "encoding/json" "fmt" "log/slog" "net/http" "os" "strconv" "strings" "time" ) // contextKey is unexported so values stashed by this package cannot collide // with values stashed by any other package using the same string. type contextKey string const ( userIDContextKey contextKey = "user_id" usernameContextKey contextKey = "username" ) var ( defaultAllowedMethods = []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"} defaultAllowedHeaders = []string{"Content-Type", "Authorization", "X-Requested-With", "X-API-Key"} ) // CORSConfig describes which cross-origin requests are permitted. // // AllowedOrigins is an explicit allowlist matched exactly against the request // Origin. The single entry "*" permits any origin, but only for unauthenticated // requests: "*" together with AllowCredentials is rejected by CORS, because it // would let any site on the internet issue credentialed requests and read the // responses. An empty AllowedOrigins denies every cross-origin request, which // leaves same-origin traffic working normally. type CORSConfig struct { AllowedOrigins []string AllowedMethods []string AllowedHeaders []string ExposedHeaders []string AllowCredentials bool MaxAge time.Duration } // LoadCORSConfigFromEnv reads CORS settings from the environment. // // CORS_ALLOWED_ORIGINS is a comma-separated allowlist and defaults to empty, // so an unconfigured service denies cross-origin requests rather than allowing // them. CORS_ALLOW_CREDENTIALS defaults to false. func LoadCORSConfigFromEnv() CORSConfig { return CORSConfig{ AllowedOrigins: splitAndTrim(getEnv("CORS_ALLOWED_ORIGINS", "")), AllowedMethods: splitAndTrim(getEnv("CORS_ALLOWED_METHODS", strings.Join(defaultAllowedMethods, ","))), AllowedHeaders: splitAndTrim(getEnv("CORS_ALLOWED_HEADERS", strings.Join(defaultAllowedHeaders, ","))), ExposedHeaders: splitAndTrim(getEnv("CORS_EXPOSED_HEADERS", "Content-Length,Content-Range")), AllowCredentials: getEnv("CORS_ALLOW_CREDENTIALS", "false") == "true", MaxAge: 24 * time.Hour, } } // CORS returns a middleware enforcing cfg. // // It returns an error for configurations that cannot be served safely, so a // service fails at startup instead of running with permissive CORS. func CORS(cfg CORSConfig) (func(http.Handler) http.Handler, error) { allowAny := len(cfg.AllowedOrigins) == 1 && cfg.AllowedOrigins[0] == "*" if allowAny && cfg.AllowCredentials { return nil, fmt.Errorf(`CORS: AllowedOrigins "*" cannot be combined with AllowCredentials; list the permitted origins explicitly`) } for _, origin := range cfg.AllowedOrigins { if origin != "*" && !strings.Contains(origin, "://") { return nil, fmt.Errorf("CORS: AllowedOrigins entry %q must include a scheme, e.g. https://%s", origin, origin) } } // Origin comparison is case-insensitive on scheme and host, so normalise the // allowlist once here rather than on every request. allowed := make(map[string]struct{}, len(cfg.AllowedOrigins)) for _, origin := range cfg.AllowedOrigins { allowed[strings.ToLower(origin)] = struct{}{} } // Defaults are applied here, not only in LoadCORSConfigFromEnv, so a // directly constructed CORSConfig still answers preflights usefully instead // of sending empty header values. if len(cfg.AllowedMethods) == 0 { cfg.AllowedMethods = defaultAllowedMethods } if len(cfg.AllowedHeaders) == 0 { cfg.AllowedHeaders = defaultAllowedHeaders } methods := strings.Join(cfg.AllowedMethods, ", ") headers := strings.Join(cfg.AllowedHeaders, ", ") exposed := strings.Join(cfg.ExposedHeaders, ", ") maxAge := "" if cfg.MaxAge > 0 { maxAge = strconv.Itoa(int(cfg.MaxAge.Seconds())) } if len(cfg.AllowedOrigins) == 0 { slog.Warn("CORS: no allowed origins configured, cross-origin requests will be denied (set CORS_ALLOWED_ORIGINS)") } return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { origin := r.Header.Get("Origin") isPreflight := r.Method == http.MethodOptions && r.Header.Get("Access-Control-Request-Method") != "" // Responses vary by Origin whenever the allowlist is consulted, so // caches must not serve one origin's response to another. if !allowAny { w.Header().Add("Vary", "Origin") } if origin != "" && originAllowed(allowed, allowAny, origin) { if allowAny { w.Header().Set("Access-Control-Allow-Origin", "*") } else { w.Header().Set("Access-Control-Allow-Origin", origin) } if cfg.AllowCredentials { w.Header().Set("Access-Control-Allow-Credentials", "true") } if exposed != "" { w.Header().Set("Access-Control-Expose-Headers", exposed) } if isPreflight { w.Header().Set("Access-Control-Allow-Methods", methods) w.Header().Set("Access-Control-Allow-Headers", headers) if maxAge != "" { w.Header().Set("Access-Control-Max-Age", maxAge) } } } // Preflights are answered here whether or not the origin passed: with // no Allow-Origin header the browser blocks the real request anyway, // and handlers never have to deal with OPTIONS. if isPreflight { w.WriteHeader(http.StatusNoContent) return } next.ServeHTTP(w, r) }) }, nil } func originAllowed(allowed map[string]struct{}, allowAny bool, origin string) bool { if allowAny { return true } _, ok := allowed[strings.ToLower(origin)] return ok } // Logs logs each request with its method, path, status and duration. func Logs(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { start := time.Now() recorder := &statusRecorder{ResponseWriter: w, status: http.StatusOK} next.ServeHTTP(recorder, r) slog.Info("http request", "method", r.Method, "path", r.URL.Path, "status", recorder.status, "duration", time.Since(start), ) }) } // statusRecorder captures the response status for logging. type statusRecorder struct { http.ResponseWriter status int wroteHeader bool } func (s *statusRecorder) WriteHeader(status int) { if s.wroteHeader { return } s.status = status s.wroteHeader = true s.ResponseWriter.WriteHeader(status) } // AuthConfig supplies the credential checks used by Auth. Every validator is // injected by the calling service — this package deliberately ships no // credentials of its own. A nil validator disables that scheme. type AuthConfig struct { // ValidateBearerToken resolves a Bearer token to a user ID. ValidateBearerToken func(token string) (userID string, err error) // ValidateAPIKey resolves an X-API-Key value to a user ID. ValidateAPIKey func(apiKey string) (userID string, ok bool) // ValidateBasicAuth reports whether HTTP basic credentials are valid. ValidateBasicAuth func(username, password string) bool // SkipPrefixes are path prefixes served without authentication, e.g. // "/swagger/" or "/health". SkipPrefixes []string } // Auth authenticates requests via Bearer token, API key or basic auth, in that // order, and rejects anything else with 401. func Auth(cfg AuthConfig) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { for _, prefix := range cfg.SkipPrefixes { if strings.HasPrefix(r.URL.Path, prefix) { next.ServeHTTP(w, r) return } } if cfg.ValidateBearerToken != nil { if token, ok := bearerToken(r); ok { if userID, err := cfg.ValidateBearerToken(token); err == nil { next.ServeHTTP(w, r.WithContext( context.WithValue(r.Context(), userIDContextKey, userID))) return } } } if cfg.ValidateAPIKey != nil { if apiKey := r.Header.Get("X-API-Key"); apiKey != "" { if userID, ok := cfg.ValidateAPIKey(apiKey); ok { next.ServeHTTP(w, r.WithContext( context.WithValue(r.Context(), userIDContextKey, userID))) return } } } if cfg.ValidateBasicAuth != nil { if username, password, ok := r.BasicAuth(); ok { if cfg.ValidateBasicAuth(username, password) { next.ServeHTTP(w, r.WithContext( context.WithValue(r.Context(), usernameContextKey, username))) return } } } slog.Warn("unauthorized request", "method", r.Method, "path", r.URL.Path) w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusUnauthorized) json.NewEncoder(w).Encode(map[string]any{ "error": "Unauthorized", "message": "Valid authentication required. Use Bearer token, API key, or basic auth.", "code": http.StatusUnauthorized, }) }) } } func bearerToken(r *http.Request) (string, bool) { header := r.Header.Get("Authorization") if !strings.HasPrefix(header, "Bearer ") { return "", false } token := strings.TrimPrefix(header, "Bearer ") return token, token != "" } // UserIDFromContext returns the user ID stored by Auth, if any. func UserIDFromContext(ctx context.Context) (string, bool) { userID, ok := ctx.Value(userIDContextKey).(string) return userID, ok } // UsernameFromContext returns the username stored by Auth, if any. func UsernameFromContext(ctx context.Context) (string, bool) { username, ok := ctx.Value(usernameContextKey).(string) return username, ok } func getEnv(key, defaultValue string) string { if value := os.Getenv(key); value != "" { return value } return defaultValue } func splitAndTrim(value string) []string { if value == "" { return nil } parts := strings.Split(value, ",") result := make([]string, 0, len(parts)) for _, part := range parts { if trimmed := strings.TrimSpace(part); trimmed != "" { result = append(result, trimmed) } } return result }