diff --git a/litellm/llms/anthropic/wif.py b/litellm/llms/anthropic/wif.py index b920a6c12fe..8fca87e8bee 100644 --- a/litellm/llms/anthropic/wif.py +++ b/litellm/llms/anthropic/wif.py @@ -82,7 +82,11 @@ _KEYCLOAK_FIELD_MAP: Final[Mapping[str, str]] = MappingProxyType( _DENIAL_HINT: Final = ( "Anthropic answers every denied exchange with the same 401; the reason (for example" " workspace_id_required or jti_reused) is only shown in the Claude Console under" - " Settings > Workload identity, in the rule's authentication history" + " Settings > Workload identity, in the rule's authentication history. jti_reused means this" + " identity token was already exchanged once: Anthropic accepts each assertion a single time, so a" + " token file or env var has to rotate before the minted token expires (the rule's" + " token_lifetime_seconds), or switch to the internal issuer or Keycloak source, which mint a" + " fresh assertion per exchange" ) _WORKSPACE_HINT: Final = ( "If the federation rule is enabled in more than one workspace, set anthropic_federation_workspace_id" diff --git a/litellm/llms/base_llm/auth/shared_token_store.py b/litellm/llms/base_llm/auth/shared_token_store.py new file mode 100644 index 00000000000..92d0fbc2949 --- /dev/null +++ b/litellm/llms/base_llm/auth/shared_token_store.py @@ -0,0 +1,163 @@ +"""Same-host token store for the JWT-bearer exchange engine. + +Anthropic accepts an assertion carrying a ``jti`` once per issuer, so every uvicorn worker that +reads the same projected token file must share the token the first exchange minted instead of +re-sending the same assertion. The engine keys the store by cache key and only reuses a stored +token minted from the assertion it currently holds; a rotated assertion always buys a fresh token. +""" + +import contextlib +import os +import sys +import tempfile +import threading +from collections.abc import Generator +from dataclasses import dataclass +from pathlib import Path +from typing import Final, Protocol + +from pydantic import BaseModel, SecretStr, ValidationError + +from litellm._logging import verbose_logger + +CACHE_DIR_ENV: Final = "LITELLM_TOKEN_EXCHANGE_CACHE_DIR" + + +@dataclass(frozen=True, slots=True) +class StoredToken: + access_token: SecretStr + expires_at_epoch: float | None + assertion_sha256: str + + +class SharedTokenStore(Protocol): + """Every method is best-effort: a store that cannot read, write, or lock degrades to a per-process + cache and never raises into the mint path.""" + + def load(self, key: str) -> StoredToken | None: ... + + def save(self, key: str, token: StoredToken) -> None: ... + + def delete(self, key: str) -> None: ... + + def lock(self, key: str) -> contextlib.AbstractContextManager[None]: ... + + +class _StoredTokenFile(BaseModel): + access_token: str + expires_at_epoch: float | None + assertion_sha256: str + + +def _directory_is_private(directory: Path) -> bool: + try: + directory.mkdir(mode=0o700, exist_ok=True) + stat: Final = directory.stat() + except OSError as e: + verbose_logger.warning("Token exchange cache directory %s is unusable (%s); caching per process", directory, e) + return False + if stat.st_uid != os.getuid() or stat.st_mode & 0o077: + verbose_logger.warning( + "Token exchange cache directory %s must be owned by uid %d with mode 0700; caching per process", + directory, + os.getuid(), + ) + return False + return True + + +class FileTokenStore: + """One ``.json`` (mode 0600) and one ``.lock`` (flock) per identity under a + directory only the proxy's uid can enter; the directory is checked on first use, not at import.""" + + def __init__(self, directory: Path) -> None: + self._directory: Final = directory + self._ready_lock: Final = threading.Lock() + self._ready: bool | None = None + + @property + def directory(self) -> Path: + return self._directory + + def _usable(self) -> bool: + with self._ready_lock: + if self._ready is None: + self._ready = _directory_is_private(self._directory) + return self._ready + + def load(self, key: str) -> StoredToken | None: + if not self._usable(): + return None + try: + raw: Final = (self._directory / f"{key}.json").read_bytes() + parsed: Final = _StoredTokenFile.model_validate_json(raw) + except FileNotFoundError: + return None + except (OSError, ValidationError) as e: + verbose_logger.debug("Ignoring unreadable token exchange cache entry: %s", e) + return None + return StoredToken( + access_token=SecretStr(parsed.access_token), + expires_at_epoch=parsed.expires_at_epoch, + assertion_sha256=parsed.assertion_sha256, + ) + + def save(self, key: str, token: StoredToken) -> None: + if not self._usable(): + return + body: Final = ( + _StoredTokenFile( + access_token=token.access_token.get_secret_value(), + expires_at_epoch=token.expires_at_epoch, + assertion_sha256=token.assertion_sha256, + ) + .model_dump_json() + .encode() + ) + try: + with tempfile.NamedTemporaryFile(dir=self._directory, prefix=f"{key}.", delete=False) as handle: + handle.write(body) + os.replace(handle.name, self._directory / f"{key}.json") + except OSError as e: + verbose_logger.debug("Token exchange cache entry not written: %s", e) + + def delete(self, key: str) -> None: + if not self._usable(): + return + with contextlib.suppress(FileNotFoundError, OSError): + (self._directory / f"{key}.json").unlink() + + @contextlib.contextmanager + def lock(self, key: str) -> Generator[None]: + if sys.platform == "win32" or not self._usable(): + yield + return + import fcntl + + try: + fd: Final = os.open(self._directory / f"{key}.lock", os.O_RDWR | os.O_CREAT, 0o600) + except OSError as e: + verbose_logger.debug("Token exchange cache lock unavailable (%s); minting without it", e) + yield + return + try: + fcntl.flock(fd, fcntl.LOCK_EX) + yield + finally: + with contextlib.suppress(OSError): + fcntl.flock(fd, fcntl.LOCK_UN) + os.close(fd) + + +def default_shared_token_store() -> SharedTokenStore | None: + """``LITELLM_TOKEN_EXCHANGE_CACHE_DIR`` relocates the store; setting it empty disables it. Without + it the store lives under the temp directory, keyed by uid, so the workers of one proxy share it and + other users on the host cannot read it. Windows has no ``flock``, so it caches per process there.""" + if sys.platform == "win32": + return None + configured: Final = os.environ.get(CACHE_DIR_ENV) + if configured == "": + return None + if configured is not None: + return FileTokenStore(Path(configured)) + return FileTokenStore(Path(tempfile.gettempdir()) / f"litellm-token-exchange-{os.getuid()}") diff --git a/litellm/llms/base_llm/auth/token_exchange.py b/litellm/llms/base_llm/auth/token_exchange.py index be2fb46d1b3..35b78eef7e9 100644 --- a/litellm/llms/base_llm/auth/token_exchange.py +++ b/litellm/llms/base_llm/auth/token_exchange.py @@ -26,6 +26,7 @@ from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError from typing_extensions import assert_never from litellm._logging import verbose_logger +from litellm.llms.base_llm.auth.shared_token_store import SharedTokenStore, StoredToken, default_shared_token_store from litellm.llms.base_llm.auth.types import ( AssertionReader, AssertionSource, @@ -290,6 +291,10 @@ def _cache_key(spec: TokenExchangeSpec) -> str: ).hexdigest() +def _assertion_digest(assertion: SecretStr) -> str: + return hashlib.sha256(assertion.get_secret_value().encode()).hexdigest() + + def _assertion_fetch(reader: AssertionReader, spec: TokenExchangeSpec) -> AssertionSource: """``spec.assertion_source`` (an identity source's own fetch/mint closure) takes priority over the engine-level reader when set; either way, failures are reported against ``spec.assertion_ref``.""" @@ -609,6 +614,10 @@ class _Unauthorized: assertion: SecretStr +def _denied(attempt: _Unauthorized) -> TokenEndpointError: + return redact_oauth_error_body(attempt.response.status_code, _capped_body_text(attempt.response), attempt.assertion) + + class JwtBearerTokenExchangeEngine: def __init__( self, @@ -618,6 +627,8 @@ class JwtBearerTokenExchangeEngine: refresh_executor: Executor | None = None, max_entries: int = 64, metrics_sink: TokenExchangeMetricsSink | None = None, + shared_store: SharedTokenStore | None = None, + wall_clock: Callable[[], float] = time.time, ) -> None: self._poster: Final[SyncTokenPoster] = poster if poster is not None else _HttpxSyncTokenPoster() self._assertion_reader: Final[AssertionReader] = ( @@ -629,6 +640,8 @@ class JwtBearerTokenExchangeEngine: self._metrics_sink: Final[TokenExchangeMetricsSink] = ( metrics_sink if metrics_sink is not None else ServiceLoggingMetricsSink() ) + self._shared_store: Final = shared_store + self._wall_clock: Final = wall_clock self._lock: Final = threading.Lock() self._entries: Final[dict[str, _Entry]] = {} # mutable-ok: engine-owned map guarded by _lock @@ -666,6 +679,8 @@ class JwtBearerTokenExchangeEngine: with self._lock: if key in self._entries: self._entries[key] = _Entry(force_refresh=True) + if self._shared_store is not None: + self._shared_store.delete(key) def _get_or_create_entry_locked(self, spec: TokenExchangeSpec) -> _Entry: key: Final = _cache_key(spec) @@ -809,23 +824,69 @@ class JwtBearerTokenExchangeEngine: return TokenTransportError(detail=f"{type(e).__name__}: {e}"[:_REDACTION_CAP]) def _exchange(self, spec: TokenExchangeSpec) -> ExchangeResult: - first: Final = self._attempt_exchange(spec) - if not isinstance(first, _Unauthorized): - return first - second: Final = self._attempt_exchange(spec) - if isinstance(second, _Unauthorized): - return redact_oauth_error_body( - second.response.status_code, _capped_body_text(second.response), second.assertion - ) - return second - - def _attempt_exchange(self, spec: TokenExchangeSpec) -> "ExchangeResult | _Unauthorized": - assertion: Final = _read_assertion(_assertion_fetch(self._assertion_reader, spec), spec.assertion_ref) + fetch: Final = _assertion_fetch(self._assertion_reader, spec) + assertion: Final = _read_assertion(fetch, spec.assertion_ref) if isinstance(assertion, AssertionSourceError): return assertion url_check: Final = validate_token_endpoint_url(spec.token_url) if isinstance(url_check, InsecureTokenUrl): return url_check + if self._shared_store is None: + return self._mint(spec, fetch, assertion) + key: Final = _cache_key(spec) + with self._shared_store.lock(key): + shared: Final = self._shared_token(self._shared_store.load(key), _assertion_digest(assertion)) + if shared is not None: + return shared + minted: Final = self._mint(spec, fetch, assertion) + if isinstance(minted, MintedToken): + self._shared_store.save(key, self._stored_token(minted)) + return minted + + def _shared_token(self, stored: StoredToken | None, assertion_sha256: str) -> MintedToken | None: + """A stored token minted from the very assertion this process holds is the token that assertion + bought: another worker sharing the token file already exchanged it, and an issuer enforcing + single-use ``jti`` would only deny a second exchange.""" + if stored is None or stored.assertion_sha256 != assertion_sha256: + return None + if stored.expires_at_epoch is None: + return MintedToken(access_token=stored.access_token, expires_at=None, assertion_sha256=assertion_sha256) + remaining: Final = stored.expires_at_epoch - self._wall_clock() + if remaining <= 0.0: + return None + return MintedToken( + access_token=stored.access_token, + expires_at=self._clock() + remaining, + assertion_sha256=assertion_sha256, + ) + + def _stored_token(self, token: MintedToken) -> StoredToken: + return StoredToken( + access_token=token.access_token, + expires_at_epoch=( + None if token.expires_at is None else self._wall_clock() + (token.expires_at - self._clock()) + ), + assertion_sha256=token.assertion_sha256, + ) + + def _mint(self, spec: TokenExchangeSpec, fetch: AssertionSource, assertion: SecretStr) -> ExchangeResult: + """One 401 earns one retry, and only with an assertion that changed since the first attempt: a + token file rotated between the read and the POST is worth resending, the same assertion is not, + since an issuer that already consumed its ``jti`` denies it again.""" + first: Final = self._post_assertion(spec, assertion) + if not isinstance(first, _Unauthorized): + return first + reread: Final = _read_assertion(fetch, spec.assertion_ref) + if isinstance(reread, AssertionSourceError): + return reread + if reread.get_secret_value() == assertion.get_secret_value(): + return _denied(first) + second: Final = self._post_assertion(spec, reread) + if isinstance(second, _Unauthorized): + return _denied(second) + return second + + def _post_assertion(self, spec: TokenExchangeSpec, assertion: SecretStr) -> "ExchangeResult | _Unauthorized": try: response: Final = self._poster.post( spec.token_url, @@ -839,7 +900,7 @@ class JwtBearerTokenExchangeEngine: return _Unauthorized(response=response, assertion=assertion) return self._parse_response(response, assertion) - def _parse_response(self, response: httpx.Response, assertion: SecretStr | None = None) -> ExchangeResult: + def _parse_response(self, response: httpx.Response, assertion: SecretStr) -> ExchangeResult: if not 200 <= response.status_code < 300: return redact_oauth_error_body(response.status_code, _capped_body_text(response), assertion) if len(response.content) > MAX_RESPONSE_BYTES: @@ -855,7 +916,8 @@ class JwtBearerTokenExchangeEngine: return MintedToken( access_token=SecretStr(parsed.access_token), expires_at=self._clock() + _sanitize_expires_in(parsed.expires_in), + assertion_sha256=_assertion_digest(assertion), ) -default_token_exchange_engine: Final = JwtBearerTokenExchangeEngine() +default_token_exchange_engine: Final = JwtBearerTokenExchangeEngine(shared_store=default_shared_token_store()) diff --git a/litellm/llms/base_llm/auth/types.py b/litellm/llms/base_llm/auth/types.py index 9d3a7a5012b..2cb59cf6f20 100644 --- a/litellm/llms/base_llm/auth/types.py +++ b/litellm/llms/base_llm/auth/types.py @@ -41,6 +41,7 @@ class TokenExchangeSpec: class MintedToken: access_token: SecretStr expires_at: float | None + assertion_sha256: str @dataclass(frozen=True, slots=True) diff --git a/tests/test_litellm/llms/base_llm/auth/test_shared_token_store.py b/tests/test_litellm/llms/base_llm/auth/test_shared_token_store.py new file mode 100644 index 00000000000..b2ec509d91b --- /dev/null +++ b/tests/test_litellm/llms/base_llm/auth/test_shared_token_store.py @@ -0,0 +1,261 @@ +"""Two engines standing in for two uvicorn workers that read the same assertion and share one +``FileTokenStore``: an issuer that accepts each assertion once must see one exchange per assertion.""" + +import json +import os +import stat +import threading +from collections.abc import Callable, Mapping +from pathlib import Path +from typing import Final + +import httpx +import pytest + +from litellm.llms.base_llm.auth.shared_token_store import ( + CACHE_DIR_ENV, + FileTokenStore, + default_shared_token_store, +) +from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine +from litellm.llms.base_llm.auth.types import MintedToken, TokenEndpointError +from tests.test_litellm.llms.base_llm.auth.test_token_exchange import ( + DEFAULT_ASSERTION, + DEFAULT_REF, + FakeClock, + ManualExecutor, + RecordingMetricsSink, + ScriptedPoster, + make_spec, + token_response, +) + + +class SingleUsePoster: + """Mints for an assertion it has never seen and answers 401 to any assertion sent a second time, + which is how an issuer enforcing single-use ``jti`` behaves.""" + + def __init__(self, token: str = "sk-ant-oat01-minted", expires_in: int = 3600) -> None: + self.requests: list[dict] = [] + self._token = token + self._expires_in = expires_in + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + body = json.loads(content) + seen_before = any(prior["assertion"] == body["assertion"] for prior in self.requests) + self.requests.append(body) + if seen_before: + return httpx.Response(401, json={"error": "invalid_grant"}) + return token_response(f"{self._token}-{len(self.requests)}", expires_in=self._expires_in) + + +def store_engine( + poster, + store: FileTokenStore, + *, + reader: Mapping[str, str] | None = None, + clock: FakeClock | None = None, + wall_clock: Callable[[], float] | None = None, +) -> JwtBearerTokenExchangeEngine: + return JwtBearerTokenExchangeEngine( + poster=poster, + assertion_reader=(reader if reader is not None else {DEFAULT_REF: DEFAULT_ASSERTION}).get, + clock=clock if clock is not None else FakeClock(), + refresh_executor=ManualExecutor(), + metrics_sink=RecordingMetricsSink(), + shared_store=store, + wall_clock=wall_clock if wall_clock is not None else FakeClock(1_700_000_000.0), + ) + + +def minted(result: object) -> MintedToken: + assert isinstance(result, MintedToken), result + return result + + +def stored_files(directory: Path) -> list[Path]: + return sorted(directory.glob("*.json")) + + +def test_second_worker_reuses_the_first_workers_token_without_a_post(tmp_path: Path): + poster = SingleUsePoster() + store = FileTokenStore(tmp_path) + first = minted(store_engine(poster, store).get_token(make_spec())) + + second = minted(store_engine(poster, store).get_token(make_spec())) + + assert second.access_token.get_secret_value() == first.access_token.get_secret_value() + assert len(poster.requests) == 1 + + +def test_a_rotated_assertion_buys_a_fresh_token_that_other_workers_pick_up(tmp_path: Path): + poster = SingleUsePoster() + store = FileTokenStore(tmp_path) + assertions = {DEFAULT_REF: "jwt-v1"} + first = minted(store_engine(poster, store, reader=assertions).get_token(make_spec())) + + assertions[DEFAULT_REF] = "jwt-v2" + rotated = minted(store_engine(poster, store, reader=assertions).get_token(make_spec())) + follower = minted(store_engine(poster, store, reader=assertions).get_token(make_spec())) + + assert rotated.access_token.get_secret_value() != first.access_token.get_secret_value() + assert follower.access_token.get_secret_value() == rotated.access_token.get_secret_value() + assert [request["assertion"] for request in poster.requests] == ["jwt-v1", "jwt-v2"] + + +def test_an_expired_shared_token_is_not_reused(tmp_path: Path): + poster = ScriptedPoster([token_response("first", expires_in=60), token_response("second", expires_in=60)]) + store = FileTokenStore(tmp_path) + wall = FakeClock(1_700_000_000.0) + minted(store_engine(poster, store, wall_clock=wall).get_token(make_spec())) + + wall.advance(61) + later = minted(store_engine(poster, store, wall_clock=wall).get_token(make_spec())) + + assert later.access_token.get_secret_value() == "second" + assert len(poster.requests) == 2 + + +def test_the_remaining_lifetime_survives_different_monotonic_origins(tmp_path: Path): + poster = ScriptedPoster([token_response(expires_in=3600)]) + store = FileTokenStore(tmp_path) + wall = FakeClock(1_700_000_000.0) + minted(store_engine(poster, store, clock=FakeClock(1_000.0), wall_clock=wall).get_token(make_spec())) + + wall.advance(600) + later_clock = FakeClock(50_000.0) + later = minted(store_engine(poster, store, clock=later_clock, wall_clock=wall).get_token(make_spec())) + + assert later.expires_at == pytest.approx(50_000.0 + 3000.0) + assert len(poster.requests) == 1 + + +def test_mandatory_refresh_serves_the_shared_token_until_it_expires_then_fails_once(tmp_path: Path): + """With an unrotated assertion there is nothing new to exchange: refreshes inside the mandatory + window keep serving the shared token, and once it has expired the one allowed POST is denied + without a second identical POST behind it.""" + poster = SingleUsePoster(expires_in=3600) + store = FileTokenStore(tmp_path) + clock = FakeClock(1_000.0) + engine = store_engine(poster, store, clock=clock, wall_clock=clock) + first = minted(engine.get_token(make_spec())) + + clock.advance(3600 - 20) + refreshed = minted(engine.get_token(make_spec())) + assert refreshed.access_token.get_secret_value() == first.access_token.get_secret_value() + assert len(poster.requests) == 1 + + clock.advance(25) + failed = engine.get_token(make_spec()) + + assert isinstance(failed, TokenEndpointError) + assert failed.status_code == 401 + assert len(poster.requests) == 2 + + +def test_a_corrupt_cache_entry_is_treated_as_absent(tmp_path: Path): + poster = ScriptedPoster([token_response("first"), token_response("second")]) + store = FileTokenStore(tmp_path) + minted(store_engine(poster, store).get_token(make_spec())) + (entry,) = stored_files(tmp_path) + entry.write_text("{not json") + + later = minted(store_engine(poster, store).get_token(make_spec())) + + assert later.access_token.get_secret_value() == "second" + assert json.loads(entry.read_text())["access_token"] == "second" + + +def test_cache_entries_are_private_to_the_owner(tmp_path: Path): + store = FileTokenStore(tmp_path / "cache") + minted(store_engine(ScriptedPoster([token_response()]), store).get_token(make_spec())) + + (entry,) = stored_files(tmp_path / "cache") + assert stat.S_IMODE((tmp_path / "cache").stat().st_mode) == 0o700 + assert stat.S_IMODE(entry.stat().st_mode) == 0o600 + + +def test_a_group_readable_cache_directory_is_refused_and_the_engine_still_mints(tmp_path: Path): + loose = tmp_path / "loose" + loose.mkdir(mode=0o750) + os.chmod(loose, 0o750) + poster = ScriptedPoster([token_response("first"), token_response("second")]) + store = FileTokenStore(loose) + + minted(store_engine(poster, store).get_token(make_spec())) + later = minted(store_engine(poster, store).get_token(make_spec())) + + assert later.access_token.get_secret_value() == "second" + assert stored_files(loose) == [] + + +def test_default_store_follows_the_cache_dir_env(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv(CACHE_DIR_ENV, "") + assert default_shared_token_store() is None + + monkeypatch.setenv(CACHE_DIR_ENV, str(tmp_path / "configured")) + configured = default_shared_token_store() + assert isinstance(configured, FileTokenStore) + assert configured.directory == tmp_path / "configured" + + monkeypatch.delenv(CACHE_DIR_ENV) + default = default_shared_token_store() + assert isinstance(default, FileTokenStore) + assert default.directory.name == f"litellm-token-exchange-{os.getuid()}" + + +class GatedSingleUsePoster(SingleUsePoster): + def __init__(self) -> None: + super().__init__() + self.entered = threading.Event() + self.release = threading.Event() + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.entered.set() + assert self.release.wait(timeout=10) + return super().post(url, content=content, headers=headers, timeout=timeout) + + +def test_a_worker_arriving_mid_exchange_waits_for_the_leader_instead_of_posting(tmp_path: Path): + poster = GatedSingleUsePoster() + store = FileTokenStore(tmp_path) + leader = store_engine(poster, store) + follower = store_engine(poster, store) + results: dict[str, object] = {} + + def lead() -> None: + results["leader"] = leader.get_token(make_spec()) + + def follow() -> None: + results["follower"] = follower.get_token(make_spec()) + + leader_thread = threading.Thread(target=lead, daemon=True) + leader_thread.start() + assert poster.entered.wait(timeout=10) + follower_thread = threading.Thread(target=follow, daemon=True) + follower_thread.start() + follower_thread.join(timeout=0.5) + assert follower_thread.is_alive() + poster.release.set() + leader_thread.join(timeout=10) + follower_thread.join(timeout=10) + + assert not follower_thread.is_alive() + assert ( + minted(results["follower"]).access_token.get_secret_value() + == minted(results["leader"]).access_token.get_secret_value() + ) + assert len(poster.requests) == 1 + + +def test_invalidate_drops_the_shared_entry(tmp_path: Path): + poster = ScriptedPoster([token_response("first"), token_response("second")]) + store = FileTokenStore(tmp_path) + engine = store_engine(poster, store) + spec: Final = make_spec() + minted(engine.get_token(spec)) + + engine.invalidate(spec) + + assert stored_files(tmp_path) == [] + assert minted(store_engine(poster, store).get_token(spec)).access_token.get_secret_value() == "second" 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 129eec88c2b..8121d9bdbf1 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 @@ -526,7 +526,9 @@ def test_401_retry_redacts_the_assertion_actually_sent_not_a_fresh_reread(): assert "assertion-v3" not in result.redacted_body -def test_401_twice_is_endpoint_error(): +def test_401_with_an_unchanged_assertion_is_not_resent(): + """An issuer that consumed the assertion's ``jti`` denies the identical assertion again, so the + retry only happens when the re-read assertion differs from the one the 401 came back for.""" poster = ScriptedPoster([httpx.Response(401, json={"error": "invalid_grant"})]) engine = make_engine(poster) @@ -535,7 +537,7 @@ def test_401_twice_is_endpoint_error(): assert isinstance(result, TokenEndpointError) assert result.status_code == 401 assert "invalid_grant" in result.redacted_body - assert len(poster.requests) == 2 + assert len(poster.requests) == 1 class TestRedactionAndCaps: @@ -1220,7 +1222,8 @@ class TwoAttemptGatedPoster: def test_follower_budget_outlasts_slow_two_attempt_leader(): poster = TwoAttemptGatedPoster() - engine = make_engine(poster) + rotating_reads = iter(["jwt-before-rotation", "jwt-after-rotation"]) + engine = make_engine(poster, reader=lambda ref: next(rotating_reads, "jwt-after-rotation")) spec = make_spec(timeout_seconds=1.0) leader_results: list[ExchangeResult] = []