Skip to content

Commit 3a2391c

Browse files
AchoArnoldCopilot
andcommitted
fix(auth): bound token metadata caches
Delegated MCP verification now prefilters bearer tokens on their unverified "iss" claim -- parsed only to discard tokens that cannot be ours and never trusted for authentication -- so Firebase ID tokens and other non-MCP credentials seen by the pre-BearerAuth middleware never reach the JWKS cache. The JWKS cache itself collapses concurrent refreshes into one in-flight fetch, refuses a new fetch until MinRefreshInterval (default one minute) has elapsed, and keeps serving an already known key while throttled, mirroring the MCP Firebase certificate cache. The 2s HTTP timeout, key rotation, and fail-open middleware behavior are unchanged. The MCP CIMD client cache is bounded at 1024 entries, purging expired entries and then deterministically evicting the entry closest to expiring under the existing mutex, preserving its 15-minute TTL. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 9ff1e38a-b018-4cf7-a5e9-5044a2efd03c
1 parent a44f222 commit 3a2391c

6 files changed

Lines changed: 621 additions & 20 deletions

File tree

‎api/pkg/auth/mcp_jwks.go‎

Lines changed: 102 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ import (
55
"crypto/rsa"
66
"encoding/base64"
77
"encoding/json"
8+
"errors"
89
"fmt"
910
"io"
1011
"math/big"
@@ -24,8 +25,19 @@ const (
2425

2526
// mcpJWKSMaxResponseBytes bounds the size of the JWKS document read from the network.
2627
mcpJWKSMaxResponseBytes = 1 << 20 // 1 MiB
28+
29+
// mcpJWKSDefaultMinRefreshInterval is the default minimum delay between two outbound
30+
// fetches of the JWKS endpoint. It bounds refresh amplification: without it, a flood of
31+
// tokens carrying random unknown "kid" headers would cause one outbound fetch per
32+
// request. The MCP service publishes a rotated signing key before it starts signing with
33+
// it, so a legitimate rotation is still picked up -- at worst one interval late.
34+
mcpJWKSDefaultMinRefreshInterval = time.Minute
2735
)
2836

37+
// errMCPJWKSRefreshThrottled reports that a JWKS refresh was skipped because the minimum
38+
// refresh interval has not elapsed yet.
39+
var errMCPJWKSRefreshThrottled = errors.New("MCP JWKS refresh is rate limited")
40+
2941
// mcpJWK is a single JSON Web Key as published by a JWKS endpoint. Only the
3042
// fields required to build an RSA public key are decoded.
3143
type mcpJWK struct {
@@ -41,20 +53,44 @@ type mcpJWKSet struct {
4153
}
4254

4355
// mcpJWKSCache fetches and caches the RSA public keys published by a JWKS
44-
// endpoint, keyed by "kid". It refreshes the cache once when a requested "kid"
45-
// cannot be found, and otherwise refreshes only after CacheTTL has elapsed.
56+
// endpoint, keyed by "kid".
57+
//
58+
// Two bounds keep an attacker from turning a stream of tokens carrying random unknown "kid"
59+
// headers into a stream of outbound fetches:
60+
//
61+
// - concurrent refreshes are collapsed into a single in-flight fetch that every waiting
62+
// caller shares, and
63+
// - a new fetch is never started until minRefreshInterval has elapsed since the previous
64+
// attempt (successful or not); until then, callers either reuse an already cached key or
65+
// fail closed.
66+
//
67+
// A legitimate key rotation is still picked up: the MCP service publishes a rotated key before
68+
// signing with it, and a missing "kid" triggers a real refresh as soon as the interval has
69+
// elapsed.
4670
type mcpJWKSCache struct {
47-
url string
48-
httpClient *http.Client
49-
cacheTTL time.Duration
71+
url string
72+
httpClient *http.Client
73+
cacheTTL time.Duration
74+
minRefreshInterval time.Duration
75+
76+
mu sync.Mutex
77+
keys map[string]*rsa.PublicKey
78+
fetchedAt time.Time
79+
lastAttemptAt time.Time
80+
inflight *mcpJWKSRefresh
81+
}
5082

51-
mu sync.Mutex
52-
keys map[string]*rsa.PublicKey
53-
fetchedAt time.Time
83+
// mcpJWKSRefresh is a single in-flight JWKS refresh shared by every caller that arrives while
84+
// it is running. err is written before done is closed, so a waiter that observes done may
85+
// safely read it.
86+
type mcpJWKSRefresh struct {
87+
done chan struct{}
88+
err error
5489
}
5590

56-
// newMCPJWKSCache creates a new mcpJWKSCache for the given JWKS URL.
57-
func newMCPJWKSCache(url string, httpClient *http.Client, cacheTTL time.Duration) *mcpJWKSCache {
91+
// newMCPJWKSCache creates a new mcpJWKSCache for the given JWKS URL. minRefreshInterval may be
92+
// <= 0, in which case mcpJWKSDefaultMinRefreshInterval is used.
93+
func newMCPJWKSCache(url string, httpClient *http.Client, cacheTTL time.Duration, minRefreshInterval time.Duration) *mcpJWKSCache {
5894
if httpClient == nil {
5995
httpClient = http.DefaultClient
6096
}
@@ -69,17 +105,22 @@ func newMCPJWKSCache(url string, httpClient *http.Client, cacheTTL time.Duration
69105
if cacheTTL <= 0 {
70106
cacheTTL = mcpJWKSDefaultCacheTTL
71107
}
108+
if minRefreshInterval <= 0 {
109+
minRefreshInterval = mcpJWKSDefaultMinRefreshInterval
110+
}
72111

73112
return &mcpJWKSCache{
74-
url: url,
75-
httpClient: client,
76-
cacheTTL: cacheTTL,
77-
keys: map[string]*rsa.PublicKey{},
113+
url: url,
114+
httpClient: client,
115+
cacheTTL: cacheTTL,
116+
minRefreshInterval: minRefreshInterval,
117+
keys: map[string]*rsa.PublicKey{},
78118
}
79119
}
80120

81-
// key returns the cached RSA public key for kid, refreshing the JWKS document
82-
// at most once per call when the cache is stale or the key is not yet known.
121+
// key returns the cached RSA public key for kid, refreshing the JWKS document when the cache is
122+
// stale or the key is not yet known -- subject to the collapsing and rate limiting described on
123+
// mcpJWKSCache.
83124
func (cache *mcpJWKSCache) key(ctx context.Context, kid string) (*rsa.PublicKey, error) {
84125
cache.mu.Lock()
85126
key, ok := cache.keys[kid]
@@ -90,7 +131,14 @@ func (cache *mcpJWKSCache) key(ctx context.Context, kid string) (*rsa.PublicKey,
90131
return key, nil
91132
}
92133

93-
if err := cache.refresh(ctx); err != nil {
134+
if err := cache.refreshOnce(ctx); err != nil {
135+
// A rate-limited refresh must not invalidate a key we already hold: serving the
136+
// (stale but still published) cached key is strictly better than failing a
137+
// legitimate request because the cache TTL elapsed moments after the last fetch
138+
// attempt.
139+
if errors.Is(err, errMCPJWKSRefreshThrottled) && ok {
140+
return key, nil
141+
}
94142
return nil, stacktrace.Propagatef(err, "cannot refresh MCP JWKS from [%s]", cache.url)
95143
}
96144

@@ -104,6 +152,43 @@ func (cache *mcpJWKSCache) key(ctx context.Context, kid string) (*rsa.PublicKey,
104152
return key, nil
105153
}
106154

155+
// refreshOnce performs at most one outbound JWKS fetch on behalf of every caller that needs one
156+
// at the same time, and refuses to start a new fetch until minRefreshInterval has elapsed since
157+
// the previous attempt.
158+
func (cache *mcpJWKSCache) refreshOnce(ctx context.Context) error {
159+
cache.mu.Lock()
160+
161+
if inflight := cache.inflight; inflight != nil {
162+
cache.mu.Unlock()
163+
select {
164+
case <-inflight.done:
165+
return inflight.err
166+
case <-ctx.Done():
167+
return ctx.Err()
168+
}
169+
}
170+
171+
if !cache.lastAttemptAt.IsZero() && time.Since(cache.lastAttemptAt) < cache.minRefreshInterval {
172+
cache.mu.Unlock()
173+
return errMCPJWKSRefreshThrottled
174+
}
175+
176+
inflight := &mcpJWKSRefresh{done: make(chan struct{})}
177+
cache.inflight = inflight
178+
cache.lastAttemptAt = time.Now()
179+
cache.mu.Unlock()
180+
181+
err := cache.refresh(ctx)
182+
inflight.err = err
183+
184+
cache.mu.Lock()
185+
cache.inflight = nil
186+
cache.mu.Unlock()
187+
close(inflight.done)
188+
189+
return err
190+
}
191+
107192
// refresh fetches and replaces the cached JWKS key set.
108193
func (cache *mcpJWKSCache) refresh(ctx context.Context) error {
109194
req, err := http.NewRequestWithContext(ctx, http.MethodGet, cache.url, nil)

0 commit comments

Comments
 (0)