From 8aff5f2a69c22c3ba815dbfbac9b753fb93324f6 Mon Sep 17 00:00:00 2001 From: qdivan <77005282+qdivan@users.noreply.github.com> Date: Mon, 17 Aug 2026 13:34:38 +0800 Subject: [PATCH 1/2] feat(rerank): add Xinference support --- litellm/_lazy_imports_registry.py | 5 + litellm/llms/xinference/rerank/__init__.py | 1 + .../llms/xinference/rerank/transformation.py | 155 ++++++++++++++++++ litellm/rerank_api/main.py | 34 +++- litellm/utils.py | 2 + .../test_xinference_rerank_transformation.py | 122 ++++++++++++++ 6 files changed, 318 insertions(+), 1 deletion(-) create mode 100644 litellm/llms/xinference/rerank/__init__.py create mode 100644 litellm/llms/xinference/rerank/transformation.py create mode 100644 tests/test_litellm/llms/xinference/rerank/test_xinference_rerank_transformation.py diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 89c72acc06d..f7fa42f127e 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -159,6 +159,7 @@ LLM_CONFIG_NAMES: Final = ( "FireworksAIRerankConfig", "VoyageRerankConfig", "IBMWatsonXRerankConfig", + "XinferenceRerankConfig", "ClarifaiConfig", "AI21ChatConfig", "LlamaAPIConfig", @@ -698,6 +699,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"), diff --git a/litellm/llms/xinference/rerank/__init__.py b/litellm/llms/xinference/rerank/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/litellm/llms/xinference/rerank/__init__.py @@ -0,0 +1 @@ + diff --git a/litellm/llms/xinference/rerank/transformation.py b/litellm/llms/xinference/rerank/transformation.py new file mode 100644 index 00000000000..93650115d9f --- /dev/null +++ b/litellm/llms/xinference/rerank/transformation.py @@ -0,0 +1,155 @@ +from collections.abc import Mapping +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 ( + OptionalRerankParams, + 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 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, + ) -> dict[str, object]: + resolved_api_key: Final = api_key or get_secret_str("XINFERENCE_API_KEY") or "stub-xinference-key" + default_headers: Final = { + "Authorization": f"Bearer {resolved_api_key}", + "accept": "application/json", + "content-type": "application/json", + } + return {**default_headers, **headers} + + def get_supported_cohere_rerank_params(self, model: str) -> list[str]: + return ["query", "documents", "top_n"] + + def map_cohere_rerank_params( + self, + non_default_params: Mapping[str, object], + model: str, + drop_params: bool, + query: str, + documents: list[str | dict[str, object]], + custom_llm_provider: str | None = None, + top_n: int | None = None, + rank_fields: list[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, + ) -> dict[str, object]: + params: Final[OptionalRerankParams] = OptionalRerankParams( + query=query, + documents=documents, + ) + if top_n is not None: + params["top_n"] = top_n + return dict(params) + + def transform_rerank_request( + self, + model: str, + optional_rerank_params: Mapping[str, object], + headers: Mapping[str, object], + litellm_params: Mapping[str, object] | None = None, + ) -> dict[str, object]: + 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") + + request: Final[dict[str, object]] = { + "model": model, + "query": optional_rerank_params["query"], + "documents": optional_rerank_params["documents"], + } + if optional_rerank_params.get("top_n") is not None: + request["top_n"] = optional_rerank_params["top_n"] + return request + + 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=list(transformed_results), + meta=meta, + ) diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 15a6f18a6bb..7420e0a5a81 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -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, @@ -482,6 +483,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: Final = ( + 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: Final = ( + 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, diff --git a/litellm/utils.py b/litellm/utils.py index d91d3092624..b66c08c6964 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8199,6 +8199,8 @@ class ProviderConfigManager: return litellm.VoyageRerankConfig() elif litellm.LlmProviders.WATSONX == provider: return litellm.IBMWatsonXRerankConfig() + elif litellm.LlmProviders.XINFERENCE == provider: + return litellm.XinferenceRerankConfig() elif litellm.LlmProviders.DASHSCOPE == provider: from litellm.llms.dashscope.rerank.transformation import ( DashScopeRerankConfig, diff --git a/tests/test_litellm/llms/xinference/rerank/test_xinference_rerank_transformation.py b/tests/test_litellm/llms/xinference/rerank/test_xinference_rerank_transformation.py new file mode 100644 index 00000000000..b9f1f8a7ce1 --- /dev/null +++ b/tests/test_litellm/llms/xinference/rerank/test_xinference_rerank_transformation.py @@ -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."}, + }, + ] From 15e10d35f4f528b055adeaa0d9a68b8fbc3e0204 Mon Sep 17 00:00:00 2001 From: qdivan <77005282+qdivan@users.noreply.github.com> Date: Mon, 17 Aug 2026 14:33:00 +0800 Subject: [PATCH 2/2] fix(rerank): resolve Xinference lint violations --- .../llms/xinference/rerank/transformation.py | 82 ++++++++++++------- litellm/rerank_api/main.py | 4 +- litellm/utils.py | 19 +++-- 3 files changed, 66 insertions(+), 39 deletions(-) diff --git a/litellm/llms/xinference/rerank/transformation.py b/litellm/llms/xinference/rerank/transformation.py index 93650115d9f..016f10b46cb 100644 --- a/litellm/llms/xinference/rerank/transformation.py +++ b/litellm/llms/xinference/rerank/transformation.py @@ -1,4 +1,5 @@ -from collections.abc import Mapping +from collections.abc import Mapping, Sequence +from types import MappingProxyType from typing import Final import httpx @@ -9,7 +10,6 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig from litellm.secret_managers.main import get_secret_str from litellm.types.rerank import ( - OptionalRerankParams, RerankBilledUnits, RerankResponse, RerankResponseDocument, @@ -39,6 +39,18 @@ class _XinferenceRerankResponse(BaseModel): _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, @@ -58,17 +70,21 @@ class XinferenceRerankConfig(BaseRerankConfig): model: str, api_key: str | None = None, optional_params: Mapping[str, object] | None = None, - ) -> dict[str, object]: + ) -> _RerankPayload: resolved_api_key: Final = api_key or get_secret_str("XINFERENCE_API_KEY") or "stub-xinference-key" - default_headers: Final = { - "Authorization": f"Bearer {resolved_api_key}", - "accept": "application/json", - "content-type": "application/json", - } - return {**default_headers, **headers} + 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) -> list[str]: - return ["query", "documents", "top_n"] + def get_supported_cohere_rerank_params(self, model: str) -> _SupportedRerankParams: + return _SupportedRerankParams(("query", "documents", "top_n")) def map_cohere_rerank_params( self, @@ -76,22 +92,18 @@ class XinferenceRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, object]], + documents: Sequence[str | Mapping[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, - rank_fields: list[str] | 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, - ) -> dict[str, object]: - params: Final[OptionalRerankParams] = OptionalRerankParams( - query=query, - documents=documents, - ) + ) -> _RerankPayload: if top_n is not None: - params["top_n"] = top_n - return dict(params) + return _RerankPayload(MappingProxyType({"query": query, "documents": documents, "top_n": top_n})) + return _RerankPayload(MappingProxyType({"query": query, "documents": documents})) def transform_rerank_request( self, @@ -99,20 +111,32 @@ class XinferenceRerankConfig(BaseRerankConfig): optional_rerank_params: Mapping[str, object], headers: Mapping[str, object], litellm_params: Mapping[str, object] | None = None, - ) -> dict[str, object]: + ) -> _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") - request: Final[dict[str, object]] = { - "model": model, - "query": optional_rerank_params["query"], - "documents": optional_rerank_params["documents"], - } if optional_rerank_params.get("top_n") is not None: - request["top_n"] = optional_rerank_params["top_n"] - return request + 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, @@ -150,6 +174,6 @@ class XinferenceRerankConfig(BaseRerankConfig): return RerankResponse( id=response_json.id or str(uuid.uuid4()), - results=list(transformed_results), + results=_RerankResults(transformed_results), meta=meta, ) diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 7420e0a5a81..c7677da86d1 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -499,14 +499,14 @@ def rerank( litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.XINFERENCE: - api_key: Final = ( + 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: Final = ( + api_base = ( dynamic_api_base or optional_params.api_base or litellm.api_base diff --git a/litellm/utils.py b/litellm/utils.py index b66c08c6964..99373c50a20 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8161,6 +8161,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, @@ -8199,14 +8208,8 @@ class ProviderConfigManager: return litellm.VoyageRerankConfig() elif litellm.LlmProviders.WATSONX == provider: return litellm.IBMWatsonXRerankConfig() - elif litellm.LlmProviders.XINFERENCE == provider: - return litellm.XinferenceRerankConfig() - 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