Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py: 44%
78 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"""Startup backfill for oauth2 MCP server rows persisted before oauth2_flow was written.
3Rows created before the write-side stamps (DCR persist, UI create, REST create) carry a
4null ``oauth2_flow`` and rely on read-time field-shape inference, which cannot tell a
5DCR-registered interactive server from an M2M server unless endpoint discovery succeeds
6first. This backfill classifies each null row once, at rest, using signals inference
7never had, and persists the result so the read path never has to infer again.
9Signal order, strongest first:
111. Per-user OAuth token rows exist for the server: only the interactive flow mints
12 per-user tokens, so this is definitive and immune to the discovery trap. BYOK API
13 keys share the same table (``LiteLLM_MCPUserCredentials``), so only rows whose
14 payload decodes as a ``type: oauth2`` token count as proof; bare keys and
15 undecodable rows prove nothing about the flow.
162. ``authorization_url`` persisted: interactive needs a user-facing authorization
17 endpoint; M2M (RFC 6749 section 4.4) never has one.
183. ``registration_url`` persisted: dynamic client registration (RFC 7591) exists to mint
19 clients for the interactive flow; M2M servers are configured with static credentials.
204. ``token_url`` plus decryptable ``client_id`` and ``client_secret``: ambiguous, left
21 unstamped. The shape is shared by M2M servers and DCR-registered interactive servers
22 whose authorization endpoint lives only in discovery (registered but never signed
23 in), so stamping client_credentials here could permanently route per-user traffic
24 through the proxy's stored client credential. The row keeps working through the
25 request-time backstop and a warning names it with the one-line fix (set oauth2_flow
26 via the dashboard or ``PUT /v1/mcp/server``); a completed interactive sign-in also
27 heals it via rule 1 at the next boot.
285. Anything else is interactive: matching how ``needs_user_oauth_token`` treats a null
29 flow, so the stamp never changes runtime routing for rows no rule recognizes.
31The backfill never stamps client_credentials: M2M is asserted by a human (config
32requires it, the API accepts it, the dashboard sets it), mirroring the config-level
33validation error. Runs before the first registry load on every boot and is idempotent:
34a healed fleet has no null rows and the backfill exits after one query.
35"""
37import json
38from collections import Counter
39from collections.abc import Mapping, Sequence
40from typing import Final, Literal, Protocol
42from pydantic import JsonValue
44from litellm._logging import verbose_proxy_logger
45from litellm.proxy._experimental.mcp_server.db import _decode_oauth_payload, decrypt_credentials
46from litellm.proxy.utils import PrismaClient
47from litellm.types.mcp import MCPCredentials
49OAuth2Flow = Literal["client_credentials", "authorization_code"]
50BackfillRule = Literal[
51 "per_user_tokens",
52 "authorization_url",
53 "registration_url",
54 "ambiguous_m2m_shape",
55 "interactive_default",
56]
58_BACKFILL_AUDIT_ACTOR: Final = "oauth2_flow_backfill"
61class _MCPServerRow(Protocol):
62 """The ``LiteLLM_MCPServerTable`` columns this backfill reads."""
64 @property
65 def server_id(self) -> str: ... 65 ↛ exitline 65 didn't return from function 'server_id' because
67 @property
68 def authorization_url(self) -> str | None: ... 68 ↛ exitline 68 didn't return from function 'authorization_url' because
70 @property
71 def registration_url(self) -> str | None: ... 71 ↛ exitline 71 didn't return from function 'registration_url' because
73 @property
74 def token_url(self) -> str | None: ... 74 ↛ exitline 74 didn't return from function 'token_url' because
76 @property
77 def credentials(self) -> str | Mapping[str, JsonValue] | None: ... 77 ↛ exitline 77 didn't return from function 'credentials' because
80class _MCPUserCredentialRow(Protocol):
81 """The ``LiteLLM_MCPUserCredentials`` columns this backfill reads."""
83 @property
84 def server_id(self) -> str: ... 84 ↛ exitline 84 didn't return from function 'server_id' because
86 @property
87 def credential_b64(self) -> str: ... 87 ↛ exitline 87 didn't return from function 'credential_b64' because
90class _MCPServerTable(Protocol):
91 async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_MCPServerRow]: ... 91 ↛ exitline 91 didn't return from function 'find_many' because
93 async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, str]) -> object: ... 93 ↛ exitline 93 didn't return from function 'update_many' because
96class _MCPUserCredentialsTable(Protocol):
97 async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_MCPUserCredentialRow]: ... 97 ↛ exitline 97 didn't return from function 'find_many' because
100def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable:
101 """The MCP server table, typed so the untyped prisma client surface stops here."""
102 return prisma_client.db.litellm_mcpservertable
105def _mcp_user_credentials_table(prisma_client: PrismaClient) -> _MCPUserCredentialsTable:
106 """The per-user MCP credential table, typed so the untyped prisma client surface stops here."""
107 return prisma_client.db.litellm_mcpusercredentials
110def _decrypted_credentials(raw_credentials: str | Mapping[str, JsonValue] | None) -> MCPCredentials | None:
111 if raw_credentials is None:
112 return None
113 parsed: JsonValue | Mapping[str, JsonValue]
114 if isinstance(raw_credentials, str):
115 try:
116 parsed = json.loads(raw_credentials)
117 except (ValueError, TypeError):
118 return None
119 else:
120 parsed = raw_credentials
121 if not isinstance(parsed, dict):
122 return None
123 return decrypt_credentials(credentials=dict(parsed))
126def classify_null_flow_row(
127 *,
128 has_per_user_tokens: bool,
129 authorization_url: str | None,
130 registration_url: str | None,
131 token_url: str | None,
132 credentials: MCPCredentials | None,
133) -> tuple[OAuth2Flow | None, BackfillRule]:
134 if has_per_user_tokens:
135 return "authorization_code", "per_user_tokens"
136 if authorization_url:
137 return "authorization_code", "authorization_url"
138 if registration_url:
139 return "authorization_code", "registration_url"
140 if token_url and credentials and credentials.get("client_id") and credentials.get("client_secret"):
141 return None, "ambiguous_m2m_shape"
142 return "authorization_code", "interactive_default"
145async def backfill_null_oauth2_flows(prisma_client: PrismaClient) -> dict[BackfillRule, int]:
146 """Classify every ``auth_type=oauth2`` row whose ``oauth2_flow`` is null; stamp the provable
147 ones, warn on the ambiguous ones, and return counts per rule."""
148 null_rows: Final[Sequence[_MCPServerRow]] = await _mcp_server_table(prisma_client).find_many(
149 where={"auth_type": "oauth2", "oauth2_flow": None},
150 )
151 if not null_rows: 151 ↛ 154line 151 didn't jump to line 154 because the condition on line 151 was always true
152 return {}
154 server_ids: Final = [row.server_id for row in null_rows]
155 token_rows: Final[Sequence[_MCPUserCredentialRow]] = await _mcp_user_credentials_table(prisma_client).find_many(
156 where={"server_id": {"in": server_ids}},
157 )
158 server_ids_with_oauth_tokens: Final[set[str]] = {
159 token_row.server_id for token_row in token_rows if _decode_oauth_payload(token_row.credential_b64) is not None
160 }
162 classified: Final = tuple(
163 (
164 row,
165 classify_null_flow_row(
166 has_per_user_tokens=row.server_id in server_ids_with_oauth_tokens,
167 authorization_url=row.authorization_url,
168 registration_url=row.registration_url,
169 token_url=row.token_url,
170 credentials=_decrypted_credentials(row.credentials),
171 ),
172 )
173 for row in null_rows
174 )
176 for row, (flow, rule) in classified:
177 if flow is None:
178 verbose_proxy_logger.warning(
179 "oauth2_flow backfill: server_id=%s is ambiguous (client credentials + token_url, "
180 "no interactive signal); left unstamped. Set oauth2_flow explicitly via the "
181 "dashboard or PUT /v1/mcp/server: client_credentials if this server is M2M, or "
182 "complete an interactive sign-in and it will be stamped authorization_code at the "
183 "next boot.",
184 row.server_id,
185 )
186 else:
187 verbose_proxy_logger.info(
188 "oauth2_flow backfill: server_id=%s stamped %s (rule=%s)",
189 row.server_id,
190 flow,
191 rule,
192 )
194 stamped_flows: Final = {flow for _, (flow, _) in classified if flow is not None}
195 for stamped_flow in stamped_flows:
196 server_ids_for_flow = [row.server_id for row, (row_flow, _) in classified if row_flow == stamped_flow]
197 await _mcp_server_table(prisma_client).update_many(
198 where={"server_id": {"in": server_ids_for_flow}, "oauth2_flow": None},
199 data={"oauth2_flow": stamped_flow, "updated_by": _BACKFILL_AUDIT_ACTOR},
200 )
202 counts: Final[dict[BackfillRule, int]] = dict(Counter(rule for _, (_, rule) in classified))
203 verbose_proxy_logger.info(
204 "oauth2_flow backfill: processed %d oauth2 server row(s): %s",
205 len(null_rows),
206 counts,
207 )
208 return counts