From b024950353efd55e33082f7f5b267f662478e5bb Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 17:33:01 +0000 Subject: [PATCH] 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> --- litellm/__init__.py | 1 + litellm/harness/__init__.py | 4 +- litellm/harness/handlers/__init__.py | 6 +- litellm/harness/handlers/tool_loop_handler.py | 345 +++++++++++++ litellm/harness/options.py | 7 +- litellm/harness/sync.py | 2 +- litellm/harness/types.py | 1 + litellm/llms/base_llm/harness/utils.py | 30 +- .../llms/deepagents/harness/transformation.py | 13 +- litellm/llms/tool_loop/__init__.py | 1 + litellm/llms/tool_loop/harness/__init__.py | 1 + .../llms/tool_loop/harness/transformation.py | 119 +++++ litellm/utils.py | 6 +- .../check_provider_folders_documented.py | 1 + .../handlers/test_tool_loop_handler.py | 476 ++++++++++++++++++ tests/unit/harness/test_init.py | 5 +- tests/unit/harness/test_types.py | 1 + tests/unit/llms/tool_loop/__init__.py | 0 tests/unit/llms/tool_loop/harness/__init__.py | 0 .../tool_loop/harness/test_transformation.py | 154 ++++++ 20 files changed, 1148 insertions(+), 25 deletions(-) create mode 100644 litellm/harness/handlers/tool_loop_handler.py create mode 100644 litellm/llms/tool_loop/__init__.py create mode 100644 litellm/llms/tool_loop/harness/__init__.py create mode 100644 litellm/llms/tool_loop/harness/transformation.py create mode 100644 tests/unit/harness/handlers/test_tool_loop_handler.py create mode 100644 tests/unit/llms/tool_loop/__init__.py create mode 100644 tests/unit/llms/tool_loop/harness/__init__.py create mode 100644 tests/unit/llms/tool_loop/harness/test_transformation.py diff --git a/litellm/__init__.py b/litellm/__init__.py index b0761da7f7c..30a9a80e9b2 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -2302,6 +2302,7 @@ _AGENT_EXPORTS: Final = frozenset( "CodexOptions", "OpenCodeOptions", "DeepAgentsOptions", + "ToolLoopOptions", } ) diff --git a/litellm/harness/__init__.py b/litellm/harness/__init__.py index 79322cdc3e5..213591c5949 100644 --- a/litellm/harness/__init__.py +++ b/litellm/harness/__init__.py @@ -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", diff --git a/litellm/harness/handlers/__init__.py b/litellm/harness/handlers/__init__.py index 4b32d7ca649..5b374a942fa 100644 --- a/litellm/harness/handlers/__init__.py +++ b/litellm/harness/handlers/__init__.py @@ -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}") diff --git a/litellm/harness/handlers/tool_loop_handler.py b/litellm/harness/handlers/tool_loop_handler.py new file mode 100644 index 00000000000..33efcba4ff7 --- /dev/null +++ b/litellm/harness/handlers/tool_loop_handler.py @@ -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") diff --git a/litellm/harness/options.py b/litellm/harness/options.py index 2e3ec5d0fd5..b0014fea393 100644 --- a/litellm/harness/options.py +++ b/litellm/harness/options.py @@ -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 diff --git a/litellm/harness/sync.py b/litellm/harness/sync.py index 1cb8fbcbca2..cdcf251af51 100644 --- a/litellm/harness/sync.py +++ b/litellm/harness/sync.py @@ -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 diff --git a/litellm/harness/types.py b/litellm/harness/types.py index 88ad7f711ea..475568a9979 100644 --- a/litellm/harness/types.py +++ b/litellm/harness/types.py @@ -27,6 +27,7 @@ class Harness(Enum): CODEX = "codex" OPENCODE = "opencode" DEEPAGENTS = "deepagents" + TOOL_LOOP = "tool_loop" def require_harness(harness: object) -> Harness: diff --git a/litellm/llms/base_llm/harness/utils.py b/litellm/llms/base_llm/harness/utils.py index 7c0278d5e36..8dcf784c90f 100644 --- a/litellm/llms/base_llm/harness/utils.py +++ b/litellm/llms/base_llm/harness/utils.py @@ -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: diff --git a/litellm/llms/deepagents/harness/transformation.py b/litellm/llms/deepagents/harness/transformation.py index 82b2c044eac..14f87f0753b 100644 --- a/litellm/llms/deepagents/harness/transformation.py +++ b/litellm/llms/deepagents/harness/transformation.py @@ -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 diff --git a/litellm/llms/tool_loop/__init__.py b/litellm/llms/tool_loop/__init__.py new file mode 100644 index 00000000000..f163f698841 --- /dev/null +++ b/litellm/llms/tool_loop/__init__.py @@ -0,0 +1 @@ +"""In-process tool loop harness.""" diff --git a/litellm/llms/tool_loop/harness/__init__.py b/litellm/llms/tool_loop/harness/__init__.py new file mode 100644 index 00000000000..49ca50b61c2 --- /dev/null +++ b/litellm/llms/tool_loop/harness/__init__.py @@ -0,0 +1 @@ +"""Tool Loop harness configuration.""" diff --git a/litellm/llms/tool_loop/harness/transformation.py b/litellm/llms/tool_loop/harness/transformation.py new file mode 100644 index 00000000000..16461916f1a --- /dev/null +++ b/litellm/llms/tool_loop/harness/transformation.py @@ -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=") diff --git a/litellm/utils.py b/litellm/utils.py index 9200844a2e3..f80a2851ba6 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 diff --git a/tests/code_coverage_tests/check_provider_folders_documented.py b/tests/code_coverage_tests/check_provider_folders_documented.py index 08fbde3d979..6f58b393cd1 100644 --- a/tests/code_coverage_tests/check_provider_folders_documented.py +++ b/tests/code_coverage_tests/check_provider_folders_documented.py @@ -33,6 +33,7 @@ EXCLUDED_FOLDERS = { "codex", "opencode", "deepagents", + "tool_loop", } diff --git a/tests/unit/harness/handlers/test_tool_loop_handler.py b/tests/unit/harness/handlers/test_tool_loop_handler.py new file mode 100644 index 00000000000..2ca265cda67 --- /dev/null +++ b/tests/unit/harness/handlers/test_tool_loop_handler.py @@ -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" diff --git a/tests/unit/harness/test_init.py b/tests/unit/harness/test_init.py index effdc238249..75ec3b60c74 100644 --- a/tests/unit/harness/test_init.py +++ b/tests/unit/harness/test_init.py @@ -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 diff --git a/tests/unit/harness/test_types.py b/tests/unit/harness/test_types.py index 036d6dd625c..e59c4e983fb 100644 --- a/tests/unit/harness/test_types.py +++ b/tests/unit/harness/test_types.py @@ -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) diff --git a/tests/unit/llms/tool_loop/__init__.py b/tests/unit/llms/tool_loop/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/tool_loop/harness/__init__.py b/tests/unit/llms/tool_loop/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/tool_loop/harness/test_transformation.py b/tests/unit/llms/tool_loop/harness/test_transformation.py new file mode 100644 index 00000000000..11bce41bb56 --- /dev/null +++ b/tests/unit/llms/tool_loop/harness/test_transformation.py @@ -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))