Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 17 additions & 9 deletions haystack/tools/parameters_schema_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand All @@ -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.

Expand All @@ -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
Expand All @@ -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
Original file line number Diff line number Diff line change
@@ -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.
77 changes: 75 additions & 2 deletions test/tools/test_component_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -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(),
Expand Down