mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
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:
parent
adbc4c277f
commit
e45eda44b5
3 changed files with 106 additions and 107 deletions
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue