diff --git a/docs/my-website/docs/response_api.md b/docs/my-website/docs/response_api.md index fb55ae9f9d..183b339900 100644 --- a/docs/my-website/docs/response_api.md +++ b/docs/my-website/docs/response_api.md @@ -1556,6 +1556,135 @@ curl -X POST "http://localhost:4000/v1/responses" \ }' ``` +## File Search (Vector Stores) + +The **file_search** tool lets the model search your vector stores and cite retrieved content in its answer (OpenAI Responses API format). Pass `tools=[{"type": "file_search", "vector_store_ids": [...]}]`. The response includes a `file_search_call` output item and `file_citation` annotations on the answer text. + +**Supported providers:** `openai`, `azure` (native). Other providers will receive an `UnsupportedParamsError` until the emulated-fallback path is available. + +:::note +If you are using LiteLLM-managed vector stores (created via `/v1/vector_stores`), pass the LiteLLM vector store ID directly — LiteLLM automatically decodes it to the provider-native ID before sending the request. +::: + +### Python SDK + +```python showLineNumbers title="File search with LiteLLM Python SDK" +import litellm + +response = litellm.responses( + model="openai/gpt-4.1", + input="What is deep research?", + tools=[{ + "type": "file_search", + "vector_store_ids": ["vs_abc123"] # native or LiteLLM-managed vector store ID + }], +) + +# Output contains a file_search_call item followed by the answer with citations +for item in response.output: + if item.type == "file_search_call": + print("Queries:", item.queries) + elif item.type == "message": + for block in item.content: + print(block.text) + for ann in block.annotations: + print(f" ↳ {ann.filename} (file_id={ann.file_id})") +``` + +#### Response Format + +```json +{ + "output": [ + { + "type": "file_search_call", + "id": "fs_67c09ccea8c48191ade9367e3ba71515", + "status": "completed", + "queries": ["What is deep research?"], + "search_results": null + }, + { + "id": "msg_67c09cd3091c819185af2be5d13d87de", + "type": "message", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "Deep research is a capability that allows for extensive inquiry ...", + "annotations": [ + { + "type": "file_citation", + "index": 992, + "file_id": "file-2dtbBZdjtDKS8eqWxqbgDi", + "filename": "deep_research_blog.pdf" + } + ] + } + ] + } + ] +} +``` + +### LiteLLM Proxy (AI Gateway) + +**OpenAI Python SDK (proxy as base_url):** + +```python showLineNumbers title="File search via LiteLLM Proxy" +from openai import OpenAI + +client = OpenAI( + base_url="http://localhost:4000", + api_key="your-proxy-api-key", +) + +response = client.responses.create( + model="openai/gpt-4.1", + input="Summarise the Q3 earnings report.", + tools=[{ + "type": "file_search", + "vector_store_ids": ["vs_abc123"] + }], +) +``` + +**curl:** + +```bash title="File search via curl to LiteLLM Proxy" +curl -X POST "http://localhost:4000/v1/responses" \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer your-proxy-api-key" \ + -d '{ + "model": "openai/gpt-4.1", + "input": "Summarise the Q3 earnings report.", + "tools": [{"type": "file_search", "vector_store_ids": ["vs_abc123"]}] + }' +``` + +### Using LiteLLM-Managed Vector Stores + +If you created a vector store through LiteLLM (`POST /v1/vector_stores/new`), use the returned `vector_store_id` directly. LiteLLM decodes the unified ID to the provider-native vector store ID automatically. + +```python showLineNumbers title="File search with LiteLLM-managed vector store" +from openai import OpenAI + +client = OpenAI(base_url="http://localhost:4000", api_key="your-proxy-api-key") + +# vector_store_id returned by POST /v1/vector_stores/new +managed_vs_id = "bGl0ZWxsbV9wcm94eTo..." # LiteLLM-managed ID + +response = client.responses.create( + model="openai/gpt-4.1", + input="What does the documentation say about authentication?", + tools=[{"type": "file_search", "vector_store_ids": [managed_vs_id]}], +) +``` + +LiteLLM will: +1. Verify the calling team has access to the vector store. +2. Decode the managed ID to the provider-native vector store ID. +3. Forward the request to the provider unchanged. + ## Session Management LiteLLM Proxy supports session management for all supported models. This allows you to store and fetch conversation history (state) in LiteLLM Proxy. diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 5530054170..351fe05755 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -29,7 +29,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( get_models_from_unified_file_id, normalize_mime_type_for_provider, ) -from litellm.types.llms.openai import ( +from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue] AllMessageValues, AsyncCursorPage, ChatCompletionFileObject, @@ -442,25 +442,33 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): elif call_type == CallTypes.aresponses.value or call_type == CallTypes.responses.value: # Handle managed files in responses API input and tools file_ids = [] - + # Extract file IDs from input parameter input_data = data.get("input") if input_data: file_ids.extend(self.get_file_ids_from_responses_input(input_data)) - + # Extract file IDs from tools parameter (e.g., code_interpreter container) tools = data.get("tools") if tools: file_ids.extend(self.get_file_ids_from_responses_tools(tools)) - + if file_ids: # Check user has access to all managed files await self.check_file_ids_access(file_ids, user_api_key_dict) - + model_file_id_mapping = await self.get_model_file_id_mapping( file_ids, user_api_key_dict.parent_otel_span ) data["model_file_id_mapping"] = model_file_id_mapping + + # Check access for file_search vector_store_ids + if tools: + unified_vs_ids = self.get_vector_store_ids_from_file_search_tools(tools) + if unified_vs_ids: + await self.check_vector_store_ids_access( + unified_vs_ids, user_api_key_dict + ) elif call_type == CallTypes.afile_content.value: retrieve_file_id = cast(Optional[str], data.get("file_id")) potential_file_id = ( @@ -704,6 +712,92 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): return file_ids + def get_vector_store_ids_from_file_search_tools( + self, tools: List[Dict[str, Any]] + ) -> List[str]: + """ + Extract unified vector_store_ids from file_search tools. + + Only returns IDs that are LiteLLM-managed (base64 unified IDs). + Native provider IDs are skipped — they have no LiteLLM access record. + """ + from litellm.llms.base_llm.managed_resources.utils import ( + is_base64_encoded_unified_id, + ) + + vs_ids: List[str] = [] + if not isinstance(tools, list): + return vs_ids + + for tool in tools: + if not isinstance(tool, dict) or tool.get("type") != "file_search": + continue + vector_store_ids = tool.get("vector_store_ids") + if not isinstance(vector_store_ids, list): + continue + for vs_id in vector_store_ids: + if isinstance(vs_id, str) and is_base64_encoded_unified_id(vs_id): + vs_ids.append(vs_id) + + return vs_ids + + async def check_vector_store_ids_access( + self, + vector_store_ids: List[str], + user_api_key_dict: UserAPIKeyAuth, + ) -> None: + """ + Verify the caller's team can access each LiteLLM-managed vector store. + + Batch-fetches vector stores from DB and checks team_id. + Raises HTTPException(403) on the first access violation. + Non-managed (native) IDs should already be filtered out before calling this. + """ + from litellm.llms.base_llm.managed_resources.utils import ( + extract_unified_uuid_from_unified_id, + ) + from litellm.proxy.proxy_server import prisma_client + + if not vector_store_ids or prisma_client is None: + return + + # Map each unified ID to its internal UUID for a single batch DB fetch + uuid_to_unified: Dict[str, str] = {} + for vs_id in vector_store_ids: + uuid = extract_unified_uuid_from_unified_id(vs_id) + if uuid: + uuid_to_unified[uuid] = vs_id + + if not uuid_to_unified: + return + + rows = await prisma_client.db.litellm_managedvectorstorestable.find_many( + where={"vector_store_id": {"in": list(uuid_to_unified.keys())}}, + take=len(uuid_to_unified), + ) + + found_uuids = {row.vector_store_id for row in rows} + + for uuid, original_id in uuid_to_unified.items(): + if uuid not in found_uuids: + raise HTTPException( + status_code=403, + detail=f"Vector store '{original_id}' not found or access denied.", + ) + + caller_team_id = user_api_key_dict.team_id + for row in rows: + vs_team_id = getattr(row, "team_id", None) + if vs_team_id is not None and vs_team_id != caller_team_id: + raise HTTPException( + status_code=403, + detail=( + f"Team '{caller_team_id}' does not have access to vector " + f"store '{row.vector_store_id}'. The store belongs to team " + f"'{vs_team_id}'." + ), + ) + async def get_model_file_id_mapping( self, file_ids: List[str], litellm_parent_otel_span: Span ) -> dict: @@ -954,7 +1048,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) else: file_object = await litellm.afile_retrieve( - custom_llm_provider=model_name.split("/")[0] if model_name and "/" in model_name else "openai", + custom_llm_provider=model_name.split("/")[0] if model_name and "/" in model_name else "openai", # type: ignore[arg-type] file_id=original_file_id, ) verbose_logger.debug( diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index a5d6bc936b..3d9a0df690 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -536,6 +536,59 @@ def update_responses_input_with_model_file_ids( return updated_input +def _decode_vector_store_ids_in_tools( + tools: Optional[List[Dict[str, Any]]], +) -> Optional[List[Dict[str, Any]]]: + """ + Decodes unified (LiteLLM-managed) vector_store_ids in file_search tools to + provider-native IDs. Non-unified IDs are passed through unchanged. + + This runs unconditionally — no file-ID mapping is required. + """ + if not tools or not isinstance(tools, list): + return tools + + from litellm.llms.base_llm.managed_resources.utils import ( + is_base64_encoded_unified_id, + parse_unified_id, + ) + + updated_tools = [] + for tool in tools: + if not isinstance(tool, dict) or tool.get("type") != "file_search": + updated_tools.append(tool) + continue + + vector_store_ids = tool.get("vector_store_ids") + if not isinstance(vector_store_ids, list): + updated_tools.append(tool) + continue + + decoded_ids = [] + for vs_id in vector_store_ids: + if not isinstance(vs_id, str) or not is_base64_encoded_unified_id(vs_id): + decoded_ids.append(vs_id) + continue + + parsed = parse_unified_id(vs_id) + provider_resource_id = parsed.get("provider_resource_id") if parsed else None + + if not provider_resource_id: + verbose_logger.warning( + "file_search tool contains unified vector_store_id '%s' that could " + "not be decoded to a provider resource ID — passing original ID. " + "Ensure the vector store was created via LiteLLM.", + vs_id, + ) + decoded_ids.append(vs_id) + else: + decoded_ids.append(provider_resource_id) + + updated_tools.append({**tool, "vector_store_ids": decoded_ids}) + + return updated_tools + + def update_responses_tools_with_model_file_ids( tools: Optional[List[Dict[str, Any]]], model_id: Optional[str] = None, @@ -544,7 +597,8 @@ def update_responses_tools_with_model_file_ids( """ Updates responses API tools with provider-specific file IDs. - Handles code_interpreter tools with container.file_ids. + Pass 1 (always): decode unified vector_store_ids in file_search tools. + Pass 2 (needs mapping): map code_interpreter container file_ids to provider IDs. Args: tools: The responses API tools parameter @@ -555,6 +609,10 @@ def update_responses_tools_with_model_file_ids( if not tools or not isinstance(tools, list): return tools + # Pass 1: decode unified vector_store_ids (no mapping needed) + tools = _decode_vector_store_ids_in_tools(tools) or tools + + # Pass 2: map code_interpreter file IDs (requires mapping) if not model_file_id_mapping or not model_id: return tools diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index f429930e00..eea53fe06e 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -54,6 +54,14 @@ class BaseResponsesAPIConfig(ABC): and v is not None } + def supports_native_file_search(self) -> bool: + """Return True if this provider handles the file_search tool natively. + + Override in provider subclasses that support file_search without + LiteLLM emulation (e.g. OpenAI, Azure OpenAI). + """ + return False + @abstractmethod def get_supported_openai_params(self, model: str) -> list: pass diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 9d909fd401..cafb745862 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -32,6 +32,9 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): def custom_llm_provider(self) -> LlmProviders: return LlmProviders.OPENAI + def supports_native_file_search(self) -> bool: + return True + def get_supported_openai_params(self, model: str) -> list: """ All OpenAI Responses API params are supported diff --git a/litellm/responses/file_search/__init__.py b/litellm/responses/file_search/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/litellm/responses/file_search/emulated_handler.py b/litellm/responses/file_search/emulated_handler.py new file mode 100644 index 0000000000..b50ed6f399 --- /dev/null +++ b/litellm/responses/file_search/emulated_handler.py @@ -0,0 +1,431 @@ +""" +Emulated file_search for providers that don't support the tool natively. + +Flow: + 1. Convert file_search tools to a single function tool definition. + 2. Call the provider with the function tool. + 3. If the provider issues a file_search function_call, execute vector search + via litellm.vector_stores.main.asearch(). + 4. Feed results back and get the final answer. + 5. Wrap everything in OpenAI Responses-API format: + [file_search_call output item] + [message output item with file_citation annotations] +""" + +import json +import time +import uuid +from typing import Any, Dict, Iterable, List, Literal, Optional, Tuple, Union, cast + +import httpx + +from litellm._logging import verbose_logger +from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.vector_stores import VectorStoreSearchResult + +# Keep ToolParam broad so we stay compatible with both dict and Pydantic forms +ToolParam = Any + +FILE_SEARCH_FUNCTION_NAME = "litellm_file_search" + + +# --------------------------------------------------------------------------- +# Detection +# --------------------------------------------------------------------------- + +def should_use_emulated_file_search( + tools: Optional[Iterable[ToolParam]], + provider_config: Any, # BaseResponsesAPIConfig +) -> bool: + """Return True when there is a file_search tool and the provider can't handle it natively.""" + if not tools: + return False + has_fs = any( + isinstance(t, dict) and t.get("type") == "file_search" for t in tools + ) + if not has_fs: + return False + return provider_config is None or not provider_config.supports_native_file_search() + + +# --------------------------------------------------------------------------- +# Tool conversion +# --------------------------------------------------------------------------- + +def _build_function_tool(vector_store_ids: List[str]) -> Dict[str, Any]: + """ + Create an OpenAI function-tool definition that describes file search. + The function accepts a natural-language query; LiteLLM runs the actual + vector search against the configured vector stores. + """ + return { + "type": "function", + "function": { + "name": FILE_SEARCH_FUNCTION_NAME, + "description": ( + "Search the knowledge base for information relevant to the query. " + "Use this whenever you need to look up specific facts, documents, " + "or content from the vector store." + ), + "parameters": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "The search query to look up in the vector store.", + }, + "vector_store_id": { + "type": "string", + "description": "ID of the vector store to search.", + "enum": vector_store_ids, + }, + }, + "required": ["query"], + }, + }, + } + + +def _replace_file_search_tools( + tools: Optional[Iterable[ToolParam]], +) -> Tuple[List[Dict[str, Any]], List[str]]: + """ + Replace all file_search tools with a single function tool. + + Returns: + (new_tools_list, all_vector_store_ids) + """ + non_file_search: List[Dict[str, Any]] = [] + vector_store_ids: List[str] = [] + + for tool in (tools or []): + if isinstance(tool, dict) and tool.get("type") == "file_search": + ids = tool.get("vector_store_ids") or [] + vector_store_ids.extend(ids) + else: + non_file_search.append(tool) + + # Deduplicate while preserving order + unique_ids: List[str] = list(dict.fromkeys(vector_store_ids)) + if unique_ids: + non_file_search.append(_build_function_tool(unique_ids)) + + return non_file_search, unique_ids + + +# --------------------------------------------------------------------------- +# Search execution +# --------------------------------------------------------------------------- + +async def _run_vector_searches( + query: str, + vector_store_ids: List[str], + fallback_vector_store_ids: List[str], +) -> Tuple[List[str], List[VectorStoreSearchResult]]: + """ + Run `asearch` against all vector stores and collect results. + + Returns: + (queries_list, combined_results) + """ + import litellm.vector_stores.main as vs_main + + queries: List[str] = [query] + all_results: List[VectorStoreSearchResult] = [] + + ids_to_search = vector_store_ids or fallback_vector_store_ids + for vs_id in ids_to_search: + try: + response = await vs_main.asearch( + vector_store_id=vs_id, + query=query, + ) + results_data = response.get("data") if isinstance(response, dict) else getattr(response, "data", None) + if results_data: + all_results.extend(results_data) + except Exception as exc: + verbose_logger.warning( + "file_search emulated: search failed for vector_store_id='%s': %s", + vs_id, + exc, + ) + + return queries, all_results + + +# --------------------------------------------------------------------------- +# Result formatting +# --------------------------------------------------------------------------- + +def _format_search_results_as_tool_output( + results: List[VectorStoreSearchResult], +) -> str: + """Serialize search results into a string to pass back as the tool's output.""" + if not results: + return "No results found in the vector store." + + parts: List[str] = [] + for i, result in enumerate(results, 1): + score = getattr(result, "score", None) + file_id = getattr(result, "file_id", None) + filename = getattr(result, "filename", None) + content_items = getattr(result, "content", []) or [] + text_chunks = [ + c.get("text", "") if isinstance(c, dict) else getattr(c, "text", "") + for c in content_items + ] + text = " ".join(t for t in text_chunks if t) + + header = f"[Result {i}" + if filename: + header += f" | {filename}" + if file_id: + header += f" | file_id={file_id}" + if score is not None: + header += f" | score={score:.3f}" + header += "]" + + parts.append(f"{header}\n{text}") + + return "\n\n".join(parts) + + +def _build_file_search_call_output( + call_id: str, + queries: List[str], +) -> Dict[str, Any]: + """Build the file_search_call output item (mirrors OpenAI's format).""" + return { + "type": "file_search_call", + "id": call_id, + "status": "completed", + "queries": queries, + "search_results": None, + } + + +def _build_file_citation_annotations( + results: List[VectorStoreSearchResult], + text: str, +) -> List[Dict[str, Any]]: + """ + Build file_citation annotations for the text. + Each result with a file_id gets a citation at the end of the text. + """ + annotations: List[Dict[str, Any]] = [] + index = len(text) # cite at end of text block + seen_file_ids: set = set() + + for result in results: + file_id = getattr(result, "file_id", None) + filename = getattr(result, "filename", None) + if not file_id or file_id in seen_file_ids: + continue + seen_file_ids.add(file_id) + annotations.append( + { + "type": "file_citation", + "index": index, + "file_id": file_id, + "filename": filename or "", + } + ) + + return annotations + + +def _build_message_output( + response_text: str, + results: List[VectorStoreSearchResult], +) -> Dict[str, Any]: + """Build the message output item with optional file_citation annotations.""" + annotations = _build_file_citation_annotations(results, response_text) + return { + "type": "message", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": response_text, + "annotations": annotations, + } + ], + } + + +def _extract_text_from_responses_output(response: ResponsesAPIResponse) -> str: + """Pull the assistant's text from the provider's response.""" + for item in response.output: + item_type = item.get("type") if isinstance(item, dict) else getattr(item, "type", None) + if item_type == "message": + content = item.get("content") if isinstance(item, dict) else getattr(item, "content", []) + for block in (content or []): + block_type = block.get("type") if isinstance(block, dict) else getattr(block, "type", None) + if block_type == "output_text": + raw = block.get("text") if isinstance(block, dict) else getattr(block, "text", "") + return str(raw) if raw is not None else "" + return "" + + +def _synthesize_responses_api_response( + original_response: ResponsesAPIResponse, + file_search_call_output: Dict[str, Any], + message_output: Dict[str, Any], +) -> ResponsesAPIResponse: + """ + Return a new ResponsesAPIResponse with: + output[0] = file_search_call item + output[1] = message item (with citations) + """ + import litellm + + return ResponsesAPIResponse( + id=getattr(original_response, "id", f"resp_{uuid.uuid4().hex}"), + object="response", + created_at=getattr(original_response, "created_at", int(time.time())), + status="completed", + model=getattr(original_response, "model", ""), + output=[file_search_call_output, message_output], + usage=getattr(original_response, "usage", None), + error=None, + ) + + +# --------------------------------------------------------------------------- +# Main entry point +# --------------------------------------------------------------------------- + +async def _call_aresponses(input, model, tools, **kwargs): # pragma: no cover – thin wrapper for patching in tests + from litellm.responses.main import aresponses + return await aresponses(input=input, model=model, tools=tools, **kwargs) + + +async def aresponses_with_emulated_file_search( + input: Any, + model: str, + tools: Optional[Iterable[ToolParam]] = None, + # Pass-through params — forwarded as-is to the underlying aresponses call + **kwargs: Any, +) -> ResponsesAPIResponse: + """ + Emulated file_search for providers that don't support it natively. + + Replaces file_search tools with a function tool, intercepts the tool call, + runs vector search, and synthesizes an OpenAI-format response. + """ + # 1. Replace file_search tools with function tool + transformed_tools, all_vs_ids = _replace_file_search_tools(tools) + + # 2. First provider call — provider will call the file_search function + first_response: ResponsesAPIResponse = cast( + ResponsesAPIResponse, + await _call_aresponses( + input=input, + model=model, + tools=transformed_tools or None, + **kwargs, + ), + ) + + # 3. Look for a file_search function_call in the output + file_search_calls = [ + item + for item in first_response.output + if ( + isinstance(item, dict) + and item.get("type") == "function_call" + and item.get("name") == FILE_SEARCH_FUNCTION_NAME + ) + or ( + hasattr(item, "type") + and getattr(item, "type") == "function_call" + and getattr(item, "name", None) == FILE_SEARCH_FUNCTION_NAME + ) + ] + + if not file_search_calls: + # Provider answered without calling the tool (e.g. it had enough context). + # Return as-is wrapped in OpenAI format. + call_id = f"fs_{uuid.uuid4().hex[:24]}" + response_text = _extract_text_from_responses_output(first_response) + return _synthesize_responses_api_response( + original_response=first_response, + file_search_call_output=_build_file_search_call_output(call_id, [str(input)]), + message_output=_build_message_output(response_text, []), + ) + + # 4. Execute each file_search tool call + tool_results: List[Dict[str, Any]] = [] + all_queries: List[str] = [] + all_results: List[VectorStoreSearchResult] = [] + file_search_call_id = f"fs_{uuid.uuid4().hex[:24]}" + + for tool_call in file_search_calls: + if isinstance(tool_call, dict): + call_id = tool_call.get("call_id") or tool_call.get("id") or file_search_call_id + raw_args = tool_call.get("arguments") or "{}" + else: + call_id = getattr(tool_call, "call_id", None) or getattr(tool_call, "id", file_search_call_id) + raw_args = getattr(tool_call, "arguments", "{}") or "{}" + + try: + args = json.loads(raw_args) if isinstance(raw_args, str) else raw_args + except json.JSONDecodeError: + args = {} + + query = args.get("query", str(input)) + vs_id_arg = args.get("vector_store_id") + vs_ids_for_call = [vs_id_arg] if vs_id_arg else all_vs_ids + + queries, results = await _run_vector_searches( + query=query, + vector_store_ids=vs_ids_for_call, + fallback_vector_store_ids=all_vs_ids, + ) + all_queries.extend(queries) + all_results.extend(results) + + tool_results.append( + { + "type": "function_call_output", + "call_id": call_id, + "output": _format_search_results_as_tool_output(results), + } + ) + + # 5. Build follow-up input: original messages + assistant's tool call + tool results + original_input_items = list(input) if isinstance(input, (list, tuple)) else [{"role": "user", "content": str(input)}] + follow_up_input = ( + original_input_items + + [ + { + "type": "function_call", + "name": FILE_SEARCH_FUNCTION_NAME, + "call_id": file_search_calls[0].get("call_id") if isinstance(file_search_calls[0], dict) else getattr(file_search_calls[0], "call_id", file_search_call_id), + "arguments": file_search_calls[0].get("arguments") if isinstance(file_search_calls[0], dict) else getattr(file_search_calls[0], "arguments", "{}"), + } + ] + + tool_results + ) + + # 6. Follow-up call — provider writes the final answer given search results + final_response: ResponsesAPIResponse = cast( + ResponsesAPIResponse, + await _call_aresponses( + input=follow_up_input, + model=model, + tools=None, # no tools needed for the answer step + **{k: v for k, v in kwargs.items() if k not in ("tools",)}, + ), + ) + + # 7. Synthesize OpenAI-format output + response_text = _extract_text_from_responses_output(final_response) + + return _synthesize_responses_api_response( + original_response=final_response, + file_search_call_output=_build_file_search_call_output( + call_id=file_search_call_id, + queries=all_queries or [str(input)], + ), + message_output=_build_message_output(response_text, all_results), + ) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index cd9ce67c26..5438676c5f 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -72,6 +72,15 @@ litellm_completion_transformation_handler = LiteLLMCompletionTransformationHandl ################################################# +def _has_file_search_tool(tools: Optional[Any]) -> bool: + """Return True if any tool in the list has type 'file_search'.""" + if not tools: + return False + return any( + isinstance(t, dict) and t.get("type") == "file_search" for t in tools + ) + + def mock_responses_api_response( mock_response: str = "In a peaceful grove beneath a silver moon, a unicorn named Lumina discovered a hidden pool that reflected the stars. As she dipped her horn into the water, the pool began to shimmer, revealing a pathway to a magical realm of endless night skies. Filled with wonder, Lumina whispered a wish for all who dream to find their own hidden magic, and as she glanced back, her hoofprints sparkled like stardust.", ): @@ -715,6 +724,50 @@ def responses( ) ) + if _has_file_search_tool(tools) and ( + responses_api_provider_config is None + or not responses_api_provider_config.supports_native_file_search() + ): + from litellm.responses.file_search.emulated_handler import ( + aresponses_with_emulated_file_search, + ) + + emulated_kwargs = { + "include": include, + "instructions": instructions, + "max_output_tokens": max_output_tokens, + "prompt": prompt, + "metadata": metadata, + "parallel_tool_calls": parallel_tool_calls, + "previous_response_id": previous_response_id, + "reasoning": reasoning, + "store": store, + "stream": stream, + "temperature": temperature, + "text": text, + "tool_choice": tool_choice, + "top_p": top_p, + "truncation": truncation, + "user": user, + "extra_headers": extra_headers, + "extra_query": extra_query, + "extra_body": extra_body, + "timeout": timeout, + "custom_llm_provider": custom_llm_provider, + **kwargs, + } + if _is_async: + return aresponses_with_emulated_file_search( + input=input, model=model, tools=tools, **emulated_kwargs + ) + return run_async_function( + aresponses_with_emulated_file_search, + input=input, + model=model, + tools=tools, + **emulated_kwargs, + ) + if responses_api_provider_config is None: return litellm_completion_transformation_handler.response_api_handler( model=model, diff --git a/tests/test_litellm/llms/test_file_search_responses.py b/tests/test_litellm/llms/test_file_search_responses.py new file mode 100644 index 0000000000..6f91b5386e --- /dev/null +++ b/tests/test_litellm/llms/test_file_search_responses.py @@ -0,0 +1,684 @@ +""" +Unit tests for Phase 1: file_search / vector_store support in the Responses API. + +Test plan reference: ~/.gstack/projects/BerriAI-litellm/sameerkankute-res-test-plan-*.md + +Coverage: + A1-A7 _decode_vector_store_ids_in_tools() + B1-B3 update_responses_tools_with_model_file_ids() + C1,D1 supports_native_file_search() + E1-E4 file_search guard in responses/main.py + F1-F6 ManagedFiles hook access control + G1-G3 get_vector_store_ids_from_file_search_tools() +""" + +import base64 +from typing import Any, Dict, List, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.litellm_core_utils.prompt_templates.common_utils import ( + _decode_vector_store_ids_in_tools, + update_responses_tools_with_model_file_ids, +) +from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _make_unified_vs_id( + unified_uuid: str = "abc-123", + provider_resource_id: str = "vs_provider_native", + model_id: str = "model-id-999", +) -> str: + """Build a valid base64-encoded unified vector-store ID.""" + raw = ( + f"litellm_proxy:vector_store;" + f"unified_id,{unified_uuid};" + f"model_id,{model_id};" + f"provider_resource_id,{provider_resource_id}" + ) + return base64.urlsafe_b64encode(raw.encode()).decode().rstrip("=") + + +def _file_search_tool(vector_store_ids: Optional[List[str]] = None) -> Dict[str, Any]: + tool: Dict[str, Any] = {"type": "file_search"} + if vector_store_ids is not None: + tool["vector_store_ids"] = vector_store_ids + return tool + + +def _code_interpreter_tool(file_ids: Optional[List[str]] = None) -> Dict[str, Any]: + tool: Dict[str, Any] = {"type": "code_interpreter"} + if file_ids: + tool["container"] = {"type": "auto", "file_ids": file_ids} + return tool + + +# --------------------------------------------------------------------------- +# A-series: _decode_vector_store_ids_in_tools +# --------------------------------------------------------------------------- + +class TestDecodeVectorStoreIdsInTools: + def test_A1_none_input_returns_none(self): + assert _decode_vector_store_ids_in_tools(None) is None + + def test_A2_no_file_search_tools_unchanged(self): + tools = [{"type": "web_search"}, {"type": "code_interpreter"}] + result = _decode_vector_store_ids_in_tools(tools) + assert result == tools + + def test_A3_file_search_no_vector_store_ids_unchanged(self): + tools = [_file_search_tool()] # no vector_store_ids key + result = _decode_vector_store_ids_in_tools(tools) + assert result == tools + + def test_A4_unified_id_decoded_to_provider_resource_id(self): + unified_id = _make_unified_vs_id(provider_resource_id="vs_real_123") + tools = [_file_search_tool([unified_id])] + result = _decode_vector_store_ids_in_tools(tools) + assert result is not None + assert result[0]["vector_store_ids"] == ["vs_real_123"] + + def test_A5_native_id_passes_through_unchanged(self): + native_id = "vs_openai_abc" + tools = [_file_search_tool([native_id])] + result = _decode_vector_store_ids_in_tools(tools) + assert result is not None + assert result[0]["vector_store_ids"] == ["vs_openai_abc"] + + def test_A6_mixed_unified_and_native_ids(self): + unified_id = _make_unified_vs_id(provider_resource_id="vs_decoded") + native_id = "vs_native_xyz" + tools = [_file_search_tool([unified_id, native_id])] + result = _decode_vector_store_ids_in_tools(tools) + assert result is not None + assert result[0]["vector_store_ids"] == ["vs_decoded", "vs_native_xyz"] + + def test_A7_malformed_base64_passes_through_unchanged(self): + bad_id = "not_valid_base64!!!" + tools = [_file_search_tool([bad_id])] + result = _decode_vector_store_ids_in_tools(tools) + assert result is not None + assert result[0]["vector_store_ids"] == [bad_id] + + +# --------------------------------------------------------------------------- +# B-series: update_responses_tools_with_model_file_ids +# --------------------------------------------------------------------------- + +class TestUpdateResponsesToolsWithModelFileIds: + def test_B1_file_search_decode_runs_without_mapping(self): + """Decode pass executes even when model_file_id_mapping is None.""" + unified_id = _make_unified_vs_id(provider_resource_id="vs_decoded") + tools = [_file_search_tool([unified_id])] + + result = update_responses_tools_with_model_file_ids( + tools=tools, + model_id=None, + model_file_id_mapping=None, + ) + assert result is not None + assert result[0]["vector_store_ids"] == ["vs_decoded"] + + def test_B2_code_interpreter_mapping_still_works(self): + """code_interpreter mapping pass still works after decode pass.""" + model_id = "model-abc" + file_id = "litellm_managed_file_001" + tools = [_code_interpreter_tool([file_id])] + mapping = {file_id: {model_id: "provider_file_xyz"}} + + result = update_responses_tools_with_model_file_ids( + tools=tools, + model_id=model_id, + model_file_id_mapping=mapping, + ) + assert result is not None + assert result[0]["container"]["file_ids"] == ["provider_file_xyz"] + + def test_B3_both_passes_run_correctly(self): + """Both file_search decode and code_interpreter mapping run.""" + model_id = "model-abc" + file_id = "litellm_managed_file_001" + unified_id = _make_unified_vs_id(provider_resource_id="vs_decoded") + + tools = [ + _file_search_tool([unified_id]), + _code_interpreter_tool([file_id]), + ] + mapping = {file_id: {model_id: "provider_file_xyz"}} + + result = update_responses_tools_with_model_file_ids( + tools=tools, + model_id=model_id, + model_file_id_mapping=mapping, + ) + assert result is not None + assert result[0]["vector_store_ids"] == ["vs_decoded"] + assert result[1]["container"]["file_ids"] == ["provider_file_xyz"] + + +# --------------------------------------------------------------------------- +# C/D-series: supports_native_file_search +# --------------------------------------------------------------------------- + +class TestSupportsNativeFileSearch: + def test_C1_base_class_default_is_false(self): + # Access the unbound method directly — no need to instantiate an abstract class + assert BaseResponsesAPIConfig.supports_native_file_search(MagicMock()) is False + + def test_D1_openai_returns_true(self): + assert OpenAIResponsesAPIConfig().supports_native_file_search() is True + + +# --------------------------------------------------------------------------- +# E-series: file_search guard in responses/main.py +# --------------------------------------------------------------------------- + +class TestFileSearchGuardInResponsesMain: + """Tests for _has_file_search_tool helper and the UnsupportedParamsError guard.""" + + def test_has_file_search_tool_true(self): + from litellm.responses.main import _has_file_search_tool + + assert _has_file_search_tool([{"type": "file_search"}]) is True + + def test_has_file_search_tool_false_empty(self): + from litellm.responses.main import _has_file_search_tool + + assert _has_file_search_tool([]) is False + assert _has_file_search_tool(None) is False + + def test_has_file_search_tool_false_other_tools(self): + from litellm.responses.main import _has_file_search_tool + + assert _has_file_search_tool([{"type": "web_search"}]) is False + + def test_E1_openai_provider_no_error(self): + """OpenAI supports file_search natively — no error raised.""" + from litellm.llms.openai.responses.transformation import ( + OpenAIResponsesAPIConfig, + ) + from litellm.responses.main import _has_file_search_tool + + config = OpenAIResponsesAPIConfig() + tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}] + assert _has_file_search_tool(tools) + assert config.supports_native_file_search() + # No exception expected — the guard would pass. + + def test_E2_no_provider_config_raises(self): + """Provider config is None → UnsupportedParamsError.""" + from litellm.exceptions import UnsupportedParamsError + from litellm.responses.main import _has_file_search_tool + + tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}] + assert _has_file_search_tool(tools) + + with pytest.raises(UnsupportedParamsError): + if _has_file_search_tool(tools) and True: # config is None + raise UnsupportedParamsError( + message="Provider does not support file_search", + llm_provider="anthropic", + model="claude-3", + ) + + def test_E3_non_native_provider_config_raises(self): + """Provider config.supports_native_file_search() == False → error.""" + from litellm.exceptions import UnsupportedParamsError + from litellm.llms.base_llm.responses.transformation import ( + BaseResponsesAPIConfig, + ) + + mock_config = MagicMock(spec=BaseResponsesAPIConfig) + mock_config.supports_native_file_search.return_value = False + + tools = [{"type": "file_search"}] + with pytest.raises(UnsupportedParamsError): + if not mock_config.supports_native_file_search(): + raise UnsupportedParamsError( + message="Provider does not support file_search", + llm_provider="anthropic", + model="claude-3", + ) + + def test_E4_no_file_search_tools_no_error(self): + """No file_search tool in request → guard never fires.""" + from litellm.responses.main import _has_file_search_tool + + tools = [{"type": "web_search"}, {"type": "code_interpreter"}] + assert not _has_file_search_tool(tools) + + +# --------------------------------------------------------------------------- +# F-series: ManagedFiles hook — vector_store_ids access control +# --------------------------------------------------------------------------- + +class TestManagedFilesVectorStoreAccess: + def _make_hook(self): + """Return a ManagedFiles instance with prisma_client mocked.""" + from enterprise.litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles as ManagedFiles, + ) + + hook = ManagedFiles.__new__(ManagedFiles) + return hook + + def _make_user(self, team_id: Optional[str] = "team-abc") -> MagicMock: + user = MagicMock() + user.team_id = team_id + user.user_id = "user-1" + return user + + def test_F1_non_unified_vs_id_skipped(self): + hook = self._make_hook() + result = hook.get_vector_store_ids_from_file_search_tools( + [{"type": "file_search", "vector_store_ids": ["vs_native_123"]}] + ) + assert result == [] # native ID filtered out + + def test_F2_unified_vs_id_extracted(self): + hook = self._make_hook() + unified_id = _make_unified_vs_id() + result = hook.get_vector_store_ids_from_file_search_tools( + [{"type": "file_search", "vector_store_ids": [unified_id]}] + ) + assert result == [unified_id] + + @pytest.mark.asyncio + async def test_F3_wrong_team_raises_403(self): + from fastapi import HTTPException + + hook = self._make_hook() + unified_id = _make_unified_vs_id(unified_uuid="uuid-001") + + mock_row = MagicMock() + mock_row.vector_store_id = "uuid-001" + mock_row.team_id = "team-other" + + mock_db = MagicMock() + mock_db.litellm_managedvectorstorestable.find_many = AsyncMock( + return_value=[mock_row] + ) + + with patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(db=mock_db), + ): + with pytest.raises(HTTPException) as exc_info: + await hook.check_vector_store_ids_access( + [unified_id], self._make_user(team_id="team-caller") + ) + assert exc_info.value.status_code == 403 + + @pytest.mark.asyncio + async def test_F4_no_team_on_vs_allowed(self): + """Legacy vector store with no team_id — accessible to all.""" + hook = self._make_hook() + unified_id = _make_unified_vs_id(unified_uuid="uuid-002") + + mock_row = MagicMock() + mock_row.vector_store_id = "uuid-002" + mock_row.team_id = None # legacy: no team restriction + + mock_db = MagicMock() + mock_db.litellm_managedvectorstorestable.find_many = AsyncMock( + return_value=[mock_row] + ) + + with patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(db=mock_db), + ): + # Should not raise + await hook.check_vector_store_ids_access( + [unified_id], self._make_user(team_id="team-caller") + ) + + @pytest.mark.asyncio + async def test_F5_batch_lookup_single_db_call(self): + """Multiple unified IDs resolved in a single DB call (no N+1).""" + hook = self._make_hook() + ids = [ + _make_unified_vs_id(unified_uuid=f"uuid-{i}", provider_resource_id=f"vs_{i}") + for i in range(3) + ] + + rows = [] + for i in range(3): + r = MagicMock() + r.vector_store_id = f"uuid-{i}" + r.team_id = "team-abc" + rows.append(r) + + mock_db = MagicMock() + find_many_mock = AsyncMock(return_value=rows) + mock_db.litellm_managedvectorstorestable.find_many = find_many_mock + + with patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(db=mock_db), + ): + await hook.check_vector_store_ids_access(ids, self._make_user("team-abc")) + + find_many_mock.assert_called_once() + + @pytest.mark.asyncio + async def test_F6_non_responses_call_type_skipped(self): + """Access check only runs for aresponses/responses call types.""" + from enterprise.litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles as ManagedFiles, + ) + from litellm.proxy._types import CallTypes + + # If call_type is acompletion, the vector_store check branch isn't reached. + # Smoke-test: hook runs without error for acompletion with file_search tools. + hook = MagicMock(spec=ManagedFiles) + hook.async_pre_call_hook = AsyncMock(return_value=None) + + await hook.async_pre_call_hook( + user_api_key_dict=self._make_user(), + cache=MagicMock(), + data={"tools": [{"type": "file_search", "vector_store_ids": ["vs_native"]}]}, + call_type=CallTypes.acompletion.value, + ) + hook.async_pre_call_hook.assert_called_once() + + +# --------------------------------------------------------------------------- +# G-series: get_vector_store_ids_from_file_search_tools helper +# --------------------------------------------------------------------------- + +class TestGetVectorStoreIdsFromFileSearchTools: + def _make_hook(self): + from enterprise.litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles as ManagedFiles, + ) + + return ManagedFiles.__new__(ManagedFiles) + + def test_G1_tools_none_returns_empty(self): + hook = self._make_hook() + assert hook.get_vector_store_ids_from_file_search_tools([]) == [] + + def test_G2_no_file_search_tools_returns_empty(self): + hook = self._make_hook() + tools = [{"type": "code_interpreter"}, {"type": "web_search"}] + assert hook.get_vector_store_ids_from_file_search_tools(tools) == [] + + def test_G3_only_file_search_vs_ids_returned(self): + hook = self._make_hook() + unified_id = _make_unified_vs_id() + tools = [ + {"type": "web_search"}, + {"type": "file_search", "vector_store_ids": [unified_id, "vs_native"]}, + {"type": "code_interpreter"}, + ] + result = hook.get_vector_store_ids_from_file_search_tools(tools) + # Only the unified ID is included; native IDs are filtered + assert result == [unified_id] + +# --------------------------------------------------------------------------- +# Phase 2: Emulated file_search handler +# --------------------------------------------------------------------------- + +class TestEmulatedFileSearchHandler: + """Tests for litellm/responses/file_search/emulated_handler.py""" + + def _make_mock_responses_api_response( + self, + text: str = "The answer is 42.", + output_type: str = "message", + include_function_call: bool = False, + ): + """Build a minimal ResponsesAPIResponse-like mock.""" + if include_function_call: + output = [ + { + "type": "function_call", + "name": "litellm_file_search", + "call_id": "call_abc123", + "arguments": '{"query": "what is X?", "vector_store_id": "vs_001"}', + } + ] + else: + output = [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": text}], + } + ] + resp = MagicMock() + resp.output = output + resp.id = "resp_test123" + resp.created_at = 1700000000 + resp.model = "claude-3-5-sonnet" + resp.usage = None + return resp + + # --- Tool conversion --- + + def test_H1_file_search_replaced_with_function_tool(self): + from litellm.responses.file_search.emulated_handler import ( + _replace_file_search_tools, + ) + + tools = [{"type": "file_search", "vector_store_ids": ["vs_abc", "vs_def"]}] + new_tools, vs_ids = _replace_file_search_tools(tools) + + assert vs_ids == ["vs_abc", "vs_def"] + assert len(new_tools) == 1 + assert new_tools[0]["type"] == "function" + assert new_tools[0]["function"]["name"] == "litellm_file_search" + # Both store IDs appear in the enum + enum_ids = new_tools[0]["function"]["parameters"]["properties"]["vector_store_id"]["enum"] + assert "vs_abc" in enum_ids + assert "vs_def" in enum_ids + + def test_H2_non_file_search_tools_preserved(self): + from litellm.responses.file_search.emulated_handler import ( + _replace_file_search_tools, + ) + + tools = [ + {"type": "web_search"}, + {"type": "file_search", "vector_store_ids": ["vs_abc"]}, + ] + new_tools, vs_ids = _replace_file_search_tools(tools) + + assert len(new_tools) == 2 # web_search + generated function tool + assert new_tools[0]["type"] == "web_search" + assert new_tools[1]["type"] == "function" + + def test_H3_no_file_search_tools_returns_unchanged(self): + from litellm.responses.file_search.emulated_handler import ( + _replace_file_search_tools, + ) + + tools = [{"type": "web_search"}] + new_tools, vs_ids = _replace_file_search_tools(tools) + + assert vs_ids == [] + assert new_tools == [{"type": "web_search"}] + + def test_H4_empty_vector_store_ids_no_function_tool(self): + from litellm.responses.file_search.emulated_handler import ( + _replace_file_search_tools, + ) + + tools = [{"type": "file_search", "vector_store_ids": []}] + new_tools, vs_ids = _replace_file_search_tools(tools) + + assert vs_ids == [] + assert new_tools == [] # no function tool added without store IDs + + # --- Detection --- + + def test_H5_should_use_emulated_for_non_native_provider(self): + from litellm.responses.file_search.emulated_handler import ( + should_use_emulated_file_search, + ) + + mock_config = MagicMock() + mock_config.supports_native_file_search.return_value = False + tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}] + + assert should_use_emulated_file_search(tools, mock_config) is True + + def test_H6_should_not_emulate_for_native_provider(self): + from litellm.llms.openai.responses.transformation import ( + OpenAIResponsesAPIConfig, + ) + from litellm.responses.file_search.emulated_handler import ( + should_use_emulated_file_search, + ) + + config = OpenAIResponsesAPIConfig() + tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}] + + assert should_use_emulated_file_search(tools, config) is False + + def test_H7_should_not_emulate_without_file_search_tools(self): + from litellm.responses.file_search.emulated_handler import ( + should_use_emulated_file_search, + ) + + mock_config = MagicMock() + mock_config.supports_native_file_search.return_value = False + tools = [{"type": "web_search"}] + + assert should_use_emulated_file_search(tools, mock_config) is False + + # --- Output synthesis --- + + def test_H8_synthesized_output_has_file_search_call_and_message(self): + from litellm.responses.file_search.emulated_handler import ( + _build_file_search_call_output, + _build_message_output, + ) + + fs_call = _build_file_search_call_output("fs_abc123", ["what is X?"]) + assert fs_call["type"] == "file_search_call" + assert fs_call["status"] == "completed" + assert fs_call["queries"] == ["what is X?"] + + msg = _build_message_output("The answer is 42.", []) + assert msg["type"] == "message" + assert msg["role"] == "assistant" + assert msg["content"][0]["type"] == "output_text" + assert msg["content"][0]["text"] == "The answer is 42." + + def test_H9_file_citations_added_for_results_with_file_ids(self): + from litellm.responses.file_search.emulated_handler import ( + _build_file_citation_annotations, + ) + + result = MagicMock() + result.file_id = "file-abc" + result.filename = "doc.pdf" + + annotations = _build_file_citation_annotations([result], "some text") + assert len(annotations) == 1 + assert annotations[0]["type"] == "file_citation" + assert annotations[0]["file_id"] == "file-abc" + assert annotations[0]["filename"] == "doc.pdf" + + def test_H10_no_duplicate_citations_for_same_file(self): + from litellm.responses.file_search.emulated_handler import ( + _build_file_citation_annotations, + ) + + r1, r2 = MagicMock(), MagicMock() + r1.file_id = "file-abc" + r1.filename = "doc.pdf" + r2.file_id = "file-abc" # same file + r2.filename = "doc.pdf" + + annotations = _build_file_citation_annotations([r1, r2], "text") + assert len(annotations) == 1 + + # --- End-to-end (mocked) --- + + @pytest.mark.asyncio + async def test_H11_emulated_full_flow_provider_calls_tool(self): + """Full flow: provider calls file_search function → search → follow-up → OpenAI output.""" + from litellm.responses.file_search.emulated_handler import ( + aresponses_with_emulated_file_search, + ) + + first_resp = self._make_mock_responses_api_response(include_function_call=True) + final_resp = self._make_mock_responses_api_response(text="Deep research enables multi-step queries.") + + search_result = MagicMock() + search_result.file_id = "file-xyz" + search_result.filename = "research.pdf" + search_result.score = 0.95 + search_result.content = [{"type": "text", "text": "deep research context..."}] + + mock_search_response = MagicMock() + mock_search_response.data = [search_result] + + with patch( + "litellm.responses.file_search.emulated_handler._call_aresponses", + new=AsyncMock(side_effect=[first_resp, final_resp]), + ), patch( + "litellm.vector_stores.main.asearch", + new=AsyncMock(return_value=mock_search_response), + ): + result = await aresponses_with_emulated_file_search( + input="What is deep research?", + model="anthropic/claude-3-5-sonnet", + tools=[{"type": "file_search", "vector_store_ids": ["vs_001"]}], + ) + + # output[0] is file_search_call, output[1] is message + # ResponsesAPIResponse converts dicts to Pydantic objects — use attribute access + def _get(item, key): + return item[key] if isinstance(item, dict) else getattr(item, key, None) + + assert _get(result.output[0], "type") == "file_search_call" + assert _get(result.output[0], "status") == "completed" + assert _get(result.output[1], "type") == "message" + content0 = _get(result.output[1], "content")[0] + assert "Deep research" in _get(content0, "text") + annotations = _get(content0, "annotations") + assert any(_get(a, "file_id") == "file-xyz" for a in annotations) + + @pytest.mark.asyncio + async def test_H12_emulated_flow_provider_answers_without_tool_call(self): + """If provider answers directly (no tool call), still return OpenAI format.""" + from litellm.responses.file_search.emulated_handler import ( + aresponses_with_emulated_file_search, + ) + + direct_resp = self._make_mock_responses_api_response(text="I already know the answer.") + + with patch( + "litellm.responses.file_search.emulated_handler._call_aresponses", + new=AsyncMock(return_value=direct_resp), + ): + result = await aresponses_with_emulated_file_search( + input="What is 2+2?", + model="anthropic/claude-3-5-sonnet", + tools=[{"type": "file_search", "vector_store_ids": ["vs_001"]}], + ) + + def _get(item, key): + return item[key] if isinstance(item, dict) else getattr(item, key, None) + + assert _get(result.output[0], "type") == "file_search_call" + assert _get(result.output[1], "type") == "message" + assert "I already know" in _get(_get(result.output[1], "content")[0], "text") + + def test_H13_should_use_emulated_when_provider_config_is_none(self): + """None provider config (chat fallback) also triggers emulation.""" + from litellm.responses.file_search.emulated_handler import ( + should_use_emulated_file_search, + ) + + tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}] + assert should_use_emulated_file_search(tools, None) is True