mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
fb1fb4f729
commit
8b20b93309
4 changed files with 66 additions and 5 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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(()))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue