mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(responses): detect emulated file search from context, not a tool name
A client function named litellm_file_search turned the tool_search lowering off even when no file search was emulated. Emulated file_search now marks its nested calls in the internal request context, next to is_internal_call, and the lowering only steps aside during that phase
This commit is contained in:
parent
e9c34d6d36
commit
fbd5417d50
5 changed files with 57 additions and 30 deletions
|
|
@ -24,6 +24,23 @@ _billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", d
|
|||
_post_response: Final[ContextVar[bool]] = ContextVar("post_response", default=False)
|
||||
|
||||
|
||||
_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."""
|
||||
|
|
|
|||
|
|
@ -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,7 +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 declares_function, needs_tool_search_lowering
|
||||
from litellm.responses.tool_search.lowering import needs_tool_search_lowering
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import (
|
||||
PromptObject,
|
||||
|
|
@ -1147,13 +1148,6 @@ 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,
|
||||
|
|
@ -1180,7 +1174,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)
|
||||
or in_emulated_file_search()
|
||||
):
|
||||
return None
|
||||
from litellm.responses.tool_search.handler import (
|
||||
|
|
|
|||
|
|
@ -217,7 +217,7 @@ 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:
|
||||
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 ()))
|
||||
|
||||
|
||||
|
|
@ -227,7 +227,7 @@ def lower_tool_search_request(
|
|||
tool_choice: ToolChoice | None,
|
||||
) -> ToolSearchLowering:
|
||||
declared: Final = tuple(tools or ())
|
||||
if _declares_client_tool_search(declared) and declares_function(declared, TOOL_SEARCH_FUNCTION_NAME):
|
||||
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)))
|
||||
|
|
|
|||
|
|
@ -200,6 +200,20 @@ async def test_emulated_file_search_leaves_client_tool_search_as_declared():
|
|||
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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue