mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(sdk): add run_tool_loop and arun_tool_loop helpers (#44381)
* feat(sdk): add run_tool_loop and arun_tool_loop helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(tests): allow-list bounded tool-loop recursion in recursive detector Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(sdk): harden run_tool_loop per review Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(deps): keep uv.lock at revision 3 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(tests): authorize typing-extensions PSF-2.0 license Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Krrish Dholakia <krrishdholakia@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
5d513d8053
commit
b61376a99e
7 changed files with 599 additions and 0 deletions
|
|
@ -1474,6 +1474,7 @@ from .rust_bridge import rust
|
|||
from .rag.main import *
|
||||
from .sandbox.main import *
|
||||
from .decisions.main import *
|
||||
from .tool_loop import ToolLoopMaxRoundsExceeded, arun_tool_loop, run_tool_loop
|
||||
from .search.main import *
|
||||
from .realtime_api.main import (
|
||||
_arealtime,
|
||||
|
|
|
|||
|
|
@ -2234,3 +2234,5 @@ HARNESS_SNAPSHOT_SKIP_DIRS: Final = frozenset(
|
|||
".ruff_cache",
|
||||
}
|
||||
)
|
||||
|
||||
DEFAULT_TOOL_LOOP_MAX_ROUNDS: Final = 20
|
||||
|
|
|
|||
172
litellm/tool_loop.py
Normal file
172
litellm/tool_loop.py
Normal file
|
|
@ -0,0 +1,172 @@
|
|||
"""Client-side tool-calling loop helpers for litellm.completion."""
|
||||
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from typing import TypeAlias, cast
|
||||
|
||||
from typing_extensions import TypedDict, Unpack
|
||||
|
||||
from litellm.constants import DEFAULT_TOOL_LOOP_MAX_ROUNDS
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionAssistantMessage,
|
||||
ChatCompletionAssistantToolCall,
|
||||
ChatCompletionToolCallFunctionChunk,
|
||||
ChatCompletionToolMessage,
|
||||
ChatCompletionToolParam,
|
||||
)
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Message, ModelResponse
|
||||
|
||||
ToolExecutor: TypeAlias = Callable[[ChatCompletionMessageToolCall], ChatCompletionToolMessage]
|
||||
AsyncToolExecutor: TypeAlias = Callable[[ChatCompletionMessageToolCall], Awaitable[ChatCompletionToolMessage]]
|
||||
|
||||
|
||||
class _ToolLoopCompletionKwargs(TypedDict, total=False, extra_items=object):
|
||||
"""Extra keywords forwarded verbatim to ``litellm.completion``, which owns their contract."""
|
||||
|
||||
|
||||
class ToolLoopMaxRoundsExceeded(RuntimeError):
|
||||
max_rounds: int
|
||||
|
||||
def __init__(self, max_rounds: int) -> None:
|
||||
self.max_rounds = max_rounds
|
||||
super().__init__(f"model still requested tool calls on round {max_rounds} of {max_rounds}")
|
||||
|
||||
|
||||
def _validate_tool_loop_args(max_rounds: int, completion_kwargs: _ToolLoopCompletionKwargs) -> None:
|
||||
if max_rounds < 1:
|
||||
raise ValueError(f"max_rounds must be >= 1, got {max_rounds}")
|
||||
if completion_kwargs.get("stream"):
|
||||
raise ValueError("run_tool_loop requires whole responses; stream=True is not supported")
|
||||
|
||||
|
||||
def _expect_model_response(response: object) -> ModelResponse:
|
||||
if not isinstance(response, ModelResponse):
|
||||
raise TypeError(f"run_tool_loop requires completion to return a ModelResponse, got {type(response).__name__}")
|
||||
return response
|
||||
|
||||
|
||||
def _assistant_tool_call(tool_call: ChatCompletionMessageToolCall) -> ChatCompletionAssistantToolCall:
|
||||
return ChatCompletionAssistantToolCall(
|
||||
id=tool_call.id,
|
||||
type="function",
|
||||
function=ChatCompletionToolCallFunctionChunk(
|
||||
name=tool_call.function.name, arguments=tool_call.function.arguments
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _assistant_message(
|
||||
message: Message, tool_calls: tuple[ChatCompletionMessageToolCall, ...]
|
||||
) -> ChatCompletionAssistantMessage:
|
||||
return cast( # cast-ok: dict literal with thinking_blocks and reasoning_items spread in only when set
|
||||
"ChatCompletionAssistantMessage",
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": message.content,
|
||||
"tool_calls": [_assistant_tool_call(tc) for tc in tool_calls],
|
||||
**{
|
||||
key: value
|
||||
for key, value in (
|
||||
("thinking_blocks", getattr(message, "thinking_blocks", None)),
|
||||
("reasoning_items", getattr(message, "reasoning_items", None)),
|
||||
)
|
||||
if value is not None
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _function_tool_call(tool_call: object) -> ChatCompletionMessageToolCall:
|
||||
if not isinstance(tool_call, ChatCompletionMessageToolCall):
|
||||
raise TypeError(
|
||||
f"run_tool_loop only executes function tool calls, got custom tool call {getattr(tool_call, 'id', None)}"
|
||||
)
|
||||
return tool_call
|
||||
|
||||
|
||||
def _function_tool_calls(message: Message) -> tuple[ChatCompletionMessageToolCall, ...]:
|
||||
return tuple(_function_tool_call(tool_call) for tool_call in message.tool_calls or ())
|
||||
|
||||
|
||||
def run_tool_loop(
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
tools: Sequence[ChatCompletionToolParam],
|
||||
execute_tool: ToolExecutor,
|
||||
max_rounds: int = DEFAULT_TOOL_LOOP_MAX_ROUNDS,
|
||||
**completion_kwargs: Unpack[_ToolLoopCompletionKwargs], # kwargs-ok: forwarded verbatim to litellm.completion
|
||||
) -> str | None:
|
||||
"""Call completion, execute each requested tool, and repeat until the model answers.
|
||||
|
||||
Returns the final assistant message content. Raises ToolLoopMaxRoundsExceeded when the
|
||||
model is still requesting tools after max_rounds completions.
|
||||
|
||||
Example:
|
||||
execute = functools.partial(run_repo_tool, repository="litellm", revision="main")
|
||||
answer = litellm.run_tool_loop(
|
||||
model="anthropic/claude-sonnet-5-5", messages=messages, tools=tools, execute_tool=execute
|
||||
)
|
||||
"""
|
||||
import litellm
|
||||
|
||||
_validate_tool_loop_args(max_rounds, completion_kwargs)
|
||||
history: tuple[AllMessageValues, ...] = tuple(messages) # rebind-ok: rounds append new turns
|
||||
for round_number in range(1, max_rounds + 1):
|
||||
message = (
|
||||
_expect_model_response(
|
||||
litellm.completion(model=model, messages=list(history), tools=list(tools), **completion_kwargs)
|
||||
)
|
||||
.choices[0]
|
||||
.message
|
||||
)
|
||||
if not message.tool_calls:
|
||||
return message.content
|
||||
tool_calls = _function_tool_calls(message)
|
||||
if round_number == max_rounds:
|
||||
break
|
||||
tool_results = tuple(execute_tool(tool_call) for tool_call in tool_calls)
|
||||
history = (*history, _assistant_message(message, tool_calls), *tool_results)
|
||||
raise ToolLoopMaxRoundsExceeded(max_rounds)
|
||||
|
||||
|
||||
async def arun_tool_loop(
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
tools: Sequence[ChatCompletionToolParam],
|
||||
execute_tool: AsyncToolExecutor,
|
||||
max_rounds: int = DEFAULT_TOOL_LOOP_MAX_ROUNDS,
|
||||
**completion_kwargs: Unpack[_ToolLoopCompletionKwargs], # kwargs-ok: forwarded verbatim to litellm.acompletion
|
||||
) -> str | None:
|
||||
"""Async version of run_tool_loop, awaiting each execute_tool call in order.
|
||||
|
||||
Returns the final assistant message content. Raises ToolLoopMaxRoundsExceeded when the
|
||||
model is still requesting tools after max_rounds completions.
|
||||
|
||||
Example:
|
||||
execute = functools.partial(arun_repo_tool, repository="litellm", revision="main")
|
||||
answer = await litellm.arun_tool_loop(
|
||||
model="anthropic/claude-sonnet-5-5", messages=messages, tools=tools, execute_tool=execute
|
||||
)
|
||||
"""
|
||||
import litellm
|
||||
|
||||
_validate_tool_loop_args(max_rounds, completion_kwargs)
|
||||
history: tuple[AllMessageValues, ...] = tuple(messages) # rebind-ok: rounds append new turns
|
||||
for round_number in range(1, max_rounds + 1):
|
||||
message = (
|
||||
_expect_model_response(
|
||||
await litellm.acompletion(model=model, messages=list(history), tools=list(tools), **completion_kwargs)
|
||||
)
|
||||
.choices[0]
|
||||
.message
|
||||
)
|
||||
if not message.tool_calls:
|
||||
return message.content
|
||||
tool_calls = _function_tool_calls(message)
|
||||
if round_number == max_rounds:
|
||||
break
|
||||
tool_results = tuple([await execute_tool(tool_call) for tool_call in tool_calls])
|
||||
history = (*history, _assistant_message(message, tool_calls), *tool_results)
|
||||
raise ToolLoopMaxRoundsExceeded(max_rounds)
|
||||
|
|
@ -34,6 +34,7 @@ dependencies = [
|
|||
"pydantic-settings>=2.14.1,<3.0",
|
||||
"jsonschema>=4.0.0,<5.0",
|
||||
"boto3>=1.43.1,<2.0",
|
||||
"typing-extensions>=4.13.0,<5.0",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
|
|
|
|||
|
|
@ -177,3 +177,4 @@ hypothesis: >=6.165.10 # MPL 2.0 license
|
|||
pytest-rerunfailures: >=15.1 # MPL 2.0 license
|
||||
pytest-recording: >=0.13.4 # MIT license
|
||||
expression: >=5.6.0 # MIT License - https://github.com/cognitedata/Expression/blob/main/LICENSE
|
||||
typing-extensions: >=4.13.0 # PSF-2.0 license - https://github.com/python/typing_extensions/blob/main/LICENSE
|
||||
|
|
|
|||
420
tests/unit/test_tool_loop.py
Normal file
420
tests/unit/test_tool_loop.py
Normal file
|
|
@ -0,0 +1,420 @@
|
|||
import json
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.tool_loop import ToolLoopMaxRoundsExceeded
|
||||
from litellm.types.llms.openai import ChatCompletionToolMessage
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall
|
||||
|
||||
OPENAI_CHAT_COMPLETIONS_URL: Final = "https://api.openai.com/v1/chat/completions"
|
||||
WEATHER_TOOLS: Final = (
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get the weather for a city",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _openai_response(content: str | None, tool_calls: list | None = None) -> dict:
|
||||
return {
|
||||
"id": "chatcmpl-tool-loop",
|
||||
"object": "chat.completion",
|
||||
"created": 1739462947,
|
||||
"model": "gpt-5-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"finish_reason": "tool_calls" if tool_calls else "stop",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
"tool_calls": tool_calls,
|
||||
},
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
||||
}
|
||||
|
||||
|
||||
def _tool_call(call_id: str, name: str, arguments: dict) -> dict:
|
||||
return {
|
||||
"id": call_id,
|
||||
"type": "function",
|
||||
"function": {"name": name, "arguments": json.dumps(arguments)},
|
||||
}
|
||||
|
||||
|
||||
def _tool_result(tool_call: ChatCompletionMessageToolCall) -> ChatCompletionToolMessage:
|
||||
return ChatCompletionToolMessage(role="tool", content='{"temp": "72F"}', tool_call_id=tool_call.id or "")
|
||||
|
||||
|
||||
def _request_bodies(respx_mock: respx.MockRouter) -> list[dict]:
|
||||
return [json.loads(call.request.content) for call in respx_mock.calls]
|
||||
|
||||
|
||||
def test_final_answer_without_tool_calls_returns_content(respx_mock: respx.MockRouter) -> None:
|
||||
route: Final = respx_mock.post(OPENAI_CHAT_COMPLETIONS_URL).mock(
|
||||
return_value=httpx.Response(200, json=_openai_response("done"))
|
||||
)
|
||||
executor_called: Final = []
|
||||
|
||||
def executor(tc: ChatCompletionMessageToolCall) -> ChatCompletionToolMessage:
|
||||
executor_called.append(tc)
|
||||
return _tool_result(tc)
|
||||
|
||||
answer: Final = litellm.run_tool_loop(
|
||||
model="openai/gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
tools=WEATHER_TOOLS,
|
||||
execute_tool=executor,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
assert answer == "done"
|
||||
assert executor_called == []
|
||||
assert route.call_count == 1
|
||||
|
||||
|
||||
def test_two_rounds_appends_assistant_and_tool_messages_in_order(respx_mock: respx.MockRouter) -> None:
|
||||
tool_calls: Final = [
|
||||
_tool_call("call_1", "get_weather", {"city": "Paris"}),
|
||||
_tool_call("call_2", "get_weather", {"city": "Tokyo"}),
|
||||
]
|
||||
route: Final = respx_mock.post(OPENAI_CHAT_COMPLETIONS_URL).mock(
|
||||
side_effect=[
|
||||
httpx.Response(200, json=_openai_response(None, tool_calls)),
|
||||
httpx.Response(200, json=_openai_response("Paris 72F, Tokyo 60F")),
|
||||
]
|
||||
)
|
||||
executed: Final = []
|
||||
|
||||
def executor(tc: ChatCompletionMessageToolCall) -> ChatCompletionToolMessage:
|
||||
executed.append(tc)
|
||||
return _tool_result(tc)
|
||||
|
||||
messages: Final = [{"role": "user", "content": "weather in Paris and Tokyo?"}]
|
||||
messages_snapshot: Final = [dict(message) for message in messages]
|
||||
|
||||
answer: Final = litellm.run_tool_loop(
|
||||
model="openai/gpt-5-mini",
|
||||
messages=messages,
|
||||
tools=WEATHER_TOOLS,
|
||||
execute_tool=executor,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
assert answer == "Paris 72F, Tokyo 60F"
|
||||
assert route.call_count == 2
|
||||
assert [tc.id for tc in executed] == ["call_1", "call_2"]
|
||||
assert [tc.function.name for tc in executed] == ["get_weather", "get_weather"]
|
||||
assert [tc.function.arguments for tc in executed] == [
|
||||
'{"city": "Paris"}',
|
||||
'{"city": "Tokyo"}',
|
||||
]
|
||||
|
||||
second_body: Final = _request_bodies(respx_mock)[1]
|
||||
assert second_body["messages"] == [
|
||||
{"role": "user", "content": "weather in Paris and Tokyo?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"city": "Paris"}'},
|
||||
},
|
||||
{
|
||||
"id": "call_2",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"city": "Tokyo"}'},
|
||||
},
|
||||
],
|
||||
},
|
||||
{"role": "tool", "content": '{"temp": "72F"}', "tool_call_id": "call_1"},
|
||||
{"role": "tool", "content": '{"temp": "72F"}', "tool_call_id": "call_2"},
|
||||
]
|
||||
|
||||
assert len(messages) == len(messages_snapshot)
|
||||
assert messages == messages_snapshot
|
||||
|
||||
|
||||
def test_response_format_and_tools_forwarded_every_round(respx_mock: respx.MockRouter) -> None:
|
||||
respx_mock.post(OPENAI_CHAT_COMPLETIONS_URL).mock(
|
||||
side_effect=[
|
||||
httpx.Response(200, json=_openai_response(None, [_tool_call("call_1", "get_weather", {"city": "Paris"})])),
|
||||
httpx.Response(200, json=_openai_response('{"summary": "sunny"}')),
|
||||
]
|
||||
)
|
||||
response_format: Final = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "weather_report",
|
||||
"schema": {
|
||||
"type": "object",
|
||||
"properties": {"summary": {"type": "string"}},
|
||||
"required": ["summary"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
litellm.run_tool_loop(
|
||||
model="openai/gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "weather?"}],
|
||||
tools=WEATHER_TOOLS,
|
||||
execute_tool=_tool_result,
|
||||
response_format=response_format,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
bodies: Final = _request_bodies(respx_mock)
|
||||
assert len(bodies) == 2
|
||||
for body in bodies:
|
||||
assert body["response_format"] == response_format
|
||||
assert body["tools"] == list(WEATHER_TOOLS)
|
||||
|
||||
|
||||
def test_max_rounds_exceeded_raises_without_executing_last_round(respx_mock: respx.MockRouter) -> None:
|
||||
tool_call: Final = _tool_call("call_1", "get_weather", {"city": "Paris"})
|
||||
route: Final = respx_mock.post(OPENAI_CHAT_COMPLETIONS_URL).mock(
|
||||
return_value=httpx.Response(200, json=_openai_response(None, [tool_call]))
|
||||
)
|
||||
executed: Final = []
|
||||
|
||||
def executor(tc: ChatCompletionMessageToolCall) -> ChatCompletionToolMessage:
|
||||
executed.append(tc)
|
||||
return _tool_result(tc)
|
||||
|
||||
with pytest.raises(ToolLoopMaxRoundsExceeded) as exc_info:
|
||||
litellm.run_tool_loop(
|
||||
model="openai/gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "weather?"}],
|
||||
tools=WEATHER_TOOLS,
|
||||
execute_tool=executor,
|
||||
max_rounds=2,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
assert exc_info.value.max_rounds == 2
|
||||
assert route.call_count == 2
|
||||
assert [tc.id for tc in executed] == ["call_1"]
|
||||
|
||||
|
||||
def test_max_rounds_below_one_rejected_before_any_request(respx_mock: respx.MockRouter) -> None:
|
||||
route: Final = respx_mock.post(OPENAI_CHAT_COMPLETIONS_URL).mock(
|
||||
return_value=httpx.Response(200, json=_openai_response("done"))
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="max_rounds must be >= 1"):
|
||||
litellm.run_tool_loop(
|
||||
model="openai/gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
tools=WEATHER_TOOLS,
|
||||
execute_tool=_tool_result,
|
||||
max_rounds=0,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
assert route.call_count == 0
|
||||
|
||||
|
||||
def test_stream_rejected_before_any_request(respx_mock: respx.MockRouter) -> None:
|
||||
route: Final = respx_mock.post(OPENAI_CHAT_COMPLETIONS_URL).mock(
|
||||
return_value=httpx.Response(200, json=_openai_response("done"))
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="stream=True is not supported"):
|
||||
litellm.run_tool_loop(
|
||||
model="openai/gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
tools=WEATHER_TOOLS,
|
||||
execute_tool=_tool_result,
|
||||
stream=True,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
assert route.call_count == 0
|
||||
|
||||
|
||||
def test_custom_tool_call_raises_type_error_without_executing(respx_mock: respx.MockRouter) -> None:
|
||||
custom_response: Final = _openai_response(
|
||||
None,
|
||||
[{"id": "call_custom", "type": "custom", "custom": {"name": "apply_patch", "input": "*** patch"}}],
|
||||
)
|
||||
route: Final = respx_mock.post(OPENAI_CHAT_COMPLETIONS_URL).mock(
|
||||
return_value=httpx.Response(200, json=custom_response)
|
||||
)
|
||||
executor_called: Final = []
|
||||
|
||||
def executor(tc: ChatCompletionMessageToolCall) -> ChatCompletionToolMessage:
|
||||
executor_called.append(tc)
|
||||
return _tool_result(tc)
|
||||
|
||||
with pytest.raises(TypeError, match="custom tool call call_custom"):
|
||||
litellm.run_tool_loop(
|
||||
model="openai/gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
tools=WEATHER_TOOLS,
|
||||
execute_tool=executor,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
assert executor_called == []
|
||||
assert route.call_count == 1
|
||||
|
||||
|
||||
def _responses_payload(response_id: str, output: list) -> dict:
|
||||
return {
|
||||
"id": response_id,
|
||||
"object": "response",
|
||||
"created_at": 1734366691,
|
||||
"status": "completed",
|
||||
"model": "gpt-5.5",
|
||||
"output": output,
|
||||
"parallel_tool_calls": True,
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||
"error": None,
|
||||
"incomplete_details": None,
|
||||
"instructions": None,
|
||||
"metadata": None,
|
||||
"temperature": None,
|
||||
"tool_choice": "auto",
|
||||
"tools": [],
|
||||
"top_p": None,
|
||||
"max_output_tokens": None,
|
||||
"previous_response_id": None,
|
||||
"reasoning": None,
|
||||
"truncation": None,
|
||||
"user": None,
|
||||
}
|
||||
|
||||
|
||||
def test_responses_bridge_replays_reasoning_items_across_rounds(respx_mock: respx.MockRouter) -> None:
|
||||
round_one: Final = _responses_payload(
|
||||
"resp_1",
|
||||
[
|
||||
{"type": "reasoning", "id": "rs_abc123", "summary": [], "encrypted_content": "enc_xyz"},
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": "fc_1",
|
||||
"call_id": "call_1",
|
||||
"name": "get_weather",
|
||||
"arguments": '{"city": "Paris"}',
|
||||
"status": "completed",
|
||||
},
|
||||
],
|
||||
)
|
||||
round_two: Final = _responses_payload(
|
||||
"resp_2",
|
||||
[
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_1",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Paris is 72F", "annotations": []}],
|
||||
}
|
||||
],
|
||||
)
|
||||
route: Final = respx_mock.post("https://api.openai.com/v1/responses").mock(
|
||||
side_effect=[httpx.Response(200, json=round_one), httpx.Response(200, json=round_two)]
|
||||
)
|
||||
executed: Final = []
|
||||
|
||||
def executor(tc: ChatCompletionMessageToolCall) -> ChatCompletionToolMessage:
|
||||
executed.append(tc)
|
||||
return _tool_result(tc)
|
||||
|
||||
answer: Final = litellm.run_tool_loop(
|
||||
model="openai/responses/gpt-5.5",
|
||||
messages=[{"role": "user", "content": "weather in Paris?"}],
|
||||
tools=WEATHER_TOOLS,
|
||||
execute_tool=executor,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
assert answer == "Paris is 72F"
|
||||
assert route.call_count == 2
|
||||
assert [tc.id for tc in executed] == ["fc_1"]
|
||||
|
||||
second_input: Final = _request_bodies(respx_mock)[1]["input"]
|
||||
item_types: Final = [item.get("type") for item in second_input]
|
||||
reasoning_index: Final = next(i for i, item in enumerate(second_input) if item.get("type") == "reasoning")
|
||||
function_call_index: Final = next(
|
||||
i for i, item in enumerate(second_input) if item.get("type") == "function_call"
|
||||
)
|
||||
reasoning_item: Final = second_input[reasoning_index]
|
||||
assert reasoning_item["id"] == "rs_abc123"
|
||||
assert reasoning_item["encrypted_content"] == "enc_xyz"
|
||||
assert reasoning_index < function_call_index, f"reasoning item must precede function_call: {item_types}"
|
||||
|
||||
|
||||
async def test_arun_tool_loop_two_rounds(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "module_level_aclient", AsyncHTTPHandler())
|
||||
tool_calls: Final = [
|
||||
_tool_call("call_1", "get_weather", {"city": "Paris"}),
|
||||
_tool_call("call_2", "get_weather", {"city": "Tokyo"}),
|
||||
]
|
||||
route: Final = respx_mock.post(OPENAI_CHAT_COMPLETIONS_URL).mock(
|
||||
side_effect=[
|
||||
httpx.Response(200, json=_openai_response(None, tool_calls)),
|
||||
httpx.Response(200, json=_openai_response("Paris 72F, Tokyo 60F")),
|
||||
]
|
||||
)
|
||||
executed: Final = []
|
||||
|
||||
async def executor(tc: ChatCompletionMessageToolCall) -> ChatCompletionToolMessage:
|
||||
executed.append(tc)
|
||||
return _tool_result(tc)
|
||||
|
||||
messages: Final = [{"role": "user", "content": "weather in Paris and Tokyo?"}]
|
||||
messages_snapshot: Final = [dict(message) for message in messages]
|
||||
|
||||
answer: Final = await litellm.arun_tool_loop(
|
||||
model="openai/gpt-5-mini",
|
||||
messages=messages,
|
||||
tools=WEATHER_TOOLS,
|
||||
execute_tool=executor,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
assert answer == "Paris 72F, Tokyo 60F"
|
||||
assert route.call_count == 2
|
||||
assert [tc.id for tc in executed] == ["call_1", "call_2"]
|
||||
|
||||
second_body: Final = _request_bodies(respx_mock)[1]
|
||||
assert second_body["messages"] == [
|
||||
{"role": "user", "content": "weather in Paris and Tokyo?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"city": "Paris"}'},
|
||||
},
|
||||
{
|
||||
"id": "call_2",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"city": "Tokyo"}'},
|
||||
},
|
||||
],
|
||||
},
|
||||
{"role": "tool", "content": '{"temp": "72F"}', "tool_call_id": "call_1"},
|
||||
{"role": "tool", "content": '{"temp": "72F"}', "tool_call_id": "call_2"},
|
||||
]
|
||||
assert messages == messages_snapshot
|
||||
2
uv.lock
generated
2
uv.lock
generated
|
|
@ -4521,6 +4521,7 @@ dependencies = [
|
|||
{ name = "pyyaml" },
|
||||
{ name = "tiktoken" },
|
||||
{ name = "tokenizers" },
|
||||
{ name = "typing-extensions" },
|
||||
]
|
||||
|
||||
[package.optional-dependencies]
|
||||
|
|
@ -4851,6 +4852,7 @@ requires-dist = [
|
|||
{ name = "tokenizers", specifier = ">=0.21.0,<1.0" },
|
||||
{ name = "tomlkit", marker = "extra == 'cli'", specifier = ">=0.13.3,<1.0" },
|
||||
{ name = "tomlkit", marker = "extra == 'proxy'", specifier = ">=0.13.3,<1.0" },
|
||||
{ name = "typing-extensions", specifier = ">=4.13.0,<5.0" },
|
||||
{ name = "uvicorn", marker = "extra == 'proxy'", specifier = ">=0.33.0,<1.0" },
|
||||
{ name = "uvloop", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.22.1,<1.0" },
|
||||
{ name = "websockets", marker = "extra == 'proxy'", specifier = ">=15.0.1,<16.0" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue