From da8aa59e14ffaad7dc06826af8488d9c1ea8cb75 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 2 Oct 2026 20:08:17 +0000 Subject: [PATCH] feat(mcp): track MCP tool versions, changelogs and deprecations Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../migration.sql | 21 + .../litellm_proxy_extras/schema.prisma | 20 + litellm/proxy/_experimental/mcp_server/db.py | 136 ++++- .../mcp_server/tool_versioning.py | 262 ++++++++++ litellm/proxy/_lazy_openapi_snapshot.json | 391 +++++++++++++++ .../mcp_management_endpoints.py | 218 ++++++-- litellm/proxy/schema.prisma | 20 + litellm/repositories/table_repositories.py | 4 + .../types/mcp_server/mcp_server_manager.py | 41 ++ schema.prisma | 20 + .../mcp_server/test_mcp_env_vars.py | 33 ++ .../mcp_server/test_mcp_partial_update.py | 337 ++++++++++++- .../mcp_server/test_tool_versioning.py | 359 ++++++++++++++ .../test_mcp_management_endpoints.py | 468 ++++++++++++++++-- .../_components/MCPToolVersionsPanel.test.tsx | 224 +++++++++ .../_components/MCPToolVersionsPanel.tsx | 379 ++++++++++++++ .../_components/mcp_server_view.test.tsx | 12 +- .../_components/mcp_server_view.tsx | 1 + .../mcp-servers/_components/mcp_tools.tsx | 10 + .../src/components/mcp_tools/types.tsx | 1 + .../src/components/networking.test.ts | 24 + .../src/components/networking.tsx | 48 ++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 213 +++++++- 23 files changed, 3132 insertions(+), 110 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20261002000000_mcp_tool_versions/migration.sql create mode 100644 litellm/proxy/_experimental/mcp_server/tool_versioning.py create mode 100644 tests/unit/proxy/_experimental/mcp_server/test_tool_versioning.py create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolVersionsPanel.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolVersionsPanel.tsx diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261002000000_mcp_tool_versions/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261002000000_mcp_tool_versions/migration.sql new file mode 100644 index 00000000000..3652c099e2a --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261002000000_mcp_tool_versions/migration.sql @@ -0,0 +1,21 @@ +CREATE TABLE IF NOT EXISTS "LiteLLM_MCPToolVersion" ( + "id" TEXT NOT NULL, + "server_id" TEXT NOT NULL, + "tool_name" TEXT NOT NULL, + "version" INTEGER NOT NULL, + "description" TEXT NOT NULL DEFAULT '', + "input_schema" JSONB NOT NULL DEFAULT '{}', + "change_kind" TEXT NOT NULL, + "changes" JSONB NOT NULL DEFAULT '[]', + "changelog" TEXT, + "deprecated_at" TIMESTAMP(3), + "sunset_date" TIMESTAMP(3), + "deprecation_note" TEXT, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT, + + CONSTRAINT "LiteLLM_MCPToolVersion_pkey" PRIMARY KEY ("id") +); +CREATE INDEX IF NOT EXISTS "LiteLLM_MCPToolVersion_server_id_idx" ON "LiteLLM_MCPToolVersion"("server_id"); +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_MCPToolVersion_server_id_tool_name_version_key" +ON "LiteLLM_MCPToolVersion"("server_id", "tool_name", "version"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index aba89526cf6..33de27cd893 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -427,6 +427,26 @@ model LiteLLM_MCPServerTable { @@index([approval_status]) } +model LiteLLM_MCPToolVersion { + id String @id @default(uuid()) + server_id String + tool_name String + version Int + description String @default("") + input_schema Json @default("{}") + change_kind String + changes Json @default("[]") + changelog String? + deprecated_at DateTime? + sunset_date DateTime? + deprecation_note String? + created_at DateTime @default(now()) + created_by String? + + @@unique([server_id, tool_name, version]) + @@index([server_id]) +} + // Named collection of {server_id, tool_name} pairs that can be granted to keys/teams model LiteLLM_MCPToolsetTable { toolset_id String @id @default(uuid()) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index bf1f6fa90c1..77814a44d1b 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -45,6 +45,7 @@ from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( MCPServerOAuthClientRepository, MCPServerRepository, + MCPToolVersionRepository, MCPUserCredentialsRepository, PrismaTableRepository, ) @@ -54,7 +55,12 @@ from litellm.repositories.verification_token_repository import ( ) from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.mcp import MCPCredentials, MCPTransportType, MCPUpstreamProtocol, validate_mcp_protocol_transport -from litellm.types.mcp_server.mcp_server_manager import MCPInfo, PinnedMCPTool +from litellm.types.mcp_server.mcp_server_manager import ( + MCPInfo, + MCPToolDeprecationRequest, + MCPToolVersion, + PinnedMCPTool, +) if TYPE_CHECKING: from prisma import models as prisma_db_models @@ -484,6 +490,13 @@ def _mcp_server_table_actions( return table +def _mcp_tool_version_table_actions( + prisma_client: PrismaClient, +) -> "TableActions[prisma_db_models.LiteLLM_MCPToolVersion]": + table: Final[TableActions[prisma_db_models.LiteLLM_MCPToolVersion]] = MCPToolVersionRepository(prisma_client).table + return table + + def _verification_token_table_actions( prisma_client: PrismaClient, ) -> "TableActions[prisma_db_models.LiteLLM_VerificationToken]": @@ -514,6 +527,9 @@ def _db_transaction_manager(prisma_client: PrismaClient) -> _UserEnvVarsTransact return manager +_MCP_PIN_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" + + 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}]},) if exclude_server_id is not None else () where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = { @@ -918,6 +934,15 @@ async def delete_mcp_server( }, ) if deleted_server is not None: + try: + await _mcp_tool_version_table_actions(prisma_client).delete_many(where={"server_id": server_id}) + except Exception as e: + verbose_proxy_logger.warning( + "MCP server %s deleted but tool version cleanup failed; " + "orphaned rows can be removed on a later delete: %s", + server_id, + e, + ) credential_user_ids: list[str] = [] try: credential_rows: Sequence[prisma_db_models.LiteLLM_MCPUserCredentials] = await _user_credential_actions( @@ -2177,23 +2202,112 @@ async def set_mcp_server_pinned_tools( server_id: str, pinned_tools: Mapping[str, PinnedMCPTool] | None, touched_by: str, + changelog: str | None = None, ) -> LiteLLM_MCPServerTable | None: """Replace the server's pinned catalog; ``None`` unpins. Only this write path sets the pin.""" from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - if await _db_find_mcp_server_row(prisma_client, server_id) is None: - return None - snapshot: Final = {name: tool.model_dump() for name, tool in (pinned_tools or {}).items()} - updated: Final = await _db_update_mcp_server_row( - prisma_client, - server_id, - {"pinned_tools": safe_dumps(snapshot), "updated_by": touched_by}, - ) - table: Final = LiteLLM_MCPServerTable.model_validate(updated.model_dump()) - decrypt_global_env_var_values(table.env_vars) + async with prisma_client.tx() as tx: + _ = await tx.execute_raw(_MCP_PIN_ADVISORY_LOCK_SQL, f"mcp_pin:{server_id}") + server_row: Final = await tx.litellm_mcpservertable.find_unique(where={"server_id": server_id}) + if server_row is None: + return None + if pinned_tools is not None: + from litellm.proxy._experimental.mcp_server.tool_versioning import plan_tool_versions + + existing_versions: Final = await _list_mcp_tool_versions_from_actions(tx.litellm_mcptoolversion, server_id) + latest_versions: Final = {row.tool_name: row for row in reversed(existing_versions)} + planned_versions: Final = plan_tool_versions(latest_versions, pinned_tools) + version_rows: Final[tuple[dict[str, object], ...]] = tuple( + { + "server_id": server_id, + "tool_name": planned.tool_name, + "version": planned.version, + "description": planned.tool.description, + "input_schema": safe_dumps(planned.tool.input_schema), + "change_kind": planned.change_kind, + "changes": safe_dumps([change.model_dump() for change in planned.changes]), + "changelog": changelog, + "created_by": touched_by, + } + for planned in planned_versions + ) + if version_rows: + await tx.litellm_mcptoolversion.create_many(data=version_rows) + snapshot: Final = {name: tool.model_dump() for name, tool in (pinned_tools or {}).items()} + updated: Final = await tx.litellm_mcpservertable.update( + where={"server_id": server_id}, + data={"pinned_tools": safe_dumps(snapshot), "updated_by": touched_by}, + ) + if updated is None: + raise ValueError(f"MCP server not found, passed server_id={server_id}") + table: Final = LiteLLM_MCPServerTable.model_validate(updated.model_dump()) + decrypt_global_env_var_values(table.env_vars) + from litellm.proxy.common_utils.config_sync_pubsub import publish_config_change_for_object_type + + await publish_config_change_for_object_type("litellm_mcpservertable") return table +async def _list_mcp_tool_versions_from_actions( + table: "TableActions[prisma_db_models.LiteLLM_MCPToolVersion]", + server_id: str, +) -> list[MCPToolVersion]: + rows: Final[Sequence[prisma_db_models.LiteLLM_MCPToolVersion]] = await table.find_many( + where={"server_id": server_id}, + order=[{"tool_name": "asc"}, {"version": "desc"}], + ) + return [MCPToolVersion.model_validate(row.model_dump()) for row in rows] + + +async def list_mcp_tool_versions(prisma_client: PrismaClient, server_id: str) -> list[MCPToolVersion]: + return await _list_mcp_tool_versions_from_actions( + _mcp_tool_version_table_actions(prisma_client), + server_id, + ) + + +async def set_mcp_tool_version_deprecation( + prisma_client: PrismaClient, + server_id: str, + tool_name: str, + version: int, + deprecation: MCPToolDeprecationRequest | None, +) -> MCPToolVersion | None: + unique_where: Final = { + "server_id_tool_name_version": { + "server_id": server_id, + "tool_name": tool_name, + "version": version, + } + } + async with prisma_client.tx() as tx: + _ = await tx.execute_raw(_MCP_PIN_ADVISORY_LOCK_SQL, f"mcp_pin:{server_id}") + row: Final[prisma_db_models.LiteLLM_MCPToolVersion | None] = await tx.litellm_mcptoolversion.find_unique( + where=unique_where + ) + if row is None: + return None + data: Final[dict[str, object]] = ( + { + "deprecated_at": None, + "sunset_date": None, + "deprecation_note": None, + } + if deprecation is None + else { + "deprecated_at": row.deprecated_at or datetime.now(timezone.utc), + "sunset_date": deprecation.sunset_date, + "deprecation_note": deprecation.deprecation_note, + } + ) + updated: Final[prisma_db_models.LiteLLM_MCPToolVersion | None] = await tx.litellm_mcptoolversion.update( + where=unique_where, + data=data, + ) + return MCPToolVersion.model_validate(updated.model_dump()) if updated is not None else None + + async def reject_mcp_server( prisma_client: PrismaClient, server_id: str, diff --git a/litellm/proxy/_experimental/mcp_server/tool_versioning.py b/litellm/proxy/_experimental/mcp_server/tool_versioning.py new file mode 100644 index 00000000000..139304def9a --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/tool_versioning.py @@ -0,0 +1,262 @@ +import json +from collections.abc import Mapping +from dataclasses import dataclass +from itertools import chain +from typing import Final, cast + +from litellm.types.mcp_server.mcp_server_manager import ( + MCPToolChange, + MCPToolChangeKind, + MCPToolVersion, + PinnedMCPTool, +) + + +@dataclass(frozen=True, slots=True) +class PlannedToolVersion: + tool_name: str + version: int + tool: PinnedMCPTool + change_kind: MCPToolChangeKind + changes: tuple[MCPToolChange, ...] + + +def _properties(input_schema: Mapping[str, object]) -> Mapping[str, object]: + properties: Final[object] = input_schema.get("properties") + return properties if isinstance(properties, Mapping) else {} + + +def _required(input_schema: Mapping[str, object]) -> frozenset[str]: + required: Final[object] = input_schema.get("required") + if not isinstance(required, list): + return frozenset() + return frozenset(value for value in required if isinstance(value, str)) + + +def _mapping_without_documentation(schema: Mapping[str, object]) -> dict[str, object]: + mapping_keys: Final = frozenset(("properties", "patternProperties", "$defs", "definitions", "dependentSchemas")) + list_keys: Final = frozenset(("items", "prefixItems", "anyOf", "oneOf", "allOf")) + schema_keys: Final = frozenset( + ( + "items", + "additionalProperties", + "unevaluatedProperties", + "contains", + "not", + "if", + "then", + "else", + "propertyNames", + ) + ) + return { + key: ( + { + child_key: _schema_without_documentation(child) + for child_key, child in cast( # cast-ok: schema maps contain nested schemas + Mapping[str, object], item + ).items() + } + if key in mapping_keys and isinstance(item, Mapping) + else [ + _schema_without_documentation(child) + for child in cast( # cast-ok: schema array positions contain objects + list[object], item + ) + ] + if key in list_keys and isinstance(item, list) + else _schema_without_documentation( + cast( # cast-ok: schema keyword values are object schemas + Mapping[str, object], item + ) + ) + if key in schema_keys and isinstance(item, Mapping) + else item + ) + for key, item in schema.items() + if key not in ("description", "title") + } + + +def _schema_without_documentation(value: object) -> object: + if not isinstance(value, Mapping): + return value + return _mapping_without_documentation( + cast(Mapping[str, object], value) # cast-ok: MCP schemas use string-keyed JSON objects + ) + + +def _type_label(parameter_schema: object) -> str: + if not isinstance(parameter_schema, Mapping) or "type" not in parameter_schema: + return "any" + return json.dumps(parameter_schema["type"]) + + +def _classify_parameter( + name: str, + previous_properties: Mapping[str, object], + current_properties: Mapping[str, object], + previous_required: frozenset[str], + current_required: frozenset[str], +) -> tuple[MCPToolChange, ...]: + in_previous: Final = name in previous_properties + in_current: Final = name in current_properties + required_change: Final[tuple[MCPToolChange, ...]] = ( + (MCPToolChange(breaking=True, summary=f'Parameter "{name}" is now required'),) + if name not in previous_required and name in current_required + else ( + (MCPToolChange(breaking=False, summary=f'Parameter "{name}" is now optional'),) + if name in previous_required and name not in current_required + else () + ) + ) + if not in_previous and not in_current: + return required_change + if in_previous and not in_current: + return (MCPToolChange(breaking=True, summary=f'Removed parameter "{name}"'),) + if in_current and not in_previous: + is_required: Final = name in current_required + summary: Final = f'Added required parameter "{name}"' if is_required else f'Added optional parameter "{name}"' + return (MCPToolChange(breaking=is_required, summary=summary),) + + previous_parameter: Final = previous_properties[name] + current_parameter: Final = current_properties[name] + previous_type: Final = _type_label(previous_parameter) + current_type: Final = _type_label(current_parameter) + schema_change: Final[tuple[MCPToolChange, ...]] = ( + ( + MCPToolChange( + breaking=True, + summary=f'Parameter "{name}" type changed from {previous_type} to {current_type}', + ), + ) + if previous_type != current_type + else ( + (MCPToolChange(breaking=True, summary=f'Parameter "{name}" schema changed'),) + if _schema_without_documentation(previous_parameter) != _schema_without_documentation(current_parameter) + else ( + (MCPToolChange(breaking=False, summary=f'Parameter "{name}" description changed'),) + if previous_parameter != current_parameter + else () + ) + ) + ) + return schema_change + required_change + + +def classify_tool_change(previous: PinnedMCPTool, current: PinnedMCPTool) -> tuple[MCPToolChange, ...]: + previous_properties: Final = _properties(previous.input_schema) + current_properties: Final = _properties(current.input_schema) + previous_required: Final = _required(previous.input_schema) + current_required: Final = _required(current.input_schema) + parameter_names: Final = sorted( + set(previous_properties) | set(current_properties) | previous_required | current_required + ) + parameter_changes: Final = tuple( + chain.from_iterable( + _classify_parameter( + name, + previous_properties, + current_properties, + previous_required, + current_required, + ) + for name in parameter_names + ) + ) + + excluded_top_level_keys: Final = frozenset(("properties", "required", "description", "title")) + previous_top_level_raw: Final = { + key: value for key, value in previous.input_schema.items() if key not in excluded_top_level_keys + } + current_top_level_raw: Final = { + key: value for key, value in current.input_schema.items() if key not in excluded_top_level_keys + } + previous_top_level: Final = { + key: value + for key, value in _mapping_without_documentation(previous.input_schema).items() + if key not in excluded_top_level_keys + } + current_top_level: Final = { + key: value + for key, value in _mapping_without_documentation(current.input_schema).items() + if key not in excluded_top_level_keys + } + description_changes: Final = ( + (MCPToolChange(breaking=False, summary="Description changed"),) + if previous.description != current.description + else () + ) + input_schema_description_changes: Final = ( + (MCPToolChange(breaking=False, summary="Input schema description changed"),) + if previous.input_schema.get("description") != current.input_schema.get("description") + else () + ) + input_schema_title_changes: Final = ( + (MCPToolChange(breaking=False, summary="Input schema title changed"),) + if previous.input_schema.get("title") != current.input_schema.get("title") + else () + ) + schema_changes: Final = ( + (MCPToolChange(breaking=True, summary="Input schema changed"),) + if previous_top_level != current_top_level + else ( + (MCPToolChange(breaking=False, summary="Input schema documentation changed"),) + if previous_top_level_raw != current_top_level_raw + else () + ) + ) + return ( + description_changes + + input_schema_description_changes + + input_schema_title_changes + + parameter_changes + + schema_changes + ) + + +def _plan_tool_version( + name: str, + latest: Mapping[str, MCPToolVersion], + snapshot: Mapping[str, PinnedMCPTool], +) -> tuple[PlannedToolVersion, ...]: + has_snapshot: Final = name in snapshot + has_latest: Final = name in latest + if has_snapshot and not has_latest: + return (PlannedToolVersion(name, 1, snapshot[name], "initial", ()),) + if has_snapshot and has_latest: + current_tool: Final = snapshot[name] + latest_version: Final = latest[name] + is_restored: Final = latest_version.change_kind == "removed" + changes: Final = ( + (MCPToolChange(breaking=False, summary="Tool restored"),) if is_restored else () + ) + classify_tool_change( + PinnedMCPTool(description=latest_version.description, input_schema=latest_version.input_schema), + current_tool, + ) + if not changes: + return () + change_kind: Final[MCPToolChangeKind] = ( + "breaking" if any(change.breaking for change in changes) else "non_breaking" + ) + return (PlannedToolVersion(name, latest_version.version + 1, current_tool, change_kind, changes),) + if has_latest and latest[name].change_kind != "removed": + previous_version: Final = latest[name] + return ( + PlannedToolVersion( + name, + previous_version.version + 1, + PinnedMCPTool(description=previous_version.description, input_schema=previous_version.input_schema), + "removed", + (MCPToolChange(breaking=True, summary="Tool removed"),), + ), + ) + return () + + +def plan_tool_versions( + latest: Mapping[str, MCPToolVersion], + snapshot: Mapping[str, PinnedMCPTool], +) -> tuple[PlannedToolVersion, ...]: + names: Final = sorted(set(latest) | set(snapshot)) + return tuple(chain.from_iterable(_plan_tool_version(name, latest, snapshot) for name in names)) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 3346b0c9ff8..93a6dc9ff46 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -37953,6 +37953,170 @@ "title": "MCPSubmissionsSummary", "type": "object" }, + "MCPToolChange": { + "additionalProperties": false, + "properties": { + "breaking": { + "title": "Breaking", + "type": "boolean" + }, + "summary": { + "title": "Summary", + "type": "string" + } + }, + "required": [ + "breaking", + "summary" + ], + "title": "MCPToolChange", + "type": "object" + }, + "MCPToolDeprecationRequest": { + "additionalProperties": false, + "properties": { + "deprecation_note": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Deprecation Note" + }, + "sunset_date": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Sunset Date" + } + }, + "title": "MCPToolDeprecationRequest", + "type": "object" + }, + "MCPToolVersion": { + "properties": { + "change_kind": { + "enum": [ + "initial", + "non_breaking", + "breaking", + "removed" + ], + "title": "Change Kind", + "type": "string" + }, + "changelog": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Changelog" + }, + "changes": { + "default": [], + "items": { + "$ref": "#/components/schemas/MCPToolChange" + }, + "title": "Changes", + "type": "array" + }, + "created_at": { + "format": "date-time", + "title": "Created At", + "type": "string" + }, + "created_by": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Created By" + }, + "deprecated_at": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Deprecated At" + }, + "deprecation_note": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Deprecation Note" + }, + "description": { + "default": "", + "title": "Description", + "type": "string" + }, + "input_schema": { + "additionalProperties": true, + "title": "Input Schema", + "type": "object" + }, + "server_id": { + "title": "Server Id", + "type": "string" + }, + "sunset_date": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Sunset Date" + }, + "tool_name": { + "title": "Tool Name", + "type": "string" + }, + "version": { + "title": "Version", + "type": "integer" + } + }, + "required": [ + "server_id", + "tool_name", + "version", + "change_kind", + "created_at" + ], + "title": "MCPToolVersion", + "type": "object" + }, "MCPToolsetTool": { "properties": { "server_id": { @@ -38717,6 +38881,24 @@ "title": "NewMCPToolsetRequest", "type": "object" }, + "PinMCPServerToolsRequest": { + "additionalProperties": false, + "properties": { + "changelog": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Changelog" + } + }, + "title": "PinMCPServerToolsRequest", + "type": "object" + }, "PinnedMCPTool": { "additionalProperties": false, "description": "One tool of an admin-pinned catalog: the description and input schema tools/list keeps serving.", @@ -40367,6 +40549,23 @@ } } ], + "requestBody": { + "content": { + "application/json": { + "schema": { + "anyOf": [ + { + "$ref": "#/components/schemas/PinMCPServerToolsRequest" + }, + { + "type": "null" + } + ], + "title": "Payload" + } + } + } + }, "responses": { "200": { "content": { @@ -40462,6 +40661,198 @@ ] } }, + "/v1/mcp/server/{server_id}/tool-versions": { + "get": { + "description": "Returns the version history of the tools recorded for an MCP server.", + "operationId": "fetch_mcp_server_tool_versions_v1_mcp_server__server_id__tool_versions_get", + "parameters": [ + { + "in": "path", + "name": "server_id", + "required": true, + "schema": { + "title": "Server Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "items": { + "$ref": "#/components/schemas/MCPToolVersion" + }, + "title": "Response Fetch Mcp Server Tool Versions V1 Mcp Server Server Id Tool Versions Get", + "type": "array" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Fetch Mcp Server Tool Versions", + "tags": [ + "mcp_management" + ] + } + }, + "/v1/mcp/server/{server_id}/tools/{tool_name}/versions/{version}/deprecation": { + "delete": { + "description": "Clear deprecation and sunset metadata for an MCP tool version.", + "operationId": "clear_tool_version_deprecation_v1_mcp_server__server_id__tools__tool_name__versions__version__deprecation_delete", + "parameters": [ + { + "in": "path", + "name": "server_id", + "required": true, + "schema": { + "title": "Server Id", + "type": "string" + } + }, + { + "in": "path", + "name": "tool_name", + "required": true, + "schema": { + "title": "Tool Name", + "type": "string" + } + }, + { + "in": "path", + "name": "version", + "required": true, + "schema": { + "title": "Version", + "type": "integer" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/MCPToolVersion" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Clear Tool Version Deprecation", + "tags": [ + "mcp_management" + ] + }, + "put": { + "description": "Set deprecation and sunset metadata for an MCP tool version.", + "operationId": "set_tool_version_deprecation_v1_mcp_server__server_id__tools__tool_name__versions__version__deprecation_put", + "parameters": [ + { + "in": "path", + "name": "server_id", + "required": true, + "schema": { + "title": "Server Id", + "type": "string" + } + }, + { + "in": "path", + "name": "tool_name", + "required": true, + "schema": { + "title": "Tool Name", + "type": "string" + } + }, + { + "in": "path", + "name": "version", + "required": true, + "schema": { + "title": "Version", + "type": "integer" + } + } + ], + "requestBody": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/MCPToolDeprecationRequest" + } + } + }, + "required": true + }, + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/MCPToolVersion" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Set Tool Version Deprecation", + "tags": [ + "mcp_management" + ] + } + }, "/v1/mcp/server/{server_id}/user-credential": { "delete": { "description": "Delete the calling user's stored API key for a BYOK MCP server. A proxy admin may pass user_id to revoke another user's stored key.", diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 0d358302fa2..2bf5b54e9be 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -155,6 +155,7 @@ if MCP_AVAILABLE: get_user_env_vars, get_user_env_vars_bulk, get_user_oauth_credential, + list_mcp_tool_versions, list_server_user_credentials, list_user_oauth_credentials, mcp_oauth_token_identity, @@ -162,6 +163,7 @@ if MCP_AVAILABLE: purge_user_oauth_credentials_for_server, reject_mcp_server, set_mcp_server_pinned_tools, + set_mcp_tool_version_deprecation, store_user_credential, store_user_oauth_credential, update_mcp_server, @@ -181,6 +183,7 @@ if MCP_AVAILABLE: ) from litellm.proxy._experimental.mcp_server.server_resolution import ( MCPServerTargetCatalog, + ResolvedMCPServer, authorize_mcp_server, resolve_mcp_server, ) @@ -244,7 +247,24 @@ if MCP_AVAILABLE: normalize_upstream_header_name, validate_mcp_protocol_transport, ) - from litellm.types.mcp_server.mcp_server_manager import MCPServer, PinnedMCPTool + from litellm.types.mcp_server.mcp_server_manager import ( + MCPServer, + MCPToolDeprecationRequest, + MCPToolVersion, + PinMCPServerToolsRequest, + PinnedMCPTool, + ) + + async def _get_allowed_tool_names_for_server( + server_id: str, + user_api_key_dict: UserAPIKeyAuth, + ) -> list[str] | None: + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + + return await MCPRequestHandler.get_allowed_tools_for_server( + server_id=server_id, + user_api_key_auth=user_api_key_dict, + ) @dataclass class _TemporaryMCPServerEntry: @@ -1553,7 +1573,8 @@ if MCP_AVAILABLE: async def pin_mcp_server_tools( server_id: str, request: Request, - user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI dependency injection + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + payload: PinMCPServerToolsRequest | None = None, ) -> dict[str, PinnedMCPTool]: if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role: raise HTTPException( @@ -1578,7 +1599,8 @@ if MCP_AVAILABLE: "error": f"MCP server '{server_id}' exposes no tools that pass the guardrails; nothing to pin." }, ) - await _store_pinned_tools(server_id, snapshot, user_api_key_dict) + changelog: Final = (payload.changelog or "").strip() if payload is not None else "" + await _store_pinned_tools(server_id, snapshot, user_api_key_dict, changelog=changelog or None) return snapshot @router.delete( @@ -1607,15 +1629,30 @@ if MCP_AVAILABLE: return {"server_id": server_id, "status": "unpinned"} async def _store_pinned_tools( - server_id: str, pinned_tools: Mapping[str, PinnedMCPTool] | None, user_api_key_dict: UserAPIKeyAuth + server_id: str, + pinned_tools: Mapping[str, PinnedMCPTool] | None, + user_api_key_dict: UserAPIKeyAuth, + changelog: str | None = None, ) -> None: prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - record: Final = await set_mcp_server_pinned_tools( - prisma_client, - server_id, - pinned_tools, - touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, - ) + touched_by: Final = user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME + try: + record: Final = await set_mcp_server_pinned_tools( + prisma_client, + server_id, + pinned_tools, + touched_by=touched_by, + changelog=changelog, + ) + except Exception as e: + if not isinstance(e, UniqueViolationError): + raise + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail={ + "error": f"Another pin of MCP server '{server_id}' recorded tool versions at the same time; retry the pin." + }, + ) from e if record is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -1715,32 +1752,12 @@ if MCP_AVAILABLE: await global_mcp_server_manager.reload_servers_from_database() return _redact_mcp_credentials(rejected) - @router.get( - "/server/{server_id}", - description="Returns the mcp server info", - dependencies=[Depends(user_api_key_auth)], - response_model=LiteLLM_MCPServerTable, - ) - async def fetch_mcp_server( + async def _resolve_and_authorize_mcp_server( request: Request, server_id: str, - user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - include_reachability: Annotated[ - bool, - Query(description="Allow the 'reachable' status for responding servers whose authentication is unchecked."), - ] = False, - ): - """ - Get the info on the mcp server specified by the `server_id` - Parameters: - - server_id: str - Required. The unique identifier of the mcp server to get info on. - ``` - curl --location 'http://localhost:4000/v1/mcp/server/server_id' \ - --header 'Authorization: Bearer your_api_key_here' - ``` - """ - prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - + user_api_key_dict: UserAPIKeyAuth, + prisma_client: "PrismaClient", + ) -> tuple[ResolvedMCPServer, bool]: from litellm.proxy.auth.ip_address_utils import IPAddressUtils client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) @@ -1775,6 +1792,139 @@ if MCP_AVAILABLE: and global_mcp_server_manager.get_mcp_server_by_id(resolved.table.server_id) is not None ), ) + return authorized, is_restricted_virtual_key + + @router.get( + "/server/{server_id}/tool-versions", + description="Returns the version history of the tools recorded for an MCP server.", + dependencies=[Depends(user_api_key_auth)], + response_model=list[MCPToolVersion], + ) + async def fetch_mcp_server_tool_versions( + request: Request, + server_id: str, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + ) -> list[MCPToolVersion]: + prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") + authorized: Final[tuple[ResolvedMCPServer, bool]] = await _resolve_and_authorize_mcp_server( + request, + server_id, + user_api_key_dict, + prisma_client, + ) + if authorized[0].source != "db": + return [] + resolved_server_id: Final = authorized[0].table.server_id + allowed: Final = await _get_allowed_tool_names_for_server( + server_id=resolved_server_id, + user_api_key_dict=user_api_key_dict, + ) + versions: Final = await list_mcp_tool_versions(prisma_client, resolved_server_id) + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + + return [version for version in versions if MCPRequestHandler.tool_is_granted(version.tool_name, allowed)] + + @router.put( + "/server/{server_id}/tools/{tool_name}/versions/{version}/deprecation", + description="Set deprecation and sunset metadata for an MCP tool version.", + dependencies=[Depends(user_api_key_auth)], + response_model=MCPToolVersion, + ) + @management_endpoint_wrapper + async def set_tool_version_deprecation( + server_id: str, + tool_name: str, + version: int, + payload: MCPToolDeprecationRequest, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + ) -> MCPToolVersion: + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={"error": "Admin access required to update MCP tool version deprecation."}, + ) + prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") + result: Final[MCPToolVersion | None] = await set_mcp_tool_version_deprecation( + prisma_client, + server_id, + tool_name, + version, + payload, + ) + if result is None: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail={"error": f"MCP tool '{tool_name}' version {version} not found for server '{server_id}'."}, + ) + return result + + @router.delete( + "/server/{server_id}/tools/{tool_name}/versions/{version}/deprecation", + description="Clear deprecation and sunset metadata for an MCP tool version.", + dependencies=[Depends(user_api_key_auth)], + response_model=MCPToolVersion, + ) + @management_endpoint_wrapper + async def clear_tool_version_deprecation( + server_id: str, + tool_name: str, + version: int, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + ) -> MCPToolVersion: + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={"error": "Admin access required to update MCP tool version deprecation."}, + ) + prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") + result: Final[MCPToolVersion | None] = await set_mcp_tool_version_deprecation( + prisma_client, + server_id, + tool_name, + version, + None, + ) + if result is None: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail={"error": f"MCP tool '{tool_name}' version {version} not found for server '{server_id}'."}, + ) + return result + + @router.get( + "/server/{server_id}", + description="Returns the mcp server info", + dependencies=[Depends(user_api_key_auth)], + response_model=LiteLLM_MCPServerTable, + ) + async def fetch_mcp_server( + request: Request, + server_id: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + include_reachability: Annotated[ + bool, + Query(description="Allow the 'reachable' status for responding servers whose authentication is unchecked."), + ] = False, + ): + """ + Get the info on the mcp server specified by the `server_id` + Parameters: + - server_id: str - Required. The unique identifier of the mcp server to get info on. + ``` + curl --location 'http://localhost:4000/v1/mcp/server/server_id' \ + --header 'Authorization: Bearer your_api_key_here' + ``` + """ + prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") + + authorized_result: Final[tuple[ResolvedMCPServer, bool]] = await _resolve_and_authorize_mcp_server( + request, + server_id, + user_api_key_dict, + prisma_client, + ) + authorized: Final = authorized_result[0] + is_restricted_virtual_key: Final = authorized_result[1] mcp_server: Final = authorized.table from_db: Final = authorized.source == "db" diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index aba89526cf6..33de27cd893 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -427,6 +427,26 @@ model LiteLLM_MCPServerTable { @@index([approval_status]) } +model LiteLLM_MCPToolVersion { + id String @id @default(uuid()) + server_id String + tool_name String + version Int + description String @default("") + input_schema Json @default("{}") + change_kind String + changes Json @default("[]") + changelog String? + deprecated_at DateTime? + sunset_date DateTime? + deprecation_note String? + created_at DateTime @default(now()) + created_by String? + + @@unique([server_id, tool_name, version]) + @@index([server_id]) +} + // Named collection of {server_id, tool_name} pairs that can be granted to keys/teams model LiteLLM_MCPToolsetTable { toolset_id String @id @default(uuid()) diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index 4e511a2ec93..802789af5ef 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -71,6 +71,10 @@ class MCPServerRepository(PrismaTableRepository["prisma_models.LiteLLM_MCPServer table_name = "litellm_mcpservertable" +class MCPToolVersionRepository(PrismaTableRepository["prisma_models.LiteLLM_MCPToolVersion"]): + table_name = "litellm_mcptoolversion" + + class ManagedObjectRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedObjectTable"]): table_name = "litellm_managedobjecttable" diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 2ee19b3e59a..3b5bf9ebc3e 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -78,6 +78,47 @@ class PinnedMCPTool(BaseModel): input_schema: dict[str, object] = Field(default_factory=dict) +MCPToolChangeKind = Literal["initial", "non_breaking", "breaking", "removed"] + + +class MCPToolChange(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + breaking: bool + summary: str + + +class MCPToolVersion(BaseModel): + model_config = ConfigDict(frozen=True) + + server_id: str + tool_name: str + version: int + description: str = "" + input_schema: dict[str, object] = Field(default_factory=dict) + change_kind: MCPToolChangeKind + changes: tuple[MCPToolChange, ...] = () + changelog: str | None = None + deprecated_at: datetime | None = None + sunset_date: datetime | None = None + deprecation_note: str | None = None + created_at: datetime + created_by: str | None = None + + +class PinMCPServerToolsRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + changelog: str | None = None + + +class MCPToolDeprecationRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + sunset_date: datetime | None = None + deprecation_note: str | None = None + + _PINNED_TOOLS: Final[TypeAdapter[dict[str, PinnedMCPTool] | None]] = TypeAdapter(dict[str, PinnedMCPTool] | None) diff --git a/schema.prisma b/schema.prisma index aba89526cf6..33de27cd893 100644 --- a/schema.prisma +++ b/schema.prisma @@ -427,6 +427,26 @@ model LiteLLM_MCPServerTable { @@index([approval_status]) } +model LiteLLM_MCPToolVersion { + id String @id @default(uuid()) + server_id String + tool_name String + version Int + description String @default("") + input_schema Json @default("{}") + change_kind String + changes Json @default("[]") + changelog String? + deprecated_at DateTime? + sunset_date DateTime? + deprecation_note String? + created_at DateTime @default(now()) + created_by String? + + @@unique([server_id, tool_name, version]) + @@index([server_id]) +} + // Named collection of {server_id, tool_name} pairs that can be granted to keys/teams model LiteLLM_MCPToolsetTable { toolset_id String @id @default(uuid()) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py index fff4221f243..7616d4134c7 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -874,6 +874,7 @@ def _mock_env_vars_prisma(row=None): prisma.db.litellm_mcpuserenvvars.upsert = AsyncMock() prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock() prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock() + prisma.db.litellm_mcptoolversion.delete_many = AsyncMock() return prisma @@ -1254,6 +1255,38 @@ async def test_delete_mcp_server_succeeds_when_orphan_cleanup_fails(): prisma.db.litellm_mcpuserenvvars.delete_many.assert_awaited_once() +@pytest.mark.asyncio +async def test_delete_mcp_server_removes_tool_version_rows(): + from unittest.mock import AsyncMock + + from litellm.proxy._experimental.mcp_server.db import delete_mcp_server + + prisma: Final = _mock_env_vars_prisma() + prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=object()) + + await delete_mcp_server(prisma, "srv-1") + + prisma.db.litellm_mcptoolversion.delete_many.assert_awaited_once_with(where={"server_id": "srv-1"}) + + +@pytest.mark.asyncio +async def test_delete_mcp_server_succeeds_when_tool_version_cleanup_fails(): + from unittest.mock import AsyncMock + + from litellm.proxy._experimental.mcp_server.db import delete_mcp_server + + deleted: Final = object() + prisma: Final = _mock_env_vars_prisma() + prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=deleted) + prisma.db.litellm_mcptoolversion.delete_many = AsyncMock(side_effect=RuntimeError("connection pool exhausted")) + + result: Final = await delete_mcp_server(prisma, "srv-1") + + assert result is deleted + prisma.db.litellm_mcptoolversion.delete_many.assert_awaited_once_with(where={"server_id": "srv-1"}) + prisma.db.litellm_mcpuserenvvars.delete_many.assert_awaited_once() + + @pytest.mark.asyncio async def test_delete_mcp_server_removes_orphaned_user_credentials(): """Deleting a server must also drop every user's stored BYOK/OAuth credential diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py index a3c52dc16b7..95694c66238 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py @@ -9,19 +9,28 @@ must be cleared rather than left at its stored value. """ import json -from unittest.mock import AsyncMock, MagicMock +from datetime import datetime, timezone +from typing import Final +from unittest.mock import AsyncMock, MagicMock, call import pytest -from prisma import Json, models from fastapi import HTTPException +from prisma import Json, models +from prisma.errors import UniqueViolationError from litellm.proxy._experimental.mcp_server.db import ( + _MCP_PIN_ADVISORY_LOCK_SQL, create_mcp_server, set_mcp_server_pinned_tools, update_mcp_server, ) from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest -from litellm.types.mcp_server.mcp_server_manager import PinnedMCPTool +from litellm.types.mcp_server.mcp_server_manager import ( + MCPToolChangeKind, + MCPToolDeprecationRequest, + MCPToolVersion, + PinnedMCPTool, +) def _credentials_cleared(value) -> bool: @@ -37,16 +46,44 @@ def _mock_prisma(): mock_prisma.db.litellm_mcpservertable.create = AsyncMock(return_value=row) mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=None) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=row) + mock_prisma.db.litellm_mcptoolversion = AsyncMock() + mock_prisma.db.litellm_mcptoolversion.create_many = AsyncMock(return_value=0) + mock_prisma.db.litellm_mcptoolversion.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_mcptoolversion.find_unique = AsyncMock(return_value=None) + mock_prisma.db.litellm_mcptoolversion.update = AsyncMock(return_value=None) + mock_prisma.db.litellm_mcptoolversion.delete_many = AsyncMock(return_value=0) tx_client = MagicMock() tx_client.execute_raw = AsyncMock() - tx_client.litellm_mcpservertable = mock_prisma.db.litellm_mcpservertable + tx_client.litellm_mcpservertable = AsyncMock() + tx_client.litellm_mcpservertable.update = AsyncMock(return_value=row) + tx_client.litellm_mcpservertable.find_unique = AsyncMock(return_value=None) + tx_client.litellm_mcptoolversion = AsyncMock() + tx_client.litellm_mcptoolversion.create_many = AsyncMock(return_value=0) + tx_client.litellm_mcptoolversion.find_many = AsyncMock(return_value=[]) tx = MagicMock() tx.__aenter__ = AsyncMock(return_value=tx_client) tx.__aexit__ = AsyncMock(return_value=False) - mock_prisma.db.tx = MagicMock(return_value=tx) + mock_prisma.tx = MagicMock(return_value=tx) + db_tx_client = MagicMock() + db_tx_client.execute_raw = AsyncMock() + db_tx_client.litellm_mcpservertable = mock_prisma.db.litellm_mcpservertable + db_tx = MagicMock() + db_tx.__aenter__ = AsyncMock(return_value=db_tx_client) + db_tx.__aexit__ = AsyncMock(return_value=False) + mock_prisma.db.tx = MagicMock(return_value=db_tx) return mock_prisma +@pytest.fixture(autouse=True) +def mock_config_sync_publish(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: + publish: Final = AsyncMock() + monkeypatch.setattr( + "litellm.proxy.common_utils.config_sync_pubsub.publish_config_change_for_object_type", + publish, + ) + return publish + + async def _run_update(data: UpdateMCPServerRequest, fields_set=None) -> dict: mock_prisma = _mock_prisma() await update_mcp_server(mock_prisma, data, "test-user", fields_set=fields_set) @@ -1122,30 +1159,106 @@ async def test_register_and_update_bodies_never_write_pinned_tools(): @pytest.mark.asyncio async def test_set_mcp_server_pinned_tools_writes_the_snapshot_and_null_clears_it(): mock_prisma = _mock_prisma() - mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=MagicMock()) + tx_client = mock_prisma.tx.return_value.__aenter__.return_value + tx_client.litellm_mcpservertable.find_unique = AsyncMock(return_value=MagicMock()) pinned = {"list_notes": PinnedMCPTool(description="List notes", input_schema={"type": "object"})} record = await set_mcp_server_pinned_tools(mock_prisma, "test-server", pinned, "admin") - written = mock_prisma.db.litellm_mcpservertable.update.call_args[1] + written = tx_client.litellm_mcpservertable.update.call_args[1] assert written["where"] == {"server_id": "test-server"} assert json.loads(written["data"]["pinned_tools"]) == { "list_notes": {"description": "List notes", "input_schema": {"type": "object"}} } assert written["data"]["updated_by"] == "admin" assert record is not None and record.server_id == "test-server" + version_row: Final = tx_client.litellm_mcptoolversion.create_many.call_args.kwargs["data"][0] + assert isinstance(version_row["input_schema"], str) + assert json.loads(version_row["input_schema"]) == {"type": "object"} + assert isinstance(version_row["changes"], str) + assert json.loads(version_row["changes"]) == [] await set_mcp_server_pinned_tools(mock_prisma, "test-server", None, "admin") - assert mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]["pinned_tools"] == "{}" + assert tx_client.litellm_mcpservertable.update.call_args[1]["data"]["pinned_tools"] == "{}" + + +@pytest.mark.asyncio +async def test_set_mcp_server_pinned_tools_uses_one_transaction_for_the_pin( + mock_config_sync_publish: AsyncMock, +) -> None: + mock_prisma: Final = _mock_prisma() + calls: Final = MagicMock() + tx: Final = mock_prisma.tx.return_value + tx_client: Final = tx.__aenter__.return_value + updated_row: Final = tx_client.litellm_mcpservertable.update.return_value + tx_client.execute_raw.return_value = 1 + tx_client.litellm_mcpservertable.find_unique.return_value = MagicMock() + tx_client.litellm_mcptoolversion.find_many.return_value = [] + tx_client.litellm_mcptoolversion.create_many.return_value = 1 + tx_client.litellm_mcpservertable.update.return_value = updated_row + tx.__aexit__.return_value = False + calls.attach_mock(tx_client.execute_raw, "lock") + calls.attach_mock(tx_client.litellm_mcpservertable.find_unique, "server_lookup") + calls.attach_mock(tx_client.litellm_mcptoolversion.find_many, "history_read") + calls.attach_mock(tx_client.litellm_mcptoolversion.create_many, "version_insert") + calls.attach_mock(tx_client.litellm_mcpservertable.update, "server_update") + calls.attach_mock(tx.__aexit__, "transaction_exit") + calls.attach_mock(mock_config_sync_publish, "publish") + + await set_mcp_server_pinned_tools( + mock_prisma, + "test-server", + {"new": PinnedMCPTool()}, + "admin", + ) + + tx_client.execute_raw.assert_awaited_once_with(_MCP_PIN_ADVISORY_LOCK_SQL, "mcp_pin:test-server") + tx_client.litellm_mcpservertable.find_unique.assert_awaited_once_with(where={"server_id": "test-server"}) + tx_client.litellm_mcptoolversion.find_many.assert_awaited_once_with( + where={"server_id": "test-server"}, + order=[{"tool_name": "asc"}, {"version": "desc"}], + ) + tx_client.litellm_mcptoolversion.create_many.assert_awaited_once() + tx_client.litellm_mcpservertable.update.assert_awaited_once() + assert [entry[0] for entry in calls.mock_calls] == [ + "lock", + "server_lookup", + "history_read", + "version_insert", + "server_update", + "transaction_exit", + "publish", + ] + assert calls.mock_calls[0] == call.lock(_MCP_PIN_ADVISORY_LOCK_SQL, "mcp_pin:test-server") + assert mock_prisma.db.litellm_mcpservertable.find_unique.await_count == 0 + assert mock_prisma.db.litellm_mcpservertable.update.await_count == 0 + assert mock_prisma.db.litellm_mcptoolversion.find_many.await_count == 0 + assert mock_prisma.db.litellm_mcptoolversion.create_many.await_count == 0 + mock_config_sync_publish.assert_awaited_once_with("litellm_mcpservertable") + + +@pytest.mark.asyncio +async def test_set_mcp_server_pinned_tools_does_not_publish_when_snapshot_update_fails( + mock_config_sync_publish: AsyncMock, +) -> None: + mock_prisma: Final = _mock_prisma() + tx_client: Final = mock_prisma.tx.return_value.__aenter__.return_value + tx_client.litellm_mcpservertable.find_unique = AsyncMock(return_value=MagicMock()) + tx_client.litellm_mcpservertable.update = AsyncMock(side_effect=RuntimeError("update failed")) + + with pytest.raises(RuntimeError, match="update failed"): + await set_mcp_server_pinned_tools(mock_prisma, "test-server", None, "admin") + + mock_config_sync_publish.assert_not_awaited() @pytest.mark.asyncio async def test_set_mcp_server_pinned_tools_on_a_missing_server_writes_nothing(): mock_prisma = _mock_prisma() - mock_prisma.db.litellm_mcpservertable.find_unique.return_value = None + tx_client = mock_prisma.tx.return_value.__aenter__.return_value assert await set_mcp_server_pinned_tools(mock_prisma, "ghost", None, "admin") is None - mock_prisma.db.litellm_mcpservertable.update.assert_not_awaited() + tx_client.litellm_mcpservertable.update.assert_not_awaited() @pytest.mark.asyncio @@ -1181,3 +1294,207 @@ async def test_protocol_update_preserves_missing_server_without_writing(clear_al result = await update_mcp_server(prisma, payload, "admin") assert result is None table.update.assert_not_awaited() + + +def _mcp_tool_version( + tool_name: str, + version: int, + change_kind: MCPToolChangeKind, + description: str = "", + input_schema: dict[str, object] | None = None, + deprecated_at: datetime | None = None, +) -> MCPToolVersion: + return MCPToolVersion( + server_id="test-server", + tool_name=tool_name, + version=version, + description=description, + input_schema=input_schema if input_schema is not None else {}, + change_kind=change_kind, + changes=(), + created_at=datetime(2026, 10, 2, tzinfo=timezone.utc), + deprecated_at=deprecated_at, + ) + + +@pytest.mark.asyncio +async def test_set_mcp_server_pinned_tools_records_only_changed_versions_before_updating_pin(): + mock_prisma: Final = _mock_prisma() + tx_client: Final = mock_prisma.tx.return_value.__aenter__.return_value + tx_client.litellm_mcpservertable.find_unique = AsyncMock(return_value=MagicMock()) + tx_client.litellm_mcptoolversion.find_many = AsyncMock( + return_value=[ + _mcp_tool_version("changed", 2, "non_breaking", input_schema={"properties": {"q": {"type": "string"}}}), + _mcp_tool_version("changed", 1, "initial"), + _mcp_tool_version("stable", 1, "initial", description="Stable"), + ] + ) + pinned: Final = { + "changed": PinnedMCPTool( + description="Changed", + input_schema={"properties": {"q": {"type": "integer"}}}, + ), + "stable": PinnedMCPTool(description="Stable"), + } + + await set_mcp_server_pinned_tools( + mock_prisma, + "test-server", + pinned, + "admin-user", + changelog="Updated the changed tool", + ) + + inserted_rows: Final = tx_client.litellm_mcptoolversion.create_many.call_args.kwargs["data"] + expected_changes: Final = [ + {"breaking": False, "summary": "Description changed"}, + { + "breaking": True, + "summary": 'Parameter "q" type changed from "string" to "integer"', + }, + ] + assert inserted_rows == ( + { + "server_id": "test-server", + "tool_name": "changed", + "version": 3, + "description": "Changed", + "input_schema": inserted_rows[0]["input_schema"], + "change_kind": "breaking", + "changes": inserted_rows[0]["changes"], + "changelog": "Updated the changed tool", + "created_by": "admin-user", + }, + ) + assert isinstance(inserted_rows[0]["input_schema"], str) + assert json.loads(inserted_rows[0]["input_schema"]) == {"properties": {"q": {"type": "integer"}}} + assert isinstance(inserted_rows[0]["changes"], str) + assert json.loads(inserted_rows[0]["changes"]) == expected_changes + assert tx_client.litellm_mcpservertable.update.await_count == 1 + + +@pytest.mark.asyncio +async def test_set_mcp_server_pinned_tools_does_not_update_pin_when_version_insert_fails(): + mock_prisma: Final = _mock_prisma() + tx_client: Final = mock_prisma.tx.return_value.__aenter__.return_value + tx_client.litellm_mcpservertable.find_unique = AsyncMock(return_value=MagicMock()) + tx_client.litellm_mcptoolversion.create_many = AsyncMock( + side_effect=UniqueViolationError({}, message="duplicate version") + ) + + with pytest.raises(UniqueViolationError): + await set_mcp_server_pinned_tools( + mock_prisma, + "test-server", + {"new": PinnedMCPTool()}, + "admin", + ) + + tx_client.litellm_mcptoolversion.create_many.assert_awaited_once() + tx_client.litellm_mcpservertable.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_list_mcp_tool_versions_orders_tools_and_versions(): + mock_prisma: Final = _mock_prisma() + rows: Final = [ + _mcp_tool_version("alpha", 2, "breaking"), + _mcp_tool_version("alpha", 1, "initial"), + _mcp_tool_version("beta", 1, "initial"), + ] + mock_prisma.db.litellm_mcptoolversion.find_many = AsyncMock(return_value=rows) + + from litellm.proxy._experimental.mcp_server.db import list_mcp_tool_versions + + assert await list_mcp_tool_versions(mock_prisma, "test-server") == rows + mock_prisma.db.litellm_mcptoolversion.find_many.assert_awaited_once_with( + where={"server_id": "test-server"}, + order=[{"tool_name": "asc"}, {"version": "desc"}], + ) + + +@pytest.mark.asyncio +async def test_set_mcp_tool_version_deprecation_preserves_existing_timestamp(): + mock_prisma: Final = _mock_prisma() + calls: Final = MagicMock() + tx_client: Final = mock_prisma.tx.return_value.__aenter__.return_value + deprecated_at: Final = datetime(2026, 10, 1, tzinfo=timezone.utc) + existing: Final = _mcp_tool_version("alpha", 2, "breaking", deprecated_at=deprecated_at) + updated: Final = existing.model_copy( + update={ + "sunset_date": datetime(2026, 12, 31, tzinfo=timezone.utc), + "deprecation_note": "Switch to v3", + } + ) + tx_client.execute_raw = AsyncMock(return_value=1) + tx_client.litellm_mcptoolversion.find_unique = AsyncMock(return_value=existing) + tx_client.litellm_mcptoolversion.update = AsyncMock(return_value=updated) + calls.attach_mock(tx_client.execute_raw, "lock") + calls.attach_mock(tx_client.litellm_mcptoolversion.find_unique, "lookup") + calls.attach_mock(tx_client.litellm_mcptoolversion.update, "update") + + from litellm.proxy._experimental.mcp_server.db import set_mcp_tool_version_deprecation + + result: Final = await set_mcp_tool_version_deprecation( + mock_prisma, + "test-server", + "alpha", + 2, + MCPToolDeprecationRequest( + sunset_date=datetime(2026, 12, 31, tzinfo=timezone.utc), + deprecation_note="Switch to v3", + ), + ) + + assert result == updated + tx_client.litellm_mcptoolversion.update.assert_awaited_once_with( + where={ + "server_id_tool_name_version": { + "server_id": "test-server", + "tool_name": "alpha", + "version": 2, + } + }, + data={ + "deprecated_at": deprecated_at, + "sunset_date": datetime(2026, 12, 31, tzinfo=timezone.utc), + "deprecation_note": "Switch to v3", + }, + ) + assert [entry[0] for entry in calls.mock_calls] == ["lock", "lookup", "update"] + assert calls.mock_calls[0] == call.lock(_MCP_PIN_ADVISORY_LOCK_SQL, "mcp_pin:test-server") + mock_prisma.db.litellm_mcptoolversion.find_unique.assert_not_awaited() + mock_prisma.db.litellm_mcptoolversion.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_set_mcp_tool_version_deprecation_clears_fields_or_returns_none_for_missing_row(): + mock_prisma: Final = _mock_prisma() + tx_client: Final = mock_prisma.tx.return_value.__aenter__.return_value + existing: Final = _mcp_tool_version( + "alpha", + 2, + "breaking", + deprecated_at=datetime(2026, 10, 1, tzinfo=timezone.utc), + ) + tx_client.litellm_mcptoolversion.find_unique = AsyncMock(return_value=existing) + tx_client.litellm_mcptoolversion.update = AsyncMock( + return_value=existing.model_copy(update={"deprecated_at": None, "sunset_date": None, "deprecation_note": None}) + ) + + from litellm.proxy._experimental.mcp_server.db import set_mcp_tool_version_deprecation + + result: Final = await set_mcp_tool_version_deprecation(mock_prisma, "test-server", "alpha", 2, None) + + assert result is not None + assert (result.deprecated_at, result.sunset_date, result.deprecation_note) == (None, None, None) + assert tx_client.litellm_mcptoolversion.update.call_args.kwargs["data"] == { + "deprecated_at": None, + "sunset_date": None, + "deprecation_note": None, + } + + tx_client.litellm_mcptoolversion.find_unique = AsyncMock(return_value=None) + assert await set_mcp_tool_version_deprecation(mock_prisma, "test-server", "missing", 3, None) is None + tx_client.litellm_mcptoolversion.update.assert_awaited_once() + mock_prisma.db.litellm_mcptoolversion.find_unique.assert_not_awaited() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_tool_versioning.py b/tests/unit/proxy/_experimental/mcp_server/test_tool_versioning.py new file mode 100644 index 00000000000..d9409f8fc9d --- /dev/null +++ b/tests/unit/proxy/_experimental/mcp_server/test_tool_versioning.py @@ -0,0 +1,359 @@ +from datetime import datetime, timezone +from typing import Final + +import pytest + +from litellm.proxy._experimental.mcp_server.tool_versioning import ( + PlannedToolVersion, + classify_tool_change, + plan_tool_versions, +) +from litellm.types.mcp_server.mcp_server_manager import ( + MCPToolChange, + MCPToolChangeKind, + MCPToolVersion, + PinnedMCPTool, +) + + +def _tool( + description: str = "", + input_schema: dict[str, object] | None = None, +) -> PinnedMCPTool: + return PinnedMCPTool(description=description, input_schema=input_schema if input_schema is not None else {}) + + +def _version( + tool_name: str, + version: int, + change_kind: MCPToolChangeKind, + description: str = "", + input_schema: dict[str, object] | None = None, +) -> MCPToolVersion: + return MCPToolVersion( + server_id="srv", + tool_name=tool_name, + version=version, + description=description, + input_schema=input_schema if input_schema is not None else {}, + change_kind=change_kind, + created_at=datetime(2026, 10, 2, tzinfo=timezone.utc), + ) + + +@pytest.mark.parametrize( + ("previous", "current", "expected"), + [ + ( + _tool("Before"), + _tool("After"), + (MCPToolChange(breaking=False, summary="Description changed"),), + ), + ( + _tool(input_schema={"properties": {"a": {"type": "string"}}}), + _tool(input_schema={"properties": {"a": {"type": "string"}}, "required": ["a"]}), + (MCPToolChange(breaking=True, summary='Parameter "a" is now required'),), + ), + ( + _tool(input_schema={"type": "object"}), + _tool(input_schema={"type": "object", "required": ["token"]}), + (MCPToolChange(breaking=True, summary='Parameter "token" is now required'),), + ), + ( + _tool(input_schema={"type": "object", "required": ["token"]}), + _tool(input_schema={"type": "object"}), + (MCPToolChange(breaking=False, summary='Parameter "token" is now optional'),), + ), + ( + _tool(input_schema={"properties": {"a": {"type": "string"}}}), + _tool(input_schema={"properties": {"b": {"type": "string"}}}), + ( + MCPToolChange(breaking=True, summary='Removed parameter "a"'), + MCPToolChange(breaking=False, summary='Added optional parameter "b"'), + ), + ), + ( + _tool(input_schema={"properties": {}}), + _tool(input_schema={"properties": {"a": {"type": "string"}}, "required": ["a"]}), + (MCPToolChange(breaking=True, summary='Added required parameter "a"'),), + ), + ( + _tool(input_schema={"properties": {"a": {"type": "string"}}}), + _tool(input_schema={"properties": {"a": {"type": "integer"}}}), + ( + MCPToolChange( + breaking=True, + summary='Parameter "a" type changed from "string" to "integer"', + ), + ), + ), + ( + _tool(input_schema={"properties": {"a": {}}}), + _tool(input_schema={"properties": {"a": {"type": "string"}}}), + (MCPToolChange(breaking=True, summary='Parameter "a" type changed from any to "string"'),), + ), + ( + _tool(input_schema={"properties": {"a": {"type": "string", "enum": ["a"]}}}), + _tool(input_schema={"properties": {"a": {"type": "string", "enum": ["a", "b"]}}}), + (MCPToolChange(breaking=True, summary='Parameter "a" schema changed'),), + ), + ( + _tool(input_schema={"properties": {"a": {"type": "string", "description": "Before"}}}), + _tool(input_schema={"properties": {"a": {"type": "string", "description": "After"}}}), + (MCPToolChange(breaking=False, summary='Parameter "a" description changed'),), + ), + ( + _tool(input_schema={"properties": {"a": {"type": "string"}}, "required": ["a"]}), + _tool(input_schema={"properties": {"a": {"type": "string"}}}), + (MCPToolChange(breaking=False, summary='Parameter "a" is now optional'),), + ), + ( + _tool(input_schema={"properties": {"a": {"type": "string", "title": "Before"}}}), + _tool(input_schema={"properties": {"a": {"type": "string", "title": "After"}}}), + (MCPToolChange(breaking=False, summary='Parameter "a" description changed'),), + ), + ( + _tool( + input_schema={ + "properties": { + "filter": { + "type": "object", + "properties": {"term": {"type": "string", "description": "Before"}}, + } + } + } + ), + _tool( + input_schema={ + "properties": { + "filter": {"type": "object", "properties": {"term": {"type": "string", "description": "After"}}} + } + } + ), + (MCPToolChange(breaking=False, summary='Parameter "filter" description changed'),), + ), + ( + _tool( + input_schema={"properties": {"filter": {"type": "object", "properties": {"term": {"type": "string"}}}}} + ), + _tool( + input_schema={"properties": {"filter": {"type": "object", "properties": {"term": {"type": "integer"}}}}} + ), + (MCPToolChange(breaking=True, summary='Parameter "filter" schema changed'),), + ), + ( + _tool( + input_schema={"properties": {"filter": {"type": "object", "properties": {"title": {"type": "string"}}}}} + ), + _tool(input_schema={"properties": {"filter": {"type": "object", "properties": {}}}}), + (MCPToolChange(breaking=True, summary='Parameter "filter" schema changed'),), + ), + ( + _tool( + input_schema={ + "properties": {"items": {"type": "array", "items": {"type": "string", "description": "Before"}}} + } + ), + _tool( + input_schema={ + "properties": {"items": {"type": "array", "items": {"type": "string", "description": "After"}}} + } + ), + (MCPToolChange(breaking=False, summary='Parameter "items" description changed'),), + ), + ( + _tool(input_schema={"properties": None}), + _tool(input_schema={"properties": ["invalid"]}), + (), + ), + ( + _tool(input_schema={"type": "object"}), + _tool(input_schema={"type": "array"}), + (MCPToolChange(breaking=True, summary="Input schema changed"),), + ), + ( + _tool(input_schema={"type": "object", "description": "Before"}), + _tool(input_schema={"type": "object", "description": "After"}), + (MCPToolChange(breaking=False, summary="Input schema description changed"),), + ), + ( + _tool(input_schema={"type": "object", "title": "Before"}), + _tool(input_schema={"type": "object", "title": "After"}), + (MCPToolChange(breaking=False, summary="Input schema title changed"),), + ), + ( + _tool(input_schema={"type": "object", "$defs": {"filter": {"type": "string", "description": "Before"}}}), + _tool(input_schema={"type": "object", "$defs": {"filter": {"type": "string", "description": "After"}}}), + (MCPToolChange(breaking=False, summary="Input schema documentation changed"),), + ), + ( + _tool(input_schema={"type": "object", "$defs": {"filter": {"type": "string"}}}), + _tool(input_schema={"type": "object", "$defs": {"filter": {"type": "integer"}}}), + (MCPToolChange(breaking=True, summary="Input schema changed"),), + ), + ], +) +def test_classify_tool_change( + previous: PinnedMCPTool, + current: PinnedMCPTool, + expected: tuple[MCPToolChange, ...], +) -> None: + assert classify_tool_change(previous, current) == expected + + +def test_schema_draft_change_is_breaking_and_recorded() -> None: + previous_schema: Final = {"$schema": "http://json-schema.org/draft-07/schema#"} + current: Final = _tool(input_schema={"$schema": "https://json-schema.org/draft/2020-12/schema"}) + expected_change: Final = MCPToolChange(breaking=True, summary="Input schema changed") + previous: Final = _tool(input_schema=previous_schema) + latest: Final = {"tool": _version("tool", 3, "initial", input_schema=previous_schema)} + + assert classify_tool_change(previous, current) == (expected_change,) + assert plan_tool_versions(latest, {"tool": current}) == ( + PlannedToolVersion("tool", 4, current, "breaking", (expected_change,)), + ) + + +def test_classify_tool_change_reports_type_changes_before_schema_changes(): + previous: Final = _tool( + description="Before", + input_schema={ + "properties": { + "z": {"type": "string"}, + "a": {"type": "string", "enum": ["a"]}, + }, + "type": "object", + }, + ) + current: Final = _tool( + description="After", + input_schema={ + "properties": { + "z": {"type": "integer"}, + "a": {"type": "integer", "enum": ["b"]}, + }, + "type": "array", + }, + ) + + assert classify_tool_change(previous, current) == ( + MCPToolChange(breaking=False, summary="Description changed"), + MCPToolChange( + breaking=True, + summary='Parameter "a" type changed from "string" to "integer"', + ), + MCPToolChange( + breaking=True, + summary='Parameter "z" type changed from "string" to "integer"', + ), + MCPToolChange(breaking=True, summary="Input schema changed"), + ) + + +def test_plan_tool_versions_creates_initial_versions_and_sorts_by_name(): + snapshot: Final = {"z": _tool("Z"), "a": _tool("A")} + + assert plan_tool_versions({}, snapshot) == ( + PlannedToolVersion("a", 1, _tool("A"), "initial", ()), + PlannedToolVersion("z", 1, _tool("Z"), "initial", ()), + ) + + +def test_plan_tool_versions_skips_unchanged_tools(): + latest: Final = {"a": _version("a", 3, "non_breaking", "A")} + + assert plan_tool_versions(latest, {"a": _tool("A")}) == () + + +@pytest.mark.parametrize( + ("current", "expected_kind", "expected_change"), + [ + ( + _tool(input_schema={"properties": {"a": {"type": "integer"}}}), + "breaking", + MCPToolChange( + breaking=True, + summary='Parameter "a" type changed from "string" to "integer"', + ), + ), + ( + _tool("After", input_schema={"properties": {"a": {"type": "string"}}}), + "non_breaking", + MCPToolChange(breaking=False, summary="Description changed"), + ), + ], +) +def test_plan_tool_versions_bumps_only_changed_tools( + current: PinnedMCPTool, + expected_kind: MCPToolChangeKind, + expected_change: MCPToolChange, +) -> None: + latest: Final = { + "a": _version("a", 4, "initial", input_schema={"properties": {"a": {"type": "string"}}}), + "stable": _version("stable", 2, "initial"), + } + snapshot: Final = {"a": current, "stable": _tool()} + + assert plan_tool_versions(latest, snapshot) == ( + PlannedToolVersion("a", 5, current, expected_kind, (expected_change,)), + ) + + +def test_plan_tool_versions_marks_required_only_parameter_change_as_breaking(): + latest: Final = { + "a": _version("a", 1, "initial", input_schema={"type": "object"}), + } + snapshot: Final = { + "a": _tool(input_schema={"type": "object", "required": ["token"]}), + } + + planned: Final = plan_tool_versions(latest, snapshot) + + assert planned[0].change_kind == "breaking" + + +def test_plan_tool_versions_marks_removed_tools_only_once_and_restores_changed_tools(): + latest: Final = { + "gone": _version( + "gone", + 1, + "initial", + input_schema={"properties": {"a": {"type": "string"}}}, + ) + } + + assert plan_tool_versions(latest, {}) == ( + PlannedToolVersion( + "gone", + 2, + _tool(input_schema={"properties": {"a": {"type": "string"}}}), + "removed", + (MCPToolChange(breaking=True, summary="Tool removed"),), + ), + ) + assert plan_tool_versions({"gone": _version("gone", 2, "removed", "Before")}, {}) == () + assert plan_tool_versions( + { + "gone": _version( + "gone", + 2, + "removed", + input_schema={"properties": {"a": {"type": "string"}}}, + ) + }, + {"gone": _tool(input_schema={"properties": {"a": {"type": "integer"}}})}, + ) == ( + PlannedToolVersion( + "gone", + 3, + _tool(input_schema={"properties": {"a": {"type": "integer"}}}), + "breaking", + ( + MCPToolChange(breaking=False, summary="Tool restored"), + MCPToolChange( + breaking=True, + summary='Parameter "a" type changed from "string" to "integer"', + ), + ), + ), + ) diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index 2cbf1d578b2..b6f689fbedc 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1,41 +1,38 @@ import asyncio +import json +import logging import os import sys import types -import json -import logging from collections.abc import Iterator, Mapping from contextlib import ExitStack, contextmanager from dataclasses import dataclass, field -from datetime import datetime, timedelta +from datetime import datetime, timedelta, timezone from types import SimpleNamespace -from typing import Final, List, Literal, Optional, cast +from typing import Final, Literal, cast from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest -from pydantic import BaseModel, TypeAdapter, ValidationError -from respx import MockRouter from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient +from prisma.errors import UniqueViolationError +from pydantic import BaseModel, TypeAdapter, ValidationError +from respx import MockRouter import litellm from litellm._uuid import uuid from litellm.caching.caching import DualCache -from litellm.integrations.custom_guardrail import CustomGuardrail -from litellm.proxy.utils import ProxyLogging from litellm.constants import UI_SESSION_TOKEN_TEAM_ID +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.models.access_group import LiteLLM_AccessGroupTable from litellm.models.organization import LiteLLM_OrganizationTable from litellm.models.team import LiteLLM_TeamTable from litellm.models.user import LiteLLM_UserTable -from litellm.proxy.management_endpoints import ( - mcp_management_endpoints as mgmt_endpoints, -) - +from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerConfig, MCPServerManager from litellm.proxy._types import ( - LiteLLM_ObjectPermissionTable, LiteLLM_MCPServerTable, + LiteLLM_ObjectPermissionTable, LitellmUserRoles, MakeMCPServersPublicRequest, MCPTransport, @@ -44,17 +41,26 @@ from litellm.proxy._types import ( UpdateMCPServerRequest, UserAPIKeyAuth, ) -from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerConfig, MCPServerManager +from litellm.proxy.management_endpoints import ( + mcp_management_endpoints as mgmt_endpoints, +) +from litellm.proxy.utils import ProxyLogging from litellm.types.mcp import MCPAuth, MCPCredentials -from litellm.types.mcp_server.mcp_server_manager import MCPServer, PinnedMCPTool +from litellm.types.mcp_server.mcp_server_manager import ( + MCPServer, + MCPToolDeprecationRequest, + MCPToolVersion, + PinMCPServerToolsRequest, + PinnedMCPTool, +) def generate_mock_mcp_server_db_record( - server_id: Optional[str] = None, + server_id: str | None = None, alias: str = "Test DB Server", url: str = "https://db-server.example.com/mcp", transport: str = "sse", - auth_type: Optional[str] = None, + auth_type: str | None = None, ) -> LiteLLM_MCPServerTable: """Generate a mock MCP server record from database""" now = datetime.now() @@ -72,11 +78,11 @@ def generate_mock_mcp_server_db_record( def generate_mock_mcp_server_config_record( - server_id: Optional[str] = None, + server_id: str | None = None, name: str = "Test Config Server", url: str = "https://config-server.example.com/mcp", transport: str = "http", - auth_type: Optional[str] = None, + auth_type: str | None = None, ) -> MCPServer: """Generate a mock MCP server record from config.yaml""" return MCPServer( @@ -107,7 +113,7 @@ def generate_mock_user_api_key_auth( user_role: LitellmUserRoles = LitellmUserRoles.PROXY_ADMIN, user_id: str = "test_user_id", api_key: str = "test_api_key", - team_id: Optional[str] = None, + team_id: str | None = None, ) -> UserAPIKeyAuth: """Generate a mock UserAPIKeyAuth object""" return UserAPIKeyAuth( @@ -118,7 +124,7 @@ def generate_mock_user_api_key_auth( ) -def generate_mock_team_record(team_id: str, team_alias: str, organization_id: str, mcp_servers: List[str]): +def generate_mock_team_record(team_id: str, team_alias: str, organization_id: str, mcp_servers: list[str]): """Generate a mock team record with object permissions""" return MagicMock( team_id=team_id, @@ -131,8 +137,8 @@ def generate_mock_team_record(team_id: str, team_alias: str, organization_id: st def setup_mock_prisma_client( mock_prisma_client: MagicMock, - team_records: List[MagicMock], - mcp_servers: List[LiteLLM_MCPServerTable], + team_records: list[MagicMock], + mcp_servers: list[LiteLLM_MCPServerTable], ): """Helper to set up a mock prisma client with proper async behavior""" mock_prisma_client.db = MagicMock() @@ -219,9 +225,7 @@ async def test_mcp_publication_list_and_detail_derive_current_status( @pytest.mark.parametrize("approval_status", ("pending_review", "rejected", "draft", "active")) @pytest.mark.parametrize("strict", (False, True)) -def test_mcp_publication_projection_excludes_unregistered_lifecycle_records( - approval_status: str, strict: bool -) -> None: +def test_mcp_publication_projection_excludes_unregistered_lifecycle_records(approval_status: str, strict: bool) -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager record: Final = LiteLLM_MCPServerTable( @@ -2238,10 +2242,10 @@ class TestTemporaryMCPSessionEndpoints: async def test_create_draft_mcp_server_prunes_drafts_past_their_lifetime(self): """Regression: abandoned OAuth sessions accumulated forever. Verified against a live proxy, where 12 drafts aged past the lifetime were still present and a 13th was added.""" - from litellm.proxy._experimental.mcp_server.db import create_draft_mcp_server - from datetime import timezone + from litellm.proxy._experimental.mcp_server.db import create_draft_mcp_server + now = datetime.now(timezone.utc) stale_one = generate_mock_mcp_server_db_record(server_id="stale-1") stale_one.updated_at = now - timedelta(hours=1) @@ -4512,8 +4516,7 @@ async def test_health_checks_probe_shared_servers_once_across_auth_contexts( for server_id in ("shared", "first", "second", "denied") } routes: Final = { - server_id: respx_mock.get(server.url).respond(401) - for server_id, server in manager.registry.items() + server_id: respx_mock.get(server.url).respond(401) for server_id, server in manager.registry.items() } contexts: Final = [ UserAPIKeyAuth( @@ -4526,15 +4529,9 @@ async def test_health_checks_probe_shared_servers_once_across_auth_contexts( for index, grants in enumerate((("shared", "first"), ("shared", "second"))) ] with ( - patch.object( - mgmt_endpoints, "global_mcp_server_manager", manager - ), - patch.object( - mcp_server_manager, "global_mcp_server_manager", manager - ), - patch.object( - mgmt_endpoints, "build_effective_auth_contexts", AsyncMock(return_value=contexts) - ), + patch.object(mgmt_endpoints, "global_mcp_server_manager", manager), + patch.object(mcp_server_manager, "global_mcp_server_manager", manager), + patch.object(mgmt_endpoints, "build_effective_auth_contexts", AsyncMock(return_value=contexts)), patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": "restricted"}), ): result: Final = await mgmt_endpoints.health_check_servers( @@ -4624,14 +4621,14 @@ async def test_health_reachability_requires_explicit_api_opt_in( assert response.status_code == 200, response.text rows: Final = ( [HealthResponse.model_validate_json(response.content)] - if detail else TypeAdapter(list[HealthResponse]).validate_json(response.content) + if detail + else TypeAdapter(list[HealthResponse]).validate_json(response.content) ) expected_status: Final = "reachable" if flag == "true" else "unknown" assert [row.model_dump() for row in rows] == [{"server_id": server.server_id, "status": expected_status}] assert route.call_count == 1 legacy_parser: Final = ( - LegacyHealthResponse.model_validate_json - if detail else TypeAdapter(list[LegacyHealthResponse]).validate_json + LegacyHealthResponse.model_validate_json if detail else TypeAdapter(list[LegacyHealthResponse]).validate_json ) if flag == "true": with pytest.raises(ValidationError, match="literal_error"): @@ -8549,7 +8546,11 @@ class TestPinMCPServerTools: """POST/DELETE /v1/mcp/server/{server_id}/pin snapshot and clear the served tool catalog.""" @staticmethod - def _pin_patches(stored, store_mock, manager): + def _pin_patches( + stored: LiteLLM_MCPServerTable | None, + store_mock: AsyncMock, + manager: MagicMock, + ): return ( patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True), patch( @@ -8560,14 +8561,18 @@ class TestPinMCPServerTools: "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", AsyncMock(return_value=stored), ), - patch("litellm.proxy.management_endpoints.mcp_management_endpoints.set_mcp_server_pinned_tools", store_mock), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.set_mcp_server_pinned_tools", store_mock + ), patch("litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", manager), patch("litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager", manager), patch.dict( sys.modules, { "litellm.proxy.proxy_server": types.SimpleNamespace( - proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), general_settings={}, llm_router=None + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + general_settings={}, + llm_router=None, ) }, ), @@ -8643,7 +8648,7 @@ class TestPinMCPServerTools: assert listing["raw_headers"] == request.headers assert listing["client_ip"] == "10.1.2.3" assert store_mock.await_args.args[1:] == ("srv-1", expected) - assert store_mock.await_args.kwargs == {"touched_by": "admin"} + assert store_mock.await_args.kwargs == {"touched_by": "admin", "changelog": None} manager.update_server.assert_awaited_once_with(stored) manager.reload_servers_from_database.assert_awaited_once() @@ -8663,7 +8668,7 @@ class TestPinMCPServerTools: assert result == {"server_id": "srv-1", "status": "unpinned"} assert store_mock.await_args.args[1:] == ("srv-1", None) - assert store_mock.await_args.kwargs == {"touched_by": "admin"} + assert store_mock.await_args.kwargs == {"touched_by": "admin", "changelog": None} manager._get_tools_from_server.assert_not_awaited() manager.reload_servers_from_database.assert_awaited_once() @@ -8752,6 +8757,363 @@ class TestPinMCPServerTools: store_mock.assert_not_awaited() +class TestMCPToolVersionEndpoints: + @pytest.mark.asyncio + async def test_pin_passes_trimmed_changelog_and_pin_without_body_remains_valid(self): + stored: Final = generate_mock_mcp_server_db_record(server_id="srv-1") + store_mock: Final = AsyncMock(return_value=stored) + manager: Final = TestPinMCPServerTools._manager([("list_notes", "List notes", {})]) + admin: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin") + + with ExitStack() as stack: + for patcher in TestPinMCPServerTools._pin_patches(stored, store_mock, manager): + stack.enter_context(patcher) + result: Final = await mgmt_endpoints.pin_mcp_server_tools( + server_id="srv-1", + request=_make_mock_request(), + payload=PinMCPServerToolsRequest(changelog=" Updated schema "), + user_api_key_dict=admin, + ) + + assert result == {"list_notes": PinnedMCPTool(description="List notes", input_schema={})} + assert store_mock.await_args.kwargs == {"touched_by": "admin", "changelog": "Updated schema"} + + @pytest.mark.asyncio + async def test_get_tool_versions_returns_rows_for_authorized_database_server(self): + prisma: Final = MagicMock() + server: Final = generate_mock_mcp_server_db_record(server_id="srv-1") + resolved: Final = mgmt_endpoints.ResolvedMCPServer(table=server, runtime=None, source="db") + version: Final = MCPToolVersion( + server_id="srv-1", + tool_name="tool_a", + version=1, + description="List notes", + input_schema={}, + change_kind="initial", + created_at=datetime(2026, 10, 2, tzinfo=timezone.utc), + ) + rows: Final = [version, version.model_copy(update={"tool_name": "tool_b"})] + list_mock: Final = AsyncMock(return_value=rows) + admin: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + allowed_mock: Final = AsyncMock(return_value=None) + + with ( + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=prisma), + patch.object(mgmt_endpoints, "resolve_mcp_server", AsyncMock(return_value=resolved)), + patch.object(mgmt_endpoints, "authorize_mcp_server", AsyncMock(return_value=resolved)), + patch.object(mgmt_endpoints, "list_mcp_tool_versions", list_mock), + patch.object(mgmt_endpoints, "_get_allowed_tool_names_for_server", allowed_mock), + patch.object(mgmt_endpoints, "_user_has_admin_view", return_value=True), + patch.object(mgmt_endpoints, "_is_restricted_virtual_key_request", return_value=False), + patch.object(mgmt_endpoints, "_get_user_mcp_management_mode", return_value="restricted"), + patch.object(mgmt_endpoints, "global_mcp_server_manager", MagicMock()), + ): + result: Final = await mgmt_endpoints.fetch_mcp_server_tool_versions( + request=_make_mock_request(), + server_id="srv-1", + user_api_key_dict=admin, + ) + + assert result == rows + list_mock.assert_awaited_once_with(prisma, "srv-1") + allowed_mock.assert_awaited_once_with(server_id="srv-1", user_api_key_dict=admin) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("allowed", "expected_names"), + ((["tool_a"], ["tool_a"]), ([], []), (None, ["tool_a", "tool_b"])), + ) + async def test_get_tool_versions_filters_rows_by_tool_grants( + self, + allowed: list[str] | None, + expected_names: list[str], + ) -> None: + prisma: Final = MagicMock() + server: Final = generate_mock_mcp_server_db_record(server_id="srv-1") + resolved: Final = mgmt_endpoints.ResolvedMCPServer(table=server, runtime=None, source="db") + rows: Final = [ + MCPToolVersion( + server_id="srv-1", + tool_name="tool_a", + version=1, + description="Tool A", + input_schema={}, + change_kind="initial", + created_at=datetime(2026, 10, 2, tzinfo=timezone.utc), + ), + MCPToolVersion( + server_id="srv-1", + tool_name="tool_b", + version=1, + description="Tool B", + input_schema={}, + change_kind="initial", + created_at=datetime(2026, 10, 2, tzinfo=timezone.utc), + ), + ] + list_mock: Final = AsyncMock(return_value=rows) + admin: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + allowed_mock: Final = AsyncMock(return_value=allowed) + + with ( + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=prisma), + patch.object( + mgmt_endpoints, "_resolve_and_authorize_mcp_server", AsyncMock(return_value=(resolved, False)) + ), + patch.object(mgmt_endpoints, "list_mcp_tool_versions", list_mock), + patch.object(mgmt_endpoints, "_get_allowed_tool_names_for_server", allowed_mock), + ): + result: Final = await mgmt_endpoints.fetch_mcp_server_tool_versions( + request=_make_mock_request(), + server_id="srv-1", + user_api_key_dict=admin, + ) + + assert [version.tool_name for version in result] == expected_names + allowed_mock.assert_awaited_once_with(server_id="srv-1", user_api_key_dict=admin) + + @pytest.mark.asyncio + async def test_get_tool_versions_returns_empty_for_authorized_config_server(self): + resolved: Final = mgmt_endpoints.ResolvedMCPServer( + table=generate_mock_mcp_server_db_record(server_id="config-server"), + runtime=None, + source="registry", + ) + list_mock: Final = AsyncMock() + admin: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()), + patch.object( + mgmt_endpoints, + "_resolve_and_authorize_mcp_server", + AsyncMock(return_value=(resolved, False)), + ), + patch.object(mgmt_endpoints, "list_mcp_tool_versions", list_mock), + ): + result: Final = await mgmt_endpoints.fetch_mcp_server_tool_versions( + request=_make_mock_request(), + server_id="config-server", + user_api_key_dict=admin, + ) + + assert result == [] + list_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_get_tool_versions_has_same_non_admin_missing_server_response_as_fetch(self): + prisma: Final = MagicMock() + not_found: Final = HTTPException( + status_code=404, + detail={"error": "MCP Server with id missing not found"}, + ) + resolve_mock: Final = AsyncMock(return_value=None) + authorize_mock: Final = AsyncMock(side_effect=not_found) + user: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.INTERNAL_USER) + + with ( + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=prisma), + patch.object(mgmt_endpoints, "resolve_mcp_server", resolve_mock), + patch.object(mgmt_endpoints, "authorize_mcp_server", authorize_mock), + patch.object(mgmt_endpoints, "_user_has_admin_view", return_value=False), + patch.object(mgmt_endpoints, "_is_restricted_virtual_key_request", return_value=False), + patch.object(mgmt_endpoints, "_get_user_mcp_management_mode", return_value="restricted"), + patch.object(mgmt_endpoints, "global_mcp_server_manager", MagicMock()), + ): + with pytest.raises(HTTPException) as versions_error: + await mgmt_endpoints.fetch_mcp_server_tool_versions( + request=_make_mock_request(), + server_id="missing", + user_api_key_dict=user, + ) + with pytest.raises(HTTPException) as server_error: + await mgmt_endpoints.fetch_mcp_server( + request=_make_mock_request(), + server_id="missing", + user_api_key_dict=user, + ) + + assert (versions_error.value.status_code, versions_error.value.detail) == ( + server_error.value.status_code, + server_error.value.detail, + ) + assert versions_error.value.status_code == 404 + assert authorize_mock.await_count == 2 + + @pytest.mark.asyncio + async def test_put_tool_version_deprecation_is_admin_only_and_returns_updated_version(self): + prisma: Final = MagicMock() + request: Final = MCPToolDeprecationRequest( + sunset_date=datetime(2026, 12, 31, tzinfo=timezone.utc), + deprecation_note="Use version 2", + ) + version: Final = MCPToolVersion( + server_id="srv-1", + tool_name="list_notes", + version=1, + change_kind="initial", + created_at=datetime(2026, 10, 2, tzinfo=timezone.utc), + deprecated_at=datetime(2026, 10, 2, tzinfo=timezone.utc), + sunset_date=request.sunset_date, + deprecation_note=request.deprecation_note, + ) + set_mock: Final = AsyncMock(return_value=version) + admin: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=prisma), + patch.object(mgmt_endpoints, "set_mcp_tool_version_deprecation", set_mock), + ): + result: Final = await mgmt_endpoints.set_tool_version_deprecation( + server_id="srv-1", + tool_name="list_notes", + version=1, + payload=request, + user_api_key_dict=admin, + ) + + assert result == version + set_mock.assert_awaited_once_with(prisma, "srv-1", "list_notes", 1, request) + + @pytest.mark.asyncio + async def test_put_tool_version_deprecation_rejects_non_admin_and_unknown_versions(self): + non_admin: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.INTERNAL_USER) + with pytest.raises(HTTPException) as forbidden: + await mgmt_endpoints.set_tool_version_deprecation( + server_id="srv-1", + tool_name="list_notes", + version=1, + payload=MCPToolDeprecationRequest(), + user_api_key_dict=non_admin, + ) + assert forbidden.value.status_code == 403 + assert forbidden.value.detail == {"error": "Admin access required to update MCP tool version deprecation."} + + with ( + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()), + patch.object(mgmt_endpoints, "set_mcp_tool_version_deprecation", AsyncMock(return_value=None)), + ): + with pytest.raises(HTTPException) as missing: + await mgmt_endpoints.set_tool_version_deprecation( + server_id="srv-1", + tool_name="missing", + version=9, + payload=MCPToolDeprecationRequest(), + user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert missing.value.status_code == 404 + assert missing.value.detail == {"error": "MCP tool 'missing' version 9 not found for server 'srv-1'."} + + @pytest.mark.asyncio + async def test_delete_tool_version_deprecation_is_admin_only_and_clears_metadata(self): + prisma: Final = MagicMock() + version: Final = MCPToolVersion( + server_id="srv-1", + tool_name="list_notes", + version=1, + change_kind="initial", + created_at=datetime(2026, 10, 2, tzinfo=timezone.utc), + ) + clear_mock: Final = AsyncMock(return_value=version) + admin: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=prisma), + patch.object(mgmt_endpoints, "set_mcp_tool_version_deprecation", clear_mock), + ): + result: Final = await mgmt_endpoints.clear_tool_version_deprecation( + server_id="srv-1", + tool_name="list_notes", + version=1, + user_api_key_dict=admin, + ) + + assert result == version + clear_mock.assert_awaited_once_with(prisma, "srv-1", "list_notes", 1, None) + + with pytest.raises(HTTPException) as forbidden: + await mgmt_endpoints.clear_tool_version_deprecation( + server_id="srv-1", + tool_name="list_notes", + version=1, + user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.INTERNAL_USER), + ) + + assert forbidden.value.status_code == 403 + + @pytest.mark.asyncio + async def test_delete_tool_version_deprecation_returns_404_for_unknown_version(self): + admin: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + with ( + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()), + patch.object(mgmt_endpoints, "set_mcp_tool_version_deprecation", AsyncMock(return_value=None)), + ): + with pytest.raises(HTTPException) as missing: + await mgmt_endpoints.clear_tool_version_deprecation( + server_id="srv-1", + tool_name="missing", + version=9, + user_api_key_dict=admin, + ) + + assert missing.value.status_code == 404 + assert missing.value.detail == {"error": "MCP tool 'missing' version 9 not found for server 'srv-1'."} + + @pytest.mark.asyncio + async def test_pin_concurrent_version_insert_conflict_returns_409_without_updating_server_row(self): + prisma: Final = MagicMock() + tx_client: Final = MagicMock() + tx_client.execute_raw = AsyncMock() + tx_client.litellm_mcpservertable.find_unique = AsyncMock(return_value=MagicMock()) + tx_client.litellm_mcpservertable.update = AsyncMock() + tx_client.litellm_mcptoolversion.find_many = AsyncMock(return_value=[]) + tx_client.litellm_mcptoolversion.create_many = AsyncMock( + side_effect=UniqueViolationError({}, message="unique version") + ) + tx: Final = MagicMock() + tx.__aenter__ = AsyncMock(return_value=tx_client) + tx.__aexit__ = AsyncMock(return_value=False) + prisma.tx = MagicMock(return_value=tx) + stored: Final = generate_mock_mcp_server_db_record(server_id="srv-1") + manager: Final = TestPinMCPServerTools._manager([("list_notes", "List notes", {})]) + admin: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin") + + try: + with ( + patch.object(mgmt_endpoints, "MCP_AVAILABLE", True), + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=prisma), + patch.object(mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=stored)), + patch.object(mgmt_endpoints, "global_mcp_server_manager", manager), + patch("litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager", manager), + patch.dict( + sys.modules, + { + "litellm.proxy.proxy_server": types.SimpleNamespace( + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + general_settings={}, + llm_router=None, + ) + }, + ), + ): + with pytest.raises(HTTPException) as conflict: + await mgmt_endpoints.pin_mcp_server_tools( + server_id="srv-1", + request=_make_mock_request(), + user_api_key_dict=admin, + ) + finally: + ProxyLogging._callback_capabilities_cache.clear() + + assert conflict.value.status_code == 409 + assert conflict.value.detail == { + "error": "Another pin of MCP server 'srv-1' recorded tool versions at the same time; retry the pin." + } + tx_client.litellm_mcptoolversion.create_many.assert_awaited_once() + tx_client.litellm_mcpservertable.update.assert_not_awaited() + + @dataclass(frozen=True) class _ResolutionEffects: byok_store: AsyncMock = field(default_factory=AsyncMock) @@ -9406,14 +9768,14 @@ class TestMCPServerResolutionCharacterization: second_id: Final = "lit3974_second_credential" prisma, manager, caller = await self._resolution_case("db_runtime", "allowed", first_id) ids: Final = (first_id, second_id) - rows: Final = tuple( - generate_mock_mcp_server_db_record(server_id=sid, alias=f"alias-{sid}") for sid in ids - ) + rows: Final = tuple(generate_mock_mcp_server_db_record(server_id=sid, alias=f"alias-{sid}") for sid in ids) prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=list(rows)) auth: Final = caller.model_copy( - update={"object_permission": LiteLLM_ObjectPermissionTable( - object_permission_id="lit3974_multiple_credentials", mcp_servers=list(ids) - )} + update={ + "object_permission": LiteLLM_ObjectPermissionTable( + object_permission_id="lit3974_multiple_credentials", mcp_servers=list(ids) + ) + } ) manager.config_mcp_servers = { **manager.config_mcp_servers, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolVersionsPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolVersionsPanel.test.tsx new file mode 100644 index 00000000000..be0f22e44d2 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolVersionsPanel.test.tsx @@ -0,0 +1,224 @@ +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import * as networking from "@/components/networking"; +import type { MCPToolVersion } from "@/components/networking"; +import { mcpServersKeys } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers"; +import { MCPToolVersionsPanel } from "./MCPToolVersionsPanel"; + +vi.mock("@/components/networking", () => ({ + getMCPToolVersions: vi.fn(), + pinMCPServerTools: vi.fn(), + setMCPToolVersionDeprecation: vi.fn(), + clearMCPToolVersionDeprecation: vi.fn(), +})); + +const toolVersions: MCPToolVersion[] = [ + { + server_id: "srv-1", + tool_name: "list_notes", + version: 2, + description: "List notes", + input_schema: {}, + change_kind: "breaking", + changes: [{ breaking: true, summary: 'Parameter "limit" type changed from "string" to "integer"' }], + changelog: "Updated the limit parameter", + deprecated_at: "2026-10-02T00:00:00Z", + sunset_date: "2026-12-31T00:00:00Z", + deprecation_note: "Move to v3", + created_at: "2026-10-02T00:00:00Z", + created_by: "admin", + }, + { + server_id: "srv-1", + tool_name: "list_notes", + version: 1, + description: "List notes", + input_schema: {}, + change_kind: "initial", + changes: [], + changelog: null, + deprecated_at: null, + sunset_date: null, + deprecation_note: null, + created_at: "2026-10-01T00:00:00Z", + created_by: "admin", + }, + { + server_id: "srv-1", + tool_name: "read_note", + version: 1, + description: "Read a note", + input_schema: {}, + change_kind: "non_breaking", + changes: [{ breaking: false, summary: "Description changed" }], + changelog: null, + deprecated_at: null, + sunset_date: null, + deprecation_note: null, + created_at: "2026-10-01T00:00:00Z", + created_by: "admin", + }, +]; + +const renderPanel = (isProxyAdmin = true, customHeaders?: Record) => + render( + + + , + ); + +describe("MCPToolVersionsPanel", () => { + afterEach(() => { + vi.restoreAllMocks(); + }); + + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(networking.getMCPToolVersions).mockResolvedValue(toolVersions); + vi.mocked(networking.pinMCPServerTools).mockResolvedValue({}); + vi.mocked(networking.setMCPToolVersionDeprecation).mockResolvedValue(toolVersions[0]); + vi.mocked(networking.clearMCPToolVersionDeprecation).mockResolvedValue(toolVersions[0]); + }); + + it("shows breaking, non-breaking and deprecated sunset badges", async () => { + renderPanel(); + + expect(await screen.findByText("Breaking")).toBeInTheDocument(); + expect(screen.getAllByText("Breaking").length).toBeGreaterThan(0); + expect(screen.getByText("Non-breaking")).toBeInTheDocument(); + expect(screen.getByText("Deprecated, sunset 2026-12-31")).toBeInTheDocument(); + }); + + it("expands version history newest first with changes and changelog", async () => { + renderPanel(false); + + const toggle = await screen.findByRole("button", { name: "Show history for list_notes" }); + expect(toggle).toHaveAttribute("aria-expanded", "false"); + await userEvent.click(toggle); + + expect(screen.getByRole("button", { name: "Hide history for list_notes" })).toHaveAttribute( + "aria-expanded", + "true", + ); + expect(screen.getByText('Parameter "limit" type changed from "string" to "integer"')).toBeInTheDocument(); + expect(screen.getByText("Updated the limit parameter")).toBeInTheDocument(); + expect(screen.getAllByRole("heading", { level: 4 }).map((heading) => heading.textContent)).toEqual(["v2", "v1"]); + expect(screen.getByText("Created 2026-10-02")).toBeInTheDocument(); + }); + + it("hides pin and deprecation controls from non-admins", async () => { + renderPanel(false); + + expect(await screen.findByRole("button", { name: "Show history for list_notes" })).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Pin and record versions" })).not.toBeInTheDocument(); + await userEvent.click(screen.getByRole("button", { name: "Show history for list_notes" })); + expect(screen.queryByRole("button", { name: "Deprecate list_notes v1" })).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Edit deprecation for list_notes v2" })).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Remove deprecation for list_notes v2" })).not.toBeInTheDocument(); + }); + + it("sends the changelog when the admin pins tools", async () => { + vi.mocked(networking.getMCPToolVersions).mockResolvedValue([]); + renderPanel(); + + fireEvent.change(await screen.findByLabelText("Changelog"), { + target: { value: " Published first release " }, + }); + await userEvent.click(screen.getByRole("button", { name: "Pin and record versions" })); + + expect(networking.pinMCPServerTools).toHaveBeenCalledWith( + "token", + "srv-1", + " Published first release ", + undefined, + ); + }); + + it("forwards upstream headers when the admin pins tools", async () => { + vi.mocked(networking.getMCPToolVersions).mockResolvedValue([]); + const customHeaders: Record = { + "x-mcp-weather-authorization": "Bearer t", + }; + renderPanel(true, customHeaders); + + await screen.findByLabelText("Changelog"); + await userEvent.click(screen.getByRole("button", { name: "Pin and record versions" })); + + expect(networking.pinMCPServerTools).toHaveBeenCalledWith("token", "srv-1", "", customHeaders); + }); + + it("invalidates all pinned catalog queries after pinning succeeds", async () => { + vi.mocked(networking.getMCPToolVersions).mockResolvedValue([]); + const invalidateQueries = vi.spyOn(QueryClient.prototype, "invalidateQueries"); + renderPanel(); + + fireEvent.change(await screen.findByLabelText("Changelog"), { + target: { value: "Initial catalog" }, + }); + await userEvent.click(screen.getByRole("button", { name: "Pin and record versions" })); + + await waitFor(() => expect(invalidateQueries).toHaveBeenCalledTimes(3)); + expect(invalidateQueries).toHaveBeenCalledWith({ queryKey: ["mcpToolVersions", "srv-1"] }); + expect(invalidateQueries).toHaveBeenCalledWith({ queryKey: ["mcpTools", "srv-1"] }); + expect(invalidateQueries).toHaveBeenCalledWith({ queryKey: mcpServersKeys.all }); + }); + + it("sends sunset dates at UTC midnight", async () => { + renderPanel(); + + await userEvent.click(await screen.findByRole("button", { name: "Show history for read_note" })); + await userEvent.click(screen.getByRole("button", { name: "Deprecate read_note v1" })); + fireEvent.change(screen.getByLabelText("Sunset date"), { target: { value: "2026-12-31" } }); + fireEvent.change(screen.getByLabelText("Note"), { target: { value: "Move to v2" } }); + await userEvent.click(screen.getByRole("button", { name: "Save" })); + + expect(networking.setMCPToolVersionDeprecation).toHaveBeenCalledWith( + "token", + { serverId: "srv-1", toolName: "read_note", version: 1 }, + { + sunset_date: "2026-12-31T00:00:00Z", + deprecation_note: "Move to v2", + }, + ); + }); + + it("prefills and edits an existing deprecation", async () => { + renderPanel(); + + await userEvent.click(await screen.findByRole("button", { name: "Show history for list_notes" })); + await userEvent.click(screen.getByRole("button", { name: "Edit deprecation for list_notes v2" })); + + expect(screen.getByLabelText("Sunset date")).toHaveValue("2026-12-31"); + expect(screen.getByLabelText("Note")).toHaveValue("Move to v3"); + + fireEvent.change(screen.getByLabelText("Sunset date"), { target: { value: "2027-01-15" } }); + fireEvent.change(screen.getByLabelText("Note"), { target: { value: "Use v4" } }); + await userEvent.click(screen.getByRole("button", { name: "Save" })); + + expect(networking.setMCPToolVersionDeprecation).toHaveBeenCalledWith( + "token", + { serverId: "srv-1", toolName: "list_notes", version: 2 }, + { + sunset_date: "2027-01-15T00:00:00Z", + deprecation_note: "Use v4", + }, + ); + }); + + it("shows loading separately from the empty state", () => { + vi.mocked(networking.getMCPToolVersions).mockReturnValue(new Promise(() => {})); + renderPanel(); + + expect(screen.getByRole("status", { name: "Loading tool versions" })).toBeInTheDocument(); + expect( + screen.queryByText("No versions recorded yet. Pin the tool list to start tracking versions."), + ).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolVersionsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolVersionsPanel.tsx new file mode 100644 index 00000000000..70650aba927 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolVersionsPanel.tsx @@ -0,0 +1,379 @@ +"use client"; + +import { type FormEvent, useState } from "react"; +import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query"; +import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert"; +import { Badge } from "@/components/ui/badge"; +import { Button } from "@/components/ui/button"; +import { Card } from "@/components/ui/card"; +import { Input } from "@/components/ui/input"; +import { Skeleton } from "@/components/ui/skeleton"; +import { Textarea } from "@/components/ui/textarea"; +import { mcpServersKeys } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers"; +import { + clearMCPToolVersionDeprecation, + getMCPToolVersions, + pinMCPServerTools, + setMCPToolVersionDeprecation, + type MCPToolVersion, +} from "@/components/networking"; + +type MCPToolChangeKind = MCPToolVersion["change_kind"]; + +const CHANGE_KIND_LABELS: Record = { + initial: "Initial", + non_breaking: "Non-breaking", + breaking: "Breaking", + removed: "Removed", +}; + +const CHANGE_KIND_VARIANTS: Record = { + initial: "outline", + non_breaking: "secondary", + breaking: "destructive", + removed: "destructive", +}; + +function formatUTCDate(value: string): string { + const parsed: Date = new Date(value); + return Number.isNaN(parsed.getTime()) ? value.slice(0, 10) : parsed.toISOString().slice(0, 10); +} + +function ChangeKindBadge({ kind }: { kind: MCPToolChangeKind }) { + return {CHANGE_KIND_LABELS[kind]}; +} + +function errorMessage(error: unknown): string { + return error instanceof Error ? error.message : String(error); +} + +export interface MCPToolVersionsPanelProps { + serverId: string; + accessToken: string | null; + isProxyAdmin: boolean; + customHeaders?: Record; +} + +export function MCPToolVersionsPanel({ + serverId, + accessToken, + isProxyAdmin, + customHeaders, +}: MCPToolVersionsPanelProps) { + const queryClient = useQueryClient(); + const queryKey: readonly ["mcpToolVersions", string] = ["mcpToolVersions", serverId]; + const [changelog, setChangelog] = useState(""); + const [expandedTool, setExpandedTool] = useState(null); + const [editingVersion, setEditingVersion] = useState(null); + const [sunsetDate, setSunsetDate] = useState(""); + const [deprecationNote, setDeprecationNote] = useState(""); + const { + data: versions, + error: loadError, + isLoading, + } = useQuery({ + queryKey, + queryFn: () => { + if (!accessToken) throw new Error("An access token is required to load tool versions."); + return getMCPToolVersions(accessToken, serverId); + }, + enabled: Boolean(accessToken), + }); + const invalidateVersions = async () => queryClient.invalidateQueries({ queryKey }); + const invalidatePinnedToolQueries = async (): Promise => { + await Promise.all([ + invalidateVersions(), + queryClient.invalidateQueries({ queryKey: ["mcpTools", serverId] }), + queryClient.invalidateQueries({ queryKey: mcpServersKeys.all }), + ]); + }; + const pinMutation = useMutation({ + mutationFn: () => { + if (!accessToken) throw new Error("An access token is required to pin tools."); + return pinMCPServerTools(accessToken, serverId, changelog, customHeaders); + }, + onSuccess: invalidatePinnedToolQueries, + }); + const deprecationMutation = useMutation({ + mutationFn: ({ + toolName, + version, + sunsetDate: value, + note, + }: { + toolName: string; + version: number; + sunsetDate: string; + note: string; + }) => { + if (!accessToken) throw new Error("An access token is required to update deprecation."); + return setMCPToolVersionDeprecation( + accessToken, + { serverId, toolName, version }, + { + sunset_date: value ? `${value}T00:00:00Z` : null, + deprecation_note: note.trim() || null, + }, + ); + }, + onSuccess: async () => { + setEditingVersion(null); + await invalidateVersions(); + }, + }); + const clearDeprecationMutation = useMutation({ + mutationFn: ({ toolName, version }: { toolName: string; version: number }) => { + if (!accessToken) throw new Error("An access token is required to clear deprecation."); + return clearMCPToolVersionDeprecation(accessToken, { serverId, toolName, version }); + }, + onSuccess: invalidateVersions, + }); + const beginDeprecation = (version: MCPToolVersion): void => { + setEditingVersion(`${version.tool_name}-${version.version}`); + setSunsetDate(version.sunset_date ? formatUTCDate(version.sunset_date) : ""); + setDeprecationNote(version.deprecation_note ?? ""); + }; + const submitDeprecation = (event: FormEvent, toolName: string, version: number): void => { + event.preventDefault(); + const variables = { toolName, version, sunsetDate, note: deprecationNote }; + deprecationMutation.mutate(variables); + }; + + if (isLoading) { + return ( +
+

+ Tool versions +

+
+
+ + +
+
+ + +
+
+
+ ); + } + + if (loadError || !accessToken) { + return ( +
+

+ Tool versions +

+ + Could not load tool versions + + {loadError ? errorMessage(loadError) : "An access token is required to load tool versions."} + + +
+ ); + } + + const orderedVersions: MCPToolVersion[] = versions ?? []; + const toolNames: string[] = [...new Set(orderedVersions.map((version) => version.tool_name))].sort((a, b) => + a.localeCompare(b), + ); + + return ( +
+

+ Tool versions +

+ + {isProxyAdmin && ( + + +