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
« 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).
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"""
11from __future__ import annotations
13import asyncio
14from collections.abc import Callable, Mapping
15from functools import partial
16from typing import TYPE_CHECKING, Final
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)
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
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
52ServerLookup = Callable[[str], "MCPServer | None"]
53StoreBuilder = Callable[[ServerLookup], tuple[InvalidatableOAuthTokenStore, bool]]
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
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)
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
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 )
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
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
119def _redis_cache_is_available() -> bool:
120 from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415
122 return user_api_key_cache.redis_cache is not None
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
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
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
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
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
179class LazyPerUserOAuthTokenStore:
180 """``OAuthTokenStore`` that builds the v2-native chain on first ``fetch``.
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 """
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
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
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)
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()
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
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()
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
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()