mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 8b20b93309 into 9a5e828310
This commit is contained in:
commit
bde0e164f2
4 changed files with 317 additions and 31 deletions
|
|
@ -46,10 +46,12 @@ class _FunctionToolCall:
|
|||
AsyncCompletion: TypeAlias = Callable[..., Awaitable[ModelResponse]]
|
||||
_ARGUMENTS_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_MAPPING_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
_OBJECT_ADAPTER: Final = TypeAdapter(object)
|
||||
_MODEL_RESPONSE_ADAPTER: Final = TypeAdapter(ModelResponse)
|
||||
_RAW_DECODE_ADAPTER: Final = TypeAdapter(tuple[object, int])
|
||||
_JSON_DECODER: Final = json.JSONDecoder()
|
||||
_HISTORY_ADAPTER: Final = TypeAdapter(list[dict[str, object]])
|
||||
_TOOL_CALLS_ADAPTER: Final = TypeAdapter(tuple[Mapping[str, object], ...])
|
||||
|
||||
|
||||
class _Usage(BaseModel):
|
||||
|
|
@ -88,6 +90,41 @@ def _normalize_tool_call(
|
|||
return _FunctionToolCall(id=call.id, name=call.custom.name, arguments=call.custom.input)
|
||||
|
||||
|
||||
def _tool_call_id(call: Mapping[str, object]) -> str | None:
|
||||
call_id: Final[object] = call.get("id")
|
||||
return call_id if isinstance(call_id, str) else None
|
||||
|
||||
|
||||
def _tool_call_ids(message: ChatCompletionMessageParam) -> tuple[str, ...]:
|
||||
message_mapping: Final = _as_mapping(message)
|
||||
if message_mapping is None or message_mapping.get("role") != "assistant":
|
||||
return ()
|
||||
tool_calls_value: Final = message_mapping.get("tool_calls")
|
||||
if tool_calls_value is None:
|
||||
return ()
|
||||
try:
|
||||
tool_calls: Final = _TOOL_CALLS_ADAPTER.validate_python(tool_calls_value)
|
||||
except ValidationError:
|
||||
return ()
|
||||
return tuple(call_id for call_id in map(_tool_call_id, tool_calls) if call_id is not None)
|
||||
|
||||
|
||||
def _tool_message_call_id(message: ChatCompletionMessageParam) -> str | None:
|
||||
message_mapping: Final = _as_mapping(message)
|
||||
if message_mapping is None or message_mapping.get("role") != "tool":
|
||||
return None
|
||||
call_id: Final[object] = message_mapping.get("tool_call_id")
|
||||
return call_id if isinstance(call_id, str) else None
|
||||
|
||||
|
||||
def _interrupted_tool_message(call_id: str) -> ChatCompletionMessageParam:
|
||||
return {
|
||||
"role": "tool",
|
||||
"tool_call_id": call_id,
|
||||
"content": "interrupted: the turn ended before this tool ran",
|
||||
}
|
||||
|
||||
|
||||
def _parse_arguments(raw: str) -> tuple[dict[str, object], str | None]: # mutable-ok: event input uses a dict
|
||||
text: Final = raw.strip()
|
||||
try:
|
||||
|
|
@ -173,24 +210,31 @@ async def _tool_outcome(
|
|||
|
||||
|
||||
def _record_usage(ctx: SessionContext, response: ModelResponse) -> None:
|
||||
ctx.calls += 1 # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
input_tokens, output_tokens = _usage_counts(response)
|
||||
ctx.input_tokens += input_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
ctx.output_tokens += output_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
ctx.cost += _response_cost(response) # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
|
||||
|
||||
def _usage_counts(response: ModelResponse) -> tuple[int, int]:
|
||||
try:
|
||||
usage_value: Final[object] = getattr(response, "usage", None)
|
||||
usage: Final = _USAGE_ADAPTER.validate_python(usage_value)
|
||||
input_tokens: Final = usage.prompt_tokens or 0
|
||||
output_tokens: Final = usage.completion_tokens or 0
|
||||
ctx.calls += 1 # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
ctx.input_tokens += input_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
ctx.output_tokens += output_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
ctx.cost += _response_cost(response) # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
except Exception:
|
||||
return
|
||||
return 0, 0
|
||||
return usage.prompt_tokens or 0, usage.completion_tokens or 0
|
||||
|
||||
|
||||
async def _execute_tool(tool: FunctionTool, arguments: Mapping[str, object]) -> object:
|
||||
validated_model: Final = tool.args_model.model_validate(arguments)
|
||||
values_object: Final[object] = validated_model.model_dump()
|
||||
validated: Final = _MAPPING_ADAPTER.validate_python(values_object)
|
||||
parameters: Final = tuple(inspect.signature(tool.fn).parameters.values())
|
||||
validated: Final[Mapping[str, object]] = MappingProxyType(
|
||||
{
|
||||
parameter.name: _OBJECT_ADAPTER.validate_python(getattr(validated_model, parameter.name))
|
||||
for parameter in parameters
|
||||
}
|
||||
)
|
||||
positional_args: Final = tuple(
|
||||
validated[parameter.name] for parameter in parameters if parameter.kind is inspect.Parameter.POSITIONAL_ONLY
|
||||
)
|
||||
|
|
@ -253,11 +297,31 @@ class ToolLoopHandler(BaseHarnessHandler):
|
|||
return _HISTORY_ADAPTER.validate_python(history_object) # pyright: ignore[reportIncompatibleMethodOverride] # base history uses Any
|
||||
|
||||
async def stop(self, ctx: SessionContext) -> None:
|
||||
return None
|
||||
self._close_pending_calls()
|
||||
|
||||
def _close_pending_calls(self) -> None:
|
||||
assistant_offset: Final[int | None] = next(
|
||||
(offset for offset, message in enumerate(reversed(self._messages)) if _tool_call_ids(message)),
|
||||
None,
|
||||
)
|
||||
if assistant_offset is None:
|
||||
return
|
||||
assistant_message: Final = self._messages[-assistant_offset - 1]
|
||||
messages_after: Final = self._messages[len(self._messages) - assistant_offset :]
|
||||
completed_call_ids: Final = frozenset(
|
||||
call_id for call_id in map(_tool_message_call_id, messages_after) if call_id is not None
|
||||
)
|
||||
pending_call_ids: Final = tuple(
|
||||
call_id for call_id in _tool_call_ids(assistant_message) if call_id not in completed_call_ids
|
||||
)
|
||||
pending_messages: Final[tuple[ChatCompletionMessageParam, ...]] = tuple(
|
||||
_interrupted_tool_message(call_id) for call_id in pending_call_ids
|
||||
)
|
||||
self._messages = (*self._messages, *pending_messages)
|
||||
|
||||
async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]:
|
||||
self._close_pending_calls()
|
||||
ctx.final_text = "" # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
ctx.output_json = None # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
user_message: Final[ChatCompletionMessageParam] = {"role": "user", "content": prompt}
|
||||
self._messages = (*self._messages, user_message)
|
||||
for _ in range(TOOL_LOOP_MAX_MODEL_CALLS):
|
||||
|
|
@ -289,7 +353,6 @@ class ToolLoopHandler(BaseHarnessHandler):
|
|||
if not tool_calls:
|
||||
final_text = content or ""
|
||||
ctx.final_text = final_text # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
ctx.output_json = content if ctx.output is not None else None # rebind-ok: per-turn output sink
|
||||
final_message: ChatCompletionMessageParam = {
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
|
|
@ -335,11 +398,11 @@ class ToolLoopHandler(BaseHarnessHandler):
|
|||
parse_error,
|
||||
approval_error,
|
||||
)
|
||||
yield ToolResult(id=call.id, output=outcome.output, is_error=outcome.is_error)
|
||||
tool_message: ChatCompletionMessageParam = {
|
||||
"role": "tool",
|
||||
"tool_call_id": call.id,
|
||||
"content": outcome.output,
|
||||
}
|
||||
self._messages = (*self._messages, tool_message)
|
||||
yield ToolResult(id=call.id, output=outcome.output, is_error=outcome.is_error)
|
||||
raise HarnessTurnError(f"Harness.TOOL_LOOP exceeded {TOOL_LOOP_MAX_MODEL_CALLS} model calls in one turn")
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from litellm.types.utils import ChatCompletionToolParam
|
|||
TOOL_LOOP_MAX_MODEL_CALLS: Final = 100
|
||||
_ANNOTATIONS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
_OBJECT_ADAPTER: Final = TypeAdapter(object)
|
||||
_STRING_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str])
|
||||
_MODEL_FACTORY: Final[Callable[..., type[BaseModel]]] = create_model
|
||||
|
||||
|
||||
|
|
@ -72,15 +73,25 @@ def function_tool(fn: Callable[..., object]) -> FunctionTool:
|
|||
return FunctionTool(name=fn.__name__, fn=fn, args_model=args_model, spec=spec)
|
||||
|
||||
|
||||
def _routing_kwargs(ctx: SessionContext) -> Mapping[str, object]:
|
||||
def _routing_kwargs(ctx: SessionContext, options: ToolLoopOptions) -> Mapping[str, object]:
|
||||
if not ctx.model:
|
||||
raise ValueError("Harness.TOOL_LOOP needs model=")
|
||||
if ctx.gateway is not None:
|
||||
caller_headers_value: Final[object] = _OBJECT_ADAPTER.validate_python(
|
||||
options.completion_kwargs.get("extra_headers")
|
||||
)
|
||||
caller_headers: Final = _STRING_HEADERS_ADAPTER.validate_python(
|
||||
MappingProxyType({}) if caller_headers_value is None else caller_headers_value
|
||||
)
|
||||
extra_headers: Final[dict[str, str]] = { # mutable-ok: acompletion(extra_headers=) requires a dict
|
||||
**caller_headers,
|
||||
**gateway_headers(ctx),
|
||||
}
|
||||
return {
|
||||
"model": f"litellm_proxy/{ctx.model}",
|
||||
"api_base": ctx.gateway.api_base,
|
||||
"api_key": ctx.gateway.api_key,
|
||||
"extra_headers": gateway_headers(ctx),
|
||||
"extra_headers": extra_headers,
|
||||
}
|
||||
return {
|
||||
"model": ctx.model,
|
||||
|
|
@ -91,7 +102,7 @@ def _routing_kwargs(ctx: SessionContext) -> Mapping[str, object]:
|
|||
|
||||
def completion_kwargs(ctx: SessionContext) -> Mapping[str, object]:
|
||||
options: Final = ToolLoopHarnessConfig().get_options(ctx)
|
||||
routing: Final = _routing_kwargs(ctx)
|
||||
routing: Final = _routing_kwargs(ctx, options)
|
||||
kwargs: Final[Mapping[str, object]] = MappingProxyType({**options.completion_kwargs, **routing})
|
||||
if ctx.output is None:
|
||||
return kwargs
|
||||
|
|
@ -114,6 +125,8 @@ class ToolLoopHarnessConfig(BaseHarnessConfig[ToolLoopOptions]):
|
|||
)
|
||||
|
||||
def validate_environment(self, ctx: SessionContext) -> None:
|
||||
self.get_options(ctx)
|
||||
options: Final = self.get_options(ctx)
|
||||
if options.completion_kwargs.get("stream"):
|
||||
raise ValueError("Harness.TOOL_LOOP does not support completion_kwargs['stream']")
|
||||
if not ctx.model:
|
||||
raise ValueError("Harness.TOOL_LOOP needs model=")
|
||||
|
|
|
|||
|
|
@ -2,15 +2,16 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from collections.abc import AsyncGenerator, Callable, Iterator, Mapping
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal
|
||||
from typing import Final, Literal, cast
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm import sandbox
|
||||
from litellm.harness import runtime
|
||||
from litellm.harness.context import GatewayTarget, SessionContext
|
||||
from litellm.harness.handlers.tool_loop_handler import ToolLoopHandler
|
||||
from litellm.harness.options import ToolLoopOptions
|
||||
|
|
@ -351,21 +352,56 @@ async def test_async_tool_is_awaited(tmp_path: Path) -> None:
|
|||
assert ToolResult(id="call-1", output="12", is_error=False) in events
|
||||
|
||||
|
||||
class Answer(BaseModel):
|
||||
value: int
|
||||
class Range(BaseModel):
|
||||
start: int
|
||||
end: int
|
||||
|
||||
|
||||
async def test_structured_output_is_forwarded_and_retained(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion((model_response(content='{"value": 7}'),))
|
||||
ctx: Final = make_context(tmp_path, output=Answer)
|
||||
async def test_nested_pydantic_arguments_reach_the_tool_as_models(tmp_path: Path) -> None:
|
||||
received: list[Range] = [] # mutable-ok: captures injected tool arguments
|
||||
|
||||
def describe_range(value: Range) -> str:
|
||||
received.append(value)
|
||||
return f"{value.start}:{value.end}"
|
||||
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("describe_range", '{"value": {"start": 2, "end": 5}}'),)),
|
||||
model_response(content="Range received"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(describe_range,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "Return a value")
|
||||
events: Final = await run_turn(handler, ctx, "Describe a range")
|
||||
|
||||
assert completion.calls[0]["response_format"] is Answer
|
||||
assert ctx.output_json == '{"value": 7}'
|
||||
assert ctx.final_text == '{"value": 7}'
|
||||
assert received == [Range(start=2, end=5)]
|
||||
assert isinstance(received[0], Range)
|
||||
assert ToolResult(id="call-1", output="2:5", is_error=False) in events
|
||||
|
||||
|
||||
class Review(BaseModel):
|
||||
summary: str
|
||||
|
||||
|
||||
async def test_structured_output_parses_fenced_json_and_forwards_response_format(tmp_path: Path) -> None:
|
||||
content: Final = '```json\n{"summary": "Looks good"}\n```'
|
||||
completion: Final = ScriptedCompletion((model_response(content=content),))
|
||||
session: Final = runtime.aagent_session(
|
||||
Harness.TOOL_LOOP,
|
||||
sandbox=sandbox.local(tmp_path),
|
||||
model="gpt-4o-mini",
|
||||
output=Review,
|
||||
)
|
||||
session.handler = make_handler(completion)
|
||||
async with session:
|
||||
result: Final = await session.arun("Return a review")
|
||||
|
||||
assert isinstance(result.output, Review)
|
||||
assert result.output.summary == "Looks good"
|
||||
assert completion.calls[0]["response_format"] is Review
|
||||
assert session.ctx.output_json is None
|
||||
|
||||
|
||||
async def test_gateway_routing_uses_proxy_model_and_tool_loop_tag(tmp_path: Path) -> None:
|
||||
|
|
@ -418,6 +454,22 @@ async def test_usage_and_cost_accumulate_across_model_calls(tmp_path: Path) -> N
|
|||
assert ctx.cost == pytest.approx(0.6)
|
||||
|
||||
|
||||
async def test_missing_usage_still_records_call_and_response_cost(tmp_path: Path) -> None:
|
||||
response: Final = model_response(content="done", hidden_params={"response_cost": 0.5})
|
||||
response.usage = None
|
||||
completion: Final = ScriptedCompletion((response,))
|
||||
ctx: Final = make_context(tmp_path)
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "Hi")
|
||||
|
||||
assert ctx.calls == 1
|
||||
assert ctx.input_tokens == 0
|
||||
assert ctx.output_tokens == 0
|
||||
assert ctx.cost == pytest.approx(0.5)
|
||||
|
||||
|
||||
async def test_history_survives_stop_and_start(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion((model_response(content="first"), model_response(content="second")))
|
||||
ctx: Final = make_context(tmp_path, instructions="Keep answers concise")
|
||||
|
|
@ -440,6 +492,116 @@ async def test_history_survives_stop_and_start(tmp_path: Path) -> None:
|
|||
assert (await handler.history(ctx))[0]["content"] == "Keep answers concise"
|
||||
|
||||
|
||||
async def test_interrupted_tool_calls_are_closed_before_the_next_session_turn(tmp_path: Path) -> None:
|
||||
executed: list[str] = [] # mutable-ok: records which injected tools ran
|
||||
|
||||
def first_tool() -> str:
|
||||
executed.append("first")
|
||||
return "first result"
|
||||
|
||||
def second_tool() -> str:
|
||||
executed.append("second")
|
||||
return "second result"
|
||||
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(
|
||||
tool_calls=(
|
||||
function_call("first_tool", "{}", "call-1"),
|
||||
function_call("second_tool", "{}", "call-2"),
|
||||
)
|
||||
),
|
||||
model_response(content="Continued"),
|
||||
)
|
||||
)
|
||||
session: Final = runtime.aagent_session(
|
||||
Harness.TOOL_LOOP,
|
||||
sandbox=sandbox.local(tmp_path),
|
||||
model="gpt-4o-mini",
|
||||
tools=(first_tool, second_tool),
|
||||
max_turns=1,
|
||||
)
|
||||
session.handler = make_handler(completion)
|
||||
async with session:
|
||||
first_result: Final = await session.arun("first turn")
|
||||
assert first_result.stop_reason == "max_turns"
|
||||
await session.arun("second turn")
|
||||
|
||||
assert executed == ["first"]
|
||||
assert completion.calls[1]["messages"] == [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call-1",
|
||||
"type": "function",
|
||||
"function": {"name": "first_tool", "arguments": "{}"},
|
||||
},
|
||||
{
|
||||
"id": "call-2",
|
||||
"type": "function",
|
||||
"function": {"name": "second_tool", "arguments": "{}"},
|
||||
},
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call-1", "content": "first result"},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call-2",
|
||||
"content": "interrupted: the turn ended before this tool ran",
|
||||
},
|
||||
{"role": "user", "content": "second turn"},
|
||||
]
|
||||
|
||||
|
||||
async def test_closing_after_tool_result_keeps_result_and_interrupts_remaining_call(tmp_path: Path) -> None:
|
||||
executed: list[str] = [] # mutable-ok: records which injected tools ran
|
||||
|
||||
def first_tool() -> str:
|
||||
executed.append("first")
|
||||
return "first result"
|
||||
|
||||
def second_tool() -> str:
|
||||
executed.append("second")
|
||||
return "second result"
|
||||
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(
|
||||
tool_calls=(
|
||||
function_call("first_tool", "{}", "call-1"),
|
||||
function_call("second_tool", "{}", "call-2"),
|
||||
)
|
||||
),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(first_tool, second_tool))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
turn: Final = cast(AsyncGenerator[Event, None], handler.turn(ctx, "first turn"))
|
||||
tool_results: list[ToolResult] = [] # mutable-ok: records the first yielded tool result
|
||||
async for event in turn:
|
||||
if isinstance(event, ToolResult):
|
||||
tool_results.append(event)
|
||||
break
|
||||
await turn.aclose()
|
||||
await handler.stop(ctx)
|
||||
|
||||
assert tool_results == [ToolResult(id="call-1", output="first result", is_error=False)]
|
||||
assert executed == ["first"]
|
||||
assert (await handler.history(ctx))[-2:] == [
|
||||
{"role": "tool", "tool_call_id": "call-1", "content": "first result"},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call-2",
|
||||
"content": "interrupted: the turn ended before this tool ran",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
async def test_duplicate_tool_names_are_rejected(tmp_path: Path) -> None:
|
||||
ctx: Final = make_context(tmp_path, tools=(add, duplicate_add()))
|
||||
handler: Final = make_handler(ScriptedCompletion(()))
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ from pathlib import Path
|
|||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from litellm import sandbox
|
||||
from litellm.harness.context import GatewayTarget, SessionContext
|
||||
|
|
@ -126,6 +126,10 @@ def test_gateway_routing_and_response_format_override_options(tmp_path: Path) ->
|
|||
"api_key": "wrong-key",
|
||||
"api_base": "https://wrong",
|
||||
"response_format": "wrong-format",
|
||||
"extra_headers": {
|
||||
"x-caller-header": "preserved",
|
||||
"x-litellm-tags": "caller-tag",
|
||||
},
|
||||
}
|
||||
),
|
||||
output=OutputModel,
|
||||
|
|
@ -137,11 +141,39 @@ def test_gateway_routing_and_response_format_override_options(tmp_path: Path) ->
|
|||
"model": "litellm_proxy/anthropic/claude",
|
||||
"api_base": "https://gateway",
|
||||
"api_key": "virtual-key",
|
||||
"extra_headers": {"x-litellm-tags": "harness,tool_loop"},
|
||||
"extra_headers": {
|
||||
"x-caller-header": "preserved",
|
||||
"x-litellm-tags": "harness,tool_loop",
|
||||
},
|
||||
"response_format": OutputModel,
|
||||
}
|
||||
|
||||
|
||||
def test_gateway_extra_headers_must_be_string_mappings(tmp_path: Path) -> None:
|
||||
gateway: Final = GatewayTarget(api_base="https://gateway", api_key="virtual-key")
|
||||
ctx: Final = make_context(
|
||||
tmp_path,
|
||||
gateway=gateway,
|
||||
options=ToolLoopOptions(completion_kwargs={"extra_headers": {"x-invalid": 1}}),
|
||||
)
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
completion_kwargs(ctx)
|
||||
|
||||
|
||||
def test_gateway_none_extra_headers_uses_only_gateway_headers(tmp_path: Path) -> None:
|
||||
gateway: Final = GatewayTarget(api_base="https://gateway", api_key="virtual-key")
|
||||
ctx: Final = make_context(
|
||||
tmp_path,
|
||||
gateway=gateway,
|
||||
options=ToolLoopOptions(completion_kwargs={"extra_headers": None}),
|
||||
)
|
||||
|
||||
kwargs: Final = completion_kwargs(ctx)
|
||||
|
||||
assert kwargs["extra_headers"] == {"x-litellm-tags": "harness,tool_loop"}
|
||||
|
||||
|
||||
def test_configuration_requires_model_and_declares_capabilities(tmp_path: Path) -> None:
|
||||
config: Final = ToolLoopHarnessConfig()
|
||||
assert config.uses_model_endpoint is False
|
||||
|
|
@ -152,3 +184,19 @@ def test_configuration_requires_model_and_declares_capabilities(tmp_path: Path)
|
|||
assert config.capabilities.permission_modes == frozenset({"ask", "full"})
|
||||
with pytest.raises(ValueError, match=r"Harness\.TOOL_LOOP needs model="):
|
||||
config.validate_environment(make_context(tmp_path, model=None))
|
||||
|
||||
|
||||
def test_configuration_rejects_truthy_completion_stream(tmp_path: Path) -> None:
|
||||
ctx: Final = make_context(tmp_path, options=ToolLoopOptions(completion_kwargs={"stream": True}))
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match=r"Harness\.TOOL_LOOP does not support completion_kwargs\['stream'\]",
|
||||
):
|
||||
ToolLoopHarnessConfig().validate_environment(ctx)
|
||||
|
||||
|
||||
def test_configuration_allows_false_completion_stream(tmp_path: Path) -> None:
|
||||
ctx: Final = make_context(tmp_path, options=ToolLoopOptions(completion_kwargs={"stream": False}))
|
||||
|
||||
ToolLoopHarnessConfig().validate_environment(ctx)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue