mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
96ef4310fb
commit
f2b332f769
6 changed files with 111 additions and 39 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue