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

1"""Store for the enterprise IdP identity assertion captured at SSO login (EMA). 

2 

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. 

10 

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""" 

17 

18from __future__ import annotations 

19 

20import json 

21from collections.abc import Mapping, Sequence 

22from datetime import datetime, timezone 

23from typing import TYPE_CHECKING, Final, Protocol 

24 

25import jwt 

26from pydantic import BaseModel, ConfigDict, SecretStr, TypeAdapter, ValidationError 

27 

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 

31 

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 

34 

35 from litellm.proxy.utils import PrismaClient 

36 

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) 

40 

41 

42class _SSOAssertionTable(Protocol): 

43 """The ``LiteLLM_SSOIdentityAssertion`` table operations this store calls.""" 

44 

45 async def find_unique(self, *, where: Mapping[str, str]) -> LiteLLM_SSOIdentityAssertion | None: ... 45 ↛ exitline 45 didn't return from function 'find_unique' because

46 

47 async def find_many(self) -> Sequence[LiteLLM_SSOIdentityAssertion]: ... 47 ↛ exitline 47 didn't return from function 'find_many' because

48 

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

50 

51 async def update(self, *, where: Mapping[str, str], data: Mapping[str, str]) -> object: ... 51 ↛ exitline 51 didn't return from function 'update' because

52 

53 

54class _MCPServerTable(Protocol): 

55 """The ``LiteLLM_MCPServerTable`` lookup the retention gate calls.""" 

56 

57 async def find_first(self, *, where: Mapping[str, str]) -> object | None: ... 57 ↛ exitline 57 didn't return from function 'find_first' because

58 

59 

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 

63 

64 

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 

68 

69 

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.""" 

73 

74 model_config = ConfigDict(frozen=True) 

75 

76 id_token: SecretStr 

77 refresh_token: SecretStr | None = None 

78 issuer: str | None = None 

79 expires_at: datetime | None = None 

80 

81 

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.""" 

85 

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 

92 

93 def epoch(self) -> int: 

94 return self._epoch 

95 

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 

101 

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 ) 

108 

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 ) 

114 

115 def flush(self) -> None: 

116 self._entries.flush_cache() # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped 

117 

118 

119_ASSERTION_CACHE: Final = SSOAssertionCache() 

120 

121 

122class _IdTokenClaims(BaseModel): 

123 exp: float | None = None 

124 iss: str | None = None 

125 

126 

127class _StoredAssertionPayload(BaseModel): 

128 id_token: str 

129 refresh_token: str | None = None 

130 issuer: str | None = None 

131 expires_at: datetime | None = None 

132 

133 

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 ) 

159 

160 

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. 

166 

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 

176 

177 

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 

191 

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 

199 

200 

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 

206 

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) 

224 

225 

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 

229 

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 ) 

253 

254 

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 

268 

269 

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. 

273 

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 """ 

280 

281 

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. 

285 

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 """ 

289 

290 async def fetch(self, user_id: str) -> SSOIdentityAssertion | None: ... 290 ↛ exitline 290 didn't return from function 'fetch' because

291 

292 

293class DbSSOAssertionStore: 

294 """The live store: the row the SSO callback wrote, read back by ``user_id``. 

295 

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 """ 

300 

301 def __init__(self, cache: SSOAssertionCache = _ASSERTION_CACHE) -> None: 

302 self._cache = cache 

303 

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 

309 

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 

315 

316 

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 

323 

324 from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 # runtime global 

325 decrypt_value_helper, 

326 encrypt_value_helper, 

327 ) 

328 

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 

347 

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 ) 

355 

356 

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 )