mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(responses): only load client tools from replayed tool searches
A replayed tool_search_output could carry hosted tools, such as a file_search pointing at another team's vector store, and the lowering added them to the request after the proxy had checked its tools. Only function, custom and namespace tools are loaded now, and a namespace keeps only its function and custom members Emulated file_search keeps only its own calls in the response, so a lowered tool_search call made inside it would be lost. Client tool_search is left as declared there, as it was before The tool_search function name check now applies only when the turn declares client tool_search, so a replay turn keeps its own function with that name
This commit is contained in:
parent
bb51372ae8
commit
e9c34d6d36
4 changed files with 88 additions and 7 deletions
|
|
@ -34,7 +34,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.tool_search.lowering import declares_function, needs_tool_search_lowering
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import (
|
||||
PromptObject,
|
||||
|
|
@ -1147,6 +1147,13 @@ def _supports_tool_search_natively(responses_api_provider_config: BaseResponsesA
|
|||
return declared
|
||||
|
||||
|
||||
def _emulates_file_search(tools: Sequence[object] | None) -> bool:
|
||||
# emulated file_search answers with its own calls only, so it would drop a lowered tool_search call
|
||||
from litellm.responses.file_search.emulated_handler import FILE_SEARCH_FUNCTION_NAME
|
||||
|
||||
return declares_function(tools, FILE_SEARCH_FUNCTION_NAME)
|
||||
|
||||
|
||||
def _responses_try_dispatch_lowered_tool_search(
|
||||
*,
|
||||
tools: Iterable[ToolParam] | None,
|
||||
|
|
@ -1173,6 +1180,7 @@ def _responses_try_dispatch_lowered_tool_search(
|
|||
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 _emulates_file_search(declared_tools)
|
||||
):
|
||||
return None
|
||||
from litellm.responses.tool_search.handler import (
|
||||
|
|
|
|||
|
|
@ -13,6 +13,9 @@ _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)
|
||||
|
||||
|
|
@ -84,8 +87,12 @@ 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 any(_parsed(_ClientToolSearchDeclaration, tool) is not None for tool in tools or ()):
|
||||
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)
|
||||
|
||||
|
|
@ -124,11 +131,23 @@ def _loaded_definition(tool: dict[str, object]) -> dict[str, object]:
|
|||
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)
|
||||
return {**_loaded_definition(tool), "tools": [_loaded_definition(member) for member in namespace.tools]}
|
||||
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]:
|
||||
|
|
@ -150,7 +169,7 @@ def _lowered_item(item: object) -> object:
|
|||
output: Final = _parsed(_ReplayedToolSearchOutput, item)
|
||||
if output is None:
|
||||
return item
|
||||
visible_tools: Final = [_visible_loaded_tool(_loaded_tool(tool)) for tool in output.tools]
|
||||
visible_tools: Final = [_visible_loaded_tool(tool) for tool in _loadable_tools(output)]
|
||||
return {
|
||||
"type": "function_call_output",
|
||||
"call_id": output.call_id,
|
||||
|
|
@ -160,8 +179,7 @@ def _lowered_item(item: object) -> object:
|
|||
|
||||
def _loaded_tools(items: Sequence[object]) -> tuple[dict[str, object], ...]:
|
||||
outputs: Final = tuple(_parsed(_ReplayedToolSearchOutput, item) for item in items)
|
||||
replayed: Final = (output.tools for output in outputs if output is not None)
|
||||
return tuple(_loaded_tool(tool) for tool in chain.from_iterable(replayed))
|
||||
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]:
|
||||
|
|
@ -199,13 +217,17 @@ def _lowered_tool_choice(tool_choice: ToolChoice | None) -> ToolChoice | None:
|
|||
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 ("function", TOOL_SEARCH_FUNCTION_NAME) in (_merge_key(index, tool) for index, tool in enumerate(declared)):
|
||||
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)))
|
||||
|
|
|
|||
|
|
@ -184,6 +184,22 @@ async def test_a_client_function_named_tool_search_is_rejected_before_any_upstre
|
|||
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_the_call_is_logged_once_with_the_tool_search_call(monkeypatch: pytest.MonkeyPatch):
|
||||
success_log: Final = _SuccessLog("tool-search-logged-once")
|
||||
|
|
|
|||
|
|
@ -183,3 +183,38 @@ def test_a_client_function_already_named_tool_search_is_rejected():
|
|||
)
|
||||
|
||||
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