From 87bd995ec437b9a09e2aeb3577f27a95cc090f40 Mon Sep 17 00:00:00 2001 From: naaa760 Date: Wed, 11 Feb 2026 10:22:55 +0530 Subject: [PATCH] add mcp prompts REST (list + get), e2e and unit tests --- tests/mcp_tests/mcp_server.py | 12 + tests/mcp_tests/test_proxy_mcp_e2e.py | 127 +++++++ .../mcp_server/test_rest_endpoints.py | 319 ++++++++++++++++++ 3 files changed, 458 insertions(+) diff --git a/tests/mcp_tests/mcp_server.py b/tests/mcp_tests/mcp_server.py index bc6accbb721..8fcf7817ae5 100644 --- a/tests/mcp_tests/mcp_server.py +++ b/tests/mcp_tests/mcp_server.py @@ -40,6 +40,18 @@ def multiply(a: int, b: int) -> int: return a * b +@mcp.prompt() +def code_review(code: str) -> str: + """Ask the LLM to review code quality and suggest improvements""" + return f"Please review this code:\n{code}" + + +@mcp.prompt() +def summarize(text: str, style: str = "concise") -> str: + """Summarize text in the requested style""" + return f"Please provide a {style} summary of the following text:\n{text}" + + def main() -> None: args = _parse_args() transport = (args.transport or "stdio").lower() diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index 2b8cde54710..86b7dc8ffbc 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -234,3 +234,130 @@ class TestProxyMcpSimpleConnections: ) assert stdio_result == "5" assert streamable_result == "9" + + +class TestProxyMcpPrompts: + @pytest.mark.asyncio + async def test_list_prompts_single_server_stdio(self, proxy_server_url: str) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", + headers={ + "Authorization": PROXY_AUTHORIZATION_HEADER, + "x-mcp-servers": "math_stdio", + }, + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + result = await session.list_prompts() + prompt_names = {p.name for p in result.prompts} + assert "code_review" in prompt_names + assert "summarize" in prompt_names + + @pytest.mark.asyncio + async def test_list_prompts_single_server_streamable_http( + self, proxy_server_url: str + ) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", + headers={ + "Authorization": PROXY_AUTHORIZATION_HEADER, + "x-mcp-servers": "math_streamable_http", + }, + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + result = await session.list_prompts() + prompt_names = {p.name for p in result.prompts} + assert "code_review" in prompt_names + assert "summarize" in prompt_names + + @pytest.mark.asyncio + async def test_list_prompts_all_servers_prefixed( + self, proxy_server_url: str + ) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", + headers={"Authorization": PROXY_AUTHORIZATION_HEADER}, + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + result = await session.list_prompts() + prompt_names = {p.name for p in result.prompts} + expected = { + "math_stdio-code_review", + "math_stdio-summarize", + "math_streamable_http-code_review", + "math_streamable_http-summarize", + } + assert expected <= prompt_names + + @pytest.mark.asyncio + async def test_get_prompt_single_server_stdio(self, proxy_server_url: str) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", + headers={ + "Authorization": PROXY_AUTHORIZATION_HEADER, + "x-mcp-servers": "math_stdio", + }, + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + result = await session.get_prompt( + "code_review", arguments={"code": "def hello(): pass"} + ) + assert result.messages + first_msg = result.messages[0] + assert first_msg.role == "user" + text = getattr(first_msg.content, "text", None) + assert text is not None + assert "def hello(): pass" in text + + @pytest.mark.asyncio + async def test_get_prompt_single_server_streamable_http( + self, proxy_server_url: str + ) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", + headers={ + "Authorization": PROXY_AUTHORIZATION_HEADER, + "x-mcp-servers": "math_streamable_http", + }, + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + result = await session.get_prompt( + "summarize", + arguments={"text": "Hello world", "style": "brief"}, + ) + assert result.messages + first_msg = result.messages[0] + assert first_msg.role == "user" + text = getattr(first_msg.content, "text", None) + assert text is not None + assert "Hello world" in text + assert "brief" in text + + @pytest.mark.asyncio + async def test_get_prompt_prefixed_name_all_servers( + self, proxy_server_url: str + ) -> None: + async with asyncio.timeout(20): + async with streamablehttp_client( + url=f"{proxy_server_url}/mcp", + headers={"Authorization": PROXY_AUTHORIZATION_HEADER}, + ) as (read, write, _get_session_id): + async with ClientSession(read, write) as session: + await session.initialize() + result = await session.get_prompt( + "math_stdio-code_review", + arguments={"code": "x = 1"}, + ) + assert result.messages + text = getattr(result.messages[0].content, "text", None) + assert text is not None + assert "x = 1" in text 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 9c77edc6743..8ad4e795ecf 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 @@ -526,6 +526,325 @@ class TestListToolsRestAPI: assert result["message"] == "Successfully retrieved tools" +class TestListPromptsRestAPI: + pytestmark = pytest.mark.asyncio + + async def test_rejects_disallowed_server(self, monkeypatch): + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return [] + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + + request = _build_request(path="/mcp-rest/prompts/list", method="GET") + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.list_prompts_rest_api( + request, + server_id="server-1", + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "access_denied" + assert "server server-1" in exc_info.value.detail["message"] + + async def test_lists_prompts_for_allowed_server(self, monkeypatch): + from mcp.types import Prompt + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + alias = "server-1" + server_name = "server-1" + name = "stub" + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + + stub_server = StubServer() + + fake_prompt = Prompt( + name="code_review", + description="Review code", + arguments=[], + ) + + async def fake_get_prompts_from_server( + server, mcp_auth_header=None, add_prefix=True, raw_headers=None, **kwargs + ): + return [fake_prompt] + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_prompts_from_server", + fake_get_prompts_from_server, + raising=False, + ) + + request = _build_request(path="/mcp-rest/prompts/list", method="GET") + result = await rest_endpoints.list_prompts_rest_api( + request, + server_id="server-1", + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert len(result["prompts"]) == 1 + assert result["prompts"][0]["name"] == "code_review" + assert result["nextCursor"] is None + assert result["error"] is None + assert result["message"] == "Successfully retrieved prompts" + + async def test_returns_nextcursor_in_response(self, monkeypatch): + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + alias = "server-1" + server_name = "server-1" + name = "stub" + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + + stub_server = StubServer() + + async def fake_get_prompts_from_server( + server, mcp_auth_header=None, add_prefix=True, raw_headers=None, **kwargs + ): + return [] + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_prompts_from_server", + fake_get_prompts_from_server, + raising=False, + ) + + request = _build_request(path="/mcp-rest/prompts/list", method="GET") + result = await rest_endpoints.list_prompts_rest_api( + request, + server_id="server-1", + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert "nextCursor" in result + assert result["nextCursor"] is None + + +class TestGetPromptRestAPI: + pytestmark = pytest.mark.asyncio + + async def test_rejects_missing_server_id(self, monkeypatch): + request = _build_request( + path="/mcp-rest/prompts/get", + method="POST", + json_body={"name": "code_review"}, + ) + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.get_prompt_rest_api( + request, + user_api_key_dict=UserAPIKeyAuth(), + ) + assert exc_info.value.status_code == 400 + assert exc_info.value.detail["error"] == "missing_parameter" + assert "server_id" in exc_info.value.detail["message"] + + async def test_rejects_missing_name(self, monkeypatch): + request = _build_request( + path="/mcp-rest/prompts/get", + method="POST", + json_body={"server_id": "server-1"}, + ) + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.get_prompt_rest_api( + request, + user_api_key_dict=UserAPIKeyAuth(), + ) + assert exc_info.value.status_code == 400 + assert exc_info.value.detail["error"] == "missing_parameter" + assert "name" in exc_info.value.detail["message"] + + async def test_rejects_disallowed_server(self, monkeypatch): + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return [] + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + + request = _build_request( + path="/mcp-rest/prompts/get", + method="POST", + json_body={"server_id": "server-1", "name": "code_review"}, + ) + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.get_prompt_rest_api( + request, + user_api_key_dict=UserAPIKeyAuth(), + ) + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "access_denied" + + async def test_gets_prompt_when_allowed(self, monkeypatch): + from mcp.types import GetPromptResult, PromptMessage, TextContent + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + alias = "server-1" + server_name = "server-1" + name = "stub" + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + + stub_server = StubServer() + + fake_result = GetPromptResult( + description="Code review prompt", + messages=[ + PromptMessage( + role="user", + content=TextContent(type="text", text="Please review: x = 1"), + ) + ], + ) + + captured = {} + + async def fake_get_prompt_from_server( + server, prompt_name, arguments=None, mcp_auth_header=None, raw_headers=None, **kwargs + ): + captured["prompt_name"] = prompt_name + captured["arguments"] = arguments + return fake_result + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_prompt_from_server", + fake_get_prompt_from_server, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "filter_server_ids_by_ip", + lambda ids, ip: ids, + raising=False, + ) + + request = _build_request( + path="/mcp-rest/prompts/get", + method="POST", + json_body={ + "server_id": "server-1", + "name": "code_review", + "arguments": {"code": "x = 1"}, + }, + ) + result = await rest_endpoints.get_prompt_rest_api( + request, + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert result["description"] == "Code review prompt" + assert len(result["messages"]) == 1 + assert result["messages"][0]["role"] == "user" + assert result["error"] is None + assert result["message"] == "Successfully retrieved prompt" + assert captured["prompt_name"] == "code_review" + assert captured["arguments"] == {"code": "x = 1"} + + class TestCallToolRestAPI: pytestmark = pytest.mark.asyncio