mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-10-10 23:13:20 +02:00
171 lines
5.2 KiB
Python
171 lines
5.2 KiB
Python
"""Generator tool bodies follow the synchronous thread-dispatch policy."""
|
|
|
|
import asyncio
|
|
import functools
|
|
import json
|
|
import threading
|
|
from collections.abc import AsyncIterator, Callable, Iterator
|
|
from contextlib import contextmanager
|
|
from typing import Any
|
|
|
|
import pytest
|
|
from mcp_types import TextContent
|
|
|
|
from fastmcp import Client, Context, FastMCP
|
|
from fastmcp.dependencies import Depends
|
|
from fastmcp.exceptions import ToolError
|
|
|
|
|
|
def _sync_wrapper(fn: Callable[..., Any]) -> Callable[..., Any]:
|
|
@functools.wraps(fn)
|
|
def wrapped(*args: Any, **kwargs: Any) -> Any:
|
|
return fn(*args, **kwargs)
|
|
|
|
return wrapped
|
|
|
|
|
|
@pytest.mark.parametrize("run_in_thread", [True, False])
|
|
@pytest.mark.parametrize("inject_context", [True, False])
|
|
async def test_sync_generator_obeys_thread_dispatch(
|
|
run_in_thread: bool, inject_context: bool
|
|
) -> None:
|
|
mcp = FastMCP()
|
|
loop_thread = threading.get_ident()
|
|
|
|
if inject_context:
|
|
|
|
@mcp.tool(run_in_thread=run_in_thread)
|
|
def thread_ids(ctx: Context) -> Iterator[int]:
|
|
assert ctx.fastmcp is mcp
|
|
yield threading.get_ident()
|
|
yield threading.get_ident()
|
|
|
|
else:
|
|
|
|
@mcp.tool(run_in_thread=run_in_thread)
|
|
def thread_ids() -> Iterator[int]:
|
|
yield threading.get_ident()
|
|
yield threading.get_ident()
|
|
|
|
async with Client(mcp) as client:
|
|
result = await client.call_tool("thread_ids")
|
|
|
|
assert isinstance(result.content[0], TextContent)
|
|
ids = json.loads(result.content[0].text)
|
|
assert len(ids) == 2
|
|
assert all((tid != loop_thread) == run_in_thread for tid in ids)
|
|
|
|
|
|
@pytest.mark.parametrize("run_in_thread", [True, False])
|
|
async def test_sync_factory_generator_obeys_thread_dispatch(
|
|
run_in_thread: bool,
|
|
) -> None:
|
|
mcp = FastMCP()
|
|
loop_thread = threading.get_ident()
|
|
|
|
@mcp.tool(run_in_thread=run_in_thread)
|
|
def thread_ids() -> Iterator[int]:
|
|
return (threading.get_ident() for _ in range(2))
|
|
|
|
async with Client(mcp) as client:
|
|
result = await client.call_tool("thread_ids")
|
|
|
|
assert isinstance(result.content[0], TextContent)
|
|
ids = json.loads(result.content[0].text)
|
|
assert len(ids) == 2
|
|
assert all((tid != loop_thread) == run_in_thread for tid in ids)
|
|
|
|
|
|
@pytest.mark.parametrize("run_in_thread", [True, False])
|
|
async def test_async_generator_stays_on_event_loop(run_in_thread: bool) -> None:
|
|
mcp = FastMCP()
|
|
loop_thread = threading.get_ident()
|
|
|
|
@mcp.tool(run_in_thread=run_in_thread)
|
|
async def thread_ids() -> AsyncIterator[int]:
|
|
yield threading.get_ident()
|
|
|
|
async with Client(mcp) as client:
|
|
result = await client.call_tool("thread_ids")
|
|
|
|
assert isinstance(result.content[0], TextContent)
|
|
assert json.loads(result.content[0].text) == [loop_thread]
|
|
|
|
|
|
@pytest.mark.parametrize("run_in_thread", [True, False])
|
|
async def test_async_factory_generator_stays_on_event_loop(run_in_thread: bool) -> None:
|
|
mcp = FastMCP()
|
|
loop_thread = threading.get_ident()
|
|
|
|
@mcp.tool(run_in_thread=run_in_thread)
|
|
async def thread_ids() -> Iterator[int]:
|
|
return (threading.get_ident() for _ in range(2))
|
|
|
|
async with Client(mcp) as client:
|
|
result = await client.call_tool("thread_ids")
|
|
|
|
assert isinstance(result.content[0], TextContent)
|
|
assert json.loads(result.content[0].text) == [loop_thread, loop_thread]
|
|
|
|
|
|
@pytest.mark.parametrize("run_in_thread", [True, False])
|
|
async def test_generator_iteration_error_is_tool_error(run_in_thread: bool) -> None:
|
|
mcp = FastMCP()
|
|
|
|
@mcp.tool(run_in_thread=run_in_thread)
|
|
def broken() -> Iterator[int]:
|
|
yield 1
|
|
raise ValueError("generator failed")
|
|
|
|
async with Client(mcp) as client:
|
|
with pytest.raises(ToolError, match="generator failed"):
|
|
await client.call_tool("broken")
|
|
|
|
|
|
@pytest.mark.parametrize("inject_context", [True, False])
|
|
async def test_sync_wrapper_of_async_generator_factory_stays_on_loop(
|
|
inject_context: bool,
|
|
) -> None:
|
|
mcp = FastMCP()
|
|
|
|
if inject_context:
|
|
|
|
async def thread_ids(ctx: Context) -> Iterator[bool]:
|
|
assert ctx.fastmcp is mcp
|
|
return (asyncio.get_running_loop().is_running() for _ in range(2))
|
|
|
|
else:
|
|
|
|
async def thread_ids() -> Iterator[bool]:
|
|
return (asyncio.get_running_loop().is_running() for _ in range(2))
|
|
|
|
mcp.tool(_sync_wrapper(thread_ids))
|
|
async with Client(mcp) as client:
|
|
result = await client.call_tool("thread_ids")
|
|
|
|
assert isinstance(result.content[0], TextContent)
|
|
assert json.loads(result.content[0].text) == [True, True]
|
|
|
|
|
|
async def test_sync_generator_consumed_before_dependency_cleanup() -> None:
|
|
mcp = FastMCP()
|
|
state = {"open": False}
|
|
|
|
@contextmanager
|
|
def connection() -> Iterator[dict[str, bool]]:
|
|
state["open"] = True
|
|
try:
|
|
yield state
|
|
finally:
|
|
state["open"] = False
|
|
|
|
@mcp.tool
|
|
def query(resource: dict[str, bool] = Depends(connection)) -> Iterator[bool]:
|
|
yield resource["open"]
|
|
|
|
async with Client(mcp) as client:
|
|
result = await client.call_tool("query")
|
|
|
|
assert isinstance(result.content[0], TextContent)
|
|
assert json.loads(result.content[0].text) == [True]
|
|
assert state["open"] is False
|