From 7705f0b975cdbc813ac3c6a7183f4278e89c7286 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Mon, 6 Jul 2026 12:48:56 -0700 Subject: [PATCH] fix(mcp): bound the M2M lock dict and cap short-lived token TTL at real expiry Greptile P2s: the per-server lock dict now evicts its oldest entry past max_locks so ephemeral server ids (REST tools preview) cannot grow it unbounded, and the min-cache floor is capped at the token's actual lifetime so an expires_in below the skew is never served past expiry --- .../client_credentials.py | 14 +++++++++- .../test_client_credentials.py | 28 +++++++++++++++++++ 2 files changed, 41 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py index 3f6c329118f..332305db3c6 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py @@ -168,6 +168,7 @@ class ClientCredentialsTokenSource: default_ttl_seconds: float = 300.0, expiry_skew_seconds: float = 60.0, min_cache_seconds: float = 10.0, + max_locks: int = 1024, clock: Callable[[], float] = time.time, ) -> None: self._post = post @@ -175,10 +176,18 @@ class ClientCredentialsTokenSource: self._default_ttl_seconds = default_ttl_seconds self._expiry_skew_seconds = expiry_skew_seconds self._min_cache_seconds = min_cache_seconds + self._max_locks = max_locks self._clock = clock self._locks: dict[str, asyncio.Lock] = {} def _lock(self, server_id: str) -> asyncio.Lock: + """Per-server single-flight lock, bounded so ephemeral server ids (e.g. the REST tools + preview mints a fresh id per call) cannot grow the dict for the life of the process. + Evicting the oldest entry while a task still holds it only means a concurrent caller for + that server may run its own grant — single-flight is an optimization, not correctness. + """ + if server_id not in self._locks and len(self._locks) >= self._max_locks: + self._locks.pop(next(iter(self._locks)), None) return self._locks.setdefault(server_id, asyncio.Lock()) async def get(self, server_id: str, config: ClientCredentialsConfig) -> Result[OAuthToken, CredError]: @@ -243,8 +252,11 @@ class ClientCredentialsTokenSource: expires_at=self._clock() + expires_in if expires_in is not None else None, scopes=_parse_granted_scopes(body.get("scope")) or (), ) + # The min-cache floor is itself capped at the token's real lifetime, so a token whose + # expires_in is below the skew is never served past its actual expiry; a non-positive + # expires_in caches nothing (every request re-fetches, serialized by the per-server lock). ttl = ( - max(expires_in - self._expiry_skew_seconds, self._min_cache_seconds) + max(expires_in - self._expiry_skew_seconds, min(float(expires_in), self._min_cache_seconds), 0.0) if expires_in is not None else self._default_ttl_seconds ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py index bd00bb77abd..db7e6208e76 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py @@ -134,6 +134,34 @@ async def test_expires_in_bounds_the_cache_lifetime(): assert len(poster.calls) == 2 +@pytest.mark.asyncio +async def test_short_lived_token_is_never_served_past_its_expiry(): + # expires_in below the skew must not be floored into serving an expired token: the cache + # entry lapses with the token itself, and the next get re-fetches. + clock = _Clock(1000.0) + poster = _FakePoster([_success("t1", expires_in=5), _success("t2", expires_in=5)]) + source = ClientCredentialsTokenSource(poster, expiry_skew_seconds=60.0, min_cache_seconds=10.0, clock=clock) + first = await source.get("s", _config()) + assert isinstance(first, Ok) and first.ok.access_token == "t1" + clock.t = 1004.0 # still within the token's real lifetime + within = await source.get("s", _config()) + assert isinstance(within, Ok) and within.ok.access_token == "t1" + clock.t = 1006.0 # past expires_at: the floor must not keep serving t1 + lapsed = await source.get("s", _config()) + assert isinstance(lapsed, Ok) and lapsed.ok.access_token == "t2" + assert len(poster.calls) == 2 + + +@pytest.mark.asyncio +async def test_lock_dict_is_bounded_for_ephemeral_server_ids(): + poster = _FakePoster([_success()]) + source = ClientCredentialsTokenSource(poster, max_locks=8) + for index in range(20): + result = await source.get(f"ephemeral-{index}", _config()) + assert isinstance(result, Ok) + assert len(source._locks) <= 8 + + @pytest.mark.asyncio async def test_missing_expires_in_is_cached_briefly_not_an_hour(): clock = _Clock(1000.0)