From df8171a6eb79db48edbcfc147e6f0daf9895885b Mon Sep 17 00:00:00 2001 From: Reza Rahemtola Date: Sun, 2 Aug 2026 16:32:17 +0200 Subject: [PATCH] fix(v2): keep raw Iterable typehint for PARALLEL_TOOLS --- instructor/v2/core/patch.py | 12 ++- tests/v2/test_parallel_tools_wrapper.py | 103 ++++++++++++++++++++++++ 2 files changed, 111 insertions(+), 4 deletions(-) create mode 100644 tests/v2/test_parallel_tools_wrapper.py diff --git a/instructor/v2/core/patch.py b/instructor/v2/core/patch.py index ac436d93e..676c5f9f7 100644 --- a/instructor/v2/core/patch.py +++ b/instructor/v2/core/patch.py @@ -209,13 +209,15 @@ def new_create_sync( # Get handlers from registry handlers = mode_registry.get_handlers(provider, mode) - if response_model is not None: + if response_model is not None and mode not in Mode.parallel_modes(): response_model = prepare_response_model(response_model) # Prepare request kwargs using registry handler - response_model, new_kwargs = handlers.request_handler( + prepared_model, new_kwargs = handlers.request_handler( response_model=response_model, kwargs=kwargs ) + if mode not in Mode.parallel_modes(): + response_model = prepared_model new_kwargs.pop("autodetect_images", None) if handlers.message_converter and "messages" in new_kwargs: new_kwargs["messages"] = handlers.message_converter( @@ -323,13 +325,15 @@ async def new_create_async( # Get handlers from registry handlers = mode_registry.get_handlers(provider, mode) - if response_model is not None: + if response_model is not None and mode not in Mode.parallel_modes(): response_model = prepare_response_model(response_model) # Prepare request kwargs using registry handler - response_model, new_kwargs = handlers.request_handler( + prepared_model, new_kwargs = handlers.request_handler( response_model=response_model, kwargs=kwargs ) + if mode not in Mode.parallel_modes(): + response_model = prepared_model new_kwargs.pop("autodetect_images", None) if handlers.message_converter and "messages" in new_kwargs: new_kwargs["messages"] = handlers.message_converter( diff --git a/tests/v2/test_parallel_tools_wrapper.py b/tests/v2/test_parallel_tools_wrapper.py new file mode 100644 index 000000000..b2ae07ddf --- /dev/null +++ b/tests/v2/test_parallel_tools_wrapper.py @@ -0,0 +1,103 @@ +"""Regression tests for PARALLEL_TOOLS through the v2 patch wrapper. + +The existing handler-level tests call ``handlers.request_handler(...)`` directly +and bypass ``patch_v2``; these exercise the wrapper path ``from_openai`` uses. +""" + +from __future__ import annotations + +from collections.abc import Iterable +from typing import Any, Union, get_args, get_origin + +import pytest +from pydantic import BaseModel + +from instructor.v2.core.mode import Mode +from instructor.v2.core.patch import patch_v2 +from instructor.v2.core.providers import Provider +from tests.coverage._openai import chat_completion, tool_call + + +class A(BaseModel): + a: str + + +class B(BaseModel): + b: str + + +def _member_types(response_model: Any) -> tuple[type[BaseModel], ...]: + inner = get_args(response_model)[0] + origin = get_origin(inner) + if origin is Union: + return get_args(inner) + return (inner,) + + +def _parallel_completion(response_model: Any) -> Any: + members = _member_types(response_model) + payload = {"A": {"a": "alpha"}, "B": {"b": "beta"}} + return chat_completion( + tool_calls=[ + tool_call( + m.__name__, payload[m.__name__], call_id=f"call_{m.__name__.lower()}" + ) + for m in members + ], + finish_reason="tool_calls", + ) + + +def _assert_parallel_result(result: Any, response_model: Any) -> None: + members = _member_types(response_model) + expected_names = [m.__name__ for m in members] + items = list(result) + assert [type(x).__name__ for x in items] == expected_names + assert items[0].a == "alpha" + if B in members: + assert items[1].b == "beta" + + +@pytest.mark.parametrize( + "response_model", + [ + pytest.param(Iterable[A], id="Iterable[A]"), + pytest.param(Iterable[Union[A, B]], id="Iterable[Union[A,B]]"), + ], +) +def test_parallel_tools_sync_wrapper(response_model: Any) -> None: + calls: list[dict[str, Any]] = [] + + def create(**kwargs: Any) -> Any: + calls.append(kwargs) + return _parallel_completion(response_model) + + patched = patch_v2(create, Provider.OPENAI, Mode.PARALLEL_TOOLS) + result = patched( + response_model=response_model, + messages=[{"role": "user", "content": "run both"}], + ) + + assert calls, "create was never called" + _assert_parallel_result(result, response_model) + tool_names = {t["function"]["name"] for t in calls[0]["tools"]} + assert tool_names == {m.__name__ for m in _member_types(response_model)} + + +@pytest.mark.asyncio +async def test_parallel_tools_async_wrapper() -> None: + response_model = Iterable[Union[A, B]] + calls: list[dict[str, Any]] = [] + + async def create(**kwargs: Any) -> Any: + calls.append(kwargs) + return _parallel_completion(response_model) + + patched = patch_v2(create, Provider.OPENAI, Mode.PARALLEL_TOOLS) + result = await patched( + response_model=response_model, + messages=[{"role": "user", "content": "run both"}], + ) + + assert calls, "create was never called" + _assert_parallel_result(result, response_model)