fix(harness): record tool results before yielding and accept null extra_headers

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-10-03 18:02:51 +00:00
parent fb1fb4f729
commit 8b20b93309
4 changed files with 66 additions and 5 deletions

View file

@ -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")

View file

@ -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),

View file

@ -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(()))

View file

@ -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