mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
113f6d8e24
commit
ae026d2a4d
4 changed files with 210 additions and 38 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue