fix(mcp): serialize per-user env var merge to avoid lost concurrent updates

store_mcp_user_env_vars read existing values, merged the submitted updates,
and wrote the result in three separate steps. Two simultaneous POSTs from the
same user to the same server could both read the same snapshot and the second
write would silently clobber the first. Move the read-modify-write into a new
merge_user_env_vars DB helper that runs inside a transaction guarded by a
(user_id, server_id) advisory lock, so concurrent writes are serialized and no
update is lost.
This commit is contained in:
mateo-berri 2026-06-04 18:54:16 +00:00 • committed by Claude
parent 113f6d8e24
commit ae026d2a4d
No known key found for this signature in database
4 changed files with 210 additions and 38 deletions

View file

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

View file

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

View file

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

View file

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