diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 1c833256598..9393717dfb0 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", @@ -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"), 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..016f10b46cb --- /dev/null +++ b/litellm/llms/xinference/rerank/transformation.py @@ -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, + ) diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index c8f7842aebf..c87514d79d9 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, @@ -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, diff --git a/litellm/utils.py b/litellm/utils.py index 802dc151428..e1e07f44bc1 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 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."}, + }, + ]