Compare commits

...

3 commits

Author SHA1 Message Date
strawgate
8100fc713f Fix ruff SIM105: use contextlib.suppress
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
2026-05-12 23:00:15 -05:00
strawgate
9fd60025f8 fix: preserve annotation metadata when cloning wrapped partials
Addresses Codex P1: when reconstructing a partial to strip __wrapped__,
copy over functools.WRAPPER_ASSIGNMENTS (__module__, __qualname__, etc.)
so deferred annotation resolution works correctly.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
2026-05-12 22:46:54 -05:00
William Easton
e3f02374c7 Support functools.partial and centralize callable utilities
functools.partial objects failed at registration (@mcp.tool didn't
recognize them) and at call time (update_wrapper set __wrapped__
causing Pydantic to ignore bound arguments).

Introduces centralized utilities replacing scattered patterns across
11 files:

callable_utils.py:
- is_callable_object(): TypeGuard replacing inspect.isroutine() in
  7 decorator entry points — recognizes partials as callables
- get_callable_name(): Extracts useful names from any callable type,
  including partials without update_wrapper
- prepare_callable(): Strips __wrapped__, unwraps callable classes
  and staticmethod — replaces 4 duplicated blocks

decorators.py:
- set_fastmcp_meta(): Attaches __fastmcp__ metadata through __func__
  for bound methods — replaces 5 identical 2-line blocks

TaskConfig:
- normalize(): Converts bool|TaskConfig|None to TaskConfig — replaces
  4 identical 6-line if/elif/else blocks

No behavior changes beyond the bug fix: existing lambda rejection,
validation, and error handling remain per-module policy.

Closes #3266

🤖 Generated with Claude Code

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-05-12 22:46:13 -05:00
14 changed files with 419 additions and 102 deletions

View file

@ -27,7 +27,6 @@ Usage::
from __future__ import annotations from __future__ import annotations
import inspect
from collections.abc import AsyncIterator, Callable, Sequence from collections.abc import AsyncIterator, Callable, Sequence
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from typing import Any, Literal, TypeVar, overload from typing import Any, Literal, TypeVar, overload
@ -38,6 +37,7 @@ from fastmcp.server.auth.authorization import AuthCheck
from fastmcp.server.providers.base import Provider from fastmcp.server.providers.base import Provider
from fastmcp.server.providers.local_provider import LocalProvider from fastmcp.server.providers.local_provider import LocalProvider
from fastmcp.tools.base import Tool from fastmcp.tools.base import Tool
from fastmcp.utilities.callable_utils import is_callable_object
from fastmcp.utilities.logging import get_logger from fastmcp.utilities.logging import get_logger
logger = get_logger(__name__) logger = get_logger(__name__)
@ -111,7 +111,7 @@ def _dispatch_decorator(
decorator_name: str, decorator_name: str,
) -> Any: ) -> Any:
"""Shared dispatch logic for @app.tool() and @app.ui() calling patterns.""" """Shared dispatch logic for @app.tool() and @app.ui() calling patterns."""
if inspect.isroutine(name_or_fn): if is_callable_object(name_or_fn):
return register(name_or_fn, name) return register(name_or_fn, name)
if isinstance(name_or_fn, str): if isinstance(name_or_fn, str):

View file

@ -39,3 +39,13 @@ def get_fastmcp_meta(fn: Any) -> Any | None:
except ValueError: except ValueError:
pass pass
return None return None
def set_fastmcp_meta(fn: Any, metadata: Any) -> None:
"""Attach FastMCP metadata to a function, handling bound methods.
For bound methods and staticmethods, the metadata is attached to the
underlying ``__func__`` so that ``get_fastmcp_meta`` can find it.
"""
target = fn.__func__ if hasattr(fn, "__func__") else fn
target.__fastmcp__ = metadata

View file

@ -2,7 +2,6 @@
from __future__ import annotations from __future__ import annotations
import functools
import inspect import inspect
import json import json
import warnings import warnings
@ -23,7 +22,7 @@ from mcp.types import Icon
from pydantic.json_schema import SkipJsonSchema from pydantic.json_schema import SkipJsonSchema
import fastmcp import fastmcp
from fastmcp.decorators import resolve_task_config from fastmcp.decorators import resolve_task_config, set_fastmcp_meta
from fastmcp.exceptions import FastMCPDeprecationWarning, FastMCPError, PromptError from fastmcp.exceptions import FastMCPDeprecationWarning, FastMCPError, PromptError
from fastmcp.prompts.base import Prompt, PromptArgument, PromptResult from fastmcp.prompts.base import Prompt, PromptArgument, PromptResult
from fastmcp.server.auth.authorization import AuthCheck from fastmcp.server.auth.authorization import AuthCheck
@ -36,6 +35,11 @@ from fastmcp.utilities.async_utils import (
call_sync_fn_in_threadpool, call_sync_fn_in_threadpool,
is_coroutine_function, is_coroutine_function,
) )
from fastmcp.utilities.callable_utils import (
get_callable_name,
is_callable_object,
prepare_callable,
)
from fastmcp.utilities.docstring_parsing import ParsedDocstring, parse_docstring from fastmcp.utilities.docstring_parsing import ParsedDocstring, parse_docstring
from fastmcp.utilities.json_schema import compress_schema from fastmcp.utilities.json_schema import compress_schema
from fastmcp.utilities.logging import get_logger from fastmcp.utilities.logging import get_logger
@ -138,9 +142,7 @@ class FunctionPrompt(Prompt):
auth=auth, auth=auth,
) )
func_name = ( func_name = metadata.name or get_callable_name(fn)
metadata.name or getattr(fn, "__name__", None) or fn.__class__.__name__
)
if func_name == "<lambda>": if func_name == "<lambda>":
raise ValueError("You must provide a name for lambda functions") raise ValueError("You must provide a name for lambda functions")
@ -157,22 +159,10 @@ class FunctionPrompt(Prompt):
# docstring as the prompt description for callable class instances. # docstring as the prompt description for callable class instances.
outer_docstring = parse_docstring(fn) outer_docstring = parse_docstring(fn)
# Normalize task to TaskConfig and validate task_config = TaskConfig.normalize(metadata.task)
task_value = metadata.task
if task_value is None:
task_config = TaskConfig(mode="forbidden")
elif isinstance(task_value, bool):
task_config = TaskConfig.from_bool(task_value)
else:
task_config = task_value
task_config.validate_function(fn, func_name) task_config.validate_function(fn, func_name)
# if the fn is a callable class, we need to get the __call__ method from here out fn = prepare_callable(fn)
if not inspect.isroutine(fn) and not isinstance(fn, functools.partial):
fn = fn.__call__
# if the fn is a staticmethod, we need to work with the underlying function
if isinstance(fn, staticmethod):
fn = fn.__func__
# For callable classes, argument descriptions must come from # For callable classes, argument descriptions must come from
# __call__'s docstring — where the exposed parameters are actually # __call__'s docstring — where the exposed parameters are actually
@ -483,8 +473,7 @@ def prompt(
task=task, task=task,
auth=auth, auth=auth,
) )
target = fn.__func__ if hasattr(fn, "__func__") else fn set_fastmcp_meta(fn, metadata)
target.__fastmcp__ = metadata
return fn return fn
def decorator(fn: F, prompt_name: str | None) -> F: def decorator(fn: F, prompt_name: str | None) -> F:
@ -498,7 +487,7 @@ def prompt(
return create_prompt(fn, prompt_name) # type: ignore[return-value] # ty:ignore[invalid-return-type] return create_prompt(fn, prompt_name) # type: ignore[return-value] # ty:ignore[invalid-return-type]
return attach_metadata(fn, prompt_name) return attach_metadata(fn, prompt_name)
if inspect.isroutine(name_or_fn): if is_callable_object(name_or_fn):
return decorator(name_or_fn, name) return decorator(name_or_fn, name)
elif isinstance(name_or_fn, str): elif isinstance(name_or_fn, str):
if name is not None: if name is not None:

View file

@ -2,7 +2,6 @@
from __future__ import annotations from __future__ import annotations
import functools
import inspect import inspect
import warnings import warnings
from collections.abc import Callable from collections.abc import Callable
@ -14,7 +13,7 @@ from pydantic import AnyUrl
from pydantic.json_schema import SkipJsonSchema from pydantic.json_schema import SkipJsonSchema
import fastmcp import fastmcp
from fastmcp.decorators import resolve_task_config from fastmcp.decorators import resolve_task_config, set_fastmcp_meta
from fastmcp.exceptions import FastMCPDeprecationWarning from fastmcp.exceptions import FastMCPDeprecationWarning
from fastmcp.resources.base import Resource, ResourceResult from fastmcp.resources.base import Resource, ResourceResult
from fastmcp.server.auth.authorization import AuthCheck from fastmcp.server.auth.authorization import AuthCheck
@ -27,6 +26,11 @@ from fastmcp.utilities.async_utils import (
call_sync_fn_in_threadpool, call_sync_fn_in_threadpool,
is_coroutine_function, is_coroutine_function,
) )
from fastmcp.utilities.callable_utils import (
get_callable_name,
is_callable_object,
prepare_callable,
)
from fastmcp.utilities.mime import resolve_ui_mime_type from fastmcp.utilities.mime import resolve_ui_mime_type
if TYPE_CHECKING: if TYPE_CHECKING:
@ -159,27 +163,12 @@ class FunctionResource(Resource):
uri_obj = AnyUrl(metadata.uri) uri_obj = AnyUrl(metadata.uri)
# Get function name - use class name for callable objects func_name = metadata.name or get_callable_name(fn)
func_name = (
metadata.name or getattr(fn, "__name__", None) or fn.__class__.__name__
)
# Normalize task to TaskConfig and validate task_config = TaskConfig.normalize(metadata.task)
task_value = metadata.task
if task_value is None:
task_config = TaskConfig(mode="forbidden")
elif isinstance(task_value, bool):
task_config = TaskConfig.from_bool(task_value)
else:
task_config = task_value
task_config.validate_function(fn, func_name) task_config.validate_function(fn, func_name)
# if the fn is a callable class, we need to get the __call__ method from here out fn = prepare_callable(fn)
if not inspect.isroutine(fn) and not isinstance(fn, functools.partial):
fn = fn.__call__
# if the fn is a staticmethod, we need to work with the underlying function
if isinstance(fn, staticmethod):
fn = fn.__func__
# Transform Context type annotations to Depends() for unified DI # Transform Context type annotations to Depends() for unified DI
fn = transform_context_annotations(fn) fn = transform_context_annotations(fn)
@ -259,7 +248,7 @@ def resource(
if isinstance(annotations, dict): if isinstance(annotations, dict):
annotations = Annotations(**annotations) annotations = Annotations(**annotations)
if inspect.isroutine(uri): if is_callable_object(uri):
raise TypeError( raise TypeError(
"The @resource decorator requires a URI. " "The @resource decorator requires a URI. "
"Use @resource('uri') instead of @resource" "Use @resource('uri') instead of @resource"
@ -325,8 +314,7 @@ def resource(
task=task, task=task,
auth=auth, auth=auth,
) )
target = fn.__func__ if hasattr(fn, "__func__") else fn set_fastmcp_meta(fn, metadata)
target.__fastmcp__ = metadata
return fn return fn
def decorator(fn: F) -> F: def decorator(fn: F) -> F:

View file

@ -2,7 +2,6 @@
from __future__ import annotations from __future__ import annotations
import functools
import inspect import inspect
import re import re
from collections.abc import Callable from collections.abc import Callable
@ -30,6 +29,7 @@ from fastmcp.server.dependencies import (
without_injected_parameters, without_injected_parameters,
) )
from fastmcp.server.tasks.config import TaskConfig, TaskMeta from fastmcp.server.tasks.config import TaskConfig, TaskMeta
from fastmcp.utilities.callable_utils import get_callable_name, prepare_callable
from fastmcp.utilities.components import FastMCPComponent from fastmcp.utilities.components import FastMCPComponent
from fastmcp.utilities.json_schema import compress_schema from fastmcp.utilities.json_schema import compress_schema
from fastmcp.utilities.mime import resolve_ui_mime_type from fastmcp.utilities.mime import resolve_ui_mime_type
@ -539,7 +539,7 @@ class FunctionResourceTemplate(ResourceTemplate):
) -> FunctionResourceTemplate: ) -> FunctionResourceTemplate:
"""Create a template from a function.""" """Create a template from a function."""
func_name = name or getattr(fn, "__name__", None) or fn.__class__.__name__ func_name = name or get_callable_name(fn)
if func_name == "<lambda>": if func_name == "<lambda>":
raise ValueError("You must provide a name for lambda functions") raise ValueError("You must provide a name for lambda functions")
@ -625,21 +625,10 @@ class FunctionResourceTemplate(ResourceTemplate):
description = description if description is not None else inspect.getdoc(fn) description = description if description is not None else inspect.getdoc(fn)
# Normalize task to TaskConfig and validate task_config = TaskConfig.normalize(task)
if task is None:
task_config = TaskConfig(mode="forbidden")
elif isinstance(task, bool):
task_config = TaskConfig.from_bool(task)
else:
task_config = task
task_config.validate_function(fn, func_name) task_config.validate_function(fn, func_name)
# if the fn is a callable class, we need to get the __call__ method from here out fn = prepare_callable(fn)
if not inspect.isroutine(fn) and not isinstance(fn, functools.partial):
fn = fn.__call__
# if the fn is a staticmethod, we need to work with the underlying function
if isinstance(fn, staticmethod):
fn = fn.__func__
# Transform Context type annotations to Depends() for unified DI # Transform Context type annotations to Depends() for unified DI
fn = transform_context_annotations(fn) fn = transform_context_annotations(fn)

View file

@ -15,10 +15,12 @@ import mcp.types
from mcp.types import AnyFunction from mcp.types import AnyFunction
import fastmcp import fastmcp
from fastmcp.decorators import set_fastmcp_meta
from fastmcp.prompts.base import Prompt from fastmcp.prompts.base import Prompt
from fastmcp.prompts.function_prompt import FunctionPrompt from fastmcp.prompts.function_prompt import FunctionPrompt
from fastmcp.server.auth.authorization import AuthCheck from fastmcp.server.auth.authorization import AuthCheck
from fastmcp.server.tasks.config import TaskConfig from fastmcp.server.tasks.config import TaskConfig
from fastmcp.utilities.callable_utils import is_callable_object
if TYPE_CHECKING: if TYPE_CHECKING:
from fastmcp.server.providers.local_provider import LocalProvider from fastmcp.server.providers.local_provider import LocalProvider
@ -223,12 +225,11 @@ class PromptDecoratorMixin:
auth=auth, auth=auth,
enabled=enabled, enabled=enabled,
) )
target = fn.__func__ if hasattr(fn, "__func__") else fn set_fastmcp_meta(fn, metadata)
target.__fastmcp__ = metadata # type: ignore[attr-defined] # ty:ignore[unresolved-attribute]
self.add_prompt(fn) self.add_prompt(fn)
return fn return fn
if inspect.isroutine(name_or_fn): if is_callable_object(name_or_fn):
return decorate_and_register(name_or_fn, name) return decorate_and_register(name_or_fn, name)
elif isinstance(name_or_fn, str): elif isinstance(name_or_fn, str):

View file

@ -14,11 +14,13 @@ import mcp.types
from mcp.types import Annotations, AnyFunction from mcp.types import Annotations, AnyFunction
import fastmcp import fastmcp
from fastmcp.decorators import set_fastmcp_meta
from fastmcp.resources.base import Resource from fastmcp.resources.base import Resource
from fastmcp.resources.function_resource import resource as standalone_resource from fastmcp.resources.function_resource import resource as standalone_resource
from fastmcp.resources.template import ResourceTemplate from fastmcp.resources.template import ResourceTemplate
from fastmcp.server.auth.authorization import AuthCheck from fastmcp.server.auth.authorization import AuthCheck
from fastmcp.server.tasks.config import TaskConfig from fastmcp.server.tasks.config import TaskConfig
from fastmcp.utilities.callable_utils import is_callable_object
if TYPE_CHECKING: if TYPE_CHECKING:
from fastmcp.server.providers.local_provider import LocalProvider from fastmcp.server.providers.local_provider import LocalProvider
@ -159,7 +161,7 @@ class ResourceDecoratorMixin:
if isinstance(annotations, dict): if isinstance(annotations, dict):
annotations = Annotations(**annotations) annotations = Annotations(**annotations)
if inspect.isroutine(uri): if is_callable_object(uri):
raise TypeError( raise TypeError(
"The @resource decorator was used incorrectly. " "The @resource decorator was used incorrectly. "
"It requires a URI as the first argument. " "It requires a URI as the first argument. "
@ -234,8 +236,7 @@ class ResourceDecoratorMixin:
auth=auth, auth=auth,
enabled=enabled, enabled=enabled,
) )
target = fn.__func__ if hasattr(fn, "__func__") else fn set_fastmcp_meta(fn, metadata)
target.__fastmcp__ = metadata # type: ignore[attr-defined] # ty:ignore[unresolved-attribute]
self.add_resource(fn) self.add_resource(fn)
return fn return fn

View file

@ -27,11 +27,13 @@ import mcp.types
from mcp.types import AnyFunction, ToolAnnotations from mcp.types import AnyFunction, ToolAnnotations
import fastmcp import fastmcp
from fastmcp.decorators import set_fastmcp_meta
from fastmcp.exceptions import FastMCPDeprecationWarning from fastmcp.exceptions import FastMCPDeprecationWarning
from fastmcp.server.auth.authorization import AuthCheck from fastmcp.server.auth.authorization import AuthCheck
from fastmcp.server.tasks.config import TaskConfig from fastmcp.server.tasks.config import TaskConfig
from fastmcp.tools.base import Tool from fastmcp.tools.base import Tool
from fastmcp.tools.function_tool import FunctionTool from fastmcp.tools.function_tool import FunctionTool
from fastmcp.utilities.callable_utils import is_callable_object
from fastmcp.utilities.types import NotSet, NotSetT from fastmcp.utilities.types import NotSet, NotSetT
try: try:
@ -377,12 +379,11 @@ class ToolDecoratorMixin:
enabled=enabled, enabled=enabled,
run_in_thread=run_in_thread, run_in_thread=run_in_thread,
) )
target = fn.__func__ if hasattr(fn, "__func__") else fn set_fastmcp_meta(fn, metadata)
target.__fastmcp__ = metadata # type: ignore[attr-defined] # ty:ignore[unresolved-attribute]
tool_obj = self.add_tool(fn) tool_obj = self.add_tool(fn)
return fn return fn
if inspect.isroutine(name_or_fn): if is_callable_object(name_or_fn):
return decorate_and_register(name_or_fn, name) return decorate_and_register(name_or_fn, name)
elif isinstance(name_or_fn, str): elif isinstance(name_or_fn, str):

View file

@ -6,14 +6,13 @@ handle task-augmented execution as specified in SEP-1686.
from __future__ import annotations from __future__ import annotations
import functools
import inspect
from collections.abc import Callable from collections.abc import Callable
from dataclasses import dataclass from dataclasses import dataclass
from datetime import timedelta from datetime import timedelta
from typing import Any, Literal from typing import Any, Literal
from fastmcp.utilities.async_utils import is_coroutine_function from fastmcp.utilities.async_utils import is_coroutine_function
from fastmcp.utilities.callable_utils import prepare_callable
# Task execution modes per SEP-1686 / MCP ToolExecution.taskSupport # Task execution modes per SEP-1686 / MCP ToolExecution.taskSupport
TaskMode = Literal["forbidden", "optional", "required"] TaskMode = Literal["forbidden", "optional", "required"]
@ -90,6 +89,23 @@ class TaskConfig:
""" """
return cls(mode="optional" if value else "forbidden") return cls(mode="optional" if value else "forbidden")
@classmethod
def normalize(cls, task: bool | TaskConfig | None) -> TaskConfig:
"""Convert a task parameter to a TaskConfig.
Args:
task: True/False for simple enable/disable, TaskConfig for full
control, or None for the default (forbidden).
Returns:
A TaskConfig instance.
"""
if task is None:
return cls(mode="forbidden")
if isinstance(task, bool):
return cls.from_bool(task)
return task
def supports_tasks(self) -> bool: def supports_tasks(self) -> bool:
"""Check if this component supports task execution. """Check if this component supports task execution.
@ -126,15 +142,7 @@ class TaskConfig:
require_docket(f"`task=True` on function '{name}'") require_docket(f"`task=True` on function '{name}'")
# Unwrap callable classes and staticmethods # Unwrap callable classes and staticmethods
fn_to_check = fn fn_to_check = prepare_callable(fn)
if (
not inspect.isroutine(fn)
and not isinstance(fn, functools.partial)
and callable(fn)
):
fn_to_check = fn.__call__
if isinstance(fn_to_check, staticmethod):
fn_to_check = fn_to_check.__func__
if not is_coroutine_function(fn_to_check): if not is_coroutine_function(fn_to_check):
raise ValueError( raise ValueError(

View file

@ -2,7 +2,6 @@
from __future__ import annotations from __future__ import annotations
import functools
import inspect import inspect
import types import types
from collections.abc import Callable from collections.abc import Callable
@ -18,6 +17,7 @@ from fastmcp.server.dependencies import (
without_injected_parameters, without_injected_parameters,
) )
from fastmcp.tools.base import ToolResult from fastmcp.tools.base import ToolResult
from fastmcp.utilities.callable_utils import get_callable_name, prepare_callable
from fastmcp.utilities.docstring_parsing import ParsedDocstring, parse_docstring from fastmcp.utilities.docstring_parsing import ParsedDocstring, parse_docstring
from fastmcp.utilities.json_schema import compress_schema from fastmcp.utilities.json_schema import compress_schema
from fastmcp.utilities.logging import get_logger from fastmcp.utilities.logging import get_logger
@ -164,15 +164,10 @@ class ParsedFunction:
) )
# collect name and description before we potentially modify the function # collect name and description before we potentially modify the function
fn_name = getattr(fn, "__name__", None) or fn.__class__.__name__ fn_name = get_callable_name(fn)
outer_docstring = parse_docstring(fn) outer_docstring = parse_docstring(fn)
# if the fn is a callable class, we need to get the __call__ method from here out fn = prepare_callable(fn)
if not inspect.isroutine(fn) and not isinstance(fn, functools.partial):
fn = fn.__call__
# if the fn is a staticmethod, we need to work with the underlying function
if isinstance(fn, staticmethod):
fn = fn.__func__
# For callable classes, parameter descriptions must come from # For callable classes, parameter descriptions must come from
# __call__'s docstring — where the exposed parameters are actually # __call__'s docstring — where the exposed parameters are actually

View file

@ -24,7 +24,7 @@ from pydantic import Field
from pydantic.json_schema import SkipJsonSchema from pydantic.json_schema import SkipJsonSchema
import fastmcp import fastmcp
from fastmcp.decorators import get_fastmcp_meta, resolve_task_config from fastmcp.decorators import get_fastmcp_meta, resolve_task_config, set_fastmcp_meta
from fastmcp.exceptions import FastMCPDeprecationWarning from fastmcp.exceptions import FastMCPDeprecationWarning
from fastmcp.server.auth.authorization import AuthCheck from fastmcp.server.auth.authorization import AuthCheck
from fastmcp.server.dependencies import without_injected_parameters from fastmcp.server.dependencies import without_injected_parameters
@ -39,6 +39,7 @@ from fastmcp.utilities.async_utils import (
call_sync_fn_in_threadpool, call_sync_fn_in_threadpool,
is_coroutine_function, is_coroutine_function,
) )
from fastmcp.utilities.callable_utils import is_callable_object
from fastmcp.utilities.logging import get_logger from fastmcp.utilities.logging import get_logger
from fastmcp.utilities.types import ( from fastmcp.utilities.types import (
NotSet, NotSet,
@ -237,14 +238,7 @@ class FunctionTool(Tool):
"accept worker-thread dispatch." "accept worker-thread dispatch."
) )
# Normalize task to TaskConfig task_config = TaskConfig.normalize(metadata.task)
task_value = metadata.task
if task_value is None:
task_config = TaskConfig(mode="forbidden")
elif isinstance(task_value, bool):
task_config = TaskConfig.from_bool(task_value)
else:
task_config = task_value
task_config.validate_function(fn, func_name) task_config.validate_function(fn, func_name)
# Handle output_schema # Handle output_schema
@ -516,8 +510,7 @@ def tool(
auth=auth, auth=auth,
run_in_thread=run_in_thread, run_in_thread=run_in_thread,
) )
target = fn.__func__ if hasattr(fn, "__func__") else fn set_fastmcp_meta(fn, metadata)
target.__fastmcp__ = metadata
return fn return fn
def decorator(fn: F, tool_name: str | None) -> F: def decorator(fn: F, tool_name: str | None) -> F:
@ -531,7 +524,7 @@ def tool(
return create_tool(fn, tool_name) # type: ignore[return-value] # ty:ignore[invalid-return-type] return create_tool(fn, tool_name) # type: ignore[return-value] # ty:ignore[invalid-return-type]
return attach_metadata(fn, tool_name) return attach_metadata(fn, tool_name)
if inspect.isroutine(name_or_fn): if is_callable_object(name_or_fn):
return decorator(name_or_fn, name) return decorator(name_or_fn, name)
elif isinstance(name_or_fn, str): elif isinstance(name_or_fn, str):
if name is not None: if name is not None:

View file

@ -0,0 +1,92 @@
"""Utilities for handling callables, including functools.partial objects.
Provides centralized helpers for the shared steps in the tool/prompt/resource
``from_function`` pipelines, avoiding duplicated ``isinstance`` checks, name
extraction logic, and callable unwrapping across the codebase.
"""
from __future__ import annotations
import contextlib
import functools
import inspect
from collections.abc import Callable
from typing import Any, TypeGuard
def is_callable_object(obj: Any) -> TypeGuard[Callable[..., Any]]:
"""Check if an object is a callable suitable for use as a tool, resource, or prompt.
Returns True for functions, methods, builtins, and functools.partial objects.
This is a broader check than ``inspect.isroutine`` which returns False for
functools.partial.
"""
return inspect.isroutine(obj) or isinstance(obj, functools.partial)
def get_callable_name(fn: Any) -> str:
"""Extract a human-readable name from a callable.
Handles functions, callable classes, and functools.partial:
- Regular functions: returns ``fn.__name__`` (e.g. ``"add"``)
- Callable classes: returns the class name (e.g. ``"MyTool"``)
- Partial with ``update_wrapper``: returns the wrapped name (e.g. ``"add"``)
- Partial without ``update_wrapper``: returns the underlying function name
(e.g. ``"add"`` instead of ``"partial"``)
"""
name = getattr(fn, "__name__", None)
if name is not None:
return name
# functools.partial without update_wrapper — use the underlying function's name
if isinstance(fn, functools.partial):
return getattr(fn.func, "__name__", None) or fn.__class__.__name__
return fn.__class__.__name__
def prepare_callable(fn: Callable[..., Any]) -> Callable[..., Any]:
"""Prepare a callable for introspection by ``inspect.signature()`` and Pydantic.
This handles three cases that would otherwise require special-casing in every
``from_function`` method:
1. **functools.partial with __wrapped__**: ``functools.update_wrapper`` sets
``__wrapped__`` which causes ``inspect.signature()`` and Pydantic to follow
it back to the original function, ignoring the partial's bound arguments.
We strip ``__wrapped__`` by reconstructing the partial.
2. **Callable classes**: Non-routine callables (classes with ``__call__``) need
to be unwrapped to their ``__call__`` method so ``inspect.signature()`` sees
the right parameters.
3. **staticmethod**: Needs unwrapping to the underlying function.
Call this AFTER extracting name/doc from the original callable, since this
may change what ``__name__`` and ``__doc__`` return.
"""
# Strip __wrapped__ from partials so Pydantic sees the partial's own
# signature with bound args removed, not the original function's signature.
if isinstance(fn, functools.partial) and hasattr(fn, "__wrapped__"):
old = fn
fn = functools.partial(fn.func, *fn.args, **fn.keywords)
# Preserve annotation metadata copied by functools.update_wrapper
# (e.g. __module__, __qualname__) needed for deferred annotation
# resolution when `from __future__ import annotations` is active.
for attr in functools.WRAPPER_ASSIGNMENTS:
try:
val = getattr(old, attr)
except AttributeError:
pass
else:
with contextlib.suppress(AttributeError):
setattr(fn, attr, val)
# Callable classes (not routines, not partials) → unwrap to __call__
if not inspect.isroutine(fn) and not isinstance(fn, functools.partial):
fn = fn.__call__
# staticmethod → unwrap to underlying function
if isinstance(fn, staticmethod):
fn = fn.__func__
return fn

View file

@ -0,0 +1,129 @@
"""Tests for functools.partial support as tools, prompts, and resources.
See https://github.com/PrefectHQ/fastmcp/issues/3266
"""
import functools
from mcp.types import TextContent
from fastmcp import Client, FastMCP
from fastmcp.tools.function_tool import FunctionTool as Tool
class TestPartialTool:
"""Test tools created from functools.partial objects."""
async def test_partial_sync(self):
def add(x: int, y: int) -> int:
return x + y
partial_add = functools.partial(add, y=10)
functools.update_wrapper(partial_add, add)
tool = Tool.from_function(partial_add)
result = await tool.run({"x": 5})
assert result.content == [TextContent(type="text", text="15")]
async def test_partial_async(self):
async def multiply(x: int, factor: int) -> int:
return x * factor
partial_mul = functools.partial(multiply, factor=3)
functools.update_wrapper(partial_mul, multiply)
tool = Tool.from_function(partial_mul)
result = await tool.run({"x": 7})
assert result.content == [TextContent(type="text", text="21")]
async def test_partial_preserves_name(self):
def greet(name: str, greeting: str = "Hello") -> str:
"""Greet someone."""
return f"{greeting}, {name}!"
partial_greet = functools.partial(greet, greeting="Hi")
functools.update_wrapper(partial_greet, greet)
tool = Tool.from_function(partial_greet)
assert tool.name == "greet"
assert tool.description == "Greet someone."
async def test_partial_without_update_wrapper(self):
def add(x: int, y: int) -> int:
return x + y
partial_add = functools.partial(add, y=10)
tool = Tool.from_function(partial_add, name="add_ten")
result = await tool.run({"x": 5})
assert result.content == [TextContent(type="text", text="15")]
async def test_partial_with_add_tool(self):
mcp = FastMCP("test")
def greet(name: str, greeting: str = "Hello") -> str:
return f"{greeting}, {name}!"
partial_greet = functools.partial(greet, greeting="Hey")
functools.update_wrapper(partial_greet, greet)
mcp.add_tool(partial_greet)
result = await mcp.call_tool("greet", {"name": "World"})
assert result.content == [TextContent(type="text", text="Hey, World!")]
async def test_partial_with_server_tool_decorator(self):
mcp = FastMCP("test")
def add(x: int, y: int) -> int:
return x + y
partial_add = functools.partial(add, y=100)
functools.update_wrapper(partial_add, add)
mcp.tool(partial_add)
result = await mcp.call_tool("add", {"x": 5})
assert result.content == [TextContent(type="text", text="105")]
class TestPartialPrompt:
"""Test prompts created from functools.partial objects."""
async def test_partial_prompt_with_decorator(self):
"""Partial can be registered via @mcp.prompt() decorator."""
mcp = FastMCP("test")
def greet_prompt(name: str, lang: str) -> str:
return f"Say hello to {name} in {lang}."
partial_greet = functools.partial(greet_prompt, lang="French")
functools.update_wrapper(partial_greet, greet_prompt)
mcp.prompt(partial_greet)
async with Client(mcp) as client:
result = await client.get_prompt("greet_prompt", {"name": "Alice"})
assert "Alice" in str(result.messages[0])
assert "French" in str(result.messages[0])
class TestPartialResource:
"""Test resources created from functools.partial objects."""
async def test_partial_resource_with_decorator(self):
"""Partial can be registered via @mcp.resource() decorator."""
mcp = FastMCP("test")
def get_data(key: str, fmt: str = "text") -> str:
return f"{key} in {fmt} format"
partial_data = functools.partial(get_data, fmt="json")
functools.update_wrapper(partial_data, get_data)
mcp.resource("data://{key}")(partial_data)
async with Client(mcp) as client:
content = await client.read_resource("data://users")
assert "users" in str(content)
assert "json" in str(content)

View file

@ -0,0 +1,121 @@
"""Tests for callable utility functions."""
import functools
from fastmcp.utilities.callable_utils import (
get_callable_name,
is_callable_object,
prepare_callable,
)
class TestIsCallableObject:
def test_function(self):
def fn():
pass
assert is_callable_object(fn) is True
def test_async_function(self):
async def fn():
pass
assert is_callable_object(fn) is True
def test_partial(self):
def fn(x, y):
return x + y
assert is_callable_object(functools.partial(fn, y=1)) is True
def test_callable_class(self):
class MyCallable:
def __call__(self):
pass
assert is_callable_object(MyCallable()) is False
def test_string(self):
assert is_callable_object("not a callable") is False
def test_none(self):
assert is_callable_object(None) is False
class TestGetCallableName:
def test_function(self):
def my_function():
pass
assert get_callable_name(my_function) == "my_function"
def test_lambda(self):
assert get_callable_name(lambda: None) == "<lambda>"
def test_partial_with_update_wrapper(self):
def add(x, y):
return x + y
p = functools.partial(add, y=10)
functools.update_wrapper(p, add)
assert get_callable_name(p) == "add"
def test_partial_without_update_wrapper(self):
def add(x, y):
return x + y
p = functools.partial(add, y=10)
assert get_callable_name(p) == "add"
def test_callable_class(self):
class MyTool:
def __call__(self):
pass
assert get_callable_name(MyTool()) == "MyTool"
class TestPrepareCallable:
def test_regular_function_unchanged(self):
def fn(x):
return x
assert prepare_callable(fn) is fn
def test_strips_wrapped_from_partial(self):
def add(x, y):
return x + y
p = functools.partial(add, y=10)
functools.update_wrapper(p, add)
assert hasattr(p, "__wrapped__")
prepared = prepare_callable(p)
assert isinstance(prepared, functools.partial)
assert not hasattr(prepared, "__wrapped__")
assert prepared.keywords == {"y": 10}
def test_partial_without_wrapper_unchanged(self):
def add(x, y):
return x + y
p = functools.partial(add, y=10)
prepared = prepare_callable(p)
assert isinstance(prepared, functools.partial)
assert prepared.func is add
def test_callable_class_unwrapped(self):
class MyCallable:
def __call__(self, x):
return x
obj = MyCallable()
prepared = prepare_callable(obj)
assert prepared == obj.__call__
def test_staticmethod_unwrapped(self):
def fn(x):
return x
sm = staticmethod(fn)
assert prepare_callable(sm) is fn