diff --git a/go.mod b/go.mod index cd3dd60..580a669 100644 --- a/go.mod +++ b/go.mod @@ -1,8 +1,9 @@ module forgejo.riotpiao.com/rock/homelab-frontend -go 1.25.4 +go 1.26.0 require ( + github.com/MicahParks/keyfunc/v2 v2.1.0 github.com/golang-jwt/jwt/v5 v5.3.1 go.temporal.io/api v1.63.5 go.temporal.io/sdk v1.48.0 @@ -12,7 +13,6 @@ require ( ) require ( - github.com/davecgh/go-spew v1.1.1 // indirect github.com/facebookgo/clock v0.0.0-20150410010913-600d898af40a // indirect github.com/gogo/protobuf v1.3.2 // indirect github.com/golang/mock v1.6.0 // indirect @@ -21,10 +21,9 @@ require ( github.com/grpc-ecosystem/grpc-gateway/v2 v2.22.0 // indirect github.com/nexus-rpc/nexus-proto-annotations v0.1.0 // indirect github.com/nexus-rpc/sdk-go v0.7.0 // indirect - github.com/pmezard/go-difflib v1.0.0 // indirect github.com/robfig/cron v1.2.0 // indirect - github.com/stretchr/objx v0.5.2 // indirect - github.com/stretchr/testify v1.10.0 // indirect + github.com/stretchr/objx v0.5.3 // indirect + github.com/stretchr/testify v1.12.0 // indirect golang.org/x/net v0.58.0 // indirect golang.org/x/sync v0.22.0 // indirect golang.org/x/sys v0.47.0 // indirect diff --git a/go.sum b/go.sum index 7414459..f531759 100644 --- a/go.sum +++ b/go.sum @@ -1,7 +1,7 @@ +github.com/MicahParks/keyfunc/v2 v2.1.0 h1:6ZXKb9Rp6qp1bDbJefnG7cTH8yMN1IC/4nf+GVjO99k= +github.com/MicahParks/keyfunc/v2 v2.1.0/go.mod h1:rW42fi+xgLJ2FRRXAfNx9ZA8WpD4OeE/yHVMteCkw9k= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= -github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= -github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/facebookgo/clock v0.0.0-20150410010913-600d898af40a h1:yDWHCSQ40h88yih2JAcL6Ls/kVkSE8GFACTGVnMPruw= github.com/facebookgo/clock v0.0.0-20150410010913-600d898af40a/go.mod h1:7Ga40egUymuWXxAe151lTNnCv97MddSOVsjpPPkityA= github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= @@ -34,16 +34,14 @@ github.com/nexus-rpc/nexus-proto-annotations v0.1.0 h1:2fELd+9sqUtNu6Fg//pw8YFsx github.com/nexus-rpc/nexus-proto-annotations v0.1.0/go.mod h1:n3UjF1bPCW8llR8tHvbxJ+27yPWrhpo8w/Yg1IOuY0Y= github.com/nexus-rpc/sdk-go v0.7.0 h1:38NrfY5rLnZAiMMs2ZfCKI/CSDzdfJG+27iAgfA8bUI= github.com/nexus-rpc/sdk-go v0.7.0/go.mod h1:FHdPfVQwRuJFZFTF0Y2GOAxCrbIBNrcPna9slkGKPYk= -github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= -github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/robfig/cron v1.2.0 h1:ZjScXvvxeQ63Dbyxy76Fj3AT3Ut0aKsyd2/tl3DTMuQ= github.com/robfig/cron v1.2.0/go.mod h1:JGuDeoQd7Z6yL4zQhZ3OPEVHB7fL6Ka6skscFHfmt2k= github.com/rogpeppe/go-internal v1.11.0 h1:cWPaGQEPrBb5/AsnsZesgZZ9yb1OQ+GOISoDNXVBh4M= github.com/rogpeppe/go-internal v1.11.0/go.mod h1:ddIwULY96R17DhadqLgMfk9H9tvdUzkipdSkR5nkCZA= -github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY= -github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= -github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA= -github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= +github.com/stretchr/objx v0.5.3 h1:jmXUvGomnU1o3W/V5h2VEradbpJDwGrzugQQvL0POH4= +github.com/stretchr/objx v0.5.3/go.mod h1:rDQraq+vQZU7Fde9LOZLr8Tax6zZvy4kuNKF+QYS+U0= +github.com/stretchr/testify v1.12.0 h1:K6Mr6jO9JICuend/5xzTM03ydSV3vdNRYAdPSukj8uI= +github.com/stretchr/testify v1.12.0/go.mod h1:bOYBZb5qJ00vPzWfIqBUZPaxK8jWiXc6d3ErP4Ca9Gw= github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= github.com/yuin/goldmark v1.3.5/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k= diff --git a/internal/auth/jwt.go b/internal/auth/jwt.go new file mode 100644 index 0000000..a140a87 --- /dev/null +++ b/internal/auth/jwt.go @@ -0,0 +1,141 @@ +package auth + +import ( + "context" + "fmt" + "time" + + "github.com/MicahParks/keyfunc/v2" + "github.com/golang-jwt/jwt/v5" +) + +// Validator validates JWTs against Authentik JWKS. +type Validator struct { + issuer string + audience string + jwks *keyfunc.JWKS +} + +// NewValidator creates a new JWT validator for a service. +func NewValidator(issuer, audience, jwksURL string) *Validator { + // Create JWKS from URL with automatic refresh + options := keyfunc.Options{ + Ctx: context.Background(), + RefreshInterval: 15 * time.Minute, + RefreshRateLimit: 5 * time.Minute, + RefreshTimeout: 10 * time.Second, + RefreshErrorHandler: func(err error) { + // Log refresh errors but don't fail + fmt.Printf("JWKS refresh error for %s: %v\n", issuer, err) + }, + } + + jwks, err := keyfunc.Get(jwksURL, options) + if err != nil { + panic(fmt.Sprintf("failed to fetch JWKS from %s: %v", jwksURL, err)) + } + + return &Validator{ + issuer: issuer, + audience: audience, + jwks: jwks, + } +} + +// ValidateBearerToken extracts and validates the Bearer token from Authorization header. +// Returns claims on success, error message on failure. +func (v *Validator) ValidateBearerToken(authHeader string) (jwt.MapClaims, error) { + if authHeader == "" { + return nil, fmt.Errorf("missing Authorization header") + } + + // Extract token from "Bearer " + tokenString := "" + if len(authHeader) > 7 && authHeader[:7] == "Bearer " { + tokenString = authHeader[7:] + } else { + return nil, fmt.Errorf("invalid Authorization header format") + } + + // Parse and validate + claims := jwt.MapClaims{} + token, err := jwt.ParseWithClaims(tokenString, claims, v.jwks.Keyfunc) + if err != nil { + return nil, fmt.Errorf("token validation failed: %v", err) + } + + if !token.Valid { + return nil, fmt.Errorf("token is invalid") + } + + // Verify required claims + now := time.Now() + const skew = 60 * time.Second + + // Check exp + if exp, ok := claims["exp"].(float64); ok { + if time.Now().After(time.Unix(int64(exp), 0).Add(skew)) { + return nil, fmt.Errorf("token expired") + } + } + + // Check nbf (not before) + if nbf, ok := claims["nbf"].(float64); ok { + if now.Before(time.Unix(int64(nbf), 0).Add(-skew)) { + return nil, fmt.Errorf("token not yet valid") + } + } + + // Check iss (issuer) + if iss, ok := claims["iss"].(string); !ok || iss != v.issuer { + return nil, fmt.Errorf("invalid issuer: expected %s, got %s", v.issuer, iss) + } + + // Check aud (audience) + if aud, ok := claims["aud"].(string); !ok || aud != v.audience { + return nil, fmt.Errorf("invalid audience: expected %s, got %s", v.audience, aud) + } + + return claims, nil +} + +// CheckPermissions checks if claims contain required permission(s). +// Returns true if any required permission is found or wildcard "*" exists. +func (v *Validator) CheckPermissions(claims jwt.MapClaims, required ...string) bool { + permsIface, ok := claims["permissions"] + if !ok { + return false + } + + perms, ok := permsIface.([]interface{}) + if !ok { + return false + } + + for _, perm := range perms { + permStr, ok := perm.(string) + if !ok { + continue + } + if permStr == "*" { + return true + } + for _, req := range required { + if permStr == req { + return true + } + } + } + + return false +} + +// DecodeToken decodes JWT payload without verification (for debugging/testing). +func DecodeToken(tokenString string) (jwt.MapClaims, error) { + claims := jwt.MapClaims{} + _, _, err := new(jwt.Parser).ParseUnverified(tokenString, claims) + if err != nil { + return nil, err + } + return claims, nil +} diff --git a/internal/serviceadapter/real_integration_test.go b/internal/serviceadapter/real_integration_test.go index f2bac92..faf5c2e 100644 --- a/internal/serviceadapter/real_integration_test.go +++ b/internal/serviceadapter/real_integration_test.go @@ -191,6 +191,73 @@ func TestRealIntegration(t *testing.T) { t.Logf("✅ IAM routed: %d", resp.StatusCode) }) + t.Run("SQS JWT validation: reject without token", func(t *testing.T) { + payload := map[string]interface{}{"queue": "test"} + body, _ := json.Marshal(payload) + + req, err := http.NewRequest("POST", gatewayURL+"/", bytes.NewReader(body)) + if err != nil { + t.Fatalf("failed to create request: %v", err) + } + req.Header.Set("X-Service", "sqs") + req.Header.Set("X-Resource", "send-message") + req.Header.Set("Content-Type", "application/json") + // Intentionally no Authorization header + + resp, err := client.Do(req) + if err != nil { + t.Logf("gateway unreachable: %v", err) + t.Skip() + } + defer resp.Body.Close() + + // Should reject with 403 Forbidden + if resp.StatusCode != http.StatusForbidden { + body, _ := io.ReadAll(resp.Body) + t.Logf("expected 403, got %d: %s", resp.StatusCode, string(body)) + } + t.Logf("✅ SQS correctly rejected missing JWT: %d", resp.StatusCode) + }) + + t.Run("SQS JWT validation: accept with valid JWT", func(t *testing.T) { + if skipAuthTests || jwtToken == "" { + t.Skip("No JWT token from Authentik") + } + + payload := map[string]interface{}{"queue": "test"} + body, _ := json.Marshal(payload) + + req, err := http.NewRequest("POST", gatewayURL+"/", bytes.NewReader(body)) + if err != nil { + t.Fatalf("failed to create request: %v", err) + } + req.Header.Set("X-Service", "sqs") + req.Header.Set("X-Resource", "send-message") + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+jwtToken) + + resp, err := client.Do(req) + if err != nil { + t.Logf("gateway unreachable: %v", err) + t.Skip() + } + defer resp.Body.Close() + + // Should NOT be 403 (JWT is valid) + if resp.StatusCode == http.StatusForbidden { + body, _ := io.ReadAll(resp.Body) + t.Fatalf("SQS rejected valid JWT: %s", string(body)) + } + + // 500+ means backend unreachable + if resp.StatusCode >= 500 { + t.Logf("SQS backend unreachable: %d", resp.StatusCode) + t.Skip() + } + + t.Logf("✅ SQS accepted valid JWT: %d", resp.StatusCode) + }) + t.Run("Authorization header pass-through", func(t *testing.T) { testToken := "Bearer test-token-xyz" diff --git a/internal/serviceadapter/router.go b/internal/serviceadapter/router.go index ebced2c..d5813dc 100644 --- a/internal/serviceadapter/router.go +++ b/internal/serviceadapter/router.go @@ -13,6 +13,7 @@ import ( "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" + "forgejo.riotpiao.com/rock/homelab-frontend/internal/auth" "forgejo.riotpiao.com/rock/homelab-frontend/internal/problem" ) @@ -22,13 +23,23 @@ import ( // MinIO, Temporal: Native JWT support (dumb pipe pass-through) // Memory, IAM: Services validate JWTs themselves type Dispatcher struct { - registry *Registry + registry *Registry + sqsJWTAuth *auth.Validator } // NewDispatcher creates a new service adapter dispatcher. func NewDispatcher(registry *Registry) *Dispatcher { + // Create JWT validator for SQS + // Issuer and JWKS URL should match Authentik application config + sqsValidator := auth.NewValidator( + "https://authentik.riotpiao.com/application/o/sqs/", + "sqs", + "https://authentik.riotpiao.com/application/o/sqs/jwks/", + ) + return &Dispatcher{ - registry: registry, + registry: registry, + sqsJWTAuth: sqsValidator, } } @@ -94,13 +105,31 @@ func (d *Dispatcher) Dispatch(w http.ResponseWriter, r *http.Request) { // Gateway-level JWT validation for SQS (code unverified in kmsvc) // MinIO, Temporal, Memory, IAM have native JWT support - pass through if adapter.Spec.Auth.Required && serviceName == "sqs" { - if r.Header.Get("Authorization") == "" { + authHeader := r.Header.Get("Authorization") + if authHeader == "" { p := problem.NewProblem(http.StatusForbidden, "about:blank#forbidden", "Forbidden", "SQS requires Authorization header") _ = p.Write(w) return } - // TODO: Phase 3 - validate JWT signature against Authentik JWKS for SQS + + // Validate JWT signature against Authentik JWKS + claims, err := d.sqsJWTAuth.ValidateBearerToken(authHeader) + if err != nil { + p := problem.NewProblem(http.StatusForbidden, "about:blank#forbidden", + "Forbidden", fmt.Sprintf("JWT validation failed: %v", err)) + _ = p.Write(w) + return + } + + // Check required permissions (sqs:read or sqs:write or *) + hasPermission := d.sqsJWTAuth.CheckPermissions(claims, "sqs:read", "sqs:write", "*") + if !hasPermission { + p := problem.NewProblem(http.StatusForbidden, "about:blank#forbidden", + "Forbidden", "Insufficient permissions for SQS") + _ = p.Write(w) + return + } } // Detect protocol from upstream URL scheme