This commit is contained in:
IvanShang 2026-08-26 15:35:33 +08:00 committed by GitHub
commit 1c3a61e1df
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 351 additions and 7 deletions

View file

@ -159,6 +159,7 @@ LLM_CONFIG_NAMES: Final = (
"FireworksAIRerankConfig",
"VoyageRerankConfig",
"IBMWatsonXRerankConfig",
"XinferenceRerankConfig",
"ClarifaiConfig",
"AI21ChatConfig",
"LlamaAPIConfig",
@ -700,6 +701,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
".llms.watsonx.rerank.transformation",
"IBMWatsonXRerankConfig",
),
"XinferenceRerankConfig": (
".llms.xinference.rerank.transformation",
"XinferenceRerankConfig",
),
"ClarifaiConfig": (".llms.clarifai.chat.transformation", "ClarifaiConfig"),
"AI21ChatConfig": (".llms.ai21.chat.transformation", "AI21ChatConfig"),
"LlamaAPIConfig": (".llms.meta_llama.chat.transformation", "LlamaAPIConfig"),

View file

@ -0,0 +1 @@

View file

@ -0,0 +1,179 @@
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final
import httpx
from pydantic import BaseModel, ConfigDict, TypeAdapter
from litellm._uuid import uuid
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.rerank import (
RerankBilledUnits,
RerankResponse,
RerankResponseDocument,
RerankResponseMeta,
RerankResponseResult,
RerankTokens,
)
DEFAULT_XINFERENCE_API_BASE: Final = "http://127.0.0.1:9997/v1"
class _XinferenceRerankResult(BaseModel):
model_config = ConfigDict(frozen=True)
index: int
relevance_score: float
document: str | None = None
class _XinferenceRerankResponse(BaseModel):
model_config = ConfigDict(frozen=True)
id: str | None = None
results: tuple[_XinferenceRerankResult, ...]
_XINFERENCE_RERANK_RESPONSE_ADAPTER: Final = TypeAdapter(_XinferenceRerankResponse)
class _RerankPayload(dict[str, object]):
pass
class _SupportedRerankParams(list[str]):
pass
class _RerankResults(list[RerankResponseResult]):
pass
class XinferenceRerankConfig(BaseRerankConfig):
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: Mapping[str, object] | None = None,
) -> str:
resolved_api_base: Final = api_base or get_secret_str("XINFERENCE_API_BASE") or DEFAULT_XINFERENCE_API_BASE
cleaned_api_base: Final = resolved_api_base.rstrip("/")
if cleaned_api_base.endswith("/rerank"):
return cleaned_api_base
return f"{cleaned_api_base}/rerank"
def validate_environment(
self,
headers: Mapping[str, object],
model: str,
api_key: str | None = None,
optional_params: Mapping[str, object] | None = None,
) -> _RerankPayload:
resolved_api_key: Final = api_key or get_secret_str("XINFERENCE_API_KEY") or "stub-xinference-key"
return _RerankPayload(
MappingProxyType(
{
"Authorization": f"Bearer {resolved_api_key}",
"accept": "application/json",
"content-type": "application/json",
**headers,
}
)
)
def get_supported_cohere_rerank_params(self, model: str) -> _SupportedRerankParams:
return _SupportedRerankParams(("query", "documents", "top_n"))
def map_cohere_rerank_params(
self,
non_default_params: Mapping[str, object],
model: str,
drop_params: bool,
query: str,
documents: Sequence[str | Mapping[str, object]],
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: Sequence[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
instruction: str | None = None,
) -> _RerankPayload:
if top_n is not None:
return _RerankPayload(MappingProxyType({"query": query, "documents": documents, "top_n": top_n}))
return _RerankPayload(MappingProxyType({"query": query, "documents": documents}))
def transform_rerank_request(
self,
model: str,
optional_rerank_params: Mapping[str, object],
headers: Mapping[str, object],
litellm_params: Mapping[str, object] | None = None,
) -> _RerankPayload:
if "query" not in optional_rerank_params:
raise ValueError("query is required for Xinference rerank")
if "documents" not in optional_rerank_params:
raise ValueError("documents is required for Xinference rerank")
if optional_rerank_params.get("top_n") is not None:
return _RerankPayload(
MappingProxyType(
{
"model": model,
"query": optional_rerank_params["query"],
"documents": optional_rerank_params["documents"],
"top_n": optional_rerank_params["top_n"],
}
)
)
return _RerankPayload(
MappingProxyType(
{
"model": model,
"query": optional_rerank_params["query"],
"documents": optional_rerank_params["documents"],
}
)
)
def transform_rerank_response(
self,
model: str,
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: str | None = None,
request_data: Mapping[str, object] | None = None,
optional_params: Mapping[str, object] | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> RerankResponse:
try:
response_json: Final = _XINFERENCE_RERANK_RESPONSE_ADAPTER.validate_python(raw_response.json())
except ValueError:
raise ValueError(f"Error parsing Xinference rerank response: {raw_response.text}")
transformed_results: Final = tuple(
RerankResponseResult(
index=result.index,
relevance_score=result.relevance_score,
document=RerankResponseDocument(text=result.document),
)
if result.document is not None
else RerankResponseResult(
index=result.index,
relevance_score=result.relevance_score,
)
for result in response_json.results
)
meta: Final = RerankResponseMeta(
billed_units=RerankBilledUnits(total_tokens=0),
tokens=RerankTokens(input_tokens=0),
)
return RerankResponse(
id=response_json.id or str(uuid.uuid4()),
results=_RerankResults(transformed_results),
meta=meta,
)

View file

@ -32,7 +32,7 @@ async def arerank(
query: str,
documents: list[str | dict[str, Any]],
custom_llm_provider: (
Literal["cohere", "together_ai", "deepinfra", "fireworks_ai", "voyage", "watsonx"] | None
Literal["cohere", "together_ai", "deepinfra", "fireworks_ai", "voyage", "watsonx", "xinference"] | None
) = None,
top_n: int | None = None,
rank_fields: list[str] | None = None,
@ -90,6 +90,7 @@ def rerank(
"fireworks_ai",
"voyage",
"watsonx",
"xinference",
]
| None
) = None,
@ -485,6 +486,37 @@ def rerank(
if credentials.get("token") is not None:
optional_rerank_params["token"] = credentials["token"]
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=api_key,
api_base=api_base,
_is_async=_is_async,
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.XINFERENCE:
api_key = (
dynamic_api_key
or optional_params.api_key
or litellm.api_key
or get_secret_str("XINFERENCE_API_KEY")
or "stub-xinference-key"
)
api_base = (
dynamic_api_base
or optional_params.api_base
or litellm.api_base
or get_secret_str("XINFERENCE_API_BASE")
or "http://127.0.0.1:9997/v1"
)
response = base_llm_http_handler.rerank(
model=model,
custom_llm_provider=_custom_llm_provider,

View file

@ -8300,6 +8300,15 @@ class ProviderConfigManager:
return litellm.PerplexityEmbeddingConfig()
return None
@staticmethod
def _get_dashscope_or_xinference_rerank_config(provider: LlmProviders) -> BaseRerankConfig:
if litellm.LlmProviders.XINFERENCE == provider:
return litellm.XinferenceRerankConfig()
from litellm.llms.dashscope.rerank.transformation import DashScopeRerankConfig
return DashScopeRerankConfig()
@staticmethod
def get_provider_rerank_config(
model: str,
@ -8338,12 +8347,8 @@ class ProviderConfigManager:
return litellm.VoyageRerankConfig()
elif litellm.LlmProviders.WATSONX == provider:
return litellm.IBMWatsonXRerankConfig()
elif litellm.LlmProviders.DASHSCOPE == provider:
from litellm.llms.dashscope.rerank.transformation import (
DashScopeRerankConfig,
)
return DashScopeRerankConfig()
elif provider in (litellm.LlmProviders.XINFERENCE, litellm.LlmProviders.DASHSCOPE):
return ProviderConfigManager._get_dashscope_or_xinference_rerank_config(provider)
return litellm.CohereRerankConfig()
@staticmethod

View file

@ -0,0 +1,122 @@
import asyncio
import json
import os
import sys
from unittest.mock import MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../../../../../.."))
import litellm
from litellm.llms.xinference.rerank.transformation import (
DEFAULT_XINFERENCE_API_BASE,
XinferenceRerankConfig,
)
def test_xinference_rerank_defaults_and_auth(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("XINFERENCE_API_BASE", raising=False)
monkeypatch.delenv("XINFERENCE_API_KEY", raising=False)
config = XinferenceRerankConfig()
assert config.get_complete_url(api_base=None, model="bge-reranker") == f"{DEFAULT_XINFERENCE_API_BASE}/rerank"
monkeypatch.setenv("XINFERENCE_API_BASE", "http://env-xinference.test/v1")
assert config.get_complete_url(api_base=None, model="bge-reranker") == "http://env-xinference.test/v1/rerank"
no_auth_headers = config.validate_environment(headers={}, model="bge-reranker")
assert no_auth_headers["Authorization"] == "Bearer stub-xinference-key"
monkeypatch.setenv("XINFERENCE_API_KEY", "env-key")
env_auth_headers = config.validate_environment(headers={}, model="bge-reranker")
assert env_auth_headers["Authorization"] == "Bearer env-key"
caller_auth_headers = config.validate_environment(
headers={"Authorization": "Bearer caller-token"},
model="bge-reranker",
)
assert caller_auth_headers["Authorization"] == "Bearer caller-token"
@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_xinference_rerank_uses_base_handler(
mock_sync_post: MagicMock,
mock_async_post: MagicMock,
sync_mode: bool,
) -> None:
response_data = {
"results": [
{"index": 1, "relevance_score": 0.92, "document": "Xinference supports rerank."},
{"index": 0, "relevance_score": 0.24, "document": "An unrelated document."},
]
}
api_base = "http://xinference.example.test/v1"
request_headers = {"Authorization": "Bearer caller-token", "x-request-id": "req-123"}
if sync_mode:
mock_response = MagicMock()
mock_response.json.return_value = response_data
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(response_data)
mock_sync_post.return_value = mock_response
response = litellm.rerank(
model="xinference/bge-reranker-large",
query="Does Xinference support rerank?",
documents=["An unrelated document.", "Xinference supports rerank."],
top_n=2,
api_base=api_base,
headers=request_headers,
)
mock_sync_post.assert_called_once()
call_kwargs = mock_sync_post.call_args.kwargs
else:
mock_response = MagicMock()
mock_response.json.return_value = response_data
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(response_data)
mock_async_post.return_value = mock_response
response = asyncio.run(
litellm.arerank(
model="xinference/bge-reranker-large",
query="Does Xinference support rerank?",
documents=["An unrelated document.", "Xinference supports rerank."],
top_n=2,
api_base=api_base,
headers=request_headers,
)
)
mock_async_post.assert_called_once()
call_kwargs = mock_async_post.call_args.kwargs
assert call_kwargs["url"] == "http://xinference.example.test/v1/rerank"
assert call_kwargs["headers"]["Authorization"] == "Bearer caller-token"
assert call_kwargs["headers"]["x-request-id"] == "req-123"
request_body = json.loads(call_kwargs["data"])
assert request_body == {
"model": "bge-reranker-large",
"query": "Does Xinference support rerank?",
"documents": ["An unrelated document.", "Xinference supports rerank."],
"top_n": 2,
}
assert response.results == [
{
"index": 1,
"relevance_score": 0.92,
"document": {"text": "Xinference supports rerank."},
},
{
"index": 0,
"relevance_score": 0.24,
"document": {"text": "An unrelated document."},
},
]