diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index ff4b72002be..b0c120ff2ce 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -3,7 +3,7 @@ import base64 import binascii import json from datetime import datetime, timedelta, timezone -from typing import Any, Dict, Iterable, List, Optional, Set, Union, cast +from typing import Any, Callable, Dict, Iterable, List, Optional, Set, Union, cast from prisma.errors import RecordNotFoundError, UniqueViolationError @@ -596,34 +596,61 @@ async def store_user_credential( BYOK, OAuth2, and user-fields payloads share the same ``credential_b64`` column. Refuse to overwrite a stored user-fields payload so saving a BYOK credential does not silently destroy the user's saved field values. + + Uses optimistic concurrency (compare-and-swap on ``credential_b64``) so + a concurrent ``store_user_field_values`` write between the read and the + write cannot be silently overwritten — mirroring the protection that + user-fields writes already have against BYOK writes. """ - # Guard against silently overwriting a user-fields payload that shares - # the same (user_id, server_id) row. - existing = await prisma_client.db.litellm_mcpusercredentials.find_unique( - where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}} - ) - if ( - existing is not None - and _decode_user_fields_payload(existing.credential_b64) is not None - ): - raise ValueError( - f"Existing credential for user {user_id} and server " - f"{server_id} holds user-fields values. Refusing to overwrite " - f"with a BYOK credential." + encoded = encrypt_value_helper(credential) + + for _attempt in range(5): + existing = await prisma_client.db.litellm_mcpusercredentials.find_unique( + where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}} ) - encoded = encrypt_value_helper(credential) - await prisma_client.db.litellm_mcpusercredentials.upsert( - where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}, - data={ - "create": { + if existing is None: + try: + await prisma_client.db.litellm_mcpusercredentials.create( + data={ + "user_id": user_id, + "server_id": server_id, + "credential_b64": encoded, + } + ) + return + except UniqueViolationError: + # A concurrent writer inserted the row first; restart so we + # can inspect what they wrote before clobbering it. + await asyncio.sleep(0) + continue + + if _decode_user_fields_payload(existing.credential_b64) is not None: + raise ValueError( + f"Existing credential for user {user_id} and server " + f"{server_id} holds user-fields values. Refusing to overwrite " + f"with a BYOK credential." + ) + + # Compare-and-swap on credential_b64: only succeed if the row still + # matches what we just inspected, so a concurrent user-fields write + # between the read and the write is not silently overwritten. + updated = await prisma_client.db.litellm_mcpusercredentials.update_many( + where={ "user_id": user_id, "server_id": server_id, - "credential_b64": encoded, + "credential_b64": existing.credential_b64, }, - "update": {"credential_b64": encoded}, - }, + data={"credential_b64": encoded}, + ) + if updated: + return + await asyncio.sleep(0) + + raise RuntimeError( + f"store_user_credential: gave up after repeated concurrent " + f"modifications for user {user_id} and server {server_id}" ) @@ -737,8 +764,10 @@ async def store_user_field_values( prisma_client: PrismaClient, user_id: str, server_id: str, - values: Dict[str, str], -) -> None: + values: Optional[Dict[str, str]] = None, + *, + merge_fn: Optional[Callable[[Dict[str, str]], Dict[str, str]]] = None, +) -> Dict[str, str]: """Persist the calling user's values for an MCP server's user fields. The full set of values is encoded into a single encrypted JSON blob in @@ -750,23 +779,56 @@ async def store_user_field_values( Refuse to overwrite a non-user-fields credential so saving user-field values does not silently destroy a stored BYOK API key or OAuth2 token. + Exactly one of ``values`` or ``merge_fn`` must be supplied: + + * ``values`` writes the dict as-is. + * ``merge_fn`` is invoked with the currently-stored user-fields values + (``{}`` if no row exists) and must return the dict to persist. It is + re-invoked on every CAS retry, so a concurrent partial save from the + same user in another tab is merged into the result instead of being + silently overwritten. + Uses optimistic concurrency to close the read-then-write race against concurrent ``store_user_credential`` / ``store_user_oauth_credential`` calls: the write is gated on the previously-observed ``credential_b64`` being unchanged, and we retry from the re-read on contention. + + Returns the dict that was ultimately persisted. """ - payload = json.dumps({"type": "user_fields", "values": values}) - encoded = encrypt_value_helper(payload) + if (values is None) == (merge_fn is None): + raise ValueError( + "store_user_field_values: exactly one of `values` or `merge_fn` " + "must be provided" + ) + + def _encode(field_values: Dict[str, str]) -> str: + return encrypt_value_helper( + json.dumps({"type": "user_fields", "values": field_values}) + ) + + # Pre-compute encoding for the ``values`` path so we don't pay encryption + # cost on every retry iteration when there is no merge function. + static_encoded: Optional[str] = None + if merge_fn is None: + assert values is not None + static_encoded = _encode(values) # Bound the retry loop; in practice contention is resolved in a single # extra round-trip but we leave headroom for pathological interleavings. - for attempt in range(5): + for _attempt in range(5): existing = await prisma_client.db.litellm_mcpusercredentials.find_unique( where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}} ) if existing is None: + if merge_fn is not None: + current_values = merge_fn({}) + encoded = _encode(current_values) + else: + assert values is not None and static_encoded is not None + current_values = values + encoded = static_encoded try: await prisma_client.db.litellm_mcpusercredentials.create( data={ @@ -775,20 +837,32 @@ async def store_user_field_values( "credential_b64": encoded, } ) - return + return current_values except UniqueViolationError: # A concurrent writer inserted the row first; restart so we # can inspect what they wrote before clobbering it. await asyncio.sleep(0) continue - if _decode_user_fields_payload(existing.credential_b64) is None: + existing_field_values = _decode_user_fields_payload(existing.credential_b64) + if existing_field_values is None: raise ValueError( f"Existing credential for user {user_id} and server " f"{server_id} is not a user-fields payload (likely BYOK or " f"OAuth2). Refusing to overwrite." ) + # When a merge function is supplied, recompute on every retry against + # the freshly-read existing values so a concurrent user-fields write + # from the same user in another tab is not silently dropped. + if merge_fn is not None: + current_values = merge_fn(existing_field_values) + encoded = _encode(current_values) + else: + assert values is not None and static_encoded is not None + current_values = values + encoded = static_encoded + # Compare-and-swap on credential_b64: only succeed if the row still # matches what we just inspected, so a concurrent BYOK/OAuth2 write # between the read and the write is not silently overwritten. @@ -801,7 +875,7 @@ async def store_user_field_values( data={"credential_b64": encoded}, ) if updated: - return + return current_values await asyncio.sleep(0) raise RuntimeError( @@ -893,38 +967,60 @@ async def store_user_oauth_credential( if scopes: payload["scopes"] = scopes - # Guard against silently overwriting a BYOK credential with an OAuth token. - # Skip the guard when the caller knows the row is already an OAuth2 credential - # (e.g. during token refresh), saving an extra DB round-trip. - if not skip_byok_guard: + encoded = encrypt_value_helper(json.dumps(payload)) + + # Optimistic concurrency: compare-and-swap on ``credential_b64`` so a + # concurrent BYOK or user-fields write between the read and the write is + # not silently overwritten. ``skip_byok_guard`` (token refresh) bypasses + # the type check but still uses CAS to avoid clobbering data. + for _attempt in range(5): existing = await prisma_client.db.litellm_mcpusercredentials.find_unique( where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}} ) + + if existing is None: + try: + await prisma_client.db.litellm_mcpusercredentials.create( + data={ + "user_id": user_id, + "server_id": server_id, + "credential_b64": encoded, + } + ) + return + except UniqueViolationError: + await asyncio.sleep(0) + continue + if ( - existing is not None + not skip_byok_guard and _decode_oauth_payload(existing.credential_b64) is None ): - # Existing row is either a BYOK secret or an OAuth2 row that no - # longer decrypts (e.g. after a salt-key rotation). In either - # case, refuse to overwrite — the caller would clobber data - # that may still be recoverable. + # Existing row is either a BYOK secret, a user-fields blob, or an + # OAuth2 row that no longer decrypts (e.g. after a salt-key + # rotation). In any case, refuse to overwrite — the caller would + # clobber data that may still be recoverable. raise ValueError( f"Existing credential for user {user_id} and server " f"{server_id} could not be verified as an OAuth2 token. " f"Refusing to overwrite." ) - encoded = encrypt_value_helper(json.dumps(payload)) - await prisma_client.db.litellm_mcpusercredentials.upsert( - where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}, - data={ - "create": { + updated = await prisma_client.db.litellm_mcpusercredentials.update_many( + where={ "user_id": user_id, "server_id": server_id, - "credential_b64": encoded, + "credential_b64": existing.credential_b64, }, - "update": {"credential_b64": encoded}, - }, + data={"credential_b64": encoded}, + ) + if updated: + return + await asyncio.sleep(0) + + raise RuntimeError( + f"store_user_oauth_credential: gave up after repeated concurrent " + f"modifications for user {user_id} and server {server_id}" ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index deb6b39555b..a2ab9d9f905 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2819,8 +2819,8 @@ class MCPServerManager: if "Authorization" in extra_headers: verbose_logger.warning( "MCPServerManager: hook_extra_headers 'Authorization' will overwrite " - "the existing Authorization header from static_headers. " - "The hook JWT will take precedence." + "the existing Authorization header (from static_headers, forwarded raw " + "headers, or user-fields). The hook JWT will take precedence." ) elif server_auth_header is not None: # server_auth_header is passed separately to _create_mcp_client as diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 3fd5c1a7f00..0209535b2c9 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -2277,19 +2277,28 @@ if MCP_AVAILABLE: ) # Merge with anything already stored — partial saves should preserve - # previously-supplied values rather than wipe them. - existing = await get_user_field_values(prisma_client, user_id, server_id) or {} - merged: Dict[str, str] = {**existing} - for key, value in payload.values.items(): - if key not in declared_keys: - continue - if value == "": - merged.pop(key, None) - else: - merged[key] = value + # previously-supplied values rather than wipe them. The merge is run + # inside ``store_user_field_values`` on every CAS retry so a + # concurrent partial save by the same user from another tab is + # folded in rather than silently overwritten. + def _merge_user_fields(existing: Dict[str, str]) -> Dict[str, str]: + merged_local: Dict[str, str] = {**existing} + for key, value in payload.values.items(): + if key not in declared_keys: + continue + if value == "": + merged_local.pop(key, None) + else: + merged_local[key] = value + return merged_local try: - await store_user_field_values(prisma_client, user_id, server_id, merged) + merged = await store_user_field_values( + prisma_client, + user_id, + server_id, + merge_fn=_merge_user_fields, + ) except ValueError as e: # The (user, server) row already holds a BYOK or OAuth2 credential. # Refuse rather than silently destroying the existing credential. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index 078adf72d4c..a533e4d79c1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -36,9 +36,13 @@ def _set_salt_key(monkeypatch): def _make_prisma_with_existing(row): """Build a MagicMock prisma_client whose user-credentials table returns ``row`` - for find_unique and behaves async-correctly for upsert/find_many.""" + for find_unique and behaves async-correctly for the CAS write paths.""" prisma = MagicMock() prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=row) + prisma.db.litellm_mcpusercredentials.create = AsyncMock() + # ``update_many`` returns the number of affected rows for CAS — 1 means + # the swap succeeded so the production loop exits. + prisma.db.litellm_mcpusercredentials.update_many = AsyncMock(return_value=1) prisma.db.litellm_mcpusercredentials.upsert = AsyncMock() prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[]) return prisma @@ -55,7 +59,20 @@ def _legacy_row(payload: str): def _stored_value(prisma) -> str: - """Pull the credential_b64 value passed to the most recent upsert call.""" + """Pull the credential_b64 value passed to the most recent write. + + The production code uses ``create`` when no row exists and + ``update_many`` (CAS) when a row already exists. Walk both call histories + and return the most recent ``credential_b64`` written. + """ + create_mock = prisma.db.litellm_mcpusercredentials.create + update_many_mock = prisma.db.litellm_mcpusercredentials.update_many + + if update_many_mock.call_args is not None: + return update_many_mock.call_args.kwargs["data"]["credential_b64"] + if create_mock.call_args is not None: + return create_mock.call_args.kwargs["data"]["credential_b64"] + # Fall back to the legacy upsert path for any tests that still rely on it. call = prisma.db.litellm_mcpusercredentials.upsert.call_args data = call.kwargs["data"] create_value = data["create"]["credential_b64"] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_user_fields.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_user_fields.py index 9c50de0c8df..325a0bb9fe7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_user_fields.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_user_fields.py @@ -545,10 +545,20 @@ async def test_post_user_field_values_rejects_undeclared_keys(): captured = {} - async def fake_store(prisma, user_id, server_id, values): + async def fake_store( + prisma, user_id, server_id, values=None, *, merge_fn=None, **kwargs + ): captured["user_id"] = user_id captured["server_id"] = server_id - captured["values"] = values + # The endpoint passes a ``merge_fn`` that closes over the request + # payload + declared_keys; resolve it against an empty existing dict + # (mirroring "no row yet") so the test can assert on the final values. + if merge_fn is not None: + result = merge_fn({}) + else: + result = values or {} + captured["values"] = result + return result async def fake_get(prisma, user_id, server_id): return None # no existing values