mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
4618e28046
commit
cbb13bf53c
5 changed files with 196 additions and 64 deletions
|
|
@ -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}"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue