diff --git a/litellm/harness/handlers/tool_loop_handler.py b/litellm/harness/handlers/tool_loop_handler.py index dda353263d8..f7fcd50a05d 100644 --- a/litellm/harness/handlers/tool_loop_handler.py +++ b/litellm/harness/handlers/tool_loop_handler.py @@ -398,11 +398,11 @@ class ToolLoopHandler(BaseHarnessHandler): parse_error, approval_error, ) - yield ToolResult(id=call.id, output=outcome.output, is_error=outcome.is_error) tool_message: ChatCompletionMessageParam = { "role": "tool", "tool_call_id": call.id, "content": outcome.output, } self._messages = (*self._messages, tool_message) + yield ToolResult(id=call.id, output=outcome.output, is_error=outcome.is_error) raise HarnessTurnError(f"Harness.TOOL_LOOP exceeded {TOOL_LOOP_MAX_MODEL_CALLS} model calls in one turn") diff --git a/litellm/llms/tool_loop/harness/transformation.py b/litellm/llms/tool_loop/harness/transformation.py index 3a4ab6eb304..9a6483f8af5 100644 --- a/litellm/llms/tool_loop/harness/transformation.py +++ b/litellm/llms/tool_loop/harness/transformation.py @@ -78,9 +78,11 @@ def _routing_kwargs(ctx: SessionContext, options: ToolLoopOptions) -> Mapping[st raise ValueError("Harness.TOOL_LOOP needs model=") if ctx.gateway is not None: caller_headers_value: Final[object] = _OBJECT_ADAPTER.validate_python( - options.completion_kwargs.get("extra_headers", {}) + options.completion_kwargs.get("extra_headers") + ) + caller_headers: Final = _STRING_HEADERS_ADAPTER.validate_python( + MappingProxyType({}) if caller_headers_value is None else caller_headers_value ) - caller_headers: Final = _STRING_HEADERS_ADAPTER.validate_python(caller_headers_value) extra_headers: Final[dict[str, str]] = { # mutable-ok: acompletion(extra_headers=) requires a dict **caller_headers, **gateway_headers(ctx), diff --git a/tests/unit/harness/handlers/test_tool_loop_handler.py b/tests/unit/harness/handlers/test_tool_loop_handler.py index df1e56404a0..7fff40015f6 100644 --- a/tests/unit/harness/handlers/test_tool_loop_handler.py +++ b/tests/unit/harness/handlers/test_tool_loop_handler.py @@ -2,9 +2,9 @@ from __future__ import annotations -from collections.abc import Callable, Iterator, Mapping +from collections.abc import AsyncGenerator, Callable, Iterator, Mapping from pathlib import Path -from typing import Final, Literal +from typing import Final, Literal, cast import pytest from pydantic import BaseModel @@ -556,6 +556,52 @@ async def test_interrupted_tool_calls_are_closed_before_the_next_session_turn(tm ] +async def test_closing_after_tool_result_keeps_result_and_interrupts_remaining_call(tmp_path: Path) -> None: + executed: list[str] = [] # mutable-ok: records which injected tools ran + + def first_tool() -> str: + executed.append("first") + return "first result" + + def second_tool() -> str: + executed.append("second") + return "second result" + + completion: Final = ScriptedCompletion( + ( + model_response( + tool_calls=( + function_call("first_tool", "{}", "call-1"), + function_call("second_tool", "{}", "call-2"), + ) + ), + ) + ) + ctx: Final = make_context(tmp_path, tools=(first_tool, second_tool)) + handler: Final = make_handler(completion) + await handler.start(ctx) + + turn: Final = cast(AsyncGenerator[Event, None], handler.turn(ctx, "first turn")) + tool_results: list[ToolResult] = [] # mutable-ok: records the first yielded tool result + async for event in turn: + if isinstance(event, ToolResult): + tool_results.append(event) + break + await turn.aclose() + await handler.stop(ctx) + + assert tool_results == [ToolResult(id="call-1", output="first result", is_error=False)] + assert executed == ["first"] + assert (await handler.history(ctx))[-2:] == [ + {"role": "tool", "tool_call_id": "call-1", "content": "first result"}, + { + "role": "tool", + "tool_call_id": "call-2", + "content": "interrupted: the turn ended before this tool ran", + }, + ] + + async def test_duplicate_tool_names_are_rejected(tmp_path: Path) -> None: ctx: Final = make_context(tmp_path, tools=(add, duplicate_add())) handler: Final = make_handler(ScriptedCompletion(())) diff --git a/tests/unit/llms/tool_loop/harness/test_transformation.py b/tests/unit/llms/tool_loop/harness/test_transformation.py index 68cfdf44a8d..8dc9a8f07f5 100644 --- a/tests/unit/llms/tool_loop/harness/test_transformation.py +++ b/tests/unit/llms/tool_loop/harness/test_transformation.py @@ -161,6 +161,19 @@ def test_gateway_extra_headers_must_be_string_mappings(tmp_path: Path) -> None: completion_kwargs(ctx) +def test_gateway_none_extra_headers_uses_only_gateway_headers(tmp_path: Path) -> None: + gateway: Final = GatewayTarget(api_base="https://gateway", api_key="virtual-key") + ctx: Final = make_context( + tmp_path, + gateway=gateway, + options=ToolLoopOptions(completion_kwargs={"extra_headers": None}), + ) + + kwargs: Final = completion_kwargs(ctx) + + assert kwargs["extra_headers"] == {"x-litellm-tags": "harness,tool_loop"} + + def test_configuration_requires_model_and_declares_capabilities(tmp_path: Path) -> None: config: Final = ToolLoopHarnessConfig() assert config.uses_model_endpoint is False