From fb1fb4f729237a5fc0ba5f2917b8401c2e80cc58 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 17:39:03 +0000 Subject: [PATCH 1/2] fix(harness): harden Harness.TOOL_LOOP history, args, output, headers and accounting Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm/harness/handlers/tool_loop_handler.py | 87 +++++++++-- .../llms/tool_loop/harness/transformation.py | 19 ++- .../handlers/test_tool_loop_handler.py | 136 ++++++++++++++++-- .../tool_loop/harness/test_transformation.py | 39 ++++- 4 files changed, 253 insertions(+), 28 deletions(-) diff --git a/litellm/harness/handlers/tool_loop_handler.py b/litellm/harness/handlers/tool_loop_handler.py index 33efcba4ff7..dda353263d8 100644 --- a/litellm/harness/handlers/tool_loop_handler.py +++ b/litellm/harness/handlers/tool_loop_handler.py @@ -46,10 +46,12 @@ class _FunctionToolCall: AsyncCompletion: TypeAlias = Callable[..., Awaitable[ModelResponse]] _ARGUMENTS_ADAPTER: Final = TypeAdapter(dict[str, object]) _MAPPING_ADAPTER: Final = TypeAdapter(Mapping[str, object]) +_OBJECT_ADAPTER: Final = TypeAdapter(object) _MODEL_RESPONSE_ADAPTER: Final = TypeAdapter(ModelResponse) _RAW_DECODE_ADAPTER: Final = TypeAdapter(tuple[object, int]) _JSON_DECODER: Final = json.JSONDecoder() _HISTORY_ADAPTER: Final = TypeAdapter(list[dict[str, object]]) +_TOOL_CALLS_ADAPTER: Final = TypeAdapter(tuple[Mapping[str, object], ...]) class _Usage(BaseModel): @@ -88,6 +90,41 @@ def _normalize_tool_call( return _FunctionToolCall(id=call.id, name=call.custom.name, arguments=call.custom.input) +def _tool_call_id(call: Mapping[str, object]) -> str | None: + call_id: Final[object] = call.get("id") + return call_id if isinstance(call_id, str) else None + + +def _tool_call_ids(message: ChatCompletionMessageParam) -> tuple[str, ...]: + message_mapping: Final = _as_mapping(message) + if message_mapping is None or message_mapping.get("role") != "assistant": + return () + tool_calls_value: Final = message_mapping.get("tool_calls") + if tool_calls_value is None: + return () + try: + tool_calls: Final = _TOOL_CALLS_ADAPTER.validate_python(tool_calls_value) + except ValidationError: + return () + return tuple(call_id for call_id in map(_tool_call_id, tool_calls) if call_id is not None) + + +def _tool_message_call_id(message: ChatCompletionMessageParam) -> str | None: + message_mapping: Final = _as_mapping(message) + if message_mapping is None or message_mapping.get("role") != "tool": + return None + call_id: Final[object] = message_mapping.get("tool_call_id") + return call_id if isinstance(call_id, str) else None + + +def _interrupted_tool_message(call_id: str) -> ChatCompletionMessageParam: + return { + "role": "tool", + "tool_call_id": call_id, + "content": "interrupted: the turn ended before this tool ran", + } + + def _parse_arguments(raw: str) -> tuple[dict[str, object], str | None]: # mutable-ok: event input uses a dict text: Final = raw.strip() try: @@ -173,24 +210,31 @@ async def _tool_outcome( def _record_usage(ctx: SessionContext, response: ModelResponse) -> None: + ctx.calls += 1 # rebind-ok: SessionContext is the runtime's per-session usage sink + input_tokens, output_tokens = _usage_counts(response) + ctx.input_tokens += input_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink + ctx.output_tokens += output_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink + ctx.cost += _response_cost(response) # rebind-ok: SessionContext is the runtime's per-session usage sink + + +def _usage_counts(response: ModelResponse) -> tuple[int, int]: try: usage_value: Final[object] = getattr(response, "usage", None) usage: Final = _USAGE_ADAPTER.validate_python(usage_value) - input_tokens: Final = usage.prompt_tokens or 0 - output_tokens: Final = usage.completion_tokens or 0 - ctx.calls += 1 # rebind-ok: SessionContext is the runtime's per-session usage sink - ctx.input_tokens += input_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink - ctx.output_tokens += output_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink - ctx.cost += _response_cost(response) # rebind-ok: SessionContext is the runtime's per-session usage sink except Exception: - return + return 0, 0 + return usage.prompt_tokens or 0, usage.completion_tokens or 0 async def _execute_tool(tool: FunctionTool, arguments: Mapping[str, object]) -> object: validated_model: Final = tool.args_model.model_validate(arguments) - values_object: Final[object] = validated_model.model_dump() - validated: Final = _MAPPING_ADAPTER.validate_python(values_object) parameters: Final = tuple(inspect.signature(tool.fn).parameters.values()) + validated: Final[Mapping[str, object]] = MappingProxyType( + { + parameter.name: _OBJECT_ADAPTER.validate_python(getattr(validated_model, parameter.name)) + for parameter in parameters + } + ) positional_args: Final = tuple( validated[parameter.name] for parameter in parameters if parameter.kind is inspect.Parameter.POSITIONAL_ONLY ) @@ -253,11 +297,31 @@ class ToolLoopHandler(BaseHarnessHandler): return _HISTORY_ADAPTER.validate_python(history_object) # pyright: ignore[reportIncompatibleMethodOverride] # base history uses Any async def stop(self, ctx: SessionContext) -> None: - return None + self._close_pending_calls() + + def _close_pending_calls(self) -> None: + assistant_offset: Final[int | None] = next( + (offset for offset, message in enumerate(reversed(self._messages)) if _tool_call_ids(message)), + None, + ) + if assistant_offset is None: + return + assistant_message: Final = self._messages[-assistant_offset - 1] + messages_after: Final = self._messages[len(self._messages) - assistant_offset :] + completed_call_ids: Final = frozenset( + call_id for call_id in map(_tool_message_call_id, messages_after) if call_id is not None + ) + pending_call_ids: Final = tuple( + call_id for call_id in _tool_call_ids(assistant_message) if call_id not in completed_call_ids + ) + pending_messages: Final[tuple[ChatCompletionMessageParam, ...]] = tuple( + _interrupted_tool_message(call_id) for call_id in pending_call_ids + ) + self._messages = (*self._messages, *pending_messages) async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]: + self._close_pending_calls() ctx.final_text = "" # rebind-ok: SessionContext is the runtime's per-turn result sink - ctx.output_json = None # rebind-ok: SessionContext is the runtime's per-turn result sink user_message: Final[ChatCompletionMessageParam] = {"role": "user", "content": prompt} self._messages = (*self._messages, user_message) for _ in range(TOOL_LOOP_MAX_MODEL_CALLS): @@ -289,7 +353,6 @@ class ToolLoopHandler(BaseHarnessHandler): if not tool_calls: final_text = content or "" ctx.final_text = final_text # rebind-ok: SessionContext is the runtime's per-turn result sink - ctx.output_json = content if ctx.output is not None else None # rebind-ok: per-turn output sink final_message: ChatCompletionMessageParam = { "role": "assistant", "content": content, diff --git a/litellm/llms/tool_loop/harness/transformation.py b/litellm/llms/tool_loop/harness/transformation.py index 16461916f1a..3a4ab6eb304 100644 --- a/litellm/llms/tool_loop/harness/transformation.py +++ b/litellm/llms/tool_loop/harness/transformation.py @@ -20,6 +20,7 @@ from litellm.types.utils import ChatCompletionToolParam TOOL_LOOP_MAX_MODEL_CALLS: Final = 100 _ANNOTATIONS_ADAPTER: Final = TypeAdapter(Mapping[str, object]) _OBJECT_ADAPTER: Final = TypeAdapter(object) +_STRING_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str]) _MODEL_FACTORY: Final[Callable[..., type[BaseModel]]] = create_model @@ -72,15 +73,23 @@ def function_tool(fn: Callable[..., object]) -> FunctionTool: return FunctionTool(name=fn.__name__, fn=fn, args_model=args_model, spec=spec) -def _routing_kwargs(ctx: SessionContext) -> Mapping[str, object]: +def _routing_kwargs(ctx: SessionContext, options: ToolLoopOptions) -> Mapping[str, object]: if not ctx.model: 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", {}) + ) + 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), + } return { "model": f"litellm_proxy/{ctx.model}", "api_base": ctx.gateway.api_base, "api_key": ctx.gateway.api_key, - "extra_headers": gateway_headers(ctx), + "extra_headers": extra_headers, } return { "model": ctx.model, @@ -91,7 +100,7 @@ def _routing_kwargs(ctx: SessionContext) -> Mapping[str, object]: def completion_kwargs(ctx: SessionContext) -> Mapping[str, object]: options: Final = ToolLoopHarnessConfig().get_options(ctx) - routing: Final = _routing_kwargs(ctx) + routing: Final = _routing_kwargs(ctx, options) kwargs: Final[Mapping[str, object]] = MappingProxyType({**options.completion_kwargs, **routing}) if ctx.output is None: return kwargs @@ -114,6 +123,8 @@ class ToolLoopHarnessConfig(BaseHarnessConfig[ToolLoopOptions]): ) def validate_environment(self, ctx: SessionContext) -> None: - self.get_options(ctx) + options: Final = self.get_options(ctx) + if options.completion_kwargs.get("stream"): + raise ValueError("Harness.TOOL_LOOP does not support completion_kwargs['stream']") if not ctx.model: raise ValueError("Harness.TOOL_LOOP needs model=") diff --git a/tests/unit/harness/handlers/test_tool_loop_handler.py b/tests/unit/harness/handlers/test_tool_loop_handler.py index 2ca265cda67..df1e56404a0 100644 --- a/tests/unit/harness/handlers/test_tool_loop_handler.py +++ b/tests/unit/harness/handlers/test_tool_loop_handler.py @@ -6,11 +6,12 @@ from collections.abc import Callable, Iterator, Mapping from pathlib import Path from typing import Final, Literal -import litellm import pytest from pydantic import BaseModel +import litellm from litellm import sandbox +from litellm.harness import runtime from litellm.harness.context import GatewayTarget, SessionContext from litellm.harness.handlers.tool_loop_handler import ToolLoopHandler from litellm.harness.options import ToolLoopOptions @@ -351,21 +352,56 @@ async def test_async_tool_is_awaited(tmp_path: Path) -> None: assert ToolResult(id="call-1", output="12", is_error=False) in events -class Answer(BaseModel): - value: int +class Range(BaseModel): + start: int + end: int -async def test_structured_output_is_forwarded_and_retained(tmp_path: Path) -> None: - completion: Final = ScriptedCompletion((model_response(content='{"value": 7}'),)) - ctx: Final = make_context(tmp_path, output=Answer) +async def test_nested_pydantic_arguments_reach_the_tool_as_models(tmp_path: Path) -> None: + received: list[Range] = [] # mutable-ok: captures injected tool arguments + + def describe_range(value: Range) -> str: + received.append(value) + return f"{value.start}:{value.end}" + + completion: Final = ScriptedCompletion( + ( + model_response(tool_calls=(function_call("describe_range", '{"value": {"start": 2, "end": 5}}'),)), + model_response(content="Range received"), + ) + ) + ctx: Final = make_context(tmp_path, tools=(describe_range,)) handler: Final = make_handler(completion) await handler.start(ctx) - await run_turn(handler, ctx, "Return a value") + events: Final = await run_turn(handler, ctx, "Describe a range") - assert completion.calls[0]["response_format"] is Answer - assert ctx.output_json == '{"value": 7}' - assert ctx.final_text == '{"value": 7}' + assert received == [Range(start=2, end=5)] + assert isinstance(received[0], Range) + assert ToolResult(id="call-1", output="2:5", is_error=False) in events + + +class Review(BaseModel): + summary: str + + +async def test_structured_output_parses_fenced_json_and_forwards_response_format(tmp_path: Path) -> None: + content: Final = '```json\n{"summary": "Looks good"}\n```' + completion: Final = ScriptedCompletion((model_response(content=content),)) + session: Final = runtime.aagent_session( + Harness.TOOL_LOOP, + sandbox=sandbox.local(tmp_path), + model="gpt-4o-mini", + output=Review, + ) + session.handler = make_handler(completion) + async with session: + result: Final = await session.arun("Return a review") + + assert isinstance(result.output, Review) + assert result.output.summary == "Looks good" + assert completion.calls[0]["response_format"] is Review + assert session.ctx.output_json is None async def test_gateway_routing_uses_proxy_model_and_tool_loop_tag(tmp_path: Path) -> None: @@ -418,6 +454,22 @@ async def test_usage_and_cost_accumulate_across_model_calls(tmp_path: Path) -> N assert ctx.cost == pytest.approx(0.6) +async def test_missing_usage_still_records_call_and_response_cost(tmp_path: Path) -> None: + response: Final = model_response(content="done", hidden_params={"response_cost": 0.5}) + response.usage = None + completion: Final = ScriptedCompletion((response,)) + ctx: Final = make_context(tmp_path) + handler: Final = make_handler(completion) + await handler.start(ctx) + + await run_turn(handler, ctx, "Hi") + + assert ctx.calls == 1 + assert ctx.input_tokens == 0 + assert ctx.output_tokens == 0 + assert ctx.cost == pytest.approx(0.5) + + async def test_history_survives_stop_and_start(tmp_path: Path) -> None: completion: Final = ScriptedCompletion((model_response(content="first"), model_response(content="second"))) ctx: Final = make_context(tmp_path, instructions="Keep answers concise") @@ -440,6 +492,70 @@ async def test_history_survives_stop_and_start(tmp_path: Path) -> None: assert (await handler.history(ctx))[0]["content"] == "Keep answers concise" +async def test_interrupted_tool_calls_are_closed_before_the_next_session_turn(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"), + ) + ), + model_response(content="Continued"), + ) + ) + session: Final = runtime.aagent_session( + Harness.TOOL_LOOP, + sandbox=sandbox.local(tmp_path), + model="gpt-4o-mini", + tools=(first_tool, second_tool), + max_turns=1, + ) + session.handler = make_handler(completion) + async with session: + first_result: Final = await session.arun("first turn") + assert first_result.stop_reason == "max_turns" + await session.arun("second turn") + + assert executed == ["first"] + assert completion.calls[1]["messages"] == [ + {"role": "user", "content": "first turn"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "first_tool", "arguments": "{}"}, + }, + { + "id": "call-2", + "type": "function", + "function": {"name": "second_tool", "arguments": "{}"}, + }, + ], + }, + {"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", + }, + {"role": "user", "content": "second turn"}, + ] + + 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 11bce41bb56..68cfdf44a8d 100644 --- a/tests/unit/llms/tool_loop/harness/test_transformation.py +++ b/tests/unit/llms/tool_loop/harness/test_transformation.py @@ -7,7 +7,7 @@ from pathlib import Path from typing import Final, Literal import pytest -from pydantic import BaseModel +from pydantic import BaseModel, ValidationError from litellm import sandbox from litellm.harness.context import GatewayTarget, SessionContext @@ -126,6 +126,10 @@ def test_gateway_routing_and_response_format_override_options(tmp_path: Path) -> "api_key": "wrong-key", "api_base": "https://wrong", "response_format": "wrong-format", + "extra_headers": { + "x-caller-header": "preserved", + "x-litellm-tags": "caller-tag", + }, } ), output=OutputModel, @@ -137,11 +141,26 @@ def test_gateway_routing_and_response_format_override_options(tmp_path: Path) -> "model": "litellm_proxy/anthropic/claude", "api_base": "https://gateway", "api_key": "virtual-key", - "extra_headers": {"x-litellm-tags": "harness,tool_loop"}, + "extra_headers": { + "x-caller-header": "preserved", + "x-litellm-tags": "harness,tool_loop", + }, "response_format": OutputModel, } +def test_gateway_extra_headers_must_be_string_mappings(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": {"x-invalid": 1}}), + ) + + with pytest.raises(ValidationError): + completion_kwargs(ctx) + + def test_configuration_requires_model_and_declares_capabilities(tmp_path: Path) -> None: config: Final = ToolLoopHarnessConfig() assert config.uses_model_endpoint is False @@ -152,3 +171,19 @@ def test_configuration_requires_model_and_declares_capabilities(tmp_path: Path) assert config.capabilities.permission_modes == frozenset({"ask", "full"}) with pytest.raises(ValueError, match=r"Harness\.TOOL_LOOP needs model="): config.validate_environment(make_context(tmp_path, model=None)) + + +def test_configuration_rejects_truthy_completion_stream(tmp_path: Path) -> None: + ctx: Final = make_context(tmp_path, options=ToolLoopOptions(completion_kwargs={"stream": True})) + + with pytest.raises( + ValueError, + match=r"Harness\.TOOL_LOOP does not support completion_kwargs\['stream'\]", + ): + ToolLoopHarnessConfig().validate_environment(ctx) + + +def test_configuration_allows_false_completion_stream(tmp_path: Path) -> None: + ctx: Final = make_context(tmp_path, options=ToolLoopOptions(completion_kwargs={"stream": False})) + + ToolLoopHarnessConfig().validate_environment(ctx) From 8b20b93309c9ea2f3dbbcaa975e1cdfda5a3578e Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 18:02:51 +0000 Subject: [PATCH 2/2] fix(harness): record tool results before yielding and accept null extra_headers Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm/harness/handlers/tool_loop_handler.py | 2 +- .../llms/tool_loop/harness/transformation.py | 6 ++- .../handlers/test_tool_loop_handler.py | 50 ++++++++++++++++++- .../tool_loop/harness/test_transformation.py | 13 +++++ 4 files changed, 66 insertions(+), 5 deletions(-) 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