diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 54deb8c202e..8122a14881a 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index efc5228335a..9e21cd42b93 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e92c68be954..588d02d3b2f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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)