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.
This commit is contained in:
mateo-berri 2026-06-04 13:16:55 +00:00
parent 4576c48660
commit 77ddd83c06
No known key found for this signature in database
4 changed files with 58 additions and 11 deletions

View file

@ -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)

View file

@ -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:

View file

@ -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

View file

@ -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