From e896ed1f289f999a463e0b9344422eafe9af3862 Mon Sep 17 00:00:00 2001 From: pei711 <199601925+pei711@users.noreply.github.com> Date: Sat, 3 Oct 2026 20:23:36 +0800 Subject: [PATCH] fix: resolve postponed dataclass fields in ComponentTool --- haystack/tools/parameters_schema_utils.py | 26 ++++--- ...ostponed-dataclasses-3c12e2db9a71ee48.yaml | 6 ++ test/tools/test_component_tool.py | 77 ++++++++++++++++++- 3 files changed, 98 insertions(+), 11 deletions(-) create mode 100644 releasenotes/notes/component-tool-postponed-dataclasses-3c12e2db9a71ee48.yaml diff --git a/haystack/tools/parameters_schema_utils.py b/haystack/tools/parameters_schema_utils.py index fcaac47c6ad..67b0a5363be 100644 --- a/haystack/tools/parameters_schema_utils.py +++ b/haystack/tools/parameters_schema_utils.py @@ -9,7 +9,7 @@ from dataclasses import MISSING, fields, is_dataclass from inspect import getdoc from types import NoneType -from typing import Any, Union, get_args, get_origin +from typing import Any, ForwardRef, Union, get_args, get_origin, get_type_hints from docstring_parser import parse from pydantic import BaseModel, Field, create_model @@ -148,21 +148,24 @@ def _get_component_param_descriptions(component: Any) -> dict[str, str]: return param_descriptions -def _dataclass_to_pydantic_model(dc_type: Any) -> type[BaseModel]: +def _dataclass_to_pydantic_model(dc_type: Any, _seen: set[type[Any]] | None = None) -> type[BaseModel]: """ Convert a Python dataclass to an equivalent Pydantic model. :param dc_type: The dataclass type to convert. + :param _seen: Dataclass types currently being converted, used to retain recursive references. :returns: A dynamically generated Pydantic model class with fields and types derived from the dataclass definition. Field descriptions are extracted from docstrings when available. """ _, param_descriptions = _get_param_descriptions(dc_type) cls = dc_type if isinstance(dc_type, type) else dc_type.__class__ + seen = (_seen or set()) | {cls} + field_types = get_type_hints(cls, include_extras=True) field_defs: dict[str, Any] = {} for field in fields(dc_type): - f_type = field.type if isinstance(field.type, str) else _resolve_type(field.type) + f_type = _resolve_type(field_types[field.name], seen) default = field.default if field.default is not MISSING else ... default = field.default_factory() if callable(field.default_factory) else default @@ -178,7 +181,7 @@ def _dataclass_to_pydantic_model(dc_type: Any) -> type[BaseModel]: return create_model(cls.__name__, **field_defs) -def _resolve_type(_type: Any) -> Any: # noqa: PLR0911 +def _resolve_type(_type: Any, _seen: set[type[Any]] | None = None) -> Any: # noqa: PLR0911 """ Recursively resolve and convert complex type annotations, transforming dataclasses into Pydantic-compatible types. @@ -187,6 +190,7 @@ def _resolve_type(_type: Any) -> Any: # noqa: PLR0911 :param _type: The type annotation to resolve. If the type is a dataclass, it will be converted to a Pydantic model. For generic types (like list[SomeDataclass]), the inner types are also resolved recursively. + :param _seen: Dataclass types currently being converted, used to retain recursive references. :returns: A fully resolved type, with all dataclass types converted to Pydantic models @@ -203,21 +207,25 @@ def _resolve_type(_type: Any) -> Any: # noqa: PLR0911 return _ToolsetSchemaPlaceholder if is_dataclass(_type): - return _dataclass_to_pydantic_model(_type) + cls = _type if isinstance(_type, type) else _type.__class__ + if _seen and cls in _seen: + # Let Pydantic resolve a back-edge to the model being built instead of expanding it again. + return ForwardRef(cls.__name__) + return _dataclass_to_pydantic_model(_type, _seen) origin = get_origin(_type) args = get_args(_type) if origin is list: - return list[_resolve_type(args[0]) if args else Any] # type: ignore[misc] + return list[_resolve_type(args[0], _seen) if args else Any] # type: ignore[misc] if origin is collections.abc.Sequence: - return Sequence[_resolve_type(args[0]) if args else Any] # type: ignore[misc] + return Sequence[_resolve_type(args[0], _seen) if args else Any] # type: ignore[misc] if _is_union_type(origin): - return Union[tuple(_resolve_type(a) for a in args)] + return Union[tuple(_resolve_type(a, _seen) for a in args)] if origin is dict: - return dict[args[0] if args else Any, _resolve_type(args[1]) if args else Any] # type: ignore[misc] + return dict[args[0] if args else Any, _resolve_type(args[1], _seen) if args else Any] # type: ignore[misc] return _type diff --git a/releasenotes/notes/component-tool-postponed-dataclasses-3c12e2db9a71ee48.yaml b/releasenotes/notes/component-tool-postponed-dataclasses-3c12e2db9a71ee48.yaml new file mode 100644 index 00000000000..83cf6a60955 --- /dev/null +++ b/releasenotes/notes/component-tool-postponed-dataclasses-3c12e2db9a71ee48.yaml @@ -0,0 +1,6 @@ +--- +fixes: + - | + ``ComponentTool`` now resolves dataclass field annotations stored as strings, including fields in modules using + ``from __future__ import annotations``. Nested dataclasses and ``Annotated`` constraints are preserved in the + generated tool schema, while recursive dataclass references continue to work. diff --git a/test/tools/test_component_tool.py b/test/tools/test_component_tool.py index ea4199476b1..6ffb4bc292a 100644 --- a/test/tools/test_component_tool.py +++ b/test/tools/test_component_tool.py @@ -4,13 +4,14 @@ import json import os -from dataclasses import dataclass -from typing import Any +from dataclasses import dataclass, field +from typing import Annotated, Any from unittest.mock import patch import pytest from openai.types.chat import ChatCompletion, ChatCompletionMessage from openai.types.chat.chat_completion import Choice +from pydantic import Field from haystack import Pipeline, SuperComponent, component from haystack.components.agents import Agent, State @@ -132,6 +133,56 @@ def run(self, person: Person) -> dict[str, str]: return {"info": f"{person.name} lives at {person.address.street}, {person.address.city}."} +@dataclass +class PostponedPerson: + """A dataclass whose field annotations are stored as strings.""" + + name: "Annotated[str, Field(min_length=1)]" + addresses: "list[Address]" + + +@component +class PostponedPersonProcessor: + """Process a dataclass with postponed annotations through sync and async tools.""" + + @component.output_types(info=str) + def run(self, person: PostponedPerson) -> dict[str, str]: + """ + :param person: The person and their addresses. + :returns: A summary of the person's addresses. + """ + return {"info": f"{person.name}: {', '.join(address.city for address in person.addresses)}"} + + @component.output_types(info=str) + async def run_async(self, person: PostponedPerson) -> dict[str, str]: + """ + :param person: The person and their addresses. + :returns: A summary of the person's addresses. + """ + return self.run(person) + + +@dataclass +class RecursivePerson: + """A recursive dataclass with a forward reference to itself.""" + + name: str + children: "list[RecursivePerson]" = field(default_factory=list) + + +@component +class RecursivePersonProcessor: + """Process a recursive dataclass.""" + + @component.output_types(names=list[str]) + def run(self, person: RecursivePerson) -> dict[str, list[str]]: + """ + :param person: The person and their children. + :returns: The children's names. + """ + return {"names": [child.name for child in person.children]} + + @component class DocumentProcessor: """A component that processes a list of Documents.""" @@ -287,6 +338,28 @@ def test_from_component_with_nested_dataclass(self): assert "info" in result assert result["info"] == "Diana lives at 123 Elm Street, Metropolis." + @pytest.mark.asyncio + @pytest.mark.parametrize("use_async", [False, True]) + async def test_postponed_dataclass_fields_keep_nested_schema_and_constraints(self, use_async: bool) -> None: + tool = ComponentTool(component=PostponedPersonProcessor()) + definitions = tool.parameters["$defs"] + person_properties = definitions["PostponedPerson"]["properties"] + assert person_properties["name"]["minLength"] == 1 + assert person_properties["addresses"]["items"] == {"$ref": "#/$defs/Address"} + assert definitions["Address"]["properties"]["city"]["description"] == "Field 'city' of 'Address'." + + restored = ComponentTool.from_dict(tool.to_dict()) + assert restored.parameters == tool.parameters + person = {"name": "Diana", "addresses": [{"street": "123 Elm Street", "city": "Metropolis"}]} + result = await restored.invoke_async(person=person) if use_async else restored.invoke(person=person) + assert result == {"info": "Diana: Metropolis"} + + def test_recursive_dataclass_schema_and_invocation(self) -> None: + tool = ComponentTool(component=RecursivePersonProcessor()) + person_schema = tool.parameters["$defs"]["RecursivePerson"] + assert person_schema["properties"]["children"]["items"] == {"$ref": "#/$defs/RecursivePerson"} + assert tool.invoke(person={"name": "Diana", "children": [{"name": "Alice"}]}) == {"names": ["Alice"]} + def test_from_component_with_list_of_documents(self): tool = ComponentTool( component=DocumentProcessor(),