diff --git a/litellm/__init__.py b/litellm/__init__.py index 922a0afcede..09961174497 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -887,6 +887,7 @@ from .llms.triton.completion.transformation import TritonConfig from .llms.triton.completion.transformation import TritonGenerateConfig from .llms.triton.completion.transformation import TritonInferConfig from .llms.triton.embedding.transformation import TritonEmbeddingConfig +from .llms.huggingface.rerank.transformation import HuggingFaceRerankConfig from .llms.databricks.chat.transformation import DatabricksConfig from .llms.databricks.embed.transformation import DatabricksEmbeddingConfig from .llms.predibase.chat.transformation import PredibaseConfig diff --git a/litellm/llms/huggingface/rerank/handler.py b/litellm/llms/huggingface/rerank/handler.py new file mode 100644 index 00000000000..a8ae15c3dae --- /dev/null +++ b/litellm/llms/huggingface/rerank/handler.py @@ -0,0 +1,5 @@ +""" +HuggingFace Rerank - uses `llm_http_handler.py` to make httpx requests + +Request/Response transformation is handled in `transformation.py` +""" diff --git a/litellm/llms/huggingface/rerank/transformation.py b/litellm/llms/huggingface/rerank/transformation.py new file mode 100644 index 00000000000..c16b1c1af09 --- /dev/null +++ b/litellm/llms/huggingface/rerank/transformation.py @@ -0,0 +1,298 @@ +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, TypedDict +import uuid +import os + +import httpx +import litellm + +from litellm.types.rerank import ( + OptionalRerankParams, + RerankBilledUnits, + RerankResponse, + RerankResponseDocument, + RerankResponseMeta, + RerankResponseResult, + RerankTokens, +) + +from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.secret_managers.main import get_secret_str +from litellm.utils import token_counter +from ..common_utils import HuggingFaceError + + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + LoggingClass = LiteLLMLoggingObj +else: + LoggingClass = Any + + +class HuggingFaceRerankResponseItem(TypedDict): + """Type definition for HuggingFace rerank API response items.""" + index: int + score: float + text: Optional[str] # Optional, included when return_text=True + + +class HuggingFaceRerankResponse(TypedDict): + """Type definition for HuggingFace rerank API complete response.""" + # The response is a list of HuggingFaceRerankResponseItem + pass + + +# Type alias for the actual response structure +HuggingFaceRerankResponseList = List[HuggingFaceRerankResponseItem] + + +class HuggingFaceRerankConfig(BaseRerankConfig): + def get_api_base(self, model: str, api_base: Optional[str]) -> str: + if api_base is not None: + return api_base + elif os.getenv("HF_API_BASE") is not None: + return os.getenv("HF_API_BASE", "") + elif os.getenv("HUGGINGFACE_API_BASE") is not None: + return os.getenv("HUGGINGFACE_API_BASE", "") + else: + return "https://api-inference.huggingface.co" + + def get_complete_url(self, api_base: Optional[str], model: str) -> str: + """ + Get the complete URL for the API call, including the /rerank suffix if necessary. + """ + # Get base URL from api_base or default + base_url = self.get_api_base(model=model, api_base=api_base) + + # Remove trailing slashes and ensure we have the /rerank endpoint + base_url = base_url.rstrip("/") + if not base_url.endswith("/rerank"): + base_url = f"{base_url}/rerank" + + return base_url + + def get_supported_cohere_rerank_params(self, model: str) -> list: + return [ + "query", + "documents", + "top_n", + "rank_fields", + "return_documents", + "max_chunks_per_doc", + "max_tokens_per_doc", + ] + + def map_cohere_rerank_params( + self, + non_default_params: Optional[dict], + model: str, + drop_params: bool, + query: str, + documents: List[Union[str, Dict[str, Any]]], + custom_llm_provider: Optional[str] = None, + top_n: Optional[int] = None, + rank_fields: Optional[List[str]] = None, + return_documents: Optional[bool] = True, + max_chunks_per_doc: Optional[int] = None, + max_tokens_per_doc: Optional[int] = None, + ) -> OptionalRerankParams: + return OptionalRerankParams( + query=query, + documents=documents, + top_n=top_n, + rank_fields=rank_fields, + return_documents=return_documents, + max_chunks_per_doc=max_chunks_per_doc, + max_tokens_per_doc=max_tokens_per_doc, + ) + + def validate_environment( + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + # Get API credentials + api_key, api_base = self.get_api_credentials(api_key=api_key, api_base=api_base) + + default_headers = { + "accept": "application/json", + "content-type": "application/json", + } + + if api_key: + default_headers["Authorization"] = f"Bearer {api_key}" + + if "Authorization" in headers: + default_headers["Authorization"] = headers["Authorization"] + + return {**default_headers, **headers} + + def transform_rerank_request( + self, + model: str, + optional_rerank_params: OptionalRerankParams, + headers: dict, + ) -> dict: + if "query" not in optional_rerank_params: + raise ValueError("query is required for HuggingFace rerank") + if "documents" not in optional_rerank_params: + raise ValueError("documents is required for HuggingFace rerank") + + # Ensure return_text is a boolean value + # HuggingFace API expects return_text parameter, corresponding to our return_documents parameter + return_documents = optional_rerank_params.get("return_documents") + # Default to returning document content unless explicitly set to False + if return_documents is False: + return_text = False + else: + return_text = True + + request_body = { + "query": optional_rerank_params["query"], + "texts": optional_rerank_params["documents"], + "raw_scores": False, + "return_text": return_text, + "truncate": False, + "truncation_direction": "Right", + } + + if optional_rerank_params.get("top_n") is not None: + request_body["top_n"] = optional_rerank_params["top_n"] + + return request_body + + def transform_rerank_response( + self, + model: str, + raw_response: httpx.Response, + model_response: RerankResponse, + logging_obj: LoggingClass, + api_key: Optional[str] = None, + request_data: dict = {}, + optional_params: dict = {}, + litellm_params: dict = {}, + ) -> RerankResponse: + try: + raw_response_json: HuggingFaceRerankResponseList = raw_response.json() + except Exception: + raise HuggingFaceError( + message=getattr(raw_response, 'text', str(raw_response)), + status_code=getattr(raw_response, 'status_code', 500) + ) + + # Use standard litellm token counter for proper token estimation + try: + # Calculate tokens for the raw response JSON string + response_text = str(raw_response_json) + estimated_output_tokens = token_counter(model=model, text=response_text) + + # Calculate input tokens from query and documents + query = request_data.get("query", "") + documents = request_data.get("texts", []) + + # Convert documents to string if they're not already + documents_text = "" + for doc in documents: + if isinstance(doc, str): + documents_text += doc + " " + elif isinstance(doc, dict) and "text" in doc: + documents_text += doc["text"] + " " + + # Calculate input tokens using the same model + input_text = query + " " + documents_text + estimated_input_tokens = token_counter(model=model, text=input_text) + except Exception: + # Fallback to reasonable estimates if token counting fails + estimated_output_tokens = len(raw_response_json) * 10 if raw_response_json else 10 + estimated_input_tokens = len(input_text) * 4 if 'input_text' in locals() else 0 + + _billed_units = RerankBilledUnits(search_units=1) + _tokens = RerankTokens( + input_tokens=estimated_input_tokens, + output_tokens=estimated_output_tokens + ) + rerank_meta = RerankResponseMeta( + api_version={"version": "1.0"}, + billed_units=_billed_units, + tokens=_tokens + ) + + # Check if documents should be returned based on request parameters + should_return_documents = request_data.get("return_text", False) or request_data.get("return_documents", False) + original_documents = request_data.get("texts", []) + + results = [] + for item in raw_response_json: + # Extract required fields with defaults to handle None values + index = item.get("index") + score = item.get("score") + + # Skip items that don't have required fields + if index is None or score is None: + continue + + # Create RerankResponseResult with required fields + result = RerankResponseResult( + index=index, + relevance_score=score + ) + + # Add optional document field if needed + if should_return_documents: + text_content = item.get("text", "") + + # 1. First try to use text returned directly from API if available + if text_content: + result["document"] = RerankResponseDocument(text=text_content) + # 2. If no text in API response but original documents are available, use those + elif original_documents and 0 <= item.get("index", -1) < len(original_documents): + doc = original_documents[item.get("index")] + if isinstance(doc, str): + result["document"] = RerankResponseDocument(text=doc) + elif isinstance(doc, dict) and "text" in doc: + result["document"] = RerankResponseDocument(text=doc["text"]) + + results.append(result) + + return RerankResponse( + id=str(uuid.uuid4()), + results=results, + meta=rerank_meta, + ) + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BaseLLMException: + return HuggingFaceError(message=error_message, status_code=status_code) + + def get_api_credentials( + self, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> tuple[Optional[str], Optional[str]]: + """ + Get API key and base URL from multiple sources. + Returns tuple of (api_key, api_base). + + Parameters: + api_key: API key provided directly to this function, takes precedence over all other sources + api_base: API base provided directly to this function, takes precedence over all other sources + """ + # Get API key from multiple sources + final_api_key = ( + api_key or + litellm.huggingface_key or + get_secret_str("HUGGINGFACE_API_KEY") + ) + + # Get API base from multiple sources + final_api_base = ( + api_base or + litellm.api_base or + get_secret_str("HF_API_BASE") or + get_secret_str("HUGGINGFACE_API_BASE") + ) + + return final_api_key, final_api_base \ No newline at end of file diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 9307ce5a550..cc80d357d11 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -324,7 +324,31 @@ def rerank( # noqa: PLR0915 client=client, ) else: - raise ValueError(f"Unsupported provider: {_custom_llm_provider}") + # Generic handler for all providers that use base_llm_http_handler + # Provider-specific logic (API key validation, URL generation, etc.) + # is handled in the respective transformation configs + + # Check if the provider is actually supported + # If rerank_provider_config is a default CohereRerankConfig but the provider is not Cohere or litellm_proxy, + # it means the provider is not supported + if (isinstance(rerank_provider_config, litellm.CohereRerankConfig) or + isinstance(rerank_provider_config, litellm.CohereRerankV2Config)) and _custom_llm_provider != "cohere" and _custom_llm_provider != "litellm_proxy": + raise ValueError(f"Unsupported provider: {_custom_llm_provider}") + + response = base_llm_http_handler.rerank( + model=model, + custom_llm_provider=_custom_llm_provider, + provider_config=rerank_provider_config, + optional_rerank_params=optional_rerank_params, + logging_obj=litellm_logging_obj, + timeout=optional_params.timeout, + api_key=dynamic_api_key or optional_params.api_key, + api_base=dynamic_api_base or optional_params.api_base, + _is_async=_is_async, + headers=headers or litellm.headers or {}, + client=client, + model_response=model_response, + ) # Placeholder return return response diff --git a/litellm/utils.py b/litellm/utils.py index cc9d2fdd116..6f45e62a983 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6728,6 +6728,8 @@ class ProviderConfigManager: return litellm.InfinityRerankConfig() elif litellm.LlmProviders.JINA_AI == provider: return litellm.JinaAIRerankConfig() + elif litellm.LlmProviders.HUGGINGFACE == provider: + return litellm.HuggingFaceRerankConfig() return litellm.CohereRerankConfig() @staticmethod diff --git a/tests/test_litellm/llms/huggingface/__init__.py b/tests/test_litellm/llms/huggingface/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/huggingface/test_rerank.py b/tests/test_litellm/llms/huggingface/test_rerank.py new file mode 100644 index 00000000000..adbbaf26deb --- /dev/null +++ b/tests/test_litellm/llms/huggingface/test_rerank.py @@ -0,0 +1,438 @@ +""" +Tests for HuggingFace rerank functionality. +Based on the test patterns from other rerank providers and the current HuggingFace implementation. +""" +import asyncio +import json +from unittest.mock import patch, MagicMock, AsyncMock + +import pytest +import litellm + + +def assert_response_shape(response, custom_llm_provider): + """Helper function to validate response structure""" + assert hasattr(response, 'id') + assert hasattr(response, 'results') + assert hasattr(response, 'meta') + assert isinstance(response.results, list) + + for result in response.results: + assert "index" in result + assert "relevance_score" in result + assert isinstance(result["index"], int) + assert isinstance(result["relevance_score"], (int, float)) + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@patch('litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post') +@patch('litellm.llms.custom_httpx.http_handler.HTTPHandler.post') +def test_basic_rerank_huggingface(mock_sync_post, mock_async_post, sync_mode): + """Test basic HuggingFace rerank functionality.""" + # Mock response data that matches HuggingFace rerank API format + mock_response_data = [ + {"index": 0, "score": 0.9}, + {"index": 1, "score": 0.1} + ] + + def return_val(): + return mock_response_data + + api_key = "test_hf_api_key" + + if sync_mode: + # Create mock response object for sync + mock_response = MagicMock() + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_sync_post.return_value = mock_response + + response = litellm.rerank( + model="huggingface/BAAI/bge-reranker-base", + query="hello", + documents=["hello", "world"], + top_n=2, + api_key=api_key, + ) + mock_sync_post.assert_called_once() + else: + # Create mock response object for async + mock_response = AsyncMock() + + def return_val(): + return mock_response_data + + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_async_post.return_value = mock_response + + response = asyncio.run(litellm.arerank( + model="huggingface/BAAI/bge-reranker-base", + query="hello", + documents=["hello", "world"], + top_n=2, + api_key=api_key, + )) + mock_async_post.assert_called_once() + + assert response.results is not None + assert len(response.results) == 2 + assert response.results[0]["index"] == 0 + assert response.results[0]["relevance_score"] == 0.9 + + assert_response_shape(response, custom_llm_provider="huggingface") + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@patch('litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post') +@patch('litellm.llms.custom_httpx.http_handler.HTTPHandler.post') +def test_huggingface_rerank_custom_api_base(mock_sync_post, mock_async_post, sync_mode): + """Test HuggingFace rerank with custom API base.""" + mock_response_data = [ + {"index": 0, "score": 0.9}, + {"index": 1, "score": 0.1} + ] + + def return_val(): + return mock_response_data + + if sync_mode: + mock_response = MagicMock() + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_sync_post.return_value = mock_response + + response = litellm.rerank( + model="huggingface/BAAI/bge-reranker-base", + query="hello", + documents=["hello", "world"], + top_n=2, + api_base="https://my-custom-hf-endpoint.com", + api_key="test_api_key", + ) + + mock_sync_post.assert_called_once() + call_url = mock_sync_post.call_args.kwargs["url"] + assert "my-custom-hf-endpoint.com" in call_url + assert response.results is not None + assert len(response.results) == 2 + else: + mock_response = AsyncMock() + + def return_val(): + return mock_response_data + + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_async_post.return_value = mock_response + + response = asyncio.run(litellm.arerank( + model="huggingface/BAAI/bge-reranker-base", + query="hello", + documents=["hello", "world"], + top_n=2, + api_base="https://my-custom-hf-endpoint.com", + api_key="test_api_key", + )) + + mock_async_post.assert_called_once() + call_url = mock_async_post.call_args.kwargs["url"] + assert "my-custom-hf-endpoint.com" in call_url + assert response.results is not None + assert len(response.results) == 2 + + +@patch('litellm.llms.custom_httpx.http_handler.HTTPHandler.post') +def test_huggingface_rerank_with_env_vars(mock_post, monkeypatch): + """Test HuggingFace rerank with environment variable configuration.""" + monkeypatch.setenv("HUGGINGFACE_API_KEY", "env_test_key") + monkeypatch.setenv("HUGGINGFACE_API_BASE", "https://env-hf-endpoint.com") + + mock_response_data = [ + {"index": 0, "score": 0.9}, + {"index": 1, "score": 0.1} + ] + + def return_val(): + return mock_response_data + + mock_response = MagicMock() + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_post.return_value = mock_response + + response = litellm.rerank( + model="huggingface/BAAI/bge-reranker-base", + query="hello", + documents=["hello", "world"], + top_n=2, + ) + + mock_post.assert_called_once() + call_url = mock_post.call_args.kwargs["url"] + assert "env-hf-endpoint.com" in call_url + + headers = mock_post.call_args.kwargs.get("headers", {}) + assert "env_test_key" in str(headers.get("Authorization", "")) + + assert response.results is not None + assert len(response.results) == 2 + + +@patch('litellm.llms.custom_httpx.http_handler.HTTPHandler.post') +def test_huggingface_rerank_return_documents(mock_post): + """Test HuggingFace rerank with return_documents=True.""" + mock_response_data = [ + {"index": 0, "score": 0.9, "text": "hello"}, + {"index": 1, "score": 0.1, "text": "world"} + ] + + def return_val(): + return mock_response_data + + mock_response = MagicMock() + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_post.return_value = mock_response + + response = litellm.rerank( + model="huggingface/BAAI/bge-reranker-base", + query="hello", + documents=["hello", "world"], + top_n=2, + return_documents=True, + api_key="test_api_key", + ) + + mock_post.assert_called_once() + request_data = json.loads(mock_post.call_args.kwargs["data"]) + assert request_data.get("return_text") is True + + assert response.results is not None + assert len(response.results) == 2 + # Check that documents are included in response + for result in response.results: + if "document" in result: + assert "text" in result["document"] + + +@patch('litellm.llms.custom_httpx.http_handler.HTTPHandler.post') +def test_huggingface_rerank_error_handling(mock_post): + """Test HuggingFace rerank error handling.""" + def return_val(): + return {"error": "Unauthorized"} + + mock_response = MagicMock() + mock_response.status_code = 401 + mock_response.json = return_val + mock_response.text = "Unauthorized" + mock_post.return_value = mock_response + + with pytest.raises(Exception): + litellm.rerank( + model="huggingface/BAAI/bge-reranker-base", + query="hello", + documents=["hello", "world"], + top_n=2, + api_key="invalid_key", + ) + + +def test_huggingface_rerank_config(): + """Test HuggingFaceRerankConfig class functionality.""" + from litellm.llms.huggingface.rerank.transformation import HuggingFaceRerankConfig + + config = HuggingFaceRerankConfig() + + # Test complete URL generation + assert config.get_complete_url(None, "test") == "https://api-inference.huggingface.co/rerank" + + # Test custom API base + custom_url = config.get_complete_url("https://custom.huggingface.co", "test") + assert custom_url == "https://custom.huggingface.co/rerank" + + # Test supported parameters + supported_params = config.get_supported_cohere_rerank_params("test") + assert "query" in supported_params + assert "documents" in supported_params + assert "top_n" in supported_params + assert "return_documents" in supported_params + + # Test parameter mapping + params = config.map_cohere_rerank_params( + non_default_params={}, + model="test", + drop_params=False, + query="hello", + documents=["hello", "world"], + top_n=2, + return_documents=True, + ) + assert params["query"] == "hello" + assert params["documents"] == ["hello", "world"] + assert params["top_n"] == 2 + assert params["return_documents"] is True + + +def test_request_transformation(): + """Test request transformation logic.""" + from litellm.llms.huggingface.rerank.transformation import HuggingFaceRerankConfig + from litellm.types.rerank import OptionalRerankParams + + config = HuggingFaceRerankConfig() + + optional_params = OptionalRerankParams( + query="hello", + documents=["hello", "world"], + top_n=2, + return_documents=True + ) + + request_body = config.transform_rerank_request( + model="test", + optional_rerank_params=optional_params, + headers={} + ) + + assert request_body["query"] == "hello" + assert request_body["texts"] == ["hello", "world"] + assert request_body["top_n"] == 2 + assert request_body["return_text"] is True + assert request_body["raw_scores"] is False + assert request_body["truncate"] is False + assert request_body["truncation_direction"] == "Right" + + +def test_response_transformation(): + """Test response transformation logic.""" + from litellm.llms.huggingface.rerank.transformation import HuggingFaceRerankConfig + from litellm.types.rerank import RerankResponse + + config = HuggingFaceRerankConfig() + + # Mock HuggingFace response + hf_response_data = [ + {"index": 0, "score": 0.9, "text": "hello"}, + {"index": 1, "score": 0.1, "text": "world"} + ] + + def return_val(): + return hf_response_data + + # Create mock httpx response + mock_response = MagicMock() + mock_response.json = return_val + + model_response = RerankResponse() + + transformed_response = config.transform_rerank_response( + model="test", + raw_response=mock_response, + model_response=model_response, + logging_obj=None, + request_data={"return_text": True} + ) + + assert transformed_response.results is not None + assert len(transformed_response.results) == 2 + assert transformed_response.results[0]["index"] == 0 + assert transformed_response.results[0]["relevance_score"] == 0.9 + assert transformed_response.results[1]["index"] == 1 + assert transformed_response.results[1]["relevance_score"] == 0.1 + + # Check documents are included when return_text is True + for result in transformed_response.results: + if "document" in result: + assert "text" in result["document"] + + +def test_validate_environment(): + """Test environment validation logic.""" + from litellm.llms.huggingface.rerank.transformation import HuggingFaceRerankConfig + + config = HuggingFaceRerankConfig() + + # Test with API key + headers = config.validate_environment( + headers={}, + model="test", + api_key="test_key" + ) + + assert "Authorization" in headers + assert "Bearer test_key" in headers["Authorization"] + assert headers["accept"] == "application/json" + assert headers["content-type"] == "application/json" + + # Test headers override + custom_headers = {"custom": "header"} + headers = config.validate_environment( + headers=custom_headers, + model="test", + api_key="test_key" + ) + + assert "custom" in headers + assert headers["custom"] == "header" + + +@patch('litellm.llms.custom_httpx.http_handler.HTTPHandler.post') +def test_huggingface_rerank_request_payload(mock_post): + """Test that the request payload is correctly formatted for HuggingFace API.""" + mock_response_data = [ + {"index": 0, "score": 0.9}, + {"index": 1, "score": 0.1} + ] + + def return_val(): + return mock_response_data + + mock_response = MagicMock() + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_post.return_value = mock_response + + response = litellm.rerank( + model="huggingface/BAAI/bge-reranker-base", + query="hello", + documents=["hello", "world"], + top_n=2, + return_documents=True, + api_key="test_api_key", + ) + + mock_post.assert_called_once() + + # Verify URL + call_url = mock_post.call_args.kwargs["url"] + assert call_url == "https://api-inference.huggingface.co/rerank" + + # Verify headers + headers = mock_post.call_args.kwargs["headers"] + assert "Bearer test_api_key" in headers["Authorization"] + assert headers["content-type"] == "application/json" + + # Verify request body + request_data = json.loads(mock_post.call_args.kwargs["data"]) + expected_request = { + "query": "hello", + "texts": ["hello", "world"], + "raw_scores": False, + "return_text": True, + "truncate": False, + "truncation_direction": "Right", + "top_n": 2 + } + + for key, value in expected_request.items(): + assert request_data[key] == value + + assert response.results is not None + assert len(response.results) == 2