From 77ddd83c069d5a082b26c371d6441ae04290ccfe Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 4 Jun 2026 13:16:55 +0000 Subject: [PATCH] fix(mcp): serialize env_vars from data_dict and document reload decryption invariant Read env_vars from the exclude_unset-filtered data_dict (like every other JSON column) so a partial update that omits env_vars can never overwrite the stored values. Make the reload path's env-var encryption state explicit by passing env_vars_are_encrypted=True, since raw DB rows are still encrypted there unlike the already-decrypted records add_server/update_server receive. --- litellm/proxy/_experimental/mcp_server/db.py | 13 +++--- .../mcp_server/mcp_server_manager.py | 8 +++- .../mcp_server/test_mcp_env_vars.py | 42 +++++++++++++++++++ .../mcp_server/test_mcp_server.py | 6 +-- 4 files changed, 58 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index c7c8b0cbef7..a7796bc7a7c 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -160,14 +160,13 @@ def _prepare_mcp_server_data( if data_dict.get("static_headers") is not None: data_dict["static_headers"] = safe_dumps(data_dict["static_headers"]) - # 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) + # env_vars is read from ``data_dict`` (not ``data``) like every other JSON + # column so the exclude_unset filter is respected: a partial update that + # omits env_vars never overwrites the stored value. Global values are + # encrypted at rest before serialization. + env_vars = data_dict.get("env_vars") if env_vars is not None: - serialized_env_vars = [ - v.model_dump() if hasattr(v, "model_dump") else dict(v) for v in env_vars - ] + serialized_env_vars = [dict(v) for v in env_vars] _encrypt_global_env_var_values(serialized_env_vars) data_dict["env_vars"] = safe_dumps(serialized_env_vars) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1c41df70a5f..d3d6de5c174 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3691,7 +3691,13 @@ class MCPServerManager: verbose_logger.debug( f"Building server from DB: {server.server_id} ({server.server_name})" ) - new_server = await self.build_mcp_server_from_table(server) + # raw_rows come straight from the DB, so their global env var + # values (like credentials) are still encrypted here, unlike the + # already-decrypted records add_server/update_server are handed. + # Decrypt them while building the registry entry. + new_server = await self.build_mcp_server_from_table( + server, env_vars_are_encrypted=True + ) # Carry the cached short_prefix from the previous registry entry # (if any) so the prefix is stable across reloads. if existing_server is not None and existing_server.short_prefix: 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 06c0508cd78..ae7b4d95134 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 @@ -854,6 +854,48 @@ def test_prepare_mcp_server_data_encrypts_global_env_var_values(env_vars_salt_ke assert entries["CORP_USER"]["value"] == "placeholder-hint" +def test_prepare_mcp_server_data_skips_unset_env_vars_on_partial_update(): + """On a partial update, env_vars must follow the same exclude_unset filter as + every other JSON column: if the caller never set env_vars, the field must not + be written, even when the request object carries a non-None env_vars that was + never marked as set. Otherwise a partial update could silently overwrite the + stored values.""" + from litellm.proxy._experimental.mcp_server.db import _prepare_mcp_server_data + from litellm.proxy._types import MCPEnvVar, UpdateMCPServerRequest + + data = UpdateMCPServerRequest.model_construct( + _fields_set={"server_id"}, + server_id="srv-1", + env_vars=[MCPEnvVar(name="DB_PASSWORD", value="s3cr3t", scope="global")], + ) + + prepared = _prepare_mcp_server_data(data, exclude_unset=True) + + assert "env_vars" not in prepared + + +def test_prepare_mcp_server_data_writes_env_vars_when_set_on_partial_update( + env_vars_salt_key, +): + """A partial update that does set env_vars must serialize and encrypt them.""" + import json + + from litellm.proxy._experimental.mcp_server.db import _prepare_mcp_server_data + from litellm.proxy._types import MCPEnvVar, UpdateMCPServerRequest + + data = UpdateMCPServerRequest( + server_id="srv-1", + env_vars=[MCPEnvVar(name="DB_PASSWORD", value="s3cr3t", scope="global")], + ) + + prepared = _prepare_mcp_server_data(data, exclude_unset=True) + + assert "env_vars" in prepared + entries = json.loads(prepared["env_vars"]) + assert entries[0]["name"] == "DB_PASSWORD" + assert entries[0]["value"] != "s3cr3t" + + @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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 6b6c7bc37d5..22217e7f5ef 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -3788,7 +3788,7 @@ class TestMCPServerManagerReload: ): await manager.reload_servers_from_database() - mock_build.assert_awaited_once_with(db_row) + mock_build.assert_awaited_once_with(db_row, env_vars_are_encrypted=True) assert manager.registry["server-1"] is rebuilt_server @pytest.mark.asyncio @@ -3819,7 +3819,7 @@ class TestMCPServerManagerReload: updated_at=timestamp, ) - async def build_server(db_row): + async def build_server(db_row, **kwargs): if db_row.server_id == "bad-server": raise RuntimeError("transient build failure") if db_row.server_id == "healthy-server": @@ -3885,7 +3885,7 @@ class TestMCPServerManagerReload: updated_at=timestamp, ) - async def build_server(db_row): + async def build_server(db_row, **kwargs): if db_row.server_id == "healthy-server": return healthy_server return bad_openapi_server