Skip to content
Closed
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
37 changes: 37 additions & 0 deletions haystack/components/generators/chat/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,37 @@
logger = logging.getLogger(__name__)


def _merge_tools_from_kwargs(openai_tools: dict[str, Any], generation_kwargs: dict[str, Any]) -> dict[str, Any]:
"""
Merge OpenAI tool specs passed via ``generation_kwargs`` into the component's tools.

Raw OpenAI tool specs passed via ``generation_kwargs["tools"]`` are combined with the
component's own tool definitions instead of letting them override each other. For a
tool name present in both, the spec from ``generation_kwargs`` wins, mirroring how
every other ``generation_kwargs`` key takes precedence. The ``tools`` key is removed
from ``generation_kwargs`` so the caller's parameter is not mutated.
"""
kwargs_tools = generation_kwargs.pop("tools", None)
if not kwargs_tools:
return openai_tools

def _tool_name(tool: dict[str, Any]) -> str:
# chat.completions specs nest the name under "function"; Responses API specs
# put it at the top level. Accept both.
function = tool.get("function")
if isinstance(function, dict):
return str(function.get("name"))
return str(tool.get("name"))

# dict keyed by tool name so a later (kwargs) definition overrides an earlier one
merged: dict[str, dict[str, Any]] = {}
for tool in openai_tools.get("tools", []):
merged[_tool_name(tool)] = tool
for tool in kwargs_tools:
merged[_tool_name(tool)] = tool
return {"tools": list(merged.values())}


@component
class OpenAIChatGenerator:
"""
Expand Down Expand Up @@ -538,6 +569,12 @@ def _prepare_api_call( # noqa: PLR0913
tool_definitions.append({"type": "function", "function": function_spec})
openai_tools = {"tools": tool_definitions}

# Merge tools passed via generation_kwargs (raw OpenAI specs) with the
# component's own tools instead of letting them override each other. For a
# tool name present in both, the spec from generation_kwargs wins, mirroring
# how every other generation_kwargs key takes precedence.
openai_tools = _merge_tools_from_kwargs(openai_tools=openai_tools, generation_kwargs=generation_kwargs)

base_args = {
"model": self.model,
"messages": openai_formatted_messages,
Expand Down
5 changes: 5 additions & 0 deletions haystack/components/generators/chat/openai_responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from pydantic import BaseModel

from haystack import component, default_from_dict, default_to_dict, logging
from haystack.components.generators.chat.openai import _merge_tools_from_kwargs
from haystack.components.generators.utils import _normalize_messages, _serialize_object
from haystack.dataclasses import (
ChatMessage,
Expand Down Expand Up @@ -578,6 +579,10 @@ def _prepare_api_call( # noqa: PLR0913

openai_tools = {"tools": tool_definitions}

# Same merge as OpenAIChatGenerator: tools passed via generation_kwargs are
# combined with the component's own tools rather than replacing them.
openai_tools = _merge_tools_from_kwargs(openai_tools=openai_tools, generation_kwargs=generation_kwargs)

base_args = {"model": self.model, "input": openai_formatted_messages, **openai_tools, **generation_kwargs}

# if `text_format` is provided, we use the `parse` endpoint for response type parsing
Expand Down
85 changes: 85 additions & 0 deletions test/components/generators/chat/test_openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -523,6 +523,91 @@ def test_run_with_generation_kwargs(
assert kwargs["temperature"] == 0.9
assert kwargs["max_completion_tokens"] == 10

def test_run_merged_tools_from_generation_kwargs(
self, chat_messages: list[ChatMessage], openai_mock_chat_completion: MagicMock
) -> None:
"""Tools passed via generation_kwargs merge with the component's own tools
instead of silently replacing them. For a name in both, the kwargs spec wins."""

def haystack_tool() -> None: ...

component = OpenAIChatGenerator(
api_key=Secret.from_token("test-api-key"),
tools=[
Tool(
name="haystack_tool",
description="hs",
parameters={"type": "object", "properties": {}},
function=haystack_tool,
)
],
)
component.run(
chat_messages,
generation_kwargs={
"tools": [
{
"type": "function",
"function": {
"name": "raw_tool",
"description": "raw",
"parameters": {"type": "object", "properties": {}},
},
}
]
},
)

_, kwargs = openai_mock_chat_completion.call_args
tool_names = [t["function"]["name"] for t in kwargs["tools"]]
assert tool_names == ["haystack_tool", "raw_tool"]

def test_run_generation_kwargs_tool_wins_on_name_collision(
self, chat_messages: list[ChatMessage], openai_mock_chat_completion: MagicMock
) -> None:
"""When the same tool name appears in both, the generation_kwargs spec wins."""

def haystack_tool() -> None: ...

component = OpenAIChatGenerator(
api_key=Secret.from_token("test-api-key"),
tools=[
Tool(
name="shared_tool",
description="component",
parameters={"type": "object", "properties": {}},
function=haystack_tool,
)
],
)
component.run(
chat_messages,
generation_kwargs={
"tools": [
{
"type": "function",
"function": {
"name": "shared_tool",
"description": "kwargs",
"parameters": {"type": "object", "properties": {}},
},
}
]
},
)

_, kwargs = openai_mock_chat_completion.call_args
assert kwargs["tools"] == [
{
"type": "function",
"function": {
"name": "shared_tool",
"description": "kwargs",
"parameters": {"type": "object", "properties": {}},
},
}
]

def test_run_with_params_streaming(
self, chat_messages: list[ChatMessage], openai_mock_chat_completion_chunk: MagicMock
) -> None:
Expand Down
82 changes: 82 additions & 0 deletions test/components/generators/chat/test_openai_responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -603,6 +603,88 @@ def test_run_with_empty_tools_override(self, tools: list[Tool], openai_mock_resp

assert "tools" not in openai_mock_responses.call_args.kwargs

def test_run_merged_tools_from_generation_kwargs(self, openai_mock_responses: MagicMock) -> None:
"""Tools passed via generation_kwargs merge with the component's own tools
for the Responses API as well."""

def haystack_tool() -> None:
...

component = OpenAIResponsesChatGenerator(
api_key=Secret.from_token("test-api-key"),
tools=[
Tool(
name="haystack_tool",
description="hs",
parameters={"type": "object", "properties": {}},
function=haystack_tool,
)
],
)
component.run(
[ChatMessage.from_user("What's the capital of France")],
generation_kwargs={
"tools": [
{
"type": "function",
"function": {
"name": "raw_tool",
"description": "raw",
"parameters": {"type": "object", "properties": {}},
},
}
]
},
)

kwargs = openai_mock_responses.call_args.kwargs
tool_names = [
t.get("name") if isinstance(t.get("name"), str) else t["function"]["name"] for t in kwargs["tools"]
]
assert tool_names == ["haystack_tool", "raw_tool"]

def test_run_generation_kwargs_tool_wins_on_name_collision(self, openai_mock_responses: MagicMock) -> None:
"""When the same tool name appears in both, the generation_kwargs spec wins."""

def haystack_tool() -> None:
...

component = OpenAIResponsesChatGenerator(
api_key=Secret.from_token("test-api-key"),
tools=[
Tool(
name="shared_tool",
description="component",
parameters={"type": "object", "properties": {}},
function=haystack_tool,
)
],
)
component.run(
[ChatMessage.from_user("What's the capital of France")],
generation_kwargs={
"tools": [
{
"type": "function",
"function": {
"name": "shared_tool",
"description": "kwargs",
"parameters": {"type": "object", "properties": {}},
},
}
]
},
)

kwargs = openai_mock_responses.call_args.kwargs
shared = next(
t
for t in kwargs["tools"]
if t.get("name") == "shared_tool" or t.get("function", {}).get("name") == "shared_tool"
)
assert shared["function"]["name"] == "shared_tool"
assert shared["function"]["description"] == "kwargs"

def test_run_with_generation_kwargs(self, openai_mock_responses: MagicMock) -> None:

component = OpenAIResponsesChatGenerator(
Expand Down
Loading