mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
4b1c9d4498
commit
7705f0b975
2 changed files with 41 additions and 1 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue