fix(responses): satisfy the code quality gates in tool search lowering

The code quality job forbids recursive functions, so loading the members of a namespace no longer goes back through the same helper. A namespace only holds one level of tools, so the result is the same. assert_never now comes from typing_extensions because litellm still supports Python 3.10

A lifted function call needs its call_id, which the Responses API always sends, so the made up id for a call without one is gone. New test cases cover a loaded function without parameters, output items that aren't function calls, stream data that isn't JSON, and a live sync stream
This commit is contained in:
Id545 2026-10-02 21:32:17 +02:00
parent 96ef4310fb
commit f2b332f769
6 changed files with 111 additions and 39 deletions

View file

@ -1,5 +1,7 @@
from collections.abc import Mapping, Sequence
from typing import Final, assert_never
from typing import Final
from typing_extensions import assert_never
import litellm
from litellm._internal_context import is_internal_call

View file

@ -1,5 +1,4 @@
import json
import uuid
from collections.abc import AsyncIterator, Iterator, Sequence
from typing import Final, Literal, Protocol
@ -39,7 +38,7 @@ class _FunctionCall(BaseModel):
name: str
namespace: str | None = None
id: str | None = None
call_id: str | None = None
call_id: str
arguments: str = ""
@ -80,8 +79,6 @@ def _search_arguments(arguments: str) -> dict[str, object]:
def _tool_search_call_item_id(call: _FunctionCall) -> str:
source: Final = call.id or call.call_id
if source is None:
return f"{_TOOL_SEARCH_CALL_ID_PREFIX}{uuid.uuid4().hex}"
return f"{_TOOL_SEARCH_CALL_ID_PREFIX}{source.partition('_')[2] or source}"

View file

@ -116,15 +116,19 @@ def _lowered_tool(tool: object) -> object:
}
def _loaded_tool(tool: dict[str, object]) -> dict[str, object]:
def _loaded_definition(tool: dict[str, object]) -> dict[str, object]:
loaded: Final = {key: value for key, value in tool.items() if key != "defer_loading"}
kind: Final = _parsed(_Typed, tool)
if kind is not None and kind.type == "function":
return {**loaded, "parameters": _object_schema(tool.get("parameters"))}
return loaded
def _loaded_tool(tool: dict[str, object]) -> dict[str, object]:
namespace: Final = _parsed(_Namespace, tool)
if namespace is None:
return loaded
return {**loaded, "tools": [_loaded_tool(member) for member in namespace.tools]}
return _loaded_definition(tool)
return {**_loaded_definition(tool), "tools": [_loaded_definition(member) for member in namespace.tools]}
def _visible_loaded_tool(tool: dict[str, object]) -> dict[str, object]:

View file

@ -116,7 +116,12 @@ async def test_the_next_turn_replays_the_search_as_function_items_with_the_loade
input=[
{"role": "user", "content": "add a meeting"},
{"type": "tool_search_call", "call_id": "call_search", "execution": "client", "arguments": {"query": "x"}},
{"type": "tool_search_output", "call_id": "call_search", "execution": "client", "tools": [loaded_namespace]},
{
"type": "tool_search_output",
"call_id": "call_search",
"execution": "client",
"tools": [loaded_namespace],
},
],
)

View file

@ -8,13 +8,17 @@ from openai.types.responses import ResponseToolSearchCall
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.hosted_vllm.responses.transformation import HostedVLLMResponsesAPIConfig
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
ResponsesAPIStreamingIterator,
SyncResponsesAPIStreamingIterator,
)
from litellm.responses.tool_search.lifting import (
ToolSearchEventLifter,
lift_tool_search_calls,
lift_tool_search_stream,
)
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.llms.openai import ResponsesAPIResponse, ResponsesAPIStreamEvents
SEARCH_CALL: Final = {
"type": "function_call",
@ -33,6 +37,13 @@ SHELL_CALL: Final = {
"status": "completed",
}
NAMESPACED_SEARCH_CALL: Final = {**SHELL_CALL, "id": "fc_ns", "name": "tool_search", "namespace": "docs"}
MESSAGE: Final = {
"type": "message",
"id": "msg_1",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "Searching the lab tools", "annotations": []}],
}
def _response(*output: dict[str, object]) -> dict[str, object]:
@ -55,7 +66,9 @@ def _event(event_type: str, **fields: object) -> dict[str, object]:
STREAM_EVENTS: Final = [
_event("response.created", response={**_response(), "status": "in_progress"}),
_event("response.output_item.added", output_index=0, item={**SEARCH_CALL, "arguments": "", "status": "in_progress"}),
_event(
"response.output_item.added", output_index=0, item={**SEARCH_CALL, "arguments": "", "status": "in_progress"}
),
_event("response.function_call_arguments.delta", item_id="fc_search", output_index=0, delta='{"query": '),
_event("response.function_call_arguments.done", item_id="fc_search", output_index=0, arguments="{}"),
_event("response.output_item.done", output_index=0, item=SEARCH_CALL),
@ -78,13 +91,14 @@ def _lifted_search_call() -> dict[str, object]:
def test_search_function_call_comes_back_as_a_tool_search_call():
response: Final = ResponsesAPIResponse.model_validate(_response(SEARCH_CALL, SHELL_CALL, NAMESPACED_SEARCH_CALL))
response: Final = ResponsesAPIResponse.model_validate(
_response(SEARCH_CALL, SHELL_CALL, NAMESPACED_SEARCH_CALL, MESSAGE)
)
lifted: Final = lift_tool_search_calls(response)
assert lifted.output[0] == ResponseToolSearchCall.model_validate(_lifted_search_call())
assert lifted.output[1] is response.output[1]
assert lifted.output[2] is response.output[2]
assert all(after is before for after, before in zip(lifted.output[1:], response.output[1:], strict=True))
def test_response_without_a_search_call_is_returned_as_is():
@ -123,40 +137,77 @@ def test_stream_events_of_a_search_call_are_lifted_and_its_argument_deltas_dropp
assert kept[6]["response"]["output"] == [_lifted_search_call(), SHELL_CALL]
@pytest.mark.asyncio
async def test_a_live_stream_logs_the_lifted_response():
body: Final = "".join(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" for event in STREAM_EVENTS)
def test_stream_data_that_is_not_json_passes_through():
assert ToolSearchEventLifter().lift("[DONE]") == "[DONE]"
def upstream(request: httpx.Request) -> httpx.Response:
return httpx.Response(200, content=body.encode(), headers={"content-type": "text/event-stream"})
SSE_BODY: Final = "".join(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" for event in STREAM_EVENTS)
def _sse_upstream(request: httpx.Request) -> httpx.Response:
return httpx.Response(200, content=SSE_BODY.encode(), headers={"content-type": "text/event-stream"})
def _logging_obj(call_type: str) -> LiteLLMLoggingObj:
logging_obj: Final = LiteLLMLoggingObj(
model="qwen",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="aresponses",
call_type=call_type,
start_time=datetime.now(),
litellm_call_id="tool-search-stream",
function_id="tool-search-stream",
litellm_call_id=f"tool-search-{call_type}",
function_id=f"tool-search-{call_type}",
)
logging_obj.model_call_details["litellm_params"] = {"aresponses": True}
async with httpx.AsyncClient(transport=httpx.MockTransport(upstream)) as client:
upstream_response: Final = await client.send(client.build_request("POST", "http://vllm.test/v1/responses"), stream=True)
stream: Final = lift_tool_search_stream(
ResponsesAPIStreamingIterator(
response=upstream_response,
model="qwen",
responses_api_provider_config=HostedVLLMResponsesAPIConfig(),
logging_obj=logging_obj,
custom_llm_provider="hosted_vllm",
)
)
events: Final = [event async for event in stream]
logging_obj.model_call_details["litellm_params"] = {"aresponses": call_type == "aresponses"}
return logging_obj
assert [str(event.type.value) for event in events].count("response.function_call_arguments.delta") == 1
def _assert_search_lifted_and_logged(events: list[object], stream: BaseResponsesAPIStreamingIterator) -> None:
assert [str(getattr(event, "type", "")) for event in events].count(
str(ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA)
) == 1
assert stream.completed_response is not None
logged_output: Final = stream.completed_response.response.output
assert [item["type"] if isinstance(item, dict) else item.type for item in logged_output] == [
"tool_search_call",
"function_call",
]
@pytest.mark.asyncio
async def test_a_live_stream_logs_the_lifted_response():
async with httpx.AsyncClient(transport=httpx.MockTransport(_sse_upstream)) as client:
upstream_response: Final = await client.send(
client.build_request("POST", "http://vllm.test/v1/responses"), stream=True
)
stream: Final = lift_tool_search_stream(
ResponsesAPIStreamingIterator(
response=upstream_response,
model="qwen",
responses_api_provider_config=HostedVLLMResponsesAPIConfig(),
logging_obj=_logging_obj("aresponses"),
custom_llm_provider="hosted_vllm",
)
)
events: Final = [event async for event in stream]
_assert_search_lifted_and_logged(events, stream)
def test_a_live_sync_stream_logs_the_lifted_response():
with httpx.Client(transport=httpx.MockTransport(_sse_upstream)) as client:
upstream_response: Final = client.send(
client.build_request("POST", "http://vllm.test/v1/responses"), stream=True
)
stream: Final = lift_tool_search_stream(
SyncResponsesAPIStreamingIterator(
response=upstream_response,
model="qwen",
responses_api_provider_config=HostedVLLMResponsesAPIConfig(),
logging_obj=_logging_obj("responses"),
custom_llm_provider="hosted_vllm",
)
)
events: Final = list(stream)
_assert_search_lifted_and_logged(events, stream)

View file

@ -42,12 +42,18 @@ CALENDAR_NAMESPACE: Final = {
}
],
}
LIST_INSTRUMENTS: Final = {
"type": "function",
"name": "list_instruments",
"description": "List the lab instruments",
"defer_loading": True,
}
SEARCH_OUTPUT: Final = {
"type": "tool_search_output",
"call_id": "call_search",
"execution": "client",
"status": "completed",
"tools": [CALENDAR_NAMESPACE],
"tools": [CALENDAR_NAMESPACE, LIST_INSTRUMENTS],
}
@ -128,8 +134,14 @@ def test_replayed_search_output_loads_its_tools_for_the_model():
assert isinstance(function_output, dict)
assert function_output["type"] == "function_call_output"
assert function_output["call_id"] == "call_search"
loaded_function: Final = {
"type": "function",
"name": "list_instruments",
"description": "List the lab instruments",
"parameters": {"type": "object", "properties": {}},
}
assert json.loads(function_output["output"]) == {
"tools": [{"type": "namespace", "name": "calendar", "description": "Calendar tools"}]
"tools": [{"type": "namespace", "name": "calendar", "description": "Calendar tools"}, loaded_function]
}
namespace: Final = _tool_named(lowered.tools, "calendar")
assert namespace["tools"] == [
@ -140,6 +152,7 @@ def test_replayed_search_output_loads_its_tools_for_the_model():
"parameters": {"type": "object", "properties": {"title": {"type": "string"}}},
}
]
assert _tool_named(lowered.tools, "list_instruments") == loaded_function
def test_loaded_members_join_a_namespace_the_request_already_declares():