mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-15 01:59:10 +02:00
102 lines
3.2 KiB
Python
102 lines
3.2 KiB
Python
import json
|
|
import sys
|
|
from collections.abc import Generator
|
|
|
|
import pytest
|
|
import uvicorn
|
|
from mcp.types import TextContent, TextResourceContents
|
|
|
|
from fastmcp.client import Client
|
|
from fastmcp.client.transports import StreamableHttpTransport
|
|
from fastmcp.server.dependencies import get_http_request
|
|
from fastmcp.server.server import FastMCP
|
|
from fastmcp.utilities.tests import run_server_in_process
|
|
|
|
|
|
def fastmcp_server():
|
|
server = FastMCP()
|
|
|
|
# Add a tool
|
|
@server.tool()
|
|
def get_headers_tool() -> dict[str, str]:
|
|
"""Get the HTTP headers from the request."""
|
|
request = get_http_request()
|
|
|
|
return dict(request.headers)
|
|
|
|
@server.resource(uri="request://headers")
|
|
async def get_headers_resource() -> dict[str, str]:
|
|
request = get_http_request()
|
|
|
|
return dict(request.headers)
|
|
|
|
# Add a prompt
|
|
@server.prompt()
|
|
def get_headers_prompt() -> str:
|
|
"""Get the HTTP headers from the request."""
|
|
request = get_http_request()
|
|
|
|
return json.dumps(dict(request.headers))
|
|
|
|
return server
|
|
|
|
|
|
def run_server(host: str, port: int) -> None:
|
|
try:
|
|
app = fastmcp_server().http_app()
|
|
server = uvicorn.Server(
|
|
config=uvicorn.Config(
|
|
app=app,
|
|
host=host,
|
|
port=port,
|
|
log_level="error",
|
|
lifespan="on",
|
|
)
|
|
)
|
|
server.run()
|
|
except Exception as e:
|
|
print(f"Server error: {e}")
|
|
sys.exit(1)
|
|
sys.exit(0)
|
|
|
|
|
|
@pytest.fixture(autouse=True, scope="module")
|
|
def sse_server() -> Generator[str, None, None]:
|
|
with run_server_in_process(run_server) as url:
|
|
yield f"{url}/mcp"
|
|
|
|
|
|
async def test_http_headers_resource(sse_server: str):
|
|
"""Test getting HTTP headers from the server."""
|
|
async with Client(
|
|
transport=StreamableHttpTransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
|
|
) as client:
|
|
raw_result = await client.read_resource("request://headers")
|
|
assert isinstance(raw_result[0], TextResourceContents)
|
|
json_result = json.loads(raw_result[0].text)
|
|
assert "x-demo-header" in json_result
|
|
assert json_result["x-demo-header"] == "ABC"
|
|
|
|
|
|
async def test_http_headers_tool(sse_server: str):
|
|
"""Test getting HTTP headers from the server."""
|
|
async with Client(
|
|
transport=StreamableHttpTransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
|
|
) as client:
|
|
result = await client.call_tool("get_headers_tool")
|
|
assert isinstance(result[0], TextContent)
|
|
json_result = json.loads(result[0].text)
|
|
assert "x-demo-header" in json_result
|
|
assert json_result["x-demo-header"] == "ABC"
|
|
|
|
|
|
async def test_http_headers_prompt(sse_server: str):
|
|
"""Test getting HTTP headers from the server."""
|
|
async with Client(
|
|
transport=StreamableHttpTransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
|
|
) as client:
|
|
result = await client.get_prompt("get_headers_prompt")
|
|
assert isinstance(result.messages[0].content, TextContent)
|
|
json_result = json.loads(result.messages[0].content.text)
|
|
assert "x-demo-header" in json_result
|
|
assert json_result["x-demo-header"] == "ABC"
|