mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-22 05:24:18 +02:00
Enforce per-tool auth checks in SamplingTool.from_callable_tool wrapper (#3494)
Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
parent
01c57a9e04
commit
ca8069cb86
2 changed files with 169 additions and 0 deletions
|
|
@ -10,6 +10,9 @@ from mcp.types import TextContent
|
|||
from mcp.types import Tool as SDKTool
|
||||
from pydantic import ConfigDict
|
||||
|
||||
from fastmcp.exceptions import AuthorizationError
|
||||
from fastmcp.server.auth.authorization import AuthContext, run_auth_checks
|
||||
from fastmcp.server.dependencies import get_access_token
|
||||
from fastmcp.tools.function_parsing import ParsedFunction
|
||||
from fastmcp.tools.function_tool import FunctionTool
|
||||
from fastmcp.tools.tool import ToolResult
|
||||
|
|
@ -151,6 +154,24 @@ class SamplingTool(FastMCPBaseModel):
|
|||
# Both FunctionTool and TransformedTool need .run() to ensure proper
|
||||
# result processing (serializers, output_schema, wrap-result flags)
|
||||
async def wrapper(**kwargs: Any) -> Any:
|
||||
# Enforce per-tool auth checks, mirroring what the server
|
||||
# dispatcher does for direct tool calls. Without this, an
|
||||
# auth-protected tool wrapped as a SamplingTool could be
|
||||
# invoked by the LLM during sampling without authorization.
|
||||
if tool.auth is not None:
|
||||
# Late import to avoid circular import with context.py
|
||||
from fastmcp.server.context import _current_transport
|
||||
|
||||
is_stdio = _current_transport.get() == "stdio"
|
||||
if not is_stdio:
|
||||
token = get_access_token()
|
||||
ctx = AuthContext(token=token, component=tool)
|
||||
if not await run_auth_checks(tool.auth, ctx):
|
||||
raise AuthorizationError(
|
||||
f"Authorization failed for tool '{tool.name}': "
|
||||
"insufficient permissions"
|
||||
)
|
||||
|
||||
result = await tool.run(kwargs)
|
||||
# Unwrap ToolResult - extract the actual value
|
||||
if isinstance(result, ToolResult):
|
||||
|
|
|
|||
|
|
@ -1,7 +1,12 @@
|
|||
"""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
|
||||
|
|
@ -290,3 +295,146 @@ class TestSamplingToolFromCallableTool:
|
|||
|
||||
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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue