mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
* 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>
2309 lines
94 KiB
Python
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})
|