From e71d5cf8fe43ccdf6a2172217952485bf0002e86 Mon Sep 17 00:00:00 2001 From: derhornspieler <15236687+derhornspieler@users.noreply.github.com> Date: Sun, 23 Aug 2026 17:51:05 -0400 Subject: [PATCH] fix(anthropic): keep the token cache at its bound after a burst of distinct identities Eviction skipped every in-flight entry and then removed exactly one, so a burst of distinct federation identities pushed the map past max_entries and left it there: each later insert evicted one and added one, holding the high-water mark for the life of the process. The cap the parameter advertises was effectively advisory It now evicts as many as the overshoot needs, soonest-to-expire first, and treats an entry with no token as evictable rather than skipping it. An entry a leader owns or a follower waits on is still never a candidate, since dropping one would break single flight, so a moment when every entry is in flight can still over-insert; that residue is bounded by the concurrent mints themselves and clears as they finish --- litellm/llms/base_llm/auth/token_exchange.py | 20 +++++++++++++------ .../llms/base_llm/auth/test_token_exchange.py | 18 +++++++++++++++++ 2 files changed, 32 insertions(+), 6 deletions(-) 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()