This is an automated email from the ASF dual-hosted git repository. lahirujayathilake pushed a commit to branch auth-endpoints in repository https://gitbox.apache.org/repos/asf/airavata-custos.git
commit c3cf8e2f1379f157972bfad987c7bdd696a5550c Author: lahiruj <[email protected]> AuthorDate: Tue Jun 16 19:11:11 2026 -0400 Add JWT verification middleware with JWKS discovery Co-Authored-By: Claude Opus 4.7 <[email protected]> --- go.mod | 14 +- go.sum | 21 ++ internal/server/middleware/auth.go | 283 +++++++++++++++++++++++++ internal/server/middleware/auth_test.go | 363 ++++++++++++++++++++++++++++++++ 4 files changed, 679 insertions(+), 2 deletions(-) diff --git a/go.mod b/go.mod index 9da28975c..4b43efc5d 100644 --- a/go.mod +++ b/go.mod @@ -7,13 +7,14 @@ require ( github.com/golang-migrate/migrate/v4 v4.19.1 github.com/google/uuid v1.6.0 github.com/jmoiron/sqlx v1.4.0 + github.com/lestrrat-go/jwx/v2 v2.1.3 github.com/prometheus/client_golang v1.23.2 github.com/prometheus/client_model v0.6.2 github.com/stretchr/testify v1.11.1 - github.com/swaggo/swag v1.16.6 go.opentelemetry.io/otel v1.41.0 go.opentelemetry.io/otel/sdk v1.41.0 go.opentelemetry.io/otel/trace v1.41.0 + golang.org/x/sync v0.18.0 google.golang.org/protobuf v1.36.8 gopkg.in/yaml.v3 v3.0.1 ) @@ -25,27 +26,36 @@ require ( github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/cpuguy83/go-md2man/v2 v2.0.0-20190314233015-f79a8a8ca69d // indirect github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect + github.com/decred/dcrd/dcrec/secp256k1/v4 v4.3.0 // indirect github.com/go-logr/logr v1.4.3 // indirect github.com/go-logr/stdr v1.2.2 // indirect github.com/go-openapi/jsonpointer v0.21.0 // indirect github.com/go-openapi/jsonreference v0.21.0 // indirect github.com/go-openapi/spec v0.20.4 // indirect github.com/go-openapi/swag v0.23.0 // indirect + github.com/goccy/go-json v0.10.3 // indirect github.com/josharian/intern v1.0.0 // indirect + github.com/lestrrat-go/blackmagic v1.0.2 // indirect + github.com/lestrrat-go/httpcc v1.0.1 // indirect + github.com/lestrrat-go/httprc v1.0.6 // indirect + github.com/lestrrat-go/iter v1.0.2 // indirect + github.com/lestrrat-go/option v1.0.1 // indirect github.com/mailru/easyjson v0.7.7 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 // indirect github.com/prometheus/common v0.66.1 // indirect github.com/prometheus/procfs v0.16.1 // indirect github.com/russross/blackfriday/v2 v2.0.1 // indirect + github.com/segmentio/asm v1.2.0 // indirect github.com/shurcooL/sanitized_anchor_name v1.0.0 // indirect github.com/stretchr/objx v0.5.2 // indirect + github.com/swaggo/swag v1.16.6 // indirect github.com/urfave/cli/v2 v2.3.0 // indirect go.opentelemetry.io/auto/sdk v1.2.1 // indirect go.opentelemetry.io/otel/metric v1.41.0 // indirect go.yaml.in/yaml/v2 v2.4.2 // indirect + golang.org/x/crypto v0.45.0 // indirect golang.org/x/mod v0.29.0 // indirect - golang.org/x/sync v0.18.0 // indirect golang.org/x/sys v0.41.0 // indirect golang.org/x/text v0.31.0 // indirect golang.org/x/tools v0.38.0 // indirect diff --git a/go.sum b/go.sum index fc3195913..c926a0d46 100644 --- a/go.sum +++ b/go.sum @@ -25,6 +25,8 @@ github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSs github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM= github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/decred/dcrd/dcrec/secp256k1/v4 v4.3.0 h1:rpfIENRNNilwHwZeG5+P150SMrnNEcHYvcCuK6dPZSg= +github.com/decred/dcrd/dcrec/secp256k1/v4 v4.3.0/go.mod h1:v57UDF4pDQJcEfFUCRop3lJL149eHGSe9Jvczhzjo/0= github.com/dhui/dktest v0.4.6 h1:+DPKyScKSEp3VLtbMDHcUq6V5Lm5zfZZVb0Sk7Ahom4= github.com/dhui/dktest v0.4.6/go.mod h1:JHTSYDtKkvFNFHJKqCzVzqXecyv+tKt8EzceOmQOgbU= github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk= @@ -57,6 +59,8 @@ github.com/go-openapi/swag v0.23.0 h1:vsEVJDUo2hPJ2tu0/Xc+4noaxyEffXNIs3cOULZ+Gr github.com/go-openapi/swag v0.23.0/go.mod h1:esZ8ITTYEsH1V2trKHjAN8Ai7xHb8RV+YSZ577vPjgQ= github.com/go-sql-driver/mysql v1.8.1 h1:LedoTUt/eveggdHS9qUFC1EFSa8bU2+1pZjSRpvNJ1Y= github.com/go-sql-driver/mysql v1.8.1/go.mod h1:wEBSXgmK//2ZFJyE+qWnIsVGmvmEKlqwuVSjsCm7DZg= +github.com/goccy/go-json v0.10.3 h1:KZ5WoDbxAIgm2HNbYckL0se1fHD6rz5j4ywS6ebzDqA= +github.com/goccy/go-json v0.10.3/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= github.com/gogo/protobuf v1.3.2 h1:Ov1cvc58UF3b5XjBnZv7+opcTcQFZebYjWzi34vdm4Q= github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q= github.com/golang-migrate/migrate/v4 v4.19.1 h1:OCyb44lFuQfYXYLx1SCxPZQGU7mcaZ7gH9yH4jSFbBA= @@ -76,6 +80,18 @@ github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/lestrrat-go/blackmagic v1.0.2 h1:Cg2gVSc9h7sz9NOByczrbUvLopQmXrfFx//N+AkAr5k= +github.com/lestrrat-go/blackmagic v1.0.2/go.mod h1:UrEqBzIR2U6CnzVyUtfM6oZNMt/7O7Vohk2J0OGSAtU= +github.com/lestrrat-go/httpcc v1.0.1 h1:ydWCStUeJLkpYyjLDHihupbn2tYmZ7m22BGkcvZZrIE= +github.com/lestrrat-go/httpcc v1.0.1/go.mod h1:qiltp3Mt56+55GPVCbTdM9MlqhvzyuL6W/NMDA8vA5E= +github.com/lestrrat-go/httprc v1.0.6 h1:qgmgIRhpvBqexMJjA/PmwSvhNk679oqD1RbovdCGW8k= +github.com/lestrrat-go/httprc v1.0.6/go.mod h1:mwwz3JMTPBjHUkkDv/IGJ39aALInZLrhBp0X7KGUZlo= +github.com/lestrrat-go/iter v1.0.2 h1:gMXo1q4c2pHmC3dn8LzRhJfP1ceCbgSiT9lUydIzltI= +github.com/lestrrat-go/iter v1.0.2/go.mod h1:Momfcq3AnRlRjI5b5O8/G5/BvpzrhoFTZcn06fEOPt4= +github.com/lestrrat-go/jwx/v2 v2.1.3 h1:Ud4lb2QuxRClYAmRleF50KrbKIoM1TddXgBrneT5/Jo= +github.com/lestrrat-go/jwx/v2 v2.1.3/go.mod h1:q6uFgbgZfEmQrfJfrCo90QcQOcXFMfbI/fO0NqRtvZo= +github.com/lestrrat-go/option v1.0.1 h1:oAzP2fvZGQKWkvHa1/SAcFolBEca1oN+mQ7eooNBEYU= +github.com/lestrrat-go/option v1.0.1/go.mod h1:5ZHFbivi4xwXxhxY9XHDe2FHo6/Z7WWmtT7T5nBBp3I= github.com/lib/pq v1.10.9 h1:YXG7RB+JIjhP29X+OtkiDnYaXQwpS4JEWq7dtCCRUEw= github.com/lib/pq v1.10.9/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/mailru/easyjson v0.0.0-20190614124828-94de47d64c63/go.mod h1:C1wdFJiN94OJF2b5HbByQZoLdCWB1Yqtg26g4irojpc= @@ -115,6 +131,8 @@ github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0t github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/russross/blackfriday/v2 v2.0.1 h1:lPqVAte+HuHNfhJ/0LC98ESWRz8afy9tM/0RK8m9o+Q= github.com/russross/blackfriday/v2 v2.0.1/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= +github.com/segmentio/asm v1.2.0 h1:9BQrFxC+YOHJlTlHGkTrFWf59nbL3XnCoFLTwDCI7ys= +github.com/segmentio/asm v1.2.0/go.mod h1:BqMnlJP91P8d+4ibuonYZw9mfnzI9HfxselHZr5aAcs= github.com/shurcooL/sanitized_anchor_name v1.0.0 h1:PdmoCO6wvbs+7yrJyMORt4/BmY5IYyJwS/kOiWx8mHo= github.com/shurcooL/sanitized_anchor_name v1.0.0/go.mod h1:1NzhyTcUVG4SuEtjjoZeVRXNmyL/1OwPU0+IJeTBvfc= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= @@ -122,6 +140,7 @@ 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.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/swaggo/swag v1.16.6 h1:qBNcx53ZaX+M5dxVyTrgQ0PJ/ACK+NzhwcbieTt+9yI= @@ -146,6 +165,8 @@ go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.yaml.in/yaml/v2 v2.4.2 h1:DzmwEr2rDGHl7lsFgAHxmNz/1NlQ7xLIrlN2h5d1eGI= go.yaml.in/yaml/v2 v2.4.2/go.mod h1:081UH+NErpNdqlCXm3TtEran0rJZGxAYx9hb/ELlsPU= +golang.org/x/crypto v0.45.0 h1:jMBrvKuj23MTlT0bQEOBcAE0mjg8mK9RXFhRH6nyF3Q= +golang.org/x/crypto v0.45.0/go.mod h1:XTGrrkGJve7CYK7J8PEww4aY7gM3qMCElcJQ8n8JdX4= golang.org/x/mod v0.29.0 h1:HV8lRxZC4l2cr3Zq1LvtOsi/ThTgWnUk/y64QSs8GwA= golang.org/x/mod v0.29.0/go.mod h1:NyhrlYXJ2H4eJiRy/WDBO6HMqZQ6q9nk4JzS3NuCK+w= golang.org/x/net v0.0.0-20210421230115-4e50805a0758/go.mod h1:72T/g9IO56b78aLF+1Kcs5dz7/ng1VjMUvfKvpfy+jM= diff --git a/internal/server/middleware/auth.go b/internal/server/middleware/auth.go new file mode 100644 index 000000000..59dd1599f --- /dev/null +++ b/internal/server/middleware/auth.go @@ -0,0 +1,283 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +// Package middleware carries the HTTP-layer cross-cutting concerns: bearer +// token verification, CORS, and any future per-request seam that runs before +// the application handlers in internal/server. Each middleware here is +// constructed once at boot and wraps the server's mux. +package middleware + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "strings" + "sync" + "time" + + "github.com/lestrrat-go/jwx/v2/jwk" + "github.com/lestrrat-go/jwx/v2/jwt" + "golang.org/x/sync/singleflight" + + "github.com/apache/airavata-custos/internal/config" + "github.com/apache/airavata-custos/pkg/identity" +) + +// UserResolver looks up a Custos user by verified OIDC subject. The auth +// middleware calls this after signature + claim validation; a nil result +// means the sub is valid but unknown to Custos, which the middleware turns +// into a 401. +type UserResolver func(ctx context.Context, oidcSub string) (userID string, err error) + +// Auth verifies OIDC bearer tokens against the configured issuer's JWKS, +// resolves the verified `sub` to a Custos user, and attaches an +// [identity.Caller] to the request context. +type Auth struct { + issuer string + audience string + jwksOverride string + resolveUser UserResolver + skipPrefixes []string + + httpClient *http.Client + jwksCacheTTL time.Duration + + mu sync.RWMutex + jwksURL string + cache *jwksEntry + sfGroup singleflight.Group +} + +type jwksEntry struct { + set jwk.Set + fetchedAt time.Time +} + +// Option tunes an Auth middleware at construction time. +type Option func(*Auth) + +// WithJWKSURL pins the JWKS URL, bypassing discovery via the issuer's +// .well-known endpoint. Used by integration tests to point at an +// in-process JWKS server. +func WithJWKSURL(u string) Option { + return func(a *Auth) { a.jwksOverride = u; a.jwksURL = u } +} + +// WithSkipPrefixes lets routes such as /healthz bypass authentication. +// Matching is by exact prefix on the request path. +func WithSkipPrefixes(prefixes ...string) Option { + return func(a *Auth) { a.skipPrefixes = append(a.skipPrefixes, prefixes...) } +} + +// NewAuth constructs the middleware. issuer and audience must be set; the +// config block's JWKSURL is treated as an override (equivalent to passing +// [WithJWKSURL]). resolveUser is mandatory. +func NewAuth(cfg config.AuthConfig, resolveUser UserResolver, opts ...Option) (*Auth, error) { + if strings.TrimSpace(cfg.Issuer) == "" { + return nil, errors.New("auth: issuer is required") + } + if strings.TrimSpace(cfg.Audience) == "" { + return nil, errors.New("auth: audience is required") + } + if resolveUser == nil { + return nil, errors.New("auth: user resolver is required") + } + a := &Auth{ + issuer: cfg.Issuer, + audience: cfg.Audience, + jwksOverride: cfg.JWKSURL, + jwksURL: cfg.JWKSURL, + resolveUser: resolveUser, + httpClient: &http.Client{Timeout: 10 * time.Second}, + jwksCacheTTL: 10 * time.Minute, + } + for _, opt := range opts { + opt(a) + } + return a, nil +} + +// Wrap returns a handler that rejects unauthenticated requests and forwards +// authenticated ones with an [identity.Caller] in the context. +func (a *Auth) Wrap(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if a.shouldSkip(r.URL.Path) { + next.ServeHTTP(w, r) + return + } + + token := bearerFromHeader(r.Header.Get("Authorization")) + if token == "" { + writeAuthError(w, http.StatusUnauthorized, "missing bearer token") + return + } + + caller, err := a.verify(r.Context(), token) + if err != nil { + writeAuthError(w, http.StatusUnauthorized, err.Error()) + return + } + + ctx := identity.WithCaller(r.Context(), caller) + next.ServeHTTP(w, r.WithContext(ctx)) + }) +} + +func (a *Auth) shouldSkip(path string) bool { + for _, p := range a.skipPrefixes { + if strings.HasPrefix(path, p) { + return true + } + } + return false +} + +// verify validates the token signature, issuer, audience, and expiry, then +// resolves the OIDC sub to a Custos user. Any failure collapses to a single +// 401 message for the client; the specific reason is intentionally generic +// so callers cannot distinguish "bad signature" from "unknown sub". +func (a *Auth) verify(ctx context.Context, tokenString string) (*identity.Caller, error) { + keys, err := a.getJWKS(ctx) + if err != nil { + return nil, fmt.Errorf("token verification unavailable") + } + + tok, err := jwt.Parse([]byte(tokenString), + jwt.WithKeySet(keys), + jwt.WithValidate(true), + jwt.WithIssuer(a.issuer), + jwt.WithAudience(a.audience), + ) + if err != nil { + return nil, fmt.Errorf("invalid token") + } + + sub := tok.Subject() + if sub == "" { + return nil, fmt.Errorf("invalid token") + } + + userID, err := a.resolveUser(ctx, sub) + if err != nil { + return nil, fmt.Errorf("identity lookup failed") + } + if userID == "" { + return nil, fmt.Errorf("unknown caller") + } + + email, _ := tok.Get("email") + emailStr, _ := email.(string) + + return &identity.Caller{ + UserID: userID, + OIDCSub: sub, + Email: emailStr, + }, nil +} + +func (a *Auth) getJWKS(ctx context.Context) (jwk.Set, error) { + a.mu.RLock() + entry := a.cache + a.mu.RUnlock() + if entry != nil && time.Since(entry.fetchedAt) < a.jwksCacheTTL { + return entry.set, nil + } + + result, err, _ := a.sfGroup.Do("jwks", func() (interface{}, error) { + a.mu.RLock() + entry2 := a.cache + a.mu.RUnlock() + if entry2 != nil && time.Since(entry2.fetchedAt) < a.jwksCacheTTL { + return entry2.set, nil + } + set, err := a.fetchJWKS(ctx) + if err != nil { + return nil, err + } + a.mu.Lock() + a.cache = &jwksEntry{set: set, fetchedAt: time.Now()} + a.mu.Unlock() + return set, nil + }) + if err != nil { + return nil, err + } + return result.(jwk.Set), nil +} + +// fetchJWKS resolves the JWKS URL (via discovery if not overridden) and +// fetches the key set. Discovery results are cached on the Auth instance so +// subsequent fetches skip the well-known round trip. +func (a *Auth) fetchJWKS(ctx context.Context) (jwk.Set, error) { + ctx, cancel := context.WithTimeout(ctx, a.httpClient.Timeout) + defer cancel() + + if a.jwksURL == "" { + url, err := a.discoverJWKSURL(ctx) + if err != nil { + return nil, err + } + a.mu.Lock() + a.jwksURL = url + a.mu.Unlock() + } + return jwk.Fetch(ctx, a.jwksURL, jwk.WithHTTPClient(a.httpClient)) +} + +func (a *Auth) discoverJWKSURL(ctx context.Context) (string, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, strings.TrimRight(a.issuer, "/")+"/.well-known/openid-configuration", nil) + if err != nil { + return "", fmt.Errorf("auth: build discovery request: %w", err) + } + resp, err := a.httpClient.Do(req) + if err != nil { + return "", fmt.Errorf("auth: discovery request: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("auth: discovery returned %d", resp.StatusCode) + } + var doc struct { + JWKSURI string `json:"jwks_uri"` + } + if err := json.NewDecoder(resp.Body).Decode(&doc); err != nil { + return "", fmt.Errorf("auth: decode discovery: %w", err) + } + if doc.JWKSURI == "" { + return "", fmt.Errorf("auth: discovery missing jwks_uri") + } + return doc.JWKSURI, nil +} + +func bearerFromHeader(h string) string { + if h == "" { + return "" + } + const prefix = "Bearer " + if len(h) <= len(prefix) || !strings.EqualFold(h[:len(prefix)], prefix) { + return "" + } + return strings.TrimSpace(h[len(prefix):]) +} + +func writeAuthError(w http.ResponseWriter, status int, msg string) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(map[string]string{"error": msg}) +} diff --git a/internal/server/middleware/auth_test.go b/internal/server/middleware/auth_test.go new file mode 100644 index 000000000..993f7db5b --- /dev/null +++ b/internal/server/middleware/auth_test.go @@ -0,0 +1,363 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package middleware_test + +import ( + "context" + "crypto/rand" + "crypto/rsa" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/lestrrat-go/jwx/v2/jwa" + "github.com/lestrrat-go/jwx/v2/jwk" + "github.com/lestrrat-go/jwx/v2/jwt" + + "github.com/apache/airavata-custos/internal/config" + "github.com/apache/airavata-custos/internal/server/middleware" + "github.com/apache/airavata-custos/pkg/identity" +) + +const ( + testIssuer = "https://issuer.test" + testAudience = "test-audience" +) + +// signer holds a fresh RSA key pair and the JWKS server backing it. Tests +// instantiate one per case so cache state never leaks between cases. +type signer struct { + privateKey jwk.Key + publicSet jwk.Set + jwksServer *httptest.Server +} + +func newSigner(t *testing.T) *signer { + t.Helper() + raw, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("generate rsa: %v", err) + } + priv, err := jwk.FromRaw(raw) + if err != nil { + t.Fatalf("priv from raw: %v", err) + } + if err := priv.Set(jwk.KeyIDKey, "test-key"); err != nil { + t.Fatalf("set kid: %v", err) + } + if err := priv.Set(jwk.AlgorithmKey, jwa.RS256); err != nil { + t.Fatalf("set alg: %v", err) + } + + pub, err := priv.PublicKey() + if err != nil { + t.Fatalf("public from priv: %v", err) + } + set := jwk.NewSet() + if err := set.AddKey(pub); err != nil { + t.Fatalf("add public key: %v", err) + } + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(set) + })) + t.Cleanup(srv.Close) + + return &signer{privateKey: priv, publicSet: set, jwksServer: srv} +} + +type tokenOpts struct { + issuer string + audience string + subject string + email string + expiresAt time.Time +} + +func (s *signer) mint(t *testing.T, o tokenOpts) string { + t.Helper() + builder := jwt.NewBuilder(). + Issuer(o.issuer). + Audience([]string{o.audience}). + Subject(o.subject). + Expiration(o.expiresAt). + IssuedAt(time.Now()) + if o.email != "" { + builder = builder.Claim("email", o.email) + } + tok, err := builder.Build() + if err != nil { + t.Fatalf("build token: %v", err) + } + signed, err := jwt.Sign(tok, jwt.WithKey(jwa.RS256, s.privateKey)) + if err != nil { + t.Fatalf("sign token: %v", err) + } + return string(signed) +} + +func resolverAlways(userID string) middleware.UserResolver { + return func(_ context.Context, _ string) (string, error) { return userID, nil } +} + +func resolverErr(err error) middleware.UserResolver { + return func(_ context.Context, _ string) (string, error) { return "", err } +} + +func newAuth(t *testing.T, s *signer, resolver middleware.UserResolver, opts ...middleware.Option) *middleware.Auth { + t.Helper() + opts = append(opts, middleware.WithJWKSURL(s.jwksServer.URL)) + a, err := middleware.NewAuth(config.AuthConfig{Issuer: testIssuer, Audience: testAudience}, resolver, opts...) + if err != nil { + t.Fatalf("new auth: %v", err) + } + return a +} + +func newRequest(token string) *http.Request { + r := httptest.NewRequest(http.MethodGet, "/projects", nil) + if token != "" { + r.Header.Set("Authorization", "Bearer "+token) + } + return r +} + +func successHandler(captured **identity.Caller) http.HandlerFunc { + return func(_ http.ResponseWriter, r *http.Request) { + *captured = identity.CallerFromContext(r.Context()) + } +} + +func TestNewAuth_Validation(t *testing.T) { + cases := []struct { + name string + cfg config.AuthConfig + resolver middleware.UserResolver + }{ + {"missing issuer", config.AuthConfig{Audience: "a"}, resolverAlways("u")}, + {"missing audience", config.AuthConfig{Issuer: "i"}, resolverAlways("u")}, + {"missing resolver", config.AuthConfig{Issuer: "i", Audience: "a"}, nil}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if _, err := middleware.NewAuth(tc.cfg, tc.resolver); err == nil { + t.Fatalf("expected error") + } + }) + } +} + +func TestAuth_AcceptsValidToken(t *testing.T) { + s := newSigner(t) + a := newAuth(t, s, resolverAlways("user-1")) + + token := s.mint(t, tokenOpts{ + issuer: testIssuer, + audience: testAudience, + subject: "sub-1", + email: "[email protected]", + expiresAt: time.Now().Add(time.Minute), + }) + + var captured *identity.Caller + a.Wrap(successHandler(&captured)).ServeHTTP(httptest.NewRecorder(), newRequest(token)) + + if captured == nil { + t.Fatalf("caller not set on context") + } + if captured.UserID != "user-1" || captured.OIDCSub != "sub-1" || captured.Email != "[email protected]" { + t.Fatalf("unexpected caller %+v", captured) + } +} + +func TestAuth_RejectsRequests(t *testing.T) { + s := newSigner(t) + good := tokenOpts{issuer: testIssuer, audience: testAudience, subject: "sub-1", expiresAt: time.Now().Add(time.Minute)} + + cases := []struct { + name string + token func() string + resolver middleware.UserResolver + header string // overrides the Authorization header when set + }{ + { + name: "no bearer", + token: func() string { return "" }, + resolver: resolverAlways("user-1"), + }, + { + name: "non-bearer scheme", + token: func() string { return "" }, + resolver: resolverAlways("user-1"), + header: "Basic abc", + }, + { + name: "malformed jwt", + token: func() string { return "not-a-jwt" }, + resolver: resolverAlways("user-1"), + }, + { + name: "expired", + token: func() string { + bad := good + bad.expiresAt = time.Now().Add(-time.Minute) + return s.mint(t, bad) + }, + resolver: resolverAlways("user-1"), + }, + { + name: "wrong issuer", + token: func() string { + bad := good + bad.issuer = "https://other.test" + return s.mint(t, bad) + }, + resolver: resolverAlways("user-1"), + }, + { + name: "wrong audience", + token: func() string { + bad := good + bad.audience = "other-audience" + return s.mint(t, bad) + }, + resolver: resolverAlways("user-1"), + }, + { + name: "unknown sub", + token: func() string { return s.mint(t, good) }, + resolver: resolverAlways(""), + }, + { + name: "resolver error", + token: func() string { return s.mint(t, good) }, + resolver: resolverErr(errors.New("db down")), + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + a := newAuth(t, s, tc.resolver) + req := newRequest(tc.token()) + if tc.header != "" { + req.Header.Set("Authorization", tc.header) + } + rec := httptest.NewRecorder() + handlerCalled := false + a.Wrap(http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { + handlerCalled = true + })).ServeHTTP(rec, req) + + if rec.Code != http.StatusUnauthorized { + t.Fatalf("status = %d, want 401", rec.Code) + } + if handlerCalled { + t.Fatalf("handler should not run for rejected requests") + } + }) + } +} + +func TestAuth_RejectsTokenSignedByForeignKey(t *testing.T) { + good := newSigner(t) + foreign := newSigner(t) + + a := newAuth(t, good, resolverAlways("user-1")) + token := foreign.mint(t, tokenOpts{ + issuer: testIssuer, + audience: testAudience, + subject: "sub-1", + expiresAt: time.Now().Add(time.Minute), + }) + + rec := httptest.NewRecorder() + a.Wrap(http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { + t.Fatal("handler should not run") + })).ServeHTTP(rec, newRequest(token)) + + if rec.Code != http.StatusUnauthorized { + t.Fatalf("status = %d, want 401", rec.Code) + } +} + +func TestAuth_SkipPrefixBypassesVerification(t *testing.T) { + s := newSigner(t) + a := newAuth(t, s, resolverAlways("user-1"), middleware.WithSkipPrefixes("/healthz")) + + called := false + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/healthz", nil) + a.Wrap(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { + called = true + if c := identity.CallerFromContext(r.Context()); c != nil { + t.Fatalf("skipped route should not carry a caller, got %+v", c) + } + })).ServeHTTP(rec, req) + + if !called { + t.Fatalf("handler should run for skipped paths") + } + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rec.Code) + } +} + +func TestAuth_DiscoversJWKSURLFromIssuer(t *testing.T) { + s := newSigner(t) + + var issuerSrv *httptest.Server + issuerSrv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/.well-known/openid-configuration" { + http.NotFound(w, r) + return + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]string{ + "issuer": issuerSrv.URL, + "jwks_uri": s.jwksServer.URL, + }) + })) + defer issuerSrv.Close() + + a, err := middleware.NewAuth( + config.AuthConfig{Issuer: issuerSrv.URL, Audience: testAudience}, + resolverAlways("user-1"), + ) + if err != nil { + t.Fatalf("new auth: %v", err) + } + + token := s.mint(t, tokenOpts{ + issuer: issuerSrv.URL, + audience: testAudience, + subject: "sub-1", + expiresAt: time.Now().Add(time.Minute), + }) + + rec := httptest.NewRecorder() + called := false + a.Wrap(http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { called = true })).ServeHTTP(rec, newRequest(token)) + + if !called || rec.Code != http.StatusOK { + t.Fatalf("discovery path failed: code=%d called=%v", rec.Code, called) + } +}
