Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_refresher.py: 40%
170 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Renew the stored SSO identity assertion so an ID-JAG agent outlives one id_token.
3The ``oauth2_id_jag`` arm asserts the id_token captured at the user's last interactive sign-in, so
4without renewal an agent holding a brokered LiteLLM key can act for that user only until that token's
5``exp``, typically an hour, and the sole recovery is another interactive login. The assertion already
6carries the IdP refresh token beside it; this module is what redeems it.
8``RefreshingSSOAssertionStore`` wraps any ``SSOAssertionStore`` and satisfies the same protocol, so
9the egress arm is unchanged: it still reads one assertion and still judges expiry itself. Renewal is
10lazy (only a read that finds a near-expiry assertion triggers one, so IdP traffic tracks actual use,
11not the size of the user table) and single-flighted per user through the same ``RefreshCoordinator``
12the ``authorization_code`` arm uses, because an IdP that rotates refresh tokens treats two concurrent
13redemptions of one token as replay and can revoke the whole grant chain.
15The refresh is redeemed against the generic-OIDC client the login itself used
16(``GENERIC_TOKEN_ENDPOINT`` / ``GENERIC_CLIENT_ID`` / ``GENERIC_CLIENT_SECRET``, which the proxy
17reconciles from the stored SSO row into the process environment at startup), authenticated the way
18that login authenticated: the non-PKCE path always sends HTTP Basic, while the PKCE path sends the
19credentials in the body when ``GENERIC_INCLUDE_CLIENT_ID`` is set, and an IdP application may accept
20only one of the two. An assertion can only exist if that client minted it, so no other client could
21redeem its refresh token, and no other method is known to be accepted. A deployment whose
22``GENERIC_SCOPE`` omits ``offline_access`` captures no refresh token at all, which is why that miss
23logs the scope by name rather than failing silently.
25Failures are values internally (``Result[_, RefreshFailure]``). At the store boundary they collapse
26onto the protocol's existing two-outcome contract: a refusal returns the expired assertion unchanged
27so the reader's own guard challenges the user to sign in again, while a transient IdP failure raises
28``AssertionStoreUnavailable`` so the reader answers 503 instead of blaming the user for an outage.
30One ambiguity remains under Redis-coordinated renewal across replicas. A cross-replica loser that
31finds the row still expiring after the holder finished cannot tell a refused refresh from a renewal
32that could not be recorded. Redeeming itself could consume a refresh token the holder may already
33have rotated, so it answers retryable 503 rather than guessing a sign-in challenge. The next
34uncontended read settles the outcome itself: a refusal challenges, and a successful refresh persists.
35If the holder rotated the token but its write failed, that rotation is lost and the next uncontended
36read's refusal challenges, which is the only honest answer because the rotated token was never
37recorded. On the refusal path, the loser pays for one retry before that challenge.
38"""
40from __future__ import annotations
42import json
43import os
44from collections.abc import Callable, Mapping
45from dataclasses import dataclass
46from datetime import datetime, timedelta, timezone
47from typing import Final, Literal, Protocol
49import httpx
50from pydantic import SecretStr, TypeAdapter, ValidationError
51from typing_extensions import assert_never
53from litellm._logging import verbose_proxy_logger
54from litellm.exceptions import Timeout
55from litellm.llms.custom_httpx.http_handler import (
56 get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # litellm http handler is untyped
57)
58from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
59 build_token_endpoint_client_auth,
60)
61from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
62 InProcessRefreshCoordinator,
63 RefreshCoordinator,
64)
65from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
66 Error,
67 Ok,
68 Result,
69)
70from litellm.proxy._experimental.mcp_server.outbound_credentials.runtime_refresh_coordinator import (
71 runtime_refresh_coordinator,
72)
73from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
74 AssertionStoreUnavailable,
75 DbSSOAssertionStore,
76 SSOAssertionStore,
77 SSOIdentityAssertion,
78 assertion_expired,
79 assertion_from_sso_login,
80 fetch_sso_identity_assertion,
81 persist_sso_identity_assertion,
82)
83from litellm.types.llms.custom_http import httpxSpecialProvider
84from litellm.types.mcp import MCPTokenEndpointAuthMethod
86_BODY_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(dict[str, object])
88_REFRESH_GRANT_TYPE: Final = "refresh_token"
89# The lock namespace for the one assertion row a user has; the sibling arm keys the same lock by
90# server_id, and no server_id can collide with this literal.
91_SINGLE_FLIGHT_KEY: Final = "sso_identity_assertion"
92# Renew this far ahead of ``exp`` so a token that would die between resolution and the second leg of
93# the exchange is replaced first. Matches the sibling per-user token store's skew.
94_DEFAULT_EXPIRY_SKEW_SECONDS: Final = 60.0
97class AssertionRead(Protocol):
98 """Reads the user's stored assertion row."""
100 async def __call__(self, user_id: str) -> SSOIdentityAssertion | None: ... 100 ↛ exitline 100 didn't return from function '__call__' because
103class AssertionWrite(Protocol):
104 """Replaces the user's stored assertion row."""
106 async def __call__(self, user_id: str, assertion: SSOIdentityAssertion) -> None: ... 106 ↛ exitline 106 didn't return from function '__call__' because
109class CoordinatorFactory(Protocol):
110 """Builds the cross-replica coordinator, or ``None`` when there is no shared lock to build on."""
112 def __call__(self) -> RefreshCoordinator | None: ... 112 ↛ exitline 112 didn't return from function '__call__' because
115class FormPost(Protocol):
116 """POSTs an OAuth form and hands back the raw response."""
118 async def __call__( 118 ↛ exitline 118 didn't return from function '__call__' because
119 self, url: str, form: Mapping[str, str], headers: Mapping[str, str]
120 ) -> httpx.Response | None: ...
123@dataclass(frozen=True, slots=True)
124class SSOClientConfig:
125 """The generic-OIDC client credentials a refresh_token grant has to authenticate as, and how."""
127 token_endpoint: str
128 client_id: str
129 client_secret: SecretStr
130 auth_method: MCPTokenEndpointAuthMethod
133def sso_client_config(env: Mapping[str, str]) -> SSOClientConfig | None:
134 """The configured generic-OIDC client, or ``None`` when the deployment has none.
136 Read from the process environment because that is where the login path reads it
137 (``_setup_generic_sso_env_vars``) and where the proxy materializes the stored ``sso_config`` row
138 at startup, so this resolves to the same client that minted the assertion. ``None`` is an
139 ordinary state, not an error: a deployment signing in through a provider that captures no
140 assertion has nothing here to renew, and a client with no secret is not a confidential client
141 that could redeem one.
143 ``auth_method`` is derived from the same ``GENERIC_INCLUDE_CLIENT_ID`` the login reads, because
144 the two login paths do not agree: the non-PKCE path always authenticates with HTTP Basic, while
145 the PKCE path puts the credentials in the body when that flag is set. Both capture assertions, so
146 a constant here would authenticate the renewal differently from the sign-in that produced the
147 refresh token and 401 against an IdP application registered for only one of the two.
148 """
149 token_endpoint: Final = env.get("GENERIC_TOKEN_ENDPOINT")
150 client_id: Final = env.get("GENERIC_CLIENT_ID")
151 client_secret: Final = env.get("GENERIC_CLIENT_SECRET")
152 if not token_endpoint or not client_id or not client_secret:
153 return None
154 includes_client_id: Final = env.get("GENERIC_INCLUDE_CLIENT_ID", "false").lower() == "true"
155 return SSOClientConfig(
156 token_endpoint=token_endpoint,
157 client_id=client_id,
158 client_secret=SecretStr(client_secret),
159 auth_method="client_secret_post" if includes_client_id else "client_secret_basic",
160 )
163@dataclass(frozen=True, slots=True)
164class RefreshFailure:
165 """Why a renewal produced nothing, split by what the caller can do about it.
167 ``rejected`` is settled: this refresh token will never work again, so the user has to sign in.
168 ``unavailable`` is transient: the same attempt may succeed in a minute, so telling the user to
169 sign in again would be a lie about whose problem it is. Both arms carry the same payload, so
170 this is a ``Literal`` discriminant rather than a ``tagged_union``; consumers still ``match`` on
171 ``kind`` with an ``assert_never`` tail.
172 """
174 kind: Literal["rejected", "unavailable"]
175 detail: str
177 @staticmethod
178 def of_rejected(detail: str) -> RefreshFailure:
179 return RefreshFailure(kind="rejected", detail=detail)
181 @staticmethod
182 def of_unavailable(detail: str) -> RefreshFailure:
183 return RefreshFailure(kind="unavailable", detail=detail)
186class TokenEndpointTransport(Protocol):
187 """One form POST to the IdP token endpoint, with the refusal/outage split preserved.
189 That split is the whole reason this is not the resolver's ``TokenEndpointClient``: that
190 collaborator maps every non-2xx to ``upstream_unavailable``, which is right for an exchange leg
191 and wrong here, where a 400 ``invalid_grant`` means the stored refresh token is dead and the user
192 must act.
193 """
195 async def post( 195 ↛ exitline 195 didn't return from function 'post' because
196 self, url: str, form: Mapping[str, str], headers: Mapping[str, str]
197 ) -> Result[Mapping[str, object], RefreshFailure]: ...
200async def post_form(url: str, form: Mapping[str, str], headers: Mapping[str, str]) -> httpx.Response | None:
201 # litellm's httpx handler is only partially typed; nothing but the response object crosses back,
202 # and the transport below validates its body, so the untyped boundary is contained here.
203 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore[reportUnknownVariableType] # litellm http handler is untyped
204 return await client.post(url, data=form, headers=headers) # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType,reportReturnType,reportArgumentType] # litellm http handler is untyped and its stub narrows data=/headers= to dict, which httpx itself does not require
207class HttpxTokenEndpointTransport:
208 """The live transport. 4xx is the IdP refusing this grant; anything else is an outage.
210 The POST itself is injected so that split, which decides whether the user is challenged or told
211 to wait, is testable without a live IdP.
212 """
214 def __init__(self, post: FormPost = post_form) -> None:
215 self._post = post
217 async def post(
218 self, url: str, form: Mapping[str, str], headers: Mapping[str, str]
219 ) -> Result[Mapping[str, object], RefreshFailure]:
220 try:
221 response: Final = await self._post(url, form, headers)
222 if response is None:
223 return Error(RefreshFailure.of_unavailable("the IdP token endpoint returned no response"))
224 response.raise_for_status()
225 body: Final = _BODY_ADAPTER.validate_python(response.json()) # pyright: ignore[reportAny] # untyped JSON; the adapter is the type gate
226 except httpx.HTTPStatusError as exc:
227 status: Final = exc.response.status_code
228 if 400 <= status < 500:
229 return Error(RefreshFailure.of_rejected(f"the IdP refused the refresh with status {status}"))
230 return Error(RefreshFailure.of_unavailable(f"the IdP token endpoint answered with status {status}"))
231 except (httpx.RequestError, Timeout) as exc:
232 return Error(RefreshFailure.of_unavailable(f"the IdP token endpoint is unreachable ({type(exc).__name__})"))
233 except json.JSONDecodeError:
234 return Error(RefreshFailure.of_unavailable("the IdP token endpoint returned a non-JSON response"))
235 except ValidationError:
236 return Error(RefreshFailure.of_unavailable("the IdP token endpoint returned a non-object response"))
237 return Ok(body)
240class SSOAssertionRefresher:
241 """Redeems the stored refresh token for a current id_token and writes the rotation back.
243 Collaborators are injected so the orchestration, the untyped response parsing and the
244 write-back race are all testable without an IdP or a database.
245 """
247 def __init__(
248 self,
249 transport: TokenEndpointTransport,
250 *,
251 client_config: Callable[[], SSOClientConfig | None] = lambda: sso_client_config(os.environ),
252 read: AssertionRead = fetch_sso_identity_assertion,
253 write: AssertionWrite = persist_sso_identity_assertion,
254 ) -> None:
255 self._transport = transport
256 self._client_config = client_config
257 self._read = read
258 self._write = write
260 async def refresh(
261 self, user_id: str, assertion: SSOIdentityAssertion
262 ) -> Result[SSOIdentityAssertion, RefreshFailure]:
263 if assertion.refresh_token is None:
264 verbose_proxy_logger.warning(
265 "ID-JAG: the stored IdP identity assertion for user_id=%s has expired and no refresh token was "
266 "captured with it, so it cannot be renewed without another interactive sign-in. Add "
267 "'offline_access' to GENERIC_SCOPE so the SSO login captures one.",
268 user_id,
269 )
270 return Error(RefreshFailure.of_rejected("no refresh token was captured at sign-in"))
271 config: Final = self._client_config()
272 if config is None:
273 verbose_proxy_logger.warning(
274 "ID-JAG: the stored IdP identity assertion for user_id=%s has expired and cannot be renewed "
275 "because the generic SSO client is not configured (GENERIC_TOKEN_ENDPOINT, GENERIC_CLIENT_ID, "
276 "GENERIC_CLIENT_SECRET).",
277 user_id,
278 )
279 return Error(RefreshFailure.of_rejected("the generic SSO client is not configured"))
281 carried_refresh_token: Final = assertion.refresh_token.get_secret_value()
282 # Whichever method the SSO login used for this client, since that is the one the IdP
283 # application is known to accept: an assertion only exists to renew because a sign-in already
284 # authenticated this client that way.
285 client_auth: Final = build_token_endpoint_client_auth(
286 auth_method=config.auth_method,
287 client_id=config.client_id,
288 client_secret=config.client_secret.get_secret_value(),
289 )
290 form: Final = { # mutable-ok: the RFC 6749 form body is a wire format the HTTP client takes as a mapping
291 "grant_type": _REFRESH_GRANT_TYPE,
292 "refresh_token": carried_refresh_token,
293 **client_auth.body,
294 }
295 match await self._transport.post(config.token_endpoint, form, client_auth.headers):
296 case Error(failure):
297 return Error(failure)
298 case Ok(body):
299 return await self._renewed_from(user_id, assertion, body, carried_refresh_token)
301 async def _renewed_from(
302 self,
303 user_id: str,
304 previous: SSOIdentityAssertion,
305 body: Mapping[str, object],
306 carried_refresh_token: str,
307 ) -> Result[SSOIdentityAssertion, RefreshFailure]:
308 """The renewed assertion, built by the same validator the login path uses.
310 A rotated refresh token replaces the stored one; an omitted one carries forward, since an
311 IdP that does not rotate expects the original to keep working.
312 """
313 rotated: Final = body.get("refresh_token")
314 renewed: Final = assertion_from_sso_login(
315 body.get("id_token"),
316 rotated if isinstance(rotated, str) and rotated else carried_refresh_token,
317 )
318 if renewed is None:
319 verbose_proxy_logger.warning(
320 "ID-JAG: the IdP accepted the refresh for user_id=%s but returned no usable id_token, so there "
321 "is nothing to assert upstream. The SSO client's grant needs the 'openid' scope for the token "
322 "endpoint to return one on a refresh.",
323 user_id,
324 )
325 return Error(RefreshFailure.of_rejected("the IdP's refresh response carried no usable id_token"))
326 failure: Final = await self._store_renewal(user_id, previous, renewed)
327 if failure is not None:
328 return Error(failure)
329 return Ok(renewed)
331 async def _store_renewal(
332 self, user_id: str, previous: SSOIdentityAssertion, renewed: SSOIdentityAssertion
333 ) -> RefreshFailure | None:
334 """Write the renewal back, unless the row moved on while this renewal was in flight.
336 The row is one per user and last-write-wins, so an interactive sign-in landing mid-renewal
337 would otherwise be overwritten with a refresh token the IdP has already rotated away, costing
338 that user a sign-in later. Comparing against the id_token this renewal started from is what
339 detects that; skipping is safe because the newer row is the one the reader wants anyway.
341 A failed write is transient, not settled. The store, not this return value, is what every
342 caller reads, so a renewal that could not be recorded is a renewal nobody will see; saying so
343 keeps a database problem answering 503 rather than telling the user to sign in again over it.
344 """
345 try:
346 current: Final = await self._read(user_id)
347 if current is not None and current.id_token.get_secret_value() != previous.id_token.get_secret_value():
348 verbose_proxy_logger.info(
349 "ID-JAG: a newer IdP identity assertion for user_id=%s was stored while this renewal was in "
350 "flight; keeping the stored one.",
351 user_id,
352 )
353 return None
354 await self._write(user_id, renewed)
355 except Exception as exc: # noqa: BLE001 # any storage failure is transient here, never the user's fault
356 verbose_proxy_logger.warning(
357 "ID-JAG: could not persist the renewed IdP identity assertion for user_id=%s, so the rotated "
358 "refresh token is lost and this user will have to sign in again once the renewed token expires: %s",
359 user_id,
360 exc,
361 )
362 return RefreshFailure.of_unavailable("the renewed IdP identity assertion could not be persisted")
363 return None
366class RefreshingSSOAssertionStore:
367 """An ``SSOAssertionStore`` that renews a near-expiry assertion before handing it back.
369 Reads the inner store; an assertion still comfortably inside its lifetime is returned untouched,
370 so the common path costs exactly what it did before. Otherwise one renewal runs per user through
371 the injected ``RefreshCoordinator`` and every caller then re-reads the inner store, which is the
372 authority: the winner's write is what they all observe, and a renewal the write-back guard
373 skipped yields the newer assertion that displaced it rather than a private copy.
375 A refusal leaves the expired assertion in place for the reader's own guard to reject, so the user
376 sees the same sign-in-again challenge as before this store existed. A transient IdP failure
377 raises ``AssertionStoreUnavailable``, the protocol's existing signal for "this is not the user's
378 fault"; concurrent in-process callers share that outcome, while a cross-replica loser answers 503
379 when its re-read still finds the row expiring. On the refusal path that costs the loser one retry,
380 which then challenges. If the holder rotated the token but its write failed, the rotation is lost
381 and the next uncontended read's refusal challenges, the only honest answer because that token was
382 never recorded.
383 """
385 def __init__(
386 self,
387 inner: SSOAssertionStore,
388 refresher: SSOAssertionRefresher,
389 *,
390 fresh_read: AssertionRead,
391 coordinator_factory: CoordinatorFactory = runtime_refresh_coordinator,
392 expiry_skew_seconds: float = _DEFAULT_EXPIRY_SKEW_SECONDS,
393 clock: Callable[[], datetime] = lambda: datetime.now(timezone.utc),
394 ) -> None:
395 self._inner = inner
396 self._refresher = refresher
397 self._fresh_read = fresh_read
398 self._coordinator_factory = coordinator_factory
399 self._in_process_coordinator = InProcessRefreshCoordinator()
400 self._distributed_coordinator: RefreshCoordinator | None = None
401 self._skew = timedelta(seconds=expiry_skew_seconds)
402 self._clock = clock
404 async def fetch(self, user_id: str) -> SSOIdentityAssertion | None:
405 assertion: Final = await self._inner.fetch(user_id)
406 if not self._expiring(assertion):
407 return assertion
408 await self._coordinator().run(
409 user_id,
410 _SINGLE_FLIGHT_KEY,
411 refresh=lambda: self._renew(user_id),
412 reread=lambda: self._reread_renewed(user_id),
413 )
414 return await self._fresh_read(user_id)
416 def _expiring(self, assertion: SSOIdentityAssertion | None) -> bool:
417 return assertion is not None and assertion_expired(assertion, self._clock() + self._skew)
419 def _coordinator(self) -> RefreshCoordinator:
420 """The cross-replica coordinator once Redis is reachable, else the in-process one.
422 Built on first use and kept, because the proxy's Redis client is not wired at import time;
423 retried while it is absent so a proxy that gains Redis later stops electing per-worker.
424 """
425 if self._distributed_coordinator is None:
426 self._distributed_coordinator = self._coordinator_factory()
427 return self._distributed_coordinator or self._in_process_coordinator
429 async def _renew(self, user_id: str) -> None:
430 """The elected renewal, judged from a fresh read so a rotation another replica just landed is
431 never redeemed again. Returns nothing: the inner store, not this return value, is what every
432 caller reads afterwards, so the winner and the losers cannot disagree."""
433 latest: Final = await self._fresh_read(user_id)
434 if latest is None or not self._expiring(latest):
435 return
436 match await self._refresher.refresh(user_id, latest):
437 case Ok(_):
438 return
439 case Error(failure):
440 match failure.kind:
441 case "rejected":
442 return
443 case "unavailable":
444 raise AssertionStoreUnavailable(failure.detail)
445 assert_never(failure.kind)
447 async def _reread_renewed(self, user_id: str) -> None:
448 """A loser cannot distinguish refusal from an unrecorded renewal without risking token replay.
450 It answers retryable 503 instead of guessing a sign-in challenge; the retry runs uncontended
451 and settles the outcome itself.
452 """
453 latest: Final = await self._fresh_read(user_id)
454 if self._expiring(latest):
455 raise AssertionStoreUnavailable(
456 f"the IdP identity assertion for user_id={user_id} was being renewed by another replica "
457 "and is not yet current; retry shortly"
458 )
461def default_sso_assertion_store() -> SSOAssertionStore:
462 """The live read seam for the ``id_jag`` arm: the stored assertion, renewed when it is stale."""
463 db_store: Final = DbSSOAssertionStore()
464 fresh_read: Final = db_store.fetch_uncached
465 return RefreshingSSOAssertionStore(
466 db_store,
467 SSOAssertionRefresher(HttpxTokenEndpointTransport(), read=fresh_read),
468 fresh_read=fresh_read,
469 )