mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(anthropic): scale the token refresh windows to the token's own lifetime
The advisory and mandatory refresh windows were flat 120s and 30s, while the fallback lifetime for a token minted without expires_in is 60s. Such a token was therefore born inside its own advisory window, so every request armed another background exchange against the provider's token endpoint and the cache never settled. Any real expires_in of 120s or less did the same Each window is now the smaller of its flat value and a fraction of the observed lifetime (half for advisory, an eighth for mandatory), so a 60s token is served until roughly its half life and then refreshed once. Tokens of 240s and above hit the flat values and keep exactly the previous behaviour. The tradeoff is deliberate: a 60s token may now be served with as little as 7.5s left, where before the mandatory wall was 30s, which for a 60s token meant refreshing at half its life on every path
This commit is contained in:
parent
3c5b467284
commit
db268de4f7
2 changed files with 139 additions and 16 deletions
|
|
@ -44,6 +44,8 @@ if TYPE_CHECKING:
|
|||
|
||||
ADVISORY_REFRESH_SECONDS: Final = 120.0
|
||||
MANDATORY_REFRESH_SECONDS: Final = 30.0
|
||||
ADVISORY_REFRESH_LIFETIME_FRACTION: Final = 0.5
|
||||
MANDATORY_REFRESH_LIFETIME_FRACTION: Final = 0.125
|
||||
ADVISORY_REFRESH_BACKOFF_SECONDS: Final = 5.0
|
||||
FALLBACK_TOKEN_TTL_SECONDS: Final = 60.0
|
||||
MAX_ASSERTION_BYTES: Final = 16 * 1024
|
||||
|
|
@ -221,6 +223,26 @@ def _sanitize_expires_in(expires_in: int | None) -> float:
|
|||
return float(expires_in)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RefreshWindows:
|
||||
advisory: float
|
||||
mandatory: float
|
||||
|
||||
|
||||
def _refresh_windows(lifetime_seconds: float | None) -> _RefreshWindows:
|
||||
"""A token whose whole life is shorter than the flat windows sits inside them from the moment it
|
||||
is minted, so every request would arm another background exchange against the token endpoint.
|
||||
Scaling each window by a fraction of the observed lifetime makes a 60s token refresh around its
|
||||
half life instead; at a lifetime of 240s and above both fractions reach the flat windows, so
|
||||
ordinary long-lived tokens keep exactly the 120s/30s behaviour."""
|
||||
if lifetime_seconds is None or lifetime_seconds <= 0.0:
|
||||
return _RefreshWindows(advisory=ADVISORY_REFRESH_SECONDS, mandatory=MANDATORY_REFRESH_SECONDS)
|
||||
return _RefreshWindows(
|
||||
advisory=min(ADVISORY_REFRESH_SECONDS, lifetime_seconds * ADVISORY_REFRESH_LIFETIME_FRACTION),
|
||||
mandatory=min(MANDATORY_REFRESH_SECONDS, lifetime_seconds * MANDATORY_REFRESH_LIFETIME_FRACTION),
|
||||
)
|
||||
|
||||
|
||||
def _capped_body_text(response: httpx.Response) -> str:
|
||||
if len(response.content) > MAX_RESPONSE_BYTES:
|
||||
return _OVERSIZED_BODY_MESSAGE
|
||||
|
|
@ -276,10 +298,11 @@ class _Entry:
|
|||
"""Single-flight state for one cache key; mutable by design, confined to the
|
||||
engine, and only ever mutated under the engine lock."""
|
||||
|
||||
__slots__ = ("backoff_until", "done", "force_refresh", "in_flight", "last_error", "token")
|
||||
__slots__ = ("backoff_until", "done", "force_refresh", "in_flight", "last_error", "lifetime_seconds", "token")
|
||||
|
||||
def __init__(self, force_refresh: bool = False) -> None:
|
||||
self.token: MintedToken | None = None
|
||||
self.lifetime_seconds: float | None = None
|
||||
self.in_flight: bool = False
|
||||
self.done: Final = threading.Event()
|
||||
self.backoff_until: float = float("-inf")
|
||||
|
|
@ -291,27 +314,30 @@ class _Entry:
|
|||
self.last_error = None
|
||||
self.done.clear()
|
||||
|
||||
def publish(self, result: ExchangeResult, backoff_until: float) -> None:
|
||||
def _store(self, token: MintedToken, now: float) -> None:
|
||||
self.token = token
|
||||
self.lifetime_seconds = None if token.expires_at is None else max(token.expires_at - now, 0.0)
|
||||
self.last_error = None
|
||||
|
||||
def publish(self, result: ExchangeResult, now: float) -> None:
|
||||
match result:
|
||||
case MintedToken():
|
||||
self.token = result
|
||||
self.last_error = None
|
||||
self._store(result, now)
|
||||
case _:
|
||||
self.last_error = result
|
||||
self.backoff_until = backoff_until
|
||||
self.backoff_until = now + ADVISORY_REFRESH_BACKOFF_SECONDS
|
||||
self.force_refresh = False
|
||||
self.in_flight = False
|
||||
self.done.set()
|
||||
|
||||
def publish_advisory(self, result: ExchangeResult, backoff_until: float) -> None:
|
||||
def publish_advisory(self, result: ExchangeResult, now: float) -> None:
|
||||
"""A failed advisory refresh records only the backoff, never ``last_error``: a follower whose
|
||||
cached token expires while this runs must be free to re-lead a fresh mint and recover."""
|
||||
match result:
|
||||
case MintedToken():
|
||||
self.token = result
|
||||
self.last_error = None
|
||||
self._store(result, now)
|
||||
case _:
|
||||
self.backoff_until = backoff_until
|
||||
self.backoff_until = now + ADVISORY_REFRESH_BACKOFF_SECONDS
|
||||
self.in_flight = False
|
||||
self.done.set()
|
||||
|
||||
|
|
@ -434,10 +460,11 @@ class JwtBearerTokenExchangeEngine:
|
|||
if token is not None and not entry.force_refresh:
|
||||
if token.expires_at is None:
|
||||
return _Serve(token=token)
|
||||
windows: Final = _refresh_windows(entry.lifetime_seconds)
|
||||
remaining: Final = token.expires_at - self._clock()
|
||||
if remaining > ADVISORY_REFRESH_SECONDS:
|
||||
if remaining > windows.advisory:
|
||||
return _Serve(token=token)
|
||||
if remaining > MANDATORY_REFRESH_SECONDS:
|
||||
if remaining > windows.mandatory:
|
||||
if entry.in_flight or self._clock() < entry.backoff_until:
|
||||
return _Serve(token=token)
|
||||
entry.arm()
|
||||
|
|
@ -460,7 +487,7 @@ class JwtBearerTokenExchangeEngine:
|
|||
def _lead(self, spec: TokenExchangeSpec, entry: _Entry) -> ExchangeResult:
|
||||
result: Final = self._exchange_never_raises(spec)
|
||||
with self._lock:
|
||||
entry.publish(result, backoff_until=self._clock() + ADVISORY_REFRESH_BACKOFF_SECONDS)
|
||||
entry.publish(result, now=self._clock())
|
||||
return result
|
||||
|
||||
def _await_leader(self, spec: TokenExchangeSpec, entry: _Entry) -> "ExchangeResult | None":
|
||||
|
|
@ -480,14 +507,14 @@ class JwtBearerTokenExchangeEngine:
|
|||
def _advisory_refresh(self, spec: TokenExchangeSpec, entry: _Entry) -> None:
|
||||
result: Final = self._exchange_never_raises(spec)
|
||||
with self._lock:
|
||||
entry.publish_advisory(result, backoff_until=self._clock() + ADVISORY_REFRESH_BACKOFF_SECONDS)
|
||||
now: Final = self._clock()
|
||||
entry.publish_advisory(result, now=now)
|
||||
stale_expires_at: Final = entry.token.expires_at if entry.token is not None else None
|
||||
stale_mandatory: Final = _refresh_windows(entry.lifetime_seconds).mandatory
|
||||
if isinstance(result, MintedToken):
|
||||
return
|
||||
seconds_to_mandatory_wall: Final = (
|
||||
max(stale_expires_at - self._clock() - MANDATORY_REFRESH_SECONDS, 0.0)
|
||||
if stale_expires_at is not None
|
||||
else 0.0
|
||||
max(stale_expires_at - now - stale_mandatory, 0.0) if stale_expires_at is not None else 0.0
|
||||
)
|
||||
verbose_logger.warning(
|
||||
"Advisory token refresh against %s failed (%s); serving the cached token for up to "
|
||||
|
|
|
|||
|
|
@ -1220,3 +1220,99 @@ class TestNonBearerTokenType:
|
|||
|
||||
assert isinstance(result, MalformedTokenResponse)
|
||||
assert "non-bearer" in result.detail
|
||||
|
||||
|
||||
class TestShortLivedRefreshWindows:
|
||||
"""A token whose lifetime is at or below the flat 120s advisory window used to be inside that
|
||||
window from birth, so every request armed another background exchange. The windows now scale
|
||||
with the observed lifetime; long-lived tokens must keep the flat 120s/30s behaviour."""
|
||||
|
||||
@staticmethod
|
||||
def _engine_with(
|
||||
expires_in: int | None,
|
||||
) -> tuple[JwtBearerTokenExchangeEngine, ScriptedPoster, FakeClock, ManualExecutor, TokenExchangeSpec]:
|
||||
poster = ScriptedPoster([token_response("short-lived", expires_in=expires_in), token_response("reminted")])
|
||||
clock = FakeClock(start=1_000.0)
|
||||
executor = ManualExecutor()
|
||||
engine = make_engine(poster, clock=clock, executor=executor)
|
||||
return engine, poster, clock, executor, make_spec()
|
||||
|
||||
def test_fallback_ttl_token_is_served_without_arming_a_refresh(self):
|
||||
engine, poster, clock, executor, spec = self._engine_with(expires_in=None)
|
||||
|
||||
first = mint(engine, spec)
|
||||
assert first.expires_at == 1_000.0 + FALLBACK_TOKEN_TTL_SECONDS
|
||||
|
||||
for _ in range(5):
|
||||
clock.advance(1.0)
|
||||
assert mint(engine, spec).access_token.get_secret_value() == "short-lived"
|
||||
|
||||
assert executor.pending == [], "a freshly minted fallback-TTL token must not arm a refresh on every request"
|
||||
assert len(poster.requests) == 1
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"elapsed,expect_advisory_submit",
|
||||
[(29.0, False), (30.0, True), (52.0, True)],
|
||||
)
|
||||
def test_fallback_ttl_token_refreshes_around_its_half_life(self, elapsed: float, expect_advisory_submit: bool):
|
||||
engine, poster, clock, executor, spec = self._engine_with(expires_in=None)
|
||||
|
||||
mint(engine, spec)
|
||||
clock.advance(elapsed)
|
||||
served = mint(engine, spec)
|
||||
|
||||
assert served.access_token.get_secret_value() == "short-lived"
|
||||
assert len(executor.pending) == (1 if expect_advisory_submit else 0)
|
||||
executor.run_all()
|
||||
assert len(poster.requests) == (2 if expect_advisory_submit else 1)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"elapsed,expect_new_token",
|
||||
[(52.0, False), (53.0, True)],
|
||||
)
|
||||
def test_fallback_ttl_mandatory_wall_scales_with_the_lifetime(self, elapsed: float, expect_new_token: bool):
|
||||
engine, poster, clock, executor, spec = self._engine_with(expires_in=None)
|
||||
|
||||
mint(engine, spec)
|
||||
clock.advance(elapsed)
|
||||
served = mint(engine, spec)
|
||||
|
||||
assert served.access_token.get_secret_value() == ("reminted" if expect_new_token else "short-lived")
|
||||
assert len(executor.pending) == (0 if expect_new_token else 1)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"elapsed,expect_advisory_submit",
|
||||
[(89.0, False), (100.0, True)],
|
||||
)
|
||||
def test_a_200s_token_scales_its_advisory_window_too(self, elapsed: float, expect_advisory_submit: bool):
|
||||
engine, poster, clock, executor, spec = self._engine_with(expires_in=200)
|
||||
|
||||
mint(engine, spec)
|
||||
clock.advance(elapsed)
|
||||
served = mint(engine, spec)
|
||||
|
||||
assert served.access_token.get_secret_value() == "short-lived"
|
||||
assert len(executor.pending) == (1 if expect_advisory_submit else 0)
|
||||
|
||||
@pytest.mark.parametrize("expires_in", [240, 3600])
|
||||
@pytest.mark.parametrize(
|
||||
"remaining,expect_advisory_submit,expect_new_token",
|
||||
[
|
||||
(121.0, False, False),
|
||||
(120.0, True, False),
|
||||
(31.0, True, False),
|
||||
(30.0, False, True),
|
||||
],
|
||||
)
|
||||
def test_long_lived_tokens_keep_the_flat_windows(
|
||||
self, expires_in: int, remaining: float, expect_advisory_submit: bool, expect_new_token: bool
|
||||
):
|
||||
engine, poster, clock, executor, spec = self._engine_with(expires_in=expires_in)
|
||||
|
||||
mint(engine, spec)
|
||||
clock.now = 1_000.0 + expires_in - remaining
|
||||
served = mint(engine, spec)
|
||||
|
||||
assert len(executor.pending) == (1 if expect_advisory_submit else 0)
|
||||
assert served.access_token.get_secret_value() == ("reminted" if expect_new_token else "short-lived")
|
||||
assert len(poster.requests) == (2 if expect_new_token else 1)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue