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.
3143type 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.
4670type 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.
83124func (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.
108193func (cache * mcpJWKSCache ) refresh (ctx context.Context ) error {
109194 req , err := http .NewRequestWithContext (ctx , http .MethodGet , cache .url , nil )
0 commit comments