feat: add HuggingFace rerank provider support (#11438)

++ feat: add HuggingFace rerank provider support

feat: add HuggingFace rerank provider support

feat: add HuggingFace rerank provider support

feat: add HuggingFace rerank provider support

feat: add HuggingFace rerank provider support
This commit is contained in:
cainiaoit 2025-06-06 14:23:01 +08:00 • committed by GitHub
parent 5d516aace1
commit be12416863
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 769 additions and 1 deletions

View file

@ -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

View file

@ -0,0 +1,5 @@
"""
HuggingFace Rerank - uses `llm_http_handler.py` to make httpx requests
Request/Response transformation is handled in `transformation.py`
"""

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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