mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(harness): add Harness.TOOL_LOOP, a minimal in-process tool-calling loop (#44391)
* feat(harness): add Harness.TOOL_LOOP, a minimal in-process tool-calling loop Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * refactor(harness): name per-tool spec FunctionTool Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
parent
231a46e40b
commit
b024950353
20 changed files with 1148 additions and 25 deletions
|
|
@ -2302,6 +2302,7 @@ _AGENT_EXPORTS: Final = frozenset(
|
|||
"CodexOptions",
|
||||
"OpenCodeOptions",
|
||||
"DeepAgentsOptions",
|
||||
"ToolLoopOptions",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
"""Agent harnesses: run Claude Code, Codex, OpenCode or Deep Agents on any LiteLLM model.
|
||||
"""Agent harnesses: run Claude Code, Codex, OpenCode, Deep Agents or Tool Loop on any LiteLLM model.
|
||||
|
||||
The entrypoints live on the top-level package:
|
||||
|
||||
|
|
@ -30,6 +30,7 @@ from litellm.harness.options import (
|
|||
CodexOptions,
|
||||
DeepAgentsOptions,
|
||||
OpenCodeOptions,
|
||||
ToolLoopOptions,
|
||||
)
|
||||
from litellm.harness.runtime import (
|
||||
AsyncEventStream,
|
||||
|
|
@ -86,6 +87,7 @@ __all__ = (
|
|||
"StateIncompatible",
|
||||
"Text",
|
||||
"ToolCall",
|
||||
"ToolLoopOptions",
|
||||
"ToolResult",
|
||||
"Usage",
|
||||
"aagent",
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
"""Handlers run a harness config: CLI runtimes as subprocesses, Deep Agents in-process."""
|
||||
"""Handlers run a harness config: CLI runtimes and in-process harnesses."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
@ -29,6 +29,10 @@ def get_harness_handler(config: BaseHarnessConfig) -> BaseHarnessHandler:
|
|||
from litellm.harness.handlers.deepagents_handler import DeepAgentsHandler
|
||||
|
||||
return DeepAgentsHandler(config)
|
||||
if config.harness is Harness.TOOL_LOOP:
|
||||
from litellm.harness.handlers.tool_loop_handler import ToolLoopHandler
|
||||
|
||||
return ToolLoopHandler(config)
|
||||
raise HarnessError(f"No handler for Harness.{config.harness.name}")
|
||||
|
||||
|
||||
|
|
|
|||
345
litellm/harness/handlers/tool_loop_handler.py
Normal file
345
litellm/harness/handlers/tool_loop_handler.py
Normal file
|
|
@ -0,0 +1,345 @@
|
|||
"""In-process tool-calling loop handler."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import inspect
|
||||
import json
|
||||
import math
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol, TypeAlias
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.harness.context import SessionContext
|
||||
from litellm.harness.errors import CapabilityUnsupported
|
||||
from litellm.harness.handlers.base import BaseHarnessHandler
|
||||
from litellm.harness.types import Approval, Event, Reasoning, Text, ToolCall, ToolResult
|
||||
from litellm.llms.base_llm.harness.transformation import HarnessTurnError
|
||||
from litellm.llms.tool_loop.harness.transformation import (
|
||||
TOOL_LOOP_MAX_MODEL_CALLS,
|
||||
FunctionTool,
|
||||
ToolLoopHarnessConfig,
|
||||
completion_kwargs,
|
||||
function_tool,
|
||||
)
|
||||
from litellm.types.completion import ChatCompletionMessageParam
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionMessageCustomToolCall,
|
||||
ChatCompletionMessageToolCall,
|
||||
ChatCompletionToolParam,
|
||||
ModelResponse,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _FunctionToolCall:
|
||||
id: str
|
||||
name: str
|
||||
arguments: str
|
||||
|
||||
|
||||
AsyncCompletion: TypeAlias = Callable[..., Awaitable[ModelResponse]]
|
||||
_ARGUMENTS_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_MAPPING_ADAPTER: Final = TypeAdapter(Mapping[str, 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]])
|
||||
|
||||
|
||||
class _Usage(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
prompt_tokens: int | None = None
|
||||
completion_tokens: int | None = None
|
||||
|
||||
|
||||
_USAGE_ADAPTER: Final = TypeAdapter(_Usage)
|
||||
|
||||
|
||||
class _AwaitableObject(Protocol):
|
||||
def __await__(self) -> Generator[object, None, object]: ...
|
||||
|
||||
|
||||
async def _await_tool_result(result: _AwaitableObject) -> object:
|
||||
return await result
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ToolOutcome:
|
||||
output: str
|
||||
is_error: bool
|
||||
|
||||
|
||||
def _normalize_tool_call(
|
||||
call: ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall,
|
||||
) -> _FunctionToolCall:
|
||||
if isinstance(call, ChatCompletionMessageToolCall):
|
||||
return _FunctionToolCall(
|
||||
id=call.id,
|
||||
name=call.function.name or "",
|
||||
arguments=call.function.arguments,
|
||||
)
|
||||
return _FunctionToolCall(id=call.id, name=call.custom.name, arguments=call.custom.input)
|
||||
|
||||
|
||||
def _parse_arguments(raw: str) -> tuple[dict[str, object], str | None]: # mutable-ok: event input uses a dict
|
||||
text: Final = raw.strip()
|
||||
try:
|
||||
raw_decoded: object = _JSON_DECODER.raw_decode(text)
|
||||
except json.JSONDecodeError as error:
|
||||
return {}, f"{type(error).__name__}: {error}"
|
||||
parsed, end = _RAW_DECODE_ADAPTER.validate_python(raw_decoded)
|
||||
if text[end:].strip():
|
||||
return {}, "JSONDecodeError: Extra data after tool arguments"
|
||||
try:
|
||||
arguments: Final = _ARGUMENTS_ADAPTER.validate_python(parsed)
|
||||
except ValidationError as error:
|
||||
if not isinstance(parsed, dict):
|
||||
return {}, "ValueError: tool arguments must be a JSON object"
|
||||
return {}, f"{type(error).__name__}: {error}"
|
||||
return arguments, None
|
||||
|
||||
|
||||
def _cost_value(value: object) -> float | None:
|
||||
if isinstance(value, bool):
|
||||
return None
|
||||
if not isinstance(value, str | int | float):
|
||||
return None
|
||||
try:
|
||||
cost: Final = float(value)
|
||||
except (OverflowError, TypeError, ValueError):
|
||||
return None
|
||||
return cost if math.isfinite(cost) else None
|
||||
|
||||
|
||||
def _as_mapping(value: object) -> Mapping[str, object] | None:
|
||||
try:
|
||||
return _MAPPING_ADAPTER.validate_python(value)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _response_cost(response: ModelResponse) -> float:
|
||||
hidden_params_value: Final[object] = getattr(response, "_hidden_params", {})
|
||||
hidden_params: Final = _as_mapping(hidden_params_value)
|
||||
if hidden_params is not None:
|
||||
additional_headers: Final = _as_mapping(hidden_params.get("additional_headers"))
|
||||
if additional_headers is not None:
|
||||
header_cost: Final = _cost_value(additional_headers.get("llm_provider-x-litellm-response-cost"))
|
||||
if header_cost is not None:
|
||||
return header_cost
|
||||
hidden_cost: Final = _cost_value(hidden_params.get("response_cost"))
|
||||
if hidden_cost is not None:
|
||||
return hidden_cost
|
||||
try:
|
||||
calculated_cost: Final = litellm.completion_cost(completion_response=response)
|
||||
except Exception:
|
||||
return 0.0
|
||||
return _cost_value(calculated_cost) or 0.0
|
||||
|
||||
|
||||
async def _approval_error(approval: Approval | None) -> str | None:
|
||||
if approval is None:
|
||||
return None
|
||||
allowed, reason = await approval.wait()
|
||||
return None if allowed else f"denied: {reason}"
|
||||
|
||||
|
||||
async def _tool_outcome(
|
||||
tool: FunctionTool | None,
|
||||
tool_name: str,
|
||||
arguments: dict[str, object],
|
||||
parse_error: str | None,
|
||||
approval_error: str | None,
|
||||
) -> _ToolOutcome:
|
||||
if parse_error is not None:
|
||||
return _ToolOutcome(output=parse_error, is_error=True)
|
||||
if approval_error is not None:
|
||||
return _ToolOutcome(output=approval_error, is_error=True)
|
||||
if tool is None:
|
||||
return _ToolOutcome(output=f"ValueError: unknown tool {tool_name!r}", is_error=True)
|
||||
try:
|
||||
result: Final = await _execute_tool(tool, arguments)
|
||||
output: Final = result if isinstance(result, str) else json.dumps(result, default=str)
|
||||
return _ToolOutcome(output=output, is_error=False)
|
||||
except Exception as error:
|
||||
return _ToolOutcome(output=f"{type(error).__name__}: {error}", is_error=True)
|
||||
|
||||
|
||||
def _record_usage(ctx: SessionContext, response: ModelResponse) -> None:
|
||||
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
|
||||
|
||||
|
||||
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())
|
||||
positional_args: Final = tuple(
|
||||
validated[parameter.name] for parameter in parameters if parameter.kind is inspect.Parameter.POSITIONAL_ONLY
|
||||
)
|
||||
keyword_args: Final[dict[str, object]] = { # mutable-ok: tool calls need keyword arguments
|
||||
parameter.name: validated[parameter.name]
|
||||
for parameter in parameters
|
||||
if parameter.kind is not inspect.Parameter.POSITIONAL_ONLY
|
||||
}
|
||||
if inspect.iscoroutinefunction(tool.fn):
|
||||
async_result: Final[object] = tool.fn(*positional_args, **keyword_args)
|
||||
if inspect.isawaitable(async_result):
|
||||
return await _await_tool_result(async_result)
|
||||
return async_result
|
||||
sync_result: Final[object] = await asyncio.to_thread(tool.fn, *positional_args, **keyword_args)
|
||||
if inspect.isawaitable(sync_result):
|
||||
return await _await_tool_result(sync_result)
|
||||
return sync_result
|
||||
|
||||
|
||||
async def _default_acompletion(**kwargs: object) -> ModelResponse: # kwargs-ok: provider-specific completion options
|
||||
response: Final[object] = await litellm.acompletion(**kwargs)
|
||||
return _MODEL_RESPONSE_ADAPTER.validate_python(response)
|
||||
|
||||
|
||||
class ToolLoopHandler(BaseHarnessHandler):
|
||||
def __init__(
|
||||
self,
|
||||
config: ToolLoopHarnessConfig,
|
||||
acompletion: AsyncCompletion | None = None,
|
||||
) -> None:
|
||||
super().__init__(config) # pyright: ignore[reportUnknownMemberType] # base handler config is unparameterized
|
||||
self._config = config
|
||||
self._acompletion = acompletion if acompletion is not None else _default_acompletion
|
||||
self._messages: tuple[ChatCompletionMessageParam, ...] = ()
|
||||
self._tools: Mapping[str, FunctionTool] = MappingProxyType({})
|
||||
self._tool_specs: tuple[ChatCompletionToolParam, ...] = ()
|
||||
self._completion_kwargs: Mapping[str, object] = MappingProxyType({})
|
||||
|
||||
async def start(self, ctx: SessionContext) -> None:
|
||||
self._config.validate_environment(ctx)
|
||||
tools: Final = tuple(function_tool(fn) for fn in ctx.tools)
|
||||
if len({tool.name for tool in tools}) != len(tools):
|
||||
raise ValueError("Harness.TOOL_LOOP tool names must be unique")
|
||||
self._tools = MappingProxyType({tool.name: tool for tool in tools})
|
||||
self._tool_specs = tuple(tool.spec for tool in tools)
|
||||
self._completion_kwargs = MappingProxyType(completion_kwargs(ctx))
|
||||
if not self._messages and ctx.instructions:
|
||||
self._messages = ({"role": "system", "content": ctx.instructions},)
|
||||
|
||||
def native_session_id(self) -> str | None:
|
||||
return None
|
||||
|
||||
async def resume(self, ctx: SessionContext, native_session_id: str) -> None:
|
||||
raise CapabilityUnsupported("Harness.TOOL_LOOP does not support resume")
|
||||
|
||||
async def history(
|
||||
self, ctx: SessionContext
|
||||
) -> list[dict[str, object]]: # mutable-ok: public API returns copied message dictionaries
|
||||
history_object: Final[object] = copy.deepcopy(list(self._messages))
|
||||
return _HISTORY_ADAPTER.validate_python(history_object) # pyright: ignore[reportIncompatibleMethodOverride] # base history uses Any
|
||||
|
||||
async def stop(self, ctx: SessionContext) -> None:
|
||||
return None
|
||||
|
||||
async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]:
|
||||
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):
|
||||
messages: list[ChatCompletionMessageParam] = copy.deepcopy( # mutable-ok: acompletion takes list messages
|
||||
list(self._messages)
|
||||
)
|
||||
tool_specs: list[ChatCompletionToolParam] = copy.deepcopy( # mutable-ok: acompletion takes tool list
|
||||
list(self._tool_specs)
|
||||
)
|
||||
request_kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
|
||||
key: value for key, value in self._completion_kwargs.items() if key not in {"messages", "tools"}
|
||||
}
|
||||
kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
|
||||
**request_kwargs,
|
||||
"messages": messages,
|
||||
**({"tools": tool_specs} if tool_specs else {}),
|
||||
}
|
||||
response = await self._acompletion(**kwargs)
|
||||
_record_usage(ctx, response)
|
||||
message = response.choices[0].message
|
||||
reasoning_value: object = getattr(message, "reasoning_content", None)
|
||||
reasoning = reasoning_value if isinstance(reasoning_value, str) else None
|
||||
content = message.content
|
||||
if reasoning:
|
||||
yield Reasoning(reasoning)
|
||||
if content:
|
||||
yield Text(content)
|
||||
tool_calls = message.tool_calls or ()
|
||||
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,
|
||||
}
|
||||
self._messages = (*self._messages, final_message)
|
||||
return
|
||||
normalized_calls = tuple(_normalize_tool_call(call) for call in tool_calls)
|
||||
assistant_message: ChatCompletionMessageParam = {
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": call.id,
|
||||
"type": "function",
|
||||
"function": {"name": call.name, "arguments": call.arguments},
|
||||
}
|
||||
for call in normalized_calls
|
||||
],
|
||||
}
|
||||
self._messages = (*self._messages, assistant_message)
|
||||
for call in normalized_calls:
|
||||
arguments, parse_error = _parse_arguments(call.arguments)
|
||||
yield ToolCall(
|
||||
id=call.id,
|
||||
name=call.name,
|
||||
native_name=call.name,
|
||||
input=arguments,
|
||||
builtin=False,
|
||||
)
|
||||
approval = (
|
||||
Approval(tool=call.name, input=arguments)
|
||||
if parse_error is None and ctx.permissions == "ask"
|
||||
else None
|
||||
)
|
||||
if approval is not None:
|
||||
yield approval
|
||||
approval_error = await _approval_error(approval)
|
||||
tool = self._tools.get(call.name)
|
||||
outcome = await _tool_outcome(
|
||||
tool,
|
||||
call.name,
|
||||
arguments,
|
||||
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)
|
||||
raise HarnessTurnError(f"Harness.TOOL_LOOP exceeded {TOOL_LOOP_MAX_MODEL_CALLS} model calls in one turn")
|
||||
|
|
@ -34,4 +34,9 @@ class DeepAgentsOptions:
|
|||
recursion_limit: int | None = None
|
||||
|
||||
|
||||
HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions
|
||||
@dataclass(frozen=True)
|
||||
class ToolLoopOptions:
|
||||
completion_kwargs: Mapping[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions | ToolLoopOptions
|
||||
|
|
|
|||
|
|
@ -411,7 +411,7 @@ def agent(
|
|||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Result | EventStream:
|
||||
"""Run an agent harness (Claude Code, Codex, OpenCode, Deep Agents) on one prompt.
|
||||
"""Run an agent harness (Claude Code, Codex, OpenCode, Deep Agents, Tool Loop) on one prompt.
|
||||
|
||||
Returns a Result. With stream=True it returns an iterator of events instead.
|
||||
Prefix the model with `litellm_proxy/` to route every model call through your
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ class Harness(Enum):
|
|||
CODEX = "codex"
|
||||
OPENCODE = "opencode"
|
||||
DEEPAGENTS = "deepagents"
|
||||
TOOL_LOOP = "tool_loop"
|
||||
|
||||
|
||||
def require_harness(harness: object) -> Harness:
|
||||
|
|
|
|||
|
|
@ -5,17 +5,33 @@ from __future__ import annotations
|
|||
import itertools
|
||||
import json
|
||||
import os
|
||||
import typing
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, TypeAlias
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from litellm.harness.context import SessionContext
|
||||
|
||||
# A decoded JSON document: what json.loads / model_json_schema() produce.
|
||||
JSONValue: TypeAlias = "dict[str, JSONValue] | list[JSONValue] | str | int | float | bool | None"
|
||||
|
||||
SKILL_MANIFEST: Final = "SKILL.md"
|
||||
_JSON_DECODER: Final = json.JSONDecoder()
|
||||
_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_RAW_DECODE_ADAPTER: Final = TypeAdapter(tuple[object, int])
|
||||
|
||||
|
||||
def gateway_headers(ctx: SessionContext) -> dict[str, str]: # mutable-ok: acompletion(extra_headers=) requires dict
|
||||
metadata_json: Final = json.dumps(dict(ctx.metadata), default=str) if ctx.metadata else None
|
||||
return {
|
||||
"x-litellm-tags": f"harness,{ctx.harness.value}",
|
||||
**({"x-litellm-spend-logs-metadata": metadata_json} if metadata_json is not None else {}),
|
||||
}
|
||||
|
||||
|
||||
def normalize_tool_name(native_name: str, mapping: Mapping[str, str]) -> str:
|
||||
|
|
@ -35,17 +51,18 @@ def last_json_object(text: str) -> str | None:
|
|||
index = text.find("{")
|
||||
while index != -1:
|
||||
try:
|
||||
obj, end = _JSON_DECODER.raw_decode(text, index)
|
||||
raw_decoded: object = _JSON_DECODER.raw_decode(text, index)
|
||||
except json.JSONDecodeError:
|
||||
index = text.find("{", index + 1)
|
||||
continue
|
||||
obj, end = _RAW_DECODE_ADAPTER.validate_python(raw_decoded)
|
||||
if isinstance(obj, dict):
|
||||
last = json.dumps(obj)
|
||||
index = text.find("{", end)
|
||||
return last
|
||||
|
||||
|
||||
def structured_output_instruction(schema: Mapping[str, Any]) -> str:
|
||||
def structured_output_instruction(schema: Mapping[str, object]) -> str:
|
||||
return (
|
||||
"When you have finished, your final message must be a single JSON object that "
|
||||
"matches this JSON schema, with no other text before or after it:\n"
|
||||
|
|
@ -82,16 +99,15 @@ def strict_json_schema(schema: JSONValue, depth: int = 0) -> JSONValue:
|
|||
return result
|
||||
|
||||
|
||||
def decode_json_line(line: bytes | str) -> Mapping[str, Any] | None:
|
||||
def decode_json_line(line: bytes | str) -> Mapping[str, object] | None:
|
||||
"""One JSONL line as a dict, or None for blank / non-JSON / non-object lines."""
|
||||
text = line.strip()
|
||||
if not text:
|
||||
return None
|
||||
try:
|
||||
obj = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
return _JSON_OBJECT_ADAPTER.validate_json(text)
|
||||
except ValidationError:
|
||||
return None
|
||||
return obj if isinstance(obj, dict) else None
|
||||
|
||||
|
||||
def stderr_tail_text(stderr_tail: Sequence[str]) -> str:
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm.harness.types import (
|
|||
ToolResult,
|
||||
)
|
||||
from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig
|
||||
from litellm.llms.base_llm.harness.utils import gateway_headers
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.harness.context import SessionContext
|
||||
|
|
@ -70,18 +71,6 @@ APPROVAL_TOOLS: Final = WRITE_TOOLS | EXECUTE_TOOLS
|
|||
_APPROVAL_DECISIONS: Final = ("approve", "reject")
|
||||
|
||||
|
||||
def gateway_headers(
|
||||
ctx: SessionContext,
|
||||
) -> dict[str, str]: # mutable-ok: ChatLiteLLM.extra_headers is a pydantic dict field
|
||||
"""Same attribution headers the session endpoint adds for CLI harnesses."""
|
||||
metadata = ctx.metadata
|
||||
metadata_json = json.dumps(dict(metadata), default=str) if metadata else None # mutable-ok: for json.dumps
|
||||
metadata_header = (("x-litellm-spend-logs-metadata", metadata_json),) if metadata_json is not None else ()
|
||||
return dict( # mutable-ok: ChatLiteLLM.extra_headers is a pydantic dict field
|
||||
(("x-litellm-tags", f"harness,{ctx.harness.value}"), *metadata_header)
|
||||
)
|
||||
|
||||
|
||||
def chat_model_kwargs(
|
||||
ctx: SessionContext,
|
||||
) -> dict[str, Any]: # mutable-ok: ChatLiteLLM constructor kwargs, splatted as **kwargs
|
||||
|
|
|
|||
1
litellm/llms/tool_loop/__init__.py
Normal file
1
litellm/llms/tool_loop/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
"""In-process tool loop harness."""
|
||||
1
litellm/llms/tool_loop/harness/__init__.py
Normal file
1
litellm/llms/tool_loop/harness/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
"""Tool Loop harness configuration."""
|
||||
119
litellm/llms/tool_loop/harness/transformation.py
Normal file
119
litellm/llms/tool_loop/harness/transformation.py
Normal file
|
|
@ -0,0 +1,119 @@
|
|||
"""Configuration and tool-schema helpers for the in-process Tool Loop harness."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, create_model
|
||||
|
||||
from litellm.harness.context import SessionContext
|
||||
from litellm.harness.options import ToolLoopOptions
|
||||
from litellm.harness.types import Capabilities, Harness
|
||||
from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig
|
||||
from litellm.llms.base_llm.harness.utils import gateway_headers
|
||||
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)
|
||||
_MODEL_FACTORY: Final[Callable[..., type[BaseModel]]] = create_model
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class FunctionTool:
|
||||
name: str
|
||||
fn: Callable[..., object]
|
||||
args_model: type[BaseModel]
|
||||
spec: ChatCompletionToolParam
|
||||
|
||||
|
||||
def _field_definition(
|
||||
parameter: inspect.Parameter,
|
||||
annotations: Mapping[str, object],
|
||||
) -> tuple[object, object]:
|
||||
annotation: Final = annotations.get(parameter.name, object)
|
||||
raw_default: Final[object] = parameter.default # pyright: ignore[reportAny] # inspect exposes defaults as Any
|
||||
if raw_default is inspect.Parameter.empty:
|
||||
return annotation, ...
|
||||
default: Final = _OBJECT_ADAPTER.validate_python(raw_default)
|
||||
return annotation, default
|
||||
|
||||
|
||||
def function_tool(fn: Callable[..., object]) -> FunctionTool:
|
||||
signature: Final = inspect.signature(fn)
|
||||
parameters: Final = tuple(signature.parameters.values())
|
||||
raw_annotations: Final[object] = inspect.get_annotations(fn, eval_str=True)
|
||||
annotations: Final = _ANNOTATIONS_ADAPTER.validate_python(raw_annotations)
|
||||
if any(
|
||||
parameter.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD) for parameter in parameters
|
||||
):
|
||||
raise ValueError(f"Tool {fn.__name__} cannot use variadic parameters")
|
||||
|
||||
fields: Final = MappingProxyType(
|
||||
{parameter.name: _field_definition(parameter, annotations) for parameter in parameters}
|
||||
)
|
||||
args_model: Final[type[BaseModel]] = _MODEL_FACTORY(
|
||||
f"{fn.__name__}_args",
|
||||
__config__=ConfigDict(extra="forbid"),
|
||||
**fields, # pyright: ignore[reportCallIssue, reportArgumentType] # Pydantic creates fields dynamically
|
||||
)
|
||||
spec: Final[ChatCompletionToolParam] = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": fn.__name__,
|
||||
"description": inspect.getdoc(fn) or "",
|
||||
"parameters": args_model.model_json_schema(),
|
||||
},
|
||||
}
|
||||
return FunctionTool(name=fn.__name__, fn=fn, args_model=args_model, spec=spec)
|
||||
|
||||
|
||||
def _routing_kwargs(ctx: SessionContext) -> Mapping[str, object]:
|
||||
if not ctx.model:
|
||||
raise ValueError("Harness.TOOL_LOOP needs model=")
|
||||
if ctx.gateway is not None:
|
||||
return {
|
||||
"model": f"litellm_proxy/{ctx.model}",
|
||||
"api_base": ctx.gateway.api_base,
|
||||
"api_key": ctx.gateway.api_key,
|
||||
"extra_headers": gateway_headers(ctx),
|
||||
}
|
||||
return {
|
||||
"model": ctx.model,
|
||||
**({"api_key": ctx.api_key} if ctx.api_key is not None else {}),
|
||||
**({"api_base": ctx.api_base} if ctx.api_base is not None else {}),
|
||||
}
|
||||
|
||||
|
||||
def completion_kwargs(ctx: SessionContext) -> Mapping[str, object]:
|
||||
options: Final = ToolLoopHarnessConfig().get_options(ctx)
|
||||
routing: Final = _routing_kwargs(ctx)
|
||||
kwargs: Final[Mapping[str, object]] = MappingProxyType({**options.completion_kwargs, **routing})
|
||||
if ctx.output is None:
|
||||
return kwargs
|
||||
return {**kwargs, "response_format": ctx.output}
|
||||
|
||||
|
||||
class ToolLoopHarnessConfig(BaseHarnessConfig[ToolLoopOptions]):
|
||||
harness = Harness.TOOL_LOOP
|
||||
options_type = ToolLoopOptions
|
||||
uses_model_endpoint = False
|
||||
capabilities = Capabilities(
|
||||
structured_output=True,
|
||||
tool_approval=True,
|
||||
tool_filtering=False,
|
||||
history=True,
|
||||
custom_tools=True,
|
||||
skills=False,
|
||||
resume=False,
|
||||
permission_modes=frozenset({"ask", "full"}),
|
||||
)
|
||||
|
||||
def validate_environment(self, ctx: SessionContext) -> None:
|
||||
self.get_options(ctx)
|
||||
if not ctx.model:
|
||||
raise ValueError("Harness.TOOL_LOOP needs model=")
|
||||
|
|
@ -9909,7 +9909,7 @@ class ProviderConfigManager:
|
|||
@staticmethod
|
||||
def get_provider_harness_config(harness: Harness) -> BaseHarnessConfig | None:
|
||||
"""
|
||||
Get the agent-harness configuration (Claude Code, Codex, OpenCode, Deep Agents).
|
||||
Get the agent-harness configuration (Claude Code, Codex, OpenCode, Deep Agents, Tool Loop).
|
||||
"""
|
||||
from litellm.harness.types import Harness as _Harness
|
||||
|
||||
|
|
@ -9935,6 +9935,10 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return DeepAgentsHarnessConfig()
|
||||
if harness == _Harness.TOOL_LOOP:
|
||||
from litellm.llms.tool_loop.harness.transformation import ToolLoopHarnessConfig
|
||||
|
||||
return ToolLoopHarnessConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ EXCLUDED_FOLDERS = {
|
|||
"codex",
|
||||
"opencode",
|
||||
"deepagents",
|
||||
"tool_loop",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
476
tests/unit/harness/handlers/test_tool_loop_handler.py
Normal file
476
tests/unit/harness/handlers/test_tool_loop_handler.py
Normal file
|
|
@ -0,0 +1,476 @@
|
|||
"""Tests for the in-process Tool Loop handler."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm import sandbox
|
||||
from litellm.harness.context import GatewayTarget, SessionContext
|
||||
from litellm.harness.handlers.tool_loop_handler import ToolLoopHandler
|
||||
from litellm.harness.options import ToolLoopOptions
|
||||
from litellm.harness.types import Approval, Event, Harness, Text, ToolCall, ToolResult
|
||||
from litellm.llms.base_llm.harness.transformation import HarnessTurnError
|
||||
from litellm.llms.tool_loop.harness.transformation import ToolLoopHarnessConfig
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
class ScriptedCompletion:
|
||||
def __init__(self, responses: tuple[ModelResponse, ...]) -> None:
|
||||
self.responses: Iterator[ModelResponse] = iter(responses)
|
||||
self.calls: list[dict[str, object]] = [] # mutable-ok: captures injected completion requests
|
||||
|
||||
async def __call__(self, **kwargs: object) -> ModelResponse:
|
||||
self.calls.append(dict(kwargs))
|
||||
return next(self.responses)
|
||||
|
||||
|
||||
def model_response(
|
||||
*,
|
||||
content: str | None = None,
|
||||
tool_calls: tuple[dict[str, object], ...] = (),
|
||||
prompt_tokens: int = 0,
|
||||
completion_tokens: int = 0,
|
||||
hidden_params: dict[str, object] | None = None,
|
||||
) -> ModelResponse:
|
||||
response: Final = ModelResponse(
|
||||
model="gpt-test",
|
||||
choices=[
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
"tool_calls": list(tool_calls) if tool_calls else None,
|
||||
}
|
||||
}
|
||||
],
|
||||
usage={
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
"total_tokens": prompt_tokens + completion_tokens,
|
||||
},
|
||||
)
|
||||
if hidden_params is not None:
|
||||
response._hidden_params = hidden_params
|
||||
return response
|
||||
|
||||
|
||||
def function_call(name: str, arguments: str, call_id: str = "call-1") -> dict[str, object]:
|
||||
return {
|
||||
"id": call_id,
|
||||
"type": "function",
|
||||
"function": {"name": name, "arguments": arguments},
|
||||
}
|
||||
|
||||
|
||||
def make_context(
|
||||
tmp_path: Path,
|
||||
*,
|
||||
model: str | None = "gpt-4o-mini",
|
||||
gateway: GatewayTarget | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: tuple[Callable[..., object], ...] = (),
|
||||
permissions: Literal["ask", "full"] = "full",
|
||||
output: type[BaseModel] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: ToolLoopOptions | None = None,
|
||||
) -> SessionContext:
|
||||
return SessionContext(
|
||||
harness=Harness.TOOL_LOOP,
|
||||
sandbox=sandbox.local(tmp_path),
|
||||
session_id="tool-loop-session",
|
||||
model=model,
|
||||
gateway=gateway,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
instructions=instructions,
|
||||
tools=tools,
|
||||
permissions=permissions,
|
||||
output=output,
|
||||
metadata={} if metadata is None else metadata,
|
||||
options=options,
|
||||
)
|
||||
|
||||
|
||||
def make_handler(completion: ScriptedCompletion) -> ToolLoopHandler:
|
||||
return ToolLoopHandler(ToolLoopHarnessConfig(), acompletion=completion)
|
||||
|
||||
|
||||
async def run_turn(
|
||||
handler: ToolLoopHandler,
|
||||
ctx: SessionContext,
|
||||
prompt: str,
|
||||
allow: bool = True,
|
||||
) -> tuple[Event, ...]:
|
||||
events: list[Event] = [] # mutable-ok: gathers this async turn for assertions
|
||||
async for event in handler.turn(ctx, prompt):
|
||||
events.append(event)
|
||||
if isinstance(event, Approval):
|
||||
event.allow() if allow else event.deny("not approved")
|
||||
return tuple(events)
|
||||
|
||||
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
|
||||
|
||||
def duplicate_add() -> Callable[..., object]:
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
|
||||
return add
|
||||
|
||||
|
||||
async def test_tool_round_trip_appends_assistant_and_tool_messages(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)),
|
||||
model_response(content="The sum is 5"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, instructions="Use tools when needed", tools=(add,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add two and three")
|
||||
|
||||
assert events == (
|
||||
ToolCall(
|
||||
id="call-1",
|
||||
name="add",
|
||||
native_name="add",
|
||||
input={"a": 2, "b": 3},
|
||||
builtin=False,
|
||||
),
|
||||
ToolResult(id="call-1", output="5", is_error=False),
|
||||
Text(delta="The sum is 5"),
|
||||
)
|
||||
assert ctx.final_text == "The sum is 5"
|
||||
assert completion.calls[1]["messages"] == [
|
||||
{"role": "system", "content": "Use tools when needed"},
|
||||
{"role": "user", "content": "Add two and three"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call-1",
|
||||
"type": "function",
|
||||
"function": {"name": "add", "arguments": '{"a": 2, "b": 3}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call-1", "content": "5"},
|
||||
]
|
||||
|
||||
|
||||
async def test_multiple_tool_calls_run_in_order(tmp_path: Path) -> None:
|
||||
values: list[int] = [] # mutable-ok: records the order of calls from the injected model response
|
||||
|
||||
def record(value: int) -> int:
|
||||
values.append(value)
|
||||
return value
|
||||
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(
|
||||
tool_calls=(
|
||||
function_call("record", '{"value": 1}', "call-1"),
|
||||
function_call("record", '{"value": 2}', "call-2"),
|
||||
)
|
||||
),
|
||||
model_response(content="Recorded"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(record,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Record both")
|
||||
|
||||
assert values == [1, 2]
|
||||
assert tuple(event for event in events if isinstance(event, ToolResult)) == (
|
||||
ToolResult(id="call-1", output="1", is_error=False),
|
||||
ToolResult(id="call-2", output="2", is_error=False),
|
||||
)
|
||||
|
||||
|
||||
async def test_tool_exception_is_returned_to_model_and_loop_continues(tmp_path: Path) -> None:
|
||||
def fail() -> str:
|
||||
raise RuntimeError("tool failed")
|
||||
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("fail", "{}"),)),
|
||||
model_response(content="Recovered"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(fail,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Run fail")
|
||||
|
||||
assert ToolResult(id="call-1", output="RuntimeError: tool failed", is_error=True) in events
|
||||
assert completion.calls[1]["messages"][-1] == {
|
||||
"role": "tool",
|
||||
"tool_call_id": "call-1",
|
||||
"content": "RuntimeError: tool failed",
|
||||
}
|
||||
assert ctx.final_text == "Recovered"
|
||||
|
||||
|
||||
async def test_unknown_tool_is_returned_and_tools_are_omitted_when_empty(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("missing", "{}"),)),
|
||||
model_response(content="Unknown tool handled"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path)
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Call a missing tool")
|
||||
|
||||
assert "tools" not in completion.calls[0]
|
||||
assert (
|
||||
ToolResult(
|
||||
id="call-1",
|
||||
output="ValueError: unknown tool 'missing'",
|
||||
is_error=True,
|
||||
)
|
||||
in events
|
||||
)
|
||||
assert ctx.final_text == "Unknown tool handled"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"arguments",
|
||||
['{"a": 2}', '{"a": 2, "b": 3, "extra": 4}'],
|
||||
)
|
||||
async def test_invalid_tool_arguments_return_validation_error(
|
||||
tmp_path: Path,
|
||||
arguments: str,
|
||||
) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", arguments),)),
|
||||
model_response(content="Arguments were invalid"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add")
|
||||
result: Final = next(event for event in events if isinstance(event, ToolResult))
|
||||
|
||||
assert result.is_error
|
||||
assert result.output.startswith("ValidationError:")
|
||||
assert completion.calls[1]["messages"][-1]["content"] == result.output
|
||||
|
||||
|
||||
async def test_malformed_tool_arguments_return_error_without_calling_tool(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", "{"),)),
|
||||
model_response(content="Malformed arguments"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add")
|
||||
result: Final = next(event for event in events if isinstance(event, ToolResult))
|
||||
|
||||
assert result.is_error
|
||||
assert result.output.startswith("JSONDecodeError:")
|
||||
assert completion.calls[1]["messages"][-1]["content"] == result.output
|
||||
|
||||
|
||||
async def test_ask_permission_denial_skips_tool_and_returns_reason(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)),
|
||||
model_response(content="Denied"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,), permissions="ask")
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add", allow=False)
|
||||
|
||||
assert any(isinstance(event, Approval) for event in events)
|
||||
assert ToolResult(id="call-1", output="denied: not approved", is_error=True) in events
|
||||
assert completion.calls[1]["messages"][-1]["content"] == "denied: not approved"
|
||||
|
||||
|
||||
async def test_ask_permission_allow_runs_tool(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)),
|
||||
model_response(content="Allowed"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,), permissions="ask")
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add")
|
||||
|
||||
assert any(isinstance(event, Approval) for event in events)
|
||||
assert ToolResult(id="call-1", output="5", is_error=False) in events
|
||||
|
||||
|
||||
async def test_async_tool_is_awaited(tmp_path: Path) -> None:
|
||||
async def multiply(a: int, b: int) -> int:
|
||||
return a * b
|
||||
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("multiply", '{"a": 3, "b": 4}'),)),
|
||||
model_response(content="12"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(multiply,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Multiply")
|
||||
|
||||
assert ToolResult(id="call-1", output="12", is_error=False) in events
|
||||
|
||||
|
||||
class Answer(BaseModel):
|
||||
value: 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)
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "Return a value")
|
||||
|
||||
assert completion.calls[0]["response_format"] is Answer
|
||||
assert ctx.output_json == '{"value": 7}'
|
||||
assert ctx.final_text == '{"value": 7}'
|
||||
|
||||
|
||||
async def test_gateway_routing_uses_proxy_model_and_tool_loop_tag(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion((model_response(content="done"),))
|
||||
gateway: Final = GatewayTarget(api_base="http://gateway", api_key="sk-virtual")
|
||||
ctx: Final = make_context(tmp_path, gateway=gateway, metadata={"team": "test"})
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "Hi")
|
||||
|
||||
assert completion.calls[0]["model"] == "litellm_proxy/gpt-4o-mini"
|
||||
assert completion.calls[0]["api_base"] == "http://gateway"
|
||||
assert completion.calls[0]["api_key"] == "sk-virtual"
|
||||
headers: Final = completion.calls[0]["extra_headers"]
|
||||
assert isinstance(headers, dict)
|
||||
assert headers["x-litellm-tags"] == "harness,tool_loop"
|
||||
assert '"team": "test"' in headers["x-litellm-spend-logs-metadata"]
|
||||
|
||||
|
||||
async def test_usage_and_cost_accumulate_across_model_calls(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(
|
||||
tool_calls=(function_call("add", '{"a": 2, "b": 3}'),),
|
||||
prompt_tokens=10,
|
||||
completion_tokens=4,
|
||||
hidden_params={
|
||||
"additional_headers": {"llm_provider-x-litellm-response-cost": "0.4"},
|
||||
"response_cost": 0.1,
|
||||
},
|
||||
),
|
||||
model_response(
|
||||
content="done",
|
||||
prompt_tokens=20,
|
||||
completion_tokens=5,
|
||||
hidden_params={"response_cost": 0.2},
|
||||
),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "Add")
|
||||
|
||||
assert ctx.calls == 2
|
||||
assert ctx.input_tokens == 30
|
||||
assert ctx.output_tokens == 9
|
||||
assert ctx.cost == pytest.approx(0.6)
|
||||
|
||||
|
||||
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")
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
await run_turn(handler, ctx, "first prompt")
|
||||
await handler.stop(ctx)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "second prompt")
|
||||
|
||||
assert completion.calls[1]["messages"] == [
|
||||
{"role": "system", "content": "Keep answers concise"},
|
||||
{"role": "user", "content": "first prompt"},
|
||||
{"role": "assistant", "content": "first"},
|
||||
{"role": "user", "content": "second prompt"},
|
||||
]
|
||||
history: Final = await handler.history(ctx)
|
||||
history[0]["content"] = "changed"
|
||||
assert (await handler.history(ctx))[0]["content"] == "Keep answers concise"
|
||||
|
||||
|
||||
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(()))
|
||||
|
||||
with pytest.raises(ValueError, match="tool names must be unique"):
|
||||
await handler.start(ctx)
|
||||
|
||||
|
||||
async def test_model_call_limit_raises_harness_turn_error(tmp_path: Path) -> None:
|
||||
repeating_response: Final = model_response(tool_calls=(function_call("ping", "{}"),))
|
||||
completion: Final = ScriptedCompletion((repeating_response,) * 100)
|
||||
|
||||
def ping() -> str:
|
||||
return "pong"
|
||||
|
||||
ctx: Final = make_context(tmp_path, tools=(ping,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
with pytest.raises(HarnessTurnError, match="exceeded 100 model calls"):
|
||||
await run_turn(handler, ctx, "Ping repeatedly")
|
||||
|
||||
|
||||
async def test_public_aagent_uses_tool_loop_with_mock_response(tmp_path: Path) -> None:
|
||||
result: Final = await litellm.aagent(
|
||||
Harness.TOOL_LOOP,
|
||||
"Say done",
|
||||
sandbox=sandbox.local(tmp_path),
|
||||
model="gpt-4o-mini",
|
||||
options=ToolLoopOptions(completion_kwargs={"mock_response": "done"}),
|
||||
)
|
||||
|
||||
assert result.text == "done"
|
||||
assert result.stop_reason == "done"
|
||||
|
|
@ -35,6 +35,7 @@ PUBLIC_NAMES = [
|
|||
"CodexOptions",
|
||||
"OpenCodeOptions",
|
||||
"DeepAgentsOptions",
|
||||
"ToolLoopOptions",
|
||||
"HarnessError",
|
||||
"CapabilityUnsupported",
|
||||
"OptionsMismatch",
|
||||
|
|
@ -89,7 +90,9 @@ def test_adapter_registry_paths_cover_every_harness():
|
|||
def test_litellm_agent_is_top_level_and_lazy():
|
||||
code = (
|
||||
"import sys, litellm; assert 'litellm.harness' not in sys.modules; "
|
||||
"assert litellm.agent is litellm.harness.agent; assert litellm.Harness.CODEX.value == 'codex'"
|
||||
"assert litellm.agent is litellm.harness.agent; "
|
||||
"assert litellm.Harness.CODEX.value == 'codex'; "
|
||||
"assert litellm.ToolLoopOptions is litellm.harness.ToolLoopOptions"
|
||||
)
|
||||
out = run_child_interpreter(code, timeout=120)
|
||||
assert out.returncode == 0, out.stderr
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from litellm.harness.types import (
|
|||
|
||||
def test_harness_is_plain_enum():
|
||||
assert Harness.CODEX.value == "codex"
|
||||
assert Harness.TOOL_LOOP.value == "tool_loop"
|
||||
assert not isinstance(Harness.CODEX, str)
|
||||
|
||||
|
||||
|
|
|
|||
0
tests/unit/llms/tool_loop/__init__.py
Normal file
0
tests/unit/llms/tool_loop/__init__.py
Normal file
0
tests/unit/llms/tool_loop/harness/__init__.py
Normal file
0
tests/unit/llms/tool_loop/harness/__init__.py
Normal file
154
tests/unit/llms/tool_loop/harness/test_transformation.py
Normal file
154
tests/unit/llms/tool_loop/harness/test_transformation.py
Normal file
|
|
@ -0,0 +1,154 @@
|
|||
"""Tests for Tool Loop schemas and completion routing."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm import sandbox
|
||||
from litellm.harness.context import GatewayTarget, SessionContext
|
||||
from litellm.harness.options import ToolLoopOptions
|
||||
from litellm.harness.types import Harness
|
||||
from litellm.llms.tool_loop.harness.transformation import (
|
||||
ToolLoopHarnessConfig,
|
||||
completion_kwargs,
|
||||
function_tool,
|
||||
)
|
||||
|
||||
|
||||
def make_context(
|
||||
tmp_path: Path,
|
||||
*,
|
||||
model: str | None = "anthropic/claude",
|
||||
gateway: GatewayTarget | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
options: ToolLoopOptions | None = None,
|
||||
output: type[BaseModel] | None = None,
|
||||
) -> SessionContext:
|
||||
return SessionContext(
|
||||
harness=Harness.TOOL_LOOP,
|
||||
sandbox=sandbox.local(tmp_path),
|
||||
session_id="tool-loop-transform",
|
||||
model=model,
|
||||
gateway=gateway,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
options=options,
|
||||
output=output,
|
||||
)
|
||||
|
||||
|
||||
def search(
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
state: Literal["open", "closed"] = "open",
|
||||
) -> str:
|
||||
"""Search records."""
|
||||
return query
|
||||
|
||||
|
||||
def test_function_tool_schema_has_required_defaulted_and_literal_fields() -> None:
|
||||
specification: Final = function_tool(search).spec
|
||||
schema: Final = specification["function"]["parameters"]
|
||||
|
||||
assert schema["required"] == ["query"]
|
||||
assert schema["properties"]["query"] == {"title": "Query", "type": "string"}
|
||||
assert schema["properties"]["limit"] == {"default": 5, "title": "Limit", "type": "integer"}
|
||||
assert schema["properties"]["state"] == {
|
||||
"default": "open",
|
||||
"enum": ["open", "closed"],
|
||||
"title": "State",
|
||||
"type": "string",
|
||||
}
|
||||
assert specification["function"]["description"] == "Search records."
|
||||
|
||||
|
||||
def test_function_schema_rejects_unknown_arguments() -> None:
|
||||
with pytest.raises(ValueError, match="Extra inputs are not permitted"):
|
||||
function_tool(search).args_model.model_validate({"query": "owner", "unknown": "value"})
|
||||
|
||||
|
||||
def variadic_positional(*args: int) -> int:
|
||||
return len(args)
|
||||
|
||||
|
||||
def variadic_keyword(**kwargs: int) -> int:
|
||||
return len(kwargs)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("fn", [variadic_positional, variadic_keyword])
|
||||
def test_variadic_tools_are_rejected(fn: Callable[..., object]) -> None:
|
||||
with pytest.raises(ValueError, match="variadic parameters"):
|
||||
function_tool(fn)
|
||||
|
||||
|
||||
def test_sdk_routing_overrides_completion_kwargs(tmp_path: Path) -> None:
|
||||
ctx: Final = make_context(
|
||||
tmp_path,
|
||||
api_key="provided-key",
|
||||
api_base="https://provider",
|
||||
options=ToolLoopOptions(
|
||||
completion_kwargs={
|
||||
"model": "wrong-model",
|
||||
"api_key": "wrong-key",
|
||||
"api_base": "https://wrong",
|
||||
"temperature": 0.2,
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
kwargs: Final = completion_kwargs(ctx)
|
||||
|
||||
assert kwargs == {
|
||||
"model": "anthropic/claude",
|
||||
"api_key": "provided-key",
|
||||
"api_base": "https://provider",
|
||||
"temperature": 0.2,
|
||||
}
|
||||
|
||||
|
||||
def test_gateway_routing_and_response_format_override_options(tmp_path: Path) -> None:
|
||||
class OutputModel(BaseModel):
|
||||
pass
|
||||
|
||||
gateway: Final = GatewayTarget(api_base="https://gateway", api_key="virtual-key")
|
||||
ctx: Final = make_context(
|
||||
tmp_path,
|
||||
gateway=gateway,
|
||||
options=ToolLoopOptions(
|
||||
completion_kwargs={
|
||||
"model": "wrong-model",
|
||||
"api_key": "wrong-key",
|
||||
"api_base": "https://wrong",
|
||||
"response_format": "wrong-format",
|
||||
}
|
||||
),
|
||||
output=OutputModel,
|
||||
)
|
||||
|
||||
kwargs: Final = completion_kwargs(ctx)
|
||||
|
||||
assert kwargs == {
|
||||
"model": "litellm_proxy/anthropic/claude",
|
||||
"api_base": "https://gateway",
|
||||
"api_key": "virtual-key",
|
||||
"extra_headers": {"x-litellm-tags": "harness,tool_loop"},
|
||||
"response_format": OutputModel,
|
||||
}
|
||||
|
||||
|
||||
def test_configuration_requires_model_and_declares_capabilities(tmp_path: Path) -> None:
|
||||
config: Final = ToolLoopHarnessConfig()
|
||||
assert config.uses_model_endpoint is False
|
||||
assert config.capabilities.structured_output
|
||||
assert config.capabilities.tool_approval
|
||||
assert config.capabilities.history
|
||||
assert config.capabilities.custom_tools
|
||||
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))
|
||||
Loading…
Add table
Reference in a new issue