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

1"""Startup backfill for oauth2 MCP server rows persisted before oauth2_flow was written. 

2 

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. 

8 

9Signal order, strongest first: 

10 

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. 

30 

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

36 

37import json 

38from collections import Counter 

39from collections.abc import Mapping, Sequence 

40from typing import Final, Literal, Protocol 

41 

42from pydantic import JsonValue 

43 

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 

48 

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] 

57 

58_BACKFILL_AUDIT_ACTOR: Final = "oauth2_flow_backfill" 

59 

60 

61class _MCPServerRow(Protocol): 

62 """The ``LiteLLM_MCPServerTable`` columns this backfill reads.""" 

63 

64 @property 

65 def server_id(self) -> str: ... 65 ↛ exitline 65 didn't return from function 'server_id' because

66 

67 @property 

68 def authorization_url(self) -> str | None: ... 68 ↛ exitline 68 didn't return from function 'authorization_url' because

69 

70 @property 

71 def registration_url(self) -> str | None: ... 71 ↛ exitline 71 didn't return from function 'registration_url' because

72 

73 @property 

74 def token_url(self) -> str | None: ... 74 ↛ exitline 74 didn't return from function 'token_url' because

75 

76 @property 

77 def credentials(self) -> str | Mapping[str, JsonValue] | None: ... 77 ↛ exitline 77 didn't return from function 'credentials' because

78 

79 

80class _MCPUserCredentialRow(Protocol): 

81 """The ``LiteLLM_MCPUserCredentials`` columns this backfill reads.""" 

82 

83 @property 

84 def server_id(self) -> str: ... 84 ↛ exitline 84 didn't return from function 'server_id' because

85 

86 @property 

87 def credential_b64(self) -> str: ... 87 ↛ exitline 87 didn't return from function 'credential_b64' because

88 

89 

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

92 

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

94 

95 

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

98 

99 

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 

103 

104 

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 

108 

109 

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

124 

125 

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" 

143 

144 

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 {} 

153 

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 } 

161 

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 ) 

175 

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 ) 

193 

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 ) 

201 

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