Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py: 48%

129 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1"""Composition root for the v2-native authorization_code per-user OAuth token store (step 1b). 

2 

3Assembles ``Cached(Refreshing(V2PerUserTokenStore))`` and replaces ``V1PerUserTokenStore`` in the 

4resolver. The runtime collaborators (DB, HTTP, the shared cache, Redis) are LiteLLM globals not ready 

5at import time, so the chain is built lazily on first use. When Redis is wired it uses the 

6cross-replica path (DualCache-backed cache + ``SET NX PX`` coordinator); otherwise it falls back to 

7the foundation's in-process defaults (correct for a single replica). The DB read/refresh-grant/persist 

8collaborators acquire their globals per call, mirroring v1's lazy-import pattern. 

9""" 

10 

11from __future__ import annotations 

12 

13import asyncio 

14from collections.abc import Callable, Mapping 

15from functools import partial 

16from typing import TYPE_CHECKING, Final 

17 

18from litellm._logging import verbose_logger 

19from litellm.proxy._experimental.mcp_server.oauth_identity_binding import credential_binding_matches 

20from litellm.proxy._experimental.mcp_server.outbound_credentials.authz_code_refresher import ( 

21 AuthorizationCodeRefresher, 

22) 

23from litellm.proxy._experimental.mcp_server.outbound_credentials.dual_cache_token_backend import ( 

24 AsyncCache, 

25 DualCacheTokenCacheBackend, 

26) 

27from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( 

28 CachedOAuthTokenStore, 

29 InvalidatableOAuthTokenStore, 

30 OAuthToken, 

31 RefreshCoordinator, 

32 RefreshingTokenStore, 

33 TokenCacheBackend, 

34 TokenStoreUnavailable, 

35) 

36from litellm.proxy._experimental.mcp_server.outbound_credentials.runtime_refresh_coordinator import ( 

37 runtime_refresh_coordinator, 

38) 

39from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import ( 

40 OAuthTokenCacheCodec, 

41) 

42from litellm.proxy._experimental.mcp_server.outbound_credentials.v2_token_store import ( 

43 V2PerUserTokenStore, 

44) 

45 

46if TYPE_CHECKING: 46 ↛ 47line 46 didn't jump to line 47 because the condition on line 46 was never true

47 from litellm.types.mcp_server.mcp_server_manager import MCPServer 

48 

49# A token with no declared expiry is cached for this long; one with an expiry is cached until then. 

50_DEFAULT_TTL_SECONDS: Final = 300.0 

51 

52ServerLookup = Callable[[str], "MCPServer | None"] 

53StoreBuilder = Callable[[ServerLookup], tuple[InvalidatableOAuthTokenStore, bool]] 

54 

55 

56async def _read_credential(user_id: str, server_id: str) -> Mapping[str, object] | None: 

57 from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 

58 get_user_oauth_credential, 

59 ) 

60 from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 

61 

62 if prisma_client is None: 

63 raise TokenStoreUnavailable("Database not connected") 

64 return await get_user_oauth_credential(prisma_client, user_id, server_id) 

65 

66 

67async def _persist_credential( 

68 user_id: str, 

69 server_id: str, 

70 access_token: str, 

71 refresh_token: str | None, 

72 expires_in: int | None, 

73 scopes: tuple[str, ...] | None, 

74 identity_binding_proof: str | None = None, 

75) -> None: 

76 from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 

77 store_user_oauth_credential, 

78 ) 

79 from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 

80 

81 if prisma_client is None: 

82 return 

83 await store_user_oauth_credential( 

84 prisma_client=prisma_client, 

85 user_id=user_id, 

86 server_id=server_id, 

87 access_token=access_token, 

88 refresh_token=refresh_token, 

89 expires_in=expires_in, 

90 scopes=list(scopes) if scopes else None, 

91 skip_byok_guard=True, 

92 identity_binding_proof=identity_binding_proof, 

93 ) 

94 

95 

96async def _post_token_endpoint(url: str, form: dict[str, str], headers: dict[str, str]) -> dict[str, object] | None: 

97 from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 

98 get_async_httpx_client, # pyright: ignore 

99 ) 

100 from litellm.types.llms.custom_http import httpxSpecialProvider # noqa: PLC0415 

101 

102 # litellm's httpx handler and httpx.Response are only partially typed; the IdP returns a JSON 

103 # object and the refresher validates each field, so the untyped boundary is contained here. 

104 provider: Final = httpxSpecialProvider.Oauth2Check 

105 request_headers: Final = {"Accept": "application/json", **headers} 

106 # A failed refresh is a miss, not a 500 (matches v1), so any error becomes None. 

107 try: 

108 client: Final = get_async_httpx_client(llm_provider=provider) # pyright: ignore 

109 response: Final = await client.post(url, headers=request_headers, data=form) # pyright: ignore 

110 response.raise_for_status() # pyright: ignore 

111 body: Final[dict[str, object]] = response.json() # pyright: ignore 

112 except Exception as exc: # noqa: BLE001 

113 verbose_logger.warning("MCP OAuth refresh request failed: %s", exc) 

114 return None 

115 else: 

116 return body # pyright: ignore 

117 

118 

119def _redis_cache_is_available() -> bool: 

120 from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 

121 

122 return user_api_key_cache.redis_cache is not None 

123 

124 

125def _runtime_backend_and_coordinator() -> tuple[TokenCacheBackend | None, RefreshCoordinator | None, bool]: 

126 """The cross-replica cache + coordinator when Redis is wired, else ``(None, None, False)`` so the 

127 foundation's in-process defaults are used (a single replica needs no shared cache or lock). 

128 """ 

129 from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 

130 decrypt_value_helper, 

131 encrypt_value_helper, 

132 ) 

133 from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 

134 

135 coordinator: Final = runtime_refresh_coordinator() 

136 if coordinator is None: 136 ↛ 138line 136 didn't jump to line 138 because the condition on line 136 was always true

137 return None, None, False 

138 codec: Final = OAuthTokenCacheCodec( 

139 encrypt_value_helper, 

140 lambda blob: decrypt_value_helper(blob, "mcp_per_user_token", exception_type="debug"), 

141 ) 

142 # user_api_key_cache satisfies the AsyncCache slice (DualCache types ttl via **kwargs) - an 

143 # untyped-boundary cast. 

144 cache: Final[AsyncCache] = user_api_key_cache # pyright: ignore 

145 backend: Final = DualCacheTokenCacheBackend(cache, codec) 

146 return backend, coordinator, True 

147 

148 

149async def _read_bound_credential( 

150 server_lookup: ServerLookup, user_id: str, server_id: str 

151) -> Mapping[str, object] | None: 

152 credential: Final = await _read_credential(user_id, server_id) 

153 server: Final = server_lookup(server_id) 

154 binding: Final = server.oauth_identity_binding if server else None 

155 if credential is not None and binding is not None and binding.mode == "enforce": 

156 if not await credential_binding_matches(binding, user_id, server_id, credential): 

157 return None 

158 return credential 

159 

160 

161def _build_per_user_oauth_token_store( 

162 server_lookup: ServerLookup, 

163) -> tuple[CachedOAuthTokenStore, bool]: 

164 backend, coordinator, uses_redis = _runtime_backend_and_coordinator() 

165 refresher: Final = AuthorizationCodeRefresher(server_lookup, _post_token_endpoint, _persist_credential) 

166 refreshing: Final = RefreshingTokenStore( 

167 V2PerUserTokenStore(partial(_read_bound_credential, server_lookup)), refresher, coordinator=coordinator 

168 ) 

169 return CachedOAuthTokenStore(refreshing, default_ttl_seconds=_DEFAULT_TTL_SECONDS, backend=backend), uses_redis 

170 

171 

172def build_per_user_oauth_token_store( 

173 server_lookup: ServerLookup, 

174) -> CachedOAuthTokenStore: 

175 store, _uses_redis = _build_per_user_oauth_token_store(server_lookup) 

176 return store 

177 

178 

179class LazyPerUserOAuthTokenStore: 

180 """``OAuthTokenStore`` that builds the v2-native chain on first ``fetch``. 

181 

182 The chain's cache/lock collaborators are LiteLLM runtime globals not available when the resolver 

183 is constructed at import time, so construction is deferred to the first request (by when they are 

184 wired). A no-Redis chain is replaced once Redis becomes available. 

185 """ 

186 

187 def __init__( 

188 self, 

189 server_lookup: ServerLookup, 

190 *, 

191 store_builder: StoreBuilder = _build_per_user_oauth_token_store, 

192 redis_available: Callable[[], bool] = _redis_cache_is_available, 

193 ) -> None: 

194 self._server_lookup = server_lookup 

195 self._store_builder = store_builder 

196 self._redis_available = redis_available 

197 self._store: InvalidatableOAuthTokenStore | None = None 

198 self._uses_redis = False 

199 self._fetch_lock = asyncio.Condition() 

200 self._local_fetches = 0 

201 

202 async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: 

203 token: Final = await self._fetch_token(user_id, server_id) 

204 server: Final = self._server_lookup(server_id) 

205 binding: Final = server.oauth_identity_binding if server else None 

206 if token is not None and binding is not None and binding.mode == "enforce": 

207 if not await credential_binding_matches( 

208 binding, user_id, server_id, {"identity_binding_proof": token.identity_binding_proof} 

209 ): 

210 await self.invalidate(user_id, server_id) 

211 return None 

212 return token 

213 

214 async def _fetch_token(self, user_id: str, server_id: str) -> OAuthToken | None: 

215 if self._uses_redis: 

216 store = self._store 

217 if store is not None: 

218 return await store.fetch(user_id, server_id) 

219 

220 store, uses_redis = await self._store_for_fetch() 

221 try: 

222 return await store.fetch(user_id, server_id) 

223 finally: 

224 if not uses_redis: 

225 await self._finish_local_fetch() 

226 

227 async def invalidate(self, user_id: str, server_id: str) -> None: 

228 """Drop the chain's cached entry for ``(user_id, server_id)`` after the credential row 

229 changes (re-auth, revoke). Builds the chain if no fetch has run yet, so a shared (Redis) 

230 cache entry written by another worker is dropped too; the in-process case is then a no-op 

231 on an empty cache. 

232 """ 

233 if self._uses_redis: 233 ↛ 234line 233 didn't jump to line 234 because the condition on line 233 was never true

234 store = self._store 

235 if store is not None: 

236 await store.invalidate(user_id, server_id) 

237 return 

238 

239 store, uses_redis = await self._store_for_fetch() 

240 try: 

241 await store.invalidate(user_id, server_id) 

242 finally: 

243 if not uses_redis: 243 ↛ exitline 243 didn't return from function 'invalidate' because the condition on line 243 was always true

244 await self._finish_local_fetch() 

245 

246 async def _store_for_fetch(self) -> tuple[InvalidatableOAuthTokenStore, bool]: 

247 async with self._fetch_lock: 

248 while ( 248 ↛ 251line 248 didn't jump to line 251 because the condition on line 248 was never true

249 self._store is not None and not self._uses_redis and self._redis_available() and self._local_fetches > 0 

250 ): 

251 await self._fetch_lock.wait() 

252 store = self._store 

253 if store is None or (not self._uses_redis and self._redis_available()): 

254 store, self._uses_redis = self._store_builder(self._server_lookup) 

255 self._store = store 

256 uses_redis: Final = self._uses_redis 

257 if not uses_redis: 257 ↛ 259line 257 didn't jump to line 259 because the condition on line 257 was always true

258 self._local_fetches += 1 

259 return store, uses_redis 

260 

261 async def _finish_local_fetch(self) -> None: 

262 async with self._fetch_lock: 

263 self._local_fetches -= 1 

264 if self._local_fetches == 0: 264 ↛ exitline 264 didn't jump to the function exit

265 self._fetch_lock.notify_all()