Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/db.py: 51%
860 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import base64
2import binascii
3import hashlib
4import json
5from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
6from dataclasses import dataclass
7from datetime import datetime, timedelta, timezone
8from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict, cast
10from fastapi import HTTPException
11from typing_extensions import ReadOnly
13from litellm._logging import verbose_proxy_logger
14from litellm._uuid import uuid
15from litellm.constants import MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS
16from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
17from litellm.proxy._experimental.mcp_server.oauth_identity_binding import (
18 RefreshTokenPresented,
19 credential_binding_matches,
20 enforce_oauth_identity_binding,
21)
22from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request
23from litellm.proxy._types import (
24 LiteLLM_MCPServerTable,
25 MCPApprovalStatus,
26 MCPEnvVar,
27 MCPEnvVarScope,
28 MCPServerUserCredentialListItem,
29 MCPSubmissionsSummary,
30 NewMCPServerRequest,
31 SpecialMCPServerName,
32 UpdateMCPServerRequest,
33 UserAPIKeyAuth,
34)
35from litellm.proxy.common_utils.encrypt_decrypt_utils import (
36 SecretMapDecodeError,
37 _get_salt_key,
38 decode_secret_map,
39 decrypt_value_helper,
40 encrypt_secret_map,
41 encrypt_value_helper,
42)
43from litellm.proxy.utils import PrismaClient
44from litellm.repositories.object_permission_repository import ObjectPermissionRepository
45from litellm.repositories.prisma_protocols import TableActions
46from litellm.repositories.table_repositories import (
47 MCPServerOAuthClientRepository,
48 MCPServerRepository,
49 MCPUserCredentialsRepository,
50 PrismaTableRepository,
51)
52from litellm.repositories.team_repository import TeamRepository
53from litellm.repositories.verification_token_repository import (
54 VerificationTokenRepository,
55)
56from litellm.types.llms.custom_http import httpxSpecialProvider
57from litellm.types.mcp import MCPCredentials
59if TYPE_CHECKING: 59 ↛ 60line 59 didn't jump to line 60 because the condition on line 59 was never true
60 from prisma import models as prisma_db_models
61 from prisma import types as prisma_db_types
63 from litellm.types.mcp_server.mcp_server_manager import MCPServer
66class _UserEnvVarsTransactionClient(Protocol):
67 litellm_mcpuserenvvars: "TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]"
68 litellm_mcpservertable: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]"
70 async def execute_raw(self, query: str, *args: object) -> int: ... 70 ↛ exitline 70 didn't return from function 'execute_raw' because
73class _UserEnvVarsTransaction(Protocol):
74 async def __aenter__(self) -> _UserEnvVarsTransactionClient: ... 74 ↛ exitline 74 didn't return from function '__aenter__' because
76 async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ... 76 ↛ exitline 76 didn't return from function '__aexit__' because
79@dataclass(frozen=True, slots=True)
80class McpIdentifierConflict:
81 """An incoming ``server_name``/``alias`` already belongs to another MCP server row.
83 ``field`` is the incoming identifier that collided, ``value`` the submitted
84 string, and ``server_id`` the existing row that owns it.
85 """
87 field: Literal["server_name", "alias"]
88 value: str
89 server_id: str
92_AUTH_FLOW_SCOPED_FIELDS: Final["frozenset[str]"] = frozenset(
93 {
94 "issuer",
95 "authorization_url",
96 "token_url",
97 "registration_url",
98 "oauth2_flow",
99 "dcr_bridge",
100 "token_exchange_endpoint",
101 "audience",
102 "subject_token_type",
103 "token_exchange_profile",
104 }
105)
108def _blank_to_none(value: str | None) -> str | None:
109 if not isinstance(value, str): 109 ↛ 111line 109 didn't jump to line 111 because the condition on line 109 was always true
110 return None
111 return value.strip() or None
114# Token-exchange settings with dedicated columns that also exist on
115# ``MCPCredentials`` as a legacy shape (rows and REST callers that predate the
116# columns). Every write lifts blob values into the columns and strips them from
117# the stored blob, so the read-time ``column or blob`` fallback only serves rows
118# the current code has never written — a cleared column can then never be
119# silently resurrected by a stale blob copy. These keys are stored plaintext
120# (endpoints/identifiers, not secrets), so values lift as-is.
121_TOKEN_EXCHANGE_COLUMN_FIELDS: Final["frozenset[str]"] = frozenset(
122 {
123 "token_exchange_endpoint",
124 "audience",
125 "subject_token_type",
126 "token_exchange_profile",
127 }
128)
130# The client-forwarded token modes share one stored-credential shape: the admin-declared upstream
131# OAuth app (client_id/client_secret) plus the same authorize relay, and neither mints anything the
132# gateway keeps. So a switch WITHIN this class must preserve the stored app, unlike a cross-class
133# switch (e.g. an oauth2 row whose client may be DCR-minted and is not reusable elsewhere).
134_CLIENT_FORWARDED_AUTH_TYPES: Final["frozenset[str]"] = frozenset({"true_passthrough", "oauth_delegate"})
136# Minted token material that must never survive a client rotation on a persisted row.
137_MINTED_TOKEN_CREDENTIAL_FIELDS: Final["frozenset[str]"] = frozenset({"access_token", "refresh_token", "expires_in"})
140class _OAuthCredentialAccessToken(TypedDict):
141 access_token: str
144class OAuthCredentialPayload(_OAuthCredentialAccessToken, total=False):
145 identity_binding_proof: ReadOnly[str]
146 type: str
147 refresh_token: str
148 expires_at: str
149 connected_at: str
150 scopes: list[str]
151 server_id: str
154OAuthGrantState = Literal["valid", "refreshable", "absent"]
157class _OAuthTokenRefreshResponse(TypedDict, total=False):
158 access_token: str
159 refresh_token: str
160 expires_in: int
161 scope: str
164def _credential_auth_class(auth_type: str | None) -> str | None:
165 """Collapse the client-forwarded modes to one credential class; every other auth_type is its own
166 class. Used so credential handling keys off whether the stored-credential shape actually changed,
167 not off a raw auth_type inequality that treats true_passthrough<->oauth_delegate as a full reset."""
168 if auth_type in _CLIENT_FORWARDED_AUTH_TYPES:
169 return "client_forwarded"
170 return auth_type
173def _drop_stale_minted_on_client_rotation(merged: dict[str, object], new_creds: dict[str, object]) -> dict[str, object]:
174 """When the update rotates the client, drop stale minted token keys it did not itself set, so an old
175 app's access/refresh token never rides forward under the new client. A no-op when no client key changed."""
176 if "client_id" not in new_creds and "client_secret" not in new_creds:
177 return merged
178 return {
179 key: value for key, value in merged.items() if key not in _MINTED_TOKEN_CREDENTIAL_FIELDS or key in new_creds
180 }
183def _is_global_env_var_scope(scope: object) -> bool:
184 """``scope="user"`` entries are placeholders the user fills in; everything
185 else (including a missing scope) is an admin-supplied global value."""
186 return scope != MCPEnvVarScope.user and scope != "user"
189def _encrypt_global_env_var_values(env_vars: Iterable[dict[str, str]]) -> None:
190 """Encrypt ``scope="global"`` env var values in place before persisting.
192 Global values hold admin-supplied secrets (API keys, passwords) that get
193 interpolated into headers, so they are encrypted at rest like credentials
194 and the per-user ``values_b64`` column. Per-user placeholders are not
195 secrets and are stored verbatim.
196 """
197 for entry in env_vars:
198 if not _is_global_env_var_scope(entry.get("scope")):
199 continue
200 value = entry.get("value")
201 if value:
202 entry["value"] = encrypt_value_helper(value)
205def decrypt_global_env_var_values(env_vars: Iterable[MCPEnvVar | dict[str, str]] | None) -> None:
206 """Decrypt ``scope="global"`` env var values in place after reading the DB.
208 Accepts ``MCPEnvVar`` models (``LiteLLM_MCPServerTable``) or plain dicts
209 (raw rows / deserialized JSON). Global values are always stored encrypted,
210 so a value that no longer decrypts (e.g. after a salt-key change) is dropped
211 and a warning is logged rather than forwarding the ciphertext into upstream
212 ``${NAME}`` headers, where it would silently fail.
213 """
214 if not env_vars:
215 return
216 for entry in env_vars:
217 is_dict = isinstance(entry, dict)
218 scope = entry.get("scope") if is_dict else getattr(entry, "scope", None)
219 if not _is_global_env_var_scope(scope):
220 continue
221 value = entry.get("value") if is_dict else getattr(entry, "value", None)
222 if not value:
223 continue
224 decrypted = decrypt_value_helper(
225 value=value,
226 key="mcp_global_env_var",
227 exception_type="debug",
228 return_original_value=False,
229 )
230 if decrypted is None: 230 ↛ 231line 230 didn't jump to line 231 because the condition on line 230 was never true
231 name = entry.get("name") if is_dict else getattr(entry, "name", None)
232 verbose_proxy_logger.warning(
233 "MCP global env var %s failed to decrypt (LITELLM_SALT_KEY "
234 "changed?); dropping it so ciphertext is not sent upstream",
235 name,
236 )
237 decrypted = ""
238 if is_dict:
239 entry["value"] = decrypted
240 else:
241 entry.value = decrypted
244def _decrypt_env_vars_on_returned_row(row: object) -> None:
245 """Decrypt ``scope="global"`` env var values on a row returned by Prisma create/update.
247 Prisma may hand back ``env_vars`` either as a parsed list (the common case for
248 JSONB columns) or as a raw JSON string (observed for some write paths). The
249 in-place decrypt helper only mutates iterables of dicts/models, so a string
250 payload would silently skip decryption and ciphertext would leak into the
251 registry via ``add_server``/``update_server`` (which trust the caller).
252 Parse the string back to a list so the in-place decrypt actually runs, and
253 write the decrypted list back onto the row so downstream consumers see plain
254 values.
255 """
256 env_vars = getattr(row, "env_vars", None)
257 if env_vars is None:
258 return
259 if isinstance(env_vars, str): 259 ↛ 260line 259 didn't jump to line 260 because the condition on line 259 was never true
260 try:
261 env_vars = json.loads(env_vars)
262 except (json.JSONDecodeError, TypeError):
263 return
264 if not isinstance(env_vars, list):
265 return
266 try:
267 setattr(row, "env_vars", env_vars)
268 except (AttributeError, TypeError):
269 pass
270 decrypt_global_env_var_values(env_vars)
273def _reencrypt_global_env_var_values(
274 env_vars: str | Iterable[Mapping[str, str]] | None, new_encryption_key: str
275) -> list[dict[str, str]] | None:
276 """Re-encrypt ``scope="global"`` env var values for master-key rotation.
278 Each global value is decrypted with the current salt key and re-encrypted
279 under ``new_encryption_key``. Returns the rebuilt list when at least one
280 value was rotated, else ``None`` so the caller can skip the DB write. A
281 value that fails to decrypt is left untouched (and logged) so a corrupt
282 entry is preserved for recovery rather than overwritten.
283 """
284 if not env_vars:
285 return None
286 entries: Iterable[Mapping[str, str]]
287 if isinstance(env_vars, str):
288 try:
289 entries = json.loads(env_vars)
290 except (json.JSONDecodeError, TypeError):
291 return None
292 if not entries:
293 return None
294 else:
295 entries = env_vars
296 rebuilt: Final = [dict(v) for v in entries]
297 rotated = False
298 for entry in rebuilt:
299 if not _is_global_env_var_scope(entry.get("scope")):
300 continue
301 value = entry.get("value")
302 if not value:
303 continue
304 decrypted = decrypt_value_helper(
305 value=value,
306 key="mcp_global_env_var",
307 exception_type="debug",
308 return_original_value=False,
309 )
310 if decrypted is None:
311 verbose_proxy_logger.warning(
312 "rotate_mcp_server_credentials_master_key: could not decrypt global env var %s, skipping",
313 entry.get("name"),
314 )
315 continue
316 entry["value"] = encrypt_value_helper(decrypted, new_encryption_key=new_encryption_key)
317 rotated = True
318 return rebuilt if rotated else None
321def _prepare_mcp_server_data(
322 data: NewMCPServerRequest | UpdateMCPServerRequest,
323 exclude_unset: bool = False,
324 fields_set: set[str] | None = None,
325) -> dict[str, Any]:
326 """
327 Helper function to prepare MCP server data for database operations.
328 Handles JSON field serialization for mcp_info and env fields.
330 Args:
331 data: NewMCPServerRequest or UpdateMCPServerRequest object
332 exclude_unset: When True, only fields the caller explicitly provided are
333 included. Used for partial updates (PUT /v1/mcp/server) so omitted
334 fields keep their existing DB value instead of being silently reset
335 to a Pydantic schema default. ``exclude_none`` is not enough here:
336 non-Optional fields (e.g. ``transport=MCPTransport.sse``,
337 ``mcp_access_groups=[]``, ``allow_all_keys=False``) are backfilled
338 with their default when omitted, and a non-None default survives the
339 ``exclude_none`` filter and overwrites the row.
341 Returns:
342 Dict with properly serialized JSON fields
343 """
344 from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
346 # Convert model to dict.
347 # - Partial update (exclude_unset): only caller-provided keys are emitted, so
348 # omitted fields are never written and keep their existing DB value.
349 # - Create (exclude_none): drop None-valued fields and let DB defaults apply.
350 if exclude_unset:
351 if fields_set is None: 351 ↛ 352line 351 didn't jump to line 352 because the condition on line 351 was never true
352 fields_set = data.fields_set()
353 data_dict = data.model_dump(exclude_unset=True)
354 # ``validate_and_normalize_mcp_server_payload`` always assigns ``alias``
355 # on the payload, which marks it as set even when the caller omitted it.
356 # Drop it only when the original request omitted alias; an explicit
357 # ``alias=None`` is a valid request to clear the stored alias.
358 if data_dict.get("alias") is None and "alias" not in fields_set:
359 data_dict.pop("alias", None)
360 # Prisma ``allowed_tools`` is a required String[]; ``null`` is invalid.
361 # The UI sends null to clear a whitelist — treat that as ``[]``.
362 if "allowed_tools" in data_dict and data_dict["allowed_tools"] is None: 362 ↛ 363line 362 didn't jump to line 363 because the condition on line 362 was never true
363 data_dict["allowed_tools"] = []
364 # Json map fields use ``@default("{}")``; explicit null means clear overrides.
365 for json_map_field in (
366 "tool_name_to_display_name",
367 "tool_name_to_description",
368 ):
369 if json_map_field in data_dict and data_dict[json_map_field] is None:
370 data_dict[json_map_field] = {}
371 else:
372 data_dict = data.model_dump(exclude_none=True)
373 # Ensure alias is always present in the dict (even if None)
374 if "alias" not in data_dict: 374 ↛ 378line 374 didn't jump to line 378 because the condition on line 374 was always true
375 data_dict["alias"] = getattr(data, "alias", None)
377 # Handle credentials serialization
378 credentials: Final = data_dict.get("credentials")
379 if credentials is not None:
380 # Lift legacy blob-shaped token-exchange settings into their dedicated
381 # columns (an explicit top-level value wins, including an explicit
382 # null) and strip them from the blob so it never seeds the read-time
383 # fallback for rows written by current code.
384 for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
385 blob_value = credentials.pop(te_field, None)
386 if blob_value is not None and te_field not in data_dict:
387 data_dict[te_field] = blob_value
388 data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=_get_salt_key())
389 data_dict["credentials"] = safe_dumps(data_dict["credentials"])
391 # Serialize JSON fields from ``data_dict`` (not ``data``) so the
392 # exclude_unset filter is respected. Reading back from ``data`` would
393 # reintroduce defaults (e.g. ``env={}``) for fields the caller never set.
394 if data_dict.get("static_headers") is not None:
395 data_dict["static_headers"] = encrypt_secret_map(data_dict["static_headers"])
397 # env_vars is read from ``data_dict`` (not ``data``) like every other JSON
398 # column so the exclude_unset filter is respected: a partial update that
399 # omits env_vars never overwrites the stored value. Global values are
400 # encrypted at rest before serialization.
401 env_vars: Final[Sequence[Mapping[str, str]] | None] = data_dict.get("env_vars")
402 if env_vars is not None:
403 serialized_env_vars: Final = [dict(v) for v in env_vars]
404 _encrypt_global_env_var_values(serialized_env_vars)
405 data_dict["env_vars"] = safe_dumps(serialized_env_vars)
407 if data_dict.get("mcp_info") is not None:
408 data_dict["mcp_info"] = safe_dumps(data_dict["mcp_info"])
410 if data_dict.get("env") is not None:
411 data_dict["env"] = encrypt_secret_map(data_dict["env"])
413 if "tool_name_to_display_name" in data_dict:
414 data_dict["tool_name_to_display_name"] = safe_dumps(data_dict["tool_name_to_display_name"] or {})
415 if "tool_name_to_description" in data_dict:
416 data_dict["tool_name_to_description"] = safe_dumps(data_dict["tool_name_to_description"] or {})
418 # mcp_access_groups is already List[str], no serialization needed
420 # On create, force is_byok so a False value is always written to the DB. On
421 # partial update, only write it when the caller explicitly provided it.
422 if not exclude_unset:
423 data_dict["is_byok"] = getattr(data, "is_byok", False)
425 return data_dict
428def encrypt_credentials(credentials: MCPCredentials, encryption_key: str | None) -> MCPCredentials:
429 auth_value: Final = credentials.get("auth_value")
430 if auth_value is not None: 430 ↛ 431line 430 didn't jump to line 431 because the condition on line 430 was never true
431 credentials["auth_value"] = encrypt_value_helper(
432 value=auth_value,
433 new_encryption_key=encryption_key,
434 )
435 client_id: Final = credentials.get("client_id")
436 if client_id is not None: 436 ↛ 437line 436 didn't jump to line 437 because the condition on line 436 was never true
437 credentials["client_id"] = encrypt_value_helper(
438 value=client_id,
439 new_encryption_key=encryption_key,
440 )
441 client_secret: Final = credentials.get("client_secret")
442 if client_secret is not None: 442 ↛ 443line 442 didn't jump to line 443 because the condition on line 442 was never true
443 credentials["client_secret"] = encrypt_value_helper(
444 value=client_secret,
445 new_encryption_key=encryption_key,
446 )
447 client_private_key: Final = credentials.get("client_private_key")
448 if client_private_key is not None: 448 ↛ 449line 448 didn't jump to line 449 because the condition on line 448 was never true
449 credentials["client_private_key"] = encrypt_value_helper(
450 value=client_private_key,
451 new_encryption_key=encryption_key,
452 )
453 # AWS SigV4 credential fields
454 aws_access_key_id: Final = credentials.get("aws_access_key_id")
455 if aws_access_key_id is not None: 455 ↛ 456line 455 didn't jump to line 456 because the condition on line 455 was never true
456 credentials["aws_access_key_id"] = encrypt_value_helper(
457 value=aws_access_key_id,
458 new_encryption_key=encryption_key,
459 )
460 aws_secret_access_key: Final = credentials.get("aws_secret_access_key")
461 if aws_secret_access_key is not None: 461 ↛ 462line 461 didn't jump to line 462 because the condition on line 461 was never true
462 credentials["aws_secret_access_key"] = encrypt_value_helper(
463 value=aws_secret_access_key,
464 new_encryption_key=encryption_key,
465 )
466 aws_session_token: Final = credentials.get("aws_session_token")
467 if aws_session_token is not None: 467 ↛ 468line 467 didn't jump to line 468 because the condition on line 467 was never true
468 credentials["aws_session_token"] = encrypt_value_helper(
469 value=aws_session_token,
470 new_encryption_key=encryption_key,
471 )
472 # aws_region_name and aws_service_name are NOT secrets — stored as-is
473 return credentials
476def _credentials_blob_to_mutable_dict(blob: str | Mapping[str, object]) -> dict[str, object]:
477 parsed_blob: Final[dict[str, object]] = json.loads(blob) if isinstance(blob, str) else dict(blob)
478 return parsed_blob
481def _mcp_server_table_actions(
482 prisma_client: PrismaClient,
483) -> "TableActions[prisma_db_models.LiteLLM_MCPServerTable]":
484 table: Final[TableActions[prisma_db_models.LiteLLM_MCPServerTable]] = MCPServerRepository(prisma_client).table
485 return table
488def _verification_token_table_actions(
489 prisma_client: PrismaClient,
490) -> "TableActions[prisma_db_models.LiteLLM_VerificationToken]":
491 table: Final[TableActions[prisma_db_models.LiteLLM_VerificationToken]] = VerificationTokenRepository(
492 prisma_client
493 ).table
494 return table
497def _team_table_actions(
498 prisma_client: PrismaClient,
499) -> "TableActions[prisma_db_models.LiteLLM_TeamTable]":
500 table: Final[TableActions[prisma_db_models.LiteLLM_TeamTable]] = TeamRepository(prisma_client).table
501 return table
504def _oauth_client_table_actions(
505 prisma_client: PrismaClient,
506) -> "TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]":
507 table: Final[TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = MCPServerOAuthClientRepository(
508 prisma_client
509 ).table
510 return table
513def _db_transaction_manager(prisma_client: PrismaClient) -> _UserEnvVarsTransaction:
514 manager: Final[_UserEnvVarsTransaction] = prisma_client.db.tx()
515 return manager
518def _identifier_where(value: str, exclude_server_id: str | None) -> "prisma_db_types.LiteLLM_MCPServerTableWhereInput":
519 own_row_guard: Final = (
520 ({"NOT": [{"server_id": exclude_server_id}]},) # mutable-ok: prisma where-inputs must be plain dicts
521 if exclude_server_id is not None
522 else ()
523 )
524 where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = {
525 "AND": [ # mutable-ok: prisma where-inputs must be plain dicts
526 {
527 "OR": [ # mutable-ok: prisma where-inputs must be plain dicts
528 {"server_name": {"equals": value, "mode": "insensitive"}},
529 {"alias": {"equals": value, "mode": "insensitive"}},
530 ]
531 },
532 {
533 "OR": [{"approval_status": None}, {"approval_status": {"not": MCPApprovalStatus.draft}}]
534 }, # mutable-ok: prisma where-inputs must be plain dicts
535 *own_row_guard,
536 ]
537 }
538 return where
541def _identifier_field(data_dict: "Mapping[str, object]", field: str) -> str | None:
542 value: Final = data_dict.get(field)
543 return value if isinstance(value, str) else None
546async def _find_mcp_server_identifier_conflict(
547 table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]",
548 *,
549 server_name: str | None,
550 alias: str | None,
551 exclude_server_id: str | None,
552) -> McpIdentifierConflict | None:
553 """Return the collision between an incoming identifier and a stored row, else None.
555 Each non-empty incoming identifier is compared case-insensitively against
556 BOTH the ``server_name`` and ``alias`` columns, because a value that matches
557 either column would still share the tool prefix another server answers to.
558 ``alias`` is checked first so the reported field is deterministic. Draft
559 rows back the transient OAuth session flow and never reach the registry, so
560 they cannot collide. NULL ``approval_status`` predates the approval
561 workflow and is kept via the inner OR, matching ``get_all_mcp_servers``.
562 """
563 candidates: Final[tuple[tuple[Literal["alias", "server_name"], str | None], ...]] = (
564 ("alias", alias),
565 ("server_name", server_name),
566 )
567 for field_name, value in candidates:
568 if not value: 568 ↛ 570line 568 didn't jump to line 570 because the condition on line 568 was always true
569 continue
570 if (row := await table.find_first(where=_identifier_where(value, exclude_server_id))) is not None:
571 return McpIdentifierConflict(field=field_name, value=value, server_id=row.server_id)
572 return None
575async def find_mcp_server_identifier_conflict(
576 prisma_client: PrismaClient,
577 *,
578 server_name: str | None,
579 alias: str | None,
580 exclude_server_id: str | None,
581) -> McpIdentifierConflict | None:
582 """Unlocked identifier-collision check, for callers outside a write path."""
583 return await _find_mcp_server_identifier_conflict(
584 _mcp_server_table_actions(prisma_client),
585 server_name=server_name,
586 alias=alias,
587 exclude_server_id=exclude_server_id,
588 )
591def _mcp_identifier_lock_keys(*identifiers: str | None) -> tuple[int, ...]:
592 """Deterministic advisory-lock keys for the lowercased identifiers, sorted
593 so concurrent requests for the same pair always lock in the same order."""
594 return tuple(
595 int.from_bytes(
596 hashlib.blake2b(f"mcp_identifier:{normalized}".encode(), digest_size=8).digest(),
597 "big",
598 signed=True,
599 )
600 for normalized in sorted(frozenset(value.lower() for value in identifiers if value))
601 )
604async def _mcp_server_write_if_identifier_free(
605 prisma_client: PrismaClient,
606 *,
607 server_name: str | None,
608 alias: str | None,
609 exclude_server_id: str | None,
610 write: "Callable[[TableActions[prisma_db_models.LiteLLM_MCPServerTable]], Awaitable[prisma_db_models.LiteLLM_MCPServerTable | None]]",
611) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None":
612 """Run ``write`` only when no other live row owns ``server_name``/``alias``.
614 The conflict check and the write share a transaction guarded by per-identifier
615 advisory locks, so two concurrent requests for the same name cannot both
616 pass the check and both insert.
617 """
618 lock_keys: Final = _mcp_identifier_lock_keys(server_name, alias)
619 async with _db_transaction_manager(prisma_client) as tx:
620 for lock_key in lock_keys: 620 ↛ 621line 620 didn't jump to line 621 because the loop on line 620 never started
621 await tx.execute_raw("SELECT pg_advisory_xact_lock($1::bigint)", lock_key)
622 conflict: Final = await _find_mcp_server_identifier_conflict(
623 tx.litellm_mcpservertable,
624 server_name=server_name,
625 alias=alias,
626 exclude_server_id=exclude_server_id,
627 )
628 if conflict is not None: 628 ↛ 629line 628 didn't jump to line 629 because the condition on line 628 was never true
629 return conflict
630 return await write(tx.litellm_mcpservertable)
633async def _db_find_mcp_server_rows(
634 prisma_client: PrismaClient,
635 where: "prisma_db_types.LiteLLM_MCPServerTableWhereInput | None" = None,
636) -> "Sequence[prisma_db_models.LiteLLM_MCPServerTable]":
637 return await _mcp_server_table_actions(prisma_client).find_many(where=where)
640async def _db_find_mcp_server_row(
641 prisma_client: PrismaClient, server_id: str
642) -> "prisma_db_models.LiteLLM_MCPServerTable | None":
643 return await _mcp_server_table_actions(prisma_client).find_unique(where={"server_id": server_id})
646async def _db_update_mcp_server_row(
647 prisma_client: PrismaClient,
648 server_id: str,
649 data: "prisma_db_types.LiteLLM_MCPServerTableUpdateInput",
650) -> "prisma_db_models.LiteLLM_MCPServerTable":
651 row: Final[prisma_db_models.LiteLLM_MCPServerTable | None] = await _mcp_server_table_actions(prisma_client).update(
652 where={"server_id": server_id},
653 data=data,
654 )
655 if row is None: 655 ↛ 656line 655 didn't jump to line 656 because the condition on line 655 was never true
656 raise ValueError(f"MCP server not found, passed server_id={server_id}")
657 return row
660def _user_credential_actions(
661 prisma_client: PrismaClient,
662) -> "TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]":
663 table: Final[TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]] = MCPUserCredentialsRepository(
664 prisma_client
665 ).table
666 return table
669class _MCPUserEnvVarsRepository(PrismaTableRepository["prisma_db_models.LiteLLM_MCPUserEnvVars"]):
670 table_name = "litellm_mcpuserenvvars"
673def _user_env_var_actions(
674 prisma_client: PrismaClient,
675) -> "TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]":
676 return _MCPUserEnvVarsRepository(prisma_client).table
679async def _db_find_user_credential_row(
680 prisma_client: PrismaClient, user_id: str, server_id: str
681) -> "prisma_db_models.LiteLLM_MCPUserCredentials | None":
682 return await _user_credential_actions(prisma_client).find_unique(
683 where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
684 )
687async def _db_find_user_credential_rows(
688 prisma_client: PrismaClient,
689 where: "prisma_db_types.LiteLLM_MCPUserCredentialsWhereInput | None" = None,
690) -> "Sequence[prisma_db_models.LiteLLM_MCPUserCredentials]":
691 return await _user_credential_actions(prisma_client).find_many(where=where)
694async def _db_upsert_user_credential_row(
695 prisma_client: PrismaClient, user_id: str, server_id: str, credential_b64: str
696) -> None:
697 await _user_credential_actions(prisma_client).upsert(
698 where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
699 data={
700 "create": {
701 "user_id": user_id,
702 "server_id": server_id,
703 "credential_b64": credential_b64,
704 },
705 "update": {"credential_b64": credential_b64},
706 },
707 )
710async def _db_find_user_env_var_rows(
711 prisma_client: PrismaClient,
712 where: "prisma_db_types.LiteLLM_MCPUserEnvVarsWhereInput | None" = None,
713) -> "Sequence[prisma_db_models.LiteLLM_MCPUserEnvVars]":
714 return await _user_env_var_actions(prisma_client).find_many(where=where)
717def decrypt_credentials(
718 credentials: MCPCredentials,
719) -> MCPCredentials:
720 """Decrypt all secret fields in an MCPCredentials dict using the global salt key."""
721 secret_fields: Final = [
722 "auth_value",
723 "client_id",
724 "client_secret",
725 "client_private_key",
726 "aws_access_key_id",
727 "aws_secret_access_key",
728 "aws_session_token",
729 ]
730 for field in secret_fields:
731 value = credentials.get(field)
732 if value is not None and isinstance(value, str):
733 credentials[field] = decrypt_value_helper(
734 value=value,
735 key=field,
736 exception_type="debug",
737 return_original_value=True,
738 )
739 return credentials
742def _readable_mcp_servers(
743 rows: Iterable["prisma_db_models.LiteLLM_MCPServerTable"],
744) -> Iterable[LiteLLM_MCPServerTable]:
745 for row in rows:
746 try:
747 table = LiteLLM_MCPServerTable.model_validate(row.model_dump())
748 except SecretMapDecodeError:
749 verbose_proxy_logger.warning("Skipping MCP server %s: cannot decrypt secret map", row.server_id)
750 continue
751 decrypt_global_env_var_values(table.env_vars)
752 yield table
755async def get_all_mcp_servers(
756 prisma_client: PrismaClient,
757 approval_status: str | None = None,
758) -> list[LiteLLM_MCPServerTable]:
759 """
760 Returns mcp servers from the db, optionally filtered by approval_status.
761 Pass approval_status=None to return every server except drafts, which back the admin OAuth
762 session flow, are addressable only by their own server_id, and must never appear in a listing.
763 NULL approval_status predates the approval workflow, so those rows are kept explicitly rather
764 than dropped by a bare inequality, which SQL evaluates as NULL and would silently hide them.
765 """
766 where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = (
767 {"approval_status": approval_status}
768 if approval_status is not None
769 else {"OR": [{"approval_status": None}, {"approval_status": {"not": MCPApprovalStatus.draft}}]}
770 )
771 mcp_servers: Final = await _db_find_mcp_server_rows(prisma_client, where)
773 return list(_readable_mcp_servers(mcp_servers))
776async def get_mcp_server(prisma_client: PrismaClient, server_id: str) -> LiteLLM_MCPServerTable | None:
777 """
778 Returns the matching mcp server from the db iff exists
779 """
780 mcp_server: Final = await _db_find_mcp_server_row(prisma_client, server_id)
781 if mcp_server is None:
782 return None
783 table: Final = LiteLLM_MCPServerTable.model_validate(mcp_server.model_dump())
784 decrypt_global_env_var_values(table.env_vars)
785 return table
788async def get_mcp_servers(prisma_client: PrismaClient, server_ids: Iterable[str]) -> list[LiteLLM_MCPServerTable]:
789 """
790 Returns the matching mcp servers from the db with the server_ids
791 """
792 _mcp_servers: Final[Sequence[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions(
793 prisma_client
794 ).find_many(
795 where={
796 "server_id": {"in": server_ids},
797 }
798 )
799 return list(_readable_mcp_servers(_mcp_servers))
802async def get_mcp_servers_by_verificationtoken(prisma_client: PrismaClient, token: str) -> list[str]:
803 """
804 Returns the mcp servers from the db for the verification token
805 """
806 verification_token_record: (
807 prisma_db_models.LiteLLM_VerificationToken | None
808 ) = await _verification_token_table_actions(prisma_client).find_unique(
809 where={
810 "token": token,
811 },
812 include={
813 "object_permission": True,
814 },
815 )
817 mcp_servers: list[str] | None = []
818 if verification_token_record is not None and verification_token_record.object_permission is not None:
819 mcp_servers = verification_token_record.object_permission.mcp_servers
820 return mcp_servers or []
823async def get_mcp_servers_by_team(prisma_client: PrismaClient, team_id: str) -> list[str]:
824 """
825 Returns the mcp servers from the db for the team id
826 """
827 team_record: prisma_db_models.LiteLLM_TeamTable | None = await _team_table_actions(prisma_client).find_unique(
828 where={
829 "team_id": team_id,
830 },
831 include={
832 "object_permission": True,
833 },
834 )
836 mcp_servers: list[str] | None = []
837 if team_record is not None and team_record.object_permission is not None:
838 mcp_servers = team_record.object_permission.mcp_servers
839 return mcp_servers or []
842async def get_all_mcp_servers_for_user(
843 prisma_client: PrismaClient,
844 user: UserAPIKeyAuth,
845) -> list[LiteLLM_MCPServerTable]:
846 """
847 Get all the mcp servers filtered by the given user has access to.
849 Following Least-Privilege Principle - the requestor should only be able to see the mcp servers that they have access to.
850 """
852 mcp_server_ids: Final[set[str]] = set()
853 mcp_servers = []
855 # Get the mcp servers for the key
856 if user.api_key:
857 token_mcp_servers: Final = await get_mcp_servers_by_verificationtoken(prisma_client, user.api_key)
858 mcp_server_ids.update(token_mcp_servers)
860 # check for special team membership
861 if SpecialMCPServerName.all_team_servers in mcp_server_ids and user.team_id is not None:
862 team_mcp_servers: Final = await get_mcp_servers_by_team(prisma_client, user.team_id)
863 mcp_server_ids.update(team_mcp_servers)
865 if len(mcp_server_ids) > 0:
866 mcp_servers = await get_mcp_servers(prisma_client, mcp_server_ids)
868 return mcp_servers
871async def get_objectpermissions_for_mcp_server(
872 prisma_client: PrismaClient, mcp_server_id: str
873) -> "Sequence[prisma_db_models.LiteLLM_ObjectPermissionTable]":
874 """
875 Get all the object permissions records and the associated team and verficiationtoken records that have access to the mcp server
876 """
877 object_permission_records: Final[
878 Sequence[prisma_db_models.LiteLLM_ObjectPermissionTable]
879 ] = await ObjectPermissionRepository(prisma_client).table.find_many(
880 where={
881 "mcp_servers": {"has": mcp_server_id},
882 },
883 include={
884 "teams": True,
885 "verification_tokens": True,
886 },
887 )
889 return object_permission_records
892async def get_virtualkeys_for_mcp_server(
893 prisma_client: PrismaClient, server_id: str
894) -> "Sequence[prisma_db_models.LiteLLM_VerificationToken]":
895 """
896 Get all the virtual keys that have access to the mcp server
897 """
898 virtual_keys: Final[
899 Sequence[prisma_db_models.LiteLLM_VerificationToken] | None
900 ] = await VerificationTokenRepository(prisma_client).table.find_many(
901 where={
902 "mcp_servers": {"has": server_id},
903 },
904 )
906 if virtual_keys is None: # pyright: ignore[reportUnnecessaryComparison] # unreachable per seam types; kept as-is
907 return []
908 return virtual_keys
911async def delete_mcp_server_from_team(prisma_client: PrismaClient, server_id: str):
912 """
913 Remove the mcp server from the team
914 """
917async def delete_mcp_server_from_virtualkey():
918 """
919 Remove the mcp server from the virtual key
920 """
923async def delete_mcp_server(
924 prisma_client: PrismaClient,
925 server_id: str,
926 invalidate_token_cache: Callable[[str, str], Awaitable[None]] | None = None,
927) -> LiteLLM_MCPServerTable | None:
928 """
929 Delete the mcp server from the db by server_id
931 The server-row delete is the commit point. Per-user credential and env var
932 rows have no FK cascade, so they are cleaned up afterwards on a best-effort
933 basis: a transient failure there leaves only orphaned rows pointing at a
934 now-missing server and must not turn a successful delete into a
935 caller-visible error. Each table is cleaned independently so a failure on one
936 still attempts the other.
938 Each enumerated credential row's user also gets their cached per-user token
939 invalidated (legacy cache + v2 store, via invalidate_token_cache, defaulting
940 to the manager's shared invalidation): the caches are keyed by
941 (user_id, server_id), so without this a re-created server reusing the same
942 server_id would serve tokens minted for the deleted server until TTL.
944 Returns the deleted mcp server record if it exists, otherwise None
945 """
946 deleted_server: Final = await MCPServerRepository(prisma_client).table.delete(
947 where={
948 "server_id": server_id,
949 },
950 )
951 if deleted_server is not None:
952 credential_user_ids: list[str] = []
953 try:
954 credential_rows: Sequence[prisma_db_models.LiteLLM_MCPUserCredentials] = await _user_credential_actions(
955 prisma_client
956 ).find_many(where={"server_id": server_id})
957 credential_user_ids = [row.user_id for row in credential_rows]
958 except Exception as e: # noqa: BLE001 - enumeration is best-effort; cached tokens expire by TTL
959 verbose_proxy_logger.warning(
960 "MCP server %s deleted but per-user credential enumeration failed; cached tokens expire by TTL: %s",
961 server_id,
962 e,
963 )
964 for model, label in (
965 (_user_credential_actions(prisma_client), "credential"),
966 (_user_env_var_actions(prisma_client), "env var"),
967 (_oauth_client_table_actions(prisma_client), "OAuth client"),
968 ):
969 try:
970 await model.delete_many(where={"server_id": server_id})
971 except Exception as e:
972 verbose_proxy_logger.warning(
973 "MCP server %s deleted but per-user %s cleanup failed; "
974 "orphaned rows can be removed on a later delete: %s",
975 server_id,
976 label,
977 e,
978 )
979 if credential_user_ids: 979 ↛ 980line 979 didn't jump to line 980 because the condition on line 979 was never true
980 if invalidate_token_cache is None:
981 from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
982 global_mcp_server_manager,
983 )
985 invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
986 for user_id in credential_user_ids:
987 await invalidate_token_cache(user_id, server_id)
988 return deleted_server # pyright: ignore[reportReturnType] # prisma row, not domain LiteLLM_MCPServerTable
991async def create_mcp_server(
992 prisma_client: PrismaClient, data: NewMCPServerRequest, touched_by: str
993) -> LiteLLM_MCPServerTable:
994 """
995 Create a new mcp server record in the db
996 """
997 if data.server_id is None: 997 ↛ 998line 997 didn't jump to line 998 because the condition on line 997 was never true
998 data.server_id = str(uuid.uuid4())
1000 # Use helper to prepare data with proper JSON serialization
1001 data_dict: Final = _prepare_mcp_server_data(data)
1003 # Add audit fields
1004 data_dict["created_by"] = touched_by
1005 data_dict["updated_by"] = touched_by
1007 new_mcp_server: Final = await MCPServerRepository(prisma_client).table.create(data=data_dict)
1009 _decrypt_env_vars_on_returned_row(new_mcp_server)
1010 return LiteLLM_MCPServerTable.model_validate(new_mcp_server.model_dump())
1013async def create_mcp_server_if_identifier_free(
1014 prisma_client: PrismaClient, data: NewMCPServerRequest, touched_by: str
1015) -> LiteLLM_MCPServerTable | McpIdentifierConflict:
1016 """Create the row only when no other live server owns ``server_name``/``alias``.
1018 Returns the McpIdentifierConflict instead of inserting when the collision
1019 check finds an existing row; the advisory-lock transaction keeps two
1020 concurrent creates of the same identifier from both passing.
1021 """
1022 if data.server_id is None:
1023 data.server_id = str(uuid.uuid4())
1025 data_dict: Final = _prepare_mcp_server_data(data)
1026 data_dict["created_by"] = touched_by
1027 data_dict["updated_by"] = touched_by
1029 async def _create(
1030 table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]",
1031 ) -> "prisma_db_models.LiteLLM_MCPServerTable | None":
1032 return await table.create(data=data_dict)
1034 written: Final = await _mcp_server_write_if_identifier_free(
1035 prisma_client,
1036 server_name=_identifier_field(data_dict, "server_name"),
1037 alias=_identifier_field(data_dict, "alias"),
1038 exclude_server_id=None,
1039 write=_create,
1040 )
1041 if isinstance(written, McpIdentifierConflict): 1041 ↛ 1042line 1041 didn't jump to line 1042 because the condition on line 1041 was never true
1042 return written
1043 if written is None: 1043 ↛ 1044line 1043 didn't jump to line 1044 because the condition on line 1043 was never true
1044 raise RuntimeError("inserted MCP server row missing")
1046 _decrypt_env_vars_on_returned_row(written)
1047 return LiteLLM_MCPServerTable.model_validate(written.model_dump())
1050async def create_draft_mcp_server(
1051 prisma_client: PrismaClient,
1052 data: NewMCPServerRequest,
1053 touched_by: str,
1054 ttl_seconds: int,
1055 server_id: str | None = None,
1056) -> LiteLLM_MCPServerTable:
1057 """
1058 Persist a short-lived draft row backing the admin OAuth "Authorize & Fetch Token" flow.
1060 The draft lives in the database rather than in process memory so that the /register,
1061 /authorize and /token legs resolve it whichever worker or replica accepts each request.
1063 Writing is strictly create-if-absent. Any existing row for the id is returned untouched, which
1064 covers both a live draft for this same session and a real server the edit form is
1065 re-authorizing against its own id, where writing a draft would collide on the primary key.
1066 Each click of Authorize mints a fresh id, so nothing is lost by never overwriting, and it is
1067 what makes concurrent callers sharing one id safe rather than mutually destructive.
1068 """
1069 draft_id: Final = server_id or data.server_id or str(uuid.uuid4())
1070 await _prune_expired_draft_mcp_servers(prisma_client, ttl_seconds)
1072 existing: Final = await _db_find_mcp_server_row(prisma_client, draft_id)
1073 if existing is not None:
1074 # Already usable by every worker, whether it is a live draft for this same session or a
1075 # real server the edit form is re-authorizing. Either way there is nothing to write, and
1076 # not writing is what keeps concurrent callers for one server_id from racing each other.
1077 return LiteLLM_MCPServerTable.model_validate(existing.model_dump())
1079 draft_payload: Final = data.model_copy(update={"server_id": draft_id, "approval_status": MCPApprovalStatus.draft})
1080 try:
1081 return await create_mcp_server(prisma_client, draft_payload, touched_by)
1082 except Exception:
1083 # Lost the create race: the read above and this create are two statements, not one. The
1084 # winner wrote a draft for this same session, so adopt it rather than failing a caller
1085 # whose session is in fact ready. Anything else still raises.
1086 raced: Final = await _db_find_mcp_server_row(prisma_client, draft_id)
1087 if raced is None or raced.approval_status != MCPApprovalStatus.draft:
1088 raise
1089 return LiteLLM_MCPServerTable.model_validate(raced.model_dump())
1092async def _prune_expired_draft_mcp_servers(prisma_client: PrismaClient, ttl_seconds: int) -> None:
1093 """Drop drafts already past ``ttl_seconds``, so abandoned OAuth sessions do not accumulate.
1095 Runs on each draft write rather than on a schedule, mirroring the in-memory cache this
1096 replaces, which pruned on every store. Expired drafts are unreadable by then anyway, so the
1097 only thing at stake is row count, and the work is bounded by how often admins authorize.
1098 """
1099 cutoff: Final = datetime.now(timezone.utc) - timedelta(seconds=max(1, ttl_seconds))
1100 # Age is filtered here rather than in the query: the draft set is bounded by how many OAuth
1101 # authorizations are in flight, so it is a handful of rows even on a busy proxy.
1102 drafts: Final = await _db_find_mcp_server_rows(
1103 prisma_client,
1104 where={"approval_status": MCPApprovalStatus.draft},
1105 )
1106 for row in drafts:
1107 # A row without a timestamp has no age to judge, so leave it rather than guess it is stale.
1108 # Two workers sweeping the same row is harmless: prisma's delete returns None for a row
1109 # that is already gone rather than raising, so the loser of that race is a no-op.
1110 if row.updated_at is not None and row.updated_at < cutoff:
1111 await delete_mcp_server(prisma_client, row.server_id)
1114async def get_draft_mcp_server(
1115 prisma_client: PrismaClient, server_id: str, ttl_seconds: int
1116) -> LiteLLM_MCPServerTable | None:
1117 """
1118 Return the draft row for ``server_id`` if it has not yet aged past ``ttl_seconds``, else None.
1120 Age is enforced in the query rather than by a sweeper so an expired draft is unreadable the
1121 moment it lapses, regardless of which process last ran a cleanup.
1122 """
1123 cutoff: Final = datetime.now(timezone.utc) - timedelta(seconds=max(1, ttl_seconds))
1124 draft_rows: Final = await _db_find_mcp_server_rows(
1125 prisma_client,
1126 where={
1127 "server_id": server_id,
1128 "approval_status": MCPApprovalStatus.draft,
1129 "updated_at": {"gte": cutoff},
1130 },
1131 )
1132 if not draft_rows:
1133 return None
1135 table: Final = LiteLLM_MCPServerTable.model_validate(draft_rows[0].model_dump())
1136 decrypt_global_env_var_values(table.env_vars)
1137 return table
1140async def _update_mcp_server_row(
1141 prisma_client: PrismaClient,
1142 *,
1143 server_id: str,
1144 data_dict: Mapping[str, object],
1145) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None":
1146 identifier_write: Final = any(field in data_dict for field in ("server_name", "alias"))
1148 async def _update(
1149 table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]",
1150 ) -> "prisma_db_models.LiteLLM_MCPServerTable | None":
1151 return await table.update(
1152 where={"server_id": server_id}, # mutable-ok: prisma where-inputs must be plain dicts
1153 data=data_dict,
1154 )
1156 if not identifier_write:
1157 return await _update(_mcp_server_table_actions(prisma_client))
1158 if "alias" in data_dict and not data_dict["alias"] and "server_name" not in data_dict:
1159 # Clearing the alias drops the prefix to the stored server_name, which
1160 # may already belong to another row, so that name needs the check too.
1161 existing: Final = await _db_find_mcp_server_row(prisma_client, server_id)
1162 if existing is None:
1163 return await _update(_mcp_server_table_actions(prisma_client))
1164 return await _mcp_server_write_if_identifier_free(
1165 prisma_client,
1166 server_name=existing.server_name,
1167 alias=None,
1168 exclude_server_id=server_id,
1169 write=_update,
1170 )
1171 return await _mcp_server_write_if_identifier_free(
1172 prisma_client,
1173 server_name=_identifier_field(data_dict, "server_name"),
1174 alias=_identifier_field(data_dict, "alias"),
1175 exclude_server_id=server_id,
1176 write=_update,
1177 )
1180async def update_mcp_server(
1181 prisma_client: PrismaClient,
1182 data: UpdateMCPServerRequest,
1183 touched_by: str,
1184 fields_set: set[str] | None = None,
1185) -> LiteLLM_MCPServerTable | McpIdentifierConflict | None:
1186 """
1187 Update a new mcp server record in the db
1189 Returns McpIdentifierConflict instead of writing when the update would put
1190 ``server_name``/``alias`` onto identifiers another live row already owns.
1191 """
1192 from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
1194 # Use helper to prepare data with proper JSON serialization.
1195 # exclude_unset=True makes this a true partial update: fields the caller did
1196 # not provide are not written, so they keep their existing DB value instead
1197 # of being reset to a schema default (transport=sse, allow_all_keys=False...).
1198 data_dict: Final = _prepare_mcp_server_data(data, exclude_unset=True, fields_set=fields_set)
1200 # Pre-fetch existing record once if we need it for auth_type, url, or credential logic
1201 existing = None
1202 has_credentials: Final = "credentials" in data_dict and data_dict["credentials"] is not None
1203 # An explicit token-exchange column write (set or clear) also migrates the
1204 # legacy blob copies below, so the existing row is needed for those updates.
1205 explicit_te_write: Final = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & data_dict.keys())
1206 url_provided: Final = "url" in data_dict and data_dict["url"] is not None
1207 issuer_provided: Final = "issuer" in data_dict
1208 if data.auth_type or has_credentials or explicit_te_write or url_provided or issuer_provided:
1209 existing = await _db_find_mcp_server_row(prisma_client, data.server_id)
1211 auth_type_changed: Final = bool(
1212 data.auth_type
1213 and existing
1214 and _credential_auth_class(existing.auth_type) != _credential_auth_class(data.auth_type)
1215 )
1216 # A url change re-points the server at a potentially different upstream, so any discovered or
1217 # trust-on-first-use OAuth endpoints/issuer belong to the old upstream and must re-discover.
1218 url_changed: Final = bool(url_provided and existing and existing.url != data_dict["url"])
1219 old_issuer: Final = _blank_to_none(getattr(existing, "issuer", None)) if existing else None
1220 issuer_changed: Final = bool(
1221 issuer_provided and old_issuer is not None and _blank_to_none(data_dict.get("issuer")) != old_issuer
1222 )
1224 # Clear stale credentials when auth_type changes but no new credentials provided
1225 if auth_type_changed and "credentials" not in data_dict: 1225 ↛ 1226line 1225 didn't jump to line 1226 because the condition on line 1225 was never true
1226 data_dict["credentials"] = None
1228 if auth_type_changed or url_changed or issuer_changed:
1229 # Clear each auth-flow-scoped field that the caller either omitted (partial update) or
1230 # resubmitted unchanged. The edit form re-sends every field, so a stale issuer/endpoint
1231 # belonging to the old upstream would otherwise survive a url/auth_type change and win in the
1232 # resolution merge; only a genuinely new submitted value is kept.
1233 data_dict.update(
1234 {
1235 field: None
1236 for field in _AUTH_FLOW_SCOPED_FIELDS
1237 if field not in data_dict or data_dict[field] == getattr(existing, field, None)
1238 }
1239 )
1241 # An explicit column write that does not touch credentials must still migrate
1242 # the row's legacy blob copies: lift values for columns the caller left
1243 # untouched, strip every copy from the blob. Without this, clearing a column
1244 # (e.g. to re-enable RFC 9728/8414 discovery) would leave the blob copy in
1245 # place, and the next credentials update's migrate-on-write would silently
1246 # repopulate the column the admin just cleared. (When credentials ARE in the
1247 # update, the merge below performs the same migration.)
1248 if explicit_te_write and "credentials" not in data_dict and existing is not None and existing.credentials: 1248 ↛ 1249line 1248 didn't jump to line 1249 because the condition on line 1248 was never true
1249 existing_creds = _credentials_blob_to_mutable_dict(existing.credentials)
1250 if _TOKEN_EXCHANGE_COLUMN_FIELDS & existing_creds.keys():
1251 for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
1252 legacy_value = existing_creds.pop(te_field, None)
1253 if legacy_value is not None and te_field not in data_dict and getattr(existing, te_field, None) is None:
1254 data_dict[te_field] = legacy_value
1255 data_dict["credentials"] = safe_dumps(existing_creds)
1257 # Merge credentials: preserve existing fields not present in the update.
1258 # Without this, a partial credential update (e.g. changing only region)
1259 # would wipe encrypted secrets that the UI cannot display back.
1260 if "credentials" in data_dict and data_dict["credentials"] is not None: 1260 ↛ 1261line 1260 didn't jump to line 1261 because the condition on line 1260 was never true
1261 if existing and existing.credentials:
1262 # Only merge when the credential CLASS is unchanged. A cross-class switch
1263 # (e.g. oauth2 → api_key, or oauth2 → true_passthrough) replaces credentials
1264 # entirely to avoid stale secrets from the previous class lingering; a switch
1265 # within the client-forwarded class (true_passthrough ↔ oauth_delegate) keeps
1266 # the same declared app and so must merge, not replace.
1267 if not auth_type_changed:
1268 existing_creds = _credentials_blob_to_mutable_dict(existing.credentials)
1269 new_creds: Final = _credentials_blob_to_mutable_dict(data_dict["credentials"])
1270 # New values override existing; existing keys not in update are preserved. A client
1271 # rotation additionally drops the previous app's stale minted token keys.
1272 merged: Final = _drop_stale_minted_on_client_rotation({**existing_creds, **new_creds}, new_creds)
1273 # Migrate-on-write for legacy rows: token-exchange settings the
1274 # old blob shape carried move to their dedicated columns (unless
1275 # the caller set the column this update, or the row already has
1276 # one) and are never re-persisted in the blob. Stored plaintext,
1277 # so the merged value lifts as-is.
1278 for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
1279 legacy_value = merged.pop(te_field, None)
1280 if (
1281 legacy_value is not None
1282 and te_field not in data_dict
1283 and getattr(existing, te_field, None) is None
1284 ):
1285 data_dict[te_field] = legacy_value
1286 data_dict["credentials"] = safe_dumps(merged)
1288 # Add audit fields
1289 data_dict["updated_by"] = touched_by
1291 # prisma-python rejects a raw ``None`` for a ``Json?`` field ("value is required but not set"); the
1292 # clear paths above use ``None`` as the merge-skip sentinel, so translate it here to ``Json(None)``,
1293 # which writes SQL null and reads back as ``None``. Done at the edge so the merge guards stay simple.
1294 if "credentials" in data_dict and data_dict["credentials"] is None:
1295 from prisma import Json # noqa: PLC0415 # local import: prisma may be ungenerated at module load in some tools
1297 data_dict["credentials"] = Json(None)
1299 updated_mcp_server: Final = await _update_mcp_server_row(
1300 prisma_client,
1301 server_id=data.server_id,
1302 data_dict=data_dict,
1303 )
1305 if isinstance(updated_mcp_server, McpIdentifierConflict): 1305 ↛ 1306line 1305 didn't jump to line 1306 because the condition on line 1305 was never true
1306 return updated_mcp_server
1307 _decrypt_env_vars_on_returned_row(updated_mcp_server)
1308 return LiteLLM_MCPServerTable.model_validate(updated_mcp_server.model_dump()) if updated_mcp_server else None
1311async def get_mcp_server_oauth_client_credentials(prisma_client: PrismaClient, server_id: str) -> object | None:
1312 """Read the persisted (encrypted) DCR OAuth client blob for a server from the
1313 server-scoped store, or None. Config.yaml-declared servers have no
1314 LiteLLM_MCPServerTable row, so their dynamically registered client lives here keyed
1315 by server_id. The returned value is the raw credentials blob for
1316 ``_get_persisted_dcr_credentials`` to parse."""
1317 row: Final[prisma_db_models.LiteLLM_MCPServerOAuthClient | None] = await _oauth_client_table_actions(
1318 prisma_client
1319 ).find_unique(where={"server_id": server_id})
1320 if row is None:
1321 return None
1322 return row.credentials
1325async def upsert_mcp_server_oauth_client_credentials(
1326 prisma_client: PrismaClient, server_id: str, credentials: MCPCredentials
1327) -> None:
1328 """Persist a server's dynamically registered OAuth client (RFC 7591 DCR) in the
1329 server-scoped store keyed by server_id, independent of any LiteLLM_MCPServerTable row.
1330 client_id/client_secret are encrypted at rest with the same salt key used for the
1331 server row's credentials blob, so ``_apply_persisted_dcr_credentials`` decrypts them the
1332 same way regardless of which store a server's client came from."""
1333 from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
1335 encrypted: Final = encrypt_credentials(credentials=MCPCredentials(**credentials), encryption_key=_get_salt_key())
1336 blob: Final = safe_dumps(encrypted)
1337 await _oauth_client_table_actions(prisma_client).upsert(
1338 where={"server_id": server_id},
1339 data={
1340 "create": {"server_id": server_id, "credentials": blob},
1341 "update": {"credentials": blob},
1342 },
1343 )
1346def _reencrypt_mcp_credentials_blob(
1347 credentials: "str | Mapping[str, object] | None", new_master_key: str
1348) -> str | None:
1349 """Decrypt an at-rest MCP credentials blob with the current key and re-encrypt it under
1350 new_master_key, returning the serialized blob or None when there is nothing to rotate. Shared by
1351 every table that stores an encrypted MCP credentials blob so a master-key rotation covers them
1352 uniformly and cannot silently skip one."""
1353 if not credentials:
1354 return None
1355 from litellm.litellm_core_utils.safe_json_dumps import safe_dumps # noqa: PLC0415 # avoids circular import
1357 creds_dict: Final = _credentials_blob_to_mutable_dict(credentials)
1358 decrypted: Final = decrypt_credentials(credentials=cast(MCPCredentials, creds_dict))
1359 encrypted: Final = encrypt_credentials(credentials=decrypted, encryption_key=new_master_key)
1360 return safe_dumps(encrypted)
1363async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient, touched_by: str, new_master_key: str):
1364 from litellm.litellm_core_utils.safe_json_dumps import safe_dumps # noqa: PLC0415 # avoids circular import
1366 mcp_servers: Final = await _db_find_mcp_server_rows(prisma_client)
1368 updated = 0
1369 for mcp_server in mcp_servers:
1370 update_data: dict[str, str] = {}
1372 rotated_credentials = _reencrypt_mcp_credentials_blob(mcp_server.credentials, new_master_key)
1373 if rotated_credentials is not None:
1374 update_data["credentials"] = rotated_credentials
1376 rotated_env_vars = _reencrypt_global_env_var_values(mcp_server.env_vars, new_master_key)
1377 if rotated_env_vars is not None:
1378 update_data["env_vars"] = safe_dumps(rotated_env_vars)
1380 for field in ("static_headers", "env"):
1381 try:
1382 if secret_map := decode_secret_map(getattr(mcp_server, field, None), key=field):
1383 update_data[field] = encrypt_secret_map(secret_map, new_encryption_key=new_master_key)
1384 except SecretMapDecodeError:
1385 verbose_proxy_logger.warning("Cannot rotate MCP %s for server %s", field, mcp_server.server_id)
1387 if not update_data:
1388 continue
1390 update_data["updated_by"] = touched_by
1391 await _mcp_server_table_actions(prisma_client).update(
1392 where={"server_id": mcp_server.server_id},
1393 data=update_data,
1394 )
1395 updated += 1
1397 oauth_clients: Final[Sequence[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = await _oauth_client_table_actions(
1398 prisma_client
1399 ).find_many()
1400 oauth_updated = 0
1401 for oauth_client in oauth_clients:
1402 rotated_credentials = _reencrypt_mcp_credentials_blob(oauth_client.credentials, new_master_key)
1403 if rotated_credentials is None:
1404 continue
1405 await _oauth_client_table_actions(prisma_client).update(
1406 where={"server_id": oauth_client.server_id},
1407 data={"credentials": rotated_credentials},
1408 )
1409 oauth_updated += 1
1411 verbose_proxy_logger.info(
1412 "rotate_mcp_server_credentials_master_key: rotated %d MCP server row(s) and %d OAuth-client row(s)",
1413 updated,
1414 oauth_updated,
1415 )
1418def _decode_user_credential(stored: str) -> str | None:
1419 """Read back a value persisted in ``LiteLLM_MCPUserCredentials.credential_b64``.
1421 Tries nacl decryption first (current write format). Falls back to a
1422 plain ``urlsafe_b64decode`` for rows persisted by older code that wrote
1423 the credential without encryption. Returns ``None`` when neither path
1424 yields a valid string.
1425 """
1426 decrypted: Final = decrypt_value_helper(
1427 value=stored,
1428 key="mcp_user_credential",
1429 exception_type="debug",
1430 return_original_value=False,
1431 )
1432 if decrypted is not None: 1432 ↛ 1434line 1432 didn't jump to line 1434 because the condition on line 1432 was always true
1433 return decrypted
1434 try:
1435 return base64.urlsafe_b64decode(stored).decode()
1436 except (binascii.Error, UnicodeDecodeError, ValueError, TypeError):
1437 return None
1440def _warn_undecryptable_credential(user_id: str, server_id: str) -> None:
1441 """Log the one credential state that otherwise reads as "user never authorized"."""
1442 verbose_proxy_logger.warning(
1443 "MCP user credential for user=%s server=%s could not be decrypted (likely written under a "
1444 "previous LITELLM_SALT_KEY); the user is treated as not connected and must re-authorize.",
1445 user_id,
1446 server_id,
1447 )
1450def _parse_oauth_payload(decoded: str | None) -> OAuthCredentialPayload | None:
1451 """Return the OAuth2 payload dict if ``decoded`` holds one, else ``None``.
1453 A row is considered an OAuth2 credential iff its decoded value parses as
1454 a JSON object with ``"type": "oauth2"``. Plain BYOK credentials (which
1455 share the same column) decode to a non-JSON string and return ``None``.
1457 Callers that need to tell an unreadable row from a readable non-OAuth2 one
1458 pass the result of :func:`_decode_user_credential` so a single decode
1459 answers both questions: ``None`` there means the value can be neither
1460 decrypted nor base64-decoded, so no caller can ever recover it.
1461 """
1462 if decoded is None: 1462 ↛ 1463line 1462 didn't jump to line 1463 because the condition on line 1462 was never true
1463 return None
1464 parsed: OAuthCredentialPayload | None
1465 try:
1466 parsed = json.loads(decoded)
1467 except (ValueError, TypeError):
1468 return None
1469 if isinstance(parsed, dict) and parsed.get("type") == "oauth2": 1469 ↛ 1471line 1469 didn't jump to line 1471 because the condition on line 1469 was always true
1470 return parsed
1471 return None
1474def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None:
1475 """Return the OAuth2 payload dict held in ``stored``, else ``None``."""
1476 return _parse_oauth_payload(_decode_user_credential(stored))
1479async def rotate_mcp_user_credentials_master_key(prisma_client: PrismaClient, new_master_key: str):
1480 """Re-encrypt every ``LiteLLM_MCPUserCredentials`` row with ``new_master_key``.
1482 Reads each ``credential_b64`` with the current salt key (falling back to
1483 legacy plain base64 for unmigrated rows) and writes it back encrypted
1484 under the new master key. Rows that are unreadable under both paths
1485 are logged and skipped so one corrupt row does not abort the rotation.
1486 """
1487 rows: Final = await _db_find_user_credential_rows(prisma_client)
1488 rotated = 0
1489 skipped = 0
1490 for row in rows:
1491 plaintext = _decode_user_credential(row.credential_b64)
1492 if plaintext is None:
1493 verbose_proxy_logger.warning(
1494 "rotate_mcp_user_credentials_master_key: could not decode "
1495 "credential for user_id=%s server_id=%s, skipping",
1496 row.user_id,
1497 row.server_id,
1498 )
1499 skipped += 1
1500 continue
1501 re_encrypted = encrypt_value_helper(plaintext, new_encryption_key=new_master_key)
1502 await _user_credential_actions(prisma_client).update(
1503 where={
1504 "user_id_server_id": {
1505 "user_id": row.user_id,
1506 "server_id": row.server_id,
1507 }
1508 },
1509 data={"credential_b64": re_encrypted},
1510 )
1511 rotated += 1
1512 verbose_proxy_logger.info(
1513 "rotate_mcp_user_credentials_master_key: rotated %d row(s), skipped %d",
1514 rotated,
1515 skipped,
1516 )
1519async def rotate_mcp_user_env_vars_master_key(prisma_client: PrismaClient, new_master_key: str):
1520 """Re-encrypt every ``LiteLLM_MCPUserEnvVars`` row with ``new_master_key``.
1522 Reads each ``values_b64`` blob with the current salt key and writes it back
1523 encrypted under the new master key. Rows that fail to decrypt are logged and
1524 skipped so one corrupt row does not abort the rotation nor overwrite values
1525 that may still be recoverable.
1526 """
1527 rows: Final = await _db_find_user_env_var_rows(prisma_client)
1528 rotated = 0
1529 skipped = 0
1530 for row in rows:
1531 plaintext = decrypt_value_helper(
1532 value=row.values_b64,
1533 key="mcp_user_env_vars",
1534 exception_type="debug",
1535 return_original_value=False,
1536 )
1537 if plaintext is None:
1538 verbose_proxy_logger.warning(
1539 "rotate_mcp_user_env_vars_master_key: could not decrypt env vars for user_id=%s server_id=%s, skipping",
1540 row.user_id,
1541 row.server_id,
1542 )
1543 skipped += 1
1544 continue
1545 re_encrypted = encrypt_value_helper(plaintext, new_encryption_key=new_master_key)
1546 await _user_env_var_actions(prisma_client).update(
1547 where={
1548 "user_id_server_id": {
1549 "user_id": row.user_id,
1550 "server_id": row.server_id,
1551 }
1552 },
1553 data={"values_b64": re_encrypted},
1554 )
1555 rotated += 1
1556 verbose_proxy_logger.info(
1557 "rotate_mcp_user_env_vars_master_key: rotated %d row(s), skipped %d",
1558 rotated,
1559 skipped,
1560 )
1563async def store_user_credential(
1564 prisma_client: PrismaClient,
1565 user_id: str,
1566 server_id: str,
1567 credential: str,
1568) -> None:
1569 """Store a user credential for a BYOK MCP server."""
1571 encoded: Final = encrypt_value_helper(credential)
1572 await _db_upsert_user_credential_row(prisma_client, user_id, server_id, encoded)
1575async def get_user_credential(
1576 prisma_client: PrismaClient,
1577 user_id: str,
1578 server_id: str,
1579) -> str | None:
1580 """Return credential for a user+server pair, or None."""
1582 row: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
1583 if row is None: 1583 ↛ 1585line 1583 didn't jump to line 1585 because the condition on line 1583 was always true
1584 return None
1585 return _decode_user_credential(row.credential_b64)
1588async def has_user_credential(
1589 prisma_client: PrismaClient,
1590 user_id: str,
1591 server_id: str,
1592) -> bool:
1593 """Return True if the user has a stored credential for this server."""
1594 row: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
1595 return row is not None
1598async def delete_user_credential(
1599 prisma_client: PrismaClient,
1600 user_id: str,
1601 server_id: str,
1602) -> None:
1603 """Delete the user's stored credential for a BYOK MCP server."""
1604 await _user_credential_actions(prisma_client).delete(
1605 where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
1606 )
1609# ── OAuth2 user-credential helpers ────────────────────────────────────────────
1612async def store_user_oauth_credential(
1613 prisma_client: PrismaClient,
1614 user_id: str,
1615 server_id: str,
1616 access_token: str,
1617 refresh_token: str | None = None,
1618 expires_in: int | None = None,
1619 scopes: list[str] | None = None,
1620 skip_byok_guard: bool = False,
1621 identity_binding_proof: str | None = None,
1622) -> None:
1623 """Persist an OAuth2 access token for a user+server pair.
1625 The payload is JSON-serialised and stored encrypted in the same
1626 ``credential_b64`` column used by BYOK. A ``"type": "oauth2"`` key
1627 differentiates it from plain BYOK API keys.
1628 """
1630 expires_at: str | None = None
1631 if expires_in is not None:
1632 expires_at = (datetime.now(timezone.utc) + timedelta(seconds=expires_in)).isoformat()
1634 payload: Final[OAuthCredentialPayload] = {
1635 "type": "oauth2",
1636 "access_token": access_token,
1637 "connected_at": datetime.now(timezone.utc).isoformat(),
1638 **({"identity_binding_proof": identity_binding_proof} if identity_binding_proof else {}),
1639 }
1640 if refresh_token: 1640 ↛ 1641line 1640 didn't jump to line 1641 because the condition on line 1640 was never true
1641 payload["refresh_token"] = refresh_token
1642 if expires_at:
1643 payload["expires_at"] = expires_at
1644 if scopes:
1645 payload["scopes"] = scopes
1647 # Guard against silently overwriting a BYOK credential with an OAuth token.
1648 # Skip the guard when the caller knows the row is already an OAuth2 credential
1649 # (e.g. during token refresh), saving an extra DB round-trip.
1650 if not skip_byok_guard: 1650 ↛ 1673line 1650 didn't jump to line 1673 because the condition on line 1650 was always true
1651 existing: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
1652 decoded: Final = _decode_user_credential(existing.credential_b64) if existing is not None else None
1653 if existing is not None and _parse_oauth_payload(decoded) is None: 1653 ↛ 1659line 1653 didn't jump to line 1659 because the condition on line 1653 was never true
1654 # Refuse only while the row still holds readable content, which is a live BYOK
1655 # secret that overwriting would destroy. A row that does not decode was written
1656 # under a different LITELLM_SALT_KEY, and one that decodes to nothing holds no
1657 # secret at all; refusing either preserves nothing and instead wedges the user
1658 # out of the OAuth flow for good, since re-authorizing is their only recovery.
1659 if decoded:
1660 raise ValueError(
1661 f"Existing credential for user {user_id} and server "
1662 f"{server_id} could not be verified as an OAuth2 token. "
1663 f"Refusing to overwrite."
1664 )
1665 verbose_proxy_logger.warning(
1666 "store_user_oauth_credential: existing credential for user=%s server=%s could not be "
1667 "decrypted (likely written under a previous LITELLM_SALT_KEY); replacing it with the "
1668 "newly authorized OAuth2 token.",
1669 user_id,
1670 server_id,
1671 )
1673 encoded: Final = encrypt_value_helper(json.dumps(payload))
1674 await _db_upsert_user_credential_row(prisma_client, user_id, server_id, encoded)
1677def is_oauth_credential_expired(cred: OAuthCredentialPayload, buffer_seconds: int = 0) -> bool:
1678 """Return True if the OAuth2 credential's access_token has expired.
1680 Checks the ``expires_at`` ISO-format string stored in the credential payload.
1681 Returns False when ``expires_at`` is absent or unparseable (treat as non-expired).
1682 With ``buffer_seconds`` > 0, a token that is still valid but expires within the
1683 buffer is also treated as expired, so callers can refresh proactively instead of
1684 handing back a token that may lapse mid-request.
1685 """
1686 expires_at: Final = cred.get("expires_at")
1687 if not expires_at:
1688 return False
1689 try:
1690 exp_dt = datetime.fromisoformat(expires_at)
1691 if exp_dt.tzinfo is None:
1692 exp_dt = exp_dt.replace(tzinfo=timezone.utc)
1693 return datetime.now(timezone.utc) + timedelta(seconds=buffer_seconds) > exp_dt
1694 except (ValueError, TypeError):
1695 return False
1698def oauth_grant_state(cred: OAuthCredentialPayload | None) -> OAuthGrantState:
1699 """Classify local grant readiness without attempting a refresh or checking upstream revocation."""
1700 if not cred or not cred.get("access_token"):
1701 return "absent"
1702 if not is_oauth_credential_expired(cred, buffer_seconds=MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS):
1703 return "valid"
1704 return "refreshable" if cred.get("refresh_token") else "absent"
1707async def get_user_oauth_credential(
1708 prisma_client: PrismaClient,
1709 user_id: str,
1710 server_id: str,
1711) -> OAuthCredentialPayload | None:
1712 """Return the decoded OAuth2 payload dict for a user+server pair, or None."""
1714 row: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
1715 if row is None:
1716 return None
1717 decoded: Final = _decode_user_credential(row.credential_b64)
1718 if decoded is None: 1718 ↛ 1719line 1718 didn't jump to line 1719 because the condition on line 1718 was never true
1719 _warn_undecryptable_credential(user_id, server_id)
1720 return _parse_oauth_payload(decoded)
1723def _server_user_credential_item(
1724 row: "prisma_db_models.LiteLLM_MCPUserCredentials",
1725) -> MCPServerUserCredentialListItem:
1726 oauth_payload: Final = _decode_oauth_payload(row.credential_b64)
1727 if oauth_payload is None: 1727 ↛ 1728line 1727 didn't jump to line 1728 because the condition on line 1727 was never true
1728 return MCPServerUserCredentialListItem(
1729 user_id=row.user_id,
1730 credential_type="byok",
1731 updated_at=row.updated_at.isoformat(),
1732 )
1733 return MCPServerUserCredentialListItem(
1734 user_id=row.user_id,
1735 credential_type="oauth2",
1736 expires_at=oauth_payload.get("expires_at"),
1737 connected_at=oauth_payload.get("connected_at"),
1738 updated_at=row.updated_at.isoformat(),
1739 )
1742async def list_server_user_credentials(
1743 prisma_client: PrismaClient,
1744 server_id: str,
1745) -> tuple[MCPServerUserCredentialListItem, ...]:
1746 """Every user's stored credential for one server, typed but without the secret, for admins."""
1747 rows: Final = await _db_find_user_credential_rows(
1748 prisma_client,
1749 {"server_id": server_id}, # mutable-ok: prisma where-inputs must be plain dicts
1750 )
1751 return tuple(_server_user_credential_item(row) for row in rows)
1754async def list_user_oauth_credentials(
1755 prisma_client: PrismaClient,
1756 user_id: str,
1757) -> list[OAuthCredentialPayload]:
1758 """Return all OAuth2 credential payloads for a user, tagged with server_id."""
1760 rows: Final = await _db_find_user_credential_rows(prisma_client, {"user_id": user_id})
1761 results: Final[list[OAuthCredentialPayload]] = []
1762 for row in rows:
1763 decoded = _decode_user_credential(row.credential_b64)
1764 if decoded is None: 1764 ↛ 1765line 1764 didn't jump to line 1765 because the condition on line 1764 was never true
1765 _warn_undecryptable_credential(user_id, row.server_id)
1766 payload = _parse_oauth_payload(decoded)
1767 if payload is None: 1767 ↛ 1768line 1767 didn't jump to line 1768 because the condition on line 1767 was never true
1768 continue
1769 payload["server_id"] = row.server_id
1770 results.append(payload)
1771 return results
1774def _decrypted_credential_field(creds: dict[str, object], field: str) -> object:
1775 """Return one credential field decrypted with the global salt key; non-string and legacy
1776 plaintext values come back unchanged (decrypt_value_helper returns the original on failure)."""
1777 value: Final = creds.get(field)
1778 if not isinstance(value, str): 1778 ↛ 1780line 1778 didn't jump to line 1780 because the condition on line 1778 was always true
1779 return value
1780 return decrypt_value_helper(
1781 value=value,
1782 key=field,
1783 exception_type="debug",
1784 return_original_value=True,
1785 )
1788def mcp_oauth_token_identity(server: object) -> tuple[object, ...]:
1789 """The upstream-OAuth-token-determining fields of an MCP server: the resource/audience (url, or
1790 spec_path for OpenAPI servers, plus the RFC 8707 upstream_resource sent on the authorize and
1791 token legs), the OAuth mode/grant (auth_type, oauth2_flow), the authorization-server endpoints,
1792 and the OAuth client + scopes. Mirrors the dashboard's getOAuthAuthorizationIdentity. When any
1793 of these change on a server update, previously stored per-user tokens were minted for the old
1794 identity and are stale. Excludes transport and delegate_auth_to_upstream, which do not affect
1795 what token is minted (RFC 8693).
1797 client_id/client_secret are compared decrypted: stored values are NaCl-encrypted with a fresh
1798 nonce on every write, so comparing ciphertext would flag every routine save as an identity
1799 change and purge tokens that are still valid."""
1800 creds: Final = getattr(server, "credentials", None)
1801 if isinstance(creds, str): 1801 ↛ 1802line 1801 didn't jump to line 1802 because the condition on line 1801 was never true
1802 try:
1803 parsed: dict[str, object] | None = json.loads(creds)
1804 except ValueError:
1805 parsed = None
1806 else:
1807 parsed = creds
1808 creds_dict: Final[dict[str, object]] = parsed if isinstance(parsed, dict) else {}
1809 return (
1810 getattr(server, "url", None),
1811 getattr(server, "spec_path", None),
1812 getattr(server, "auth_type", None),
1813 getattr(server, "oauth2_flow", None),
1814 getattr(server, "issuer", None),
1815 getattr(server, "authorization_url", None),
1816 getattr(server, "token_url", None),
1817 getattr(server, "registration_url", None),
1818 _decrypted_credential_field(creds_dict, "client_id"),
1819 _decrypted_credential_field(creds_dict, "client_secret"),
1820 creds_dict.get("scopes"),
1821 creds_dict.get("upstream_resource"),
1822 )
1825async def purge_user_oauth_credentials_for_server(
1826 prisma_client: PrismaClient,
1827 server_id: str,
1828 invalidate_token_cache: Callable[[str, str], Awaitable[None]] | None = None,
1829) -> int:
1830 """Delete every stored per-user OAuth token for a server and invalidate each user's cached
1831 token everywhere it can be served from (the legacy per-user token cache and the v2 per-user OAuth
1832 token store), so no user keeps a token minted for a superseded configuration. Called when a server
1833 update changes a mint-relevant field (see mcp_oauth_token_identity). Returns the number of rows
1834 removed.
1836 LiteLLM_MCPUserCredentials also stores BYOK API keys in the same column; only rows whose payload
1837 decodes as an OAuth2 credential (see _decode_oauth_payload) are deleted, because a config change
1838 only invalidates minted tokens, never a user's own stored key. Rows are therefore deleted per
1839 (user_id, server_id) pair rather than by a blanket server_id filter. An OAuth row inserted while
1840 the purge runs for a user not yet enumerated survives; a re-auth completing in the window for an
1841 already-enumerated user is deleted along with the stale row (the pair delete cannot tell them
1842 apart), which costs that user one extra re-auth and nothing else.
1844 invalidate_token_cache is injectable for tests; it defaults to the manager's shared
1845 invalidate_user_oauth_token_cache, the single invalidation point for per-user tokens."""
1846 rows: Final = await _db_find_user_credential_rows(prisma_client, {"server_id": server_id})
1847 oauth_rows: Final = [row for row in rows if _decode_oauth_payload(row.credential_b64) is not None]
1848 if not oauth_rows: 1848 ↛ 1850line 1848 didn't jump to line 1850 because the condition on line 1848 was always true
1849 return 0
1850 deleted_count: Final = await _user_credential_actions(prisma_client).delete_many(
1851 where={"server_id": server_id, "user_id": {"in": [row.user_id for row in oauth_rows]}}
1852 )
1853 if invalidate_token_cache is None:
1854 from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
1855 global_mcp_server_manager,
1856 )
1858 invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
1860 for row in oauth_rows:
1861 await invalidate_token_cache(row.user_id, server_id)
1862 if deleted_count != len(oauth_rows):
1863 verbose_proxy_logger.warning(
1864 "MCP server %s: purge removed %d OAuth credential row(s) but %d were enumerated; "
1865 "row(s) were deleted concurrently during the purge",
1866 server_id,
1867 deleted_count,
1868 len(oauth_rows),
1869 )
1870 return deleted_count
1873async def refresh_user_oauth_token(
1874 prisma_client: PrismaClient,
1875 user_id: str,
1876 server: "MCPServer",
1877 cred: OAuthCredentialPayload,
1878) -> OAuthCredentialPayload | None:
1879 """Attempt to refresh a per-user OAuth2 token using its stored refresh_token.
1881 POSTs to ``server.effective_token_url`` with ``grant_type=refresh_token``.
1883 On success: persists the new credential via ``store_user_oauth_credential``
1884 and returns the updated payload dict.
1885 On failure (network error, invalid_grant, missing refresh_token, …): logs a
1886 warning and returns ``None`` — the caller is responsible for clearing the
1887 stale credential and triggering re-authentication.
1888 """
1889 binding: Final = server.oauth_identity_binding
1890 if binding is not None and binding.mode == "enforce":
1891 if not await credential_binding_matches(binding, user_id, server.server_id, cred):
1892 return None
1894 refresh_token: Final[str | None] = cred.get("refresh_token")
1895 token_url: Final[str | None] = getattr(server, "effective_token_url", None) or getattr(server, "token_url", None)
1896 server_id: Final[str] = getattr(server, "server_id", "")
1897 client_id: Final[str | None] = getattr(server, "client_id", None)
1898 client_secret: Final[str | None] = getattr(server, "client_secret", None)
1900 if not refresh_token:
1901 verbose_proxy_logger.debug(
1902 "refresh_user_oauth_token: no refresh_token stored for user=%s server=%s",
1903 user_id,
1904 server_id,
1905 )
1906 return None
1907 if not token_url:
1908 verbose_proxy_logger.debug(
1909 "refresh_user_oauth_token: server=%s has no token_url configured",
1910 server_id,
1911 )
1912 return None
1914 try:
1915 token_request: Final = build_upstream_oauth2_token_request(
1916 server,
1917 auth_method=getattr(server, "token_endpoint_auth_method", None),
1918 client_id=client_id,
1919 client_secret=client_secret,
1920 )
1921 token_data: Final[dict[str, str]] = {
1922 "grant_type": "refresh_token",
1923 "refresh_token": refresh_token,
1924 **token_request.body,
1925 }
1926 async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
1927 response: Final = await async_client.post(
1928 token_url,
1929 headers={"Accept": "application/json", **token_request.headers},
1930 data=token_data,
1931 )
1932 response.raise_for_status()
1933 body: Final[_OAuthTokenRefreshResponse] = response.json()
1934 except Exception as exc:
1935 verbose_proxy_logger.warning(
1936 "refresh_user_oauth_token: refresh request failed for user=%s server=%s: %s",
1937 user_id,
1938 server_id,
1939 exc,
1940 )
1941 return None
1943 try:
1944 binding_proof: Final = await enforce_oauth_identity_binding(
1945 server=server,
1946 token_response=body,
1947 litellm_user_id=user_id,
1948 grant_type="refresh_token",
1949 refresh_ownership=RefreshTokenPresented(refresh_token),
1950 )
1951 except HTTPException as exc:
1952 if exc.status_code != 403:
1953 raise
1954 return None
1956 access_token: Final[str | None] = body.get("access_token")
1957 if not access_token:
1958 verbose_proxy_logger.warning(
1959 "refresh_user_oauth_token: token response missing access_token for user=%s server=%s",
1960 user_id,
1961 server_id,
1962 )
1963 return None
1965 expires_in: int | None = None
1966 raw_expires: Final = body.get("expires_in")
1967 try:
1968 expires_in = int(raw_expires) if raw_expires is not None else None
1969 except (TypeError, ValueError):
1970 pass
1972 # Rotate refresh token when the provider returns a new one
1973 new_refresh_token: Final[str | None] = body.get("refresh_token") or refresh_token
1975 raw_scope: Final = body.get("scope")
1976 scopes: list[str] | None = (raw_scope.split() if isinstance(raw_scope, str) and raw_scope else None) or cred.get(
1977 "scopes"
1978 )
1980 await store_user_oauth_credential(
1981 prisma_client=prisma_client,
1982 user_id=user_id,
1983 server_id=server_id,
1984 access_token=access_token,
1985 refresh_token=new_refresh_token,
1986 expires_in=expires_in,
1987 scopes=scopes,
1988 identity_binding_proof=binding_proof,
1989 skip_byok_guard=True, # Row is already OAuth2; skip the extra find_unique check
1990 )
1992 verbose_proxy_logger.info(
1993 "refresh_user_oauth_token: refreshed token for user=%s server=%s",
1994 user_id,
1995 server_id,
1996 )
1997 return await get_user_oauth_credential(prisma_client, user_id, server_id)
2000async def resolve_valid_user_oauth_token(
2001 user_id: str,
2002 server: "MCPServer",
2003 cred: OAuthCredentialPayload | None,
2004 prisma_client: PrismaClient | None = None,
2005) -> OAuthCredentialPayload | None:
2006 """Return an OAuth2 credential whose access_token is good for the next request.
2008 Returns the credential unchanged while its token is valid for at least
2009 ``MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS``. Only when the token is expired (or
2010 expiring within that buffer) and a refresh_token is stored does it mint a new one
2011 via ``refresh_user_oauth_token``. Returns None when there is no usable token
2012 (missing token, expired with no refresh_token, or a failed refresh).
2014 The refresh_token is only ever sent to the server's token_url inside
2015 ``refresh_user_oauth_token``; it is never exposed to the caller beyond the cred
2016 dict it already holds. ``prisma_client`` is fetched lazily and only when a refresh
2017 actually happens, so the valid-token path never requires a DB handle.
2018 """
2019 grant: Final = oauth_grant_state(cred)
2020 if cred is None or grant == "absent":
2021 return None
2022 binding: Final = server.oauth_identity_binding
2023 if binding is not None and binding.mode == "enforce":
2024 if not await credential_binding_matches(binding, user_id, server.server_id, cred):
2025 return None
2026 if grant == "valid":
2027 return cred
2028 if prisma_client is None:
2029 from litellm.proxy.utils import get_prisma_client_or_throw
2031 prisma_client = get_prisma_client_or_throw("Database not connected. Cannot refresh OAuth token.")
2032 refreshed: Final = await refresh_user_oauth_token(
2033 prisma_client=prisma_client,
2034 user_id=user_id,
2035 server=server,
2036 cred=cred,
2037 )
2038 if not refreshed or not refreshed.get("access_token"):
2039 return None
2040 return refreshed
2043async def resolve_user_oauth_access_token(
2044 user_id: str | None,
2045 server: "MCPServer",
2046 prefetched_creds: Mapping[str, OAuthCredentialPayload] | None = None,
2047) -> str | None:
2048 """Resolve a user's valid OAuth2 access token for a server: Redis cache, else DB + refresh.
2050 The egress token-resolution core shared by v1's header builder and the v2 ``OAuthTokenStore``
2051 adapter. Redis fast-path (skipped when ``prefetched_creds`` is supplied), else a DB read through
2052 ``resolve_valid_user_oauth_token`` (which refreshes an expired token when a ``refresh_token`` is
2053 stored), re-warming the Redis cache with the per-server TTL. Returns ``None`` when there is no
2054 usable token; any error is swallowed to ``None`` so a transient failure reads as "not
2055 authorized" rather than raising.
2056 """
2057 server_id: Final[str | None] = getattr(server, "server_id", None)
2058 if not user_id or not server_id:
2059 return None
2060 try:
2061 from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
2062 _compute_per_user_token_ttl,
2063 mcp_per_user_token_cache,
2064 )
2066 binding: Final = server.oauth_identity_binding
2067 enforce_binding: Final = binding is not None and binding.mode == "enforce"
2068 if prefetched_creds is None and enforce_binding and binding is not None:
2069 bound_token: Final = await mcp_per_user_token_cache.get_token(user_id, server_id)
2070 if bound_token is not None:
2071 if await credential_binding_matches(
2072 binding, user_id, server_id, {"identity_binding_proof": bound_token.identity_binding_proof}
2073 ):
2074 return bound_token.access_token
2075 await mcp_per_user_token_cache.delete(user_id, server_id)
2076 if prefetched_creds is None and not enforce_binding:
2077 cached_token: Final = await mcp_per_user_token_cache.get(user_id, server_id)
2078 if cached_token is not None:
2079 return cached_token
2081 prisma_client = None
2082 if prefetched_creds is not None:
2083 cred = prefetched_creds.get(server_id)
2084 else:
2085 from litellm.proxy.utils import get_prisma_client_or_throw
2087 prisma_client = get_prisma_client_or_throw(
2088 "Database not connected. Connect a database to use OAuth2 MCP tools."
2089 )
2090 cred = await get_user_oauth_credential(prisma_client, user_id, server_id)
2092 if not cred or not cred.get("access_token"):
2093 return None
2095 cred = await resolve_valid_user_oauth_token(
2096 user_id=user_id,
2097 server=server,
2098 cred=cred,
2099 prisma_client=prisma_client,
2100 )
2101 if cred is None:
2102 # Refresh failed or token expired with no usable refresh_token — clear the stale
2103 # Redis entry so the next request doesn't reuse it.
2104 await mcp_per_user_token_cache.delete(user_id, server_id)
2105 return None
2107 access_token: Final[str] = cred["access_token"]
2108 if prefetched_creds is None:
2109 ttl: Final = _compute_per_user_token_ttl(server, _remaining_token_seconds(cred.get("expires_at")))
2110 await mcp_per_user_token_cache.set(
2111 user_id, server_id, access_token, ttl, identity_binding_proof=cred.get("identity_binding_proof")
2112 )
2113 return access_token
2114 except Exception as e:
2115 verbose_proxy_logger.warning(
2116 "resolve_user_oauth_access_token: failed for user=%s server=%s: %s",
2117 user_id,
2118 server_id,
2119 e,
2120 )
2121 return None
2124def _remaining_token_seconds(expires_at: str | None) -> int | None:
2125 """Seconds until ``expires_at`` (ISO 8601), or None when absent/past/unparseable."""
2126 if not expires_at:
2127 return None
2128 try:
2129 exp_dt = datetime.fromisoformat(expires_at)
2130 except (ValueError, TypeError):
2131 return None
2132 if exp_dt.tzinfo is None:
2133 exp_dt = exp_dt.replace(tzinfo=timezone.utc)
2134 remaining: Final = int((exp_dt - datetime.now(timezone.utc)).total_seconds())
2135 return remaining if remaining > 0 else None
2138async def get_active_submitted_mcp_server_ids_for_user(
2139 prisma_client: PrismaClient,
2140 user_id: str,
2141) -> list[str]:
2142 """Return active BYOM servers submitted by this user (creator visibility)."""
2143 if not user_id: 2143 ↛ 2144line 2143 didn't jump to line 2144 because the condition on line 2143 was never true
2144 return []
2146 rows: Final = await _db_find_mcp_server_rows(
2147 prisma_client,
2148 {
2149 "submitted_by": user_id,
2150 "approval_status": MCPApprovalStatus.active,
2151 },
2152 )
2153 return [row.server_id for row in rows]
2156async def approve_mcp_server(
2157 prisma_client: PrismaClient,
2158 server_id: str,
2159 touched_by: str,
2160) -> LiteLLM_MCPServerTable:
2161 """Set approval_status=active and record reviewed_at."""
2162 now: Final = datetime.now(timezone.utc)
2163 updated: Final = await _db_update_mcp_server_row(
2164 prisma_client,
2165 server_id,
2166 {
2167 "approval_status": MCPApprovalStatus.active,
2168 "reviewed_at": now,
2169 "updated_by": touched_by,
2170 },
2171 )
2172 table: Final = LiteLLM_MCPServerTable.model_validate(updated.model_dump())
2173 decrypt_global_env_var_values(table.env_vars)
2174 return table
2177async def reject_mcp_server(
2178 prisma_client: PrismaClient,
2179 server_id: str,
2180 touched_by: str,
2181 review_notes: str | None = None,
2182) -> LiteLLM_MCPServerTable:
2183 """Set approval_status=rejected, record reviewed_at and review_notes."""
2184 now: Final = datetime.now(timezone.utc)
2185 data: Final[prisma_db_types.LiteLLM_MCPServerTableUpdateInput] = {
2186 "approval_status": MCPApprovalStatus.rejected,
2187 "reviewed_at": now,
2188 "updated_by": touched_by,
2189 }
2190 if review_notes is not None:
2191 data["review_notes"] = review_notes
2192 updated: Final = await _db_update_mcp_server_row(prisma_client, server_id, data)
2193 table: Final = LiteLLM_MCPServerTable.model_validate(updated.model_dump())
2194 decrypt_global_env_var_values(table.env_vars)
2195 return table
2198async def get_mcp_submissions(
2199 prisma_client: PrismaClient,
2200) -> MCPSubmissionsSummary:
2201 """
2202 Returns all MCP servers that were submitted by non-admin users (submitted_at IS NOT NULL),
2203 along with a summary count breakdown by approval_status.
2204 Mirrors get_guardrail_submissions() from guardrail_endpoints.py.
2205 """
2206 rows: Final[Sequence[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions(
2207 prisma_client
2208 ).find_many(
2209 where={"submitted_at": {"not": None}},
2210 order={"submitted_at": "desc"},
2211 take=500, # safety cap; paginate if needed in a future iteration
2212 )
2213 items: Final = list(_readable_mcp_servers(rows))
2215 pending: Final = sum(1 for i in items if i.approval_status == MCPApprovalStatus.pending_review)
2216 active: Final = sum(1 for i in items if i.approval_status == MCPApprovalStatus.active)
2217 rejected: Final = sum(1 for i in items if i.approval_status == MCPApprovalStatus.rejected)
2219 return MCPSubmissionsSummary(
2220 total=len(items),
2221 pending_review=pending,
2222 active=active,
2223 rejected=rejected,
2224 items=items,
2225 )
2228# ── Per-user MCP environment variables ────────────────────────────────────
2231def _decode_user_env_vars(stored: str) -> dict[str, str]:
2232 """Decrypt a ``values_b64`` blob and parse it as a flat ``{name: value}`` dict."""
2233 decrypted: Final = decrypt_value_helper(
2234 value=stored,
2235 key="mcp_user_env_vars",
2236 exception_type="debug",
2237 return_original_value=False,
2238 )
2239 if decrypted is None: 2239 ↛ 2240line 2239 didn't jump to line 2240 because the condition on line 2239 was never true
2240 if stored:
2241 verbose_proxy_logger.warning(
2242 "MCP per-user env vars failed to decrypt (LITELLM_SALT_KEY "
2243 "changed?); treating as unset so the user is prompted to "
2244 "re-enter them rather than silently forwarding ciphertext"
2245 )
2246 return {}
2247 parsed: dict[str, object] | None
2248 try:
2249 parsed = json.loads(decrypted)
2250 except (ValueError, TypeError):
2251 return {}
2252 if not isinstance(parsed, dict): 2252 ↛ 2253line 2252 didn't jump to line 2253 because the condition on line 2252 was never true
2253 return {}
2254 return {str(k): str(v) for k, v in parsed.items()}
2257async def get_user_env_vars(
2258 prisma_client: PrismaClient,
2259 user_id: str,
2260 server_id: str,
2261) -> dict[str, str]:
2262 """Return the calling user's env var dict for ``server_id`` (empty if none)."""
2263 row: Final = await _user_env_var_actions(prisma_client).find_unique(
2264 where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
2265 )
2266 if row is None:
2267 return {}
2268 return _decode_user_env_vars(row.values_b64)
2271async def get_user_env_vars_bulk(
2272 prisma_client: PrismaClient,
2273 user_id: str,
2274 server_ids: Iterable[str],
2275) -> dict[str, dict[str, str]]:
2276 """Return ``{server_id: {var_name: value}}`` for one user across many servers.
2278 Servers with no stored row are simply absent from the result.
2279 """
2280 ids: Final = list(server_ids)
2281 if not ids: 2281 ↛ 2282line 2281 didn't jump to line 2282 because the condition on line 2281 was never true
2282 return {}
2283 rows: Final = await _db_find_user_env_var_rows(prisma_client, {"user_id": user_id, "server_id": {"in": ids}})
2284 return {row.server_id: _decode_user_env_vars(row.values_b64) for row in rows}
2287async def merge_user_env_vars(
2288 prisma_client: PrismaClient,
2289 user_id: str,
2290 server_id: str,
2291 updates: dict[str, str],
2292 allowed_names: Iterable[str],
2293) -> dict[str, str]:
2294 """Merge ``updates`` into the user's stored env vars for ``server_id`` and
2295 return the resulting set.
2297 The read-modify-write runs inside a transaction guarded by a
2298 ``(user_id, server_id)`` advisory lock so two concurrent writes from the
2299 same user can't drop one update. Names outside ``allowed_names`` are pruned,
2300 so an admin retiring a user-scoped variable also clears its stored value.
2301 """
2302 allowed: Final = set(allowed_names)
2303 lock_key: Final = int.from_bytes(
2304 hashlib.blake2b(f"{user_id}:{server_id}".encode(), digest_size=8).digest(),
2305 "big",
2306 signed=True,
2307 )
2308 async with _db_transaction_manager(prisma_client) as tx:
2309 await tx.execute_raw("SELECT pg_advisory_xact_lock($1::bigint)", lock_key)
2310 row: Final[prisma_db_models.LiteLLM_MCPUserEnvVars | None] = await tx.litellm_mcpuserenvvars.find_unique(
2311 where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
2312 )
2313 existing: Final = _decode_user_env_vars(row.values_b64) if row is not None else {}
2314 merged: Final = {k: v for k, v in {**existing, **updates}.items() if k in allowed}
2315 encoded: Final = encrypt_value_helper(json.dumps(merged))
2316 await tx.litellm_mcpuserenvvars.upsert(
2317 where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
2318 data={
2319 "create": {
2320 "user_id": user_id,
2321 "server_id": server_id,
2322 "values_b64": encoded,
2323 },
2324 "update": {"values_b64": encoded},
2325 },
2326 )
2327 return merged
2330async def delete_user_env_vars(
2331 prisma_client: PrismaClient,
2332 user_id: str,
2333 server_id: str,
2334) -> None:
2335 """Remove the calling user's env var values for ``server_id``.
2337 Uses ``delete_many`` so a missing row is a no-op; real DB errors still
2338 propagate to the caller instead of being silently swallowed.
2339 """
2340 await _user_env_var_actions(prisma_client).delete_many(where={"user_id": user_id, "server_id": server_id})