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 <yassin@berri.ai>
This commit is contained in:
Cursor Agent 2026-05-19 07:15:31 +00:00
parent e45eda44b5
commit 6b0af7f89a
No known key found for this signature in database
3 changed files with 83 additions and 27 deletions

View file

@ -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}"
)

View file

@ -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 []

View file

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