mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(auth): share one exchanged token across workers reading the same assertion
Anthropic accepts each identity assertion exactly once, so two uvicorn workers reading the same token file both minting from it means the second exchange is denied with jti_reused. Minted tokens now land in a per-user 0700 cache directory guarded by a file lock, so workers on the same host reuse one exchange until the token expires or the assertion rotates. A 401 is only retried when the re-read assertion actually differs, and the denial hint explains jti_reused. LITELLM_TOKEN_EXCHANGE_CACHE_DIR moves the cache and an empty value disables it
This commit is contained in:
parent
7a53a2bc1d
commit
66b72db818
6 changed files with 512 additions and 18 deletions
|
|
@ -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"
|
||||
|
|
|
|||
163
litellm/llms/base_llm/auth/shared_token_store.py
Normal file
163
litellm/llms/base_llm/auth/shared_token_store.py
Normal file
|
|
@ -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 ``<cache key>.json`` (mode 0600) and one ``<cache key>.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()}")
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ class TokenExchangeSpec:
|
|||
class MintedToken:
|
||||
access_token: SecretStr
|
||||
expires_at: float | None
|
||||
assertion_sha256: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
|
|||
261
tests/test_litellm/llms/base_llm/auth/test_shared_token_store.py
Normal file
261
tests/test_litellm/llms/base_llm/auth/test_shared_token_store.py
Normal file
|
|
@ -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"
|
||||
|
|
@ -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] = []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue