From 67688e1def5a61aaf671afee3cf69362bbf5c58c Mon Sep 17 00:00:00 2001 From: derhornspieler <15236687+derhornspieler@users.noreply.github.com> Date: Sun, 23 Aug 2026 16:34:56 -0400 Subject: [PATCH] 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 --- litellm/llms/base_llm/auth/token_exchange.py | 178 ++++++++++++- litellm/llms/base_llm/auth/types.py | 16 ++ litellm/types/services.py | 9 + .../integrations/test_prometheus_services.py | 25 ++ .../llms/base_llm/auth/test_token_exchange.py | 235 ++++++++++++++++++ 5 files changed, 456 insertions(+), 7 deletions(-) diff --git a/litellm/llms/base_llm/auth/token_exchange.py b/litellm/llms/base_llm/auth/token_exchange.py index 450433bc9fa..dcd49c6c56d 100644 --- a/litellm/llms/base_llm/auth/token_exchange.py +++ b/litellm/llms/base_llm/auth/token_exchange.py @@ -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 diff --git a/litellm/llms/base_llm/auth/types.py b/litellm/llms/base_llm/auth/types.py index 21a2801bdd4..9d3a7a5012b 100644 --- a/litellm/llms/base_llm/auth/types.py +++ b/litellm/llms/base_llm/auth/types.py @@ -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.""" diff --git a/litellm/types/services.py b/litellm/types/services.py index 74f908548d5..c58b5352323 100644 --- a/litellm/types/services.py +++ b/litellm/types/services.py @@ -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]}, diff --git a/tests/test_litellm/integrations/test_prometheus_services.py b/tests/test_litellm/integrations/test_prometheus_services.py index 2303061ede8..a5e95de3ea9 100644 --- a/tests/test_litellm/integrations/test_prometheus_services.py +++ b/tests/test_litellm/integrations/test_prometheus_services.py @@ -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", + } 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 6a7f139e295..0432e7f279f 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 @@ -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" + )