mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge c3ce477b3d into 6532dcb73b
This commit is contained in:
commit
7c9677e351
14 changed files with 1371 additions and 47 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
0
litellm/responses/tool_search/__init__.py
Normal file
0
litellm/responses/tool_search/__init__.py
Normal file
90
litellm/responses/tool_search/handler.py
Normal file
90
litellm/responses/tool_search/handler.py
Normal 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)
|
||||
191
litellm/responses/tool_search/lifting.py
Normal file
191
litellm/responses/tool_search/lifting.py
Normal 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
|
||||
238
litellm/responses/tool_search/lowering.py
Normal file
238
litellm/responses/tool_search/lowering.py
Normal 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),
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
0
tests/unit/responses/tool_search/__init__.py
Normal file
0
tests/unit/responses/tool_search/__init__.py
Normal file
241
tests/unit/responses/tool_search/test_handler.py
Normal file
241
tests/unit/responses/tool_search/test_handler.py
Normal 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"]
|
||||
213
tests/unit/responses/tool_search/test_lifting.py
Normal file
213
tests/unit/responses/tool_search/test_lifting.py
Normal 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)
|
||||
220
tests/unit/responses/tool_search/test_lowering.py
Normal file
220
tests/unit/responses/tool_search/test_lowering.py
Normal 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)
|
||||
Loading…
Add table
Reference in a new issue