From c9f834f0935eb04a87eaef18a35bd8926fcd8634 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 11 Mar 2026 09:05:38 -0700 Subject: [PATCH] fix(mcp): guard BYOK overwrite in oauth credential store, raise clear error when client_id absent --- litellm/proxy/_experimental/mcp_server/db.py | 18 ++++++++++++++++++ .../mcp_server/discoverable_endpoints.py | 13 ++++++++++++- 2 files changed, 30 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index be30f866eee..72875b5cf8d 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -500,6 +500,24 @@ async def store_user_oauth_credential( if scopes: payload["scopes"] = scopes + # Guard against silently overwriting a BYOK credential with an OAuth token. + # BYOK credentials lack a "type" field (or use a non-"oauth2" type). + existing = await prisma_client.db.litellm_mcpusercredentials.find_unique( + where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}} + ) + if existing is not None: + try: + raw = json.loads(base64.urlsafe_b64decode(existing.credential_b64).decode()) + if raw.get("type") != "oauth2": + raise ValueError( + f"A non-OAuth2 credential already exists for user {user_id} " + f"and server {server_id}. Refusing to overwrite." + ) + except (ValueError, KeyError): + raise + except Exception: + pass # Malformed existing record — allow overwrite + encoded = base64.urlsafe_b64encode(json.dumps(payload).encode()).decode() await prisma_client.db.litellm_mcpusercredentials.upsert( where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}, diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 1e368787dbf..ad1dadb1222 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -341,8 +341,19 @@ async def authorize( mcp_server = _resolve_oauth2_server_for_root_endpoints() if mcp_server is None: raise HTTPException(status_code=404, detail="MCP server not found") - # Use server's stored client_id when caller doesn't supply one + # Use server's stored client_id when caller doesn't supply one. + # Raise a clear error instead of passing an empty string — an empty + # client_id would silently produce a broken authorization URL. resolved_client_id: str = mcp_server.client_id or client_id or "" + if not resolved_client_id: + raise HTTPException( + status_code=400, + detail={ + "error": "client_id is required but was not supplied and is not " + "stored on the MCP server record. Provide client_id as a query " + "parameter or configure it on the server." + }, + ) return await authorize_with_server( request=request, mcp_server=mcp_server,