mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-05 08:07:05 +00:00
Merge 15e10d35f4 into ef8eb98724
This commit is contained in:
commit
1c3a61e1df
6 changed files with 351 additions and 7 deletions
|
|
@ -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"),
|
||||
|
|
|
|||
1
litellm/llms/xinference/rerank/__init__.py
Normal file
1
litellm/llms/xinference/rerank/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
|
||||
179
litellm/llms/xinference/rerank/transformation.py
Normal file
179
litellm/llms/xinference/rerank/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."},
|
||||
},
|
||||
]
|
||||
Loading…
Add table
Reference in a new issue