diff --git a/litellm/llms/base_llm/auth/token_exchange.py b/litellm/llms/base_llm/auth/token_exchange.py index 1bf29b00cd6..8518fe88905 100644 --- a/litellm/llms/base_llm/auth/token_exchange.py +++ b/litellm/llms/base_llm/auth/token_exchange.py @@ -15,6 +15,7 @@ import time from collections.abc import Callable, Coroutine, Mapping from concurrent.futures import Executor, ThreadPoolExecutor from dataclasses import dataclass +from math import inf from types import MappingProxyType from typing import TYPE_CHECKING, Final, Protocol, TypeAlias from urllib.parse import urlencode, urlsplit @@ -577,7 +578,7 @@ class JwtBearerTokenExchangeEngine: case _Lead(call_type=call_type): return self._lead(spec, entry, call_type) case _Follow(): - followed: Final = self._await_leader(spec, entry) + followed = self._await_leader(spec, entry) # rebind-ok: one leader wait per round if followed is not None: return followed case _: @@ -615,13 +616,20 @@ class JwtBearerTokenExchangeEngine: del self._entries[key] if len(self._entries) < self._max_entries: return - evictable: Final = tuple( - (entry.token.expires_at, key) + # Evict soonest-to-expire first, and take as many as the overshoot needs rather than one, so a + # burst of distinct identities does not leave the map permanently above max_entries. An entry + # a leader owns or a follower waits on is never a candidate, so a moment where every entry is + # in flight still over-inserts; that residue is bounded by the concurrent mints themselves. + evictable: Final = sorted( + ( + entry.token.expires_at if entry.token is not None and entry.token.expires_at is not None else -inf, + key, + ) for key, entry in self._entries.items() - if not entry.in_flight and entry.token is not None and entry.token.expires_at is not None + if not entry.in_flight ) - if evictable: - del self._entries[min(evictable)[1]] + for _, key in evictable[: len(self._entries) - self._max_entries + 1]: + del self._entries[key] def _classify_and_arm_locked(self, entry: _Entry) -> _Decision: token: Final = entry.token diff --git a/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py b/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py index 0d61863689d..697997a5faa 100644 --- a/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py +++ b/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py @@ -889,6 +889,24 @@ def test_cache_key_semantics(): assert cached.access_token.get_secret_value() == "sk-ant-oat01-minted" +def test_the_cache_returns_to_its_bound_after_an_all_in_flight_burst(): + """An entry a leader owns is never evictable, so a burst of distinct identities can push the map + past max_entries. It must come back down once those entries are idle, rather than holding the + high-water mark for the life of the process.""" + clock = FakeClock() + engine = make_engine(ScriptedPoster([token_response(expires_in=3600)]), clock=clock, max_entries=4) + + def spec_for(index: int) -> TokenExchangeSpec: + return make_spec(cache_key_identity=("fdrl_1", f"org-{index}", "", "")) + + for index in range(12): + mint(engine, spec_for(index)) + + assert len(engine._entries) <= 4, ( # noqa: SLF001 # the bound under test is internal state + f"the cap is enforced once entries are idle, saw {len(engine._entries)}" + ) + + def test_bounded_eviction(): clock = FakeClock()