mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(dashscope): forward rerank instructions
This commit is contained in:
parent
3fd1d4e741
commit
4ae6403be9
7 changed files with 85 additions and 503 deletions
|
|
@ -2,16 +2,11 @@
|
|||
Common utilities for the DashScope LLM provider.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
|
@ -21,7 +16,6 @@ if TYPE_CHECKING:
|
|||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
|
||||
from litellm.llms.dashscope.rerank.transformation import DashScopeRerankConfig
|
||||
|
||||
DASHSCOPE_CHAT_COMPATIBLE_PATH: Final = "/compatible-mode/v1"
|
||||
DASHSCOPE_RERANK_PATH: Final = "/compatible-api/v1/reranks"
|
||||
|
|
@ -61,60 +55,7 @@ def get_dashscope_family_embedding_config(custom_llm_provider: str) -> "BaseEmbe
|
|||
return DashScopeEmbeddingConfig()
|
||||
|
||||
|
||||
def get_dashscope_family_rerank_config(
|
||||
custom_llm_provider: str, model: str, api_base: str | None = None
|
||||
) -> "BaseRerankConfig":
|
||||
provider_config: Final = _get_dashscope_family_rerank_provider_config(custom_llm_provider)
|
||||
model_cost: Final[Mapping[str, object]] = litellm.model_cost
|
||||
runtime_api: Final = next(
|
||||
(
|
||||
declared_api
|
||||
for key in (f"{custom_llm_provider}/{model}", f"dashscope/{model}", model)
|
||||
if (declared_api := _rerank_api_from_model_info(model_cost.get(key))) is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
rerank_api: Final = runtime_api or _bundled_dashscope_rerank_apis().get(f"dashscope/{model}")
|
||||
if rerank_api == "native":
|
||||
from litellm.llms.dashscope.rerank.native_transformation import DashScopeNativeRerankConfig
|
||||
|
||||
return DashScopeNativeRerankConfig(
|
||||
provider_config, api_base=api_base or get_secret_str(f"{custom_llm_provider.upper()}_API_BASE_RERANK")
|
||||
)
|
||||
return provider_config
|
||||
|
||||
|
||||
def _rerank_api_from_model_info(raw_model_info: object) -> str | None:
|
||||
if raw_model_info is None:
|
||||
return None
|
||||
model_info: Final = TypeAdapter(Mapping[str, object]).validate_python(raw_model_info)
|
||||
provider_info: Final = model_info.get("provider_specific_entry")
|
||||
if provider_info is None:
|
||||
return None
|
||||
metadata: Final = TypeAdapter(Mapping[str, object]).validate_python(provider_info)
|
||||
rerank_api: Final[str | None] = TypeAdapter(Literal["native", "compatible"] | None).validate_python(
|
||||
metadata.get("rerank_api")
|
||||
)
|
||||
return rerank_api
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _bundled_dashscope_rerank_apis() -> Mapping[str, str]:
|
||||
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
|
||||
|
||||
model_infos: Final = TypeAdapter(Mapping[str, object]).validate_json(
|
||||
GetModelCostMap.read_local_model_cost_map_text()
|
||||
)
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: rerank_api
|
||||
for key, model_info in model_infos.items()
|
||||
if key.startswith("dashscope/") and (rerank_api := _rerank_api_from_model_info(model_info)) is not None
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _get_dashscope_family_rerank_provider_config(custom_llm_provider: str) -> "DashScopeRerankConfig":
|
||||
def get_dashscope_family_rerank_config(custom_llm_provider: str) -> "BaseRerankConfig":
|
||||
if custom_llm_provider == "qwencloud":
|
||||
from litellm.llms.dashscope.qwencloud import QwenCloudRerankConfig
|
||||
|
||||
|
|
|
|||
|
|
@ -1,77 +0,0 @@
|
|||
"""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, api_base: str | None = None) -> None:
|
||||
self._provider_config: Final = provider_config
|
||||
self._api_base: Final = api_base
|
||||
|
||||
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:
|
||||
native_base: Final = self._api_base or self._resolve_rerank_api_base(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({}))
|
||||
)
|
||||
return output.get("results"), response_json.get("request_id"), usage.get("prompt_tokens")
|
||||
|
|
@ -3,16 +3,16 @@ Transformation logic for DashScope's OpenAI-compatible /v1/reranks API.
|
|||
|
||||
Supports
|
||||
- qwen3-rerank
|
||||
- qwen3.7-text-rerank
|
||||
|
||||
(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.
|
||||
(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.)
|
||||
|
||||
Endpoint
|
||||
- https://dashscope.aliyuncs.com/compatible-api/v1/reranks
|
||||
|
||||
Note: chat/embed live under `/compatible-mode/v1/`, but qwen3-rerank's
|
||||
Note: chat/embed live under `/compatible-mode/v1/`, but DashScope's rerank
|
||||
route is exposed under `/compatible-api/v1/reranks` per the docs. A chat-shaped
|
||||
`.aliyuncs.com/compatible-mode/v1` base reaching this config (the chat default
|
||||
from `get_llm_provider`, or a `DASHSCOPE_API_BASE` env var) is redirected to
|
||||
|
|
@ -28,12 +28,9 @@ Docs - https://help.aliyun.com/zh/model-studio/text-rerank-api
|
|||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import 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
|
||||
|
|
@ -41,6 +38,7 @@ 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,
|
||||
|
|
@ -52,17 +50,12 @@ from ..common_utils import DashScopeError, resolve_dashscope_family_rerank_api_b
|
|||
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
|
||||
|
||||
Targets DashScope's qwen3-rerank model. Request fields: model, query,
|
||||
documents, top_n, return_documents. Response: results[].index,
|
||||
Targets DashScope's qwen3-rerank and qwen3.7-text-rerank models. Request fields:
|
||||
model, query, documents, top_n, return_documents, instruct. Response: results[].index,
|
||||
results[].relevance_score, optionally results[].document.text (when
|
||||
return_documents=true), plus a top-level usage.total_tokens counter.
|
||||
"""
|
||||
|
|
@ -85,7 +78,7 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object] | None = None,
|
||||
optional_params: dict | None = None,
|
||||
) -> str:
|
||||
resolved_api_base: Final = self._resolve_rerank_api_base(api_base)
|
||||
if resolved_api_base == DEFAULT_RERANK_URL:
|
||||
|
|
@ -103,12 +96,12 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, object],
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
optional_params: Mapping[str, object] | None = None,
|
||||
optional_params: dict | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
) -> dict[str, object]:
|
||||
) -> dict:
|
||||
return {
|
||||
"Authorization": f"Bearer {self._resolve_api_key(api_key)}",
|
||||
"accept": "application/json",
|
||||
|
|
@ -117,7 +110,7 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
}
|
||||
|
||||
def get_supported_cohere_rerank_params(self, model: str) -> list[str]:
|
||||
return ["query", "documents", "top_n", "return_documents"]
|
||||
return ["query", "documents", "top_n", "return_documents", "instruction"]
|
||||
|
||||
def map_cohere_rerank_params(
|
||||
self,
|
||||
|
|
@ -134,17 +127,14 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
max_tokens_per_doc: int | None = None,
|
||||
instruction: str | None = None,
|
||||
) -> dict[str, object]:
|
||||
params: Final = MappingProxyType(
|
||||
{
|
||||
"query": query,
|
||||
"documents": documents,
|
||||
"top_n": top_n,
|
||||
"return_documents": return_documents,
|
||||
"instruction": instruction,
|
||||
}
|
||||
params: Final[OptionalRerankParams] = OptionalRerankParams(
|
||||
query=query,
|
||||
documents=documents,
|
||||
top_n=top_n,
|
||||
return_documents=return_documents,
|
||||
instruction=instruction,
|
||||
)
|
||||
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}
|
||||
return {name: value for name, value in params.items() if value is not None}
|
||||
|
||||
def transform_rerank_request(
|
||||
self,
|
||||
|
|
@ -158,16 +148,18 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
if "documents" not in optional_rerank_params:
|
||||
raise ValueError("documents is required for DashScope rerank")
|
||||
|
||||
request: Final[dict[str, object]] = {
|
||||
"model": model,
|
||||
"query": optional_rerank_params["query"],
|
||||
"documents": optional_rerank_params["documents"],
|
||||
return {
|
||||
name: value
|
||||
for name, value in (
|
||||
("model", model),
|
||||
("query", optional_rerank_params["query"]),
|
||||
("documents", optional_rerank_params["documents"]),
|
||||
("top_n", optional_rerank_params.get("top_n")),
|
||||
("return_documents", optional_rerank_params.get("return_documents")),
|
||||
("instruct", optional_rerank_params.get("instruction")),
|
||||
)
|
||||
if name in ("model", "query", "documents") or value is not None
|
||||
}
|
||||
if optional_rerank_params.get("top_n") is not None:
|
||||
request["top_n"] = optional_rerank_params["top_n"]
|
||||
if optional_rerank_params.get("return_documents") is not None:
|
||||
request["return_documents"] = optional_rerank_params["return_documents"]
|
||||
return request
|
||||
|
||||
def transform_rerank_response(
|
||||
self,
|
||||
|
|
@ -177,22 +169,24 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: str | None = None,
|
||||
request_data: dict | None = None,
|
||||
optional_params: Mapping[str, object] | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
optional_params: dict | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> RerankResponse:
|
||||
request: Final = request_data or MappingProxyType({})
|
||||
request_data = request_data or {}
|
||||
optional_params = optional_params or {}
|
||||
litellm_params = litellm_params or {}
|
||||
try:
|
||||
response_json: Final = TypeAdapter(Mapping[str, object]).validate_json(raw_response.content)
|
||||
except ValidationError as exc:
|
||||
response_json: Final = raw_response.json()
|
||||
except Exception:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=raw_response.text,
|
||||
) from exc
|
||||
)
|
||||
|
||||
logging_obj.post_call(
|
||||
input=self._get_request_query(request),
|
||||
input=request_data.get("query"),
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": request},
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
original_response=response_json,
|
||||
)
|
||||
|
||||
|
|
@ -200,26 +194,23 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
if "code" in response_json and "results" not in response_json:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=str(response_json.get("message", response_json)),
|
||||
message=response_json.get("message", str(response_json)),
|
||||
)
|
||||
|
||||
usage: Final = TypeAdapter(DashScopeRerankUsage).validate_python(
|
||||
response_json.get("usage") or MappingProxyType({})
|
||||
)
|
||||
results, response_id, input_tokens = self._get_response_fields(response_json, usage)
|
||||
results: Final = response_json.get("results")
|
||||
if results is None:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=f"No results in DashScope rerank response: {response_json}",
|
||||
)
|
||||
|
||||
# Both protocols return:
|
||||
# qwen3-rerank returns:
|
||||
# {"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 TypeAdapter(tuple[Mapping[str, object], ...]).validate_python(results):
|
||||
for r in results:
|
||||
item: dict[str, object] = {
|
||||
"index": r["index"],
|
||||
"relevance_score": r["relevance_score"],
|
||||
|
|
@ -232,28 +223,18 @@ class DashScopeRerankConfig(BaseRerankConfig):
|
|||
item["document"] = {"text": doc}
|
||||
transformed_results.append(item)
|
||||
|
||||
billed_units: Final = RerankBilledUnits(total_tokens=usage.get("total_tokens"))
|
||||
tokens: Final = RerankTokens(input_tokens=input_tokens)
|
||||
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)
|
||||
meta: Final = RerankResponseMeta(billed_units=billed_units, tokens=tokens)
|
||||
|
||||
return RerankResponse.model_validate(
|
||||
MappingProxyType(
|
||||
{
|
||||
"id": response_id or str(uuid.uuid4()),
|
||||
"results": transformed_results,
|
||||
"meta": meta,
|
||||
}
|
||||
)
|
||||
return RerankResponse(
|
||||
id=response_json.get("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]:
|
||||
return response_json.get("results"), response_json.get("id"), usage.get("total_tokens")
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
|
|
|
|||
|
|
@ -15015,14 +15015,6 @@
|
|||
}
|
||||
]
|
||||
},
|
||||
"dashscope/qwen3.7-text-rerank": {
|
||||
"litellm_provider": "dashscope",
|
||||
"mode": "rerank",
|
||||
"provider_specific_entry": {
|
||||
"rerank_api": "native"
|
||||
},
|
||||
"source": "https://help.aliyun.com/zh/model-studio/text-rerank-api"
|
||||
},
|
||||
"dashscope/qwen-turbo": {
|
||||
"input_cost_per_token": 5e-08,
|
||||
"litellm_provider": "dashscope",
|
||||
|
|
|
|||
|
|
@ -8508,7 +8508,7 @@ class ProviderConfigManager:
|
|||
get_dashscope_family_rerank_config,
|
||||
)
|
||||
|
||||
return get_dashscope_family_rerank_config(provider.value, model, api_base)
|
||||
return get_dashscope_family_rerank_config(provider.value)
|
||||
return litellm.CohereRerankConfig()
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -15015,14 +15015,6 @@
|
|||
}
|
||||
]
|
||||
},
|
||||
"dashscope/qwen3.7-text-rerank": {
|
||||
"litellm_provider": "dashscope",
|
||||
"mode": "rerank",
|
||||
"provider_specific_entry": {
|
||||
"rerank_api": "native"
|
||||
},
|
||||
"source": "https://help.aliyun.com/zh/model-studio/text-rerank-api"
|
||||
},
|
||||
"dashscope/qwen-turbo": {
|
||||
"input_cost_per_token": 5e-08,
|
||||
"litellm_provider": "dashscope",
|
||||
|
|
|
|||
|
|
@ -7,10 +7,9 @@ from unittest.mock import MagicMock
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
|
||||
from litellm.llms.dashscope.common_utils import DashScopeError, get_dashscope_family_rerank_config
|
||||
from litellm.llms.dashscope.common_utils import DashScopeError
|
||||
from litellm.llms.dashscope.rerank.transformation import (
|
||||
DEFAULT_RERANK_URL,
|
||||
DashScopeRerankConfig,
|
||||
|
|
@ -112,6 +111,7 @@ class TestDashScopeRerankRequest:
|
|||
"documents",
|
||||
"top_n",
|
||||
"return_documents",
|
||||
"instruction",
|
||||
]
|
||||
|
||||
def test_map_params_drops_unsupported(self):
|
||||
|
|
@ -128,7 +128,6 @@ class TestDashScopeRerankRequest:
|
|||
return_documents=True,
|
||||
max_chunks_per_doc=5,
|
||||
max_tokens_per_doc=100,
|
||||
instruction="Unsupported on the compatible protocol",
|
||||
)
|
||||
assert params == {
|
||||
"query": "什么是文本排序模型",
|
||||
|
|
@ -300,32 +299,30 @@ class TestDashScopeRerankResponse:
|
|||
)
|
||||
assert out.id is not None and len(out.id) > 0
|
||||
|
||||
@pytest.mark.parametrize("model", ["qwen3-rerank", "qwen3.7-text-rerank"])
|
||||
def test_error_envelope_raises(self, model):
|
||||
def test_error_envelope_raises(self):
|
||||
body = {
|
||||
"code": "InvalidApiKey",
|
||||
"message": "Invalid API-key provided.",
|
||||
"request_id": "fb53",
|
||||
}
|
||||
with pytest.raises(DashScopeError) as exc_info:
|
||||
get_dashscope_family_rerank_config("dashscope", model).transform_rerank_response(
|
||||
model=model,
|
||||
self.config.transform_rerank_response(
|
||||
model="qwen3-rerank",
|
||||
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)
|
||||
|
||||
@pytest.mark.parametrize("model", ["qwen3-rerank", "qwen3.7-text-rerank"])
|
||||
def test_non_json_response_raises(self, model):
|
||||
def test_non_json_response_raises(self):
|
||||
bad = httpx.Response(
|
||||
status_code=500,
|
||||
content=b"<html>bad gateway</html>",
|
||||
request=httpx.Request("POST", "https://example.com"),
|
||||
)
|
||||
with pytest.raises(DashScopeError):
|
||||
get_dashscope_family_rerank_config("dashscope", model).transform_rerank_response(
|
||||
model=model,
|
||||
self.config.transform_rerank_response(
|
||||
model="qwen3-rerank",
|
||||
raw_response=bad,
|
||||
model_response=RerankResponse(),
|
||||
logging_obj=self.logging,
|
||||
|
|
@ -355,280 +352,36 @@ class TestProviderConfigManagerDispatch:
|
|||
|
||||
@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):
|
||||
@pytest.mark.parametrize("model", ["qwen3-rerank", "qwen3.7-text-rerank"])
|
||||
@pytest.mark.parametrize("instruction", [None, "", "Retrieve semantically similar text."])
|
||||
async def test_instruction_reaches_compatible_endpoint(provider, model, is_async, instruction, respx_mock, monkeypatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setenv(f"{provider.upper()}_API_BASE", "https://rerank.example/compatible-api/v1/reranks")
|
||||
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"},
|
||||
)
|
||||
route = respx_mock.post("https://rerank.example/compatible-api/v1/reranks")
|
||||
route.respond(200, json={"id": "ranking", "results": [{"index": 0, "relevance_score": 0.9}]})
|
||||
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)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.parametrize("provider", ["dashscope", "qwencloud", "qwen_ai_platform"])
|
||||
@pytest.mark.parametrize("base_case", ["rerank_override", "explicit", "explicit_default", "general_only"])
|
||||
async def test_native_rerank_base_precedence(provider, is_async, base_case, respx_mock, monkeypatch):
|
||||
import litellm
|
||||
|
||||
prefix = provider.upper()
|
||||
default_host = "dashscope-intl.aliyuncs.com" if provider == "qwencloud" else "dashscope.aliyuncs.com"
|
||||
monkeypatch.setenv(f"{prefix}_API_BASE", "https://chat.example/compatible-mode/v1")
|
||||
monkeypatch.delenv(f"{prefix}_API_BASE_RERANK", raising=False)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
if base_case != "general_only":
|
||||
monkeypatch.setenv(f"{prefix}_API_BASE_RERANK", "https://rerank.example/api/v1")
|
||||
explicit_base = {
|
||||
"rerank_override": None,
|
||||
"explicit": "https://explicit.example/api/v1",
|
||||
"explicit_default": f"https://{default_host}/compatible-mode/v1",
|
||||
"general_only": None,
|
||||
}[base_case]
|
||||
expected_host = {
|
||||
"rerank_override": "rerank.example",
|
||||
"explicit": "explicit.example",
|
||||
"explicit_default": default_host,
|
||||
"general_only": "chat.example",
|
||||
}[base_case]
|
||||
route = respx_mock.post(f"https://{expected_host}/api/v1/services/rerank/text-rerank/text-rerank")
|
||||
route.respond(
|
||||
200, json={"request_id": "base-precedence", "output": {"results": [{"index": 0, "relevance_score": 0.9}]}}
|
||||
)
|
||||
kwargs = {
|
||||
"model": f"{provider}/qwen3.7-text-rerank",
|
||||
"model": f"{provider}/{model}",
|
||||
"query": "question",
|
||||
"documents": ["answer"],
|
||||
"api_key": "fake-review-key",
|
||||
"api_base": explicit_base,
|
||||
"top_n": 1,
|
||||
"return_documents": False,
|
||||
"instruction": instruction,
|
||||
"api_key": "test-key",
|
||||
}
|
||||
|
||||
response = await litellm.arerank(**kwargs) if is_async else litellm.rerank(**kwargs)
|
||||
|
||||
assert json.loads(route.calls[0].request.content)["input"] == {"query": "question", "documents": ["answer"]}
|
||||
assert response.id == "base-precedence"
|
||||
assert response.results == [{"index": 0, "relevance_score": 0.9}]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ["dashscope", "qwencloud", "qwen_ai_platform"])
|
||||
@pytest.mark.parametrize("rerank_api", ["native", "compatible"])
|
||||
def test_rerank_protocol_uses_runtime_model_metadata(provider, rerank_api, respx_mock, monkeypatch):
|
||||
import litellm
|
||||
|
||||
model = "custom-rerank" if rerank_api == "native" else "qwen3.7-text-rerank"
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
model if provider == "dashscope" and rerank_api == "native" else f"{provider}/{model}",
|
||||
{
|
||||
"litellm_provider": provider,
|
||||
"mode": "rerank",
|
||||
"provider_specific_entry": {"rerank_api": rerank_api},
|
||||
},
|
||||
)
|
||||
results = [{"index": 0, "relevance_score": 0.9}]
|
||||
url = (
|
||||
"https://proxy.example/api/v1/services/rerank/text-rerank/text-rerank"
|
||||
if rerank_api == "native"
|
||||
else "https://proxy.example/api/v1/reranks"
|
||||
)
|
||||
route = respx_mock.post(url)
|
||||
route.respond(
|
||||
200,
|
||||
json={"request_id": "metadata", "output": {"results": results}}
|
||||
if rerank_api == "native"
|
||||
else {"id": "metadata", "results": results},
|
||||
)
|
||||
|
||||
response = litellm.rerank(
|
||||
model=f"{provider}/{model}",
|
||||
query="question",
|
||||
documents=["answer"],
|
||||
api_key="fake-key",
|
||||
api_base="https://proxy.example/api/v1",
|
||||
return_documents=None,
|
||||
)
|
||||
|
||||
assert json.loads(route.calls[0].request.content) == (
|
||||
{"model": model, "input": {"query": "question", "documents": ["answer"]}, "parameters": {}}
|
||||
if rerank_api == "native"
|
||||
else {"model": model, "query": "question", "documents": ["answer"]}
|
||||
)
|
||||
assert response.id == "metadata"
|
||||
assert response.results == results
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ["dashscope", "qwencloud", "qwen_ai_platform"])
|
||||
def test_rerank_uses_bundled_metadata_when_remote_map_lacks_model(provider, respx_mock, monkeypatch):
|
||||
import litellm
|
||||
|
||||
for prefix in ("dashscope", "qwencloud", "qwen_ai_platform"):
|
||||
monkeypatch.delitem(litellm.model_cost, f"{prefix}/qwen3.7-text-rerank", raising=False)
|
||||
route = respx_mock.post("https://proxy.example/api/v1/services/rerank/text-rerank/text-rerank")
|
||||
route.respond(200, json={"request_id": "bundled", "output": {"results": [{"index": 0, "relevance_score": 0.9}]}})
|
||||
|
||||
response = litellm.rerank(
|
||||
model=f"{provider}/qwen3.7-text-rerank",
|
||||
query="question",
|
||||
documents=["answer"],
|
||||
api_key="fake-key",
|
||||
api_base="https://proxy.example/api/v1",
|
||||
)
|
||||
|
||||
assert json.loads(route.calls[0].request.content)["input"] == {"query": "question", "documents": ["answer"]}
|
||||
assert response.id == "bundled"
|
||||
body = json.loads(route.calls[0].request.content)
|
||||
assert body == {
|
||||
"model": model,
|
||||
"query": "question",
|
||||
"documents": ["answer"],
|
||||
"top_n": 1,
|
||||
"return_documents": False,
|
||||
**({"instruct": instruction} if instruction is not None else {}),
|
||||
}
|
||||
assert response.id == "ranking"
|
||||
assert response.results == [{"index": 0, "relevance_score": 0.9}]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue