mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
add mcp prompts REST (list + get), e2e and unit tests
This commit is contained in:
parent
36784f7f41
commit
87bd995ec4
3 changed files with 458 additions and 0 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue