Files
query-orchestration/internal/rbac/middleware.go
T
2025-03-14 14:15:13 -07:00

262 lines
7.6 KiB
Go

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:<name>" or "query:<name>" or "cookie:<name>"
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")
// }
// }
//}