fix(mcp): tighten user-fields/BYOK collisions and dedupe value resolution

- db.get_user_credential / store_user_credential: refuse to read or
  overwrite a user-fields payload, so BYOK never injects the raw JSON
  blob as an Authorization header.
- server._enforce_user_fields: when the caller has no user_id but the
  server only declares optional user_fields, return without raising
  the spurious 401.
- Promote _get_user_field_values_cached to module level so the
  dispatch path in mcp_server_manager._resolve_user_field_values
  shares the same cache + DB lookup implementation.
- _call_regular_mcp_tool: resolve user-field values once and pass them
  to resolve_user_field_headers / resolve_user_field_env directly,
  eliminating the duplicate cache lookup and the two wrapper helpers.

Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
Cursor Agent 2026-05-19 06:59:19 +00:00
parent adbc4c277f
commit e45eda44b5
No known key found for this signature in database
3 changed files with 106 additions and 107 deletions

View file

@ -588,7 +588,27 @@ async def store_user_credential(
server_id: str,
credential: str,
) -> None:
"""Store a user credential for a BYOK MCP server."""
"""Store a user credential for a BYOK MCP server.
BYOK, OAuth2, and user-fields payloads share the same ``credential_b64``
column. Refuse to overwrite a stored user-fields payload so saving a
BYOK credential does not silently destroy the user's saved field values.
"""
# Guard against silently overwriting a user-fields payload 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 not None
):
raise ValueError(
f"Existing credential for user {user_id} and server "
f"{server_id} holds user-fields values. Refusing to overwrite "
f"with a BYOK credential."
)
encoded = encrypt_value_helper(credential)
await prisma_client.db.litellm_mcpusercredentials.upsert(
@ -609,13 +629,22 @@ async def get_user_credential(
user_id: str,
server_id: str,
) -> Optional[str]:
"""Return credential for a user+server pair, or None."""
"""Return credential for a user+server pair, or None.
The ``credential_b64`` column is multiplexed between BYOK strings,
OAuth2 blobs and user-fields blobs. A row holding a user-fields
payload is not a BYOK credential — return ``None`` so the caller
triggers the normal "no credential stored" flow instead of injecting
the raw JSON blob as an Authorization header.
"""
row = await prisma_client.db.litellm_mcpusercredentials.find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
if row is None:
return None
if _decode_user_fields_payload(row.credential_b64) is not None:
return None
return _decode_user_credential(row.credential_b64)

View file

@ -12,7 +12,6 @@ import hashlib
import json
import os
import re
import time
from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from urllib.parse import urlparse
@ -1348,6 +1347,10 @@ class MCPServerManager:
``server.execute_mcp_tool`` raises a 401 with a config_url before
we even get here, so by the time injection happens we already
know the values are present.
Delegates to ``server._get_user_field_values_cached`` so the
enforcement check and dispatch share a single cache + DB-lookup
implementation.
"""
from litellm.proxy._experimental.mcp_server.user_fields import (
coerce_user_fields,
@ -1358,70 +1361,13 @@ class MCPServerManager:
if user_api_key_auth is None or not getattr(user_api_key_auth, "user_id", None):
return {}
# Reuse the cache and lookup function from server.py so we don't
# pay a second DB round-trip after the enforcement check populated
# the same cache key.
try:
from litellm.proxy._experimental.mcp_server.server import ( # noqa: PLC0415
_USER_FIELDS_CACHE_TTL,
_user_fields_cache,
_write_user_fields_cache,
)
except ImportError:
_user_fields_cache = {} # type: ignore[assignment]
_USER_FIELDS_CACHE_TTL = 60
_write_user_fields_cache = None # type: ignore[assignment]
user_id = user_api_key_auth.user_id or ""
cache_key = (user_id, mcp_server.server_id)
cached = _user_fields_cache.get(cache_key) if _user_fields_cache else None
if cached is not None:
values, ts = cached
if time.monotonic() - ts < _USER_FIELDS_CACHE_TTL:
return values or {}
from litellm.proxy._experimental.mcp_server.db import get_user_field_values
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return {}
values = await get_user_field_values(
prisma_client=prisma_client,
user_id=user_id,
server_id=mcp_server.server_id,
from litellm.proxy._experimental.mcp_server.server import ( # noqa: PLC0415
_get_user_field_values_cached,
)
if _write_user_fields_cache is not None:
_write_user_fields_cache(user_id, mcp_server.server_id, values)
_, values = await _get_user_field_values_cached(mcp_server, user_api_key_auth)
return values or {}
async def _resolve_user_field_headers(
self,
mcp_server: MCPServer,
user_api_key_auth: Optional["UserAPIKeyAuth"],
) -> Dict[str, str]:
from litellm.proxy._experimental.mcp_server.user_fields import (
resolve_user_field_headers,
)
stored = await self._resolve_user_field_values(mcp_server, user_api_key_auth)
if not stored:
return {}
return resolve_user_field_headers(mcp_server, stored)
async def _resolve_user_field_env(
self,
mcp_server: MCPServer,
user_api_key_auth: Optional["UserAPIKeyAuth"],
) -> Dict[str, str]:
from litellm.proxy._experimental.mcp_server.user_fields import (
resolve_user_field_env,
)
stored = await self._resolve_user_field_values(mcp_server, user_api_key_auth)
if not stored:
return {}
return resolve_user_field_env(mcp_server, stored)
async def _create_mcp_client(
self,
server: MCPServer,
@ -2841,11 +2787,23 @@ class MCPServerManager:
extra_headers = {}
extra_headers.update(mcp_server.static_headers)
# User-fields: inject each declared user field's stored value either
# as a header (http/sse) or env var (stdio path; handled below).
user_field_headers = await self._resolve_user_field_headers(
# User-fields: resolve the calling user's stored values once and
# reuse them for both the header (http/sse) and env (stdio) paths
# below — calling the resolver twice would repeat cache lookups
# and coercion work for the same data.
from litellm.proxy._experimental.mcp_server.user_fields import (
resolve_user_field_env,
resolve_user_field_headers,
)
stored_user_field_values = await self._resolve_user_field_values(
mcp_server, user_api_key_auth
)
user_field_headers = (
resolve_user_field_headers(mcp_server, stored_user_field_values)
if stored_user_field_values
else {}
)
if user_field_headers:
if extra_headers is None:
extra_headers = {}
@ -2880,8 +2838,10 @@ class MCPServerManager:
if extra_headers is not None and len(extra_headers) == 0:
extra_headers = None
user_field_env = await self._resolve_user_field_env(
mcp_server, user_api_key_auth
user_field_env = (
resolve_user_field_env(mcp_server, stored_user_field_values)
if stored_user_field_values
else {}
)
stdio_env = self._build_stdio_env(
mcp_server, raw_headers, user_field_env=user_field_env or None

View file

@ -121,6 +121,46 @@ def _write_user_fields_cache(
_user_fields_cache[(user_id, server_id)] = (values, time.monotonic())
async def _get_user_field_values_cached(
mcp_server: MCPServer,
user_api_key_auth: Optional[UserAPIKeyAuth],
) -> Tuple[Optional[str], Optional[Dict[str, str]]]:
"""Return (user_id, stored_values) for the calling user.
``stored_values`` is None when either no user_id is available or
the user has not yet saved any values. Reads through a 60s
in-memory cache so back-to-back tool calls don't hammer the DB.
Module-level (not nested) so ``mcp_server_manager`` can reuse the
same cache + DB-lookup path during dispatch — keeping two copies
in sync was a maintenance hazard.
"""
user_id = (user_api_key_auth.user_id if user_api_key_auth else None) or ""
if not user_id:
return None, None
cache_key = (user_id, mcp_server.server_id)
cached = _user_fields_cache.get(cache_key)
if cached is not None:
values, ts = cached
if time.monotonic() - ts < _USER_FIELDS_CACHE_TTL:
return user_id, values
from litellm.proxy._experimental.mcp_server.db import get_user_field_values
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return user_id, None
values = await get_user_field_values(
prisma_client=prisma_client,
user_id=user_id,
server_id=mcp_server.server_id,
)
_write_user_fields_cache(user_id, mcp_server.server_id, values)
return user_id, values
# Check if MCP is available
# "mcp" requires python 3.10 or higher, but several litellm users use python 3.8
# We're making this conditional import to avoid breaking users who use python 3.8.
@ -2066,41 +2106,6 @@ if MCP_AVAILABLE:
},
)
async def _get_user_field_values_cached(
mcp_server: MCPServer,
user_api_key_auth: Optional[UserAPIKeyAuth],
) -> Tuple[Optional[str], Optional[Dict[str, str]]]:
"""Return (user_id, stored_values) for the calling user.
``stored_values`` is None when either no user_id is available or
the user has not yet saved any values. Reads through a 60s
in-memory cache so back-to-back tool calls don't hammer the DB.
"""
user_id = (user_api_key_auth.user_id if user_api_key_auth else None) or ""
if not user_id:
return None, None
cache_key = (user_id, mcp_server.server_id)
cached = _user_fields_cache.get(cache_key)
if cached is not None:
values, ts = cached
if time.monotonic() - ts < _USER_FIELDS_CACHE_TTL:
return user_id, values
from litellm.proxy._experimental.mcp_server.db import get_user_field_values
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return user_id, None
values = await get_user_field_values(
prisma_client=prisma_client,
user_id=user_id,
server_id=mcp_server.server_id,
)
_write_user_fields_cache(user_id, mcp_server.server_id, values)
return user_id, values
async def _enforce_user_fields(
mcp_server: MCPServer,
user_api_key_auth: Optional[UserAPIKeyAuth],
@ -2121,11 +2126,16 @@ if MCP_AVAILABLE:
user_id = (user_api_key_auth.user_id if user_api_key_auth else None) or ""
if not user_id:
# No identity means we can't look up stored values; treat as
# missing and let the user log in via the dashboard.
# No identity means we can't look up stored values. Compute the
# missing-required set against an empty value map and only raise
# when at least one required field is actually missing — servers
# whose user_fields are all optional must not be blocked here.
missing = compute_missing_user_fields(mcp_server, None)
if not missing:
return
detail = build_user_fields_missing_error(
mcp_server,
compute_missing_user_fields(mcp_server, None),
missing,
_resolve_proxy_base_url_env(),
)
raise HTTPException(status_code=401, detail=detail)