package rbac import ( "errors" "fmt" "log/slog" "net/http" "strings" "github.com/golang-jwt/jwt/v5" "github.com/labstack/echo/v4" ) // JWTConfig contains configuration for the JWT middleware type JWTConfig struct { KeyProvider KeyProvider TokenLookup string // Format: "header:" or "query:" or "cookie:" AuthScheme string // Usually "Bearer" RequiredGroups []string // Groups that are required for access (any one of these) ContextKey string // Key to store the claims in the context ErrorHandler func(echo.Context, error) error Logger *slog.Logger // Logger instance for middleware } // DefaultJWTConfig is the default JWT auth middleware config var DefaultJWTConfig = JWTConfig{ TokenLookup: "header:Authorization", AuthScheme: "Bearer", ContextKey: "user", ErrorHandler: defaultErrorHandler, } // defaultErrorHandler returns a 401 Unauthorized error for JWT validation failures func defaultErrorHandler(c echo.Context, err error) error { return echo.NewHTTPError(http.StatusUnauthorized, "invalid or expired jwt") } // extractToken determines which token extractor to use based on config func extractToken(config JWTConfig) func(echo.Context) (string, error) { parts := strings.Split(config.TokenLookup, ":") extractor := tokenFromHeader(parts[1], config.AuthScheme) if parts[0] == "query" { extractor = tokenFromQuery(parts[1]) } else if parts[0] == "cookie" { extractor = tokenFromCookie(parts[1]) } return extractor } // validateGroups checks if the user belongs to any of the required groups func validateGroups(claims jwt.MapClaims, requiredGroups []string, logger *slog.Logger) bool { if len(requiredGroups) == 0 { logger.Debug("no required groups specified, access granted") return true } groups, ok := claims["cognito:groups"].([]interface{}) if !ok { logger.Warn("no groups found in claims", "user", claims["sub"]) return false } userGroups := make([]string, 0) for _, group := range groups { if groupStr, ok := group.(string); ok { userGroups = append(userGroups, groupStr) } } for _, reqGroup := range requiredGroups { for _, group := range groups { if groupStr, ok := group.(string); ok && strings.EqualFold(groupStr, reqGroup) { logger.Debug("group validation successful", "user", claims["sub"], "required_group", reqGroup, "user_groups", userGroups) return true } } } logger.Warn("group validation failed", "user", claims["sub"], "required_groups", requiredGroups, "user_groups", userGroups) return false } // parseJWT parses and validates the JWT token func parseJWT(tokenString string, keyProvider KeyProvider, logger *slog.Logger) (*jwt.Token, jwt.MapClaims, error) { token, err := jwt.Parse(tokenString, func(token *jwt.Token) (interface{}, error) { if _, ok := token.Method.(*jwt.SigningMethodRSA); !ok { logger.Error("unexpected signing method", "method", token.Method.Alg()) return nil, errors.New("unexpected signing method") } // Get key ID from token header kid, ok := token.Header["kid"].(string) if !ok { logger.Error("no key ID in token header") return nil, errors.New("no key ID (kid) in token header") } // Get public key from provider return keyProvider.GetPublicKey(kid) }) if err != nil { logger.Error("failed to parse JWT", "error", err) return nil, nil, err } claims, ok := token.Claims.(jwt.MapClaims) if !ok || !token.Valid { logger.Error("invalid token claims") return nil, nil, errors.New("invalid token") } logger.Debug("JWT parsed successfully", "user", claims["sub"], "email", claims["email"]) return token, claims, nil } // JWTWithConfig returns a JWT auth middleware with config func JWTWithConfig(config JWTConfig) echo.MiddlewareFunc { // Set defaults if config.TokenLookup == "" { config.TokenLookup = DefaultJWTConfig.TokenLookup } if config.AuthScheme == "" { config.AuthScheme = DefaultJWTConfig.AuthScheme } if config.ContextKey == "" { config.ContextKey = DefaultJWTConfig.ContextKey } if config.ErrorHandler == nil { config.ErrorHandler = DefaultJWTConfig.ErrorHandler } if config.KeyProvider == nil { panic("JWT middleware requires a KeyProvider") } if config.Logger == nil { config.Logger = slog.Default() } extractor := extractToken(config) return func(next echo.HandlerFunc) echo.HandlerFunc { return func(c echo.Context) error { tokenString, err := extractor(c) if err != nil { config.Logger.Error("failed to extract token", "error", err, "path", c.Request().URL.Path, "method", c.Request().Method) return config.ErrorHandler(c, err) } _, claims, err := parseJWT(tokenString, config.KeyProvider, config.Logger) if err != nil { config.Logger.Error("failed to validate JWT", "error", err, "path", c.Request().URL.Path, "method", c.Request().Method) return config.ErrorHandler(c, err) } // Check required groups if specified if !validateGroups(claims, config.RequiredGroups, config.Logger) { config.Logger.Warn("insufficient permissions", "user", claims["sub"], "path", c.Request().URL.Path, "method", c.Request().Method) return echo.NewHTTPError(http.StatusForbidden, "insufficient permissions") } // Store user information in context c.Set(config.ContextKey, claims) config.Logger.Info("authenticated request", "user", claims["sub"], "path", c.Request().URL.Path, "method", c.Request().Method) return next(c) } } } // JWT returns a JWT auth middleware with default configuration func JWT(keyProvider KeyProvider) echo.MiddlewareFunc { config := DefaultJWTConfig config.KeyProvider = keyProvider return JWTWithConfig(config) } // tokenFromHeader extracts token from Authorization header func tokenFromHeader(header string, authScheme string) func(echo.Context) (string, error) { return func(c echo.Context) (string, error) { auth := c.Request().Header.Get(header) if auth == "" { return "", fmt.Errorf("missing auth header: %s", header) } l := len(authScheme) if len(auth) > l+1 && auth[:l] == authScheme { return auth[l+1:], nil } return "", fmt.Errorf("invalid auth header format for scheme: %s", authScheme) } } // tokenFromQuery extracts token from query parameter func tokenFromQuery(param string) func(echo.Context) (string, error) { return func(c echo.Context) (string, error) { token := c.QueryParam(param) if token == "" { return "", errors.New("missing auth query parameter") } return token, nil } } // tokenFromCookie extracts token from cookie func tokenFromCookie(name string) func(echo.Context) (string, error) { return func(c echo.Context) (string, error) { cookie, err := c.Cookie(name) if err != nil { return "", err } return cookie.Value, nil } } //// GroupGuard middleware ensures a user has at least one of the required groups //func GroupGuard(requiredGroups ...string) echo.MiddlewareFunc { // return func(next echo.HandlerFunc) echo.HandlerFunc { // return func(c echo.Context) error { // userClaims, ok := c.Get("user").(jwt.MapClaims) // if !ok { // return echo.NewHTTPError(http.StatusUnauthorized, "user not authenticated") // } // // groups, ok := userClaims["cognito:groups"].([]interface{}) // if !ok { // return echo.NewHTTPError(http.StatusForbidden, "no groups found for user") // } // // for _, requiredGroup := range requiredGroups { // for _, group := range groups { // if groupStr, ok := group.(string); ok && strings.EqualFold(groupStr, requiredGroup) { // return next(c) // } // } // } // // return echo.NewHTTPError(http.StatusForbidden, "insufficient permissions") // } // } //}