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
This commit is contained in:
Tin Chi Lo 2026-07-06 12:48:56 -07:00 • committed by Tin
parent 4b1c9d4498
commit 7705f0b975
2 changed files with 41 additions and 1 deletions

View file

@ -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
)

View file

@ -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)