Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py: 34%
163 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"""Store for the enterprise IdP identity assertion captured at SSO login (EMA).
3The ``oauth2_id_jag`` egress arm needs the user's IdP ``id_token`` as its RFC 8693
4``subject_token``. A front-door client holds an identity-only ``llm_session_`` bearer, not an
5IdP assertion, so the assertion captured at the one SSO login is the only usable subject
6source for it. This module owns both sides of that state: the SSO callback persists here
7(write-through to the DB so a login on one pod is visible to every pod) and the resolver
8seam reads back by ``user_id``. Retention is gated on an ``oauth2_id_jag`` server actually
9being registered, so a gateway with no EMA upstream never stores bearer material.
11The row is one encrypted payload per user, latest login wins. ``expires_at`` mirrors the
12id_token ``exp`` claim and is judged by the reader, never enforced by deletion here: an
13expired assertion with a refresh token is still renewable, and the DB row is the source of
14truth, the same contract as the per-user OAuth credential store. Reads use a per-process cache with
15TTL ``MCP_SSO_ASSERTION_CACHE_TTL_SECONDS``; invalidation also guards against stale in-flight reads.
16"""
18from __future__ import annotations
20import json
21from collections.abc import Mapping, Sequence
22from datetime import datetime, timezone
23from typing import TYPE_CHECKING, Final, Protocol
25import jwt
26from pydantic import BaseModel, ConfigDict, SecretStr, TypeAdapter, ValidationError
28from litellm._logging import verbose_proxy_logger
29from litellm.caching.in_memory_cache import InMemoryCache
30from litellm.constants import MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE, MCP_SSO_ASSERTION_CACHE_TTL_SECONDS
32if TYPE_CHECKING: 32 ↛ 33line 32 didn't jump to line 33 because the condition on line 32 was never true
33 from prisma.models import LiteLLM_SSOIdentityAssertion
35 from litellm.proxy.utils import PrismaClient
37_ASSERTION_DECRYPT_LOG_KEY: Final = "sso_identity_assertion"
38_STR_ADAPTER: Final[TypeAdapter[str]] = TypeAdapter(str)
39_MAYBE_STR_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
42class _SSOAssertionTable(Protocol):
43 """The ``LiteLLM_SSOIdentityAssertion`` table operations this store calls."""
45 async def find_unique(self, *, where: Mapping[str, str]) -> LiteLLM_SSOIdentityAssertion | None: ... 45 ↛ exitline 45 didn't return from function 'find_unique' because
47 async def find_many(self) -> Sequence[LiteLLM_SSOIdentityAssertion]: ... 47 ↛ exitline 47 didn't return from function 'find_many' because
49 async def upsert(self, *, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]) -> object: ... 49 ↛ exitline 49 didn't return from function 'upsert' because
51 async def update(self, *, where: Mapping[str, str], data: Mapping[str, str]) -> object: ... 51 ↛ exitline 51 didn't return from function 'update' because
54class _MCPServerTable(Protocol):
55 """The ``LiteLLM_MCPServerTable`` lookup the retention gate calls."""
57 async def find_first(self, *, where: Mapping[str, str]) -> object | None: ... 57 ↛ exitline 57 didn't return from function 'find_first' because
60def _assertion_table(prisma_client: PrismaClient) -> _SSOAssertionTable:
61 """The SSO assertion table, typed so the untyped prisma client surface stops here."""
62 return prisma_client.db.litellm_ssoidentityassertion
65def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable:
66 """The MCP server table, typed so the untyped prisma client surface stops here."""
67 return prisma_client.db.litellm_mcpservertable
70class SSOIdentityAssertion(BaseModel):
71 """The IdP material an EMA exchange needs: ``id_token`` is the RFC 8693 subject token,
72 ``expires_at`` bounds its usefulness, and the refresh token renews it without re-login."""
74 model_config = ConfigDict(frozen=True)
76 id_token: SecretStr
77 refresh_token: SecretStr | None = None
78 issuer: str | None = None
79 expires_at: datetime | None = None
82class SSOAssertionCache:
83 """Process-local read cache. ``invalidate`` bumps a process-wide epoch so a fetch that started
84 before a login cannot repopulate the old assertion after it."""
86 def __init__(self, ttl_seconds: int = MCP_SSO_ASSERTION_CACHE_TTL_SECONDS) -> None:
87 self._entries = InMemoryCache(
88 max_size_in_memory=MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE,
89 default_ttl=ttl_seconds,
90 )
91 self._epoch: int = 0
93 def epoch(self) -> int:
94 return self._epoch
96 def get(self, user_id: str) -> SSOIdentityAssertion | None:
97 cached: Final = self._entries.get_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped
98 user_id
99 )
100 return cached if isinstance(cached, SSOIdentityAssertion) else None
102 def set_if_unchanged(self, user_id: str, assertion: SSOIdentityAssertion, seen_epoch: int) -> None:
103 if self._epoch != seen_epoch:
104 return
105 self._entries.set_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped
106 user_id, assertion
107 )
109 def invalidate(self, user_id: str) -> None:
110 self._epoch += 1
111 self._entries.delete_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped
112 user_id
113 )
115 def flush(self) -> None:
116 self._entries.flush_cache() # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped
119_ASSERTION_CACHE: Final = SSOAssertionCache()
122class _IdTokenClaims(BaseModel):
123 exp: float | None = None
124 iss: str | None = None
127class _StoredAssertionPayload(BaseModel):
128 id_token: str
129 refresh_token: str | None = None
130 issuer: str | None = None
131 expires_at: datetime | None = None
134def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIdentityAssertion | None:
135 """The typed carrier built where the raw token response exists; ``None`` when the provider
136 sent no id_token or sent one that is not a decodable JWT, since neither is exchangeable
137 under EMA. Inputs are ``object`` because they come straight from the provider's untyped
138 token response; this is the one boundary that validates them. The token arrived over TLS
139 from the IdP's own token endpoint, so claims are read without signature verification,
140 matching how the SSO callback already decodes it for identity."""
141 raw_id_token: Final = id_token if isinstance(id_token, str) and id_token else None
142 if raw_id_token is None:
143 return None
144 raw_refresh_token: Final = refresh_token if isinstance(refresh_token, str) and refresh_token else None
145 try:
146 claims: Final = _IdTokenClaims.model_validate(jwt.decode(raw_id_token, options={"verify_signature": False}))
147 expires_at: Final = datetime.fromtimestamp(claims.exp, tz=timezone.utc) if claims.exp is not None else None
148 except Exception: # noqa: BLE001 # decode failure = not retainable; never raise into login
149 verbose_proxy_logger.warning(
150 "SSO id_token could not be decoded or its claims were unusable; not retaining it for EMA egress."
151 )
152 return None
153 return SSOIdentityAssertion(
154 id_token=SecretStr(raw_id_token),
155 refresh_token=SecretStr(raw_refresh_token) if raw_refresh_token else None,
156 issuer=claims.iss,
157 expires_at=expires_at,
158 )
161def assertion_expired(assertion: SSOIdentityAssertion, now: datetime) -> bool:
162 """Whether the assertion's ``exp`` has passed at ``now``. An assertion carrying no expiry is
163 treated as usable and left for the IdP to reject, since the store records what the id_token
164 claimed rather than imposing a lifetime of its own. A naive ``expires_at`` is read as UTC so a
165 stored value that lost its offset compares instead of raising.
167 Lives beside the model rather than in either reader so the egress guard and the renewal
168 trigger judge the same field the same way; passing a ``now`` in the future is how a caller
169 asks "is this about to expire" without a second, driftable predicate.
170 """
171 expires_at: Final = assertion.expires_at
172 if expires_at is None:
173 return False
174 normalized: Final = expires_at if expires_at.tzinfo is not None else expires_at.replace(tzinfo=timezone.utc)
175 return normalized <= now
178async def ema_assertion_retention_enabled() -> bool:
179 """Whether any MCP server uses ``oauth2_id_jag``, evaluated per login so the gateway only
180 retains bearer material while an EMA upstream exists to spend it on. Judged against the two
181 configuration authorities: the pod-local config declaration and the shared DB row. The
182 in-memory registry is deliberately not consulted in either direction; it is a per-process
183 snapshot of the DB state that can be stale both ways (a server added on another pod would
184 silently drop the write, one removed on another pod would keep retaining bearer material),
185 and a gate guarding a shared-DB write must judge against that storage's authority."""
186 from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids import cycle
187 global_mcp_server_manager,
188 )
189 from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
190 from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global
192 config_servers: Final = global_mcp_server_manager.config_mcp_servers.values()
193 if any(server.auth_type == MCPAuth.oauth2_id_jag for server in config_servers):
194 return True
195 if prisma_client is None:
196 return False
197 row: Final = await _mcp_server_table(prisma_client).find_first(where={"auth_type": MCPAuth.oauth2_id_jag.value})
198 return row is not None
201async def persist_sso_identity_assertion(
202 user_id: str, assertion: SSOIdentityAssertion, cache: SSOAssertionCache = _ASSERTION_CACHE
203) -> None:
204 from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper # noqa: PLC0415 # runtime global
205 from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
207 if prisma_client is None:
208 return
209 payload: Final = {
210 "id_token": assertion.id_token.get_secret_value(),
211 **({"refresh_token": assertion.refresh_token.get_secret_value()} if assertion.refresh_token else {}),
212 **({"issuer": assertion.issuer} if assertion.issuer else {}),
213 **({"expires_at": assertion.expires_at.isoformat()} if assertion.expires_at else {}),
214 }
215 encoded: Final = _STR_ADAPTER.validate_python(encrypt_value_helper(json.dumps(payload)))
216 await _assertion_table(prisma_client).upsert(
217 where={"user_id": user_id},
218 data={
219 "create": {"user_id": user_id, "assertion_b64": encoded},
220 "update": {"assertion_b64": encoded},
221 },
222 )
223 cache.invalidate(user_id)
226async def _read_assertion_from_db(user_id: str) -> SSOIdentityAssertion | None:
227 from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper # noqa: PLC0415 # runtime global
228 from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
230 if prisma_client is None:
231 return None
232 row: Final = await _assertion_table(prisma_client).find_unique(where={"user_id": user_id})
233 if row is None:
234 return None
235 raw: Final = _MAYBE_STR_ADAPTER.validate_python(
236 decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
237 )
238 if raw is None:
239 return None
240 try:
241 payload: Final = _StoredAssertionPayload.model_validate_json(raw)
242 except ValidationError:
243 verbose_proxy_logger.warning(
244 "Stored SSO identity assertion for user_id=%s could not be parsed; treating as absent.", user_id
245 )
246 return None
247 return SSOIdentityAssertion(
248 id_token=SecretStr(payload.id_token),
249 refresh_token=SecretStr(payload.refresh_token) if payload.refresh_token else None,
250 issuer=payload.issuer,
251 expires_at=payload.expires_at,
252 )
255async def fetch_sso_identity_assertion(
256 user_id: str, cache: SSOAssertionCache = _ASSERTION_CACHE
257) -> SSOIdentityAssertion | None:
258 """The stored assertion for ``user_id``, or ``None`` when absent, undecryptable (salt-key
259 rotation), or unparseable. Expiry is not judged here; the reader owns that policy."""
260 cached: Final = cache.get(user_id)
261 if cached is not None:
262 return cached
263 seen_epoch: Final = cache.epoch()
264 assertion: Final = await _read_assertion_from_db(user_id)
265 if assertion is not None:
266 cache.set_if_unchanged(user_id, assertion, seen_epoch)
267 return assertion
270class AssertionStoreUnavailable(Exception):
271 """Raised by ``fetch`` when the assertion cannot be read for a transient reason: the DB is
272 down, or the IdP behind a renewing store could not be reached.
274 Distinct from returning ``None`` for "this user has no captured assertion": an outage must not
275 read as a definite absence, which would tell the user to sign in again over a transient failure,
276 and it must not escape as an unhandled error on the egress or retry path. The message names the
277 real component for the operator log; callers get the reader's generic 503. Mirrors
278 ``TokenStoreUnavailable`` on the sibling per-user OAuth store.
279 """
282class SSOAssertionStore(Protocol):
283 """The read seam the ``id_jag`` egress arm depends on, so the arm takes a collaborator
284 rather than reaching for a module-level function and a proxy global at call time.
286 Returns the user's captured assertion, or ``None`` when they have never signed in. Raises
287 ``AssertionStoreUnavailable`` when the backing store is unreachable.
288 """
290 async def fetch(self, user_id: str) -> SSOIdentityAssertion | None: ... 290 ↛ exitline 290 didn't return from function 'fetch' because
293class DbSSOAssertionStore:
294 """The live store: the row the SSO callback wrote, read back by ``user_id``.
296 A storage failure is re-raised as ``AssertionStoreUnavailable`` so the resolver can map it to a
297 typed fail-closed result; letting the raw driver error escape would surface a DB blip as a 500
298 from credential resolution and from the upstream-401 retry.
299 """
301 def __init__(self, cache: SSOAssertionCache = _ASSERTION_CACHE) -> None:
302 self._cache = cache
304 async def fetch(self, user_id: str) -> SSOIdentityAssertion | None:
305 try:
306 return await fetch_sso_identity_assertion(user_id, cache=self._cache)
307 except Exception as exc: # noqa: BLE001 # any driver/storage failure is an outage, not an absence
308 raise AssertionStoreUnavailable(str(exc)) from exc
310 async def fetch_uncached(self, user_id: str) -> SSOIdentityAssertion | None:
311 try:
312 return await _read_assertion_from_db(user_id)
313 except Exception as exc: # noqa: BLE001 # any driver/storage failure is an outage, not an absence
314 raise AssertionStoreUnavailable(str(exc)) from exc
317async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:
318 """Re-encrypt every stored assertion under ``new_master_key`` during a salt-key rotation,
319 mirroring the sibling per-user credential tables; an unreadable row is skipped so one
320 corrupt row does not abort the rotation. Rows are decrypted one at a time inside the loop
321 so the whole table's plaintext is never held in memory at once."""
322 from prisma.models import LiteLLM_SSOIdentityAssertion as AssertionRow # noqa: PLC0415 # generated at runtime
324 from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 # runtime global
325 decrypt_value_helper,
326 encrypt_value_helper,
327 )
329 async def _rotate_row(row: AssertionRow) -> bool:
330 plaintext: Final = _MAYBE_STR_ADAPTER.validate_python(
331 decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
332 )
333 if plaintext is None:
334 verbose_proxy_logger.warning(
335 "rotate_sso_identity_assertions_master_key: could not decrypt assertion for user_id=%s, skipping",
336 row.user_id,
337 )
338 return False
339 re_encrypted: Final = _STR_ADAPTER.validate_python(
340 encrypt_value_helper(plaintext, new_encryption_key=new_master_key)
341 )
342 await _assertion_table(prisma_client).update(
343 where={"user_id": row.user_id},
344 data={"assertion_b64": re_encrypted},
345 )
346 return True
348 rows: Final = await _assertion_table(prisma_client).find_many()
349 outcomes: Final = [await _rotate_row(row) for row in rows]
350 verbose_proxy_logger.info(
351 "rotate_sso_identity_assertions_master_key: rotated %d row(s), skipped %d",
352 sum(outcomes),
353 len(outcomes) - sum(outcomes),
354 )
357async def retain_sso_identity_assertion_for_ema(user_id: str, assertion: SSOIdentityAssertion | None) -> None:
358 """The SSO-callback hook: a no-op unless there is material AND an EMA server is registered.
359 A store failure is logged and swallowed because the login itself must not fail on an
360 egress-side write; the cost of a miss is a 401 challenge at the EMA upstream, not a lockout."""
361 if assertion is None:
362 return
363 try:
364 if not await ema_assertion_retention_enabled():
365 return
366 await persist_sso_identity_assertion(user_id, assertion)
367 except Exception as exc: # noqa: BLE001 # the login itself must not fail on an egress-side write
368 verbose_proxy_logger.warning(
369 "Failed to persist the SSO identity assertion for EMA egress (user_id=%s): %s", user_id, exc
370 )