Spaces:
Running
Running
File size: 4,276 Bytes
73679eb ac33644 73679eb 224caf9 73679eb f6c5553 73679eb c14577c 73679eb ac33644 c1940a6 c29318c 73679eb ac33644 c14577c c29318c ac33644 d4fceae ac33644 73679eb ac33644 73679eb d4fceae 73679eb ac33644 73679eb ac33644 73679eb 3a9d761 73679eb ac33644 3a9d761 ac33644 d4fceae ac33644 73679eb ac33644 73679eb d4fceae 73679eb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 | import json
from collections.abc import Generator
import pytest
from fastmcp.client import Client
from fastmcp.client.transports import SSETransport, 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, **kwargs) -> None:
fastmcp_server().run(host=host, port=port, **kwargs)
@pytest.fixture(autouse=True, scope="module")
def shttp_server() -> Generator[str, None, None]:
with run_server_in_process(run_server, transport="http") as url:
yield f"{url}/mcp"
@pytest.fixture(autouse=True, scope="module")
def sse_server() -> Generator[str, None, None]:
with run_server_in_process(run_server, transport="sse") as url:
yield f"{url}/sse"
async def test_http_headers_resource_shttp(shttp_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-DEMO-HEADER": "ABC"}
)
) as client:
raw_result = await client.read_resource("request://headers")
json_result = json.loads(raw_result[0].text) # type: ignore[attr-defined]
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_resource_sse(sse_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
) as client:
raw_result = await client.read_resource("request://headers")
json_result = json.loads(raw_result[0].text) # type: ignore[attr-defined]
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_tool_shttp(shttp_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-DEMO-HEADER": "ABC"}
)
) as client:
result = await client.call_tool("get_headers_tool")
assert "x-demo-header" in result.data
assert result.data["x-demo-header"] == "ABC"
async def test_http_headers_tool_sse(sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
) as client:
result = await client.call_tool("get_headers_tool")
assert "x-demo-header" in result.data
assert result.data["x-demo-header"] == "ABC"
async def test_http_headers_prompt_shttp(shttp_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-DEMO-HEADER": "ABC"}
)
) as client:
result = await client.get_prompt("get_headers_prompt")
json_result = json.loads(result.messages[0].content.text) # type: ignore[attr-defined]
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_prompt_sse(sse_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
) as client:
result = await client.get_prompt("get_headers_prompt")
json_result = json.loads(result.messages[0].content.text) # type: ignore[attr-defined]
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
|