mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-18 03:29:11 +02:00
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com> Co-authored-by: Jeremiah Lowin <jlowin@users.noreply.github.com> Co-authored-by: Marvin Context Protocol <41898282+Marvin Context Protocol@users.noreply.github.com> Co-authored-by: voidborne-d <voidborne-d@users.noreply.github.com> Co-authored-by: marvin-context-protocol[bot] <225465937+marvin-context-protocol[bot]@users.noreply.github.com> Co-authored-by: Claude <noreply@anthropic.com> Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> Co-authored-by: d 🔹 <258577966+voidborne-d@users.noreply.github.com> Co-authored-by: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Co-authored-by: nightcityblade <nightcityblade@gmail.com> Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> Co-authored-by: Bill Easton <strawgate@users.noreply.github.com> Co-authored-by: Sumanshu Nankana <sumanshunankana@gmail.com> Co-authored-by: Eric Robinson <ericrobinson@indeed.com> Co-authored-by: Martim Santos <martimfasantos@gmail.com> Co-authored-by: d 🔹 <liusway405@gmail.com> Co-authored-by: Matthieu B <66959271+mtthidoteu@users.noreply.github.com> Co-authored-by: Sascha Buehrle <47737812+saschabuehrle@users.noreply.github.com> Co-authored-by: Hakancan <142545736+hkc5@users.noreply.github.com> Co-authored-by: nightcityblade <jackchen@haloailabs.com> Co-authored-by: Matt Hallowell <17804673+mhallo@users.noreply.github.com> Co-authored-by: nate nowack <thrast36@gmail.com> Co-authored-by: Bill Easton <williamseaston@gmail.com> Co-authored-by: Marcus Shu <46469249+shulkx@users.noreply.github.com> Co-authored-by: Rushabh Doshi <radoshi@gmail.com> Co-authored-by: AIKAWA Shigechika <shige@aikawa.jp> Co-authored-by: Jeremy Simon <simonjer805@gmail.com> Co-authored-by: Miguel Miranda Dias <7780875+pandego@users.noreply.github.com> Co-authored-by: Anthony James Padavano <padavano.anthony@gmail.com> Co-authored-by: Mostafa Kamal <hiremostafa@gmail.com> Fix auto-close MRE script posting comment without closing (#3386) Fix WorkOS token scope verification bypass 🤖 Generated with Codex (#3407) Fix initialize McpError fallthrough 🤖 Generated with Codex (#3413) Fix transform arg collisions with passthrough params (#3431) Fix get_* returning None when latest version is disabled (#3439) Fix get_* returning None when latest version is disabled (#3421) Fix server lifespan overlap teardown (#3415) Fix $ref output schema object detection regression (#3420) resolved annotations (#3429) Fix async partial callables rejected by iscoroutinefunction (#3438) Fix async partial callables rejected by iscoroutinefunction (#3423) fix: add version to components (#3458) fix: use intent-based flag for OIDC scope patch in load_access_token (#3465) Fixes #3461 fix: normalize Google scope shorthands and surface valid_scopes (#3477) fix: resolve ty 0.0.23 type-checking errors and bump pin (#3481) fix: shield lifespan teardown from cancellation (#3480) fix: forward custom_route endpoints from mounted servers (#3462) fix updates _get_additional_http_routes() to traverse providers, Fixes #3457 fix: remove hardcoded version from CLI help text (#3456) fix: monty 0.0.8 compatibility, drop external_functions from constructor (#3468) fix: task test teardown hanging 5s per test (#3499) Closes #3498 fix: validate workspace path is a directory before cursor install (#3440) Fixes #3426 fix: handle re.error from malformed URI templates in build_regex (#3501) fix: reject empty/OIDC-only required_scopes in AzureProvider (#3503) fix: restrict $ref resolution to local refs only (SSRF/LFI) (#3502) fix warnings and timeouts (#3504) close upgrade check issue when build passes (#3505) Closes #3484 fix: URL-encode path params to prevent SSRF/path traversal (GHSA-vv7q-7jx5-f767) (#3507) fix: prevent path traversal in skill download (#3493) fix: prefer IdP-granted scopes over client-requested scopes in OAuthProxy (#3492) fix: remove unrelated transform and http.py changes from PR scope fix: remove forced follow_redirects from httpx_client_factory calls (#3496) fix: stop passing follow_redirects to httpx_client_factory fix: restore follow_redirects=True for custom httpx client factories Closes #3509 fix: CSRF double-submit cookie check in consent flow (#3519) fix: validate server names in install commands (#3522) fix: use raw strings for regex in pytest.raises match (#3523) fix: reject refresh tokens used as Bearer access tokens (#3524) fix: route ResourcesAsTools/PromptsAsTools through server middleware (#3495) fix: resolve Pyright "Module is not callable" on @tool, @resource, @prompt decorators (#3540) fix: filter warnings by message in KEY_PREFIX test (#3549) fix: suppress output schema for ToolResult subclass annotations (#3548) fix: increase sleep duration in proxy cache tests (#3567) fix: store absolute token expiry to prevent stale expires_in on reload (#3572) fix: preserve tool properties named 'title' during schema compression (#3582) Fix loopback redirect URI port matching per RFC 8252 §7.3 (#3589) Fix app tool routing: visibility check and middleware propagation (#3591) Fix query parameter serialization to respect OpenAPI explode/style settings (#3595) Fix dev apps form: union types, textarea support, JSON parsing (#3597) fix(google): replace deprecated /oauth2/v1/tokeninfo with /oauth2/v3/userinfo (#3603) fix: resolve EntraOBOToken dependency injection through MultiAuth (#3609) fix(docs): correct misleading stateless_http header (#3622) fix: filesystem provider import machinery (#3626) Closes #3625 (issues 2, 3, 6) fix: recover StdioTransport after subprocess exits (#3630) fix(server): preserve mounted tool task metadata (#3632) fix: scope deprecation warning filter to FastMCPDeprecationWarning (#3649) fix imports, add PrefabAppConfig (#3650) fix: resolve CurrentFastMCP/ctx.fastmcp to child server in mounted background tasks (#3651) Fix blocking docs issues: chart imports, Select API, Rx consistency (#3652) closed by default (#3657) Fix prompt caching middleware missing wrap/unwrap round-trip (#3666) fix: serialize object query params per OpenAPI style/explode rules (#3662) Fixes #2857 fix: HTTP request headers not accessible in background task workers (#3631) fix: restore HTTP headers in worker execution path for background tasks (#3681) fix: strip discriminator after dereferencing schemas (#3682) fix: remove stale ty:ignore directives for ty 0.0.26 (#3684) Fix docs gaps in app provider pages (#3690) fix: dev apps log panel UX improvements (#3698) fix dev server empty string args (#3700)
440 lines
15 KiB
Python
440 lines
15 KiB
Python
"""Tests for SamplingTool."""
|
|
|
|
import pytest
|
|
from mcp.server.auth.middleware.auth_context import auth_context_var
|
|
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
|
|
|
|
from fastmcp.exceptions import AuthorizationError
|
|
from fastmcp.server.auth import AccessToken, require_scopes
|
|
from fastmcp.server.context import _current_transport
|
|
from fastmcp.server.sampling import SamplingTool
|
|
from fastmcp.tools.function_tool import FunctionTool
|
|
from fastmcp.tools.tool_transform import ArgTransform, TransformedTool
|
|
|
|
|
|
class TestSamplingToolFromFunction:
|
|
"""Tests for SamplingTool.from_function()."""
|
|
|
|
def test_from_simple_function(self):
|
|
def search(query: str) -> str:
|
|
"""Search the web."""
|
|
return f"Results for: {query}"
|
|
|
|
tool = SamplingTool.from_function(search)
|
|
|
|
assert tool.name == "search"
|
|
assert tool.description == "Search the web."
|
|
assert "query" in tool.parameters.get("properties", {})
|
|
assert tool.fn is search
|
|
|
|
def test_from_function_with_overrides(self):
|
|
def search(query: str) -> str:
|
|
return f"Results for: {query}"
|
|
|
|
tool = SamplingTool.from_function(
|
|
search,
|
|
name="web_search",
|
|
description="Search the internet",
|
|
)
|
|
|
|
assert tool.name == "web_search"
|
|
assert tool.description == "Search the internet"
|
|
|
|
def test_from_lambda_requires_name(self):
|
|
with pytest.raises(ValueError, match="must provide a name for lambda"):
|
|
SamplingTool.from_function(lambda x: x)
|
|
|
|
def test_from_lambda_with_name(self):
|
|
tool = SamplingTool.from_function(lambda x: x * 2, name="double")
|
|
|
|
assert tool.name == "double"
|
|
|
|
def test_from_async_function(self):
|
|
async def async_search(query: str) -> str:
|
|
"""Async search."""
|
|
return f"Async results for: {query}"
|
|
|
|
tool = SamplingTool.from_function(async_search)
|
|
|
|
assert tool.name == "async_search"
|
|
assert tool.description == "Async search."
|
|
|
|
def test_multiple_parameters(self):
|
|
def search(query: str, limit: int = 10, include_images: bool = False) -> str:
|
|
"""Search with options."""
|
|
return f"Results for: {query}"
|
|
|
|
tool = SamplingTool.from_function(search)
|
|
props = tool.parameters.get("properties", {})
|
|
|
|
assert "query" in props
|
|
assert "limit" in props
|
|
assert "include_images" in props
|
|
|
|
|
|
class TestSamplingToolRun:
|
|
"""Tests for SamplingTool.run()."""
|
|
|
|
async def test_run_sync_function(self):
|
|
def add(a: int, b: int) -> int:
|
|
"""Add two numbers."""
|
|
return a + b
|
|
|
|
tool = SamplingTool.from_function(add)
|
|
result = await tool.run({"a": 2, "b": 3})
|
|
assert result == 5
|
|
|
|
async def test_run_async_function(self):
|
|
async def async_add(a: int, b: int) -> int:
|
|
"""Add two numbers asynchronously."""
|
|
return a + b
|
|
|
|
tool = SamplingTool.from_function(async_add)
|
|
result = await tool.run({"a": 2, "b": 3})
|
|
assert result == 5
|
|
|
|
async def test_run_with_no_arguments(self):
|
|
def get_value() -> str:
|
|
"""Return a fixed value."""
|
|
return "hello"
|
|
|
|
tool = SamplingTool.from_function(get_value)
|
|
result = await tool.run()
|
|
assert result == "hello"
|
|
|
|
async def test_run_with_none_arguments(self):
|
|
def get_value() -> str:
|
|
"""Return a fixed value."""
|
|
return "hello"
|
|
|
|
tool = SamplingTool.from_function(get_value)
|
|
result = await tool.run(None)
|
|
assert result == "hello"
|
|
|
|
|
|
class TestSamplingToolSDKConversion:
|
|
"""Tests for SamplingTool._to_sdk_tool() internal method."""
|
|
|
|
def test_to_sdk_tool(self):
|
|
def search(query: str) -> str:
|
|
"""Search the web."""
|
|
return f"Results for: {query}"
|
|
|
|
tool = SamplingTool.from_function(search)
|
|
sdk_tool = tool._to_sdk_tool()
|
|
|
|
assert sdk_tool.name == "search"
|
|
assert sdk_tool.description == "Search the web."
|
|
assert "query" in sdk_tool.inputSchema.get("properties", {})
|
|
|
|
|
|
class TestSamplingToolFromCallableTool:
|
|
"""Tests for SamplingTool.from_callable_tool()."""
|
|
|
|
def test_from_function_tool(self):
|
|
"""Test converting a FunctionTool to SamplingTool."""
|
|
|
|
def search(query: str) -> str:
|
|
"""Search the web."""
|
|
return f"Results for: {query}"
|
|
|
|
function_tool = FunctionTool.from_function(search)
|
|
sampling_tool = SamplingTool.from_callable_tool(function_tool)
|
|
|
|
assert sampling_tool.name == "search"
|
|
assert sampling_tool.description == "Search the web."
|
|
assert "query" in sampling_tool.parameters.get("properties", {})
|
|
# fn is now a wrapper that calls tool.run() for proper result processing
|
|
assert callable(sampling_tool.fn)
|
|
|
|
def test_from_function_tool_with_overrides(self):
|
|
"""Test converting FunctionTool with name/description overrides."""
|
|
|
|
def search(query: str) -> str:
|
|
"""Search the web."""
|
|
return f"Results for: {query}"
|
|
|
|
function_tool = FunctionTool.from_function(search)
|
|
sampling_tool = SamplingTool.from_callable_tool(
|
|
function_tool,
|
|
name="web_search",
|
|
description="Search the internet",
|
|
)
|
|
|
|
assert sampling_tool.name == "web_search"
|
|
assert sampling_tool.description == "Search the internet"
|
|
|
|
def test_from_transformed_tool(self):
|
|
"""Test converting a TransformedTool to SamplingTool."""
|
|
|
|
def original(query: str, limit: int) -> str:
|
|
"""Original tool."""
|
|
return f"Results for: {query} (limit: {limit})"
|
|
|
|
function_tool = FunctionTool.from_function(original)
|
|
transformed_tool = TransformedTool.from_tool(
|
|
function_tool,
|
|
name="search_transformed",
|
|
transform_args={"query": ArgTransform(name="q")},
|
|
)
|
|
|
|
sampling_tool = SamplingTool.from_callable_tool(transformed_tool)
|
|
|
|
assert sampling_tool.name == "search_transformed"
|
|
assert sampling_tool.description == "Original tool."
|
|
# The transformed tool should have 'q' instead of 'query'
|
|
assert "q" in sampling_tool.parameters.get("properties", {})
|
|
assert "limit" in sampling_tool.parameters.get("properties", {})
|
|
|
|
async def test_from_function_tool_execution(self):
|
|
"""Test that converted FunctionTool executes correctly."""
|
|
|
|
def add(a: int, b: int) -> int:
|
|
"""Add two numbers."""
|
|
return a + b
|
|
|
|
function_tool = FunctionTool.from_function(add)
|
|
sampling_tool = SamplingTool.from_callable_tool(function_tool)
|
|
|
|
result = await sampling_tool.run({"a": 2, "b": 3})
|
|
assert result == 5
|
|
|
|
async def test_from_transformed_tool_execution(self):
|
|
"""Test that converted TransformedTool executes correctly."""
|
|
|
|
def multiply(x: int, y: int) -> int:
|
|
"""Multiply two numbers."""
|
|
return x * y
|
|
|
|
function_tool = FunctionTool.from_function(multiply)
|
|
transformed_tool = TransformedTool.from_tool(
|
|
function_tool,
|
|
transform_args={"x": ArgTransform(name="a"), "y": ArgTransform(name="b")},
|
|
)
|
|
|
|
sampling_tool = SamplingTool.from_callable_tool(transformed_tool)
|
|
|
|
# Use the transformed parameter names
|
|
result = await sampling_tool.run({"a": 3, "b": 4})
|
|
# Result should be unwrapped from ToolResult
|
|
assert result == 12
|
|
|
|
def test_from_invalid_tool_type(self):
|
|
"""Test that from_callable_tool rejects non-tool objects."""
|
|
|
|
class NotATool:
|
|
pass
|
|
|
|
with pytest.raises(
|
|
TypeError,
|
|
match="Expected FunctionTool or TransformedTool",
|
|
):
|
|
SamplingTool.from_callable_tool(NotATool()) # type: ignore[arg-type] # ty:ignore[invalid-argument-type]
|
|
|
|
def test_from_plain_function_fails(self):
|
|
"""Test that plain functions are rejected by from_callable_tool."""
|
|
|
|
def my_function():
|
|
pass
|
|
|
|
with pytest.raises(TypeError, match="Expected FunctionTool or TransformedTool"):
|
|
SamplingTool.from_callable_tool(my_function) # type: ignore[arg-type] # ty:ignore[invalid-argument-type]
|
|
|
|
async def test_from_function_tool_with_output_schema(self):
|
|
"""Test that FunctionTool with output_schema is handled correctly."""
|
|
|
|
def search(query: str) -> dict:
|
|
"""Search for something."""
|
|
return {"results": ["item1", "item2"], "count": 2}
|
|
|
|
# Create FunctionTool with x-fastmcp-wrap-result
|
|
function_tool = FunctionTool.from_function(
|
|
search,
|
|
output_schema={
|
|
"type": "object",
|
|
"properties": {
|
|
"results": {"type": "array"},
|
|
"count": {"type": "integer"},
|
|
},
|
|
"x-fastmcp-wrap-result": True,
|
|
},
|
|
)
|
|
|
|
sampling_tool = SamplingTool.from_callable_tool(function_tool)
|
|
|
|
# Run the tool - should unwrap the {"result": {...}} wrapper
|
|
result = await sampling_tool.run({"query": "test"})
|
|
|
|
# Should get the unwrapped dict, not ToolResult
|
|
assert isinstance(result, dict)
|
|
assert result == {"results": ["item1", "item2"], "count": 2}
|
|
|
|
async def test_from_function_tool_without_wrap_result(self):
|
|
"""Test that FunctionTool without x-fastmcp-wrap-result is handled correctly."""
|
|
|
|
def get_data() -> dict:
|
|
"""Get some data."""
|
|
return {"status": "ok", "value": 42}
|
|
|
|
# Create FunctionTool with output_schema but no wrap-result flag
|
|
function_tool = FunctionTool.from_function(
|
|
get_data,
|
|
output_schema={
|
|
"type": "object",
|
|
"properties": {
|
|
"status": {"type": "string"},
|
|
"value": {"type": "integer"},
|
|
},
|
|
},
|
|
)
|
|
|
|
sampling_tool = SamplingTool.from_callable_tool(function_tool)
|
|
|
|
# Run the tool - should return structured_content directly
|
|
result = await sampling_tool.run({})
|
|
|
|
assert isinstance(result, dict)
|
|
assert result == {"status": "ok", "value": 42}
|
|
|
|
|
|
class TestSamplingToolAuthEnforcement:
|
|
"""Tests that auth-protected tools enforce auth when used via sampling."""
|
|
|
|
async def test_auth_protected_tool_blocked_without_token(self):
|
|
"""An auth-protected tool wrapped as SamplingTool must reject
|
|
calls when no valid token is present in a non-stdio transport."""
|
|
|
|
def secret_action() -> str:
|
|
"""Do something privileged."""
|
|
return "secret"
|
|
|
|
function_tool = FunctionTool.from_function(
|
|
secret_action,
|
|
auth=require_scopes("admin"),
|
|
)
|
|
sampling_tool = SamplingTool.from_callable_tool(function_tool)
|
|
|
|
transport_token = _current_transport.set("streamable-http")
|
|
try:
|
|
with pytest.raises(AuthorizationError, match="insufficient permissions"):
|
|
await sampling_tool.run({})
|
|
finally:
|
|
_current_transport.reset(transport_token)
|
|
|
|
async def test_auth_protected_tool_blocked_with_wrong_scopes(self):
|
|
"""An auth-protected tool rejects calls when the token lacks
|
|
the required scopes."""
|
|
|
|
def secret_action() -> str:
|
|
"""Do something privileged."""
|
|
return "secret"
|
|
|
|
function_tool = FunctionTool.from_function(
|
|
secret_action,
|
|
auth=require_scopes("admin"),
|
|
)
|
|
sampling_tool = SamplingTool.from_callable_tool(function_tool)
|
|
|
|
token = AccessToken(
|
|
token="test",
|
|
client_id="c",
|
|
scopes=["read"],
|
|
expires_at=None,
|
|
claims={},
|
|
)
|
|
transport_token = _current_transport.set("streamable-http")
|
|
auth_token = auth_context_var.set(AuthenticatedUser(token))
|
|
try:
|
|
with pytest.raises(AuthorizationError, match="insufficient permissions"):
|
|
await sampling_tool.run({})
|
|
finally:
|
|
auth_context_var.reset(auth_token)
|
|
_current_transport.reset(transport_token)
|
|
|
|
async def test_auth_protected_tool_allowed_with_correct_scopes(self):
|
|
"""An auth-protected tool succeeds when the token has the
|
|
required scopes."""
|
|
|
|
def secret_action() -> str:
|
|
"""Do something privileged."""
|
|
return "secret"
|
|
|
|
function_tool = FunctionTool.from_function(
|
|
secret_action,
|
|
auth=require_scopes("admin"),
|
|
)
|
|
sampling_tool = SamplingTool.from_callable_tool(function_tool)
|
|
|
|
token = AccessToken(
|
|
token="test",
|
|
client_id="c",
|
|
scopes=["admin"],
|
|
expires_at=None,
|
|
claims={},
|
|
)
|
|
transport_token = _current_transport.set("streamable-http")
|
|
auth_token = auth_context_var.set(AuthenticatedUser(token))
|
|
try:
|
|
result = await sampling_tool.run({})
|
|
assert result == "secret"
|
|
finally:
|
|
auth_context_var.reset(auth_token)
|
|
_current_transport.reset(transport_token)
|
|
|
|
async def test_auth_protected_tool_skipped_on_stdio(self):
|
|
"""Auth checks are skipped for stdio transport, matching
|
|
server dispatcher behavior."""
|
|
|
|
def secret_action() -> str:
|
|
"""Do something privileged."""
|
|
return "secret"
|
|
|
|
function_tool = FunctionTool.from_function(
|
|
secret_action,
|
|
auth=require_scopes("admin"),
|
|
)
|
|
sampling_tool = SamplingTool.from_callable_tool(function_tool)
|
|
|
|
transport_token = _current_transport.set("stdio")
|
|
try:
|
|
result = await sampling_tool.run({})
|
|
assert result == "secret"
|
|
finally:
|
|
_current_transport.reset(transport_token)
|
|
|
|
async def test_tool_without_auth_runs_normally(self):
|
|
"""Tools without auth still run without any auth context."""
|
|
|
|
def public_action() -> str:
|
|
"""Do something public."""
|
|
return "public"
|
|
|
|
function_tool = FunctionTool.from_function(public_action)
|
|
sampling_tool = SamplingTool.from_callable_tool(function_tool)
|
|
|
|
result = await sampling_tool.run({})
|
|
assert result == "public"
|
|
|
|
async def test_auth_protected_transformed_tool_blocked(self):
|
|
"""Auth checks also apply to TransformedTools with auth."""
|
|
|
|
def secret_action(x: int) -> int:
|
|
"""Privileged computation."""
|
|
return x * 2
|
|
|
|
function_tool = FunctionTool.from_function(
|
|
secret_action,
|
|
auth=require_scopes("compute"),
|
|
)
|
|
transformed_tool = TransformedTool.from_tool(
|
|
function_tool,
|
|
transform_args={"x": ArgTransform(name="value")},
|
|
)
|
|
sampling_tool = SamplingTool.from_callable_tool(transformed_tool)
|
|
|
|
transport_token = _current_transport.set("streamable-http")
|
|
try:
|
|
with pytest.raises(AuthorizationError, match="insufficient permissions"):
|
|
await sampling_tool.run({"value": 5})
|
|
finally:
|
|
_current_transport.reset(transport_token)
|