This commit is contained in:
devin-ai-integration[bot] 2026-10-04 00:21:31 -07:00 • committed by GitHub
commit bde0e164f2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 317 additions and 31 deletions

View file

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

@ -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,25 @@ 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(
MappingProxyType({}) if caller_headers_value is None else 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 +102,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 +125,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=")

View file

@ -2,15 +2,16 @@
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 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,116 @@ 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_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

@ -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,39 @@ 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_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
@ -152,3 +184,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)