From fa671c6b327fb566228a968fea53ac217ceaed93 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Tue, 19 May 2026 05:50:21 +0000 Subject: [PATCH] fix(mcp): address security and correctness bugs in user-fields PR - Add missing 'import time' in mcp_server_manager.py; _resolve_user_field_values() was calling time.monotonic() without the module imported, raising NameError on every cached tool call against a server with user_fields. - list_mcp_user_field_values now filters via get_all_mcp_servers_for_user instead of get_all_mcp_servers so non-admin callers no longer see server IDs and user-field metadata (header_name/env_var_name) for servers they cannot access. - store_user_field_values now refuses to overwrite a non-user-fields credential (BYOK / OAuth2) sharing the same (user_id, server_id) row, mirroring the existing skip_byok_guard pattern in store_user_oauth_credential. The POST endpoint surfaces this as a 409 conflict. - Deduplicate user-fields parsing helpers: replace _coerce_user_fields_list / _has_required_user_fields / _compute_missing_user_field_keys with the existing coerce_user_fields / server_has_user_fields / compute_missing_user_fields from user_fields.py. - Remove dead lookup_cached_user_fields() (and its now-unused 'time' import) from user_fields.py. Co-authored-by: Yassin Kortam --- litellm/proxy/_experimental/mcp_server/db.py | 19 ++++ .../mcp_server/mcp_server_manager.py | 1 + .../_experimental/mcp_server/user_fields.py | 24 +---- .../mcp_management_endpoints.py | 102 ++++++------------ 4 files changed, 56 insertions(+), 90 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 94314357d2a..54deb8c202e 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -678,8 +678,27 @@ async def store_user_field_values( ``LiteLLM_MCPUserCredentials.credential_b64``. A ``"type"`` discriminator lets ``get_user_field_values`` tell user-fields rows apart from BYOK strings and OAuth2 payloads sharing the same column. + + 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." + ) + payload = json.dumps({"type": "user_fields", "values": values}) encoded = encrypt_value_helper(payload) await prisma_client.db.litellm_mcpusercredentials.upsert( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0f0bf128501..efc5228335a 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -12,6 +12,7 @@ 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 diff --git a/litellm/proxy/_experimental/mcp_server/user_fields.py b/litellm/proxy/_experimental/mcp_server/user_fields.py index 5684f004e19..a3c910fbe73 100644 --- a/litellm/proxy/_experimental/mcp_server/user_fields.py +++ b/litellm/proxy/_experimental/mcp_server/user_fields.py @@ -13,8 +13,7 @@ resolution / injection logic. from __future__ import annotations -import time -from typing import Any, Dict, List, Optional, Tuple +from typing import Any, Dict, List, Optional from litellm._logging import verbose_logger from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -171,24 +170,3 @@ def resolve_user_field_env( continue env[env_var_name] = value return env - - -def lookup_cached_user_fields( - cache: Dict[Tuple[str, str], Tuple[Optional[Dict[str, str]], float]], - user_id: str, - server_id: str, - ttl_seconds: int, -) -> Tuple[bool, Optional[Dict[str, str]]]: - """Return (cache_hit, values) for the (user, server) cache pair. - - Pulled out so server.py can pass its own cache dict in; keeping the - cache as module-level state in server.py preserves the existing - invalidation hooks called from the management endpoints. - """ - cached = cache.get((user_id, server_id)) - if cached is None: - return False, None - values, ts = cached - if time.monotonic() - ts >= ttl_seconds: - return False, None - return True, values diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 0c392791d88..320549717e0 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -113,7 +113,6 @@ if MCP_AVAILABLE: delete_user_credential, get_all_mcp_servers_for_user, delete_user_field_values, - get_all_mcp_servers, get_mcp_server, get_mcp_servers, get_mcp_submissions, @@ -126,6 +125,11 @@ if MCP_AVAILABLE: store_user_oauth_credential, update_mcp_server, ) + from litellm.proxy._experimental.mcp_server.user_fields import ( + coerce_user_fields, + compute_missing_user_fields, + server_has_user_fields, + ) from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( authorize_with_server, exchange_token_with_server, @@ -480,47 +484,6 @@ if MCP_AVAILABLE: ) -> List[LiteLLM_MCPServerTable]: return [_redact_mcp_credentials(server) for server in mcp_servers] - def _coerce_user_fields_list(raw: Any) -> List[Dict[str, Any]]: - """Normalize a server.user_fields value (JSON string or list) to a list of dicts. - - Used by the BYOK-list-annotation block and any other site that needs - to inspect declared fields without instantiating MCPUserField models. - """ - if not raw: - return [] - if isinstance(raw, list): - return [e for e in raw if isinstance(e, dict)] - if isinstance(raw, str): - try: - parsed = json.loads(raw) - except (ValueError, TypeError): - return [] - if isinstance(parsed, list): - return [e for e in parsed if isinstance(e, dict)] - return [] - - def _has_required_user_fields(raw: Any) -> bool: - """True iff the server declares at least one required user field.""" - for entry in _coerce_user_fields_list(raw): - if entry.get("required", True): - return True - return False - - def _compute_missing_user_field_keys( - raw: Any, stored_values: Dict[str, str] - ) -> List[str]: - """Field keys the user must still supply (required + currently empty).""" - missing: List[str] = [] - for entry in _coerce_user_fields_list(raw): - field_key = entry.get("field_key") - if not isinstance(field_key, str) or not field_key: - continue - if not entry.get("required", True): - continue - if not stored_values.get(field_key): - missing.append(field_key) - return missing - def _is_restricted_virtual_key_request(user_api_key_dict: UserAPIKeyAuth) -> bool: """Best-effort detection for route-restricted virtual keys. @@ -1007,8 +970,7 @@ if MCP_AVAILABLE: relevant_server_ids = [ s.server_id for s in redacted_mcp_servers - if getattr(s, "is_byok", False) - or _has_required_user_fields(getattr(s, "user_fields", None)) + if getattr(s, "is_byok", False) or server_has_user_fields(s) ] if relevant_server_ids: cred_rows = ( @@ -1036,13 +998,13 @@ if MCP_AVAILABLE: for server in redacted_mcp_servers: if getattr(server, "is_byok", False): server.has_user_credential = server.server_id in byok_set - if _has_required_user_fields(getattr(server, "user_fields", None)): + if server_has_user_fields(server): stored = user_fields_by_server.get(server.server_id, {}) - server.missing_user_field_keys = ( - _compute_missing_user_field_keys( - getattr(server, "user_fields", None), stored - ) - ) + server.missing_user_field_keys = [ + f["field_key"] + for f in compute_missing_user_fields(server, stored) + if isinstance(f.get("field_key"), str) + ] # Virtual keys only get a sanitized discovery view. if is_restricted_virtual_key: @@ -2245,16 +2207,10 @@ if MCP_AVAILABLE: # Pre-compute allowed keys from the server's declared fields, then # filter the incoming payload to only those keys. This prevents # callers from polluting the storage blob with arbitrary keys. - raw_fields = getattr(server, "user_fields", None) or [] - if isinstance(raw_fields, str): - try: - raw_fields = json.loads(raw_fields) - except (ValueError, TypeError): - raw_fields = [] declared_keys = { entry.get("field_key") - for entry in raw_fields - if isinstance(entry, dict) and entry.get("field_key") + for entry in coerce_user_fields(server) + if entry.get("field_key") } if not declared_keys: raise HTTPException( @@ -2277,7 +2233,19 @@ if MCP_AVAILABLE: else: merged[key] = value - await store_user_field_values(prisma_client, user_id, server_id, merged) + try: + await store_user_field_values(prisma_client, user_id, server_id, merged) + except ValueError as e: + # The (user, server) row already holds a BYOK or OAuth2 credential. + # Refuse rather than silently destroying the existing credential. + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail={ + "error": "credential_conflict", + "message": str(e), + "server_id": server_id, + }, + ) # Invalidate the BYOK credential cache for this (user, server) pair # so the next tool call re-reads the row. We piggyback on the @@ -2353,7 +2321,13 @@ if MCP_AVAILABLE: status_code=status.HTTP_400_BAD_REQUEST, detail={"error": "User ID not found in token"}, ) - all_servers = await get_all_mcp_servers(prisma_client) + # Only consider servers the calling user is permitted to see. This + # mirrors the per-user filtering used by the main MCP list endpoint + # and prevents leaking server IDs / user-field metadata (header + # names, env var names) to users without access. + all_servers = await get_all_mcp_servers_for_user( + prisma_client, user_api_key_dict + ) # Pre-filter to servers that actually declare user_fields, then # batch-fetch every credential row for the calling user in one # query. The N+1 alternative (per-server get_user_field_values) @@ -2361,13 +2335,7 @@ if MCP_AVAILABLE: relevant_servers: List["LiteLLM_MCPServerTable"] = [] relevant_ids: List[str] = [] for server in all_servers: - raw_fields = getattr(server, "user_fields", None) or [] - if isinstance(raw_fields, str): - try: - raw_fields = json.loads(raw_fields) - except (ValueError, TypeError): - raw_fields = [] - if not raw_fields: + if not server_has_user_fields(server): continue relevant_servers.append(server) relevant_ids.append(server.server_id)