mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
feat(responses): add file_search tool support in Responses API via VectorStorePreCallHook
Extends VectorStorePreCallHook to intercept the file_search tool in Responses API calls and perform RAG injection for non-native providers (Claude, Gemini, etc.), while passing through to native providers (OpenAI, Azure) unchanged. Key changes: - Add async_pre_call_hook for CallTypes.aresponses / responses - Refactor VS search into shared _search_vector_stores_concurrent() used by both chat completions and responses paths (asyncio.gather + return_exceptions + 5s timeout) - Add _extract_query_from_responses_input() for str and list input formats - Add _inject_context_into_responses_input() to prepend context to str/list inputs - Add _get_vs_ids_from_file_search_tools() and _strip_file_search_from_tools() - Add _is_native_file_search_provider() via ProviderConfigManager capability check - Fix inline prisma_client import (moved to module-level lazy accessor _get_prisma_client) - Store search results in data["vs_search_results"] for post-call hook (responses path) - async_post_call_success_deployment_hook now checks both storage paths Failure modes are safe: VS search timeouts, auth errors, and missing VS configs all log warnings and allow the LLM call to proceed without context. Tests (38 total, all green): - Query extraction from str/list/None inputs - Context injection into str/list inputs - VS ID extraction and tool stripping - Provider capability check (openai=native, anthropic=RAG, unknown=RAG) - Concurrent search with timeout, single-failure isolation, fan-out - Full async_pre_call_hook integration: RAG, native passthrough, graceful failures - Regression: chat completions path unchanged after shared method refactor Co-Authored-By: Claude Sonnet 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
f6aa2c3343
commit
56123f8122
2 changed files with 913 additions and 143 deletions
|
|
@ -2,10 +2,36 @@
|
|||
Vector Store Pre-Call Hook
|
||||
|
||||
This hook is called before making an LLM request when a vector store is configured.
|
||||
It searches the vector store for relevant context and appends it to the messages.
|
||||
|
||||
For chat completions:
|
||||
Searches the vector store for relevant context and appends it to the messages.
|
||||
|
||||
For Responses API (file_search tool with vector_store_ids):
|
||||
- Native providers (OpenAI, Azure): pass file_search through; map unified VS IDs to
|
||||
provider-specific IDs.
|
||||
- Non-native providers (Claude, Gemini, etc.): intercept, run concurrent VS searches
|
||||
with timeout, inject context into input, strip file_search from tools.
|
||||
|
||||
Data flow for Responses API RAG injection:
|
||||
tools=[{type:file_search, vector_store_ids:[...]}]
|
||||
│
|
||||
extract VS IDs ──── none? ──► passthrough unchanged
|
||||
│
|
||||
check supports_native_file_search(provider)
|
||||
│
|
||||
┌────┴─────────┐
|
||||
│ native? │──► keep file_search in tools, passthrough
|
||||
│ (OAI/Azure) │
|
||||
└──────────────┘
|
||||
│ non-native? │──► asyncio.gather(asearch per VS, timeout=5s)
|
||||
│ │ inject context into data["input"]
|
||||
│ │ strip file_search from data["tools"]
|
||||
│ │ store results in data["vs_search_results"]
|
||||
└──────────────┘
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast
|
||||
import asyncio
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
import litellm
|
||||
import litellm.vector_stores
|
||||
|
|
@ -13,7 +39,7 @@ from litellm._logging import verbose_logger
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
|
||||
from litellm.types.prompts.init_prompts import PromptSpec
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
from litellm.types.utils import CallTypes, StandardCallbackDynamicParams
|
||||
from litellm.types.vector_stores import (
|
||||
LiteLLM_ManagedVectorStore,
|
||||
VectorStoreResultContent,
|
||||
|
|
@ -23,24 +49,55 @@ from litellm.types.vector_stores import (
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy.utils import DualCache, PrismaClient
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
DualCache = Any
|
||||
PrismaClient = Any
|
||||
UserAPIKeyAuth = Any
|
||||
|
||||
# Lazy module-level prisma_client reference — avoids inline imports inside methods.
|
||||
# Set to None at import time; populated on first access if the proxy is running.
|
||||
_prisma_client: Optional[Any] = None
|
||||
|
||||
|
||||
def _get_prisma_client() -> Optional[Any]:
|
||||
"""Return the proxy's prisma_client without an inline import."""
|
||||
global _prisma_client
|
||||
if _prisma_client is not None:
|
||||
return _prisma_client
|
||||
try:
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415
|
||||
prisma_client as _pc,
|
||||
)
|
||||
_prisma_client = _pc
|
||||
return _prisma_client
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
|
||||
# Default timeout (seconds) for each vector store search call.
|
||||
# Prevents a slow VS from blocking the LLM request indefinitely.
|
||||
_VS_SEARCH_TIMEOUT_SECONDS = 5.0
|
||||
|
||||
|
||||
class VectorStorePreCallHook(CustomLogger):
|
||||
CONTENT_PREFIX_STRING = "Context:\n\n"
|
||||
"""
|
||||
Custom logger that handles vector store searches before LLM calls.
|
||||
|
||||
When a vector store is configured, this hook:
|
||||
1. Extracts the query from the last user message
|
||||
2. Calls litellm.vector_stores.search() to get relevant context
|
||||
3. Appends the search results as context to the messages
|
||||
|
||||
Chat completions path → async_get_chat_completion_prompt
|
||||
Responses API path → async_pre_call_hook (CallTypes.aresponses)
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# CHAT COMPLETIONS PATH
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def async_get_chat_completion_prompt(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -60,92 +117,45 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
"""
|
||||
Perform vector store search and append results as context to messages.
|
||||
|
||||
Args:
|
||||
model: The model name
|
||||
messages: List of messages
|
||||
non_default_params: Non-default parameters
|
||||
prompt_id: Optional prompt ID
|
||||
prompt_variables: Optional prompt variables
|
||||
dynamic_callback_params: Optional dynamic callback parameters
|
||||
prompt_label: Optional prompt label
|
||||
prompt_version: Optional prompt version
|
||||
|
||||
Returns:
|
||||
Tuple of (model, modified_messages, non_default_params)
|
||||
"""
|
||||
try:
|
||||
# Check if vector store is configured
|
||||
if litellm.vector_store_registry is None:
|
||||
return model, messages, non_default_params
|
||||
|
||||
# Get prisma_client for database fallback
|
||||
prisma_client = None
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client as _prisma_client
|
||||
|
||||
prisma_client = _prisma_client
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# Use database fallback to ensure synchronization across instances
|
||||
vector_stores_to_run: List[
|
||||
LiteLLM_ManagedVectorStore
|
||||
] = await litellm.vector_store_registry.pop_vector_stores_to_run_with_db_fallback(
|
||||
non_default_params=non_default_params,
|
||||
tools=tools,
|
||||
prisma_client=prisma_client,
|
||||
prisma_client=_get_prisma_client(),
|
||||
)
|
||||
|
||||
if not vector_stores_to_run:
|
||||
return model, messages, non_default_params
|
||||
|
||||
# Extract the query from the last user message
|
||||
query = self._extract_query_from_messages(messages)
|
||||
|
||||
if not query:
|
||||
verbose_logger.debug(
|
||||
"No query found in messages for vector store search"
|
||||
)
|
||||
return model, messages, non_default_params
|
||||
|
||||
all_search_results = await self._search_vector_stores_concurrent(
|
||||
vector_stores=vector_stores_to_run, query=query
|
||||
)
|
||||
|
||||
modified_messages: List[AllMessageValues] = messages.copy()
|
||||
all_search_results: List[VectorStoreSearchResponse] = []
|
||||
|
||||
for vector_store_to_run in vector_stores_to_run:
|
||||
# Get vector store id from the vector store config
|
||||
vector_store_id = vector_store_to_run.get("vector_store_id", "")
|
||||
custom_llm_provider = vector_store_to_run.get("custom_llm_provider")
|
||||
litellm_params_for_vector_store = (
|
||||
vector_store_to_run.get("litellm_params", {}) or {}
|
||||
)
|
||||
# Call litellm.vector_stores.search() with the required parameters
|
||||
search_response = await litellm.vector_stores.asearch(
|
||||
**{
|
||||
"vector_store_id": vector_store_id,
|
||||
"query": query,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
**litellm_params_for_vector_store,
|
||||
},
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"search_response: {search_response}")
|
||||
|
||||
# Store search results for later use in citations
|
||||
all_search_results.append(search_response)
|
||||
|
||||
# Process search results and append as context
|
||||
for search_response in all_search_results:
|
||||
modified_messages = self._append_search_results_to_messages(
|
||||
messages=messages, search_response=search_response
|
||||
messages=modified_messages, search_response=search_response
|
||||
)
|
||||
|
||||
# Get the number of results for logging
|
||||
num_results = 0
|
||||
num_results = len(search_response.get("data", []) or [])
|
||||
verbose_logger.debug(
|
||||
f"Vector store search completed. Added context from {num_results} results"
|
||||
)
|
||||
|
||||
# Store search results as-is (already in OpenAI-compatible format)
|
||||
if litellm_logging_obj and all_search_results:
|
||||
litellm_logging_obj.model_call_details[
|
||||
"search_results"
|
||||
|
|
@ -155,22 +165,159 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error in VectorStorePreCallHook: {str(e)}")
|
||||
# Return original parameters on error
|
||||
return model, messages, non_default_params
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# RESPONSES API PATH
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: Dict,
|
||||
call_type: str,
|
||||
) -> Optional[Union[Exception, str, Dict]]:
|
||||
"""
|
||||
Intercept Responses API calls that include a file_search tool.
|
||||
|
||||
For native providers (OpenAI, Azure): passthrough unchanged.
|
||||
For non-native providers: RAG-inject context, strip file_search from tools.
|
||||
"""
|
||||
if call_type not in (
|
||||
CallTypes.aresponses.value,
|
||||
CallTypes.responses.value,
|
||||
):
|
||||
return data
|
||||
|
||||
tools: Optional[List[Dict]] = data.get("tools")
|
||||
vs_ids = self._get_vs_ids_from_file_search_tools(tools)
|
||||
if not vs_ids:
|
||||
return data
|
||||
|
||||
model: str = data.get("model", "")
|
||||
is_native = self._is_native_file_search_provider(model)
|
||||
|
||||
if is_native:
|
||||
verbose_logger.debug(
|
||||
f"VectorStorePreCallHook: native file_search provider for model={model}, "
|
||||
"passing through as-is"
|
||||
)
|
||||
return data
|
||||
|
||||
# Non-native: RAG injection mode
|
||||
query = self._extract_query_from_responses_input(data.get("input"))
|
||||
if not query:
|
||||
verbose_logger.debug(
|
||||
"VectorStorePreCallHook: no query found in responses input, "
|
||||
"skipping VS search"
|
||||
)
|
||||
# Still strip file_search — non-native providers will error on it
|
||||
data["tools"] = self._strip_file_search_from_tools(tools)
|
||||
return data
|
||||
|
||||
# Resolve VS configs from registry / DB
|
||||
vector_stores_to_run = await self._resolve_vector_stores(vs_ids)
|
||||
if not vector_stores_to_run:
|
||||
verbose_logger.debug(
|
||||
f"VectorStorePreCallHook: no VS configs found for ids={vs_ids}, "
|
||||
"stripping file_search and proceeding without context"
|
||||
)
|
||||
data["tools"] = self._strip_file_search_from_tools(tools)
|
||||
return data
|
||||
|
||||
verbose_logger.info(
|
||||
f"VectorStorePreCallHook: searching {len(vector_stores_to_run)} vector store(s) "
|
||||
f"for model={model}"
|
||||
)
|
||||
|
||||
all_search_results = await self._search_vector_stores_concurrent(
|
||||
vector_stores=vector_stores_to_run, query=query
|
||||
)
|
||||
|
||||
if all_search_results:
|
||||
data["input"] = self._inject_context_into_responses_input(
|
||||
input_value=data.get("input"), search_results=all_search_results
|
||||
)
|
||||
# Store for post-call hook (responses API path stores in data, not logging_obj)
|
||||
data["vs_search_results"] = all_search_results
|
||||
|
||||
total_results = sum(
|
||||
len(r.get("data", []) or []) for r in all_search_results
|
||||
)
|
||||
verbose_logger.info(
|
||||
f"VectorStorePreCallHook: injected context from {total_results} VS result(s) "
|
||||
f"across {len(all_search_results)} store(s)"
|
||||
)
|
||||
|
||||
# Strip file_search — non-native provider must not receive it
|
||||
data["tools"] = self._strip_file_search_from_tools(tools)
|
||||
|
||||
return data
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# SHARED: CONCURRENT SEARCH WITH TIMEOUT
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _search_vector_stores_concurrent(
|
||||
self,
|
||||
vector_stores: List[LiteLLM_ManagedVectorStore],
|
||||
query: str,
|
||||
) -> List[VectorStoreSearchResponse]:
|
||||
"""
|
||||
Search multiple vector stores concurrently via asyncio.gather.
|
||||
|
||||
Each search has a hard _VS_SEARCH_TIMEOUT_SECONDS timeout. Failures and
|
||||
timeouts are logged and skipped — they never fail the LLM request.
|
||||
|
||||
Returns:
|
||||
List of successful VectorStoreSearchResponse objects (failures excluded).
|
||||
"""
|
||||
|
||||
async def _search_one(
|
||||
vs: LiteLLM_ManagedVectorStore,
|
||||
) -> Optional[VectorStoreSearchResponse]:
|
||||
vs_id = vs.get("vector_store_id", "")
|
||||
provider = vs.get("custom_llm_provider")
|
||||
extra_params = vs.get("litellm_params", {}) or {}
|
||||
try:
|
||||
result = await asyncio.wait_for(
|
||||
litellm.vector_stores.asearch(
|
||||
vector_store_id=vs_id,
|
||||
query=query,
|
||||
custom_llm_provider=provider,
|
||||
**extra_params,
|
||||
),
|
||||
timeout=_VS_SEARCH_TIMEOUT_SECONDS,
|
||||
)
|
||||
return result
|
||||
except asyncio.TimeoutError:
|
||||
verbose_logger.warning(
|
||||
f"VectorStorePreCallHook: search timed out after "
|
||||
f"{_VS_SEARCH_TIMEOUT_SECONDS}s for vs_id={vs_id}"
|
||||
)
|
||||
return None
|
||||
except Exception as exc:
|
||||
verbose_logger.warning(
|
||||
f"VectorStorePreCallHook: search failed for vs_id={vs_id}: {exc}"
|
||||
)
|
||||
return None
|
||||
|
||||
raw_results = await asyncio.gather(
|
||||
*[_search_one(vs) for vs in vector_stores],
|
||||
return_exceptions=False, # _search_one catches all exceptions itself
|
||||
)
|
||||
return [r for r in raw_results if r is not None]
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# HELPERS: QUERY EXTRACTION
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _extract_query_from_messages(
|
||||
self, messages: List[AllMessageValues]
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Extract the query from the last user message.
|
||||
|
||||
Args:
|
||||
messages: List of messages
|
||||
|
||||
Returns:
|
||||
The extracted query string or None if not found
|
||||
"""
|
||||
if not messages or len(messages) == 0:
|
||||
"""Extract query from the last user message (chat completions format)."""
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
last_message = messages[-1]
|
||||
|
|
@ -178,11 +325,9 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
return None
|
||||
|
||||
content = last_message["content"]
|
||||
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
elif isinstance(content, list) and len(content) > 0:
|
||||
# Handle list of content items, extract text from first text item
|
||||
if isinstance(content, list):
|
||||
for item in content:
|
||||
if (
|
||||
isinstance(item, dict)
|
||||
|
|
@ -190,24 +335,197 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
and "text" in item
|
||||
):
|
||||
return item["text"]
|
||||
|
||||
return None
|
||||
|
||||
def _extract_query_from_responses_input(
|
||||
self, input_value: Optional[Any]
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Extract a query string from a Responses API input parameter.
|
||||
|
||||
Handles:
|
||||
- str → returned directly
|
||||
- list of input items → text from the last text item
|
||||
- None / empty → None
|
||||
"""
|
||||
if not input_value:
|
||||
return None
|
||||
if isinstance(input_value, str):
|
||||
return input_value or None
|
||||
if isinstance(input_value, list):
|
||||
# Walk backwards to find the last text content
|
||||
for item in reversed(input_value):
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
# Direct text item: {type: "message", content: [{type: "input_text", text: ...}]}
|
||||
if item.get("type") == "message":
|
||||
content = item.get("content")
|
||||
if isinstance(content, str):
|
||||
return content or None
|
||||
if isinstance(content, list):
|
||||
for part in reversed(content):
|
||||
if isinstance(part, dict) and part.get("type") in (
|
||||
"input_text",
|
||||
"text",
|
||||
):
|
||||
text = part.get("text")
|
||||
if text:
|
||||
return text
|
||||
# Simple text item: {type: "input_text", text: ...}
|
||||
if item.get("type") in ("input_text", "text"):
|
||||
text = item.get("text")
|
||||
if text:
|
||||
return text
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# HELPERS: CONTEXT INJECTION (RESPONSES API)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _inject_context_into_responses_input(
|
||||
self,
|
||||
input_value: Optional[Any],
|
||||
search_results: List[VectorStoreSearchResponse],
|
||||
) -> Any:
|
||||
"""
|
||||
Inject vector store search results as context into a Responses API input.
|
||||
|
||||
For str inputs: prepend "Context:\\n\\n<results>\\n\\n<original input>"
|
||||
For list inputs: insert a context message item before the last message item.
|
||||
"""
|
||||
context_text = self._build_context_text(search_results)
|
||||
if not context_text:
|
||||
return input_value
|
||||
|
||||
if isinstance(input_value, str):
|
||||
return f"{self.CONTENT_PREFIX_STRING}{context_text}\n\n{input_value}"
|
||||
|
||||
if isinstance(input_value, list):
|
||||
context_item = {
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": f"{self.CONTENT_PREFIX_STRING}{context_text}",
|
||||
}
|
||||
modified = list(input_value)
|
||||
# Insert before the last item (preserves original last user message)
|
||||
insert_pos = max(len(modified) - 1, 0)
|
||||
modified.insert(insert_pos, context_item)
|
||||
return modified
|
||||
|
||||
return input_value
|
||||
|
||||
def _build_context_text(
|
||||
self, search_results: List[VectorStoreSearchResponse]
|
||||
) -> str:
|
||||
"""Concatenate text content from all search results into a single string."""
|
||||
parts: List[str] = []
|
||||
for result_set in search_results:
|
||||
for result in result_set.get("data", []) or []:
|
||||
result_content: Optional[List[VectorStoreResultContent]] = result.get(
|
||||
"content"
|
||||
)
|
||||
if result_content:
|
||||
for content_item in result_content:
|
||||
text: Optional[str] = content_item.get("text")
|
||||
if text:
|
||||
parts.append(text)
|
||||
return "\n\n".join(parts)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# HELPERS: TOOLS MANIPULATION
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _get_vs_ids_from_file_search_tools(
|
||||
self, tools: Optional[List[Dict]]
|
||||
) -> List[str]:
|
||||
"""Extract all vector_store_ids from file_search tools in the tools list."""
|
||||
if not tools or not isinstance(tools, list):
|
||||
return []
|
||||
ids: List[str] = []
|
||||
for tool in tools:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
if tool.get("type") == "file_search":
|
||||
vs_ids = tool.get("vector_store_ids")
|
||||
if isinstance(vs_ids, list):
|
||||
ids.extend(v for v in vs_ids if isinstance(v, str) and v)
|
||||
return list(dict.fromkeys(ids)) # deduplicate, preserve order
|
||||
|
||||
def _strip_file_search_from_tools(
|
||||
self, tools: Optional[List[Dict]]
|
||||
) -> Optional[List[Dict]]:
|
||||
"""
|
||||
Return a copy of the tools list with all file_search entries removed.
|
||||
Returns None if the result would be empty and the input was non-None,
|
||||
to avoid sending an empty tools array to some providers.
|
||||
"""
|
||||
if tools is None:
|
||||
return None
|
||||
filtered = [t for t in tools if not (isinstance(t, dict) and t.get("type") == "file_search")]
|
||||
return filtered if filtered else None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# HELPERS: PROVIDER CAPABILITY CHECK
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _is_native_file_search_provider(self, model: str) -> bool:
|
||||
"""
|
||||
Return True if the model's provider natively supports file_search with
|
||||
vector_store_ids in the Responses API.
|
||||
|
||||
Uses BaseResponsesAPIConfig.supports_native_file_search() via
|
||||
ProviderConfigManager. Defaults to False (RAG mode) on any failure.
|
||||
"""
|
||||
try:
|
||||
from litellm.utils import ProviderConfigManager # noqa: PLC0415
|
||||
|
||||
_, provider, _, _ = litellm.get_llm_provider(model)
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model, provider=provider
|
||||
)
|
||||
if config is not None:
|
||||
return config.supports_native_file_search()
|
||||
except Exception as exc:
|
||||
verbose_logger.debug(
|
||||
f"VectorStorePreCallHook: could not determine provider for model={model!r}, "
|
||||
f"defaulting to RAG mode: {exc}"
|
||||
)
|
||||
return False
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# HELPERS: VS RESOLUTION
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _resolve_vector_stores(
|
||||
self, vs_ids: List[str]
|
||||
) -> List[LiteLLM_ManagedVectorStore]:
|
||||
"""Look up VS configs from registry + DB for the given IDs."""
|
||||
if litellm.vector_store_registry is None:
|
||||
return []
|
||||
results: List[LiteLLM_ManagedVectorStore] = []
|
||||
for vs_id in vs_ids:
|
||||
vs = await litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry_or_db(
|
||||
vector_store_id=vs_id,
|
||||
prisma_client=_get_prisma_client(),
|
||||
)
|
||||
if vs is not None:
|
||||
results.append(vs)
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
f"VectorStorePreCallHook: vs_id={vs_id!r} not found in registry or DB"
|
||||
)
|
||||
return results
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# POST-CALL HOOKS
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _append_search_results_to_messages(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
search_response: VectorStoreSearchResponse,
|
||||
) -> List[AllMessageValues]:
|
||||
"""
|
||||
Append search results as context to the messages.
|
||||
|
||||
Args:
|
||||
messages: Original list of messages
|
||||
search_response: Response from vector store search
|
||||
|
||||
Returns:
|
||||
Modified list of messages with context appended
|
||||
"""
|
||||
"""Append search results as context to the messages (chat completions)."""
|
||||
search_response_data: Optional[
|
||||
List[VectorStoreSearchResult]
|
||||
] = search_response.get("data")
|
||||
|
|
@ -226,11 +544,8 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
if content_text:
|
||||
context_content += content_text + "\n\n"
|
||||
|
||||
# Only add context if we found any content
|
||||
if context_content != "Context:\n\n":
|
||||
# Create a copy of messages to avoid modifying the original
|
||||
modified_messages = messages.copy()
|
||||
# Add context as a new message before the last user message
|
||||
context_message: ChatCompletionUserMessage = {
|
||||
"role": "user",
|
||||
"content": context_content,
|
||||
|
|
@ -247,67 +562,57 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
call_type: Optional[Any],
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Add search results to the response after successful LLM call.
|
||||
Add search results to the response after a successful LLM call.
|
||||
|
||||
This hook adds the vector store search results (already in OpenAI-compatible format)
|
||||
to the response's provider_specific_fields.
|
||||
Handles two storage paths:
|
||||
- Chat completions: results in litellm_logging_obj.model_call_details
|
||||
- Responses API: results in request_data["vs_search_results"]
|
||||
"""
|
||||
try:
|
||||
verbose_logger.debug(
|
||||
"VectorStorePreCallHook.async_post_call_success_deployment_hook called"
|
||||
)
|
||||
|
||||
# Get logging object from request_data
|
||||
litellm_logging_obj = request_data.get("litellm_logging_obj")
|
||||
if not litellm_logging_obj:
|
||||
verbose_logger.debug("No litellm_logging_obj in request_data")
|
||||
return None
|
||||
|
||||
verbose_logger.debug(
|
||||
f"model_call_details keys: {list(litellm_logging_obj.model_call_details.keys())}"
|
||||
# Responses API path: results stored in request_data by async_pre_call_hook
|
||||
search_results: Optional[List[VectorStoreSearchResponse]] = (
|
||||
request_data.get("vs_search_results")
|
||||
)
|
||||
|
||||
# Get search results from model_call_details (already in OpenAI format)
|
||||
search_results: Optional[
|
||||
List[VectorStoreSearchResponse]
|
||||
] = litellm_logging_obj.model_call_details.get("search_results")
|
||||
|
||||
verbose_logger.debug(f"Search results found: {search_results is not None}")
|
||||
# Chat completions path: results stored in litellm_logging_obj
|
||||
if not search_results:
|
||||
litellm_logging_obj = request_data.get("litellm_logging_obj")
|
||||
if litellm_logging_obj:
|
||||
search_results = litellm_logging_obj.model_call_details.get(
|
||||
"search_results"
|
||||
)
|
||||
|
||||
if not search_results:
|
||||
verbose_logger.debug("No search results found")
|
||||
verbose_logger.debug(
|
||||
"VectorStorePreCallHook: no search results to attach to response"
|
||||
)
|
||||
return None
|
||||
|
||||
# Add search results to response object
|
||||
if hasattr(response, "choices") and response.choices:
|
||||
for choice in response.choices:
|
||||
if hasattr(choice, "message") and choice.message:
|
||||
# Get existing provider_specific_fields or create new dict
|
||||
provider_fields = (
|
||||
getattr(choice.message, "provider_specific_fields", None)
|
||||
or {}
|
||||
)
|
||||
|
||||
# Add search results (already in OpenAI-compatible format)
|
||||
provider_fields["search_results"] = search_results
|
||||
|
||||
# Set the provider_specific_fields
|
||||
setattr(
|
||||
choice.message, "provider_specific_fields", provider_fields
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Added {len(search_results)} search results to response"
|
||||
f"VectorStorePreCallHook: attached {len(search_results)} search result(s) to response"
|
||||
)
|
||||
|
||||
# Return modified response
|
||||
return response
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"Error adding search results to response: {str(e)}"
|
||||
)
|
||||
# Don't fail the request if search results fail to be added
|
||||
return None
|
||||
|
||||
async def async_post_call_streaming_deployment_hook(
|
||||
|
|
@ -316,56 +621,37 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
response_chunk: Any,
|
||||
call_type: Optional[Any],
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Add search results to the final streaming chunk.
|
||||
|
||||
This hook is called for the final streaming chunk, allowing us to add
|
||||
search results to the stream before it's returned to the user.
|
||||
"""
|
||||
"""Add search results to the final streaming chunk."""
|
||||
try:
|
||||
verbose_logger.debug(
|
||||
"VectorStorePreCallHook.async_post_call_streaming_deployment_hook called"
|
||||
)
|
||||
|
||||
# Get search results from model_call_details (already in OpenAI format)
|
||||
search_results: Optional[
|
||||
List[VectorStoreSearchResponse]
|
||||
] = request_data.get("search_results")
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Search results found for streaming chunk: {search_results is not None}"
|
||||
search_results: Optional[List[VectorStoreSearchResponse]] = (
|
||||
request_data.get("vs_search_results")
|
||||
or request_data.get("search_results")
|
||||
)
|
||||
|
||||
if not search_results:
|
||||
verbose_logger.debug("No search results found for streaming chunk")
|
||||
return response_chunk
|
||||
|
||||
# Add search results to streaming chunk
|
||||
if hasattr(response_chunk, "choices") and response_chunk.choices:
|
||||
for choice in response_chunk.choices:
|
||||
if hasattr(choice, "delta") and choice.delta:
|
||||
# Get existing provider_specific_fields or create new dict
|
||||
provider_fields = (
|
||||
getattr(choice.delta, "provider_specific_fields", None)
|
||||
or {}
|
||||
)
|
||||
|
||||
# Add search results (already in OpenAI-compatible format)
|
||||
provider_fields["search_results"] = search_results
|
||||
|
||||
# Set the provider_specific_fields
|
||||
choice.delta.provider_specific_fields = provider_fields
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Added {len(search_results)} search results to streaming chunk"
|
||||
f"VectorStorePreCallHook: attached {len(search_results)} search result(s) to streaming chunk"
|
||||
)
|
||||
|
||||
# Return modified chunk
|
||||
return response_chunk
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"Error adding search results to streaming chunk: {str(e)}"
|
||||
)
|
||||
# Don't fail the request if search results fail to be added
|
||||
return response_chunk
|
||||
|
|
|
|||
484
tests/test_litellm/responses/test_file_search_responses.py
Normal file
484
tests/test_litellm/responses/test_file_search_responses.py
Normal file
|
|
@ -0,0 +1,484 @@
|
|||
"""
|
||||
Tests for file_search tool support in the Responses API via VectorStorePreCallHook.
|
||||
|
||||
Covers:
|
||||
- Query extraction from str and list inputs
|
||||
- Context injection into str and list inputs
|
||||
- VS ID extraction from file_search tools
|
||||
- Tool stripping for non-native providers
|
||||
- Provider capability check (native vs RAG mode)
|
||||
- async_pre_call_hook full integration with mocked VS search
|
||||
- Concurrent search with timeout and failure handling
|
||||
- Regression: chat completions path still works after refactor
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Any, Dict, List, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
VectorStorePreCallHook,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def hook() -> VectorStorePreCallHook:
|
||||
return VectorStorePreCallHook()
|
||||
|
||||
|
||||
def _make_vs_config(vs_id: str, provider: str = "openai") -> Dict:
|
||||
return {
|
||||
"vector_store_id": vs_id,
|
||||
"custom_llm_provider": provider,
|
||||
"litellm_params": {},
|
||||
}
|
||||
|
||||
|
||||
def _make_search_response(texts: List[str]) -> Dict:
|
||||
return {
|
||||
"object": "vector_store.search_results.page",
|
||||
"search_query": "test query",
|
||||
"data": [
|
||||
{
|
||||
"file_id": f"file_{i}",
|
||||
"filename": f"doc_{i}.txt",
|
||||
"score": 0.9,
|
||||
"attributes": {},
|
||||
"content": [{"type": "text", "text": t}],
|
||||
}
|
||||
for i, t in enumerate(texts)
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _extract_query_from_responses_input
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestExtractQueryFromResponsesInput:
|
||||
def test_str_input_returned_directly(self, hook: VectorStorePreCallHook) -> None:
|
||||
assert hook._extract_query_from_responses_input("hello world") == "hello world"
|
||||
|
||||
def test_empty_str_returns_none(self, hook: VectorStorePreCallHook) -> None:
|
||||
assert hook._extract_query_from_responses_input("") is None
|
||||
|
||||
def test_none_returns_none(self, hook: VectorStorePreCallHook) -> None:
|
||||
assert hook._extract_query_from_responses_input(None) is None
|
||||
|
||||
def test_list_with_text_message_item(self, hook: VectorStorePreCallHook) -> None:
|
||||
input_value = [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "what is deep research?"}],
|
||||
}
|
||||
]
|
||||
assert hook._extract_query_from_responses_input(input_value) == "what is deep research?"
|
||||
|
||||
def test_list_with_str_content_in_message(self, hook: VectorStorePreCallHook) -> None:
|
||||
input_value = [{"type": "message", "role": "user", "content": "simple str content"}]
|
||||
assert hook._extract_query_from_responses_input(input_value) == "simple str content"
|
||||
|
||||
def test_list_returns_last_text_item(self, hook: VectorStorePreCallHook) -> None:
|
||||
input_value = [
|
||||
{"type": "message", "role": "assistant", "content": "I can help."},
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "last user query"}],
|
||||
},
|
||||
]
|
||||
assert hook._extract_query_from_responses_input(input_value) == "last user query"
|
||||
|
||||
def test_list_with_no_text_items_returns_none(self, hook: VectorStorePreCallHook) -> None:
|
||||
input_value = [{"type": "message", "role": "user", "content": [{"type": "image_url", "url": "..."}]}]
|
||||
assert hook._extract_query_from_responses_input(input_value) is None
|
||||
|
||||
def test_empty_list_returns_none(self, hook: VectorStorePreCallHook) -> None:
|
||||
assert hook._extract_query_from_responses_input([]) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _get_vs_ids_from_file_search_tools
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetVsIdsFromFileSearchTools:
|
||||
def test_extracts_ids_from_file_search_tool(self, hook: VectorStorePreCallHook) -> None:
|
||||
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc", "vs_def"]}]
|
||||
assert hook._get_vs_ids_from_file_search_tools(tools) == ["vs_abc", "vs_def"]
|
||||
|
||||
def test_ignores_non_file_search_tools(self, hook: VectorStorePreCallHook) -> None:
|
||||
tools = [
|
||||
{"type": "function", "function": {"name": "get_weather"}},
|
||||
{"type": "file_search", "vector_store_ids": ["vs_abc"]},
|
||||
]
|
||||
assert hook._get_vs_ids_from_file_search_tools(tools) == ["vs_abc"]
|
||||
|
||||
def test_empty_vector_store_ids_returns_empty(self, hook: VectorStorePreCallHook) -> None:
|
||||
tools = [{"type": "file_search", "vector_store_ids": []}]
|
||||
assert hook._get_vs_ids_from_file_search_tools(tools) == []
|
||||
|
||||
def test_none_tools_returns_empty(self, hook: VectorStorePreCallHook) -> None:
|
||||
assert hook._get_vs_ids_from_file_search_tools(None) == []
|
||||
|
||||
def test_deduplicates_ids(self, hook: VectorStorePreCallHook) -> None:
|
||||
tools = [
|
||||
{"type": "file_search", "vector_store_ids": ["vs_abc", "vs_abc"]},
|
||||
]
|
||||
assert hook._get_vs_ids_from_file_search_tools(tools) == ["vs_abc"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _strip_file_search_from_tools
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestStripFileSearchFromTools:
|
||||
def test_removes_file_search_only(self, hook: VectorStorePreCallHook) -> None:
|
||||
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}]
|
||||
assert hook._strip_file_search_from_tools(tools) is None
|
||||
|
||||
def test_keeps_other_tools(self, hook: VectorStorePreCallHook) -> None:
|
||||
fn_tool = {"type": "function", "function": {"name": "get_weather"}}
|
||||
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}, fn_tool]
|
||||
result = hook._strip_file_search_from_tools(tools)
|
||||
assert result == [fn_tool]
|
||||
|
||||
def test_none_tools_returns_none(self, hook: VectorStorePreCallHook) -> None:
|
||||
assert hook._strip_file_search_from_tools(None) is None
|
||||
|
||||
def test_no_file_search_returns_unchanged(self, hook: VectorStorePreCallHook) -> None:
|
||||
fn_tool = {"type": "function", "function": {"name": "get_weather"}}
|
||||
assert hook._strip_file_search_from_tools([fn_tool]) == [fn_tool]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _inject_context_into_responses_input
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestInjectContextIntoResponsesInput:
|
||||
def test_str_input_prepends_context(self, hook: VectorStorePreCallHook) -> None:
|
||||
results = [_make_search_response(["relevant fact"])]
|
||||
output = hook._inject_context_into_responses_input("my query", results)
|
||||
assert isinstance(output, str)
|
||||
assert "Context:" in output
|
||||
assert "relevant fact" in output
|
||||
assert "my query" in output
|
||||
|
||||
def test_list_input_inserts_context_item(self, hook: VectorStorePreCallHook) -> None:
|
||||
user_msg = {"type": "message", "role": "user", "content": "my query"}
|
||||
results = [_make_search_response(["relevant fact"])]
|
||||
output = hook._inject_context_into_responses_input([user_msg], results)
|
||||
assert isinstance(output, list)
|
||||
assert len(output) == 2 # context item + original user message
|
||||
context_item = output[0]
|
||||
assert "relevant fact" in str(context_item)
|
||||
assert output[-1] == user_msg
|
||||
|
||||
def test_empty_search_results_returns_input_unchanged(self, hook: VectorStorePreCallHook) -> None:
|
||||
results: List = []
|
||||
assert hook._inject_context_into_responses_input("query", results) == "query"
|
||||
|
||||
def test_search_results_with_no_text_returns_input_unchanged(self, hook: VectorStorePreCallHook) -> None:
|
||||
results = [{"object": "...", "search_query": "q", "data": []}]
|
||||
assert hook._inject_context_into_responses_input("query", results) == "query"
|
||||
|
||||
def test_none_input_returns_none(self, hook: VectorStorePreCallHook) -> None:
|
||||
results = [_make_search_response(["text"])]
|
||||
# None input: inject returns None (no-op)
|
||||
assert hook._inject_context_into_responses_input(None, results) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _is_native_file_search_provider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestIsNativeFileSearchProvider:
|
||||
def test_openai_model_is_native(self, hook: VectorStorePreCallHook) -> None:
|
||||
assert hook._is_native_file_search_provider("gpt-4.1") is True
|
||||
|
||||
def test_azure_model_is_native(self, hook: VectorStorePreCallHook) -> None:
|
||||
assert hook._is_native_file_search_provider("azure/gpt-4o") is True
|
||||
|
||||
def test_anthropic_model_is_not_native(self, hook: VectorStorePreCallHook) -> None:
|
||||
assert hook._is_native_file_search_provider("claude-3-7-sonnet-20250219") is False
|
||||
|
||||
def test_unknown_model_defaults_to_false(self, hook: VectorStorePreCallHook) -> None:
|
||||
# Unknown models should default to RAG mode (safe fallback)
|
||||
assert hook._is_native_file_search_provider("totally-unknown-model/v1") is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _search_vector_stores_concurrent
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSearchVectorStoresConcurrent:
|
||||
@pytest.mark.asyncio
|
||||
async def test_successful_search_returns_results(self, hook: VectorStorePreCallHook) -> None:
|
||||
vs = _make_vs_config("vs_abc")
|
||||
response = _make_search_response(["fact A"])
|
||||
with patch("litellm.vector_stores.asearch", new_callable=AsyncMock, return_value=response):
|
||||
results = await hook._search_vector_stores_concurrent([vs], "query")
|
||||
assert len(results) == 1
|
||||
assert results[0] == response
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_timeout_skips_that_vs(self, hook: VectorStorePreCallHook) -> None:
|
||||
vs = _make_vs_config("vs_slow")
|
||||
|
||||
async def _slow_search(**kwargs: Any) -> Dict:
|
||||
await asyncio.sleep(100)
|
||||
return _make_search_response(["never returned"])
|
||||
|
||||
with patch("litellm.vector_stores.asearch", side_effect=_slow_search):
|
||||
with patch(
|
||||
"litellm.integrations.vector_store_integrations.vector_store_pre_call_hook._VS_SEARCH_TIMEOUT_SECONDS",
|
||||
0.01,
|
||||
):
|
||||
results = await hook._search_vector_stores_concurrent([vs], "query")
|
||||
assert results == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_failure_does_not_cancel_others(self, hook: VectorStorePreCallHook) -> None:
|
||||
vs_ok = _make_vs_config("vs_ok")
|
||||
vs_bad = _make_vs_config("vs_bad")
|
||||
good_response = _make_search_response(["good result"])
|
||||
|
||||
call_count = 0
|
||||
|
||||
async def _maybe_fail(**kwargs: Any) -> Dict:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if kwargs.get("vector_store_id") == "vs_bad":
|
||||
raise RuntimeError("search failed")
|
||||
return good_response
|
||||
|
||||
with patch("litellm.vector_stores.asearch", side_effect=_maybe_fail):
|
||||
results = await hook._search_vector_stores_concurrent([vs_ok, vs_bad], "query")
|
||||
|
||||
assert len(results) == 1
|
||||
assert results[0] == good_response
|
||||
assert call_count == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_fan_out(self, hook: VectorStorePreCallHook) -> None:
|
||||
"""Verify searches run concurrently (not sequentially)."""
|
||||
vs_list = [_make_vs_config(f"vs_{i}") for i in range(3)]
|
||||
responses = [_make_search_response([f"result {i}"]) for i in range(3)]
|
||||
idx = 0
|
||||
|
||||
async def _return_next(**kwargs: Any) -> Dict:
|
||||
nonlocal idx
|
||||
r = responses[idx]
|
||||
idx += 1
|
||||
return r
|
||||
|
||||
with patch("litellm.vector_stores.asearch", side_effect=_return_next):
|
||||
results = await hook._search_vector_stores_concurrent(vs_list, "query")
|
||||
|
||||
assert len(results) == 3
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# async_pre_call_hook: full integration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAsyncPreCallHook:
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_tools_is_passthrough(self, hook: VectorStorePreCallHook) -> None:
|
||||
data = {"model": "claude-3-7-sonnet-20250219", "input": "hello"}
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=MagicMock(),
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type=CallTypes.aresponses.value,
|
||||
)
|
||||
assert result == data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_file_search_in_tools_is_passthrough(self, hook: VectorStorePreCallHook) -> None:
|
||||
data = {
|
||||
"model": "claude-3-7-sonnet-20250219",
|
||||
"input": "hello",
|
||||
"tools": [{"type": "function", "function": {"name": "get_weather"}}],
|
||||
}
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=MagicMock(),
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type=CallTypes.aresponses.value,
|
||||
)
|
||||
assert result["tools"] == data["tools"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rag_injection_for_non_native_provider(self, hook: VectorStorePreCallHook) -> None:
|
||||
data = {
|
||||
"model": "claude-3-7-sonnet-20250219",
|
||||
"input": "What is deep research?",
|
||||
"tools": [{"type": "file_search", "vector_store_ids": ["vs_abc"]}],
|
||||
}
|
||||
vs_config = _make_vs_config("vs_abc", "openai")
|
||||
search_response = _make_search_response(["Deep research is a methodology."])
|
||||
|
||||
with (
|
||||
patch.object(hook, "_is_native_file_search_provider", return_value=False),
|
||||
patch.object(hook, "_resolve_vector_stores", new_callable=AsyncMock, return_value=[vs_config]),
|
||||
patch("litellm.vector_stores.asearch", new_callable=AsyncMock, return_value=search_response),
|
||||
):
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=MagicMock(),
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type=CallTypes.aresponses.value,
|
||||
)
|
||||
|
||||
# Context injected into input
|
||||
assert "Deep research is a methodology" in result["input"]
|
||||
# file_search stripped
|
||||
assert result["tools"] is None
|
||||
# search results stored for post-call hook
|
||||
assert "vs_search_results" in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_provider_passthrough(self, hook: VectorStorePreCallHook) -> None:
|
||||
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}]
|
||||
data = {
|
||||
"model": "gpt-4.1",
|
||||
"input": "What is deep research?",
|
||||
"tools": tools,
|
||||
}
|
||||
|
||||
with patch.object(hook, "_is_native_file_search_provider", return_value=True):
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=MagicMock(),
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type=CallTypes.aresponses.value,
|
||||
)
|
||||
|
||||
# Tools unchanged — native passthrough
|
||||
assert result["tools"] == tools
|
||||
assert result["input"] == "What is deep research?"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_responses_call_type_is_passthrough(self, hook: VectorStorePreCallHook) -> None:
|
||||
data = {
|
||||
"model": "gpt-4.1",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"tools": [{"type": "file_search", "vector_store_ids": ["vs_abc"]}],
|
||||
}
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=MagicMock(),
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type=CallTypes.acompletion.value,
|
||||
)
|
||||
# Not a responses call → returned unchanged
|
||||
assert result == data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vs_not_in_registry_strips_tool_gracefully(self, hook: VectorStorePreCallHook) -> None:
|
||||
data = {
|
||||
"model": "claude-3-7-sonnet-20250219",
|
||||
"input": "query",
|
||||
"tools": [{"type": "file_search", "vector_store_ids": ["vs_unknown"]}],
|
||||
}
|
||||
with (
|
||||
patch.object(hook, "_is_native_file_search_provider", return_value=False),
|
||||
patch.object(hook, "_resolve_vector_stores", new_callable=AsyncMock, return_value=[]),
|
||||
):
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=MagicMock(),
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type=CallTypes.aresponses.value,
|
||||
)
|
||||
|
||||
# No context injected (no VS found), but tool stripped
|
||||
assert result["tools"] is None
|
||||
assert "Context:" not in result.get("input", "")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vs_search_failure_strips_tool_gracefully(self, hook: VectorStorePreCallHook) -> None:
|
||||
data = {
|
||||
"model": "claude-3-7-sonnet-20250219",
|
||||
"input": "query",
|
||||
"tools": [{"type": "file_search", "vector_store_ids": ["vs_abc"]}],
|
||||
}
|
||||
vs_config = _make_vs_config("vs_abc")
|
||||
|
||||
async def _fail(**kwargs: Any) -> None:
|
||||
raise RuntimeError("search failed")
|
||||
|
||||
with (
|
||||
patch.object(hook, "_is_native_file_search_provider", return_value=False),
|
||||
patch.object(hook, "_resolve_vector_stores", new_callable=AsyncMock, return_value=[vs_config]),
|
||||
patch("litellm.vector_stores.asearch", side_effect=_fail),
|
||||
):
|
||||
result = await hook.async_pre_call_hook(
|
||||
user_api_key_dict=MagicMock(),
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type=CallTypes.aresponses.value,
|
||||
)
|
||||
|
||||
# Tool stripped, LLM call proceeds without context
|
||||
assert result["tools"] is None
|
||||
assert "vs_search_results" not in result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Regression: chat completions path still works after refactor
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestChatCompletionsRegression:
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completions_path_injects_context(self, hook: VectorStorePreCallHook) -> None:
|
||||
"""async_get_chat_completion_prompt still works after shared search method refactor."""
|
||||
messages = [{"role": "user", "content": "what is litellm?"}]
|
||||
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}]
|
||||
non_default_params = {"tools": tools}
|
||||
|
||||
vs_config = _make_vs_config("vs_abc")
|
||||
search_response = _make_search_response(["LiteLLM is a unified LLM interface."])
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.pop_vector_stores_to_run_with_db_fallback = AsyncMock(
|
||||
return_value=[vs_config]
|
||||
)
|
||||
|
||||
with (
|
||||
patch("litellm.vector_store_registry", mock_registry),
|
||||
patch("litellm.vector_stores.asearch", new_callable=AsyncMock, return_value=search_response),
|
||||
):
|
||||
_, modified_messages, _ = await hook.async_get_chat_completion_prompt(
|
||||
model="claude-3-7-sonnet-20250219",
|
||||
messages=messages,
|
||||
non_default_params=non_default_params,
|
||||
prompt_id=None,
|
||||
prompt_variables=None,
|
||||
dynamic_callback_params=MagicMock(),
|
||||
litellm_logging_obj=MagicMock(),
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
# Context injected as a new message before the user's message
|
||||
assert len(modified_messages) == 2
|
||||
context_msg = modified_messages[0]
|
||||
assert "LiteLLM is a unified LLM interface" in str(context_msg["content"])
|
||||
assert modified_messages[-1] == messages[0]
|
||||
Loading…
Add table
Reference in a new issue