Fix concurrent write race + warning message in MCP user-fields

- store_user_credential / store_user_oauth_credential now use the same
  compare-and-swap (create + update_many) pattern as store_user_field_values
  so a concurrent user-fields write between the find_unique check and the
  upsert cannot silently overwrite the user-fields payload.

- store_user_field_values accepts an optional merge_fn callable that is
  re-evaluated on every CAS retry. The user-fields POST endpoint now
  passes a merge_fn so a concurrent partial save by the same user from
  another tab is folded into the result instead of being clobbered with a
  stale pre-loop merge.

- The hook_extra_headers Authorization-overwrite warning no longer claims
  the existing header came specifically from static_headers — it can now
  come from static_headers, forwarded raw headers, or user-fields.

Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
Cursor Agent 2026-05-19 08:03:00 +00:00
parent 4618e28046
commit cbb13bf53c
No known key found for this signature in database
5 changed files with 196 additions and 64 deletions

View file

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

View file

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

View file

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

View file

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

View file

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