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,