fix(rerank): resolve Xinference lint violations

This commit is contained in:
qdivan 2026-08-17 14:33:00 +08:00
parent 8aff5f2a69
commit 15e10d35f4
3 changed files with 66 additions and 39 deletions

View file

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

View file

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

View file

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