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

1"""Renew the stored SSO identity assertion so an ID-JAG agent outlives one id_token. 

2 

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. 

7 

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. 

14 

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. 

24 

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. 

29 

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

39 

40from __future__ import annotations 

41 

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 

48 

49import httpx 

50from pydantic import SecretStr, TypeAdapter, ValidationError 

51from typing_extensions import assert_never 

52 

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 

85 

86_BODY_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(dict[str, object]) 

87 

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 

95 

96 

97class AssertionRead(Protocol): 

98 """Reads the user's stored assertion row.""" 

99 

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

101 

102 

103class AssertionWrite(Protocol): 

104 """Replaces the user's stored assertion row.""" 

105 

106 async def __call__(self, user_id: str, assertion: SSOIdentityAssertion) -> None: ... 106 ↛ exitline 106 didn't return from function '__call__' because

107 

108 

109class CoordinatorFactory(Protocol): 

110 """Builds the cross-replica coordinator, or ``None`` when there is no shared lock to build on.""" 

111 

112 def __call__(self) -> RefreshCoordinator | None: ... 112 ↛ exitline 112 didn't return from function '__call__' because

113 

114 

115class FormPost(Protocol): 

116 """POSTs an OAuth form and hands back the raw response.""" 

117 

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

121 

122 

123@dataclass(frozen=True, slots=True) 

124class SSOClientConfig: 

125 """The generic-OIDC client credentials a refresh_token grant has to authenticate as, and how.""" 

126 

127 token_endpoint: str 

128 client_id: str 

129 client_secret: SecretStr 

130 auth_method: MCPTokenEndpointAuthMethod 

131 

132 

133def sso_client_config(env: Mapping[str, str]) -> SSOClientConfig | None: 

134 """The configured generic-OIDC client, or ``None`` when the deployment has none. 

135 

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. 

142 

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 ) 

161 

162 

163@dataclass(frozen=True, slots=True) 

164class RefreshFailure: 

165 """Why a renewal produced nothing, split by what the caller can do about it. 

166 

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

173 

174 kind: Literal["rejected", "unavailable"] 

175 detail: str 

176 

177 @staticmethod 

178 def of_rejected(detail: str) -> RefreshFailure: 

179 return RefreshFailure(kind="rejected", detail=detail) 

180 

181 @staticmethod 

182 def of_unavailable(detail: str) -> RefreshFailure: 

183 return RefreshFailure(kind="unavailable", detail=detail) 

184 

185 

186class TokenEndpointTransport(Protocol): 

187 """One form POST to the IdP token endpoint, with the refusal/outage split preserved. 

188 

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

194 

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]: ... 

198 

199 

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 

205 

206 

207class HttpxTokenEndpointTransport: 

208 """The live transport. 4xx is the IdP refusing this grant; anything else is an outage. 

209 

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

213 

214 def __init__(self, post: FormPost = post_form) -> None: 

215 self._post = post 

216 

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) 

238 

239 

240class SSOAssertionRefresher: 

241 """Redeems the stored refresh token for a current id_token and writes the rotation back. 

242 

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

246 

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 

259 

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

280 

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) 

300 

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. 

309 

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) 

330 

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. 

335 

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. 

340 

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 

364 

365 

366class RefreshingSSOAssertionStore: 

367 """An ``SSOAssertionStore`` that renews a near-expiry assertion before handing it back. 

368 

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. 

374 

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

384 

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 

403 

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) 

415 

416 def _expiring(self, assertion: SSOIdentityAssertion | None) -> bool: 

417 return assertion is not None and assertion_expired(assertion, self._clock() + self._skew) 

418 

419 def _coordinator(self) -> RefreshCoordinator: 

420 """The cross-replica coordinator once Redis is reachable, else the in-process one. 

421 

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 

428 

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) 

446 

447 async def _reread_renewed(self, user_id: str) -> None: 

448 """A loser cannot distinguish refusal from an unrecorded renewal without risking token replay. 

449 

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 ) 

459 

460 

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 )