litellm/litellm/proxy/_experimental/mcp_server/db.py
joshua-berri 3743c8563e
fix(mcp): adopt shared server resolution and caller authorization (#43263)
* test(mcp): characterize server resolution and authorization

* refactor(mcp): extract shared server resolution

* fix(mcp): adopt shared resolution in management endpoints

* fix(mcp): scope credential metadata resolution outside loop

* test(mcp): pin catalog isolation and batched credential permissions

* fix(mcp): restrict catalog detail and batch credential permissions

* test(mcp): enforce identity isolation in database fixtures

* test(mcp): name resolution tests by behavior

* test(mcp): describe detail access assertion failures

* chore: keep agent naming discipline local

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
2026-09-26 14:37:32 -07:00

2309 lines
94 KiB
Python

import base64
import binascii
import hashlib
import json
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict, cast
from fastapi import HTTPException
from typing_extensions import ReadOnly
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.oauth_identity_binding import (
RefreshTokenPresented,
credential_binding_matches,
enforce_oauth_identity_binding,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
MCPApprovalStatus,
MCPEnvVar,
MCPEnvVarScope,
MCPServerUserCredentialListItem,
MCPSubmissionsSummary,
NewMCPServerRequest,
UpdateMCPServerRequest,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
SecretMapDecodeError,
_get_salt_key,
decode_secret_map,
decrypt_value_helper,
encrypt_secret_map,
encrypt_value_helper,
)
from litellm.proxy.utils import PrismaClient
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import (
MCPServerOAuthClientRepository,
MCPServerRepository,
MCPUserCredentialsRepository,
PrismaTableRepository,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPCredentials
if TYPE_CHECKING:
from prisma import models as prisma_db_models
from prisma import types as prisma_db_types
from litellm.types.mcp_server.mcp_server_manager import MCPServer
class _UserEnvVarsTransactionClient(Protocol):
litellm_mcpuserenvvars: "TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]"
litellm_mcpservertable: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]"
async def execute_raw(self, query: str, *args: object) -> int: ...
class _UserEnvVarsTransaction(Protocol):
async def __aenter__(self) -> _UserEnvVarsTransactionClient: ...
async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ...
@dataclass(frozen=True, slots=True)
class McpIdentifierConflict:
"""An incoming ``server_name``/``alias`` already belongs to another MCP server row.
``field`` is the incoming identifier that collided, ``value`` the submitted
string, and ``server_id`` the existing row that owns it.
"""
field: Literal["server_name", "alias"]
value: str
server_id: str
_AUTH_FLOW_SCOPED_FIELDS: Final["frozenset[str]"] = frozenset(
{
"issuer",
"authorization_url",
"token_url",
"registration_url",
"oauth2_flow",
"dcr_bridge",
"token_exchange_endpoint",
"audience",
"subject_token_type",
"token_exchange_profile",
}
)
def _blank_to_none(value: str | None) -> str | None:
if not isinstance(value, str):
return None
return value.strip() or None
# Token-exchange settings with dedicated columns that also exist on
# ``MCPCredentials`` as a legacy shape (rows and REST callers that predate the
# columns). Every write lifts blob values into the columns and strips them from
# the stored blob, so the read-time ``column or blob`` fallback only serves rows
# the current code has never written — a cleared column can then never be
# silently resurrected by a stale blob copy. These keys are stored plaintext
# (endpoints/identifiers, not secrets), so values lift as-is.
_TOKEN_EXCHANGE_COLUMN_FIELDS: Final["frozenset[str]"] = frozenset(
{
"token_exchange_endpoint",
"audience",
"subject_token_type",
"token_exchange_profile",
}
)
# The client-forwarded token modes share one stored-credential shape: the admin-declared upstream
# OAuth app (client_id/client_secret) plus the same authorize relay, and neither mints anything the
# gateway keeps. So a switch WITHIN this class must preserve the stored app, unlike a cross-class
# switch (e.g. an oauth2 row whose client may be DCR-minted and is not reusable elsewhere).
_CLIENT_FORWARDED_AUTH_TYPES: Final["frozenset[str]"] = frozenset({"true_passthrough", "oauth_delegate"})
# Minted token material that must never survive a client rotation on a persisted row.
_MINTED_TOKEN_CREDENTIAL_FIELDS: Final["frozenset[str]"] = frozenset({"access_token", "refresh_token", "expires_in"})
class _OAuthCredentialAccessToken(TypedDict):
access_token: str
class OAuthCredentialPayload(_OAuthCredentialAccessToken, total=False):
identity_binding_proof: ReadOnly[str]
type: str
refresh_token: str
expires_at: str
connected_at: str
scopes: list[str]
server_id: str
OAuthGrantState = Literal["valid", "refreshable", "absent"]
class _OAuthTokenRefreshResponse(TypedDict, total=False):
access_token: str
refresh_token: str
expires_in: int
scope: str
def _credential_auth_class(auth_type: str | None) -> str | None:
"""Collapse the client-forwarded modes to one credential class; every other auth_type is its own
class. Used so credential handling keys off whether the stored-credential shape actually changed,
not off a raw auth_type inequality that treats true_passthrough<->oauth_delegate as a full reset."""
if auth_type in _CLIENT_FORWARDED_AUTH_TYPES:
return "client_forwarded"
return auth_type
def _drop_stale_minted_on_client_rotation(merged: dict[str, object], new_creds: dict[str, object]) -> dict[str, object]:
"""When the update rotates the client, drop stale minted token keys it did not itself set, so an old
app's access/refresh token never rides forward under the new client. A no-op when no client key changed."""
if "client_id" not in new_creds and "client_secret" not in new_creds:
return merged
return {
key: value for key, value in merged.items() if key not in _MINTED_TOKEN_CREDENTIAL_FIELDS or key in new_creds
}
def _is_global_env_var_scope(scope: object) -> bool:
"""``scope="user"`` entries are placeholders the user fills in; everything
else (including a missing scope) is an admin-supplied global value."""
return scope != MCPEnvVarScope.user and scope != "user"
def _encrypt_global_env_var_values(env_vars: Iterable[dict[str, str]]) -> None:
"""Encrypt ``scope="global"`` env var values in place before persisting.
Global values hold admin-supplied secrets (API keys, passwords) that get
interpolated into headers, so they are encrypted at rest like credentials
and the per-user ``values_b64`` column. Per-user placeholders are not
secrets and are stored verbatim.
"""
for entry in env_vars:
if not _is_global_env_var_scope(entry.get("scope")):
continue
value = entry.get("value")
if value:
entry["value"] = encrypt_value_helper(value)
def decrypt_global_env_var_values(env_vars: Iterable[MCPEnvVar | dict[str, str]] | None) -> None:
"""Decrypt ``scope="global"`` env var values in place after reading the DB.
Accepts ``MCPEnvVar`` models (``LiteLLM_MCPServerTable``) or plain dicts
(raw rows / deserialized JSON). Global values are always stored encrypted,
so a value that no longer decrypts (e.g. after a salt-key change) is dropped
and a warning is logged rather than forwarding the ciphertext into upstream
``${NAME}`` headers, where it would silently fail.
"""
if not env_vars:
return
for entry in env_vars:
is_dict = isinstance(entry, dict)
scope = entry.get("scope") if is_dict else getattr(entry, "scope", None)
if not _is_global_env_var_scope(scope):
continue
value = entry.get("value") if is_dict else getattr(entry, "value", None)
if not value:
continue
decrypted = decrypt_value_helper(
value=value,
key="mcp_global_env_var",
exception_type="debug",
return_original_value=False,
)
if decrypted is None:
name = entry.get("name") if is_dict else getattr(entry, "name", None)
verbose_proxy_logger.warning(
"MCP global env var %s failed to decrypt (LITELLM_SALT_KEY "
"changed?); dropping it so ciphertext is not sent upstream",
name,
)
decrypted = ""
if is_dict:
entry["value"] = decrypted
else:
entry.value = decrypted
def _decrypt_env_vars_on_returned_row(row: object) -> None:
"""Decrypt ``scope="global"`` env var values on a row returned by Prisma create/update.
Prisma may hand back ``env_vars`` either as a parsed list (the common case for
JSONB columns) or as a raw JSON string (observed for some write paths). The
in-place decrypt helper only mutates iterables of dicts/models, so a string
payload would silently skip decryption and ciphertext would leak into the
registry via ``add_server``/``update_server`` (which trust the caller).
Parse the string back to a list so the in-place decrypt actually runs, and
write the decrypted list back onto the row so downstream consumers see plain
values.
"""
env_vars = getattr(row, "env_vars", None)
if env_vars is None:
return
if isinstance(env_vars, str):
try:
env_vars = json.loads(env_vars)
except (json.JSONDecodeError, TypeError):
return
if not isinstance(env_vars, list):
return
try:
setattr(row, "env_vars", env_vars)
except (AttributeError, TypeError):
pass
decrypt_global_env_var_values(env_vars)
def _reencrypt_global_env_var_values(
env_vars: str | Iterable[Mapping[str, str]] | None, new_encryption_key: str
) -> list[dict[str, str]] | None:
"""Re-encrypt ``scope="global"`` env var values for master-key rotation.
Each global value is decrypted with the current salt key and re-encrypted
under ``new_encryption_key``. Returns the rebuilt list when at least one
value was rotated, else ``None`` so the caller can skip the DB write. A
value that fails to decrypt is left untouched (and logged) so a corrupt
entry is preserved for recovery rather than overwritten.
"""
if not env_vars:
return None
entries: Iterable[Mapping[str, str]]
if isinstance(env_vars, str):
try:
entries = json.loads(env_vars)
except (json.JSONDecodeError, TypeError):
return None
if not entries:
return None
else:
entries = env_vars
rebuilt: Final = [dict(v) for v in entries]
rotated = False
for entry in rebuilt:
if not _is_global_env_var_scope(entry.get("scope")):
continue
value = entry.get("value")
if not value:
continue
decrypted = decrypt_value_helper(
value=value,
key="mcp_global_env_var",
exception_type="debug",
return_original_value=False,
)
if decrypted is None:
verbose_proxy_logger.warning(
"rotate_mcp_server_credentials_master_key: could not decrypt global env var %s, skipping",
entry.get("name"),
)
continue
entry["value"] = encrypt_value_helper(decrypted, new_encryption_key=new_encryption_key)
rotated = True
return rebuilt if rotated else None
def _prepare_mcp_server_data(
data: NewMCPServerRequest | UpdateMCPServerRequest,
exclude_unset: bool = False,
fields_set: set[str] | None = None,
) -> dict[str, Any]:
"""
Helper function to prepare MCP server data for database operations.
Handles JSON field serialization for mcp_info and env fields.
Args:
data: NewMCPServerRequest or UpdateMCPServerRequest object
exclude_unset: When True, only fields the caller explicitly provided are
included. Used for partial updates (PUT /v1/mcp/server) so omitted
fields keep their existing DB value instead of being silently reset
to a Pydantic schema default. ``exclude_none`` is not enough here:
non-Optional fields (e.g. ``transport=MCPTransport.sse``,
``mcp_access_groups=[]``, ``allow_all_keys=False``) are backfilled
with their default when omitted, and a non-None default survives the
``exclude_none`` filter and overwrites the row.
Returns:
Dict with properly serialized JSON fields
"""
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
# Convert model to dict.
# - Partial update (exclude_unset): only caller-provided keys are emitted, so
# omitted fields are never written and keep their existing DB value.
# - Create (exclude_none): drop None-valued fields and let DB defaults apply.
if exclude_unset:
if fields_set is None:
fields_set = data.fields_set()
data_dict = data.model_dump(exclude_unset=True)
# ``validate_and_normalize_mcp_server_payload`` always assigns ``alias``
# on the payload, which marks it as set even when the caller omitted it.
# Drop it only when the original request omitted alias; an explicit
# ``alias=None`` is a valid request to clear the stored alias.
if data_dict.get("alias") is None and "alias" not in fields_set:
data_dict.pop("alias", None)
# Prisma ``allowed_tools`` is a required String[]; ``null`` is invalid.
# The UI sends null to clear a whitelist — treat that as ``[]``.
if "allowed_tools" in data_dict and data_dict["allowed_tools"] is None:
data_dict["allowed_tools"] = []
# Json map fields use ``@default("{}")``; explicit null means clear overrides.
for json_map_field in (
"tool_name_to_display_name",
"tool_name_to_description",
):
if json_map_field in data_dict and data_dict[json_map_field] is None:
data_dict[json_map_field] = {}
else:
data_dict = data.model_dump(exclude_none=True)
# Ensure alias is always present in the dict (even if None)
if "alias" not in data_dict:
data_dict["alias"] = getattr(data, "alias", None)
# Handle credentials serialization
credentials: Final = data_dict.get("credentials")
if credentials is not None:
# Lift legacy blob-shaped token-exchange settings into their dedicated
# columns (an explicit top-level value wins, including an explicit
# null) and strip them from the blob so it never seeds the read-time
# fallback for rows written by current code.
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
blob_value = credentials.pop(te_field, None)
if blob_value is not None and te_field not in data_dict:
data_dict[te_field] = blob_value
data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=_get_salt_key())
data_dict["credentials"] = safe_dumps(data_dict["credentials"])
# Serialize JSON fields from ``data_dict`` (not ``data``) so the
# exclude_unset filter is respected. Reading back from ``data`` would
# reintroduce defaults (e.g. ``env={}``) for fields the caller never set.
if data_dict.get("static_headers") is not None:
data_dict["static_headers"] = encrypt_secret_map(data_dict["static_headers"])
# env_vars is read from ``data_dict`` (not ``data``) like every other JSON
# column so the exclude_unset filter is respected: a partial update that
# omits env_vars never overwrites the stored value. Global values are
# encrypted at rest before serialization.
env_vars: Final[Sequence[Mapping[str, str]] | None] = data_dict.get("env_vars")
if env_vars is not None:
serialized_env_vars: Final = [dict(v) for v in env_vars]
_encrypt_global_env_var_values(serialized_env_vars)
data_dict["env_vars"] = safe_dumps(serialized_env_vars)
if data_dict.get("mcp_info") is not None:
data_dict["mcp_info"] = safe_dumps(data_dict["mcp_info"])
if data_dict.get("env") is not None:
data_dict["env"] = encrypt_secret_map(data_dict["env"])
if "tool_name_to_display_name" in data_dict:
data_dict["tool_name_to_display_name"] = safe_dumps(data_dict["tool_name_to_display_name"] or {})
if "tool_name_to_description" in data_dict:
data_dict["tool_name_to_description"] = safe_dumps(data_dict["tool_name_to_description"] or {})
# mcp_access_groups is already List[str], no serialization needed
# On create, force is_byok so a False value is always written to the DB. On
# partial update, only write it when the caller explicitly provided it.
if not exclude_unset:
data_dict["is_byok"] = getattr(data, "is_byok", False)
return data_dict
def encrypt_credentials(credentials: MCPCredentials, encryption_key: str | None) -> MCPCredentials:
auth_value: Final = credentials.get("auth_value")
if auth_value is not None:
credentials["auth_value"] = encrypt_value_helper(
value=auth_value,
new_encryption_key=encryption_key,
)
client_id: Final = credentials.get("client_id")
if client_id is not None:
credentials["client_id"] = encrypt_value_helper(
value=client_id,
new_encryption_key=encryption_key,
)
client_secret: Final = credentials.get("client_secret")
if client_secret is not None:
credentials["client_secret"] = encrypt_value_helper(
value=client_secret,
new_encryption_key=encryption_key,
)
client_private_key: Final = credentials.get("client_private_key")
if client_private_key is not None:
credentials["client_private_key"] = encrypt_value_helper(
value=client_private_key,
new_encryption_key=encryption_key,
)
# AWS SigV4 credential fields
aws_access_key_id: Final = credentials.get("aws_access_key_id")
if aws_access_key_id is not None:
credentials["aws_access_key_id"] = encrypt_value_helper(
value=aws_access_key_id,
new_encryption_key=encryption_key,
)
aws_secret_access_key: Final = credentials.get("aws_secret_access_key")
if aws_secret_access_key is not None:
credentials["aws_secret_access_key"] = encrypt_value_helper(
value=aws_secret_access_key,
new_encryption_key=encryption_key,
)
aws_session_token: Final = credentials.get("aws_session_token")
if aws_session_token is not None:
credentials["aws_session_token"] = encrypt_value_helper(
value=aws_session_token,
new_encryption_key=encryption_key,
)
# aws_region_name and aws_service_name are NOT secrets — stored as-is
return credentials
def _credentials_blob_to_mutable_dict(blob: str | Mapping[str, object]) -> dict[str, object]:
parsed_blob: Final[dict[str, object]] = json.loads(blob) if isinstance(blob, str) else dict(blob)
return parsed_blob
def _mcp_server_table_actions(
prisma_client: PrismaClient,
) -> "TableActions[prisma_db_models.LiteLLM_MCPServerTable]":
table: Final[TableActions[prisma_db_models.LiteLLM_MCPServerTable]] = MCPServerRepository(prisma_client).table
return table
def _verification_token_table_actions(
prisma_client: PrismaClient,
) -> "TableActions[prisma_db_models.LiteLLM_VerificationToken]":
table: Final[TableActions[prisma_db_models.LiteLLM_VerificationToken]] = VerificationTokenRepository(
prisma_client
).table
return table
def _team_table_actions(
prisma_client: PrismaClient,
) -> "TableActions[prisma_db_models.LiteLLM_TeamTable]":
table: Final[TableActions[prisma_db_models.LiteLLM_TeamTable]] = TeamRepository(prisma_client).table
return table
def _oauth_client_table_actions(
prisma_client: PrismaClient,
) -> "TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]":
table: Final[TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = MCPServerOAuthClientRepository(
prisma_client
).table
return table
def _db_transaction_manager(prisma_client: PrismaClient) -> _UserEnvVarsTransaction:
manager: Final[_UserEnvVarsTransaction] = prisma_client.db.tx()
return manager
def _identifier_where(value: str, exclude_server_id: str | None) -> "prisma_db_types.LiteLLM_MCPServerTableWhereInput":
own_row_guard: Final = (
({"NOT": [{"server_id": exclude_server_id}]},) # mutable-ok: prisma where-inputs must be plain dicts
if exclude_server_id is not None
else ()
)
where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = {
"AND": [ # mutable-ok: prisma where-inputs must be plain dicts
{
"OR": [ # mutable-ok: prisma where-inputs must be plain dicts
{"server_name": {"equals": value, "mode": "insensitive"}},
{"alias": {"equals": value, "mode": "insensitive"}},
]
},
{
"OR": [{"approval_status": None}, {"approval_status": {"not": MCPApprovalStatus.draft}}]
}, # mutable-ok: prisma where-inputs must be plain dicts
*own_row_guard,
]
}
return where
def _identifier_field(data_dict: "Mapping[str, object]", field: str) -> str | None:
value: Final = data_dict.get(field)
return value if isinstance(value, str) else None
async def _find_mcp_server_identifier_conflict(
table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]",
*,
server_name: str | None,
alias: str | None,
exclude_server_id: str | None,
) -> McpIdentifierConflict | None:
"""Return the collision between an incoming identifier and a stored row, else None.
Each non-empty incoming identifier is compared case-insensitively against
BOTH the ``server_name`` and ``alias`` columns, because a value that matches
either column would still share the tool prefix another server answers to.
``alias`` is checked first so the reported field is deterministic. Draft
rows back the transient OAuth session flow and never reach the registry, so
they cannot collide. NULL ``approval_status`` predates the approval
workflow and is kept via the inner OR, matching ``get_all_mcp_servers``.
"""
candidates: Final[tuple[tuple[Literal["alias", "server_name"], str | None], ...]] = (
("alias", alias),
("server_name", server_name),
)
for field_name, value in candidates:
if not value:
continue
if (row := await table.find_first(where=_identifier_where(value, exclude_server_id))) is not None:
return McpIdentifierConflict(field=field_name, value=value, server_id=row.server_id)
return None
async def find_mcp_server_identifier_conflict(
prisma_client: PrismaClient,
*,
server_name: str | None,
alias: str | None,
exclude_server_id: str | None,
) -> McpIdentifierConflict | None:
"""Unlocked identifier-collision check, for callers outside a write path."""
return await _find_mcp_server_identifier_conflict(
_mcp_server_table_actions(prisma_client),
server_name=server_name,
alias=alias,
exclude_server_id=exclude_server_id,
)
def _mcp_identifier_lock_keys(*identifiers: str | None) -> tuple[int, ...]:
"""Deterministic advisory-lock keys for the lowercased identifiers, sorted
so concurrent requests for the same pair always lock in the same order."""
return tuple(
int.from_bytes(
hashlib.blake2b(f"mcp_identifier:{normalized}".encode(), digest_size=8).digest(),
"big",
signed=True,
)
for normalized in sorted(frozenset(value.lower() for value in identifiers if value))
)
async def _mcp_server_write_if_identifier_free(
prisma_client: PrismaClient,
*,
server_name: str | None,
alias: str | None,
exclude_server_id: str | None,
write: "Callable[[TableActions[prisma_db_models.LiteLLM_MCPServerTable]], Awaitable[prisma_db_models.LiteLLM_MCPServerTable | None]]",
) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None":
"""Run ``write`` only when no other live row owns ``server_name``/``alias``.
The conflict check and the write share a transaction guarded by per-identifier
advisory locks, so two concurrent requests for the same name cannot both
pass the check and both insert.
"""
lock_keys: Final = _mcp_identifier_lock_keys(server_name, alias)
async with _db_transaction_manager(prisma_client) as tx:
for lock_key in lock_keys:
await tx.execute_raw("SELECT pg_advisory_xact_lock($1::bigint)", lock_key)
conflict: Final = await _find_mcp_server_identifier_conflict(
tx.litellm_mcpservertable,
server_name=server_name,
alias=alias,
exclude_server_id=exclude_server_id,
)
if conflict is not None:
return conflict
return await write(tx.litellm_mcpservertable)
async def _db_find_mcp_server_rows(
prisma_client: PrismaClient,
where: "prisma_db_types.LiteLLM_MCPServerTableWhereInput | None" = None,
) -> "Sequence[prisma_db_models.LiteLLM_MCPServerTable]":
return await _mcp_server_table_actions(prisma_client).find_many(where=where)
async def _db_find_mcp_server_row(
prisma_client: PrismaClient, server_id: str
) -> "prisma_db_models.LiteLLM_MCPServerTable | None":
return await _mcp_server_table_actions(prisma_client).find_unique(where={"server_id": server_id})
async def _db_update_mcp_server_row(
prisma_client: PrismaClient,
server_id: str,
data: "prisma_db_types.LiteLLM_MCPServerTableUpdateInput",
) -> "prisma_db_models.LiteLLM_MCPServerTable":
row: Final[prisma_db_models.LiteLLM_MCPServerTable | None] = await _mcp_server_table_actions(prisma_client).update(
where={"server_id": server_id},
data=data,
)
if row is None:
raise ValueError(f"MCP server not found, passed server_id={server_id}")
return row
def _user_credential_actions(
prisma_client: PrismaClient,
) -> "TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]":
table: Final[TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]] = MCPUserCredentialsRepository(
prisma_client
).table
return table
class _MCPUserEnvVarsRepository(PrismaTableRepository["prisma_db_models.LiteLLM_MCPUserEnvVars"]):
table_name = "litellm_mcpuserenvvars"
def _user_env_var_actions(
prisma_client: PrismaClient,
) -> "TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]":
return _MCPUserEnvVarsRepository(prisma_client).table
async def _db_find_user_credential_row(
prisma_client: PrismaClient, user_id: str, server_id: str
) -> "prisma_db_models.LiteLLM_MCPUserCredentials | None":
return await _user_credential_actions(prisma_client).find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
async def _db_find_user_credential_rows(
prisma_client: PrismaClient,
where: "prisma_db_types.LiteLLM_MCPUserCredentialsWhereInput | None" = None,
) -> "Sequence[prisma_db_models.LiteLLM_MCPUserCredentials]":
return await _user_credential_actions(prisma_client).find_many(where=where)
async def _db_upsert_user_credential_row(
prisma_client: PrismaClient, user_id: str, server_id: str, credential_b64: str
) -> None:
await _user_credential_actions(prisma_client).upsert(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
data={
"create": {
"user_id": user_id,
"server_id": server_id,
"credential_b64": credential_b64,
},
"update": {"credential_b64": credential_b64},
},
)
async def _db_find_user_env_var_rows(
prisma_client: PrismaClient,
where: "prisma_db_types.LiteLLM_MCPUserEnvVarsWhereInput | None" = None,
) -> "Sequence[prisma_db_models.LiteLLM_MCPUserEnvVars]":
return await _user_env_var_actions(prisma_client).find_many(where=where)
def decrypt_credentials(
credentials: MCPCredentials,
) -> MCPCredentials:
"""Decrypt all secret fields in an MCPCredentials dict using the global salt key."""
secret_fields: Final = [
"auth_value",
"client_id",
"client_secret",
"client_private_key",
"aws_access_key_id",
"aws_secret_access_key",
"aws_session_token",
]
for field in secret_fields:
value = credentials.get(field)
if value is not None and isinstance(value, str):
credentials[field] = decrypt_value_helper(
value=value,
key=field,
exception_type="debug",
return_original_value=True,
)
return credentials
def _readable_mcp_servers(
rows: Iterable["prisma_db_models.LiteLLM_MCPServerTable"],
) -> Iterable[LiteLLM_MCPServerTable]:
for row in rows:
try:
table = LiteLLM_MCPServerTable.model_validate(row.model_dump())
except SecretMapDecodeError:
verbose_proxy_logger.warning("Skipping MCP server %s: cannot decrypt secret map", row.server_id)
continue
decrypt_global_env_var_values(table.env_vars)
yield table
async def get_all_mcp_servers(
prisma_client: PrismaClient,
approval_status: str | None = None,
) -> list[LiteLLM_MCPServerTable]:
"""
Returns mcp servers from the db, optionally filtered by approval_status.
Pass approval_status=None to return every server except drafts, which back the admin OAuth
session flow, are addressable only by their own server_id, and must never appear in a listing.
NULL approval_status predates the approval workflow, so those rows are kept explicitly rather
than dropped by a bare inequality, which SQL evaluates as NULL and would silently hide them.
"""
where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = (
{"approval_status": approval_status}
if approval_status is not None
else {"OR": [{"approval_status": None}, {"approval_status": {"not": MCPApprovalStatus.draft}}]}
)
mcp_servers: Final = await _db_find_mcp_server_rows(prisma_client, where)
return list(_readable_mcp_servers(mcp_servers))
async def get_mcp_server(prisma_client: PrismaClient, server_id: str) -> LiteLLM_MCPServerTable | None:
"""
Returns the matching mcp server from the db iff exists
"""
mcp_server: Final = await _db_find_mcp_server_row(prisma_client, server_id)
if mcp_server is None:
return None
table: Final = LiteLLM_MCPServerTable.model_validate(mcp_server.model_dump())
decrypt_global_env_var_values(table.env_vars)
return table
async def get_mcp_servers(prisma_client: PrismaClient, server_ids: Iterable[str]) -> list[LiteLLM_MCPServerTable]:
"""
Returns the matching mcp servers from the db with the server_ids
"""
_mcp_servers: Final[Sequence[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions(
prisma_client
).find_many(
where={
"server_id": {"in": server_ids},
}
)
return list(_readable_mcp_servers(_mcp_servers))
async def get_mcp_servers_by_verificationtoken(prisma_client: PrismaClient, token: str) -> list[str]:
"""
Returns the mcp servers from the db for the verification token
"""
verification_token_record: (
prisma_db_models.LiteLLM_VerificationToken | None
) = await _verification_token_table_actions(prisma_client).find_unique(
where={
"token": token,
},
include={
"object_permission": True,
},
)
mcp_servers: list[str] | None = []
if verification_token_record is not None and verification_token_record.object_permission is not None:
mcp_servers = verification_token_record.object_permission.mcp_servers
return mcp_servers or []
async def get_mcp_servers_by_team(prisma_client: PrismaClient, team_id: str) -> list[str]:
"""
Returns the mcp servers from the db for the team id
"""
team_record: prisma_db_models.LiteLLM_TeamTable | None = await _team_table_actions(prisma_client).find_unique(
where={
"team_id": team_id,
},
include={
"object_permission": True,
},
)
mcp_servers: list[str] | None = []
if team_record is not None and team_record.object_permission is not None:
mcp_servers = team_record.object_permission.mcp_servers
return mcp_servers or []
async def get_objectpermissions_for_mcp_server(
prisma_client: PrismaClient, mcp_server_id: str
) -> "Sequence[prisma_db_models.LiteLLM_ObjectPermissionTable]":
"""
Get all the object permissions records and the associated team and verficiationtoken records that have access to the mcp server
"""
object_permission_records: Final[
Sequence[prisma_db_models.LiteLLM_ObjectPermissionTable]
] = await ObjectPermissionRepository(prisma_client).table.find_many(
where={
"mcp_servers": {"has": mcp_server_id},
},
include={
"teams": True,
"verification_tokens": True,
},
)
return object_permission_records
async def get_virtualkeys_for_mcp_server(
prisma_client: PrismaClient, server_id: str
) -> "Sequence[prisma_db_models.LiteLLM_VerificationToken]":
"""
Get all the virtual keys that have access to the mcp server
"""
virtual_keys: Final[
Sequence[prisma_db_models.LiteLLM_VerificationToken] | None
] = await VerificationTokenRepository(prisma_client).table.find_many(
where={
"mcp_servers": {"has": server_id},
},
)
if virtual_keys is None: # pyright: ignore[reportUnnecessaryComparison] # unreachable per seam types; kept as-is
return []
return virtual_keys
async def delete_mcp_server_from_team(prisma_client: PrismaClient, server_id: str):
"""
Remove the mcp server from the team
"""
async def delete_mcp_server_from_virtualkey():
"""
Remove the mcp server from the virtual key
"""
async def delete_mcp_server(
prisma_client: PrismaClient,
server_id: str,
invalidate_token_cache: Callable[[str, str], Awaitable[None]] | None = None,
) -> LiteLLM_MCPServerTable | None:
"""
Delete the mcp server from the db by server_id
The server-row delete is the commit point. Per-user credential and env var
rows have no FK cascade, so they are cleaned up afterwards on a best-effort
basis: a transient failure there leaves only orphaned rows pointing at a
now-missing server and must not turn a successful delete into a
caller-visible error. Each table is cleaned independently so a failure on one
still attempts the other.
Each enumerated credential row's user also gets their cached per-user token
invalidated (legacy cache + v2 store, via invalidate_token_cache, defaulting
to the manager's shared invalidation): the caches are keyed by
(user_id, server_id), so without this a re-created server reusing the same
server_id would serve tokens minted for the deleted server until TTL.
Returns the deleted mcp server record if it exists, otherwise None
"""
deleted_server: Final = await MCPServerRepository(prisma_client).table.delete(
where={
"server_id": server_id,
},
)
if deleted_server is not None:
credential_user_ids: list[str] = []
try:
credential_rows: Sequence[prisma_db_models.LiteLLM_MCPUserCredentials] = await _user_credential_actions(
prisma_client
).find_many(where={"server_id": server_id})
credential_user_ids = [row.user_id for row in credential_rows]
except Exception as e: # noqa: BLE001 - enumeration is best-effort; cached tokens expire by TTL
verbose_proxy_logger.warning(
"MCP server %s deleted but per-user credential enumeration failed; cached tokens expire by TTL: %s",
server_id,
e,
)
for model, label in (
(_user_credential_actions(prisma_client), "credential"),
(_user_env_var_actions(prisma_client), "env var"),
(_oauth_client_table_actions(prisma_client), "OAuth client"),
):
try:
await model.delete_many(where={"server_id": server_id})
except Exception as e:
verbose_proxy_logger.warning(
"MCP server %s deleted but per-user %s cleanup failed; "
"orphaned rows can be removed on a later delete: %s",
server_id,
label,
e,
)
if credential_user_ids:
if invalidate_token_cache is None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
for user_id in credential_user_ids:
await invalidate_token_cache(user_id, server_id)
return deleted_server # pyright: ignore[reportReturnType] # prisma row, not domain LiteLLM_MCPServerTable
async def create_mcp_server(
prisma_client: PrismaClient, data: NewMCPServerRequest, touched_by: str
) -> LiteLLM_MCPServerTable:
"""
Create a new mcp server record in the db
"""
if data.server_id is None:
data.server_id = str(uuid.uuid4())
# Use helper to prepare data with proper JSON serialization
data_dict: Final = _prepare_mcp_server_data(data)
# Add audit fields
data_dict["created_by"] = touched_by
data_dict["updated_by"] = touched_by
new_mcp_server: Final = await MCPServerRepository(prisma_client).table.create(data=data_dict)
_decrypt_env_vars_on_returned_row(new_mcp_server)
return LiteLLM_MCPServerTable.model_validate(new_mcp_server.model_dump())
async def create_mcp_server_if_identifier_free(
prisma_client: PrismaClient, data: NewMCPServerRequest, touched_by: str
) -> LiteLLM_MCPServerTable | McpIdentifierConflict:
"""Create the row only when no other live server owns ``server_name``/``alias``.
Returns the McpIdentifierConflict instead of inserting when the collision
check finds an existing row; the advisory-lock transaction keeps two
concurrent creates of the same identifier from both passing.
"""
if data.server_id is None:
data.server_id = str(uuid.uuid4())
data_dict: Final = _prepare_mcp_server_data(data)
data_dict["created_by"] = touched_by
data_dict["updated_by"] = touched_by
async def _create(
table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]",
) -> "prisma_db_models.LiteLLM_MCPServerTable | None":
return await table.create(data=data_dict)
written: Final = await _mcp_server_write_if_identifier_free(
prisma_client,
server_name=_identifier_field(data_dict, "server_name"),
alias=_identifier_field(data_dict, "alias"),
exclude_server_id=None,
write=_create,
)
if isinstance(written, McpIdentifierConflict):
return written
if written is None:
raise RuntimeError("inserted MCP server row missing")
_decrypt_env_vars_on_returned_row(written)
return LiteLLM_MCPServerTable.model_validate(written.model_dump())
async def create_draft_mcp_server(
prisma_client: PrismaClient,
data: NewMCPServerRequest,
touched_by: str,
ttl_seconds: int,
server_id: str | None = None,
) -> LiteLLM_MCPServerTable:
"""
Persist a short-lived draft row backing the admin OAuth "Authorize & Fetch Token" flow.
The draft lives in the database rather than in process memory so that the /register,
/authorize and /token legs resolve it whichever worker or replica accepts each request.
Writing is strictly create-if-absent. Any existing row for the id is returned untouched, which
covers both a live draft for this same session and a real server the edit form is
re-authorizing against its own id, where writing a draft would collide on the primary key.
Each click of Authorize mints a fresh id, so nothing is lost by never overwriting, and it is
what makes concurrent callers sharing one id safe rather than mutually destructive.
"""
draft_id: Final = server_id or data.server_id or str(uuid.uuid4())
await _prune_expired_draft_mcp_servers(prisma_client, ttl_seconds)
existing: Final = await _db_find_mcp_server_row(prisma_client, draft_id)
if existing is not None:
# Already usable by every worker, whether it is a live draft for this same session or a
# real server the edit form is re-authorizing. Either way there is nothing to write, and
# not writing is what keeps concurrent callers for one server_id from racing each other.
return LiteLLM_MCPServerTable.model_validate(existing.model_dump())
draft_payload: Final = data.model_copy(update={"server_id": draft_id, "approval_status": MCPApprovalStatus.draft})
try:
return await create_mcp_server(prisma_client, draft_payload, touched_by)
except Exception:
# Lost the create race: the read above and this create are two statements, not one. The
# winner wrote a draft for this same session, so adopt it rather than failing a caller
# whose session is in fact ready. Anything else still raises.
raced: Final = await _db_find_mcp_server_row(prisma_client, draft_id)
if raced is None or raced.approval_status != MCPApprovalStatus.draft:
raise
return LiteLLM_MCPServerTable.model_validate(raced.model_dump())
async def _prune_expired_draft_mcp_servers(prisma_client: PrismaClient, ttl_seconds: int) -> None:
"""Drop drafts already past ``ttl_seconds``, so abandoned OAuth sessions do not accumulate.
Runs on each draft write rather than on a schedule, mirroring the in-memory cache this
replaces, which pruned on every store. Expired drafts are unreadable by then anyway, so the
only thing at stake is row count, and the work is bounded by how often admins authorize.
"""
cutoff: Final = datetime.now(timezone.utc) - timedelta(seconds=max(1, ttl_seconds))
# Age is filtered here rather than in the query: the draft set is bounded by how many OAuth
# authorizations are in flight, so it is a handful of rows even on a busy proxy.
drafts: Final = await _db_find_mcp_server_rows(
prisma_client,
where={"approval_status": MCPApprovalStatus.draft},
)
for row in drafts:
# A row without a timestamp has no age to judge, so leave it rather than guess it is stale.
# Two workers sweeping the same row is harmless: prisma's delete returns None for a row
# that is already gone rather than raising, so the loser of that race is a no-op.
if row.updated_at is not None and row.updated_at < cutoff:
await delete_mcp_server(prisma_client, row.server_id)
async def get_draft_mcp_server(
prisma_client: PrismaClient, server_id: str, ttl_seconds: int
) -> LiteLLM_MCPServerTable | None:
"""
Return the draft row for ``server_id`` if it has not yet aged past ``ttl_seconds``, else None.
Age is enforced in the query rather than by a sweeper so an expired draft is unreadable the
moment it lapses, regardless of which process last ran a cleanup.
"""
cutoff: Final = datetime.now(timezone.utc) - timedelta(seconds=max(1, ttl_seconds))
draft_rows: Final = await _db_find_mcp_server_rows(
prisma_client,
where={
"server_id": server_id,
"approval_status": MCPApprovalStatus.draft,
"updated_at": {"gte": cutoff},
},
)
if not draft_rows:
return None
table: Final = LiteLLM_MCPServerTable.model_validate(draft_rows[0].model_dump())
decrypt_global_env_var_values(table.env_vars)
return table
async def _update_mcp_server_row(
prisma_client: PrismaClient,
*,
server_id: str,
data_dict: Mapping[str, object],
) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None":
identifier_write: Final = any(field in data_dict for field in ("server_name", "alias"))
async def _update(
table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]",
) -> "prisma_db_models.LiteLLM_MCPServerTable | None":
return await table.update(
where={"server_id": server_id}, # mutable-ok: prisma where-inputs must be plain dicts
data=data_dict,
)
if not identifier_write:
return await _update(_mcp_server_table_actions(prisma_client))
if "alias" in data_dict and not data_dict["alias"] and "server_name" not in data_dict:
# Clearing the alias drops the prefix to the stored server_name, which
# may already belong to another row, so that name needs the check too.
existing: Final = await _db_find_mcp_server_row(prisma_client, server_id)
if existing is None:
return await _update(_mcp_server_table_actions(prisma_client))
return await _mcp_server_write_if_identifier_free(
prisma_client,
server_name=existing.server_name,
alias=None,
exclude_server_id=server_id,
write=_update,
)
return await _mcp_server_write_if_identifier_free(
prisma_client,
server_name=_identifier_field(data_dict, "server_name"),
alias=_identifier_field(data_dict, "alias"),
exclude_server_id=server_id,
write=_update,
)
async def update_mcp_server(
prisma_client: PrismaClient,
data: UpdateMCPServerRequest,
touched_by: str,
fields_set: set[str] | None = None,
) -> LiteLLM_MCPServerTable | McpIdentifierConflict | None:
"""
Update a new mcp server record in the db
Returns McpIdentifierConflict instead of writing when the update would put
``server_name``/``alias`` onto identifiers another live row already owns.
"""
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
# Use helper to prepare data with proper JSON serialization.
# exclude_unset=True makes this a true partial update: fields the caller did
# not provide are not written, so they keep their existing DB value instead
# of being reset to a schema default (transport=sse, allow_all_keys=False...).
data_dict: Final = _prepare_mcp_server_data(data, exclude_unset=True, fields_set=fields_set)
# Pre-fetch existing record once if we need it for auth_type, url, or credential logic
existing = None
has_credentials: Final = "credentials" in data_dict and data_dict["credentials"] is not None
# An explicit token-exchange column write (set or clear) also migrates the
# legacy blob copies below, so the existing row is needed for those updates.
explicit_te_write: Final = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & data_dict.keys())
url_provided: Final = "url" in data_dict and data_dict["url"] is not None
issuer_provided: Final = "issuer" in data_dict
if data.auth_type or has_credentials or explicit_te_write or url_provided or issuer_provided:
existing = await _db_find_mcp_server_row(prisma_client, data.server_id)
auth_type_changed: Final = bool(
data.auth_type
and existing
and _credential_auth_class(existing.auth_type) != _credential_auth_class(data.auth_type)
)
# A url change re-points the server at a potentially different upstream, so any discovered or
# trust-on-first-use OAuth endpoints/issuer belong to the old upstream and must re-discover.
url_changed: Final = bool(url_provided and existing and existing.url != data_dict["url"])
old_issuer: Final = _blank_to_none(getattr(existing, "issuer", None)) if existing else None
issuer_changed: Final = bool(
issuer_provided and old_issuer is not None and _blank_to_none(data_dict.get("issuer")) != old_issuer
)
# Clear stale credentials when auth_type changes but no new credentials provided
if auth_type_changed and "credentials" not in data_dict:
data_dict["credentials"] = None
if auth_type_changed or url_changed or issuer_changed:
# Clear each auth-flow-scoped field that the caller either omitted (partial update) or
# resubmitted unchanged. The edit form re-sends every field, so a stale issuer/endpoint
# belonging to the old upstream would otherwise survive a url/auth_type change and win in the
# resolution merge; only a genuinely new submitted value is kept.
data_dict.update(
{
field: None
for field in _AUTH_FLOW_SCOPED_FIELDS
if field not in data_dict or data_dict[field] == getattr(existing, field, None)
}
)
# An explicit column write that does not touch credentials must still migrate
# the row's legacy blob copies: lift values for columns the caller left
# untouched, strip every copy from the blob. Without this, clearing a column
# (e.g. to re-enable RFC 9728/8414 discovery) would leave the blob copy in
# place, and the next credentials update's migrate-on-write would silently
# repopulate the column the admin just cleared. (When credentials ARE in the
# update, the merge below performs the same migration.)
if explicit_te_write and "credentials" not in data_dict and existing is not None and existing.credentials:
existing_creds = _credentials_blob_to_mutable_dict(existing.credentials)
if _TOKEN_EXCHANGE_COLUMN_FIELDS & existing_creds.keys():
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
legacy_value = existing_creds.pop(te_field, None)
if legacy_value is not None and te_field not in data_dict and getattr(existing, te_field, None) is None:
data_dict[te_field] = legacy_value
data_dict["credentials"] = safe_dumps(existing_creds)
# Merge credentials: preserve existing fields not present in the update.
# Without this, a partial credential update (e.g. changing only region)
# would wipe encrypted secrets that the UI cannot display back.
if "credentials" in data_dict and data_dict["credentials"] is not None:
if existing and existing.credentials:
# Only merge when the credential CLASS is unchanged. A cross-class switch
# (e.g. oauth2 → api_key, or oauth2 → true_passthrough) replaces credentials
# entirely to avoid stale secrets from the previous class lingering; a switch
# within the client-forwarded class (true_passthrough ↔ oauth_delegate) keeps
# the same declared app and so must merge, not replace.
if not auth_type_changed:
existing_creds = _credentials_blob_to_mutable_dict(existing.credentials)
new_creds: Final = _credentials_blob_to_mutable_dict(data_dict["credentials"])
# New values override existing; existing keys not in update are preserved. A client
# rotation additionally drops the previous app's stale minted token keys.
merged: Final = _drop_stale_minted_on_client_rotation({**existing_creds, **new_creds}, new_creds)
# Migrate-on-write for legacy rows: token-exchange settings the
# old blob shape carried move to their dedicated columns (unless
# the caller set the column this update, or the row already has
# one) and are never re-persisted in the blob. Stored plaintext,
# so the merged value lifts as-is.
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
legacy_value = merged.pop(te_field, None)
if (
legacy_value is not None
and te_field not in data_dict
and getattr(existing, te_field, None) is None
):
data_dict[te_field] = legacy_value
data_dict["credentials"] = safe_dumps(merged)
# Add audit fields
data_dict["updated_by"] = touched_by
# prisma-python rejects a raw ``None`` for a ``Json?`` field ("value is required but not set"); the
# clear paths above use ``None`` as the merge-skip sentinel, so translate it here to ``Json(None)``,
# which writes SQL null and reads back as ``None``. Done at the edge so the merge guards stay simple.
if "credentials" in data_dict and data_dict["credentials"] is None:
from prisma import Json # noqa: PLC0415 # local import: prisma may be ungenerated at module load in some tools
data_dict["credentials"] = Json(None)
updated_mcp_server: Final = await _update_mcp_server_row(
prisma_client,
server_id=data.server_id,
data_dict=data_dict,
)
if isinstance(updated_mcp_server, McpIdentifierConflict):
return updated_mcp_server
_decrypt_env_vars_on_returned_row(updated_mcp_server)
return LiteLLM_MCPServerTable.model_validate(updated_mcp_server.model_dump()) if updated_mcp_server else None
async def get_mcp_server_oauth_client_credentials(prisma_client: PrismaClient, server_id: str) -> object | None:
"""Read the persisted (encrypted) DCR OAuth client blob for a server from the
server-scoped store, or None. Config.yaml-declared servers have no
LiteLLM_MCPServerTable row, so their dynamically registered client lives here keyed
by server_id. The returned value is the raw credentials blob for
``_get_persisted_dcr_credentials`` to parse."""
row: Final[prisma_db_models.LiteLLM_MCPServerOAuthClient | None] = await _oauth_client_table_actions(
prisma_client
).find_unique(where={"server_id": server_id})
if row is None:
return None
return row.credentials
async def upsert_mcp_server_oauth_client_credentials(
prisma_client: PrismaClient, server_id: str, credentials: MCPCredentials
) -> None:
"""Persist a server's dynamically registered OAuth client (RFC 7591 DCR) in the
server-scoped store keyed by server_id, independent of any LiteLLM_MCPServerTable row.
client_id/client_secret are encrypted at rest with the same salt key used for the
server row's credentials blob, so ``_apply_persisted_dcr_credentials`` decrypts them the
same way regardless of which store a server's client came from."""
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
encrypted: Final = encrypt_credentials(credentials=MCPCredentials(**credentials), encryption_key=_get_salt_key())
blob: Final = safe_dumps(encrypted)
await _oauth_client_table_actions(prisma_client).upsert(
where={"server_id": server_id},
data={
"create": {"server_id": server_id, "credentials": blob},
"update": {"credentials": blob},
},
)
def _reencrypt_mcp_credentials_blob(
credentials: "str | Mapping[str, object] | None", new_master_key: str
) -> str | None:
"""Decrypt an at-rest MCP credentials blob with the current key and re-encrypt it under
new_master_key, returning the serialized blob or None when there is nothing to rotate. Shared by
every table that stores an encrypted MCP credentials blob so a master-key rotation covers them
uniformly and cannot silently skip one."""
if not credentials:
return None
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps # noqa: PLC0415 # avoids circular import
creds_dict: Final = _credentials_blob_to_mutable_dict(credentials)
decrypted: Final = decrypt_credentials(credentials=cast(MCPCredentials, creds_dict))
encrypted: Final = encrypt_credentials(credentials=decrypted, encryption_key=new_master_key)
return safe_dumps(encrypted)
async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient, touched_by: str, new_master_key: str):
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps # noqa: PLC0415 # avoids circular import
mcp_servers: Final = await _db_find_mcp_server_rows(prisma_client)
updated = 0
for mcp_server in mcp_servers:
update_data: dict[str, str] = {}
rotated_credentials = _reencrypt_mcp_credentials_blob(mcp_server.credentials, new_master_key)
if rotated_credentials is not None:
update_data["credentials"] = rotated_credentials
rotated_env_vars = _reencrypt_global_env_var_values(mcp_server.env_vars, new_master_key)
if rotated_env_vars is not None:
update_data["env_vars"] = safe_dumps(rotated_env_vars)
for field in ("static_headers", "env"):
try:
if secret_map := decode_secret_map(getattr(mcp_server, field, None), key=field):
update_data[field] = encrypt_secret_map(secret_map, new_encryption_key=new_master_key)
except SecretMapDecodeError:
verbose_proxy_logger.warning("Cannot rotate MCP %s for server %s", field, mcp_server.server_id)
if not update_data:
continue
update_data["updated_by"] = touched_by
await _mcp_server_table_actions(prisma_client).update(
where={"server_id": mcp_server.server_id},
data=update_data,
)
updated += 1
oauth_clients: Final[Sequence[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = await _oauth_client_table_actions(
prisma_client
).find_many()
oauth_updated = 0
for oauth_client in oauth_clients:
rotated_credentials = _reencrypt_mcp_credentials_blob(oauth_client.credentials, new_master_key)
if rotated_credentials is None:
continue
await _oauth_client_table_actions(prisma_client).update(
where={"server_id": oauth_client.server_id},
data={"credentials": rotated_credentials},
)
oauth_updated += 1
verbose_proxy_logger.info(
"rotate_mcp_server_credentials_master_key: rotated %d MCP server row(s) and %d OAuth-client row(s)",
updated,
oauth_updated,
)
def _decode_user_credential(stored: str) -> str | None:
"""Read back a value persisted in ``LiteLLM_MCPUserCredentials.credential_b64``.
Tries nacl decryption first (current write format). Falls back to a
plain ``urlsafe_b64decode`` for rows persisted by older code that wrote
the credential without encryption. Returns ``None`` when neither path
yields a valid string.
"""
decrypted: Final = decrypt_value_helper(
value=stored,
key="mcp_user_credential",
exception_type="debug",
return_original_value=False,
)
if decrypted is not None:
return decrypted
try:
return base64.urlsafe_b64decode(stored).decode()
except (binascii.Error, UnicodeDecodeError, ValueError, TypeError):
return None
def _warn_undecryptable_credential(user_id: str, server_id: str) -> None:
"""Log the one credential state that otherwise reads as "user never authorized"."""
verbose_proxy_logger.warning(
"MCP user credential for user=%s server=%s could not be decrypted (likely written under a "
"previous LITELLM_SALT_KEY); the user is treated as not connected and must re-authorize.",
user_id,
server_id,
)
def _parse_oauth_payload(decoded: str | None) -> OAuthCredentialPayload | None:
"""Return the OAuth2 payload dict if ``decoded`` holds one, else ``None``.
A row is considered an OAuth2 credential iff its decoded value parses as
a JSON object with ``"type": "oauth2"``. Plain BYOK credentials (which
share the same column) decode to a non-JSON string and return ``None``.
Callers that need to tell an unreadable row from a readable non-OAuth2 one
pass the result of :func:`_decode_user_credential` so a single decode
answers both questions: ``None`` there means the value can be neither
decrypted nor base64-decoded, so no caller can ever recover it.
"""
if decoded is None:
return None
parsed: OAuthCredentialPayload | None
try:
parsed = json.loads(decoded)
except (ValueError, TypeError):
return None
if isinstance(parsed, dict) and parsed.get("type") == "oauth2":
return parsed
return None
def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None:
"""Return the OAuth2 payload dict held in ``stored``, else ``None``."""
return _parse_oauth_payload(_decode_user_credential(stored))
async def rotate_mcp_user_credentials_master_key(prisma_client: PrismaClient, new_master_key: str):
"""Re-encrypt every ``LiteLLM_MCPUserCredentials`` row with ``new_master_key``.
Reads each ``credential_b64`` with the current salt key (falling back to
legacy plain base64 for unmigrated rows) and writes it back encrypted
under the new master key. Rows that are unreadable under both paths
are logged and skipped so one corrupt row does not abort the rotation.
"""
rows: Final = await _db_find_user_credential_rows(prisma_client)
rotated = 0
skipped = 0
for row in rows:
plaintext = _decode_user_credential(row.credential_b64)
if plaintext is None:
verbose_proxy_logger.warning(
"rotate_mcp_user_credentials_master_key: could not decode "
"credential for user_id=%s server_id=%s, skipping",
row.user_id,
row.server_id,
)
skipped += 1
continue
re_encrypted = encrypt_value_helper(plaintext, new_encryption_key=new_master_key)
await _user_credential_actions(prisma_client).update(
where={
"user_id_server_id": {
"user_id": row.user_id,
"server_id": row.server_id,
}
},
data={"credential_b64": re_encrypted},
)
rotated += 1
verbose_proxy_logger.info(
"rotate_mcp_user_credentials_master_key: rotated %d row(s), skipped %d",
rotated,
skipped,
)
async def rotate_mcp_user_env_vars_master_key(prisma_client: PrismaClient, new_master_key: str):
"""Re-encrypt every ``LiteLLM_MCPUserEnvVars`` row with ``new_master_key``.
Reads each ``values_b64`` blob with the current salt key and writes it back
encrypted under the new master key. Rows that fail to decrypt are logged and
skipped so one corrupt row does not abort the rotation nor overwrite values
that may still be recoverable.
"""
rows: Final = await _db_find_user_env_var_rows(prisma_client)
rotated = 0
skipped = 0
for row in rows:
plaintext = decrypt_value_helper(
value=row.values_b64,
key="mcp_user_env_vars",
exception_type="debug",
return_original_value=False,
)
if plaintext is None:
verbose_proxy_logger.warning(
"rotate_mcp_user_env_vars_master_key: could not decrypt env vars for user_id=%s server_id=%s, skipping",
row.user_id,
row.server_id,
)
skipped += 1
continue
re_encrypted = encrypt_value_helper(plaintext, new_encryption_key=new_master_key)
await _user_env_var_actions(prisma_client).update(
where={
"user_id_server_id": {
"user_id": row.user_id,
"server_id": row.server_id,
}
},
data={"values_b64": re_encrypted},
)
rotated += 1
verbose_proxy_logger.info(
"rotate_mcp_user_env_vars_master_key: rotated %d row(s), skipped %d",
rotated,
skipped,
)
async def store_user_credential(
prisma_client: PrismaClient,
user_id: str,
server_id: str,
credential: str,
) -> None:
"""Store a user credential for a BYOK MCP server."""
encoded: Final = encrypt_value_helper(credential)
await _db_upsert_user_credential_row(prisma_client, user_id, server_id, encoded)
async def get_user_credential(
prisma_client: PrismaClient,
user_id: str,
server_id: str,
) -> str | None:
"""Return credential for a user+server pair, or None."""
row: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
if row is None:
return None
return _decode_user_credential(row.credential_b64)
async def has_user_credential(
prisma_client: PrismaClient,
user_id: str,
server_id: str,
) -> bool:
"""Return True if the user has a stored credential for this server."""
row: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
return row is not None
async def delete_user_credential(
prisma_client: PrismaClient,
user_id: str,
server_id: str,
) -> None:
"""Delete the user's stored credential for a BYOK MCP server."""
await _user_credential_actions(prisma_client).delete(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
# ── OAuth2 user-credential helpers ────────────────────────────────────────────
async def store_user_oauth_credential(
prisma_client: PrismaClient,
user_id: str,
server_id: str,
access_token: str,
refresh_token: str | None = None,
expires_in: int | None = None,
scopes: list[str] | None = None,
skip_byok_guard: bool = False,
identity_binding_proof: str | None = None,
) -> None:
"""Persist an OAuth2 access token for a user+server pair.
The payload is JSON-serialised and stored encrypted in the same
``credential_b64`` column used by BYOK. A ``"type": "oauth2"`` key
differentiates it from plain BYOK API keys.
"""
expires_at: str | None = None
if expires_in is not None:
expires_at = (datetime.now(timezone.utc) + timedelta(seconds=expires_in)).isoformat()
payload: Final[OAuthCredentialPayload] = {
"type": "oauth2",
"access_token": access_token,
"connected_at": datetime.now(timezone.utc).isoformat(),
**({"identity_binding_proof": identity_binding_proof} if identity_binding_proof else {}),
}
if refresh_token:
payload["refresh_token"] = refresh_token
if expires_at:
payload["expires_at"] = expires_at
if scopes:
payload["scopes"] = scopes
# Guard against silently overwriting a BYOK credential with an OAuth token.
# Skip the guard when the caller knows the row is already an OAuth2 credential
# (e.g. during token refresh), saving an extra DB round-trip.
if not skip_byok_guard:
existing: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
decoded: Final = _decode_user_credential(existing.credential_b64) if existing is not None else None
if existing is not None and _parse_oauth_payload(decoded) is None:
# Refuse only while the row still holds readable content, which is a live BYOK
# secret that overwriting would destroy. A row that does not decode was written
# under a different LITELLM_SALT_KEY, and one that decodes to nothing holds no
# secret at all; refusing either preserves nothing and instead wedges the user
# out of the OAuth flow for good, since re-authorizing is their only recovery.
if decoded:
raise ValueError(
f"Existing credential for user {user_id} and server "
f"{server_id} could not be verified as an OAuth2 token. "
f"Refusing to overwrite."
)
verbose_proxy_logger.warning(
"store_user_oauth_credential: existing credential for user=%s server=%s could not be "
"decrypted (likely written under a previous LITELLM_SALT_KEY); replacing it with the "
"newly authorized OAuth2 token.",
user_id,
server_id,
)
encoded: Final = encrypt_value_helper(json.dumps(payload))
await _db_upsert_user_credential_row(prisma_client, user_id, server_id, encoded)
def is_oauth_credential_expired(cred: OAuthCredentialPayload, buffer_seconds: int = 0) -> bool:
"""Return True if the OAuth2 credential's access_token has expired.
Checks the ``expires_at`` ISO-format string stored in the credential payload.
Returns False when ``expires_at`` is absent or unparseable (treat as non-expired).
With ``buffer_seconds`` > 0, a token that is still valid but expires within the
buffer is also treated as expired, so callers can refresh proactively instead of
handing back a token that may lapse mid-request.
"""
expires_at: Final = cred.get("expires_at")
if not expires_at:
return False
try:
exp_dt = datetime.fromisoformat(expires_at)
if exp_dt.tzinfo is None:
exp_dt = exp_dt.replace(tzinfo=timezone.utc)
return datetime.now(timezone.utc) + timedelta(seconds=buffer_seconds) > exp_dt
except (ValueError, TypeError):
return False
def oauth_grant_state(cred: OAuthCredentialPayload | None) -> OAuthGrantState:
"""Classify local grant readiness without attempting a refresh or checking upstream revocation."""
if not cred or not cred.get("access_token"):
return "absent"
if not is_oauth_credential_expired(cred, buffer_seconds=MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS):
return "valid"
return "refreshable" if cred.get("refresh_token") else "absent"
async def get_user_oauth_credential(
prisma_client: PrismaClient,
user_id: str,
server_id: str,
) -> OAuthCredentialPayload | None:
"""Return the decoded OAuth2 payload dict for a user+server pair, or None."""
row: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
if row is None:
return None
decoded: Final = _decode_user_credential(row.credential_b64)
if decoded is None:
_warn_undecryptable_credential(user_id, server_id)
return _parse_oauth_payload(decoded)
def _server_user_credential_item(
row: "prisma_db_models.LiteLLM_MCPUserCredentials",
) -> MCPServerUserCredentialListItem:
oauth_payload: Final = _decode_oauth_payload(row.credential_b64)
if oauth_payload is None:
return MCPServerUserCredentialListItem(
user_id=row.user_id,
credential_type="byok",
updated_at=row.updated_at.isoformat(),
)
return MCPServerUserCredentialListItem(
user_id=row.user_id,
credential_type="oauth2",
expires_at=oauth_payload.get("expires_at"),
connected_at=oauth_payload.get("connected_at"),
updated_at=row.updated_at.isoformat(),
)
async def list_server_user_credentials(
prisma_client: PrismaClient,
server_id: str,
) -> tuple[MCPServerUserCredentialListItem, ...]:
"""Every user's stored credential for one server, typed but without the secret, for admins."""
rows: Final = await _db_find_user_credential_rows(
prisma_client,
{"server_id": server_id}, # mutable-ok: prisma where-inputs must be plain dicts
)
return tuple(_server_user_credential_item(row) for row in rows)
async def list_user_oauth_credentials(
prisma_client: PrismaClient,
user_id: str,
) -> list[OAuthCredentialPayload]:
"""Return all OAuth2 credential payloads for a user, tagged with server_id."""
rows: Final = await _db_find_user_credential_rows(prisma_client, {"user_id": user_id})
results: Final[list[OAuthCredentialPayload]] = []
for row in rows:
decoded = _decode_user_credential(row.credential_b64)
if decoded is None:
_warn_undecryptable_credential(user_id, row.server_id)
payload = _parse_oauth_payload(decoded)
if payload is None:
continue
payload["server_id"] = row.server_id
results.append(payload)
return results
def _decrypted_credential_field(creds: dict[str, object], field: str) -> object:
"""Return one credential field decrypted with the global salt key; non-string and legacy
plaintext values come back unchanged (decrypt_value_helper returns the original on failure)."""
value: Final = creds.get(field)
if not isinstance(value, str):
return value
return decrypt_value_helper(
value=value,
key=field,
exception_type="debug",
return_original_value=True,
)
def mcp_oauth_token_identity(server: object) -> tuple[object, ...]:
"""The upstream-OAuth-token-determining fields of an MCP server: the resource/audience (url, or
spec_path for OpenAPI servers, plus the RFC 8707 upstream_resource sent on the authorize and
token legs), the OAuth mode/grant (auth_type, oauth2_flow), the authorization-server endpoints,
and the OAuth client + scopes. Mirrors the dashboard's getOAuthAuthorizationIdentity. When any
of these change on a server update, previously stored per-user tokens were minted for the old
identity and are stale. Excludes transport and delegate_auth_to_upstream, which do not affect
what token is minted (RFC 8693).
client_id/client_secret are compared decrypted: stored values are NaCl-encrypted with a fresh
nonce on every write, so comparing ciphertext would flag every routine save as an identity
change and purge tokens that are still valid."""
creds: Final = getattr(server, "credentials", None)
if isinstance(creds, str):
try:
parsed: dict[str, object] | None = json.loads(creds)
except ValueError:
parsed = None
else:
parsed = creds
creds_dict: Final[dict[str, object]] = parsed if isinstance(parsed, dict) else {}
return (
getattr(server, "url", None),
getattr(server, "spec_path", None),
getattr(server, "auth_type", None),
getattr(server, "oauth2_flow", None),
getattr(server, "issuer", None),
getattr(server, "authorization_url", None),
getattr(server, "token_url", None),
getattr(server, "registration_url", None),
_decrypted_credential_field(creds_dict, "client_id"),
_decrypted_credential_field(creds_dict, "client_secret"),
creds_dict.get("scopes"),
creds_dict.get("upstream_resource"),
)
async def purge_user_oauth_credentials_for_server(
prisma_client: PrismaClient,
server_id: str,
invalidate_token_cache: Callable[[str, str], Awaitable[None]] | None = None,
) -> int:
"""Delete every stored per-user OAuth token for a server and invalidate each user's cached
token everywhere it can be served from (the legacy per-user token cache and the v2 per-user OAuth
token store), so no user keeps a token minted for a superseded configuration. Called when a server
update changes a mint-relevant field (see mcp_oauth_token_identity). Returns the number of rows
removed.
LiteLLM_MCPUserCredentials also stores BYOK API keys in the same column; only rows whose payload
decodes as an OAuth2 credential (see _decode_oauth_payload) are deleted, because a config change
only invalidates minted tokens, never a user's own stored key. Rows are therefore deleted per
(user_id, server_id) pair rather than by a blanket server_id filter. An OAuth row inserted while
the purge runs for a user not yet enumerated survives; a re-auth completing in the window for an
already-enumerated user is deleted along with the stale row (the pair delete cannot tell them
apart), which costs that user one extra re-auth and nothing else.
invalidate_token_cache is injectable for tests; it defaults to the manager's shared
invalidate_user_oauth_token_cache, the single invalidation point for per-user tokens."""
rows: Final = await _db_find_user_credential_rows(prisma_client, {"server_id": server_id})
oauth_rows: Final = [row for row in rows if _decode_oauth_payload(row.credential_b64) is not None]
if not oauth_rows:
return 0
deleted_count: Final = await _user_credential_actions(prisma_client).delete_many(
where={"server_id": server_id, "user_id": {"in": [row.user_id for row in oauth_rows]}}
)
if invalidate_token_cache is None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
for row in oauth_rows:
await invalidate_token_cache(row.user_id, server_id)
if deleted_count != len(oauth_rows):
verbose_proxy_logger.warning(
"MCP server %s: purge removed %d OAuth credential row(s) but %d were enumerated; "
"row(s) were deleted concurrently during the purge",
server_id,
deleted_count,
len(oauth_rows),
)
return deleted_count
async def refresh_user_oauth_token(
prisma_client: PrismaClient,
user_id: str,
server: "MCPServer",
cred: OAuthCredentialPayload,
) -> OAuthCredentialPayload | None:
"""Attempt to refresh a per-user OAuth2 token using its stored refresh_token.
POSTs to ``server.effective_token_url`` with ``grant_type=refresh_token``.
On success: persists the new credential via ``store_user_oauth_credential``
and returns the updated payload dict.
On failure (network error, invalid_grant, missing refresh_token, …): logs a
warning and returns ``None`` — the caller is responsible for clearing the
stale credential and triggering re-authentication.
"""
binding: Final = server.oauth_identity_binding
if binding is not None and binding.mode == "enforce":
if not await credential_binding_matches(binding, user_id, server.server_id, cred):
return None
refresh_token: Final[str | None] = cred.get("refresh_token")
token_url: Final[str | None] = getattr(server, "effective_token_url", None) or getattr(server, "token_url", None)
server_id: Final[str] = getattr(server, "server_id", "")
client_id: Final[str | None] = getattr(server, "client_id", None)
client_secret: Final[str | None] = getattr(server, "client_secret", None)
if not refresh_token:
verbose_proxy_logger.debug(
"refresh_user_oauth_token: no refresh_token stored for user=%s server=%s",
user_id,
server_id,
)
return None
if not token_url:
verbose_proxy_logger.debug(
"refresh_user_oauth_token: server=%s has no token_url configured",
server_id,
)
return None
try:
token_request: Final = build_upstream_oauth2_token_request(
server,
auth_method=getattr(server, "token_endpoint_auth_method", None),
client_id=client_id,
client_secret=client_secret,
)
token_data: Final[dict[str, str]] = {
"grant_type": "refresh_token",
"refresh_token": refresh_token,
**token_request.body,
}
async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
response: Final = await async_client.post(
token_url,
headers={"Accept": "application/json", **token_request.headers},
data=token_data,
)
response.raise_for_status()
body: Final[_OAuthTokenRefreshResponse] = response.json()
except Exception as exc:
verbose_proxy_logger.warning(
"refresh_user_oauth_token: refresh request failed for user=%s server=%s: %s",
user_id,
server_id,
exc,
)
return None
try:
binding_proof: Final = await enforce_oauth_identity_binding(
server=server,
token_response=body,
litellm_user_id=user_id,
grant_type="refresh_token",
refresh_ownership=RefreshTokenPresented(refresh_token),
)
except HTTPException as exc:
if exc.status_code != 403:
raise
return None
access_token: Final[str | None] = body.get("access_token")
if not access_token:
verbose_proxy_logger.warning(
"refresh_user_oauth_token: token response missing access_token for user=%s server=%s",
user_id,
server_id,
)
return None
expires_in: int | None = None
raw_expires: Final = body.get("expires_in")
try:
expires_in = int(raw_expires) if raw_expires is not None else None
except (TypeError, ValueError):
pass
# Rotate refresh token when the provider returns a new one
new_refresh_token: Final[str | None] = body.get("refresh_token") or refresh_token
raw_scope: Final = body.get("scope")
scopes: list[str] | None = (raw_scope.split() if isinstance(raw_scope, str) and raw_scope else None) or cred.get(
"scopes"
)
await store_user_oauth_credential(
prisma_client=prisma_client,
user_id=user_id,
server_id=server_id,
access_token=access_token,
refresh_token=new_refresh_token,
expires_in=expires_in,
scopes=scopes,
identity_binding_proof=binding_proof,
skip_byok_guard=True, # Row is already OAuth2; skip the extra find_unique check
)
verbose_proxy_logger.info(
"refresh_user_oauth_token: refreshed token for user=%s server=%s",
user_id,
server_id,
)
return await get_user_oauth_credential(prisma_client, user_id, server_id)
async def resolve_valid_user_oauth_token(
user_id: str,
server: "MCPServer",
cred: OAuthCredentialPayload | None,
prisma_client: PrismaClient | None = None,
) -> OAuthCredentialPayload | None:
"""Return an OAuth2 credential whose access_token is good for the next request.
Returns the credential unchanged while its token is valid for at least
``MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS``. Only when the token is expired (or
expiring within that buffer) and a refresh_token is stored does it mint a new one
via ``refresh_user_oauth_token``. Returns None when there is no usable token
(missing token, expired with no refresh_token, or a failed refresh).
The refresh_token is only ever sent to the server's token_url inside
``refresh_user_oauth_token``; it is never exposed to the caller beyond the cred
dict it already holds. ``prisma_client`` is fetched lazily and only when a refresh
actually happens, so the valid-token path never requires a DB handle.
"""
grant: Final = oauth_grant_state(cred)
if cred is None or grant == "absent":
return None
binding: Final = server.oauth_identity_binding
if binding is not None and binding.mode == "enforce":
if not await credential_binding_matches(binding, user_id, server.server_id, cred):
return None
if grant == "valid":
return cred
if prisma_client is None:
from litellm.proxy.utils import get_prisma_client_or_throw
prisma_client = get_prisma_client_or_throw("Database not connected. Cannot refresh OAuth token.")
refreshed: Final = await refresh_user_oauth_token(
prisma_client=prisma_client,
user_id=user_id,
server=server,
cred=cred,
)
if not refreshed or not refreshed.get("access_token"):
return None
return refreshed
async def resolve_user_oauth_access_token(
user_id: str | None,
server: "MCPServer",
prefetched_creds: Mapping[str, OAuthCredentialPayload] | None = None,
) -> str | None:
"""Resolve a user's valid OAuth2 access token for a server: Redis cache, else DB + refresh.
The egress token-resolution core shared by v1's header builder and the v2 ``OAuthTokenStore``
adapter. Redis fast-path (skipped when ``prefetched_creds`` is supplied), else a DB read through
``resolve_valid_user_oauth_token`` (which refreshes an expired token when a ``refresh_token`` is
stored), re-warming the Redis cache with the per-server TTL. Returns ``None`` when there is no
usable token; any error is swallowed to ``None`` so a transient failure reads as "not
authorized" rather than raising.
"""
server_id: Final[str | None] = getattr(server, "server_id", None)
if not user_id or not server_id:
return None
try:
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
_compute_per_user_token_ttl,
mcp_per_user_token_cache,
)
binding: Final = server.oauth_identity_binding
enforce_binding: Final = binding is not None and binding.mode == "enforce"
if prefetched_creds is None and enforce_binding and binding is not None:
bound_token: Final = await mcp_per_user_token_cache.get_token(user_id, server_id)
if bound_token is not None:
if await credential_binding_matches(
binding, user_id, server_id, {"identity_binding_proof": bound_token.identity_binding_proof}
):
return bound_token.access_token
await mcp_per_user_token_cache.delete(user_id, server_id)
if prefetched_creds is None and not enforce_binding:
cached_token: Final = await mcp_per_user_token_cache.get(user_id, server_id)
if cached_token is not None:
return cached_token
prisma_client = None
if prefetched_creds is not None:
cred = prefetched_creds.get(server_id)
else:
from litellm.proxy.utils import get_prisma_client_or_throw
prisma_client = get_prisma_client_or_throw(
"Database not connected. Connect a database to use OAuth2 MCP tools."
)
cred = await get_user_oauth_credential(prisma_client, user_id, server_id)
if not cred or not cred.get("access_token"):
return None
cred = await resolve_valid_user_oauth_token(
user_id=user_id,
server=server,
cred=cred,
prisma_client=prisma_client,
)
if cred is None:
# Refresh failed or token expired with no usable refresh_token — clear the stale
# Redis entry so the next request doesn't reuse it.
await mcp_per_user_token_cache.delete(user_id, server_id)
return None
access_token: Final[str] = cred["access_token"]
if prefetched_creds is None:
ttl: Final = _compute_per_user_token_ttl(server, _remaining_token_seconds(cred.get("expires_at")))
await mcp_per_user_token_cache.set(
user_id, server_id, access_token, ttl, identity_binding_proof=cred.get("identity_binding_proof")
)
return access_token
except Exception as e:
verbose_proxy_logger.warning(
"resolve_user_oauth_access_token: failed for user=%s server=%s: %s",
user_id,
server_id,
e,
)
return None
def _remaining_token_seconds(expires_at: str | None) -> int | None:
"""Seconds until ``expires_at`` (ISO 8601), or None when absent/past/unparseable."""
if not expires_at:
return None
try:
exp_dt = datetime.fromisoformat(expires_at)
except (ValueError, TypeError):
return None
if exp_dt.tzinfo is None:
exp_dt = exp_dt.replace(tzinfo=timezone.utc)
remaining: Final = int((exp_dt - datetime.now(timezone.utc)).total_seconds())
return remaining if remaining > 0 else None
async def get_active_submitted_mcp_server_ids_for_user(
prisma_client: PrismaClient,
user_id: str,
) -> list[str]:
"""Return active BYOM servers submitted by this user (creator visibility)."""
if not user_id:
return []
rows: Final = await _db_find_mcp_server_rows(
prisma_client,
{
"submitted_by": user_id,
"approval_status": MCPApprovalStatus.active,
},
)
return [row.server_id for row in rows]
async def approve_mcp_server(
prisma_client: PrismaClient,
server_id: str,
touched_by: str,
) -> LiteLLM_MCPServerTable:
"""Set approval_status=active and record reviewed_at."""
now: Final = datetime.now(timezone.utc)
updated: Final = await _db_update_mcp_server_row(
prisma_client,
server_id,
{
"approval_status": MCPApprovalStatus.active,
"reviewed_at": now,
"updated_by": touched_by,
},
)
table: Final = LiteLLM_MCPServerTable.model_validate(updated.model_dump())
decrypt_global_env_var_values(table.env_vars)
return table
async def reject_mcp_server(
prisma_client: PrismaClient,
server_id: str,
touched_by: str,
review_notes: str | None = None,
) -> LiteLLM_MCPServerTable:
"""Set approval_status=rejected, record reviewed_at and review_notes."""
now: Final = datetime.now(timezone.utc)
data: Final[prisma_db_types.LiteLLM_MCPServerTableUpdateInput] = {
"approval_status": MCPApprovalStatus.rejected,
"reviewed_at": now,
"updated_by": touched_by,
}
if review_notes is not None:
data["review_notes"] = review_notes
updated: Final = await _db_update_mcp_server_row(prisma_client, server_id, data)
table: Final = LiteLLM_MCPServerTable.model_validate(updated.model_dump())
decrypt_global_env_var_values(table.env_vars)
return table
async def get_mcp_submissions(
prisma_client: PrismaClient,
) -> MCPSubmissionsSummary:
"""
Returns all MCP servers that were submitted by non-admin users (submitted_at IS NOT NULL),
along with a summary count breakdown by approval_status.
Mirrors get_guardrail_submissions() from guardrail_endpoints.py.
"""
rows: Final[Sequence[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions(
prisma_client
).find_many(
where={"submitted_at": {"not": None}},
order={"submitted_at": "desc"},
take=500, # safety cap; paginate if needed in a future iteration
)
items: Final = list(_readable_mcp_servers(rows))
pending: Final = sum(1 for i in items if i.approval_status == MCPApprovalStatus.pending_review)
active: Final = sum(1 for i in items if i.approval_status == MCPApprovalStatus.active)
rejected: Final = sum(1 for i in items if i.approval_status == MCPApprovalStatus.rejected)
return MCPSubmissionsSummary(
total=len(items),
pending_review=pending,
active=active,
rejected=rejected,
items=items,
)
# ── Per-user MCP environment variables ────────────────────────────────────
def _decode_user_env_vars(stored: str) -> dict[str, str]:
"""Decrypt a ``values_b64`` blob and parse it as a flat ``{name: value}`` dict."""
decrypted: Final = decrypt_value_helper(
value=stored,
key="mcp_user_env_vars",
exception_type="debug",
return_original_value=False,
)
if decrypted is None:
if stored:
verbose_proxy_logger.warning(
"MCP per-user env vars failed to decrypt (LITELLM_SALT_KEY "
"changed?); treating as unset so the user is prompted to "
"re-enter them rather than silently forwarding ciphertext"
)
return {}
parsed: dict[str, object] | None
try:
parsed = json.loads(decrypted)
except (ValueError, TypeError):
return {}
if not isinstance(parsed, dict):
return {}
return {str(k): str(v) for k, v in parsed.items()}
async def get_user_env_vars(
prisma_client: PrismaClient,
user_id: str,
server_id: str,
) -> dict[str, str]:
"""Return the calling user's env var dict for ``server_id`` (empty if none)."""
row: Final = await _user_env_var_actions(prisma_client).find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
if row is None:
return {}
return _decode_user_env_vars(row.values_b64)
async def get_user_env_vars_bulk(
prisma_client: PrismaClient,
user_id: str,
server_ids: Iterable[str],
) -> dict[str, dict[str, str]]:
"""Return ``{server_id: {var_name: value}}`` for one user across many servers.
Servers with no stored row are simply absent from the result.
"""
ids: Final = list(server_ids)
if not ids:
return {}
rows: Final = await _db_find_user_env_var_rows(prisma_client, {"user_id": user_id, "server_id": {"in": ids}})
return {row.server_id: _decode_user_env_vars(row.values_b64) for row in rows}
async def merge_user_env_vars(
prisma_client: PrismaClient,
user_id: str,
server_id: str,
updates: dict[str, str],
allowed_names: Iterable[str],
) -> dict[str, str]:
"""Merge ``updates`` into the user's stored env vars for ``server_id`` and
return the resulting set.
The read-modify-write runs inside a transaction guarded by a
``(user_id, server_id)`` advisory lock so two concurrent writes from the
same user can't drop one update. Names outside ``allowed_names`` are pruned,
so an admin retiring a user-scoped variable also clears its stored value.
"""
allowed: Final = set(allowed_names)
lock_key: Final = int.from_bytes(
hashlib.blake2b(f"{user_id}:{server_id}".encode(), digest_size=8).digest(),
"big",
signed=True,
)
async with _db_transaction_manager(prisma_client) as tx:
await tx.execute_raw("SELECT pg_advisory_xact_lock($1::bigint)", lock_key)
row: Final[prisma_db_models.LiteLLM_MCPUserEnvVars | None] = await tx.litellm_mcpuserenvvars.find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
existing: Final = _decode_user_env_vars(row.values_b64) if row is not None else {}
merged: Final = {k: v for k, v in {**existing, **updates}.items() if k in allowed}
encoded: Final = encrypt_value_helper(json.dumps(merged))
await tx.litellm_mcpuserenvvars.upsert(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
data={
"create": {
"user_id": user_id,
"server_id": server_id,
"values_b64": encoded,
},
"update": {"values_b64": encoded},
},
)
return merged
async def delete_user_env_vars(
prisma_client: PrismaClient,
user_id: str,
server_id: str,
) -> None:
"""Remove the calling user's env var values for ``server_id``.
Uses ``delete_many`` so a missing row is a no-op; real DB errors still
propagate to the caller instead of being silently swallowed.
"""
await _user_env_var_actions(prisma_client).delete_many(where={"user_id": user_id, "server_id": server_id})