package mw import ( "context" "net/http" "strings" "crussell/auth" ) type contextKey string const ( UserIDKey contextKey = "user_id" UserRoleKey contextKey = "user_role" JTIKey contextKey = "jti" ) // RequireAuth middleware - validates JWT and adds user info to context func RequireAuth(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { authHeader := r.Header.Get("Authorization") if authHeader == "" || !strings.HasPrefix(authHeader, "Bearer ") { http.Error(w, "missing or invalid authorization header", http.StatusUnauthorized) return } tokenString := strings.TrimPrefix(authHeader, "Bearer ") userID, role, jti, err := auth.VerifyToken(tokenString, r.Context()) if err != nil { http.Error(w, "invalid token", http.StatusUnauthorized) return } // Add user info and JTI to context ctx := context.WithValue(r.Context(), UserIDKey, userID) ctx = context.WithValue(ctx, UserRoleKey, role) ctx = context.WithValue(ctx, JTIKey, jti) next.ServeHTTP(w, r.WithContext(ctx)) }) } // OptionalAuth middleware - extracts user info if token present, otherwise passes through func OptionalAuth(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { authHeader := r.Header.Get("Authorization") if authHeader != "" && strings.HasPrefix(authHeader, "Bearer ") { tokenString := strings.TrimPrefix(authHeader, "Bearer ") userID, role, jti, err := auth.VerifyToken(tokenString, r.Context()) if err == nil { ctx := context.WithValue(r.Context(), UserIDKey, userID) ctx = context.WithValue(ctx, UserRoleKey, role) ctx = context.WithValue(ctx, JTIKey, jti) next.ServeHTTP(w, r.WithContext(ctx)) return } } next.ServeHTTP(w, r) }) } // RequireRole middleware - checks if user has required role(s) func RequireRole(allowedRoles ...string) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { role, ok := r.Context().Value(UserRoleKey).(string) if !ok { http.Error(w, "unauthorized", http.StatusUnauthorized) return } // Check if user has one of the allowed roles hasRole := false for _, allowedRole := range allowedRoles { if role == allowedRole { hasRole = true break } } if !hasRole { http.Error(w, "forbidden", http.StatusForbidden) return } next.ServeHTTP(w, r) }) } } // RequireVerified middleware - only allows verified_email and admin func RequireVerified(next http.Handler) http.Handler { return RequireRole("verified_email", "admin")(next) } // RequireAdmin middleware - only allows admin func RequireAdmin(next http.Handler) http.Handler { return RequireRole("admin")(next) } // Helper functions to get user info from context func GetUserID(ctx context.Context) (string, bool) { userID, ok := ctx.Value(UserIDKey).(string) return userID, ok } func GetUserRole(ctx context.Context) (string, bool) { role, ok := ctx.Value(UserRoleKey).(string) return role, ok } func GetJTI(ctx context.Context) (string, bool) { jti, ok := ctx.Value(JTIKey).(string) return jti, ok }