This commit is contained in:
Idder Ghanbaja 2026-10-04 17:40:55 +02:00 • committed by GitHub
commit 7c9677e351
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 1371 additions and 47 deletions

View file

@ -37,6 +37,23 @@ REDIS_FAMILIES_METADATA_KEY: Final = "families"
_service_caller: Final[ContextVar[str | None]] = ContextVar("service_caller", default=None)
_emulated_file_search: Final[ContextVar[bool]] = ContextVar("emulated_file_search", default=False)
@contextmanager
def emulated_file_search_phase() -> Generator[None]:
"""Nested calls of emulated file_search, whose answer keeps only its own tool calls."""
token: Final = _emulated_file_search.set(True)
try:
yield
finally:
_emulated_file_search.reset(token)
def in_emulated_file_search() -> bool:
return _emulated_file_search.get()
@contextmanager
def post_response_phase() -> Generator[None]:
"""Work the caller no longer waits for (success callbacks, response-cache writes), including tasks it spawns."""

View file

@ -66,6 +66,9 @@ class BaseResponsesAPIConfig(ABC):
def supports_encrypted_agent_messages(self) -> bool:
return False
def supports_native_tool_search(self) -> bool:
return False
def sign_request(
self,
headers: dict,

View file

@ -116,6 +116,9 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
def supports_encrypted_agent_messages(self) -> bool:
return self.custom_llm_provider in (LlmProviders.OPENAI, LlmProviders.AZURE)
def supports_native_tool_search(self) -> bool:
return self.custom_llm_provider in (LlmProviders.OPENAI, LlmProviders.AZURE, LlmProviders.CHATGPT)
@staticmethod
def _is_gpt_5_model(model: str) -> bool:
"""Return True only for actual OpenAI GPT-5 models.

View file

@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, Any, Final, TypeAlias, cast # noqa: TID251 #
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._internal_context import is_internal_call
from litellm._internal_context import emulated_file_search_phase, is_internal_call
from litellm._logging import verbose_logger
from litellm.types.llms.openai import ResponseOutputItem, ResponsesAPIResponse
from litellm.types.vector_stores import VectorStoreSearchResult
@ -537,15 +537,16 @@ async def aresponses_with_emulated_file_search(
_prev_internal: Final = is_internal_call.get()
is_internal_call.set(True)
try:
first_response: Final[ResponsesAPIResponse] = cast(
ResponsesAPIResponse,
await _call_aresponses(
input=input,
model=model,
tools=transformed_tools or None,
**call_kwargs,
),
)
with emulated_file_search_phase():
first_response: Final[ResponsesAPIResponse] = cast(
ResponsesAPIResponse,
await _call_aresponses(
input=input,
model=model,
tools=transformed_tools or None,
**call_kwargs,
),
)
finally:
is_internal_call.set(_prev_internal)
@ -601,15 +602,16 @@ async def aresponses_with_emulated_file_search(
# Also an internal sub-call; billing is suppressed so the outer call fires once.
is_internal_call.set(True)
try:
final_response: Final[ResponsesAPIResponse] = cast(
ResponsesAPIResponse,
await _call_aresponses(
input=follow_up_input,
model=model,
tools=None, # no tools needed for the answer step
**call_kwargs,
),
)
with emulated_file_search_phase():
final_response: Final[ResponsesAPIResponse] = cast(
ResponsesAPIResponse,
await _call_aresponses(
input=follow_up_input,
model=model,
tools=None, # no tools needed for the answer step
**call_kwargs,
),
)
finally:
is_internal_call.set(_prev_internal)

View file

@ -13,6 +13,7 @@ from pydantic import BaseModel, TypeAdapter, ValidationError
from typing_extensions import assert_never
import litellm
from litellm._internal_context import in_emulated_file_search
from litellm._logging import verbose_logger
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
@ -34,6 +35,7 @@ from litellm.responses.litellm_completion_transformation.handler import (
LiteLLMCompletionTransformationHandler,
)
from litellm.responses.mcp.request_context import MCPRequestContext
from litellm.responses.tool_search.lowering import needs_tool_search_lowering
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
PromptObject,
@ -1043,13 +1045,8 @@ def _responses_try_dispatch_mcp_gateway(
return run_async_function(aresponses_api_with_mcp, **mcp_call_kwargs)
def _responses_try_dispatch_emulated_file_search(
def _responses_reentry_kwargs(
*,
tools: Iterable[ToolParam] | None,
input: str | ResponseInputParam,
model: str,
responses_api_provider_config: BaseResponsesAPIConfig | None,
use_chat_completions_api: bool,
include: list[ResponseIncludable] | None,
instructions: str | None,
max_output_tokens: int | None,
@ -1076,22 +1073,12 @@ def _responses_try_dispatch_emulated_file_search(
extra_body: dict[str, object] | None,
timeout: float | httpx.Timeout | None,
custom_llm_provider: str | None,
use_chat_completions_api: bool,
kwargs: dict[str, object],
_is_async: bool,
) -> ResponsesAPIResponse | Coroutine[object, object, ResponsesAPIResponse] | None:
"""Return a response when emulated file_search handles the call; otherwise None."""
if not _has_file_search_tool(tools) or not (
responses_api_provider_config is None
or use_chat_completions_api is True
or not responses_api_provider_config.supports_native_file_search()
):
return None
from litellm.responses.file_search.emulated_handler import (
aresponses_with_emulated_file_search,
)
) -> dict[str, object]:
"""The keyword arguments an emulation passes when it calls aresponses again for this request."""
_internal_skip: Final = {"litellm_call_id", "aresponses"}
emulated_kwargs: Final = {
return {
"include": include,
"instructions": instructions,
"max_output_tokens": max_output_tokens,
@ -1121,14 +1108,94 @@ def _responses_try_dispatch_emulated_file_search(
**({"use_chat_completions_api": True} if use_chat_completions_api else {}),
**{k: v for k, v in kwargs.items() if k not in _internal_skip},
}
def _responses_try_dispatch_emulated_file_search(
*,
tools: Iterable[ToolParam] | None,
input: str | ResponseInputParam,
model: str,
responses_api_provider_config: BaseResponsesAPIConfig | None,
use_chat_completions_api: bool,
reentry_kwargs: Mapping[str, object],
_is_async: bool,
) -> ResponsesAPIResponse | Coroutine[object, object, ResponsesAPIResponse] | None:
"""Return a response when emulated file_search handles the call; otherwise None."""
if not _has_file_search_tool(tools) or not (
responses_api_provider_config is None
or use_chat_completions_api is True
or not responses_api_provider_config.supports_native_file_search()
):
return None
from litellm.responses.file_search.emulated_handler import (
aresponses_with_emulated_file_search,
)
if _is_async:
return aresponses_with_emulated_file_search(input=input, model=model, tools=tools, **emulated_kwargs)
return aresponses_with_emulated_file_search(input=input, model=model, tools=tools, **reentry_kwargs)
return run_async_function(
aresponses_with_emulated_file_search,
input=input,
model=model,
tools=tools,
**emulated_kwargs,
**reentry_kwargs,
)
class _DeploymentToolSearchSupport(BaseModel):
supports_tool_search: bool | None = None
def _supports_tool_search_natively(responses_api_provider_config: BaseResponsesAPIConfig, model_info: object) -> bool:
try:
declared: Final = _DeploymentToolSearchSupport.model_validate(model_info).supports_tool_search
except ValidationError:
return responses_api_provider_config.supports_native_tool_search()
if declared is None:
return responses_api_provider_config.supports_native_tool_search()
return declared
def _responses_try_dispatch_lowered_tool_search(
*,
tools: Iterable[ToolParam] | None,
input: str | ResponseInputParam,
model: str,
tool_choice: ToolChoice | None,
responses_api_provider_config: BaseResponsesAPIConfig | None,
use_chat_completions_api: bool,
model_info: object,
reentry_kwargs: Mapping[str, object],
_is_async: bool,
) -> (
ResponsesAPIResponse
| BaseResponsesAPIStreamingIterator
| Coroutine[object, object, ResponsesAPIResponse | BaseResponsesAPIStreamingIterator]
| None
):
"""Return a response when a client tool_search is lowered to a function tool for a native Responses
deployment that has no tool search of its own; otherwise None. The chat-completions bridge handles
tool_search itself."""
declared_tools: Final = tuple(tools) if isinstance(tools, Sequence) else None
if (
responses_api_provider_config is None
or _bridges_to_chat_completions(responses_api_provider_config, use_chat_completions_api)
or _supports_tool_search_natively(responses_api_provider_config, model_info)
or not needs_tool_search_lowering(input, declared_tools)
or in_emulated_file_search()
):
return None
from litellm.responses.tool_search.handler import (
aresponses_with_lowered_tool_search,
responses_with_lowered_tool_search,
)
if _is_async:
return aresponses_with_lowered_tool_search(
input=input, model=model, tools=declared_tools, tool_choice=tool_choice, call_kwargs=reentry_kwargs
)
return responses_with_lowered_tool_search(
input=input, model=model, tools=declared_tools, tool_choice=tool_choice, call_kwargs=reentry_kwargs
)
@ -1336,12 +1403,7 @@ def responses(
)
)
_file_search_dispatch: Final = _responses_try_dispatch_emulated_file_search(
tools=tools,
input=input,
model=model,
responses_api_provider_config=responses_api_provider_config,
use_chat_completions_api=use_chat_completions_api,
reentry_kwargs: Final = _responses_reentry_kwargs(
include=include,
instructions=instructions,
max_output_tokens=max_output_tokens,
@ -1368,12 +1430,35 @@ def responses(
extra_body=extra_body,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
use_chat_completions_api=use_chat_completions_api,
kwargs=kwargs,
)
_file_search_dispatch: Final = _responses_try_dispatch_emulated_file_search(
tools=tools,
input=input,
model=model,
responses_api_provider_config=responses_api_provider_config,
use_chat_completions_api=use_chat_completions_api,
reentry_kwargs=reentry_kwargs,
_is_async=_is_async,
)
if _file_search_dispatch is not None:
return _file_search_dispatch
_tool_search_dispatch: Final = _responses_try_dispatch_lowered_tool_search(
tools=tools,
input=input,
model=model,
tool_choice=tool_choice,
responses_api_provider_config=responses_api_provider_config,
use_chat_completions_api=use_chat_completions_api,
model_info=deployment_model_info,
reentry_kwargs=reentry_kwargs,
_is_async=_is_async,
)
if _tool_search_dispatch is not None:
return _tool_search_dispatch
if _bridges_to_chat_completions(responses_api_provider_config, use_chat_completions_api):
bridge_kwargs: Final = _bridge_kwargs(kwargs, responses_api_provider_config, allowed_openai_params)
return litellm_completion_transformation_handler.response_api_handler(

View file

@ -0,0 +1,90 @@
from collections.abc import Mapping, Sequence
from typing import Final
from typing_extensions import assert_never
import litellm
from litellm._internal_context import is_internal_call
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
from litellm.responses.tool_search.lifting import lift_tool_search_calls, lift_tool_search_stream
from litellm.responses.tool_search.lowering import (
TOOL_SEARCH_FUNCTION_NAME,
LoweredToolSearchRequest,
ToolSearchFunctionNameTaken,
lower_tool_search_request,
)
from litellm.types.llms.openai import ResponseInputParam, ResponsesAPIResponse, ToolChoice
def _lowered_call_kwargs(
input: str | ResponseInputParam,
model: str,
tools: Sequence[object] | None,
tool_choice: ToolChoice | None,
call_kwargs: Mapping[str, object],
) -> dict[str, object]:
custom_llm_provider: Final = call_kwargs.get("custom_llm_provider")
lowering: Final = lower_tool_search_request(input=input, tools=tools, tool_choice=tool_choice)
match lowering:
case LoweredToolSearchRequest():
return {
**call_kwargs,
"model": model,
"input": lowering.input,
"tools": list(lowering.tools),
"tool_choice": lowering.tool_choice,
}
case ToolSearchFunctionNameTaken():
raise litellm.BadRequestError(
message=(
f"A function tool named '{TOOL_SEARCH_FUNCTION_NAME}' can't be declared alongside a client "
"tool_search tool, because this deployment receives tool_search as that function"
),
model=model,
llm_provider=custom_llm_provider if isinstance(custom_llm_provider, str) else "",
)
return assert_never(lowering)
def _lifted(result: object) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator:
if isinstance(result, BaseResponsesAPIStreamingIterator):
return lift_tool_search_stream(result)
if isinstance(result, ResponsesAPIResponse):
return lift_tool_search_calls(result)
raise TypeError(f"Unexpected Responses API result: {type(result).__name__}")
async def aresponses_with_lowered_tool_search(
input: str | ResponseInputParam,
model: str,
tools: Sequence[object] | None,
tool_choice: ToolChoice | None,
call_kwargs: Mapping[str, object],
) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator:
from litellm.responses.dispatch import aresponses
lowered_kwargs: Final = _lowered_call_kwargs(input, model, tools, tool_choice, call_kwargs)
internal_call: Final = is_internal_call.set(True)
try:
result: Final = await aresponses(**lowered_kwargs)
finally:
is_internal_call.reset(internal_call)
return _lifted(result)
def responses_with_lowered_tool_search(
input: str | ResponseInputParam,
model: str,
tools: Sequence[object] | None,
tool_choice: ToolChoice | None,
call_kwargs: Mapping[str, object],
) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator:
from litellm.responses.dispatch import responses
lowered_kwargs: Final = _lowered_call_kwargs(input, model, tools, tool_choice, call_kwargs)
internal_call: Final = is_internal_call.set(True)
try:
result: Final = responses(**lowered_kwargs)
finally:
is_internal_call.reset(internal_call)
return _lifted(result)

View file

@ -0,0 +1,191 @@
import json
from collections.abc import AsyncIterator, Iterator, Sequence
from typing import Final, Literal, Protocol
from openai._streaming import ServerSentEvent
from openai.types.responses import ResponseToolSearchCall
from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
CachedResponsesAPIStreamingIterator,
MockResponsesAPIStreamingIterator,
ResponsesAPIStreamingIterator,
SyncResponsesAPIStreamingIterator,
)
from litellm.responses.tool_search.lowering import TOOL_SEARCH_FUNCTION_NAME
from litellm.types.llms.openai import (
ResponseCompletedEvent,
ResponseIncompleteEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
_TOOL_SEARCH_CALL_ID_PREFIX: Final = "tsc_"
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
_ITEM_EVENTS: Final = frozenset({ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE})
_ARGUMENT_EVENTS: Final = frozenset(
{ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE}
)
_TERMINAL_EVENTS: Final = frozenset(
{ResponsesAPIStreamEvents.RESPONSE_COMPLETED, ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE}
)
class _FunctionCall(BaseModel):
type: Literal["function_call"]
name: str
namespace: str | None = None
id: str | None = None
call_id: str
arguments: str = ""
class _Event(BaseModel):
type: str = ""
item_id: str | None = None
item: object = None
response: dict[str, object] | None = None
class _Response(BaseModel):
output: tuple[object, ...] = ()
class _HasOutput(Protocol):
@property
def output(self) -> Sequence[object]: ...
def _tool_search_function_call(item: object) -> _FunctionCall | None:
plain: Final = item.model_dump(exclude_none=True) if isinstance(item, BaseModel) else item
try:
call: Final = _FunctionCall.model_validate(plain)
except ValidationError:
return None
return call if call.name == TOOL_SEARCH_FUNCTION_NAME and not call.namespace else None
def _search_arguments(arguments: str) -> dict[str, object]:
try:
return _JSON_OBJECT.validate_json(arguments)
except ValidationError:
verbose_logger.warning(
"tool_search was called with arguments that aren't a JSON object, using them as the query"
)
return {"query": arguments}
def _tool_search_call_item_id(call: _FunctionCall) -> str:
source: Final = call.id or call.call_id
return f"{_TOOL_SEARCH_CALL_ID_PREFIX}{source.partition('_')[2] or source}"
def _tool_search_call(call: _FunctionCall, status: Literal["in_progress", "completed"]) -> dict[str, object]:
return {
"type": "tool_search_call",
"id": _tool_search_call_item_id(call),
"call_id": call.call_id,
"execution": "client",
"status": status,
"arguments": {} if status == "in_progress" else _search_arguments(call.arguments),
}
def _lifted_output_item(item: object) -> object:
call: Final = _tool_search_function_call(item)
if call is None:
return item
return ResponseToolSearchCall.model_validate(_tool_search_call(call, "completed"))
def _output_items(response: _HasOutput) -> tuple[object, ...]:
return tuple(response.output)
def lift_tool_search_calls(response: ResponsesAPIResponse) -> ResponsesAPIResponse:
output: Final = _output_items(response)
lifted: Final = [_lifted_output_item(item) for item in output]
if all(new is old for new, old in zip(lifted, output, strict=True)):
return response
return response.model_copy(update={"output": lifted})
def _lifted_output_dict(item: object) -> object:
call: Final = _tool_search_function_call(item)
return item if call is None else _tool_search_call(call, "completed")
def _lifted_response_dict(response: dict[str, object]) -> dict[str, object]:
output: Final = _Response.model_validate(response).output
return {**response, "output": [_lifted_output_dict(item) for item in output]}
class ToolSearchEventLifter:
def __init__(self) -> None:
self._tool_search_item_ids: frozenset[str] = frozenset()
def lift(self, data: str) -> str | None:
try:
raw: Final = _JSON_OBJECT.validate_json(data)
except ValidationError:
return data
event: Final = _Event.model_validate(raw)
if event.type in _ARGUMENT_EVENTS and event.item_id in self._tool_search_item_ids:
return None
if event.type in _TERMINAL_EVENTS and event.response is not None:
return json.dumps({**raw, "response": _lifted_response_dict(event.response)})
call: Final = _tool_search_function_call(event.item) if event.type in _ITEM_EVENTS else None
if call is None:
return data
if call.id is not None:
self._tool_search_item_ids = self._tool_search_item_ids | {call.id}
status: Final = "in_progress" if event.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED else "completed"
return json.dumps({**raw, "item": _tool_search_call(call, status)})
def _lifted_event(event: ServerSentEvent, lifter: ToolSearchEventLifter) -> ServerSentEvent | None:
data: Final = lifter.lift(event.data)
if data is None:
return None
return ServerSentEvent(event=event.event, data=data, id=event.id, retry=event.retry)
async def _alifted_events(
events: AsyncIterator[ServerSentEvent], lifter: ToolSearchEventLifter
) -> AsyncIterator[ServerSentEvent]:
async for event in events:
if (lifted := _lifted_event(event, lifter)) is not None:
yield lifted
def _lifted_events(events: Iterator[ServerSentEvent], lifter: ToolSearchEventLifter) -> Iterator[ServerSentEvent]:
lifted_events: Final = (_lifted_event(event, lifter) for event in events)
return (lifted for lifted in lifted_events if lifted is not None)
def _lift_synthetic_events(stream: MockResponsesAPIStreamingIterator | CachedResponsesAPIStreamingIterator) -> None:
terminal_event: Final = stream.completed_response
if not isinstance(terminal_event, (ResponseCompletedEvent, ResponseIncompleteEvent)):
return
stream._set_events_from_response( # pyright: ignore[reportPrivateUsage] # the only way to rebuild a replayed stream
transformed=lift_tool_search_calls(terminal_event.response), logging_obj=stream.logging_obj
)
def lift_tool_search_stream(stream: BaseResponsesAPIStreamingIterator) -> BaseResponsesAPIStreamingIterator:
match stream:
case ResponsesAPIStreamingIterator():
stream.stream_iterator = _alifted_events( # rebind-ok: lift events before the stream logs its own output
stream.stream_iterator, ToolSearchEventLifter()
)
case SyncResponsesAPIStreamingIterator():
stream.stream_iterator = _lifted_events( # rebind-ok: lift events before the stream logs its own output
stream.stream_iterator, ToolSearchEventLifter()
)
case MockResponsesAPIStreamingIterator() | CachedResponsesAPIStreamingIterator():
_lift_synthetic_events(stream)
case _:
verbose_logger.debug("tool_search lifting skipped for %s", type(stream).__name__)
return stream

View file

@ -0,0 +1,238 @@
import json
from collections.abc import Sequence
from dataclasses import dataclass
from itertools import chain
from typing import Final, Literal, TypeAlias, TypeVar
from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm.types.llms.openai import ResponseInputParam, ToolChoice
TOOL_SEARCH_FUNCTION_NAME: Final = "tool_search"
_DEFAULT_TOOL_SEARCH_DESCRIPTION: Final = (
"Search the client tool catalog and load the matching tools for the next call."
)
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
# a replayed search may only load tools the client runs itself, never hosted ones the proxy would run
_LOADABLE_TOOL_TYPES: Final = frozenset({"function", "custom", "namespace"})
_LOADABLE_MEMBER_TYPES: Final = frozenset({"function", "custom"})
_ModelT = TypeVar("_ModelT", bound=BaseModel)
class _Typed(BaseModel):
type: str = ""
name: str | None = None
class _ClientToolSearchDeclaration(BaseModel):
type: Literal["tool_search"]
execution: Literal["client"]
description: str | None = None
parameters: object = None
class _ReplayedToolSearchCall(BaseModel):
type: Literal["tool_search_call"]
call_id: str | None = None
arguments: object = None
status: str | None = None
class _ReplayedToolSearchOutput(BaseModel):
type: Literal["tool_search_output"]
call_id: str | None = None
tools: tuple[dict[str, object], ...] = ()
class _Namespace(BaseModel):
type: Literal["namespace"]
tools: tuple[dict[str, object], ...] = ()
@dataclass(frozen=True, slots=True)
class LoweredToolSearchRequest:
input: str | Sequence[object]
tools: tuple[object, ...]
tool_choice: ToolChoice | None
@dataclass(frozen=True, slots=True)
class ToolSearchFunctionNameTaken:
pass
ToolSearchLowering: TypeAlias = LoweredToolSearchRequest | ToolSearchFunctionNameTaken
def _plain(value: object) -> object:
return value.model_dump(exclude_none=True) if isinstance(value, BaseModel) else value
def _parsed(model: type[_ModelT], value: object) -> _ModelT | None:
try:
return model.model_validate(_plain(value))
except ValidationError:
return None
def _json_object(value: object) -> dict[str, object] | None:
try:
return _JSON_OBJECT.validate_python(_plain(value))
except ValidationError:
return None
def _is_tool_search_item(item: object) -> bool:
return _parsed(_ReplayedToolSearchCall, item) is not None or _parsed(_ReplayedToolSearchOutput, item) is not None
def _declares_client_tool_search(tools: Sequence[object]) -> bool:
return any(_parsed(_ClientToolSearchDeclaration, tool) is not None for tool in tools)
def needs_tool_search_lowering(input: str | ResponseInputParam, tools: Sequence[object] | None) -> bool:
if _declares_client_tool_search(tools or ()):
return True
return not isinstance(input, str) and any(_is_tool_search_item(item) for item in input)
def _object_schema(parameters: object) -> dict[str, object]:
schema: Final = _json_object(parameters)
if schema is None:
return {"type": "object", "properties": {}}
if "type" in schema:
return schema
return {"type": "object", **schema}
def _lowered_tool(tool: object) -> object:
declaration: Final = _parsed(_ClientToolSearchDeclaration, tool)
if declaration is None:
return tool
return {
"type": "function",
"name": TOOL_SEARCH_FUNCTION_NAME,
"description": declaration.description or _DEFAULT_TOOL_SEARCH_DESCRIPTION,
"parameters": (
{"type": "object", "properties": {"query": {"type": "string"}}, "required": ["query"]}
if declaration.parameters is None
else _object_schema(declaration.parameters)
),
"strict": False,
}
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 _has_type(tool: object, types: frozenset[str]) -> bool:
kind: Final = _parsed(_Typed, tool)
return kind is not None and kind.type in types
def _loaded_tool(tool: dict[str, object]) -> dict[str, object]:
namespace: Final = _parsed(_Namespace, tool)
if namespace is None:
return _loaded_definition(tool)
members: Final = [
_loaded_definition(member) for member in namespace.tools if _has_type(member, _LOADABLE_MEMBER_TYPES)
]
return {**_loaded_definition(tool), "tools": members}
def _loadable_tools(output: _ReplayedToolSearchOutput) -> tuple[dict[str, object], ...]:
return tuple(_loaded_tool(tool) for tool in output.tools if _has_type(tool, _LOADABLE_TOOL_TYPES))
def _visible_loaded_tool(tool: dict[str, object]) -> dict[str, object]:
if _parsed(_Namespace, tool) is None:
return tool
return {key: value for key, value in tool.items() if key in ("type", "name", "description")}
def _lowered_item(item: object) -> object:
call: Final = _parsed(_ReplayedToolSearchCall, item)
if call is not None:
return {
"type": "function_call",
"call_id": call.call_id,
"name": TOOL_SEARCH_FUNCTION_NAME,
"arguments": json.dumps({} if call.arguments is None else call.arguments),
**({} if call.status is None else {"status": call.status}),
}
output: Final = _parsed(_ReplayedToolSearchOutput, item)
if output is None:
return item
visible_tools: Final = [_visible_loaded_tool(tool) for tool in _loadable_tools(output)]
return {
"type": "function_call_output",
"call_id": output.call_id,
"output": json.dumps({"tools": visible_tools}, separators=(",", ":")),
}
def _loaded_tools(items: Sequence[object]) -> tuple[dict[str, object], ...]:
outputs: Final = tuple(_parsed(_ReplayedToolSearchOutput, item) for item in items)
return tuple(chain.from_iterable(_loadable_tools(output) for output in outputs if output is not None))
def _merge_key(index: int, tool: object) -> tuple[str, str]:
kind: Final = _parsed(_Typed, tool)
if kind is not None and kind.type in ("function", "custom", "namespace") and kind.name is not None:
return kind.type, kind.name
return "position", str(index)
def _namespace_members(group: Sequence[object]) -> tuple[dict[str, object], ...]:
namespaces: Final = (_parsed(_Namespace, tool) for tool in group)
return tuple(chain.from_iterable(namespace.tools for namespace in namespaces if namespace is not None))
def _merged_group(kind: str, group: Sequence[object]) -> object:
latest: Final = group[-1]
latest_object: Final = _json_object(latest) if kind == "namespace" and len(group) > 1 else None
if latest_object is None:
return latest
return {**latest_object, "tools": list(_merged_tools(_namespace_members(group)))}
def _merged_tools(tools: Sequence[object]) -> tuple[object, ...]:
keys: Final = tuple(_merge_key(index, tool) for index, tool in enumerate(tools))
keyed: Final = tuple(zip(keys, tools, strict=True))
return tuple(
_merged_group(key[0], tuple(tool for tool_key, tool in keyed if tool_key == key)) for key in dict.fromkeys(keys)
)
def _lowered_tool_choice(tool_choice: ToolChoice | None) -> ToolChoice | None:
kind: Final = None if isinstance(tool_choice, str) else _parsed(_Typed, tool_choice)
if kind is None or kind.type != "tool_search":
return tool_choice
return {"type": "function", "name": TOOL_SEARCH_FUNCTION_NAME}
def _declares_function(tools: Sequence[object] | None, name: str) -> bool:
return ("function", name) in (_merge_key(index, tool) for index, tool in enumerate(tools or ()))
def lower_tool_search_request(
input: str | ResponseInputParam,
tools: Sequence[object] | None,
tool_choice: ToolChoice | None,
) -> ToolSearchLowering:
declared: Final = tuple(tools or ())
if _declares_client_tool_search(declared) and _declares_function(declared, TOOL_SEARCH_FUNCTION_NAME):
return ToolSearchFunctionNameTaken()
items: Final = () if isinstance(input, str) else tuple(input)
lowered_tools: Final = _merged_tools((*(_lowered_tool(tool) for tool in declared), *_loaded_tools(items)))
return LoweredToolSearchRequest(
input=input if isinstance(input, str) else [_lowered_item(item) for item in items],
tools=lowered_tools,
tool_choice=_lowered_tool_choice(tool_choice),
)

View file

@ -9,7 +9,12 @@ import pytest
import litellm
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.azure.responses.transformation import AzureOpenAIResponsesAPIConfig
from litellm.llms.chatgpt.responses.transformation import ChatGPTResponsesAPIConfig
from litellm.llms.hosted_vllm.responses.transformation import HostedVLLMResponsesAPIConfig
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.llms.openai_like.responses.transformation import OpenAILikeResponsesConfig
from litellm.llms.openrouter.responses.transformation import OpenRouterResponsesAPIConfig
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
from litellm.responses.litellm_completion_transformation.transformation import LiteLLMCompletionResponsesConfig
from litellm.types.llms.openai import (
ImageGenerationPartialImageEvent,
@ -2322,3 +2327,19 @@ class TestReasoningFollowsModelSupport:
drop_params=True,
)
assert mapped["reasoning"] == {"effort": "medium"}
@pytest.mark.parametrize(
("config", "expected"),
[
(OpenAIResponsesAPIConfig(), True),
(AzureOpenAIResponsesAPIConfig(), True),
(ChatGPTResponsesAPIConfig(), True),
(HostedVLLMResponsesAPIConfig(), False),
(OpenRouterResponsesAPIConfig(), False),
(XAIResponsesAPIConfig(), False),
(OpenAILikeResponsesConfig(), False),
],
)
def test_only_openai_backends_run_tool_search_natively(config: BaseResponsesAPIConfig, expected: bool):
assert config.supports_native_tool_search() is expected

View file

@ -0,0 +1,241 @@
import asyncio
import json
from typing import Final
import httpx
import pytest
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
CLIENT_TOOL_SEARCH: Final = {
"type": "tool_search",
"execution": "client",
"description": "Search the deferred tools",
"parameters": {"type": "object", "properties": {"query": {"type": "string"}}, "required": ["query"]},
}
SEARCH_FUNCTION_CALL: Final = {
"type": "function_call",
"id": "fc_search",
"call_id": "call_search",
"name": "tool_search",
"arguments": '{"query": "calendar"}',
"status": "completed",
}
UPSTREAM_RESPONSE: Final = {
"id": "resp_upstream",
"object": "response",
"created_at": 1,
"model": "qwen",
"status": "completed",
"output": [SEARCH_FUNCTION_CALL],
"tools": [],
"parallel_tool_calls": True,
"tool_choice": "auto",
"usage": {"input_tokens": 12, "output_tokens": 6, "total_tokens": 18},
}
class _Upstream:
def __init__(self) -> None:
self.bodies: list[dict[str, object]] = []
def __call__(self, request: httpx.Request) -> httpx.Response:
self.bodies.append(json.loads(request.content))
return httpx.Response(200, json=UPSTREAM_RESPONSE)
class _SuccessLog(CustomLogger):
def __init__(self, call_id: str) -> None:
super().__init__()
self.call_id: Final = call_id
self.outputs: list[list[object]] = []
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if kwargs.get("litellm_call_id") == self.call_id:
self.outputs.append(list(response_obj.output))
def _types(tools: object) -> list[object]:
assert isinstance(tools, list)
return [tool.get("type") for tool in tools]
def _output_types(output: object) -> list[object]:
assert isinstance(output, list)
return [item.get("type") if isinstance(item, dict) else getattr(item, "type", None) for item in output]
async def _hosted_vllm_call(upstream: _Upstream, **overrides: object):
request: Final = {"input": "find a calendar tool", "tools": [CLIENT_TOOL_SEARCH], **overrides}
return await litellm.aresponses(
model="hosted_vllm/qwen",
api_base="http://vllm.test/v1",
client=AsyncHTTPHandler(transport=httpx.MockTransport(upstream)),
**request,
)
@pytest.mark.asyncio
async def test_hosted_vllm_receives_a_function_and_the_client_a_tool_search_call():
upstream: Final = _Upstream()
response: Final = await _hosted_vllm_call(upstream)
assert _types(upstream.bodies[0]["tools"]) == ["function"]
assert _output_types(response.output) == ["tool_search_call"]
assert response.output[0].arguments == {"query": "calendar"}
@pytest.mark.asyncio
async def test_a_streamed_hosted_vllm_search_reaches_the_client_as_tool_search_call_events():
upstream: Final = _Upstream()
stream: Final = await _hosted_vllm_call(upstream, stream=True)
events: Final = [event async for event in stream]
item_events: Final = [event for event in events if getattr(event, "item", None) is not None]
assert [event.item.type for event in item_events] == ["tool_search_call", "tool_search_call"]
assert not [event for event in events if "function_call_arguments" in str(event.type.value)]
@pytest.mark.asyncio
async def test_the_next_turn_replays_the_search_as_function_items_with_the_loaded_tools():
upstream: Final = _Upstream()
loaded_namespace: Final = {
"type": "namespace",
"name": "calendar",
"description": "Calendar tools",
"tools": [{"type": "function", "name": "create_event", "defer_loading": True, "parameters": {}}],
}
await _hosted_vllm_call(
upstream,
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],
},
],
)
sent: Final = upstream.bodies[0]
assert _output_types(sent["input"])[1:] == ["function_call", "function_call_output"]
assert _types(sent["tools"]) == ["function", "namespace"]
@pytest.mark.asyncio
async def test_a_native_openai_deployment_keeps_tool_search():
upstream: Final = _Upstream()
await litellm.aresponses(
model="openai/gpt-tool-search-test",
api_key="test-key",
api_base="http://openai.test/v1",
input="find a calendar tool",
tools=[CLIENT_TOOL_SEARCH],
client=AsyncHTTPHandler(transport=httpx.MockTransport(upstream)),
)
assert _types(upstream.bodies[0]["tools"]) == ["tool_search"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("model", "supports_tool_search", "expected_tool_type"),
[
("openai/gpt-tool-search-test", False, "function"),
("hosted_vllm/qwen", True, "tool_search"),
],
)
async def test_the_deployment_model_info_decides_whether_tool_search_is_native(
model: str, supports_tool_search: bool, expected_tool_type: str
):
upstream: Final = _Upstream()
await litellm.aresponses(
model=model,
api_key="test-key",
api_base="http://upstream.test/v1",
input="find a calendar tool",
tools=[CLIENT_TOOL_SEARCH],
model_info={"supports_tool_search": supports_tool_search},
client=AsyncHTTPHandler(transport=httpx.MockTransport(upstream)),
)
assert _types(upstream.bodies[0]["tools"]) == [expected_tool_type]
@pytest.mark.asyncio
async def test_a_client_function_named_tool_search_is_rejected_before_any_upstream_call():
upstream: Final = _Upstream()
with pytest.raises(litellm.BadRequestError, match="named 'tool_search' can't be declared alongside"):
await _hosted_vllm_call(
upstream, tools=[CLIENT_TOOL_SEARCH, {"type": "function", "name": "tool_search", "parameters": {}}]
)
assert upstream.bodies == []
@pytest.mark.asyncio
async def test_emulated_file_search_leaves_client_tool_search_as_declared():
upstream: Final = _Upstream()
await litellm.aresponses(
model="fireworks_ai/qwen",
api_key="test-key",
api_base="http://fireworks.test/v1",
input="find a calendar tool",
tools=[CLIENT_TOOL_SEARCH, {"type": "file_search", "vector_store_ids": ["vs_lab"]}],
client=AsyncHTTPHandler(transport=httpx.MockTransport(upstream)),
)
assert _types(upstream.bodies[0]["tools"]) == ["tool_search", "function"]
@pytest.mark.asyncio
async def test_a_client_function_named_like_the_file_search_emulation_keeps_tool_search_lowered():
upstream: Final = _Upstream()
own_function: Final = {
"type": "function",
"name": "litellm_file_search",
"parameters": {"type": "object", "properties": {}},
}
await _hosted_vllm_call(upstream, tools=[CLIENT_TOOL_SEARCH, own_function])
assert [tool["name"] for tool in upstream.bodies[0]["tools"]] == ["tool_search", "litellm_file_search"]
@pytest.mark.asyncio
async def test_the_call_is_logged_once_with_the_tool_search_call(monkeypatch: pytest.MonkeyPatch):
success_log: Final = _SuccessLog("tool-search-logged-once")
monkeypatch.setattr(litellm, "callbacks", [success_log])
await _hosted_vllm_call(_Upstream(), litellm_call_id="tool-search-logged-once")
await asyncio.sleep(0)
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0)
assert [_output_types(output) for output in success_log.outputs] == [["tool_search_call"]]
def test_the_sync_api_lowers_and_lifts_too():
upstream: Final = _Upstream()
response: Final = litellm.responses(
model="hosted_vllm/qwen",
api_base="http://vllm.test/v1",
input="find a calendar tool",
tools=[CLIENT_TOOL_SEARCH],
client=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(upstream))),
)
assert _types(upstream.bodies[0]["tools"]) == ["function"]
assert _output_types(response.output) == ["tool_search_call"]

View file

@ -0,0 +1,213 @@
import json
from datetime import datetime
from typing import Final
import httpx
import pytest
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 (
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, ResponsesAPIStreamEvents
SEARCH_CALL: Final = {
"type": "function_call",
"id": "fc_search",
"call_id": "call_search",
"name": "tool_search",
"arguments": '{"query": "calendar create", "limit": 1}',
"status": "completed",
}
SHELL_CALL: Final = {
"type": "function_call",
"id": "fc_shell",
"call_id": "call_shell",
"name": "exec_command",
"arguments": '{"cmd": "ls"}',
"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]:
return {
"id": "resp_1",
"object": "response",
"created_at": 1,
"model": "qwen",
"status": "completed",
"output": list(output),
"tools": [],
"parallel_tool_calls": True,
"tool_choice": "auto",
}
def _event(event_type: str, **fields: object) -> dict[str, object]:
return {"type": event_type, **fields}
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.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),
_event("response.output_item.added", output_index=1, item={**SHELL_CALL, "arguments": ""}),
_event("response.function_call_arguments.delta", item_id="fc_shell", output_index=1, delta='{"cmd": "ls"}'),
_event("response.output_item.done", output_index=1, item=SHELL_CALL),
_event("response.completed", response=_response(SEARCH_CALL, SHELL_CALL)),
]
def _lifted_search_call() -> dict[str, object]:
return {
"type": "tool_search_call",
"id": "tsc_search",
"call_id": "call_search",
"execution": "client",
"status": "completed",
"arguments": {"query": "calendar create", "limit": 1},
}
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, MESSAGE)
)
lifted: Final = lift_tool_search_calls(response)
assert lifted.output[0] == ResponseToolSearchCall.model_validate(_lifted_search_call())
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():
response: Final = ResponsesAPIResponse.model_validate(_response(SHELL_CALL))
assert lift_tool_search_calls(response) is response
def test_search_arguments_that_are_not_a_json_object_become_the_query():
response: Final = ResponsesAPIResponse.model_validate(_response({**SEARCH_CALL, "arguments": "calendar"}))
lifted_item: Final = lift_tool_search_calls(response).output[0]
assert isinstance(lifted_item, ResponseToolSearchCall)
assert lifted_item.arguments == {"query": "calendar"}
def test_stream_events_of_a_search_call_are_lifted_and_its_argument_deltas_dropped():
lifter: Final = ToolSearchEventLifter()
lifted: Final = [lifter.lift(json.dumps(event)) for event in STREAM_EVENTS]
kept: Final = [json.loads(data) for data in lifted if data is not None]
assert [event["type"] for event in kept] == [
"response.created",
"response.output_item.added",
"response.output_item.done",
"response.output_item.added",
"response.function_call_arguments.delta",
"response.output_item.done",
"response.completed",
]
assert kept[1]["item"] == {**_lifted_search_call(), "status": "in_progress", "arguments": {}}
assert kept[2]["item"] == _lifted_search_call()
assert kept[3]["item"]["name"] == "exec_command"
assert kept[6]["response"]["output"] == [_lifted_search_call(), SHELL_CALL]
def test_stream_data_that_is_not_json_passes_through():
assert ToolSearchEventLifter().lift("[DONE]") == "[DONE]"
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=call_type,
start_time=datetime.now(),
litellm_call_id=f"tool-search-{call_type}",
function_id=f"tool-search-{call_type}",
)
logging_obj.model_call_details["litellm_params"] = {"aresponses": call_type == "aresponses"}
return logging_obj
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

@ -0,0 +1,220 @@
import json
from typing import Final
import pytest
from litellm.responses.tool_search.lowering import (
LoweredToolSearchRequest,
ToolSearchFunctionNameTaken,
lower_tool_search_request,
needs_tool_search_lowering,
)
CLIENT_TOOL_SEARCH: Final = {
"type": "tool_search",
"execution": "client",
"description": "Search the deferred tools",
"parameters": {
"type": "object",
"properties": {"query": {"type": "string"}, "limit": {"type": "number"}},
"required": ["query"],
},
}
SHELL_TOOL: Final = {"type": "function", "name": "exec_command", "parameters": {"type": "object", "properties": {}}}
SEARCH_CALL: Final = {
"type": "tool_search_call",
"call_id": "call_search",
"execution": "client",
"status": "completed",
"arguments": {"query": "calendar create", "limit": 1},
}
CALENDAR_NAMESPACE: Final = {
"type": "namespace",
"name": "calendar",
"description": "Calendar tools",
"tools": [
{
"type": "function",
"name": "create_event",
"description": "Create an event",
"defer_loading": True,
"parameters": {"properties": {"title": {"type": "string"}}},
}
],
}
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, LIST_INSTRUMENTS],
}
def _lowered(input: object, tools: list[object], tool_choice: object = None) -> LoweredToolSearchRequest:
lowering: Final = lower_tool_search_request(input=input, tools=tools, tool_choice=tool_choice)
assert isinstance(lowering, LoweredToolSearchRequest)
return lowering
def _tool_named(tools: tuple[object, ...], name: str) -> dict[str, object]:
matches: Final = [tool for tool in tools if isinstance(tool, dict) and tool.get("name") == name]
assert len(matches) == 1
return matches[0]
@pytest.mark.parametrize(
("input", "tools", "expected"),
[
("hi", [SHELL_TOOL], False),
("hi", [SHELL_TOOL, CLIENT_TOOL_SEARCH], True),
("hi", [{"type": "tool_search"}], False),
([{"role": "user", "content": "hi"}, SEARCH_CALL], [SHELL_TOOL], True),
([{"role": "user", "content": "hi"}, SEARCH_OUTPUT], None, True),
],
)
def test_only_client_tool_search_state_needs_lowering(input: object, tools: list[object] | None, expected: bool):
assert needs_tool_search_lowering(input, tools) is expected
def test_client_tool_search_becomes_a_function_with_its_declared_schema():
lowered: Final = _lowered("find a calendar tool", [SHELL_TOOL, CLIENT_TOOL_SEARCH])
assert lowered.tools[0] is SHELL_TOOL
assert _tool_named(lowered.tools, "tool_search") == {
"type": "function",
"name": "tool_search",
"description": "Search the deferred tools",
"parameters": CLIENT_TOOL_SEARCH["parameters"],
"strict": False,
}
def test_client_tool_search_without_schema_gets_a_query_parameter():
lowered: Final = _lowered("hi", [{"type": "tool_search", "execution": "client"}])
search_function: Final = _tool_named(lowered.tools, "tool_search")
assert search_function["parameters"] == {
"type": "object",
"properties": {"query": {"type": "string"}},
"required": ["query"],
}
def test_hosted_tool_search_is_left_to_the_provider():
hosted: Final = {"type": "tool_search"}
lowered: Final = _lowered([SEARCH_CALL], [hosted])
assert lowered.tools == (hosted,)
def test_replayed_search_call_becomes_a_function_call_with_the_same_call_id():
lowered: Final = _lowered([{"role": "user", "content": "hi"}, SEARCH_CALL], [CLIENT_TOOL_SEARCH])
assert isinstance(lowered.input, list)
function_call: Final = lowered.input[1]
assert isinstance(function_call, dict)
assert function_call["type"] == "function_call"
assert function_call["call_id"] == "call_search"
assert function_call["name"] == "tool_search"
assert json.loads(function_call["arguments"]) == SEARCH_CALL["arguments"]
def test_replayed_search_output_loads_its_tools_for_the_model():
lowered: Final = _lowered([SEARCH_CALL, SEARCH_OUTPUT], [SHELL_TOOL, CLIENT_TOOL_SEARCH])
assert isinstance(lowered.input, list)
function_output: Final = lowered.input[1]
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"}, loaded_function]
}
namespace: Final = _tool_named(lowered.tools, "calendar")
assert namespace["tools"] == [
{
"type": "function",
"name": "create_event",
"description": "Create an event",
"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():
declared_namespace: Final = {
"type": "namespace",
"name": "calendar",
"description": "Calendar tools",
"tools": [{"type": "function", "name": "list_events", "parameters": {"type": "object", "properties": {}}}],
}
lowered: Final = _lowered([SEARCH_CALL, SEARCH_OUTPUT], [declared_namespace, CLIENT_TOOL_SEARCH])
namespace: Final = _tool_named(lowered.tools, "calendar")
assert isinstance(namespace["tools"], list)
assert [member["name"] for member in namespace["tools"]] == ["list_events", "create_event"]
def test_tool_choice_for_the_search_tool_targets_the_search_function():
lowered: Final = _lowered("hi", [CLIENT_TOOL_SEARCH], {"type": "tool_search"})
assert lowered.tool_choice == {"type": "function", "name": "tool_search"}
def test_a_client_function_already_named_tool_search_is_rejected():
lowering: Final = lower_tool_search_request(
input="hi",
tools=[CLIENT_TOOL_SEARCH, {"type": "function", "name": "tool_search", "parameters": {}}],
tool_choice=None,
)
assert isinstance(lowering, ToolSearchFunctionNameTaken)
def test_a_replay_turn_without_client_search_keeps_its_own_tool_search_function():
own_function: Final = {
"type": "function",
"name": "tool_search",
"parameters": {"type": "object", "properties": {}},
}
lowered: Final = _lowered([SEARCH_CALL, SEARCH_OUTPUT], [own_function])
assert _tool_named(lowered.tools, "tool_search") == own_function
def test_a_replayed_search_output_loads_only_tools_the_client_runs():
other_team_store: Final = {"type": "file_search", "vector_store_ids": ["vs_other_team"]}
smuggled_output: Final = {
**SEARCH_OUTPUT,
"tools": [
other_team_store,
{"type": "mcp", "server_label": "lab", "server_url": "litellm_proxy"},
{
"type": "namespace",
"name": "calendar",
"description": "Calendar tools",
"tools": [{"type": "function", "name": "create_event", "parameters": {}}, other_team_store],
},
],
}
lowered: Final = _lowered([SEARCH_CALL, smuggled_output], [CLIENT_TOOL_SEARCH])
assert [tool["type"] for tool in lowered.tools if isinstance(tool, dict)] == ["function", "namespace"]
assert [member["type"] for member in _tool_named(lowered.tools, "calendar")["tools"]] == ["function"]
assert "vs_other_team" not in json.dumps(lowered.input)