mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(anthropic): report workload identity token health as a litellm service
A deployment that mints short-lived tokens on a two tier refresh was unobservable: the exchange emitted one warning and no metrics, so the first sign of a failing token endpoint was user visible 401s. Registering the exchange as a service type puts it on the same path redis and postgres already use, so success and failure counters and a latency histogram reach /metrics and OTel spans from one emit, with no per backend wiring Each attempt reports its call type, separating a cold mint from a mandatory refresh and from a background advisory refresh, and tokens served from cache are counted on their own service so a cache hit never fakes an exchange latency. Failures carry the ExchangeError variant as a low cardinality error class The sink is a constructor parameter like the poster and the clock, so tests inject a fake rather than patching. It hands events to a single worker thread, which keeps the sync path free of the async hooks and of any need for a running loop, and every emission is wrapped so a broken metrics backend can never fail or delay a mint. Labels carry only the variant name and the redacted summary the error path already produces, never an assertion, a token or a secret
This commit is contained in:
parent
d3ee6c6fa2
commit
67688e1def
5 changed files with 456 additions and 7 deletions
|
|
@ -12,11 +12,11 @@ import hashlib
|
|||
import json
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Mapping
|
||||
from collections.abc import Callable, Coroutine, Mapping
|
||||
from concurrent.futures import Executor, ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias
|
||||
from typing import TYPE_CHECKING, Final, Protocol, TypeAlias
|
||||
from urllib.parse import urlencode, urlsplit
|
||||
|
||||
import httpx
|
||||
|
|
@ -28,6 +28,7 @@ from litellm.llms.base_llm.auth.types import (
|
|||
AssertionReader,
|
||||
AssertionSource,
|
||||
AssertionSourceError,
|
||||
ExchangeCallType,
|
||||
ExchangeError,
|
||||
ExchangeResult,
|
||||
InsecureTokenUrl,
|
||||
|
|
@ -35,13 +36,20 @@ from litellm.llms.base_llm.auth.types import (
|
|||
MintedToken,
|
||||
SyncTokenPoster,
|
||||
TokenEndpointError,
|
||||
TokenExchangeMetricsSink,
|
||||
TokenExchangeSpec,
|
||||
TokenTransportError,
|
||||
)
|
||||
from litellm.types.services import ServiceTypes
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
CALL_TYPE_COLD_MINT: Final[ExchangeCallType] = "cold_mint"
|
||||
CALL_TYPE_MANDATORY_REFRESH: Final[ExchangeCallType] = "mandatory_refresh"
|
||||
CALL_TYPE_ADVISORY_REFRESH: Final[ExchangeCallType] = "advisory_refresh"
|
||||
CALL_TYPE_CACHE_HIT: Final = "cache_hit"
|
||||
|
||||
ADVISORY_REFRESH_SECONDS: Final = 120.0
|
||||
MANDATORY_REFRESH_SECONDS: Final = 30.0
|
||||
ADVISORY_REFRESH_LIFETIME_FRACTION: Final = 0.5
|
||||
|
|
@ -166,6 +174,44 @@ def _error_summary(error: ExchangeError) -> str:
|
|||
assert_never(error)
|
||||
|
||||
|
||||
class _MetricsFailure(Exception):
|
||||
"""Never raised: typed carriers handed to the service failure hook so the prometheus
|
||||
``error_class`` label names the ``ExchangeError`` variant; the message is the redacted
|
||||
``_error_summary`` and carries no credential material."""
|
||||
|
||||
|
||||
class TokenExchangeAssertionSourceFailure(_MetricsFailure): ...
|
||||
|
||||
|
||||
class TokenExchangeInsecureUrlFailure(_MetricsFailure): ...
|
||||
|
||||
|
||||
class TokenExchangeEndpointFailure(_MetricsFailure): ...
|
||||
|
||||
|
||||
class TokenExchangeTransportFailure(_MetricsFailure): ...
|
||||
|
||||
|
||||
class TokenExchangeMalformedResponseFailure(_MetricsFailure): ...
|
||||
|
||||
|
||||
def _failure_exception(error: ExchangeError) -> _MetricsFailure:
|
||||
summary: Final = _error_summary(error)
|
||||
match error:
|
||||
case AssertionSourceError():
|
||||
return TokenExchangeAssertionSourceFailure(summary)
|
||||
case InsecureTokenUrl():
|
||||
return TokenExchangeInsecureUrlFailure(summary)
|
||||
case TokenEndpointError():
|
||||
return TokenExchangeEndpointFailure(summary)
|
||||
case TokenTransportError():
|
||||
return TokenExchangeTransportFailure(summary)
|
||||
case MalformedTokenResponse():
|
||||
return TokenExchangeMalformedResponseFailure(summary)
|
||||
case _:
|
||||
assert_never(error)
|
||||
|
||||
|
||||
def _cache_key(spec: TokenExchangeSpec) -> str:
|
||||
return hashlib.sha256(
|
||||
"\x1f".join((spec.token_url, spec.assertion_ref, *spec.cache_key_identity)).encode()
|
||||
|
|
@ -294,6 +340,94 @@ class _HttpxSyncTokenPoster:
|
|||
return response
|
||||
|
||||
|
||||
class _ServiceLoggingHooks(Protocol):
|
||||
"""The slice of ``litellm._service_logger.ServiceLogging`` the metrics sink calls; a protocol
|
||||
so tests inject a recorder instead of monkeypatching."""
|
||||
|
||||
async def async_service_success_hook(self, service: ServiceTypes, call_type: str, duration: float) -> None: ...
|
||||
|
||||
async def async_service_failure_hook(
|
||||
self, service: ServiceTypes, duration: float, error: str | Exception, call_type: str
|
||||
) -> None: ...
|
||||
|
||||
|
||||
_HooksCoroFactory: TypeAlias = Callable[
|
||||
[_ServiceLoggingHooks], # mutable-ok: Callable param-list syntax, not a list
|
||||
Coroutine[object, object, None],
|
||||
]
|
||||
|
||||
|
||||
def _default_service_logging() -> _ServiceLoggingHooks:
|
||||
from litellm._service_logger import ServiceLogging
|
||||
|
||||
return ServiceLogging()
|
||||
|
||||
|
||||
class ServiceLoggingMetricsSink:
|
||||
"""Default sink: bridges engine metrics onto litellm's ServiceTypes pattern
|
||||
(prometheus ``litellm_anthropic_wif_*`` via ``service_callback``). The engine's entry points
|
||||
are sync threads with no event loop, and the service hooks are async, so every emission is
|
||||
fire-and-forget on a dedicated single worker thread that owns its own short-lived loop --
|
||||
the mint path only ever pays for an executor queue put."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
service_logging_factory: Callable[[], _ServiceLoggingHooks] = _default_service_logging,
|
||||
executor: Executor | None = None,
|
||||
) -> None:
|
||||
self._lock: Final = threading.Lock()
|
||||
self._service_logging_factory: Final = service_logging_factory
|
||||
self._service_logging: _ServiceLoggingHooks | None = None
|
||||
self._executor: Executor | None = executor
|
||||
|
||||
def _service_logging_instance(self) -> _ServiceLoggingHooks:
|
||||
with self._lock:
|
||||
if self._service_logging is None:
|
||||
self._service_logging = self._service_logging_factory()
|
||||
return self._service_logging
|
||||
|
||||
def _executor_instance(self) -> Executor:
|
||||
with self._lock:
|
||||
if self._executor is None:
|
||||
self._executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="litellm-token-exchange-metrics")
|
||||
return self._executor
|
||||
|
||||
def _emit(self, coro_factory: _HooksCoroFactory) -> None:
|
||||
try:
|
||||
asyncio.run(coro_factory(self._service_logging_instance()))
|
||||
except Exception as e: # noqa: BLE001 # metrics are best-effort; emission failures must never surface
|
||||
verbose_logger.debug("token exchange metrics emission failed: %s", e)
|
||||
|
||||
def _submit(self, coro_factory: _HooksCoroFactory) -> None:
|
||||
self._executor_instance().submit(self._emit, coro_factory)
|
||||
|
||||
def exchange_success(self, *, call_type: ExchangeCallType, duration_seconds: float) -> None:
|
||||
def start(hooks: _ServiceLoggingHooks) -> Coroutine[object, object, None]:
|
||||
return hooks.async_service_success_hook(
|
||||
service=ServiceTypes.ANTHROPIC_WIF, call_type=call_type, duration=duration_seconds
|
||||
)
|
||||
|
||||
self._submit(start)
|
||||
|
||||
def exchange_failure(self, *, call_type: ExchangeCallType, duration_seconds: float, error: ExchangeError) -> None:
|
||||
failure: Final = _failure_exception(error)
|
||||
|
||||
def start(hooks: _ServiceLoggingHooks) -> Coroutine[object, object, None]:
|
||||
return hooks.async_service_failure_hook(
|
||||
service=ServiceTypes.ANTHROPIC_WIF, duration=duration_seconds, error=failure, call_type=call_type
|
||||
)
|
||||
|
||||
self._submit(start)
|
||||
|
||||
def cache_hit(self) -> None:
|
||||
def start(hooks: _ServiceLoggingHooks) -> Coroutine[object, object, None]:
|
||||
return hooks.async_service_success_hook(
|
||||
service=ServiceTypes.ANTHROPIC_WIF_CACHE, call_type=CALL_TYPE_CACHE_HIT, duration=0.0
|
||||
)
|
||||
|
||||
self._submit(start)
|
||||
|
||||
|
||||
class _Entry:
|
||||
"""Single-flight state for one cache key; mutable by design, confined to the
|
||||
engine, and only ever mutated under the engine lock."""
|
||||
|
|
@ -354,7 +488,7 @@ class _ServeAndRefresh:
|
|||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Lead:
|
||||
pass
|
||||
call_type: ExchangeCallType
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -384,6 +518,7 @@ class JwtBearerTokenExchangeEngine:
|
|||
clock: Callable[[], float] = time.monotonic,
|
||||
refresh_executor: Executor | None = None,
|
||||
max_entries: int = 64,
|
||||
metrics_sink: TokenExchangeMetricsSink | None = None,
|
||||
) -> None:
|
||||
self._poster: Final[SyncTokenPoster] = poster if poster is not None else _HttpxSyncTokenPoster()
|
||||
self._assertion_reader: Final[AssertionReader] = (
|
||||
|
|
@ -392,6 +527,9 @@ class JwtBearerTokenExchangeEngine:
|
|||
self._clock: Final = clock
|
||||
self._refresh_executor: Executor | None = refresh_executor
|
||||
self._max_entries: Final = max_entries
|
||||
self._metrics_sink: Final[TokenExchangeMetricsSink] = (
|
||||
metrics_sink if metrics_sink is not None else ServiceLoggingMetricsSink()
|
||||
)
|
||||
self._lock: Final = threading.Lock()
|
||||
self._entries: Final[dict[str, _Entry]] = {} # mutable-ok: engine-owned map guarded by _lock
|
||||
|
||||
|
|
@ -401,14 +539,16 @@ class JwtBearerTokenExchangeEngine:
|
|||
decision: Final = self._classify_and_arm_locked(entry)
|
||||
match decision:
|
||||
case _Serve(token=token):
|
||||
self._report_cache_hit()
|
||||
return token
|
||||
case _ServeAndRefresh(token=token):
|
||||
self._report_cache_hit()
|
||||
self._executor_instance().submit(self._advisory_refresh, spec, entry)
|
||||
return token
|
||||
case _Fail(error=error):
|
||||
return error
|
||||
case _Lead():
|
||||
return self._lead(spec, entry)
|
||||
case _Lead(call_type=call_type):
|
||||
return self._lead(spec, entry, call_type)
|
||||
case _Follow():
|
||||
followed: Final = self._await_leader(spec, entry)
|
||||
return followed if followed is not None else self.get_token(spec)
|
||||
|
|
@ -474,7 +614,7 @@ class JwtBearerTokenExchangeEngine:
|
|||
if entry.last_error is not None and self._clock() < entry.backoff_until:
|
||||
return _Fail(error=entry.last_error)
|
||||
entry.arm()
|
||||
return _Lead()
|
||||
return _Lead(call_type=CALL_TYPE_COLD_MINT if token is None else CALL_TYPE_MANDATORY_REFRESH)
|
||||
|
||||
def _executor_instance(self) -> Executor:
|
||||
with self._lock:
|
||||
|
|
@ -484,10 +624,13 @@ class JwtBearerTokenExchangeEngine:
|
|||
)
|
||||
return self._refresh_executor
|
||||
|
||||
def _lead(self, spec: TokenExchangeSpec, entry: _Entry) -> ExchangeResult:
|
||||
def _lead(self, spec: TokenExchangeSpec, entry: _Entry, call_type: ExchangeCallType) -> ExchangeResult:
|
||||
started: Final = self._clock()
|
||||
result: Final = self._exchange_never_raises(spec)
|
||||
duration: Final = self._clock() - started
|
||||
with self._lock:
|
||||
entry.publish(result, now=self._clock())
|
||||
self._report_exchange(call_type, duration, result)
|
||||
return result
|
||||
|
||||
def _await_leader(self, spec: TokenExchangeSpec, entry: _Entry) -> "ExchangeResult | None":
|
||||
|
|
@ -505,12 +648,15 @@ class JwtBearerTokenExchangeEngine:
|
|||
return TokenTransportError(detail="timed out waiting for the token exchange leader")
|
||||
|
||||
def _advisory_refresh(self, spec: TokenExchangeSpec, entry: _Entry) -> None:
|
||||
started: Final = self._clock()
|
||||
result: Final = self._exchange_never_raises(spec)
|
||||
duration: Final = self._clock() - started
|
||||
with self._lock:
|
||||
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
|
||||
self._report_exchange(CALL_TYPE_ADVISORY_REFRESH, duration, result)
|
||||
if isinstance(result, MintedToken):
|
||||
return
|
||||
seconds_to_mandatory_wall: Final = (
|
||||
|
|
@ -525,6 +671,24 @@ class JwtBearerTokenExchangeEngine:
|
|||
ADVISORY_REFRESH_BACKOFF_SECONDS,
|
||||
)
|
||||
|
||||
def _report_exchange(self, call_type: ExchangeCallType, duration_seconds: float, result: ExchangeResult) -> None:
|
||||
try:
|
||||
match result:
|
||||
case MintedToken():
|
||||
self._metrics_sink.exchange_success(call_type=call_type, duration_seconds=duration_seconds)
|
||||
case _:
|
||||
self._metrics_sink.exchange_failure(
|
||||
call_type=call_type, duration_seconds=duration_seconds, error=result
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # metrics are best-effort; a sink failure must never fail a mint
|
||||
verbose_logger.debug("token exchange metrics emission failed: %s", e)
|
||||
|
||||
def _report_cache_hit(self) -> None:
|
||||
try:
|
||||
self._metrics_sink.cache_hit()
|
||||
except Exception as e: # noqa: BLE001 # metrics are best-effort; a sink failure must never fail a serve
|
||||
verbose_logger.debug("token exchange cache-hit metric emission failed: %s", e)
|
||||
|
||||
def _exchange_never_raises(self, spec: TokenExchangeSpec) -> ExchangeResult:
|
||||
"""The single-flight leader and the advisory refresher must always publish a result: an
|
||||
unhandled exception here would leave the entry armed (in_flight, cleared event) forever, so
|
||||
|
|
|
|||
|
|
@ -76,6 +76,22 @@ ExchangeError: TypeAlias = (
|
|||
)
|
||||
ExchangeResult: TypeAlias = MintedToken | ExchangeError
|
||||
|
||||
ExchangeCallType: TypeAlias = Literal["cold_mint", "mandatory_refresh", "advisory_refresh"]
|
||||
|
||||
|
||||
class TokenExchangeMetricsSink(Protocol):
|
||||
"""Observability seam for the exchange engine. Implementations must be best-effort: never raise
|
||||
into the mint path, never block the calling thread, and never receive credential material --
|
||||
``ExchangeError`` values are redacted by construction."""
|
||||
|
||||
def exchange_success(self, *, call_type: ExchangeCallType, duration_seconds: float) -> None: ...
|
||||
|
||||
def exchange_failure(
|
||||
self, *, call_type: ExchangeCallType, duration_seconds: float, error: ExchangeError
|
||||
) -> None: ...
|
||||
|
||||
def cache_hit(self) -> None: ...
|
||||
|
||||
|
||||
class SyncTokenPoster(Protocol):
|
||||
"""Returns the response for ANY status; never raises for status."""
|
||||
|
|
|
|||
|
|
@ -25,6 +25,8 @@ class ServiceTypes(str, enum.Enum):
|
|||
AUTH = "auth"
|
||||
PROXY_PRE_CALL = "proxy_pre_call"
|
||||
POD_LOCK_MANAGER = "pod_lock_manager"
|
||||
ANTHROPIC_WIF = "anthropic_wif"
|
||||
ANTHROPIC_WIF_CACHE = "anthropic_wif_cache"
|
||||
|
||||
"""
|
||||
Operational metrics for DB Transaction Queues
|
||||
|
|
@ -65,6 +67,13 @@ DEFAULT_SERVICE_CONFIGS: Final = {
|
|||
ServiceTypes.ROUTER.value: {"metrics": [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM]},
|
||||
ServiceTypes.AUTH.value: {"metrics": [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM]},
|
||||
ServiceTypes.PROXY_PRE_CALL.value: {"metrics": [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM]},
|
||||
ServiceTypes.ANTHROPIC_WIF.value: { # mutable-ok: ServiceConfig mandates the dict-of-list shape
|
||||
"metrics": [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM] # mutable-ok: ServiceConfig mandates a list
|
||||
},
|
||||
# cache hits are counter-only: no HTTP call happens, so observing a latency would be a lie
|
||||
ServiceTypes.ANTHROPIC_WIF_CACHE.value: { # mutable-ok: ServiceConfig mandates the dict-of-list shape
|
||||
"metrics": [ServiceMetrics.COUNTER] # mutable-ok: ServiceConfig mandates a list
|
||||
},
|
||||
# Operational metrics for DB Transaction Queues
|
||||
ServiceTypes.POD_LOCK_MANAGER.value: {"metrics": [ServiceMetrics.GAUGE]},
|
||||
ServiceTypes.IN_MEMORY_DAILY_SPEND_UPDATE_QUEUE.value: {"metrics": [ServiceMetrics.GAUGE]},
|
||||
|
|
|
|||
|
|
@ -135,3 +135,28 @@ def test_services_logger_custom_latency_buckets():
|
|||
REGISTRY.unregister(collector)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def test_anthropic_wif_services_are_wired_into_the_registry():
|
||||
"""Reverting the ANTHROPIC_WIF/ANTHROPIC_WIF_CACHE ServiceTypes members or their
|
||||
DEFAULT_SERVICE_CONFIGS entries must fail here: the exchange service gets counters plus a
|
||||
latency histogram, while the cache-hit service is counter-only so a hit can never fake a latency."""
|
||||
from litellm.types.services import DEFAULT_SERVICE_CONFIGS
|
||||
|
||||
assert ServiceTypes.ANTHROPIC_WIF.value == "anthropic_wif"
|
||||
assert ServiceTypes.ANTHROPIC_WIF_CACHE.value == "anthropic_wif_cache"
|
||||
assert DEFAULT_SERVICE_CONFIGS["anthropic_wif"]["metrics"] == [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM]
|
||||
assert DEFAULT_SERVICE_CONFIGS["anthropic_wif_cache"]["metrics"] == [ServiceMetrics.COUNTER]
|
||||
|
||||
pl = PrometheusServicesLogger()
|
||||
wif_names = {obj._name for obj in pl.payload_to_prometheus_map["anthropic_wif"]}
|
||||
assert wif_names == {
|
||||
"litellm_anthropic_wif_latency",
|
||||
"litellm_anthropic_wif_failed_requests",
|
||||
"litellm_anthropic_wif_total_requests",
|
||||
}
|
||||
cache_names = {obj._name for obj in pl.payload_to_prometheus_map["anthropic_wif_cache"]}
|
||||
assert cache_names == {
|
||||
"litellm_anthropic_wif_cache_failed_requests",
|
||||
"litellm_anthropic_wif_cache_total_requests",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -15,10 +15,14 @@ from pydantic import SecretStr
|
|||
from litellm.llms.base_llm.auth.token_exchange import (
|
||||
_REDACTION_CAP,
|
||||
ADVISORY_REFRESH_BACKOFF_SECONDS,
|
||||
CALL_TYPE_CACHE_HIT,
|
||||
FALLBACK_TOKEN_TTL_SECONDS,
|
||||
MAX_ASSERTION_BYTES,
|
||||
MAX_RESPONSE_BYTES,
|
||||
JwtBearerTokenExchangeEngine,
|
||||
ServiceLoggingMetricsSink,
|
||||
TokenExchangeEndpointFailure,
|
||||
TokenExchangeTransportFailure,
|
||||
_default_assertion_reader,
|
||||
_error_summary,
|
||||
_HttpxSyncTokenPoster,
|
||||
|
|
@ -27,6 +31,7 @@ from litellm.llms.base_llm.auth.token_exchange import (
|
|||
)
|
||||
from litellm.llms.base_llm.auth.types import (
|
||||
AssertionSourceError,
|
||||
ExchangeError,
|
||||
ExchangeResult,
|
||||
InsecureTokenUrl,
|
||||
MalformedTokenResponse,
|
||||
|
|
@ -36,6 +41,7 @@ from litellm.llms.base_llm.auth.types import (
|
|||
TokenTransportError,
|
||||
)
|
||||
from litellm.secret_managers.main import OidcPathNotAllowedError, _resolve_oidc_file_path
|
||||
from litellm.types.services import ServiceTypes
|
||||
|
||||
DEFAULT_REF: Final = "oidc/env/TEST_ASSERTION"
|
||||
DEFAULT_ASSERTION: Final = "test-jwt-assertion"
|
||||
|
|
@ -147,12 +153,29 @@ def make_spec(**overrides) -> TokenExchangeSpec:
|
|||
return TokenExchangeSpec(**base)
|
||||
|
||||
|
||||
class RecordingMetricsSink:
|
||||
def __init__(self) -> None:
|
||||
self.successes: list[tuple[str, float]] = []
|
||||
self.failures: list[tuple[str, float, ExchangeError]] = []
|
||||
self.cache_hits = 0
|
||||
|
||||
def exchange_success(self, *, call_type: str, duration_seconds: float) -> None:
|
||||
self.successes.append((call_type, duration_seconds))
|
||||
|
||||
def exchange_failure(self, *, call_type: str, duration_seconds: float, error: ExchangeError) -> None:
|
||||
self.failures.append((call_type, duration_seconds, error))
|
||||
|
||||
def cache_hit(self) -> None:
|
||||
self.cache_hits += 1
|
||||
|
||||
|
||||
def make_engine(
|
||||
poster,
|
||||
reader: Mapping[str, str] | Callable[[str], str | None] | None = None,
|
||||
clock: FakeClock | None = None,
|
||||
executor: concurrent.futures.Executor | None = None,
|
||||
max_entries: int = 64,
|
||||
metrics_sink=None,
|
||||
) -> JwtBearerTokenExchangeEngine:
|
||||
resolved_reader = reader if callable(reader) else (reader or {DEFAULT_REF: DEFAULT_ASSERTION}).get
|
||||
return JwtBearerTokenExchangeEngine(
|
||||
|
|
@ -161,6 +184,7 @@ def make_engine(
|
|||
clock=clock if clock is not None else FakeClock(),
|
||||
refresh_executor=executor if executor is not None else ManualExecutor(),
|
||||
max_entries=max_entries,
|
||||
metrics_sink=metrics_sink if metrics_sink is not None else RecordingMetricsSink(),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1316,3 +1340,214 @@ class TestShortLivedRefreshWindows:
|
|||
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)
|
||||
|
||||
|
||||
class RaisingMetricsSink:
|
||||
def exchange_success(self, *, call_type: str, duration_seconds: float) -> None:
|
||||
raise RuntimeError("metrics sink down")
|
||||
|
||||
def exchange_failure(self, *, call_type: str, duration_seconds: float, error: ExchangeError) -> None:
|
||||
raise RuntimeError("metrics sink down")
|
||||
|
||||
def cache_hit(self) -> None:
|
||||
raise RuntimeError("metrics sink down")
|
||||
|
||||
|
||||
class TestMetricsEmission:
|
||||
def test_cold_mint_emits_success_with_duration(self):
|
||||
clock = FakeClock()
|
||||
sink = RecordingMetricsSink()
|
||||
poster = ScriptedPoster([token_response()], on_request=lambda _request: clock.advance(0.25))
|
||||
engine = make_engine(poster, clock=clock, metrics_sink=sink)
|
||||
|
||||
mint(engine, make_spec())
|
||||
|
||||
assert sink.successes == [("cold_mint", 0.25)]
|
||||
assert sink.failures == []
|
||||
assert sink.cache_hits == 0
|
||||
|
||||
def test_cache_hit_emits_counter_not_a_mint(self):
|
||||
clock = FakeClock()
|
||||
sink = RecordingMetricsSink()
|
||||
engine = make_engine(ScriptedPoster([token_response()]), clock=clock, metrics_sink=sink)
|
||||
spec = make_spec()
|
||||
|
||||
mint(engine, spec)
|
||||
clock.advance(100.0)
|
||||
mint(engine, spec)
|
||||
|
||||
assert sink.cache_hits == 1
|
||||
assert len(sink.successes) == 1
|
||||
|
||||
def test_advisory_refresh_call_type(self):
|
||||
clock = FakeClock(start=1_000.0)
|
||||
sink = RecordingMetricsSink()
|
||||
executor = ManualExecutor()
|
||||
poster = ScriptedPoster([token_response("old", expires_in=3600), token_response("new")])
|
||||
engine = make_engine(poster, clock=clock, executor=executor, metrics_sink=sink)
|
||||
spec = make_spec()
|
||||
|
||||
mint(engine, spec)
|
||||
clock.now = 1_000.0 + 3600 - 119.0
|
||||
mint(engine, spec)
|
||||
executor.run_all()
|
||||
|
||||
assert [call_type for call_type, _ in sink.successes] == ["cold_mint", "advisory_refresh"]
|
||||
assert sink.cache_hits == 1
|
||||
|
||||
def test_mandatory_refresh_call_type(self):
|
||||
clock = FakeClock(start=1_000.0)
|
||||
sink = RecordingMetricsSink()
|
||||
poster = ScriptedPoster([token_response("old", expires_in=3600), token_response("new")])
|
||||
engine = make_engine(poster, clock=clock, metrics_sink=sink)
|
||||
spec = make_spec()
|
||||
|
||||
mint(engine, spec)
|
||||
clock.now = 1_000.0 + 3600 - 29.0
|
||||
mint(engine, spec)
|
||||
|
||||
assert [call_type for call_type, _ in sink.successes] == ["cold_mint", "mandatory_refresh"]
|
||||
assert sink.cache_hits == 0
|
||||
|
||||
def test_failed_exchange_emits_failure_once_and_negative_cache_does_not_reemit(self):
|
||||
sink = RecordingMetricsSink()
|
||||
poster = ScriptedPoster([httpx.Response(503, json={"error": "unavailable"})])
|
||||
engine = make_engine(poster, metrics_sink=sink)
|
||||
spec = make_spec()
|
||||
|
||||
first = engine.get_token(spec)
|
||||
second = engine.get_token(spec)
|
||||
|
||||
assert isinstance(first, TokenEndpointError)
|
||||
assert isinstance(second, TokenEndpointError)
|
||||
assert len(sink.failures) == 1
|
||||
call_type, _duration, error = sink.failures[0]
|
||||
assert call_type == "cold_mint"
|
||||
assert isinstance(error, TokenEndpointError)
|
||||
assert error.status_code == 503
|
||||
assert sink.successes == []
|
||||
|
||||
def test_failure_payload_carries_no_assertion_material(self):
|
||||
sink = RecordingMetricsSink()
|
||||
engine = make_engine(EchoingUnauthorizedPoster(), metrics_sink=sink)
|
||||
|
||||
result = engine.get_token(make_spec())
|
||||
|
||||
assert isinstance(result, TokenEndpointError)
|
||||
(failure,) = sink.failures
|
||||
assert DEFAULT_ASSERTION not in repr(failure)
|
||||
assert DEFAULT_ASSERTION not in _error_summary(failure[2])
|
||||
|
||||
def test_raising_sink_never_breaks_mint_serve_or_failure(self):
|
||||
clock = FakeClock()
|
||||
engine = make_engine(ScriptedPoster([token_response()]), clock=clock, metrics_sink=RaisingMetricsSink())
|
||||
spec = make_spec()
|
||||
|
||||
minted = mint(engine, spec)
|
||||
clock.advance(100.0)
|
||||
served = mint(engine, spec)
|
||||
|
||||
assert served.access_token.get_secret_value() == minted.access_token.get_secret_value()
|
||||
|
||||
failing = make_engine(RaisingPoster(httpx.ConnectError("boom")), metrics_sink=RaisingMetricsSink())
|
||||
result = failing.get_token(make_spec())
|
||||
assert isinstance(result, TokenTransportError)
|
||||
|
||||
|
||||
class RecordingServiceHooks:
|
||||
def __init__(self) -> None:
|
||||
self.successes: list[tuple[ServiceTypes, str, float]] = []
|
||||
self.failures: list[tuple[ServiceTypes, float, str | Exception, str]] = []
|
||||
|
||||
async def async_service_success_hook(self, service: ServiceTypes, call_type: str, duration: float) -> None:
|
||||
self.successes.append((service, call_type, duration))
|
||||
|
||||
async def async_service_failure_hook(
|
||||
self, service: ServiceTypes, duration: float, error: str | Exception, call_type: str
|
||||
) -> None:
|
||||
self.failures.append((service, duration, error, call_type))
|
||||
|
||||
|
||||
class RaisingServiceHooks:
|
||||
"""Every hook raises, and each call is recorded first so a test can prove the sink kept
|
||||
calling through rather than bailing after the first failure."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.attempts: list[str] = [] # mutable-ok: a test spy accumulating calls in order
|
||||
|
||||
async def async_service_success_hook(self, service: ServiceTypes, call_type: str, duration: float) -> None:
|
||||
self.attempts.append(f"success:{call_type}")
|
||||
raise RuntimeError("hook down")
|
||||
|
||||
async def async_service_failure_hook(
|
||||
self, service: ServiceTypes, duration: float, error: str | Exception, call_type: str
|
||||
) -> None:
|
||||
self.attempts.append(f"failure:{call_type}")
|
||||
raise RuntimeError("hook down")
|
||||
|
||||
|
||||
class TestServiceLoggingMetricsSink:
|
||||
def _sink(self, hooks) -> ServiceLoggingMetricsSink:
|
||||
return ServiceLoggingMetricsSink(service_logging_factory=lambda: hooks, executor=InlineExecutor())
|
||||
|
||||
def test_success_maps_to_anthropic_wif_service(self):
|
||||
hooks = RecordingServiceHooks()
|
||||
|
||||
self._sink(hooks).exchange_success(call_type="cold_mint", duration_seconds=0.2)
|
||||
|
||||
assert hooks.successes == [(ServiceTypes.ANTHROPIC_WIF, "cold_mint", 0.2)]
|
||||
|
||||
def test_failure_maps_variant_and_redacted_summary(self):
|
||||
hooks = RecordingServiceHooks()
|
||||
error = TokenEndpointError(status_code=503, redacted_body="error: unavailable")
|
||||
|
||||
self._sink(hooks).exchange_failure(call_type="mandatory_refresh", duration_seconds=0.1, error=error)
|
||||
|
||||
((service, duration, emitted, call_type),) = hooks.failures
|
||||
assert service is ServiceTypes.ANTHROPIC_WIF
|
||||
assert duration == 0.1
|
||||
assert call_type == "mandatory_refresh"
|
||||
assert isinstance(emitted, TokenExchangeEndpointFailure)
|
||||
assert str(emitted) == _error_summary(error)
|
||||
|
||||
def test_transport_failure_gets_its_own_error_class(self):
|
||||
hooks = RecordingServiceHooks()
|
||||
|
||||
self._sink(hooks).exchange_failure(
|
||||
call_type="advisory_refresh", duration_seconds=0.05, error=TokenTransportError(detail="ConnectError: boom")
|
||||
)
|
||||
|
||||
((_service, _duration, emitted, _call_type),) = hooks.failures
|
||||
assert isinstance(emitted, TokenExchangeTransportFailure)
|
||||
|
||||
def test_cache_hit_maps_to_cache_service_with_zero_duration(self):
|
||||
hooks = RecordingServiceHooks()
|
||||
|
||||
self._sink(hooks).cache_hit()
|
||||
|
||||
assert hooks.successes == [(ServiceTypes.ANTHROPIC_WIF_CACHE, CALL_TYPE_CACHE_HIT, 0.0)]
|
||||
|
||||
def test_end_to_end_reflected_assertion_never_reaches_the_hook(self):
|
||||
hooks = RecordingServiceHooks()
|
||||
sink = self._sink(hooks)
|
||||
engine = make_engine(EchoingUnauthorizedPoster(), metrics_sink=sink)
|
||||
|
||||
result = engine.get_token(make_spec())
|
||||
|
||||
assert isinstance(result, TokenEndpointError)
|
||||
((_service, _duration, emitted, call_type),) = hooks.failures
|
||||
assert call_type == "cold_mint"
|
||||
assert DEFAULT_ASSERTION not in str(emitted)
|
||||
assert DEFAULT_ASSERTION not in repr(emitted)
|
||||
|
||||
def test_raising_hooks_are_swallowed(self):
|
||||
hooks: Final = RaisingServiceHooks()
|
||||
sink: Final = self._sink(hooks)
|
||||
|
||||
sink.exchange_success(call_type="cold_mint", duration_seconds=0.2)
|
||||
sink.cache_hit()
|
||||
sink.exchange_failure(call_type="cold_mint", duration_seconds=0.1, error=TokenTransportError(detail="boom"))
|
||||
|
||||
assert hooks.attempts == ["success:cold_mint", "success:cache_hit", "failure:cold_mint"], (
|
||||
"every event is still handed to the hooks, and one raising hook does not stop the next"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue