diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 911755a75be..5aa910deab7 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -11,6 +11,7 @@ from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, MCPApprovalStatus, + MCPEnvVarScope, MCPSubmissionsSummary, NewMCPServerRequest, SpecialMCPServerName, @@ -28,6 +29,57 @@ from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.mcp import MCPCredentials +def _is_global_env_var_scope(scope: Any) -> bool: + """``scope="user"`` entries are placeholders the user fills in; everything + else (including a missing scope) is an admin-supplied global value.""" + return scope != MCPEnvVarScope.user and scope != "user" + + +def _encrypt_global_env_var_values(env_vars: Iterable[Dict[str, Any]]) -> None: + """Encrypt ``scope="global"`` env var values in place before persisting. + + Global values hold admin-supplied secrets (API keys, passwords) that get + interpolated into headers, so they are encrypted at rest like credentials + and the per-user ``values_b64`` column. Per-user placeholders are not + secrets and are stored verbatim. + """ + for entry in env_vars: + if not _is_global_env_var_scope(entry.get("scope")): + continue + value = entry.get("value") + if value: + entry["value"] = encrypt_value_helper(value) + + +def decrypt_global_env_var_values(env_vars: Optional[Iterable[Any]]) -> None: + """Decrypt ``scope="global"`` env var values in place after reading the DB. + + Accepts ``MCPEnvVar`` models (``LiteLLM_MCPServerTable``) or plain dicts + (raw rows / deserialized JSON). Decryption is best-effort: a value that no + longer decrypts (e.g. after a salt-key change) is left untouched. + """ + if not env_vars: + return + for entry in env_vars: + is_dict = isinstance(entry, dict) + scope = entry.get("scope") if is_dict else getattr(entry, "scope", None) + if not _is_global_env_var_scope(scope): + continue + value = entry.get("value") if is_dict else getattr(entry, "value", None) + if not value: + continue + decrypted = decrypt_value_helper( + value=value, + key="mcp_global_env_var", + exception_type="debug", + return_original_value=True, + ) + if is_dict: + entry["value"] = decrypted + else: + entry.value = decrypted + + def _prepare_mcp_server_data( data: Union[NewMCPServerRequest, UpdateMCPServerRequest], exclude_unset: bool = False, @@ -100,11 +152,14 @@ def _prepare_mcp_server_data( # Handle env_vars serialization. Pydantic models are dumped to a list of # plain dicts so the JSON column receives ``[{name, value, scope, ...}]``. + # Global values are encrypted at rest before serialization. env_vars = getattr(data, "env_vars", None) if env_vars is not None: - data_dict["env_vars"] = safe_dumps( - [v.model_dump() if hasattr(v, "model_dump") else dict(v) for v in env_vars] - ) + serialized_env_vars = [ + v.model_dump() if hasattr(v, "model_dump") else dict(v) for v in env_vars + ] + _encrypt_global_env_var_values(serialized_env_vars) + data_dict["env_vars"] = safe_dumps(serialized_env_vars) if data_dict.get("mcp_info") is not None: data_dict["mcp_info"] = safe_dumps(data_dict["mcp_info"]) @@ -215,10 +270,13 @@ async def get_all_mcp_servers( where=where if where else {} ) - return [ + tables = [ LiteLLM_MCPServerTable(**mcp_server.model_dump()) for mcp_server in mcp_servers ] + for table in tables: + decrypt_global_env_var_values(table.env_vars) + return tables except Exception as e: verbose_proxy_logger.debug( "litellm.proxy._experimental.mcp_server.db.py::get_all_mcp_servers - {}".format( @@ -241,7 +299,9 @@ async def get_mcp_server( ) if mcp_server is None: return None - return LiteLLM_MCPServerTable(**mcp_server.model_dump()) + table = LiteLLM_MCPServerTable(**mcp_server.model_dump()) + decrypt_global_env_var_values(table.env_vars) + return table async def get_mcp_servers( @@ -259,7 +319,9 @@ async def get_mcp_servers( ) final_mcp_servers: List[LiteLLM_MCPServerTable] = [] for _mcp_server in _mcp_servers: - final_mcp_servers.append(LiteLLM_MCPServerTable(**_mcp_server.model_dump())) + table = LiteLLM_MCPServerTable(**_mcp_server.model_dump()) + decrypt_global_env_var_values(table.env_vars) + final_mcp_servers.append(table) return final_mcp_servers @@ -437,6 +499,7 @@ async def create_mcp_server( data=data_dict # type: ignore ) + decrypt_global_env_var_values(getattr(new_mcp_server, "env_vars", None)) return new_mcp_server @@ -514,6 +577,7 @@ async def update_mcp_server( where={"server_id": data.server_id}, data=data_dict # type: ignore ) + decrypt_global_env_var_values(getattr(updated_mcp_server, "env_vars", None)) return updated_mcp_server @@ -935,7 +999,9 @@ async def approve_mcp_server( "updated_by": touched_by, }, ) - return LiteLLM_MCPServerTable(**updated.model_dump()) + table = LiteLLM_MCPServerTable(**updated.model_dump()) + decrypt_global_env_var_values(table.env_vars) + return table async def reject_mcp_server( @@ -957,7 +1023,9 @@ async def reject_mcp_server( where={"server_id": server_id}, data=data, ) - return LiteLLM_MCPServerTable(**updated.model_dump()) + table = LiteLLM_MCPServerTable(**updated.model_dump()) + decrypt_global_env_var_values(table.env_vars) + return table async def get_mcp_submissions( @@ -974,6 +1042,8 @@ async def get_mcp_submissions( take=500, # safety cap; paginate if needed in a future iteration ) items = [LiteLLM_MCPServerTable(**r.model_dump()) for r in rows] + for item in items: + decrypt_global_env_var_values(item.env_vars) pending = sum( 1 for i in items if i.approval_status == MCPApprovalStatus.pending_review diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4af4a6aff07..a24bedb1a8f 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -880,6 +880,12 @@ class MCPServerManager: getattr(mcp_server, "static_headers", None) ) env_vars_list = _deserialize_json_list(getattr(mcp_server, "env_vars", None)) + if credentials_are_encrypted: + from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 + decrypt_global_env_var_values, + ) + + decrypt_global_env_var_values(env_vars_list) credentials_dict = _deserialize_json_dict( getattr(mcp_server, "credentials", None) ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py index 52db2f4de9a..f17a99a1756 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -588,6 +588,95 @@ async def test_delete_user_env_vars_is_idempotent_delete_many(): assert call.kwargs["where"] == {"user_id": "alice", "server_id": "srv-1"} +# ── DB helpers: global env vars encrypted at rest ───────────────────────── + + +def _global_env_var_server_request(env_vars): + from litellm.proxy._types import NewMCPServerRequest + + return NewMCPServerRequest( + alias="echo", + url="https://upstream.example.com/mcp", + transport="http", + auth_type="none", + static_headers={"X-Db": "${DB_PASSWORD}"}, + env_vars=env_vars, + ) + + +def test_prepare_mcp_server_data_encrypts_global_env_var_values(env_vars_salt_key): + """``scope="global"`` secrets must be encrypted before they reach the JSON + column, while ``scope="user"`` placeholders (not secrets) stay verbatim.""" + import json + + from litellm.proxy._experimental.mcp_server.db import ( + _prepare_mcp_server_data, + decrypt_global_env_var_values, + ) + from litellm.proxy._types import MCPEnvVar + + req = _global_env_var_server_request( + [ + MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global"), + MCPEnvVar( + name="CORP_USER", + value="placeholder-hint", + scope="user", + description="your db user", + ), + ] + ) + + stored = _prepare_mcp_server_data(req)["env_vars"] + entries = {e["name"]: e for e in json.loads(stored)} + + # The global secret is unrecoverable from the stored JSON ... + assert "s3cr3t-p@ss" not in stored + assert entries["DB_PASSWORD"]["value"] != "s3cr3t-p@ss" + # ... but the per-user placeholder is stored as-is. + assert entries["CORP_USER"]["value"] == "placeholder-hint" + + # And the encrypted global decrypts back to the original secret. + decrypt_global_env_var_values(list(entries.values())) + assert entries["DB_PASSWORD"]["value"] == "s3cr3t-p@ss" + assert entries["CORP_USER"]["value"] == "placeholder-hint" + + +@pytest.mark.asyncio +async def test_build_mcp_server_from_table_decrypts_global_env_vars(env_vars_salt_key): + """End-to-end: an encrypted global value persisted in the DB must be + decrypted when the server is built into the runtime registry, so ``${NAME}`` + headers interpolate to the real secret instead of forwarding ciphertext.""" + import json + + from litellm.proxy._experimental.mcp_server.db import _prepare_mcp_server_data + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + from litellm.proxy._types import LiteLLM_MCPServerTable, MCPEnvVar + + req = _global_env_var_server_request( + [MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")] + ) + prepared = _prepare_mcp_server_data(req) + + table = LiteLLM_MCPServerTable( + server_id="srv-global", + alias="echo", + url="https://upstream.example.com/mcp", + transport="http", + auth_type="none", + static_headers={"X-Db": "${DB_PASSWORD}"}, + env_vars=json.loads(prepared["env_vars"]), + ) + + manager = MCPServerManager() + server = await manager.build_mcp_server_from_table(table) + + headers = await manager._resolve_static_headers_with_env_vars(server, None) + assert headers == {"X-Db": "s3cr3t-p@ss"} + + # ── REST exception handling ───────────────────────────────────────────────