diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 92b5fb6552f..110062b081b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -29,7 +29,7 @@ from dataclasses import dataclass, replace from functools import lru_cache from itertools import chain from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast +from typing import TYPE_CHECKING, Any, Final, Generic, Literal, NoReturn, TypeAlias, TypedDict, TypeVar, cast from urllib.parse import ParseResult, urlparse import anyio @@ -1415,6 +1415,44 @@ def _upstream_failure_suffix(exc: BaseException) -> str: return f"\n upstream exchange: {detail}" if detail else "" +def _raise_single_server_list_failure(error: Exception, server: MCPServer, catalog: str) -> NoReturn: + """Relay a failed single-server catalog fetch: auth challenges (upstream, or a v2 resolver's + HTTPException 401/403 raised at client-build time) become ``MCPUpstreamAuthError`` with the + ``WWW-Authenticate`` kept (dropped for dcr_bridge servers, whose upstream challenge points at the + wrong metadata); anything else becomes a classified ``MCPServerListError``.""" + match error: + case MCPUpstreamAuthError() if server.is_dcr_bridge and error.www_authenticate is not None: + raise MCPUpstreamAuthError( + status_code=error.status_code, + www_authenticate=None, + server_name=error.server_name, + ) from error + case MCPUpstreamAuthError() | MCPServerListError(): + raise error + case HTTPException() if error.status_code in (401, 403): + headers: Final = error.headers or {} + challenge_header: Final = headers.get("WWW-Authenticate") or headers.get("www-authenticate") + raise MCPUpstreamAuthError( + status_code=error.status_code, + www_authenticate=None if server.is_dcr_bridge else challenge_header, + server_name=server.name, + ) from error + case HTTPException(): + verbose_logger.warning("Failed to get %s from server %s: %s", catalog, server.name, error) + raise MCPServerListError( + ServerListFault(tag="internal", status_code=error.status_code), server.name + ) from error + case _: + verbose_logger.warning( + "Failed to get %s from server %s: %s%s", + catalog, + server.name, + type(error).__name__, + _upstream_failure_suffix(error), + ) + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) + + def _obo_retry_applies(server: MCPServer, subject_token: str | None) -> bool: """Whether an upstream 401/403 should invalidate the minted credential and retry once. @@ -4451,40 +4489,8 @@ class MCPServerManager: return prefixed_or_original_tools - except MCPUpstreamAuthError as upstream_auth_error: - # Pass-through 401 must surface to single-server routes so the - # client triggers the upstream OAuth flow. The multi-server - # aggregator catches this explicitly to keep absorbing. - if server.is_dcr_bridge and upstream_auth_error.www_authenticate is not None: - raise MCPUpstreamAuthError( - status_code=upstream_auth_error.status_code, - www_authenticate=None, - server_name=upstream_auth_error.server_name, - ) from upstream_auth_error - raise - except HTTPException as e: - # A v2 resolver auth challenge (token_exchange's RFC 9728 401, authorization_code's - # browser-OAuth 401, or a 403) is raised at client-build time, inside this try. Route it - # through the same MCPUpstreamAuthError channel as pass-through so single-server routes - # surface the challenge (the client re-authenticates) while the aggregator keeps absorbing. - # Non-auth HTTP errors stay absorbed so one misconfigured server can't blank the listing. - if e.status_code in (401, 403): - headers: Final = e.headers or {} - challenge_header: Final = headers.get("WWW-Authenticate") or headers.get("www-authenticate") - raise MCPUpstreamAuthError( - status_code=e.status_code, - www_authenticate=None if server.is_dcr_bridge else challenge_header, - server_name=server.name, - ) from e - verbose_logger.warning("Failed to get tools from server %s: %s", server.name, e) - raise MCPServerListError(ServerListFault(tag="internal", status_code=e.status_code), server.name) from e - except MCPServerListError: - raise except Exception as e: - verbose_logger.warning( - "Failed to get tools from server %s: %s%s", server.name, type(e).__name__, _upstream_failure_suffix(e) - ) - raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge) + _raise_single_server_list_failure(e, server, "tools") def _invalidate_discovery_lists(self, server_id: str) -> None: self._prompt_discovery_cache.invalidate(server_id) @@ -4563,7 +4569,7 @@ class MCPServerManager: return self._create_prefixed_prompts(items, server, add_prefix=add_prefix) except Exception as error: if raise_on_error: - raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) + _raise_single_server_list_failure(error, server, "prompts") verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error) return [] @@ -4609,7 +4615,7 @@ class MCPServerManager: return self._create_prefixed_resources(items, server, add_prefix=add_prefix) except Exception as error: if raise_on_error: - raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) + _raise_single_server_list_failure(error, server, "resources") verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error) return [] @@ -4655,7 +4661,7 @@ class MCPServerManager: return self._create_prefixed_resource_templates(items, server, add_prefix=add_prefix) except Exception as error: if raise_on_error: - raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) + _raise_single_server_list_failure(error, server, "resource templates") verbose_logger.warning("Failed to get resource_templates from server %s: %s", server.name, error) return [] diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 31704aed0e4..38acf835c0b 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -199,8 +199,21 @@ if MCP_AVAILABLE: fire_mcp_tool_call_failure_logging, ) + class MCPCatalogPrompt(Prompt): + """An MCP server's prompt as the upstream reports it. Subclassed only so the OpenAPI + component gets a name distinct from the prompt-management ``Prompt`` request model.""" + class ListMCPPromptsRestAPIResponse(BaseModel): - prompts: list[Prompt] + prompts: list[MCPCatalogPrompt] + + @classmethod + def from_prompts(cls, prompts: Sequence[Prompt]) -> "ListMCPPromptsRestAPIResponse": + return cls( + prompts=[ + MCPCatalogPrompt.model_validate(prompt.model_dump(by_alias=True, exclude_unset=True)) + for prompt in prompts + ] + ) class ListMCPResourcesRestAPIResponse(BaseModel): resources: list[Resource] @@ -1138,7 +1151,7 @@ if MCP_AVAILABLE: raise _relay_upstream_auth_http_exception(e, request) from e except MCPServerListError as e: raise _catalog_list_http_exception(e, context.server, "prompts") from e - return ListMCPPromptsRestAPIResponse(prompts=prompts) + return ListMCPPromptsRestAPIResponse.from_prompts(prompts) @router.get("/resources/list", dependencies=[Depends(user_api_key_auth)]) async def list_resources_rest_api( diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 3571179cd08..54b713e48ac 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -21899,7 +21899,7 @@ "properties": { "prompts": { "items": { - "$ref": "#/components/schemas/Prompt" + "$ref": "#/components/schemas/MCPCatalogPrompt" }, "title": "Prompts", "type": "array" @@ -21935,6 +21935,83 @@ "title": "ListMCPResourcesRestAPIResponse", "type": "object" }, + "MCPCatalogPrompt": { + "additionalProperties": true, + "description": "An MCP server's prompt as the upstream reports it. Subclassed only so the OpenAPI\ncomponent gets a name distinct from the prompt-management ``Prompt`` request model.", + "properties": { + "_meta": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "Meta" + }, + "arguments": { + "anyOf": [ + { + "items": { + "$ref": "#/components/schemas/PromptArgument" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Arguments" + }, + "description": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Description" + }, + "icons": { + "anyOf": [ + { + "items": { + "$ref": "#/components/schemas/Icon" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Icons" + }, + "name": { + "title": "Name", + "type": "string" + }, + "title": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Title" + } + }, + "required": [ + "name" + ], + "title": "MCPCatalogPrompt", + "type": "object" + }, "MCPCredentials": { "properties": { "audience": { @@ -22746,83 +22823,6 @@ "title": "NewMCPServerRequest", "type": "object" }, - "Prompt": { - "additionalProperties": true, - "description": "A prompt or prompt template that the server offers.", - "properties": { - "_meta": { - "anyOf": [ - { - "additionalProperties": true, - "type": "object" - }, - { - "type": "null" - } - ], - "title": "Meta" - }, - "arguments": { - "anyOf": [ - { - "items": { - "$ref": "#/components/schemas/PromptArgument" - }, - "type": "array" - }, - { - "type": "null" - } - ], - "title": "Arguments" - }, - "description": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Description" - }, - "icons": { - "anyOf": [ - { - "items": { - "$ref": "#/components/schemas/Icon" - }, - "type": "array" - }, - { - "type": "null" - } - ], - "title": "Icons" - }, - "name": { - "title": "Name", - "type": "string" - }, - "title": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Title" - } - }, - "required": [ - "name" - ], - "title": "Prompt", - "type": "object" - }, "PromptArgument": { "additionalProperties": true, "description": "An argument for a prompt template.", @@ -32168,7 +32168,7 @@ "properties": { "prompts": { "items": { - "$ref": "#/components/schemas/Prompt" + "$ref": "#/components/schemas/MCPCatalogPrompt" }, "title": "Prompts", "type": "array" @@ -32204,6 +32204,83 @@ "title": "ListMCPResourcesRestAPIResponse", "type": "object" }, + "MCPCatalogPrompt": { + "additionalProperties": true, + "description": "An MCP server's prompt as the upstream reports it. Subclassed only so the OpenAPI\ncomponent gets a name distinct from the prompt-management ``Prompt`` request model.", + "properties": { + "_meta": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "Meta" + }, + "arguments": { + "anyOf": [ + { + "items": { + "$ref": "#/components/schemas/PromptArgument" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Arguments" + }, + "description": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Description" + }, + "icons": { + "anyOf": [ + { + "items": { + "$ref": "#/components/schemas/Icon" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Icons" + }, + "name": { + "title": "Name", + "type": "string" + }, + "title": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Title" + } + }, + "required": [ + "name" + ], + "title": "MCPCatalogPrompt", + "type": "object" + }, "MCPCredentials": { "properties": { "audience": { @@ -33015,83 +33092,6 @@ "title": "NewMCPServerRequest", "type": "object" }, - "Prompt": { - "additionalProperties": true, - "description": "A prompt or prompt template that the server offers.", - "properties": { - "_meta": { - "anyOf": [ - { - "additionalProperties": true, - "type": "object" - }, - { - "type": "null" - } - ], - "title": "Meta" - }, - "arguments": { - "anyOf": [ - { - "items": { - "$ref": "#/components/schemas/PromptArgument" - }, - "type": "array" - }, - { - "type": "null" - } - ], - "title": "Arguments" - }, - "description": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Description" - }, - "icons": { - "anyOf": [ - { - "items": { - "$ref": "#/components/schemas/Icon" - }, - "type": "array" - }, - { - "type": "null" - } - ], - "title": "Icons" - }, - "name": { - "title": "Name", - "type": "string" - }, - "title": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Title" - } - }, - "required": [ - "name" - ], - "title": "Prompt", - "type": "object" - }, "PromptArgument": { "additionalProperties": true, "description": "An argument for a prompt template.", diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 1462dde8cc2..4e315b605d6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -3964,6 +3964,39 @@ class TestMCPServerManager: assert exc_info.value.fault == ServerListFault(tag="unreachable") assert exc_info.value.server_name == server.name + @pytest.mark.asyncio + @pytest.mark.parametrize( + "manager_method", + ["get_prompts_from_server", "get_resources_from_server", "get_resource_templates_from_server"], + ) + @pytest.mark.parametrize("challenge_carrier", ["resolver_http_exception", "upstream_auth_error"]) + async def test_catalog_fetch_relays_auth_challenge_like_tools(self, manager_method, challenge_carrier): + """An auth challenge raised while building the client (a v2 resolver HTTPException 401) or by + the upstream itself must reach a single-server caller as MCPUpstreamAuthError with the + WWW-Authenticate intact, exactly as the tools listing relays it, not as a bare fault.""" + manager = MCPServerManager() + server = MCPServer( + server_id="server-1", + name="alias-server", + alias="alias-server", + server_name="alias-server", + url="https://example.com", + transport=MCPTransport.http, + ) + challenge = 'Bearer resource_metadata="https://example.com/.well-known/oauth-protected-resource"' + raised = ( + HTTPException(status_code=401, detail="Unauthorized", headers={"WWW-Authenticate": challenge}) + if challenge_carrier == "resolver_http_exception" + else MCPUpstreamAuthError(status_code=401, www_authenticate=challenge, server_name=server.name) + ) + + with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, side_effect=raised): + with pytest.raises(MCPUpstreamAuthError) as exc_info: + await getattr(manager, manager_method)(server, user_api_key_auth=None, raise_on_error=True) + + assert exc_info.value.status_code == 401 + assert exc_info.value.www_authenticate == challenge + @pytest.mark.asyncio async def test_read_resource_from_server_success(self): manager = MCPServerManager() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index fda8b1620e4..d92ebd685a5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -2550,6 +2550,33 @@ class TestListPromptsAndResourcesRestAPI: assert exc_info.value.headers is not None assert "www-authenticate" in {key.lower() for key in exc_info.value.headers} + def test_openapi_keeps_prompt_management_and_mcp_prompt_contracts_distinct(self): + """The catalog response reuses the MCP SDK prompt type, which shares its class name with the + prompt-management request model, so the two must land as separate OpenAPI components: + POST /prompts still requires prompt_id + litellm_params while the catalog item requires name.""" + from fastapi import FastAPI + + from litellm.proxy.prompts.prompt_endpoints import router as prompt_router + + app = FastAPI() + app.include_router(rest_endpoints.router) + app.include_router(prompt_router) + spec = app.openapi() + schemas = spec["components"]["schemas"] + + def component(ref: Dict[str, Any]) -> Dict[str, Any]: + return schemas[ref["$ref"].rsplit("/", 1)[1]] + + create_prompt_operation = spec["paths"]["/prompts"]["post"] + create_prompt_body = component(create_prompt_operation["requestBody"]["content"]["application/json"]["schema"]) + assert {"prompt_id", "litellm_params"} <= set(create_prompt_body["required"]) + + catalog_operation = spec["paths"]["/mcp-rest/prompts/list"]["get"] + catalog_response = component(catalog_operation["responses"]["200"]["content"]["application/json"]["schema"]) + catalog_prompt = component(catalog_response["properties"]["prompts"]["items"]) + assert catalog_prompt["required"] == ["name"] + assert "arguments" in catalog_prompt["properties"] + class TestCallToolRestAPI: pytestmark = pytest.mark.asyncio diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx index a183dbf8a89..f76613d9a5b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx @@ -1,4 +1,4 @@ -import { render, screen, waitFor, within } from "@testing-library/react"; +import { act, render, screen, waitFor, within } from "@testing-library/react"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { describe, expect, it, vi, beforeEach } from "vitest"; import MCPToolsViewer from "./mcp_tools"; @@ -24,8 +24,13 @@ vi.mock("@/utils/mcpTokenStore", () => ({ removeToken: vi.fn(), })); -const { toolsOAuthFlowSpy } = vi.hoisted(() => ({ +const { toolsOAuthFlowSpy, userMcpOAuthFlowSpy } = vi.hoisted(() => ({ toolsOAuthFlowSpy: vi.fn(() => ({ startOAuthFlow: vi.fn(), status: "idle", error: null })), + userMcpOAuthFlowSpy: vi.fn((_options: { onSuccess: () => void }) => ({ + startOAuthFlow: vi.fn(), + status: "idle", + error: null, + })), })); vi.mock("@/hooks/useToolsOAuthFlow", () => ({ @@ -33,7 +38,7 @@ vi.mock("@/hooks/useToolsOAuthFlow", () => ({ })); vi.mock("@/hooks/useUserMcpOAuthFlow", () => ({ - useUserMcpOAuthFlow: () => ({ startOAuthFlow: vi.fn(), status: "idle", error: null }), + useUserMcpOAuthFlow: userMcpOAuthFlowSpy, })); const GATE_TEXT = "Authentication required"; @@ -315,4 +320,28 @@ describe("MCPToolsViewer prompts and resources catalog", () => { expect(await within(resources).findByText("Error: upstream unreachable")).toBeInTheDocument(); expect(await screen.findByText("summarize")).toBeInTheDocument(); }); + + it("reloads prompts and resources together with tools after the user re-authorizes", async () => { + const expiredPrompts = { + prompts: [], + error: "auth_required", + message: "upstream credential expired", + status: 401, + }; + vi.mocked(listMCPPrompts).mockResolvedValue(expiredPrompts); + userMcpOAuthFlowSpy.mockClear(); + + renderViewer({ oauth2_flow: null, delegate_auth_to_upstream: false }); + + const prompts = await screen.findByRole("region", { name: "Prompts" }); + expect(await within(prompts).findByText("Error: upstream credential expired")).toBeInTheDocument(); + expect(vi.mocked(listMCPPrompts)).toHaveBeenCalledTimes(1); + + vi.mocked(listMCPPrompts).mockResolvedValue({ prompts: [{ name: "summarize" }] }); + act(() => userMcpOAuthFlowSpy.mock.calls.at(-1)?.[0].onSuccess()); + + expect(await within(prompts).findByText("summarize")).toBeInTheDocument(); + expect(vi.mocked(listMCPTools)).toHaveBeenCalledTimes(2); + expect(vi.mocked(listMCPResources)).toHaveBeenCalledTimes(2); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx index bacd2a58216..4be58462a7a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx @@ -201,26 +201,40 @@ const MCPToolsViewer = ({ }, }); - const { data: mcpPromptsResponse, isLoading: isLoadingPrompts } = useQuery({ + const { + data: mcpPromptsResponse, + isLoading: isLoadingPrompts, + refetch: refetchPrompts, + } = useQuery({ queryKey: ["mcpPrompts", serverId, passthroughHeaders, oauthToken], queryFn: () => listMCPPrompts(accessToken ?? "", serverId, buildCustomHeaders()), enabled: catalogQueriesEnabled, staleTime: 30000, }); - const { data: mcpResourcesResponse, isLoading: isLoadingResources } = useQuery({ + const { + data: mcpResourcesResponse, + isLoading: isLoadingResources, + refetch: refetchResources, + } = useQuery({ queryKey: ["mcpResources", serverId, passthroughHeaders, oauthToken], queryFn: () => listMCPResources(accessToken ?? "", serverId, buildCustomHeaders()), enabled: catalogQueriesEnabled, staleTime: 30000, }); + const refetchCatalog = useCallback(() => { + refetchTools(); + refetchPrompts(); + refetchResources(); + }, [refetchTools, refetchPrompts, refetchResources]); + // authorization_code authorize: same redirect+exchange flow as the admin "Authorize & Fetch" // and the chat "Connect" button, but persists the token to the per-user DB. const onAuthorizationCodeAuthSuccess = useCallback(() => { refetchAuthorizationCodeCred(); - refetchTools(); - }, [refetchAuthorizationCodeCred, refetchTools]); + refetchCatalog(); + }, [refetchAuthorizationCodeCred, refetchCatalog]); const { startOAuthFlow: startDbOAuthFlow, @@ -370,7 +384,7 @@ const MCPToolsViewer = ({