diff --git a/libs/core/langchain_core/tools/base.py b/libs/core/langchain_core/tools/base.py index a1cd5502f1..69020241cd 100644 --- a/libs/core/langchain_core/tools/base.py +++ b/libs/core/langchain_core/tools/base.py @@ -20,7 +20,6 @@ from typing import ( cast, get_args, get_origin, - get_type_hints, ) import typing_extensions @@ -38,7 +37,7 @@ from pydantic.fields import FieldInfo from pydantic.v1 import BaseModel as BaseModelV1 from pydantic.v1 import ValidationError as ValidationErrorV1 from pydantic.v1 import validate_arguments as validate_arguments_v1 -from typing_extensions import Self, override +from typing_extensions import Self, get_type_hints, override from langchain_core.callbacks import ( AsyncCallbackManager, @@ -184,7 +183,7 @@ def _infer_arg_descriptions( Returns: A tuple containing the function description and argument descriptions. """ - annotations = typing.get_type_hints(fn, include_extras=True) + annotations = get_type_hints(fn, include_extras=True) if parse_docstring: description, arg_descriptions = _parse_python_function_docstring( fn, annotations, error_on_invalid_docstring=error_on_invalid_docstring @@ -730,20 +729,8 @@ class ChildTool(BaseTool): # rather than by the model, so they must be excluded from the schema and # re-injected during execution. `StructuredTool` overrides this to # inspect its wrapped `func`/`coroutine` instead. - # - # Resolve annotations via `get_type_hints` (rather than reading raw - # `signature` annotations) so postponed annotations -- e.g. from - # `from __future__ import annotations` or quoted forward references -- - # are recognized. Fall back to the raw annotation per-parameter when a - # hint can't be resolved. for method in (self._run, self._arun): - params = signature(method).parameters - hints = _get_type_hints(method, include_extras=True) or {} - keys = frozenset( - name - for name, param in params.items() - if _is_injected_arg_type(hints.get(name, param.annotation)) - ) + keys = _get_injected_args_keys_from_signature(method) if keys: return keys return _EMPTY_SET @@ -1496,6 +1483,11 @@ def _stringify(content: Any) -> str: return str(content) +def _describe_callable(func: Any) -> str: + """Return a readable identifier for a callable, for use in log messages.""" + return getattr(func, "__qualname__", None) or repr(func) + + def _get_type_hints( func: Callable[..., Any], *, include_extras: bool = False ) -> dict[str, type] | None: @@ -1514,6 +1506,202 @@ def _get_type_hints( try: return get_type_hints(func, include_extras=include_extras) except Exception: + _logger.debug( + "Failed to resolve type hints for %s.", + _describe_callable(func), + exc_info=True, + ) + return None + + +def _get_class_signature_source(cls: type) -> Callable[..., Any] | None: + """Return the callable whose parameters `signature(cls)` describes. + + A class's own annotations describe its attributes, not its constructor + parameters, so a class is never its own annotation owner. Constructor + precedence mirrors `inspect.signature`: a metaclass `__call__` override, + then a `__new__`/`__init__` the class defines itself, then an inherited one. + + Args: + cls: The class being inspected. + + Returns: + The constructor whose annotations describe `signature(cls)`, or `None` + when only slot wrappers inherited from `object` are available. + """ + metaclass = type(cls) + if metaclass is not type: + metaclass_call = inspect.getattr_static(metaclass, "__call__", None) + if inspect.isfunction(metaclass_call): + return metaclass_call + for attr in ("__new__", "__init__"): + if attr in cls.__dict__: + return cast("Callable[..., Any]", getattr(cls, attr)) + for attr in ("__new__", "__init__"): + # `object.__new__`/`object.__init__` are slot wrappers carrying no + # useful annotations, so only inherited Python constructors qualify. + inherited = getattr(cls, attr, None) + if inspect.isfunction(inherited): + return cast("Callable[..., Any]", inherited) + return None + + +def _get_type_hints_source( + func: Callable[..., Any], +) -> Callable[..., Any] | None: + """Return the callable that owns the annotations for an effective signature. + + Unwrapping stops as soon as a callable declares its own `__signature__` (or + is a bound method), because that declaration -- not the annotations of + whatever it wraps -- defines the effective signature. + + Args: + func: The callable being inspected. + + Returns: + The function or method whose annotations describe `signature(func)`, or + `None` when an explicit `__signature__` supplies the annotations or no + annotation owner can be identified. + """ + func = inspect.unwrap( + func, + stop=lambda wrapped: ( + hasattr(wrapped, "__signature__") or inspect.ismethod(wrapped) + ), + ) + if inspect.ismethod(func): + return _get_type_hints_source(func.__func__) + if getattr(func, "__signature__", None) is not None: + return None + if isinstance(func, functools.partial): + return _get_type_hints_source(func.func) + if inspect.isclass(func): + constructor = _get_class_signature_source(func) + if constructor is None: + return None + return _get_type_hints_source(constructor) + if not inspect.isroutine(func): + callable_obj = cast("Any", func) + return _get_type_hints_source(cast("Callable[..., Any]", callable_obj.__call__)) + return func + + +def _get_callable_globals(func: Callable[..., Any]) -> dict[str, Any]: + """Return the globals namespace associated with a callable. + + Explicit `__signature__` objects store annotations separately from a + callable's `__annotations__`. String annotations on such signatures + still need the callable's defining globals to be resolved. + + Args: + func: The callable whose defining globals to locate. + + Returns: + The callable's globals namespace, or an empty `dict` when none can be + located (e.g. classes, builtins, and other C-implemented callables), in + which case no string annotation will resolve. + """ + if isinstance(func, functools.partial): + return _get_callable_globals(func.func) + if inspect.ismethod(func): + return _get_callable_globals(func.__func__) + globalns = getattr(func, "__globals__", None) + if isinstance(globalns, dict): + return globalns + if not inspect.isroutine(func) and not inspect.isclass(func): + callable_obj = cast("Any", func) + return _get_callable_globals(cast("Callable[..., Any]", callable_obj.__call__)) + return {} + + +def _get_injected_args_keys_from_signature(func: Callable[..., Any]) -> frozenset[str]: + """Identify injected-argument parameters of a callable. + + Resolve annotations with `get_type_hints` (rather than reading raw + `signature` annotations) so postponed annotations -- e.g. from + `from __future__ import annotations` or quoted forward references -- are + recognized. `get_type_hints` resolves every annotation at once and raises + on the first failure, so any parameter it leaves uncovered is retried on its + own: an unrelated unresolvable forward reference must not disable injection + for parameters whose annotations do resolve. Each distinct string annotation + is resolved at most once per call, failures included. + + A *string* annotation naming a type that exists only in the callable's + defining local scope cannot be resolved here, since only globals are + available. Such a parameter keeps its raw string annotation, which + `_is_injected_arg_type` never classifies as injected, so it stays visible to + the model rather than being injected at call time. Non-string annotations + are unaffected, since `signature` yields the live object. + + Args: + func: The function (or bound method) whose signature to inspect. + + Returns: + `frozenset` of parameter names annotated as injected arguments. + """ + params = signature(func).parameters + hint_source = _get_type_hints_source(func) + hints = ( + _get_type_hints(hint_source, include_extras=True) + if hint_source is not None + else None + ) or {} + globalns: dict[str, Any] | None = None + resolved_annotations: dict[str, Any] = {} + keys = set() + for name, param in params.items(): + annotation = hints.get(name, param.annotation) + if isinstance(annotation, str): + if globalns is None: + globalns = _get_callable_globals( + func if hint_source is None else hint_source + ) + if annotation not in resolved_annotations: + resolved_annotations[annotation] = _resolve_forward_ref( + annotation, globalns + ) + resolved = resolved_annotations[annotation] + if resolved is None: + _logger.debug( + "Could not resolve annotation %r for parameter %r of %s; it " + "will not be treated as an injected argument.", + annotation, + name, + _describe_callable(func), + ) + else: + annotation = resolved + if _is_injected_arg_type(annotation): + keys.add(name) + return frozenset(keys) + + +def _resolve_forward_ref(annotation: str, globalns: dict[str, Any]) -> Any: + """Resolve a single string annotation, returning `None` on failure. + + Uses a temporary annotated function so each annotation can be passed through + the public `get_type_hints` API independently. + + Args: + annotation: The raw string annotation to resolve. + globalns: The globals namespace to resolve names against. + + Returns: + The resolved type, or `None` if the annotation cannot be resolved. + """ + + def _annotation_holder() -> None: + pass + + _annotation_holder.__annotations__ = {"value": annotation} + try: + return get_type_hints( + _annotation_holder, + globalns=globalns, + include_extras=True, + )["value"] + except Exception: + _logger.debug("Failed to resolve annotation %r.", annotation, exc_info=True) return None diff --git a/libs/core/langchain_core/tools/structured.py b/libs/core/langchain_core/tools/structured.py index e9643a0ad8..57892a4464 100644 --- a/libs/core/langchain_core/tools/structured.py +++ b/libs/core/langchain_core/tools/structured.py @@ -27,8 +27,8 @@ from langchain_core.tools.base import ( FILTERED_ARGS, ArgsSchema, BaseTool, + _get_injected_args_keys_from_signature, _get_runnable_config_param, - _is_injected_arg_type, create_schema_from_function, ) from langchain_core.utils.pydantic import is_basemodel_subclass @@ -256,11 +256,7 @@ class StructuredTool(BaseTool): fn = self.func or self.coroutine if fn is None: return _EMPTY_SET - return frozenset( - k - for k, v in signature(fn).parameters.items() - if _is_injected_arg_type(v.annotation) - ) + return _get_injected_args_keys_from_signature(fn) def _filter_schema_args(func: Callable[..., Any]) -> list[str]: diff --git a/libs/core/tests/unit_tests/test_tools.py b/libs/core/tests/unit_tests/test_tools.py index 9b16fa8b5f..63fbfa4368 100644 --- a/libs/core/tests/unit_tests/test_tools.py +++ b/libs/core/tests/unit_tests/test_tools.py @@ -7,6 +7,7 @@ import pickle import sys import textwrap import threading +import warnings from collections.abc import Callable from dataclasses import dataclass from datetime import datetime @@ -2895,6 +2896,353 @@ def test_tool_injected_arg_with_custom_schema() -> None: assert captured["context"].value == "test_context" +@dataclass +class _PostponedRuntime(_DirectlyInjectedToolArg): + """Custom directly injected runtime used to exercise postponed annotations.""" + + some_obj: object + + +@pytest.mark.parametrize("schema_format", ["model", "json_schema"]) +def test_tool_allows_postponed_runtime_annotation_with_custom_schema( + schema_format: Literal["model", "json_schema"], +) -> None: + """Ensure postponed injected annotations are preserved with custom args_schema. + + Regression test: `StructuredTool._injected_args_keys` previously read raw + `signature` annotations, which are strings under `from __future__ import + annotations`. A quoted forward reference reproduces the same string-annotation + behavior without a module-wide `__future__` import. + """ + + class InputSchema(BaseModel): + query: str + + captured: dict[str, Any] = {} + + args_schema = ( + InputSchema if schema_format == "model" else InputSchema.model_json_schema() + ) + + @tool(args_schema=args_schema) + def runtime_tool(query: str, runtime: "_PostponedRuntime") -> str: + """Echo the query and capture runtime value.""" + captured["runtime"] = runtime + return query + + runtime_obj = object() + runtime = _PostponedRuntime(some_obj=runtime_obj) + + # The injected arg is detected from the postponed annotation... + assert "runtime" in runtime_tool._injected_args_keys + # ...survives input validation and reaches the function... + assert runtime_tool.invoke({"query": "hello", "runtime": runtime}) == "hello" + assert captured["runtime"] is runtime + # ...and does not leak into the model-facing tool-call schema. + tool_call_schema = runtime_tool.tool_call_schema + if isinstance(tool_call_schema, dict): + properties = tool_call_schema["properties"] + else: + properties = model_json_schema(tool_call_schema)["properties"] + assert "runtime" not in properties + + +def test_tool_partial_does_not_restore_bound_injected_arg() -> None: + """Injected args already bound by a partial are absent from its signature.""" + + class InputSchema(BaseModel): + y: int + + bound_runtime = _PostponedRuntime(some_obj=object()) + + def fn(x: int, runtime: _PostponedRuntime, y: int) -> int: + return x + y + + tool_ = StructuredTool.from_function( + func=partial(fn, 1, bound_runtime), + name="fn", + description="Add two numbers.", + args_schema=InputSchema, + ) + + assert "runtime" not in tool_._injected_args_keys + assert tool_.invoke({"y": 2, "runtime": object()}) == 3 + + +def test_tool_custom_schema_unresolvable_forward_ref_does_not_raise() -> None: + """Unresolved forward references must not break `_injected_args_keys`. + + When a postponed annotation cannot be resolved (e.g. the name is not + defined), `get_type_hints` raises. `_injected_args_keys` must fall back to + the raw (string) annotation instead of propagating the error. + """ + + class InputSchema(BaseModel): + query: str + + def fn( + query: str, + runtime: "NotDefinedAnywhere123", # type: ignore[name-defined] # noqa: F821 + ) -> str: + """Fn with unresolvable forward ref.""" + return query + + tool_ = StructuredTool.from_function( + func=fn, + name="fn", + description="desc", + args_schema=InputSchema, + ) + + # The unresolved annotation is not recognized as injected (safe fallback), + # and accessing `_injected_args_keys` must not raise. + assert "runtime" not in tool_._injected_args_keys + + +def test_tool_unresolvable_sibling_annotation_does_not_disable_injection() -> None: + """An unresolvable annotation on one parameter must not hide injected args. + + `get_type_hints` resolves all annotations at once, so a single unresolvable + forward reference (e.g. on `query`) would otherwise discard the resolvable + hints too -- including the injected `runtime` -- leaving + `_injected_args_keys` empty and dropping the injected value during custom + `args_schema` validation. + """ + + class InputSchema(BaseModel): + query: str + + captured: dict[str, Any] = {} + + @tool(args_schema=InputSchema) + def runtime_tool( + query: "NotDefinedAnywhere123", # type: ignore[name-defined] # noqa: F821 + runtime: "_PostponedRuntime", + ) -> str: + """Echo the query and capture runtime value.""" + captured["runtime"] = runtime + return query # type: ignore[no-any-return] + + runtime = _PostponedRuntime(some_obj=object()) + + # The resolvable injected annotation is still detected without relying on + # deprecated typing internals. + with warnings.catch_warnings(): + warnings.simplefilter("error", DeprecationWarning) + injected_args_keys = runtime_tool._injected_args_keys + assert "runtime" in injected_args_keys + # ...while the unresolvable one is not misclassified. + assert "query" not in injected_args_keys + + assert runtime_tool.invoke({"query": "hello", "runtime": runtime}) == "hello" + assert captured["runtime"] is runtime + + +def test_tool_fallback_preserves_annotated_injected_arg() -> None: + """Per-parameter fallback preserves `Annotated` injection metadata.""" + + class InputSchema(BaseModel): + query: str + + captured: dict[str, Any] = {} + + @tool(args_schema=InputSchema) + def runtime_tool( + query: "NotDefinedAnywhere123", # type: ignore[name-defined] # noqa: F821 + runtime: "Annotated[str, InjectedToolArg]", + ) -> str: + """Echo the query and capture the injected value.""" + captured["runtime"] = runtime + return query # type: ignore[no-any-return] + + assert runtime_tool._injected_args_keys == frozenset({"runtime"}) + assert runtime_tool.invoke({"query": "hello", "runtime": "injected"}) == "hello" + assert captured["runtime"] == "injected" + + +def test_tool_partial_fallback_uses_wrapped_function_namespace() -> None: + """Partial fallback resolves hints in the wrapped function's namespace.""" + + class InputSchema(BaseModel): + query: str + + captured: dict[str, Any] = {} + + def runtime_fn( + bound: int, + query: "NotDefinedAnywhere123", # type: ignore[name-defined] # noqa: F821 + runtime: "_PostponedRuntime", + ) -> str: + captured["bound"] = bound + captured["runtime"] = runtime + return query # type: ignore[no-any-return] + + tool_ = StructuredTool.from_function( + func=partial(runtime_fn, 1), + name="runtime_tool", + description="Echo a query.", + args_schema=InputSchema, + ) + runtime = _PostponedRuntime(some_obj=object()) + + assert tool_._injected_args_keys == frozenset({"runtime"}) + assert tool_.invoke({"query": "hello", "runtime": runtime}) == "hello" + assert captured == {"bound": 1, "runtime": runtime} + + +def test_base_tool_unresolvable_sibling_annotation_does_not_disable_injection() -> None: + """`BaseTool` subclasses get the same per-parameter hint resolution.""" + + class MultiplyInput(BaseModel): + a: int + + captured: dict[str, Any] = {} + + class Multiplier(BaseTool): + name: str = "Multiplier" + description: str = "Multiply." + args_schema: type[BaseModel] = MultiplyInput + + def _run( + self, + a: "NotDefinedAnywhere123", # type: ignore[name-defined] # noqa: F821 + runtime: "_CustomRuntime", + ) -> int: + captured["runtime"] = runtime + return a * 2 # type: ignore[no-any-return] + + tool_ = Multiplier() + assert "runtime" in tool_._injected_args_keys + assert "a" not in tool_._injected_args_keys + + runtime = _CustomRuntime(data={"scale": 10}) + assert tool_.invoke({"a": 2, "runtime": runtime}) == 4 + assert captured["runtime"] is runtime + + +def test_tool_class_callable_uses_constructor_annotations() -> None: + """A class used as `func` is inspected via its constructor, not its attributes. + + Regression test: a class's own annotations describe attributes, so a + class-level annotation sharing a name with a constructor parameter must not + shadow that parameter's annotation. + """ + + class InputSchema(BaseModel): + query: str + + class SearchTool: + runtime: str = "class attribute sharing a name with a constructor param" + + def __init__(self, query: str, runtime: _PostponedRuntime) -> None: + self.query = query + self.runtime = runtime # type: ignore[assignment] + + tool_ = StructuredTool.from_function( + func=SearchTool, + name="search_tool", + description="Search.", + args_schema=InputSchema, + ) + runtime = _PostponedRuntime(some_obj=object()) + + assert tool_._injected_args_keys == frozenset({"runtime"}) + result = tool_.invoke({"query": "hello", "runtime": runtime}) + assert result.query == "hello" + assert result.runtime is runtime + + +def test_tool_class_callable_attribute_annotation_is_not_injected() -> None: + """A class attribute annotation must not mark a constructor param injected. + + Otherwise a model-facing argument is treated as injected, and the raw input + value replaces the one validated against `args_schema`. + """ + + class InputSchema(BaseModel): + query: str + count: int + + class CountTool: + count: _PostponedRuntime = None # type: ignore[assignment] + + def __init__(self, query: str, count: int) -> None: + self.count = count # type: ignore[assignment] + + tool_ = StructuredTool.from_function( + func=CountTool, + name="count_tool", + description="Count.", + args_schema=InputSchema, + ) + + assert tool_._injected_args_keys == frozenset() + tool_call_schema = tool_.tool_call_schema + assert not isinstance(tool_call_schema, dict) + assert "count" in model_json_schema(tool_call_schema)["properties"] + # The schema-validated value reaches the constructor, not the raw input. + result = tool_.invoke({"query": "hello", "count": "5"}) + assert result.count == 5 + assert isinstance(result.count, int) + + +def test_tool_class_callable_uses_new_annotations() -> None: + """Classes that customize allocation are inspected via `__new__`.""" + + class InputSchema(BaseModel): + query: str + + captured: dict[str, Any] = {} + + class NewTool: + # Quoted so resolution must go through `__new__`'s own globals. + def __new__(cls, query: str, runtime: "_PostponedRuntime") -> str: # type: ignore[misc] + captured["runtime"] = runtime + return query + + tool_ = StructuredTool.from_function( + func=NewTool, + name="new_tool", + description="Echo a query.", + args_schema=InputSchema, + ) + runtime = _PostponedRuntime(some_obj=object()) + + assert tool_._injected_args_keys == frozenset({"runtime"}) + assert tool_.invoke({"query": "hello", "runtime": runtime}) == "hello" + assert captured["runtime"] is runtime + + +def test_tool_class_callable_uses_metaclass_call_annotations() -> None: + """A metaclass `__call__` override supplies the effective signature.""" + + class InputSchema(BaseModel): + query: str + + captured: dict[str, Any] = {} + + class _Meta(type): + def __call__(cls, query: str, runtime: _PostponedRuntime) -> str: + captured["runtime"] = runtime + return query + + class MetaTool(metaclass=_Meta): + def __init__(self, runtime: int) -> None: + """Constructor parameters are shadowed by the metaclass `__call__`.""" + + tool_ = StructuredTool.from_function( + func=MetaTool, + name="meta_tool", + description="Echo a query.", + args_schema=InputSchema, + ) + runtime = _PostponedRuntime(some_obj=object()) + + assert tool_._injected_args_keys == frozenset({"runtime"}) + assert tool_.invoke({"query": "hello", "runtime": runtime}) == "hello" + assert captured["runtime"] is runtime + + def test_tool_injected_tool_call_id() -> None: @tool def foo(x: int, tool_call_id: Annotated[str, InjectedToolCallId]) -> ToolMessage: diff --git a/libs/core/tests/unit_tests/test_tools_postponed_annotations.py b/libs/core/tests/unit_tests/test_tools_postponed_annotations.py new file mode 100644 index 0000000000..defce1678c --- /dev/null +++ b/libs/core/tests/unit_tests/test_tools_postponed_annotations.py @@ -0,0 +1,178 @@ +"""Tests for tools defined in modules with postponed annotations.""" + +from __future__ import annotations + +from functools import update_wrapper +from inspect import Parameter, Signature +from typing import Annotated, Any + +from pydantic import BaseModel + +from langchain_core.tools import InjectedToolArg, StructuredTool +from langchain_core.tools.base import _DirectlyInjectedToolArg +from langchain_core.utils.pydantic import model_json_schema + + +class _PostponedRuntime(_DirectlyInjectedToolArg): + """Runtime type whose annotation is postponed by the module future import.""" + + +class _InputSchema(BaseModel): + query: str + + +class _CallableTool: + """Callable-object tool with annotations stored on `__call__`.""" + + received_runtime: _PostponedRuntime | None = None + + def __call__(self, query: str, runtime: _PostponedRuntime) -> str: + self.received_runtime = runtime + return query + + +def test_callable_object_resolves_postponed_injected_arg() -> None: + """Callable objects resolve annotations from their `__call__` method.""" + callable_tool = _CallableTool() + tool = StructuredTool.from_function( + func=callable_tool, + name="callable_tool", + description="Echo a query.", + args_schema=_InputSchema, + ) + runtime = _PostponedRuntime() + + assert tool._injected_args_keys == frozenset({"runtime"}) + assert tool.invoke({"query": "hello", "runtime": runtime}) == "hello" + assert callable_tool.received_runtime is runtime + + tool_call_schema = tool.tool_call_schema + assert not isinstance(tool_call_schema, dict) + assert "runtime" not in model_json_schema(tool_call_schema)["properties"] + + +def test_callable_wrapper_uses_wrapped_function_annotations() -> None: + """Callable wrappers resolve hints from their effective signature source.""" + received: dict[str, Any] = {} + + def wrapped(query: str, runtime: _PostponedRuntime) -> str: + received["runtime"] = runtime + return query + + class CallableWrapper: + def __init__(self) -> None: + update_wrapper(self, wrapped) + + def __call__(self, query: str, runtime: str) -> str: + return wrapped(query, runtime) # type: ignore[arg-type] + + callable_wrapper = CallableWrapper() + tool = StructuredTool.from_function( + func=callable_wrapper, + name="callable_wrapper", + description="Echo a query.", + args_schema=_InputSchema, + ) + runtime = _PostponedRuntime() + + assert tool._injected_args_keys == frozenset({"runtime"}) + assert tool.invoke({"query": "hello", "runtime": runtime}) == "hello" + assert received["runtime"] is runtime + + +def test_callable_object_uses_explicit_signature_annotations() -> None: + """An explicit signature takes precedence over `__call__` annotations.""" + received: dict[str, Any] = {} + + class CallableWithSignature: + __signature__ = Signature( + parameters=[ + Parameter("query", Parameter.POSITIONAL_OR_KEYWORD, annotation=str), + Parameter( + "runtime", + Parameter.POSITIONAL_OR_KEYWORD, + annotation=_PostponedRuntime, + ), + ] + ) + + def __call__(self, query: str, runtime: str) -> str: + received["runtime"] = runtime + return query + + callable_with_signature = CallableWithSignature() + tool = StructuredTool.from_function( + func=callable_with_signature, + name="callable_with_signature", + description="Echo a query.", + args_schema=_InputSchema, + ) + runtime = _PostponedRuntime() + + assert tool._injected_args_keys == frozenset({"runtime"}) + assert tool.invoke({"query": "hello", "runtime": runtime}) == "hello" + assert received["runtime"] is runtime + + +def test_explicit_signature_resolves_string_injected_annotation() -> None: + """String annotations on explicit signatures use callable globals.""" + received: dict[str, Any] = {} + + class CallableWithStringSignature: + __signature__ = Signature( + parameters=[ + Parameter("query", Parameter.POSITIONAL_OR_KEYWORD, annotation=str), + Parameter( + "runtime", + Parameter.POSITIONAL_OR_KEYWORD, + annotation="Annotated[_PostponedRuntime, InjectedToolArg]", + ), + ] + ) + + def __call__( + self, query: str, runtime: Annotated[_PostponedRuntime, InjectedToolArg] + ) -> str: + received["runtime"] = runtime + return query + + tool = StructuredTool.from_function( + func=CallableWithStringSignature(), + name="callable_with_string_signature", + description="Echo a query.", + args_schema=_InputSchema, + ) + runtime = _PostponedRuntime() + + assert tool._injected_args_keys == frozenset({"runtime"}) + assert tool.invoke({"query": "hello", "runtime": runtime}) == "hello" + assert received["runtime"] is runtime + tool_call_schema = tool.tool_call_schema + assert not isinstance(tool_call_schema, dict) + assert "runtime" not in model_json_schema(tool_call_schema)["properties"] + + +def test_class_callable_resolves_postponed_injected_arg() -> None: + """Classes resolve constructor annotations postponed by the future import.""" + + class ClassTool: + def __init__(self, query: str, runtime: _PostponedRuntime) -> None: + self.query = query + self.runtime = runtime + + tool = StructuredTool.from_function( + func=ClassTool, + name="class_tool", + description="Echo a query.", + args_schema=_InputSchema, + ) + runtime = _PostponedRuntime() + + assert tool._injected_args_keys == frozenset({"runtime"}) + result = tool.invoke({"query": "hello", "runtime": runtime}) + assert result.query == "hello" + assert result.runtime is runtime + + tool_call_schema = tool.tool_call_schema + assert not isinstance(tool_call_schema, dict) + assert "runtime" not in model_json_schema(tool_call_schema)["properties"]