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 3802fd5db9cf414a674f539f3b31ad8b02860d7b
Author: lahiruj <[email protected]>
AuthorDate: Tue Jun 16 19:11:11 2026 -0400

    Add JWT verification middleware with JWKS discovery
---
 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)
+       }
+}

Reply via email to