mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(dashscope): support qwen3.7 text reranking
This commit is contained in:
parent
13df85cceb
commit
cf8f548792
8 changed files with 396 additions and 63 deletions
44
cookbook/dashscope_rerank.md
Normal file
44
cookbook/dashscope_rerank.md
Normal file
|
|
@ -0,0 +1,44 @@
|
|||
# Qwen3.7 text reranking
|
||||
|
||||
Use `dashscope/qwen3.7-text-rerank` with LiteLLM's rerank interface and a Beijing DashScope API key
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
response = litellm.rerank(
|
||||
model="dashscope/qwen3.7-text-rerank",
|
||||
query="How can I reset my password?",
|
||||
documents=[
|
||||
"The weather is sunny today.",
|
||||
"Open account settings and select Reset password.",
|
||||
"How do I change my password?",
|
||||
],
|
||||
top_n=2,
|
||||
return_documents=True,
|
||||
instruction="Retrieve semantically similar text.",
|
||||
)
|
||||
```
|
||||
|
||||
Set `DASHSCOPE_API_KEY` in the environment or pass `api_key` explicitly. The asynchronous equivalent is `await litellm.arerank(...)`
|
||||
|
||||
For the proxy, add a model entry and send the same query, documents and options to `/v1/rerank` using its configured model alias
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: qwen37-rerank
|
||||
litellm_params:
|
||||
model: dashscope/qwen3.7-text-rerank
|
||||
api_key: os.environ/DASHSCOPE_API_KEY
|
||||
model_info:
|
||||
mode: rerank
|
||||
```
|
||||
|
||||
LiteLLM sends the native DashScope request to `https://dashscope.aliyuncs.com/api/v1/services/rerank/text-rerank/text-rerank`. An explicit `api_base` or `DASHSCOPE_API_BASE_RERANK` can select a different host, an `/api/v1` base, or the complete native endpoint. Chat-compatible DashScope paths are converted to the native path for this model
|
||||
|
||||
`instruction` maps to DashScope's `parameters.instruct`. When omitted, the provider chooses its default relevance criterion. `top_n` and `return_documents` map to the corresponding native parameters. Live calls on 2026-09-08 confirmed that this model returns `document.text` when `return_documents=True`, despite the official parameter table omitting it from the supported-model list
|
||||
|
||||
The response contains the provider request ID, original document indices, relevance scores and requested document text. `meta.tokens.input_tokens` comes from `usage.prompt_tokens`; `meta.billed_units.total_tokens` comes from `usage.total_tokens`. These counters do not add a model price or dollar-cost calculation
|
||||
|
||||
The existing `dashscope/qwen3-rerank` model retains its compatible protocol. This change does not add multimodal reranking or reinterpret structured candidate documents
|
||||
|
||||
Protocol reference: [DashScope text rerank API](https://help.aliyun.com/zh/model-studio/text-rerank-api)
|
||||
|
|
@ -2,7 +2,7 @@
|
|||
Common utilities for the DashScope LLM provider.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -15,6 +15,7 @@ if TYPE_CHECKING:
|
|||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
|
||||
from litellm.llms.dashscope.rerank.transformation import DashScopeRerankConfig
|
||||
|
||||
|
||||
def get_dashscope_family_embedding_config(custom_llm_provider: str) -> "BaseEmbeddingConfig":
|
||||
|
|
@ -33,7 +34,16 @@ def get_dashscope_family_embedding_config(custom_llm_provider: str) -> "BaseEmbe
|
|||
return DashScopeEmbeddingConfig()
|
||||
|
||||
|
||||
def get_dashscope_family_rerank_config(custom_llm_provider: str) -> "BaseRerankConfig":
|
||||
def get_dashscope_family_rerank_config(custom_llm_provider: str, model: str) -> "BaseRerankConfig":
|
||||
provider_config: Final = _get_dashscope_family_rerank_provider_config(custom_llm_provider)
|
||||
if model == "qwen3.7-text-rerank":
|
||||
from litellm.llms.dashscope.rerank.native_transformation import DashScopeNativeRerankConfig
|
||||
|
||||
return DashScopeNativeRerankConfig(provider_config)
|
||||
return provider_config
|
||||
|
||||
|
||||
def _get_dashscope_family_rerank_provider_config(custom_llm_provider: str) -> "DashScopeRerankConfig":
|
||||
if custom_llm_provider == "qwencloud":
|
||||
from litellm.llms.dashscope.qwencloud import QwenCloudRerankConfig
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Final
|
||||
from typing import ClassVar, Final
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
|
@ -47,11 +47,13 @@ class QwenAIPlatformEmbeddingConfig(DashScopeEmbeddingConfig):
|
|||
|
||||
|
||||
class QwenAIPlatformRerankConfig(DashScopeRerankConfig):
|
||||
DEFAULT_RERANK_API_BASE: ClassVar[str] = QWEN_AI_PLATFORM_RERANK_API_BASE
|
||||
|
||||
def _resolve_api_key(self, api_key: str | None) -> str:
|
||||
return _require_qwen_ai_platform_api_key(api_key)
|
||||
|
||||
def _resolve_rerank_api_base(self, api_base: str | None) -> str:
|
||||
return api_base or get_secret_str("QWEN_AI_PLATFORM_API_BASE_RERANK") or QWEN_AI_PLATFORM_RERANK_API_BASE
|
||||
return api_base or get_secret_str("QWEN_AI_PLATFORM_API_BASE_RERANK") or self.DEFAULT_RERANK_API_BASE
|
||||
|
||||
|
||||
class QwenAIPlatformImageGenerationConfig(DashScopeImageGenerationConfig):
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Final
|
||||
from typing import ClassVar, Final
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
|
@ -47,11 +47,13 @@ class QwenCloudEmbeddingConfig(DashScopeEmbeddingConfig):
|
|||
|
||||
|
||||
class QwenCloudRerankConfig(DashScopeRerankConfig):
|
||||
DEFAULT_RERANK_API_BASE: ClassVar[str] = QWENCLOUD_RERANK_API_BASE
|
||||
|
||||
def _resolve_api_key(self, api_key: str | None) -> str:
|
||||
return _require_qwencloud_api_key(api_key)
|
||||
|
||||
def _resolve_rerank_api_base(self, api_base: str | None) -> str:
|
||||
return api_base or get_secret_str("QWENCLOUD_API_BASE_RERANK") or QWENCLOUD_RERANK_API_BASE
|
||||
return api_base or get_secret_str("QWENCLOUD_API_BASE_RERANK") or self.DEFAULT_RERANK_API_BASE
|
||||
|
||||
|
||||
class QwenCloudImageGenerationConfig(DashScopeImageGenerationConfig):
|
||||
|
|
|
|||
83
litellm/llms/dashscope/rerank/native_transformation.py
Normal file
83
litellm/llms/dashscope/rerank/native_transformation.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
"""DashScope native text reranking with input/parameters and output.results envelopes.
|
||||
|
||||
qwen3.7-text-rerank was verified in Beijing, including return_documents=true.
|
||||
Protocol routing through brand aliases does not establish regional model availability.
|
||||
Docs: https://help.aliyun.com/zh/model-studio/text-rerank-api
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from .transformation import DashScopeRerankConfig, DashScopeRerankUsage
|
||||
|
||||
|
||||
class DashScopeNativeRerankConfig(DashScopeRerankConfig):
|
||||
def __init__(self, provider_config: DashScopeRerankConfig) -> None:
|
||||
# Reuse brand-specific credentials and hosts without duplicating alias subclasses.
|
||||
self._provider_config: Final = provider_config
|
||||
|
||||
def _resolve_api_key(self, api_key: str | None) -> str:
|
||||
return self._provider_config._resolve_api_key(api_key)
|
||||
|
||||
def _resolve_rerank_api_base(self, api_base: str | None) -> str:
|
||||
return self._provider_config._resolve_rerank_api_base(api_base)
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object] | None = None,
|
||||
) -> str:
|
||||
# Provider discovery supplies the brand's chat default even when no api_base was passed.
|
||||
# Resolve the rerank-specific environment override before constructing the native path.
|
||||
default_chat_base: Final = self._provider_config.DEFAULT_RERANK_API_BASE.replace(
|
||||
"/compatible-api/v1/reranks", "/compatible-mode/v1"
|
||||
)
|
||||
native_base: Final = self._resolve_rerank_api_base(None if api_base == default_chat_base else api_base)
|
||||
parsed: Final = urlsplit(native_base.rstrip("/"))
|
||||
if parsed.path.endswith("/services/rerank/text-rerank/text-rerank"):
|
||||
return urlunsplit(parsed)
|
||||
native_path: Final = parsed.path.removesuffix("/compatible-mode/v1").removesuffix("/compatible-api/v1/reranks")
|
||||
api_path: Final = native_path if native_path.endswith("/api/v1") else f"{native_path}/api/v1"
|
||||
return urlunsplit(parsed._replace(path=f"{api_path}/services/rerank/text-rerank/text-rerank"))
|
||||
|
||||
def get_supported_cohere_rerank_params(self, model: str) -> list[str]:
|
||||
return [*super().get_supported_cohere_rerank_params(model), "instruction"]
|
||||
|
||||
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]:
|
||||
request: Final = super().transform_rerank_request(model, optional_rerank_params, headers, litellm_params)
|
||||
return {
|
||||
"model": model,
|
||||
"input": {"query": request["query"], "documents": request["documents"]},
|
||||
"parameters": {
|
||||
("instruct" if name == "instruction" else name): optional_rerank_params[name]
|
||||
for name in ("top_n", "return_documents", "instruction")
|
||||
if optional_rerank_params.get(name) is not None
|
||||
},
|
||||
}
|
||||
|
||||
def _get_request_query(self, request_data: Mapping[str, object]) -> object:
|
||||
return (
|
||||
TypeAdapter(Mapping[str, object])
|
||||
.validate_python(request_data.get("input", MappingProxyType({})))
|
||||
.get("query")
|
||||
)
|
||||
|
||||
def _get_response_fields(
|
||||
self, response_json: Mapping[str, object], usage: DashScopeRerankUsage
|
||||
) -> tuple[object, object, int | None]:
|
||||
output: Final = TypeAdapter(Mapping[str, object]).validate_python(
|
||||
response_json.get("output", MappingProxyType({}))
|
||||
)
|
||||
# Native usage separates prompt_tokens from total_tokens; keep both provider counters.
|
||||
return output.get("results"), response_json.get("request_id"), usage.get("prompt_tokens")
|
||||
|
|
@ -4,14 +4,15 @@ Transformation logic for DashScope's OpenAI-compatible /v1/reranks API.
|
|||
Supports
|
||||
- qwen3-rerank
|
||||
|
||||
(Other DashScope rerankers — gte-rerank-v2 / qwen3-vl-rerank — share the same
|
||||
endpoint but have not been validated against this transformer. Behavior with
|
||||
those models is undefined.)
|
||||
(Other DashScope rerankers — gte-rerank-v2 / qwen3-vl-rerank — have not been
|
||||
validated against this transformer. Behavior with those models is undefined.)
|
||||
|
||||
The native qwen3.7-text-rerank protocol is implemented in native_transformation.py.
|
||||
|
||||
Endpoint
|
||||
- https://dashscope.aliyuncs.com/compatible-api/v1/reranks
|
||||
|
||||
Note: chat/embed live under `/compatible-mode/v1/`, but DashScope's rerank
|
||||
Note: chat/embed live under `/compatible-mode/v1/`, but qwen3-rerank's
|
||||
route is exposed under `/compatible-api/v1/reranks` per the docs. Override
|
||||
with `DASHSCOPE_API_BASE_RERANK` to point at a different host or path.
|
||||
|
||||
|
|
@ -23,9 +24,12 @@ Docs - https://help.aliyun.com/zh/model-studio/text-rerank-api
|
|||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Final
|
||||
from types import MappingProxyType
|
||||
from typing import ClassVar, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -33,7 +37,6 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
|||
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,
|
||||
RerankResponseMeta,
|
||||
|
|
@ -45,6 +48,11 @@ from ..common_utils import DashScopeError
|
|||
DEFAULT_RERANK_URL: Final = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks"
|
||||
|
||||
|
||||
class DashScopeRerankUsage(TypedDict, total=False):
|
||||
prompt_tokens: ReadOnly[int | None]
|
||||
total_tokens: ReadOnly[int | None]
|
||||
|
||||
|
||||
class DashScopeRerankConfig(BaseRerankConfig):
|
||||
"""
|
||||
Reference: https://help.aliyun.com/zh/model-studio/text-rerank-api
|
||||
|
|
@ -53,8 +61,12 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
documents, top_n, return_documents. Response: results[].index,
|
||||
results[].relevance_score, optionally results[].document.text (when
|
||||
return_documents=true), plus a top-level usage.total_tokens counter.
|
||||
|
||||
Brand aliases supply their own API key and base resolvers.
|
||||
"""
|
||||
|
||||
DEFAULT_RERANK_API_BASE: ClassVar[str] = DEFAULT_RERANK_URL
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
|
|
@ -69,13 +81,13 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
def _resolve_rerank_api_base(self, api_base: str | None) -> str:
|
||||
if api_base is not None:
|
||||
return api_base
|
||||
return get_secret_str("DASHSCOPE_API_BASE_RERANK") or DEFAULT_RERANK_URL
|
||||
return get_secret_str("DASHSCOPE_API_BASE_RERANK") or self.DEFAULT_RERANK_API_BASE
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: dict | None = None,
|
||||
optional_params: Mapping[str, object] | None = None,
|
||||
) -> str:
|
||||
resolved_api_base: Final = self._resolve_rerank_api_base(api_base)
|
||||
if resolved_api_base == DEFAULT_RERANK_URL:
|
||||
|
|
@ -93,12 +105,12 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
headers: Mapping[str, object],
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
optional_params: dict | None = None,
|
||||
optional_params: Mapping[str, object] | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
) -> dict:
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"Authorization": f"Bearer {self._resolve_api_key(api_key)}",
|
||||
"accept": "application/json",
|
||||
|
|
@ -106,16 +118,16 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
**headers,
|
||||
}
|
||||
|
||||
def get_supported_cohere_rerank_params(self, model: str) -> list:
|
||||
def get_supported_cohere_rerank_params(self, model: str) -> list[str]:
|
||||
return ["query", "documents", "top_n", "return_documents"]
|
||||
|
||||
def map_cohere_rerank_params(
|
||||
self,
|
||||
non_default_params: dict | None,
|
||||
non_default_params: Mapping[str, object] | None,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
query: str,
|
||||
documents: list[str | dict[str, Any]],
|
||||
documents: list[str | dict[str, object]],
|
||||
custom_llm_provider: str | None = None,
|
||||
top_n: int | None = None,
|
||||
rank_fields: list[str] | None = None,
|
||||
|
|
@ -123,26 +135,27 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
max_chunks_per_doc: int | None = None,
|
||||
max_tokens_per_doc: int | None = None,
|
||||
instruction: str | None = None,
|
||||
) -> dict:
|
||||
# qwen3-rerank accepts query/documents/top_n/return_documents. The
|
||||
# rest (rank_fields, max_*_per_doc) are silently dropped.
|
||||
params: Final[OptionalRerankParams] = OptionalRerankParams(
|
||||
query=query,
|
||||
documents=documents,
|
||||
) -> dict[str, object]:
|
||||
# rank_fields and max_*_per_doc have no supported mapping and are omitted.
|
||||
params: Final = MappingProxyType(
|
||||
{
|
||||
"query": query,
|
||||
"documents": documents,
|
||||
"top_n": top_n,
|
||||
"return_documents": return_documents,
|
||||
"instruction": instruction,
|
||||
}
|
||||
)
|
||||
if top_n is not None:
|
||||
params["top_n"] = top_n
|
||||
if return_documents is not None:
|
||||
params["return_documents"] = return_documents
|
||||
return dict(params)
|
||||
supported_params: Final = self.get_supported_cohere_rerank_params(model)
|
||||
return {name: value for name, value in params.items() if value is not None and name in supported_params}
|
||||
|
||||
def transform_rerank_request(
|
||||
self,
|
||||
model: str,
|
||||
optional_rerank_params: dict,
|
||||
headers: dict,
|
||||
litellm_params: dict | None = None,
|
||||
) -> dict:
|
||||
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 DashScope rerank")
|
||||
if "documents" not in optional_rerank_params:
|
||||
|
|
@ -167,22 +180,20 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: str | None = None,
|
||||
request_data: dict | None = None,
|
||||
optional_params: dict | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
optional_params: Mapping[str, object] | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
) -> RerankResponse:
|
||||
request_data = request_data or {}
|
||||
optional_params = optional_params or {}
|
||||
litellm_params = litellm_params or {}
|
||||
try:
|
||||
response_json: Final = raw_response.json()
|
||||
except Exception:
|
||||
response_json: Final = TypeAdapter(Mapping[str, object]).validate_json(raw_response.content)
|
||||
except ValidationError as exc:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=raw_response.text,
|
||||
)
|
||||
) from exc
|
||||
|
||||
logging_obj.post_call(
|
||||
input=request_data.get("query"),
|
||||
input=self._get_request_query(request_data),
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
original_response=response_json,
|
||||
|
|
@ -192,23 +203,26 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
if "code" in response_json and "results" not in response_json:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=response_json.get("message", str(response_json)),
|
||||
message=str(response_json.get("message", response_json)),
|
||||
)
|
||||
|
||||
results: Final = response_json.get("results")
|
||||
usage: Final = TypeAdapter(DashScopeRerankUsage).validate_python(
|
||||
response_json.get("usage") or MappingProxyType({})
|
||||
)
|
||||
results, response_id, input_tokens = self._get_response_fields(response_json, usage)
|
||||
if results is None:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=f"No results in DashScope rerank response: {response_json}",
|
||||
)
|
||||
|
||||
# qwen3-rerank returns:
|
||||
# Both protocols return:
|
||||
# {"index": int, "relevance_score": float}
|
||||
# plus, when return_documents=true was sent:
|
||||
# "document": {"text": "..."}
|
||||
# which already matches LiteLLM's RerankResponseDocument shape.
|
||||
transformed_results: Final[list[dict]] = []
|
||||
for r in results:
|
||||
for r in TypeAdapter(tuple[Mapping[str, object], ...]).validate_python(results):
|
||||
item: dict[str, object] = {
|
||||
"index": r["index"],
|
||||
"relevance_score": r["relevance_score"],
|
||||
|
|
@ -221,18 +235,29 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
item["document"] = {"text": doc}
|
||||
transformed_results.append(item)
|
||||
|
||||
usage: Final = response_json.get("usage") or {}
|
||||
total_tokens: Final = usage.get("total_tokens")
|
||||
billed_units: Final = RerankBilledUnits(total_tokens=total_tokens)
|
||||
tokens: Final = RerankTokens(input_tokens=total_tokens)
|
||||
billed_units: Final = RerankBilledUnits(total_tokens=usage.get("total_tokens"))
|
||||
tokens: Final = RerankTokens(input_tokens=input_tokens)
|
||||
meta: Final = RerankResponseMeta(billed_units=billed_units, tokens=tokens)
|
||||
|
||||
return RerankResponse(
|
||||
id=response_json.get("id") or str(uuid.uuid4()),
|
||||
results=transformed_results,
|
||||
meta=meta,
|
||||
return RerankResponse.model_validate(
|
||||
MappingProxyType(
|
||||
{
|
||||
"id": response_id or str(uuid.uuid4()),
|
||||
"results": transformed_results,
|
||||
"meta": meta,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
def _get_request_query(self, request_data: Mapping[str, object]) -> object:
|
||||
return request_data.get("query")
|
||||
|
||||
def _get_response_fields(
|
||||
self, response_json: Mapping[str, object], usage: DashScopeRerankUsage
|
||||
) -> tuple[object, object, int | None]:
|
||||
# The compatible API reports its input count only as total_tokens.
|
||||
return response_json.get("results"), response_json.get("id"), usage.get("total_tokens")
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
|
|
|
|||
|
|
@ -8439,7 +8439,7 @@ class ProviderConfigManager:
|
|||
get_dashscope_family_rerank_config,
|
||||
)
|
||||
|
||||
return get_dashscope_family_rerank_config(provider.value)
|
||||
return get_dashscope_family_rerank_config(provider.value, model)
|
||||
return litellm.CohereRerankConfig()
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -7,9 +7,10 @@ from unittest.mock import MagicMock
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
|
||||
from litellm.llms.dashscope.common_utils import DashScopeError
|
||||
from litellm.llms.dashscope.common_utils import DashScopeError, get_dashscope_family_rerank_config
|
||||
from litellm.llms.dashscope.rerank.transformation import (
|
||||
DEFAULT_RERANK_URL,
|
||||
DashScopeRerankConfig,
|
||||
|
|
@ -103,6 +104,7 @@ class TestDashScopeRerankRequest:
|
|||
return_documents=True,
|
||||
max_chunks_per_doc=5,
|
||||
max_tokens_per_doc=100,
|
||||
instruction="Unsupported on the compatible protocol",
|
||||
)
|
||||
assert params == {
|
||||
"query": "什么是文本排序模型",
|
||||
|
|
@ -274,30 +276,32 @@ class TestDashScopeRerankResponse:
|
|||
)
|
||||
assert out.id is not None and len(out.id) > 0
|
||||
|
||||
def test_error_envelope_raises(self):
|
||||
@pytest.mark.parametrize("model", ["qwen3-rerank", "qwen3.7-text-rerank"])
|
||||
def test_error_envelope_raises(self, model):
|
||||
body = {
|
||||
"code": "InvalidApiKey",
|
||||
"message": "Invalid API-key provided.",
|
||||
"request_id": "fb53",
|
||||
}
|
||||
with pytest.raises(DashScopeError) as exc_info:
|
||||
self.config.transform_rerank_response(
|
||||
model="qwen3-rerank",
|
||||
get_dashscope_family_rerank_config("dashscope", model).transform_rerank_response(
|
||||
model=model,
|
||||
raw_response=self._resp(body, status_code=401),
|
||||
model_response=RerankResponse(),
|
||||
logging_obj=self.logging,
|
||||
)
|
||||
assert "Invalid API-key provided." in str(exc_info.value)
|
||||
|
||||
def test_non_json_response_raises(self):
|
||||
@pytest.mark.parametrize("model", ["qwen3-rerank", "qwen3.7-text-rerank"])
|
||||
def test_non_json_response_raises(self, model):
|
||||
bad = httpx.Response(
|
||||
status_code=500,
|
||||
content=b"<html>bad gateway</html>",
|
||||
request=httpx.Request("POST", "https://example.com"),
|
||||
)
|
||||
with pytest.raises(DashScopeError):
|
||||
self.config.transform_rerank_response(
|
||||
model="qwen3-rerank",
|
||||
get_dashscope_family_rerank_config("dashscope", model).transform_rerank_response(
|
||||
model=model,
|
||||
raw_response=bad,
|
||||
model_response=RerankResponse(),
|
||||
logging_obj=self.logging,
|
||||
|
|
@ -323,3 +327,166 @@ class TestProviderConfigManagerDispatch:
|
|||
present_version_params=[],
|
||||
)
|
||||
assert isinstance(cfg, DashScopeRerankConfig)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.parametrize("return_documents", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
"provider,host",
|
||||
[
|
||||
("dashscope", "dashscope.aliyuncs.com"),
|
||||
("qwencloud", "dashscope-intl.aliyuncs.com"),
|
||||
("qwen_ai_platform", "dashscope.aliyuncs.com"),
|
||||
],
|
||||
)
|
||||
async def test_qwen37_rerank_public_call(
|
||||
is_async, return_documents, provider, host, respx_mock: respx.MockRouter, monkeypatch
|
||||
):
|
||||
import litellm
|
||||
|
||||
monkeypatch.delenv("DASHSCOPE_API_BASE", raising=False)
|
||||
monkeypatch.delenv("DASHSCOPE_API_BASE_RERANK", raising=False)
|
||||
monkeypatch.delenv(f"{provider.upper()}_API_BASE", raising=False)
|
||||
monkeypatch.delenv(f"{provider.upper()}_API_BASE_RERANK", raising=False)
|
||||
monkeypatch.setenv(f"{provider.upper()}_API_KEY", "fake-brand-key")
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
route = respx_mock.post(f"https://{host}/api/v1/services/rerank/text-rerank/text-rerank")
|
||||
results = [{"index": 1, "relevance_score": 0.88, **({"document": {"text": "answer"}} if return_documents else {})}]
|
||||
route.respond(
|
||||
200,
|
||||
json={
|
||||
"output": {"results": results},
|
||||
"usage": {"prompt_tokens": 237, "total_tokens": 261, "details": {"provider_metadata": True}},
|
||||
"request_id": "qwen37-request-id",
|
||||
},
|
||||
)
|
||||
kwargs = {
|
||||
"model": f"{provider}/qwen3.7-text-rerank",
|
||||
"query": "question",
|
||||
"documents": ["unrelated", "answer"],
|
||||
"top_n": 1,
|
||||
"return_documents": return_documents,
|
||||
"instruction": "Retrieve semantically similar text.",
|
||||
}
|
||||
|
||||
response = await litellm.arerank(**kwargs) if is_async else litellm.rerank(**kwargs)
|
||||
|
||||
assert json.loads(route.calls[0].request.content) == {
|
||||
"model": "qwen3.7-text-rerank",
|
||||
"input": {"query": "question", "documents": ["unrelated", "answer"]},
|
||||
"parameters": {
|
||||
"top_n": 1,
|
||||
"return_documents": return_documents,
|
||||
"instruct": "Retrieve semantically similar text.",
|
||||
},
|
||||
}
|
||||
assert route.calls[0].request.headers["authorization"] == "Bearer fake-brand-key"
|
||||
assert response.id == "qwen37-request-id"
|
||||
assert response.results == results
|
||||
assert response.meta == {"billed_units": {"total_tokens": 261}, "tokens": {"input_tokens": 237}}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base",
|
||||
[
|
||||
"https://proxy.example/api/v1",
|
||||
"https://proxy.example/api/v1/services/rerank/text-rerank/text-rerank/",
|
||||
],
|
||||
)
|
||||
def test_qwen37_rerank_custom_url(api_base):
|
||||
assert get_dashscope_family_rerank_config("dashscope", "qwen3.7-text-rerank").get_complete_url(
|
||||
api_base, "qwen3.7-text-rerank"
|
||||
) == ("https://proxy.example/api/v1/services/rerank/text-rerank/text-rerank")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider,host",
|
||||
[("dashscope", "dashscope-intl.aliyuncs.com"), ("qwencloud", "dashscope.aliyuncs.com")],
|
||||
)
|
||||
def test_qwen37_rerank_explicit_region(provider, host):
|
||||
config = get_dashscope_family_rerank_config(provider, "qwen3.7-text-rerank")
|
||||
assert config.get_complete_url(f"https://{host}/compatible-mode/v1", "qwen3.7-text-rerank") == (
|
||||
f"https://{host}/api/v1/services/rerank/text-rerank/text-rerank"
|
||||
)
|
||||
|
||||
|
||||
def test_qwen37_rerank_response_logging():
|
||||
config = get_dashscope_family_rerank_config("dashscope", "qwen3.7-text-rerank")
|
||||
logging = MagicMock()
|
||||
request = {"model": "qwen3.7-text-rerank", "input": {"query": "question", "documents": ["answer"]}, "parameters": {}}
|
||||
payload = {"request_id": "request-id", "output": {"results": [{"index": 0, "relevance_score": 0.88}]}}
|
||||
|
||||
response = config.transform_rerank_response(
|
||||
model="qwen3.7-text-rerank",
|
||||
raw_response=httpx.Response(200, json=payload),
|
||||
model_response=RerankResponse(),
|
||||
logging_obj=logging,
|
||||
request_data=request,
|
||||
)
|
||||
|
||||
logging.post_call.assert_called_once_with(
|
||||
input="question", api_key=None, additional_args={"complete_input_dict": request}, original_response=payload
|
||||
)
|
||||
assert response.id == "request-id"
|
||||
assert response.results == [{"index": 0, "relevance_score": 0.88}]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ["dashscope", "qwencloud", "qwen_ai_platform"])
|
||||
def test_qwen37_rerank_environment_url(provider, respx_mock: respx.MockRouter, monkeypatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.delenv("DASHSCOPE_API_BASE", raising=False)
|
||||
monkeypatch.delenv(f"{provider.upper()}_API_BASE", raising=False)
|
||||
monkeypatch.setenv(f"{provider.upper()}_API_BASE_RERANK", "https://proxy.example/api/v1")
|
||||
route = respx_mock.post("https://proxy.example/api/v1/services/rerank/text-rerank/text-rerank")
|
||||
route.respond(
|
||||
200,
|
||||
json={
|
||||
"output": {"results": [{"index": 0, "relevance_score": 0.88}]},
|
||||
"request_id": "all-results",
|
||||
"usage": {"prompt_tokens": 10, "total_tokens": 10},
|
||||
},
|
||||
)
|
||||
|
||||
response = litellm.rerank(
|
||||
model=f"{provider}/qwen3.7-text-rerank",
|
||||
query="question",
|
||||
documents=["answer"],
|
||||
return_documents=None,
|
||||
api_key="fake-dashscope-key",
|
||||
)
|
||||
|
||||
assert json.loads(route.calls[0].request.content) == {
|
||||
"model": "qwen3.7-text-rerank",
|
||||
"input": {"query": "question", "documents": ["answer"]},
|
||||
"parameters": {},
|
||||
}
|
||||
assert response.results == [{"index": 0, "relevance_score": 0.88}]
|
||||
assert response.meta["tokens"]["input_tokens"] == 10
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
async def test_qwen37_rerank_preserves_provider_error(is_async, respx_mock: respx.MockRouter, monkeypatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
route = respx_mock.post("https://proxy.example/api/v1/services/rerank/text-rerank/text-rerank")
|
||||
route.respond(
|
||||
400,
|
||||
json={"code": "InvalidParameter", "message": "documents must not be empty", "request_id": "invalid-documents"},
|
||||
)
|
||||
kwargs = {
|
||||
"model": "dashscope/qwen3.7-text-rerank",
|
||||
"query": "question",
|
||||
"documents": [],
|
||||
"api_key": "fake-dashscope-key",
|
||||
"api_base": "https://proxy.example/api/v1",
|
||||
}
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="documents must not be empty") as error:
|
||||
await litellm.arerank(**kwargs) if is_async else litellm.rerank(**kwargs)
|
||||
|
||||
assert error.value.status_code == 400
|
||||
assert "DashscopeException" in str(error.value)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue