package middleware import ( "context" "log/slog" "net/http" "strings" "time" "git.arcline.it/ArclineIT/nexus/internal/auth" "git.arcline.it/ArclineIT/nexus/internal/config" "github.com/google/uuid" ) type contextKey string const ( UserIDKey contextKey = "user_id" UserEmailKey contextKey = "user_email" RequestIDKey contextKey = "request_id" ) // RequestID injects a unique ID into every request for tracing. func RequestID(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { id := r.Header.Get("X-Request-ID") if id == "" { id = uuid.New().String() } ctx := context.WithValue(r.Context(), RequestIDKey, id) w.Header().Set("X-Request-ID", id) next.ServeHTTP(w, r.WithContext(ctx)) }) } // Logger logs every HTTP request with structured fields. func Logger(logger *slog.Logger) 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() wrapped := &responseWriter{ResponseWriter: w, statusCode: http.StatusOK} next.ServeHTTP(wrapped, r) logger.Info("http request", slog.String("method", r.Method), slog.String("path", r.URL.Path), slog.Int("status", wrapped.statusCode), slog.Duration("duration", time.Since(start)), slog.String("remote_addr", r.RemoteAddr), ) }) } } // Recoverer catches panics and returns a 500 response. func Recoverer(logger *slog.Logger) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { defer func() { if rec := recover(); rec != nil { logger.Error("panic recovered", slog.Any("panic", rec), slog.String("path", r.URL.Path), ) http.Error(w, `{"error":"internal server error"}`, http.StatusInternalServerError) } }() next.ServeHTTP(w, r) }) } } // CORS sets permissive CORS headers for development. func CORS(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Access-Control-Allow-Origin", "*") w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS") w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type, X-Request-ID") if r.Method == http.MethodOptions { w.WriteHeader(http.StatusNoContent) return } next.ServeHTTP(w, r) }) } // Authenticate validates the JWT Bearer token and injects user info into context. func Authenticate(cfg *config.Config) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { token := extractBearerToken(r) if token == "" { http.Error(w, `{"error":"missing authorization header"}`, http.StatusUnauthorized) return } claims, err := auth.ValidateToken(cfg, token) if err != nil { http.Error(w, `{"error":"invalid or expired token"}`, http.StatusUnauthorized) return } ctx := r.Context() ctx = context.WithValue(ctx, UserIDKey, claims.Subject) ctx = context.WithValue(ctx, UserEmailKey, claims.Email) next.ServeHTTP(w, r.WithContext(ctx)) }) } } // WebAuth validates the JWT from a cookie or Bearer header and injects user // info into context. If the token is missing or invalid, it redirects to // /login instead of returning a JSON error — suitable for browser-based flows. func WebAuth(cfg *config.Config) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { token := "" // Check cookie first, then fall back to Bearer header if cookie, err := r.Cookie("nexus_access_token"); err == nil && cookie.Value != "" { token = cookie.Value } else { token = extractBearerToken(r) } if token == "" { http.Redirect(w, r, "/login", http.StatusSeeOther) return } claims, err := auth.ValidateToken(cfg, token) if err != nil { // Clear invalid cookies clearAuthCookies(w) http.Redirect(w, r, "/login", http.StatusSeeOther) return } ctx := r.Context() ctx = context.WithValue(ctx, UserIDKey, claims.Subject) ctx = context.WithValue(ctx, UserEmailKey, claims.Email) next.ServeHTTP(w, r.WithContext(ctx)) }) } } // extractBearerToken pulls a Bearer token from the Authorization header. func extractBearerToken(r *http.Request) string { header := r.Header.Get("Authorization") if header == "" { return "" } parts := strings.SplitN(header, " ", 2) if len(parts) != 2 || !strings.EqualFold(parts[0], "bearer") { return "" } return parts[1] } // clearAuthCookies removes Nexus auth cookies from the response. func clearAuthCookies(w http.ResponseWriter) { http.SetCookie(w, &http.Cookie{ Name: "nexus_access_token", Value: "", Path: "/", MaxAge: -1, }) http.SetCookie(w, &http.Cookie{ Name: "nexus_refresh_token", Value: "", Path: "/", MaxAge: -1, }) } // responseWriter wraps http.ResponseWriter to capture the status code. type responseWriter struct { http.ResponseWriter statusCode int } func (rw *responseWriter) WriteHeader(code int) { rw.statusCode = code rw.ResponseWriter.WriteHeader(code) }