From f2b332f769166c2d0fb796b111e05b38f3f82ac6 Mon Sep 17 00:00:00 2001 From: Id545 Date: Fri, 2 Oct 2026 21:32:17 +0200 Subject: [PATCH] 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 --- litellm/responses/tool_search/handler.py | 4 +- litellm/responses/tool_search/lifting.py | 5 +- litellm/responses/tool_search/lowering.py | 10 +- .../responses/tool_search/test_handler.py | 7 +- .../responses/tool_search/test_lifting.py | 107 +++++++++++++----- .../responses/tool_search/test_lowering.py | 17 ++- 6 files changed, 111 insertions(+), 39 deletions(-) diff --git a/litellm/responses/tool_search/handler.py b/litellm/responses/tool_search/handler.py index 8f045688113..647ccbb81ea 100644 --- a/litellm/responses/tool_search/handler.py +++ b/litellm/responses/tool_search/handler.py @@ -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 diff --git a/litellm/responses/tool_search/lifting.py b/litellm/responses/tool_search/lifting.py index 9f2fe009084..501e981a83a 100644 --- a/litellm/responses/tool_search/lifting.py +++ b/litellm/responses/tool_search/lifting.py @@ -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}" diff --git a/litellm/responses/tool_search/lowering.py b/litellm/responses/tool_search/lowering.py index aaf33e805e1..59ef2bda084 100644 --- a/litellm/responses/tool_search/lowering.py +++ b/litellm/responses/tool_search/lowering.py @@ -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]: diff --git a/tests/unit/responses/tool_search/test_handler.py b/tests/unit/responses/tool_search/test_handler.py index 304ca5d5984..93b6398ebe0 100644 --- a/tests/unit/responses/tool_search/test_handler.py +++ b/tests/unit/responses/tool_search/test_handler.py @@ -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], + }, ], ) diff --git a/tests/unit/responses/tool_search/test_lifting.py b/tests/unit/responses/tool_search/test_lifting.py index 38c35398d14..14791f6484a 100644 --- a/tests/unit/responses/tool_search/test_lifting.py +++ b/tests/unit/responses/tool_search/test_lifting.py @@ -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) diff --git a/tests/unit/responses/tool_search/test_lowering.py b/tests/unit/responses/tool_search/test_lowering.py index 0dba6161eed..cf0d22cbc89 100644 --- a/tests/unit/responses/tool_search/test_lowering.py +++ b/tests/unit/responses/tool_search/test_lowering.py @@ -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():