262 lines
7.6 KiB
Go
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")
|
|
// }
|
|
// }
|
|
//}
|