fix(mcp): guard BYOK overwrite in oauth credential store, raise clear error when client_id absent

This commit is contained in:
Ishaan Jaffer 2026-03-11 09:05:38 -07:00
parent 03cba2c27e
commit c9f834f093
2 changed files with 30 additions and 1 deletions

View file

@ -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}},

View file

@ -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,