From 6b0af7f89a5b7fe7cf1ec4f038ae3a1b0a965c77 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Tue, 19 May 2026 07:15:31 +0000 Subject: [PATCH] Fix MCP user_fields TOCTOU + handle Pydantic-typed user_fields entries - coerce_user_fields / _build_user_fields_status: accept MCPUserField Pydantic instances in addition to raw dicts, so a fully-typed LiteLLM_MCPServerTable passed through these helpers no longer silently drops every entry. - store_user_field_values: replace read-then-write with optimistic concurrency (create + UniqueViolation retry, then compare-and-swap on credential_b64 via update_many) so a concurrent BYOK / OAuth2 write to the same (user_id, server_id) row cannot be silently overwritten. Co-authored-by: Yassin Kortam --- litellm/proxy/_experimental/mcp_server/db.py | 76 +++++++++++++------ .../_experimental/mcp_server/user_fields.py | 31 ++++++-- .../mcp_management_endpoints.py | 3 + 3 files changed, 83 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 8122a14881a..f0f93846c34 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -1,9 +1,12 @@ +import asyncio import base64 import binascii import json from datetime import datetime, timedelta, timezone from typing import Any, Dict, Iterable, List, Optional, Set, Union, cast +from prisma.errors import UniqueViolationError + from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.proxy._types import ( @@ -711,35 +714,64 @@ async def store_user_field_values( BYOK and OAuth2 credentials share the same ``(user_id, server_id)`` row. Refuse to overwrite a non-user-fields credential so saving user-field values does not silently destroy a stored BYOK API key or OAuth2 token. - """ - # Guard against silently overwriting a BYOK or OAuth2 credential that - # shares the same (user_id, server_id) row. - 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 - and _decode_user_fields_payload(existing.credential_b64) is None - ): - raise ValueError( - f"Existing credential for user {user_id} and server " - f"{server_id} is not a user-fields payload (likely BYOK or " - f"OAuth2). Refusing to overwrite." - ) + Uses optimistic concurrency to close the read-then-write race against + concurrent ``store_user_credential`` / ``store_user_oauth_credential`` + calls: the write is gated on the previously-observed ``credential_b64`` + being unchanged, and we retry from the re-read on contention. + """ payload = json.dumps({"type": "user_fields", "values": values}) encoded = encrypt_value_helper(payload) - await prisma_client.db.litellm_mcpusercredentials.upsert( - where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}, - data={ - "create": { + + # Bound the retry loop; in practice contention is resolved in a single + # extra round-trip but we leave headroom for pathological interleavings. + for attempt in range(5): + 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 None: + try: + await prisma_client.db.litellm_mcpusercredentials.create( + data={ + "user_id": user_id, + "server_id": server_id, + "credential_b64": encoded, + } + ) + return + except UniqueViolationError: + # A concurrent writer inserted the row first; restart so we + # can inspect what they wrote before clobbering it. + await asyncio.sleep(0) + continue + + if _decode_user_fields_payload(existing.credential_b64) is None: + raise ValueError( + f"Existing credential for user {user_id} and server " + f"{server_id} is not a user-fields payload (likely BYOK or " + f"OAuth2). Refusing to overwrite." + ) + + # Compare-and-swap on credential_b64: only succeed if the row still + # matches what we just inspected, so a concurrent BYOK/OAuth2 write + # between the read and the write is not silently overwritten. + updated = await prisma_client.db.litellm_mcpusercredentials.update_many( + where={ "user_id": user_id, "server_id": server_id, - "credential_b64": encoded, + "credential_b64": existing.credential_b64, }, - "update": {"credential_b64": encoded}, - }, + data={"credential_b64": encoded}, + ) + if updated: + return + await asyncio.sleep(0) + + raise RuntimeError( + f"store_user_field_values: gave up after repeated concurrent " + f"modifications for user {user_id} and server {server_id}" ) diff --git a/litellm/proxy/_experimental/mcp_server/user_fields.py b/litellm/proxy/_experimental/mcp_server/user_fields.py index 9f47126e54f..af72740fff2 100644 --- a/litellm/proxy/_experimental/mcp_server/user_fields.py +++ b/litellm/proxy/_experimental/mcp_server/user_fields.py @@ -26,19 +26,40 @@ else: UserFieldServer = MCPServer +def _entry_to_dict(entry: Any) -> Optional[Dict[str, Any]]: + """Coerce a single user-field entry to a plain dict, or None if malformed. + + Accepts both raw dicts (as Prisma hands JSONB columns back) and Pydantic + ``MCPUserField`` model instances (as a fully-typed + ``LiteLLM_MCPServerTable`` would carry). + """ + if isinstance(entry, dict): + return entry + model_dump = getattr(entry, "model_dump", None) + if callable(model_dump): + try: + dumped = model_dump() + except Exception: # noqa: BLE001 — drop malformed entries silently + return None + if isinstance(dumped, dict): + return dumped + return None + + def coerce_user_fields(server: UserFieldServer) -> List[Dict[str, Any]]: """Return the server's declared user fields as a list of plain dicts. The column is stored as JSONB but Prisma sometimes hands it back as a - string; this normalises both shapes and silently drops malformed - entries (the admin form filters these out, but DB writes from - external tooling might not). + string; fully-typed ``LiteLLM_MCPServerTable`` instances carry it as a + list of ``MCPUserField`` Pydantic models. This normalises all shapes + and silently drops malformed entries (the admin form filters these + out, but DB writes from external tooling might not). """ raw = getattr(server, "user_fields", None) if not raw: return [] if isinstance(raw, list): - return [e for e in raw if isinstance(e, dict)] + return [d for d in (_entry_to_dict(e) for e in raw) if d is not None] if isinstance(raw, str): import json @@ -47,7 +68,7 @@ def coerce_user_fields(server: UserFieldServer) -> List[Dict[str, Any]]: except (ValueError, TypeError): return [] if isinstance(parsed, list): - return [e for e in parsed if isinstance(e, dict)] + return [d for d in (_entry_to_dict(e) for e in parsed) if d is not None] return [] diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index aa2ba932805..62b35b56fa2 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -2140,6 +2140,9 @@ if MCP_AVAILABLE: raw_fields = [] user_fields: List[MCPUserField] = [] for entry in raw_fields: + if isinstance(entry, MCPUserField): + user_fields.append(entry) + continue if not isinstance(entry, dict): continue try: