diff --git a/litellm/proxy/client/models.py b/litellm/proxy/client/models.py index 4b16087e15b..603597cc117 100644 --- a/litellm/proxy/client/models.py +++ b/litellm/proxy/client/models.py @@ -32,7 +32,7 @@ class ModelsManagementClient: headers["Authorization"] = f"Bearer {self._api_key}" return headers - def list(self, return_request: bool = False) -> list[dict[str, Any]] | requests.Request: + def list(self, return_request: bool = False) -> builtins.list[dict[str, Any]] | requests.Request: """ Get the list of models supported by the server. diff --git a/litellm/proxy/client/teams.py b/litellm/proxy/client/teams.py index 105060e5ca9..54a6e869fef 100644 --- a/litellm/proxy/client/teams.py +++ b/litellm/proxy/client/teams.py @@ -40,7 +40,7 @@ class TeamsManagementClient: self, user_id: str | None = None, organization_id: str | None = None, - ) -> list[dict[str, Any]]: + ) -> builtins.list[dict[str, Any]]: """ List teams that the user belongs to. diff --git a/litellm/vector_stores/main.py b/litellm/vector_stores/main.py index cd576755f5f..662b4981304 100644 --- a/litellm/vector_stores/main.py +++ b/litellm/vector_stores/main.py @@ -39,7 +39,7 @@ base_llm_http_handler = BaseLLMHTTPHandler() def mock_vector_store_search_response( - mock_results: list[VectorStoreSearchResult] | None = None, + mock_results: builtins.list[VectorStoreSearchResult] | None = None, ): """Mock response for vector store search""" if mock_results is None: @@ -93,7 +93,7 @@ def mock_vector_store_create_response( @client async def acreate( name: str | None = None, - file_ids: list[str] | None = None, + file_ids: builtins.list[str] | None = None, expires_after: dict | None = None, chunking_strategy: dict | None = None, metadata: dict[str, str] | None = None, @@ -157,7 +157,7 @@ async def acreate( @client def create( name: str | None = None, - file_ids: list[str] | None = None, + file_ids: builtins.list[str] | None = None, expires_after: dict | None = None, chunking_strategy: dict | None = None, metadata: dict[str, str] | None = None, @@ -270,7 +270,7 @@ def create( @client async def asearch( vector_store_id: str, - query: str | list[str], + query: str | builtins.list[str], filters: dict | None = None, max_num_results: int | None = None, ranking_options: dict | None = None, @@ -339,7 +339,7 @@ async def asearch( @client def search( vector_store_id: str, - query: str | list[str], + query: str | builtins.list[str], filters: dict | None = None, max_num_results: int | None = None, ranking_options: dict | None = None, diff --git a/tests/test_litellm/vector_stores/test_main.py b/tests/test_litellm/vector_stores/test_main.py index d01e696906a..21f4165bd26 100644 --- a/tests/test_litellm/vector_stores/test_main.py +++ b/tests/test_litellm/vector_stores/test_main.py @@ -9,6 +9,8 @@ serialization trap). from unittest.mock import MagicMock, patch +import pytest + import litellm.vector_stores.main as vector_stores_main from litellm.vector_stores.main import search @@ -19,7 +21,8 @@ MOCK_SEARCH_RESPONSE = { } -def test_search_threads_router_to_handler(): +@pytest.mark.parametrize("query", ["q", ["q", "another question"]]) +def test_search_threads_router_to_handler(query: str | list[str]): """search() must pass its router param through to the HTTP handler""" mock_router = MagicMock() logger = MagicMock() @@ -37,7 +40,7 @@ def test_search_threads_router_to_handler(): ): response = search( vector_store_id="bkt:idx", - query="q", + query=query, custom_llm_provider="s3_vectors", router=mock_router, litellm_logging_obj=logger, @@ -46,6 +49,7 @@ def test_search_threads_router_to_handler(): assert response == MOCK_SEARCH_RESPONSE mock_handler.assert_called_once() assert mock_handler.call_args.kwargs["router"] is mock_router + assert mock_handler.call_args.kwargs["query"] == query def test_search_router_not_in_litellm_params():