diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 34093fa7638..5c9b7c5588a 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -1250,6 +1250,47 @@ async def get_user_env_vars_bulk( return {row.server_id: _decode_user_env_vars(row.values_b64) for row in rows} +async def merge_user_env_vars( + prisma_client: PrismaClient, + user_id: str, + server_id: str, + updates: Dict[str, str], + allowed_names: Iterable[str], +) -> Dict[str, str]: + """Merge ``updates`` into the user's stored env vars for ``server_id`` and + return the resulting set. + + The read-modify-write runs inside a transaction guarded by a + ``(user_id, server_id)`` advisory lock so two concurrent writes from the + same user can't drop one update. Names outside ``allowed_names`` are pruned, + so an admin retiring a user-scoped variable also clears its stored value. + """ + allowed = set(allowed_names) + async with prisma_client.db.tx() as tx: + await tx.query_raw( + "SELECT pg_advisory_xact_lock(hashtextextended($1, 0))", + f"{user_id}:{server_id}", + ) + row = await tx.litellm_mcpuserenvvars.find_unique( + where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}} + ) + existing = _decode_user_env_vars(row.values_b64) if row is not None else {} + merged = {k: v for k, v in {**existing, **updates}.items() if k in allowed} + encoded = encrypt_value_helper(json.dumps(merged)) + await tx.litellm_mcpuserenvvars.upsert( + where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}, + data={ + "create": { + "user_id": user_id, + "server_id": server_id, + "values_b64": encoded, + }, + "update": {"values_b64": encoded}, + }, + ) + return merged + + async def delete_user_env_vars( prisma_client: PrismaClient, user_id: str, diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index ae2ffe0f89a..c06fd1ff128 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -123,9 +123,9 @@ if MCP_AVAILABLE: get_user_env_vars_bulk, get_user_oauth_credential, list_user_oauth_credentials, + merge_user_env_vars, reject_mcp_server, store_user_credential, - store_user_env_vars, store_user_oauth_credential, update_mcp_server, ) @@ -2321,11 +2321,9 @@ if MCP_AVAILABLE: updates = { k: v for k, v in payload.values.items() if k in allowed_names and v != "" } - existing = await get_user_env_vars(prisma_client, user_id, server_id) - merged = { - k: v for k, v in {**existing, **updates}.items() if k in allowed_names - } - await store_user_env_vars(prisma_client, user_id, server_id, merged) + merged = await merge_user_env_vars( + prisma_client, user_id, server_id, updates, allowed_names + ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( invalidate_user_env_vars_cache, ) 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 2dcc8a8390e..9bd726d7962 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 @@ -818,6 +818,88 @@ def _captured_values_blob(prisma) -> str: return create_value +def _transactional_env_vars_prisma(read_delay: float = 0.0): + """A prisma stand-in backed by an in-memory store that honours + ``db.tx()`` and the ``pg_advisory_xact_lock`` advisory lock. + + ``read_delay`` inserts an ``await`` point inside ``find_unique`` so two + concurrent merges interleave between their read and write; the advisory lock + is what keeps them from clobbering each other. Drop the lock and the second + write wins, losing the first update. + """ + import asyncio + from unittest.mock import MagicMock + + class _Store: + def __init__(self): + self.rows = {} + self.locks = {} + + class _Table: + def __init__(self, store, delay=0.0): + self._store = store + self._delay = delay + + async def find_unique(self, where): + ident = where["user_id_server_id"] + key = (ident["user_id"], ident["server_id"]) + blob = self._store.rows.get(key) + # Yield after capturing the read so an unserialised concurrent merge + # would race on this stale snapshot. + if self._delay: + await asyncio.sleep(self._delay) + if blob is None: + return None + row = MagicMock() + row.values_b64 = blob + return row + + async def upsert(self, where, data): + ident = where["user_id_server_id"] + key = (ident["user_id"], ident["server_id"]) + self._store.rows[key] = data["update"]["values_b64"] + + async def delete_many(self, where): + self._store.rows.pop((where["user_id"], where["server_id"]), None) + + class _Tx: + def __init__(self, store, delay): + self._store = store + self._held = None + self.litellm_mcpuserenvvars = _Table(store, delay=delay) + + async def __aenter__(self): + return self + + async def __aexit__(self, *exc): + if self._held is not None: + self._held.release() + self._held = None + return False + + async def query_raw(self, query, *args): + lock_key = args[0] + lock = self._store.locks.setdefault(lock_key, asyncio.Lock()) + await lock.acquire() + self._held = lock + return [{"pg_advisory_xact_lock": None}] + + class _DB: + def __init__(self, store, delay): + self._store = store + self._delay = delay + self.litellm_mcpuserenvvars = _Table(store) + + def tx(self): + return _Tx(self._store, self._delay) + + class _Prisma: + def __init__(self, delay): + self.db = _DB(_Store(), delay) + + return _Prisma(read_delay) + + @pytest.mark.asyncio async def test_store_user_env_vars_does_not_persist_plaintext(env_vars_salt_key): from litellm.proxy._experimental.mcp_server.db import store_user_env_vars @@ -915,6 +997,62 @@ async def test_delete_user_env_vars_is_idempotent_delete_many(): assert call.kwargs["where"] == {"user_id": "alice", "server_id": "srv-1"} +@pytest.mark.asyncio +async def test_merge_user_env_vars_preserves_existing_and_prunes_disallowed( + env_vars_salt_key, +): + """Merging one update keeps the user's other stored values and drops any + name the admin no longer declares as user-scoped.""" + from litellm.proxy._experimental.mcp_server.db import ( + merge_user_env_vars, + store_user_env_vars, + ) + + prisma = _transactional_env_vars_prisma() + await store_user_env_vars( + prisma, + "alice", + "srv-1", + {"CORP_USERNAME": "alice", "CORP_PASSWORD": "old", "RETIRED": "x"}, + ) + + merged = await merge_user_env_vars( + prisma, + "alice", + "srv-1", + {"CORP_PASSWORD": "new"}, + {"CORP_USERNAME", "CORP_PASSWORD"}, + ) + + # CORP_USERNAME survives, CORP_PASSWORD updates, RETIRED (no longer declared) + # is pruned. + assert merged == {"CORP_USERNAME": "alice", "CORP_PASSWORD": "new"} + + +@pytest.mark.asyncio +async def test_merge_user_env_vars_serializes_concurrent_writes(env_vars_salt_key): + """Two simultaneous merges for the same (user, server) must not lose an + update: the advisory-locked transaction serialises the read-modify-write so + both distinct values survive.""" + import asyncio + + from litellm.proxy._experimental.mcp_server.db import ( + get_user_env_vars, + merge_user_env_vars, + ) + + allowed = {"TOKEN_A", "TOKEN_B"} + prisma = _transactional_env_vars_prisma(read_delay=0.02) + + await asyncio.gather( + merge_user_env_vars(prisma, "alice", "srv-1", {"TOKEN_A": "a"}, allowed), + merge_user_env_vars(prisma, "alice", "srv-1", {"TOKEN_B": "b"}, allowed), + ) + + stored = await get_user_env_vars(prisma, "alice", "srv-1") + assert stored == {"TOKEN_A": "a", "TOKEN_B": "b"} + + @pytest.mark.asyncio async def test_delete_mcp_server_removes_orphaned_user_env_vars(): """Deleting a server must also drop every user's per-user env var rows for diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index ed4e80d711c..7c3e2393ee9 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -3908,7 +3908,7 @@ class TestStoreMCPUserEnvVars: server = _make_env_var_server( env_vars=_ENV_VARS_MIXED, static_headers=_STATIC_HEADERS_MIXED ) - store_mock = AsyncMock() + merge_mock = AsyncMock(return_value={"CORP_USERNAME": "alice"}) with ( patch.object( mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() @@ -3916,10 +3916,7 @@ class TestStoreMCPUserEnvVars: patch.object( mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=server) ), - patch.object( - mgmt_endpoints, "get_user_env_vars", AsyncMock(return_value={}) - ), - patch.object(mgmt_endpoints, "store_user_env_vars", store_mock), + patch.object(mgmt_endpoints, "merge_user_env_vars", merge_mock), ): result = await mgmt_endpoints.store_mcp_user_env_vars( server_id="srv-1", @@ -3932,25 +3929,30 @@ class TestStoreMCPUserEnvVars: ), user_api_key_dict=generate_mock_user_api_key_auth(user_id="alice"), ) - # Only the declared, non-empty value is persisted. - store_mock.assert_awaited_once() - _, _, _, persisted = store_mock.await_args.args - assert persisted == {"CORP_USERNAME": "alice"} + # Only the declared, non-empty value reaches the atomic merge, scoped to + # the admin-declared user vars. + merge_mock.assert_awaited_once() + _, _, _, updates, allowed_names = merge_mock.await_args.args + assert updates == {"CORP_USERNAME": "alice"} + assert set(allowed_names) == { + "CORP_USERNAME", + "CORP_PASSWORD", + "UNUSED_USER_VAR", + } # CORP_PASSWORD remains unset in the returned status. assert result.missing_count == 1 @pytest.mark.asyncio - async def test_merges_over_existing_values(self): - """Updating one credential must not wipe other already-stored values. - - The user updates only CORP_PASSWORD; their previously-stored - CORP_USERNAME (write-only, never shown back in the form) must be - preserved instead of being cleared. - """ + async def test_forwards_only_submitted_updates_and_returns_merged_status(self): + """The endpoint forwards only the user's submitted (allowed, non-empty) + update to the atomic merge and reports status from the merged result, so + a one-field edit never sends the other stored values back through.""" server = _make_env_var_server( env_vars=_ENV_VARS_MIXED, static_headers=_STATIC_HEADERS_MIXED ) - store_mock = AsyncMock() + merge_mock = AsyncMock( + return_value={"CORP_USERNAME": "alice", "CORP_PASSWORD": "new"} + ) with ( patch.object( mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() @@ -3958,14 +3960,7 @@ class TestStoreMCPUserEnvVars: patch.object( mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=server) ), - patch.object( - mgmt_endpoints, - "get_user_env_vars", - AsyncMock( - return_value={"CORP_USERNAME": "alice", "CORP_PASSWORD": "old"} - ), - ), - patch.object(mgmt_endpoints, "store_user_env_vars", store_mock), + patch.object(mgmt_endpoints, "merge_user_env_vars", merge_mock), ): result = await mgmt_endpoints.store_mcp_user_env_vars( server_id="srv-1", @@ -3974,10 +3969,10 @@ class TestStoreMCPUserEnvVars: ), user_api_key_dict=generate_mock_user_api_key_auth(user_id="alice"), ) - store_mock.assert_awaited_once() - _, _, _, persisted = store_mock.await_args.args - assert persisted == {"CORP_USERNAME": "alice", "CORP_PASSWORD": "new"} - # Both credentials report as set in the returned status. + merge_mock.assert_awaited_once() + _, _, _, updates, _ = merge_mock.await_args.args + assert updates == {"CORP_PASSWORD": "new"} + # Status reflects the merged set returned by the atomic merge. assert result.missing_count == 0 @pytest.mark.asyncio @@ -4228,7 +4223,7 @@ class TestMCPUserEnvVarsAccessControl: server = _make_env_var_server( env_vars=_ENV_VARS_MIXED, static_headers=_STATIC_HEADERS_MIXED ) - store_mock = AsyncMock() + merge_mock = AsyncMock() with ( patch.object( mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() @@ -4241,7 +4236,7 @@ class TestMCPUserEnvVarsAccessControl: "get_all_mcp_servers_for_user", AsyncMock(return_value=[]), ), - patch.object(mgmt_endpoints, "store_user_env_vars", store_mock), + patch.object(mgmt_endpoints, "merge_user_env_vars", merge_mock), ): with pytest.raises(HTTPException) as exc: await mgmt_endpoints.store_mcp_user_env_vars( @@ -4255,7 +4250,7 @@ class TestMCPUserEnvVarsAccessControl: ), ) assert exc.value.status_code == 403 - store_mock.assert_not_awaited() + merge_mock.assert_not_awaited() @pytest.mark.asyncio async def test_clear_forbidden_for_non_admin_without_access(self):