diff --git a/litellm/llms/deepinfra/rerank/transformation.py b/litellm/llms/deepinfra/rerank/transformation.py index f78e7881585..b64e025e0d3 100644 --- a/litellm/llms/deepinfra/rerank/transformation.py +++ b/litellm/llms/deepinfra/rerank/transformation.py @@ -108,7 +108,9 @@ class DeepinfraRerankConfig(BaseRerankConfig): # Start with the basic parameters optional_rerank_params = {} if query: - optional_rerank_params["queries"] = [query] # Single-element array; Deepinfra broadcasts it across all documents + optional_rerank_params["queries"] = [ + query + ] # Single-element array; Deepinfra broadcasts it across all documents if non_default_params is not None: for k, v in non_default_params.items(): diff --git a/tests/test_litellm/llms/deepinfra/test_deepinfra_rerank_transformation.py b/tests/test_litellm/llms/deepinfra/test_deepinfra_rerank_transformation.py index c69ea599c47..d0177528e6c 100644 --- a/tests/test_litellm/llms/deepinfra/test_deepinfra_rerank_transformation.py +++ b/tests/test_litellm/llms/deepinfra/test_deepinfra_rerank_transformation.py @@ -42,7 +42,6 @@ class TestDeepinfraRerankTransform: with pytest.raises(ValueError, match="Deepinfra API Base is required"): self.config.get_complete_url(None, model) - def test_map_cohere_rerank_params_basic(self): """Test basic parameter mapping for DeepInfra rerank.""" params = self.config.map_cohere_rerank_params( @@ -238,9 +237,7 @@ class TestDeepinfraRerankTransform: query="query1", documents=documents, ) - assert params["queries"] == [ - "query1" - ], f"Failed for {num_docs} documents" + assert params["queries"] == ["query1"], f"Failed for {num_docs} documents" def test_get_error_class_basic(self): """Test error class generation for basic error."""