mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
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:
parent
e45eda44b5
commit
6b0af7f89a
3 changed files with 83 additions and 27 deletions
|
|
@ -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}"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue