Enforce per-tool auth checks in SamplingTool.from_callable_tool wrapper (#3494)

Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
Jeremiah Lowin 2026-03-14 16:14:50 -04:00 committed by GitHub
commit ca8069cb86
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 169 additions and 0 deletions

View file

@ -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):

View file

@ -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)