diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 28a0097c30d..e4fe26cd564 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -1,6 +1,6 @@ # What is this? ## Helper utilities -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any, Iterable, List, Optional, Union import httpx @@ -70,6 +70,15 @@ def remove_index_from_tool_calls( return +def remove_items_at_indices(items: Optional[List[Any]], indices: Iterable[int]) -> None: + """Remove items from a list in-place by index""" + if items is None: + return + for index in sorted(set(indices), reverse=True): + if 0 <= index < len(items): + items.pop(index) + + def add_missing_spend_metadata_to_litellm_metadata( litellm_metadata: dict, metadata: dict ) -> dict: diff --git a/litellm/main.py b/litellm/main.py index dffa5d40cd5..6cb14f12d79 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -436,6 +436,15 @@ async def acompletion( tools=tools, prompt_label=kwargs.get("prompt_label", None), ) + ######################################################### + # if the chat completion logging hook removed all tools, + # set tools to None + # eg. in certain cases when users send vector stores as tools + # we don't want the tools to go to the upstream llm + # relevant issue: https://github.com/BerriAI/litellm/issues/11404 + ######################################################### + if tools is not None and len(tools) == 0: + tools = None ######################################################### ######################################################### @@ -2769,9 +2778,9 @@ def completion( # type: ignore # noqa: PLR0915 "aws_region_name" not in optional_params or optional_params["aws_region_name"] is None ): - optional_params[ - "aws_region_name" - ] = aws_bedrock_client.meta.region_name + optional_params["aws_region_name"] = ( + aws_bedrock_client.meta.region_name + ) bedrock_route = BedrockModelInfo.get_bedrock_route(model) if bedrock_route == "converse": @@ -4545,9 +4554,9 @@ def adapter_completion( new_kwargs = translation_obj.translate_completion_input_params(kwargs=kwargs) response: Union[ModelResponse, CustomStreamWrapper] = completion(**new_kwargs) # type: ignore - translated_response: Optional[ - Union[BaseModel, AdapterCompletionStreamWrapper] - ] = None + translated_response: Optional[Union[BaseModel, AdapterCompletionStreamWrapper]] = ( + None + ) if isinstance(response, ModelResponse): translated_response = translation_obj.translate_completion_output_params( response=response @@ -5505,9 +5514,9 @@ def stream_chunk_builder( # noqa: PLR0915 ] if len(content_chunks) > 0: - response["choices"][0]["message"][ - "content" - ] = processor.get_combined_content(content_chunks) + response["choices"][0]["message"]["content"] = ( + processor.get_combined_content(content_chunks) + ) reasoning_chunks = [ chunk @@ -5518,9 +5527,9 @@ def stream_chunk_builder( # noqa: PLR0915 ] if len(reasoning_chunks) > 0: - response["choices"][0]["message"][ - "reasoning_content" - ] = processor.get_combined_reasoning_content(reasoning_chunks) + response["choices"][0]["message"]["reasoning_content"] = ( + processor.get_combined_reasoning_content(reasoning_chunks) + ) audio_chunks = [ chunk diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index a045406d78a..a13543ae061 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -4,6 +4,7 @@ from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, Dict, List, Optional from litellm._logging import verbose_logger +from litellm.litellm_core_utils.core_helpers import remove_items_at_indices from litellm.types.vector_stores import ( LiteLLM_ManagedVectorStore, LiteLLM_ManagedVectorStoreListResponse, @@ -55,9 +56,49 @@ class VectorStoreRegistry: vector_store_ids = non_default_params.pop("vector_store_ids", None) or [] # 2. check if vector_store_ids is provided as a tool in the request - vector_store_ids = self._get_vector_store_ids_from_tool_calls( - tools=tools, vector_store_ids=vector_store_ids + vector_store_ids = self.get_and_pop_recognised_vector_store_tools( + tools=tools, + vector_store_ids=vector_store_ids, ) + + return vector_store_ids + + def get_and_pop_recognised_vector_store_tools( + self, tools: Optional[List[Dict]] = None, vector_store_ids: List[str] = [] + ) -> List[str]: + """ + Returns and pops the vector store ids from the tool calls + + It only pops the recognised vector store tools from the tools list. + + Args: + tools: The tools to pop the vector store ids from + vector_store_ids: The list of vector store IDs the user provided + + Returns: + The vector store ids that were popped + """ + if tools: + tools_to_remove: List[int] = [] + for i, tool in enumerate(tools): + tool_vector_store_ids: List[str] = tool.get("vector_store_ids", []) + if len(tool_vector_store_ids) == 0: + continue + # remove the tool if all vector_store_ids are recognised in the registry + recognised = all( + any(vs.get("vector_store_id") == vs_id for vs in self.vector_stores) + for vs_id in tool_vector_store_ids + ) + if recognised: + tools_to_remove.append(i) + vector_store_ids.extend(tool_vector_store_ids) + + # remove recognised tools from the original list + remove_items_at_indices( + items=tools, + indices=tools_to_remove, + ) + return vector_store_ids def get_vector_store_to_run( diff --git a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py index a8afe9ebbaa..532a5c11c52 100644 --- a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py +++ b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py @@ -210,6 +210,7 @@ async def test_openai_with_vector_store_ids_in_tool_call_mock_openai(setup_vecto # Verify the API was called mock_client.assert_called_once() request_body = mock_client.call_args.kwargs + print("request body:", json.dumps(request_body, indent=4, default=str)) # Verify the request contains messages with knowledge base context assert "messages" in request_body @@ -226,6 +227,48 @@ async def test_openai_with_vector_store_ids_in_tool_call_mock_openai(setup_vecto assert messages[1]["role"] == "user" assert BedrockVectorStore.CONTENT_PREFIX_STRING in messages[1]["content"] + # assert that the tool call was not sent to the upstream llm API if it's a litellm vector store + assert "tools" not in request_body + + +@pytest.mark.asyncio +async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_registry): + """Ensure unrecognized vector store tools are forwarded to the provider""" + litellm.callbacks = [BedrockVectorStore(aws_region_name="us-west-2")] + from openai import AsyncOpenAI + + client = AsyncOpenAI(api_key="fake-api-key") + + with patch.object( + client.chat.completions.with_raw_response, "create" + ) as mock_client: + try: + await litellm.acompletion( + model="gpt-4", + messages=[{"role": "user", "content": "what is litellm?"}], + tools=[ + {"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]}, + {"type": "file_search", "vector_store_ids": ["unknownVS"]}, + ], + client=client, + ) + except Exception as e: + print(f"Error: {e}") + + mock_client.assert_called_once() + request_body = mock_client.call_args.kwargs + + assert "messages" in request_body + messages = request_body["messages"] + assert len(messages) >= 2 + assert messages[1]["role"] == "user" + assert BedrockVectorStore.CONTENT_PREFIX_STRING in messages[1]["content"] + + assert "tools" in request_body + tools = request_body["tools"] + assert len(tools) == 1 + assert tools[0]["vector_store_ids"] == ["unknownVS"] + @pytest.mark.asyncio async def test_logging_with_knowledge_base_hook(setup_vector_store_registry):