Files
query-orchestration/internal/rbac/middleware.go
T
2025-03-14 12:36:52 -07:00

213 lines
5.9 KiB
Go

package rbac
import (
"errors"
"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
}
// 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) bool {
if len(requiredGroups) == 0 {
return true
}
groups, ok := claims["cognito:groups"].([]interface{})
if !ok {
return false
}
for _, reqGroup := range requiredGroups {
for _, group := range groups {
if groupStr, ok := group.(string); ok && strings.EqualFold(groupStr, reqGroup) {
return true
}
}
}
return false
}
// parseJWT parses and validates the JWT token
func parseJWT(tokenString string, keyProvider KeyProvider) (*jwt.Token, jwt.MapClaims, error) {
token, err := jwt.Parse(tokenString, func(token *jwt.Token) (interface{}, error) {
if _, ok := token.Method.(*jwt.SigningMethodRSA); !ok {
return nil, errors.New("unexpected signing method")
}
// Get key ID from token header
kid, ok := token.Header["kid"].(string)
if !ok {
return nil, errors.New("no key ID (kid) in token header")
}
// Get public key from provider
return keyProvider.GetPublicKey(kid)
})
if err != nil {
return nil, nil, err
}
claims, ok := token.Claims.(jwt.MapClaims)
if !ok || !token.Valid {
return nil, nil, errors.New("invalid token")
}
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")
}
extractor := extractToken(config)
return func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error {
tokenString, err := extractor(c)
if err != nil {
return config.ErrorHandler(c, err)
}
_, claims, err := parseJWT(tokenString, config.KeyProvider)
if err != nil {
return config.ErrorHandler(c, err)
}
// Check required groups if specified
if !validateGroups(claims, config.RequiredGroups) {
return echo.NewHTTPError(http.StatusForbidden, "insufficient permissions")
}
// Store user information in context
c.Set(config.ContextKey, claims)
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 "", errors.New("missing auth header")
}
l := len(authScheme)
if len(auth) > l+1 && auth[:l] == authScheme {
return auth[l+1:], nil
}
return "", errors.New("invalid auth header")
}
}
// 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")
// }
// }
//}