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:
derhornspieler 2026-08-23 15:17:42 -04:00
parent 3c5b467284
commit db268de4f7
2 changed files with 139 additions and 16 deletions

View file

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

View file

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