From 2d29dfc6b269073f975759905fc85fe408669d17 Mon Sep 17 00:00:00 2001 From: Vincent Date: Wed, 13 May 2026 11:29:08 +0800 Subject: [PATCH] feat(dashscope): add embeddings and reranks(qwen3-rerank) support via OpenAI-compatible endpoint (#27508) Squash-merged by litellm-agent from yimao's PR. --- README.md | 2 +- litellm/__init__.py | 6 + litellm/llms/dashscope/common_utils.py | 28 ++ litellm/llms/dashscope/embed/__init__.py | 7 + .../llms/dashscope/embed/transformation.py | 191 +++++++++++ litellm/llms/dashscope/rerank/__init__.py | 7 + .../llms/dashscope/rerank/transformation.py | 238 +++++++++++++ litellm/main.py | 27 ++ litellm/utils.py | 12 + ...test_dashscope_embedding_transformation.py | 141 ++++++++ .../test_dashscope_rerank_transformation.py | 322 ++++++++++++++++++ 11 files changed, 980 insertions(+), 1 deletion(-) create mode 100644 litellm/llms/dashscope/common_utils.py create mode 100644 litellm/llms/dashscope/embed/__init__.py create mode 100644 litellm/llms/dashscope/embed/transformation.py create mode 100644 litellm/llms/dashscope/rerank/__init__.py create mode 100644 litellm/llms/dashscope/rerank/transformation.py create mode 100644 tests/test_litellm/llms/dashscope/test_dashscope_embedding_transformation.py create mode 100644 tests/test_litellm/llms/dashscope/test_dashscope_rerank_transformation.py diff --git a/README.md b/README.md index 72fd43925c9..8df351e9303 100644 --- a/README.md +++ b/README.md @@ -292,7 +292,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \ | [CompactifAI (`compactifai`)](https://docs.litellm.ai/docs/providers/compactifai) | ✅ | ✅ | ✅ | | | | | | | | | [Custom (`custom`)](https://docs.litellm.ai/docs/providers/custom_llm_server) | ✅ | ✅ | ✅ | | | | | | | | | [Custom OpenAI (`custom_openai`)](https://docs.litellm.ai/docs/providers/openai_compatible) | ✅ | ✅ | ✅ | | | ✅ | ✅ | ✅ | ✅ | | -| [Dashscope (`dashscope`)](https://docs.litellm.ai/docs/providers/dashscope) | ✅ | ✅ | ✅ | | | | | | | | +| [Dashscope (`dashscope`)](https://docs.litellm.ai/docs/providers/dashscope) | ✅ | ✅ | ✅ | ✅ | | | | | | ✅ | | [Databricks (`databricks`)](https://docs.litellm.ai/docs/providers/databricks) | ✅ | ✅ | ✅ | | | | | | | | | [DataRobot (`datarobot`)](https://docs.litellm.ai/docs/providers/datarobot) | ✅ | ✅ | ✅ | | | | | | | | | [Deepgram (`deepgram`)](https://docs.litellm.ai/docs/providers/deepgram) | ✅ | ✅ | ✅ | | | ✅ | | | | | diff --git a/litellm/__init__.py b/litellm/__init__.py index e1b367fb234..b9da0524095 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1880,6 +1880,12 @@ if TYPE_CHECKING: from .llms.dashscope.chat.transformation import ( DashScopeChatConfig as DashScopeChatConfig, ) + from .llms.dashscope.embed.transformation import ( + DashScopeEmbeddingConfig as DashScopeEmbeddingConfig, + ) + from .llms.dashscope.rerank.transformation import ( + DashScopeRerankConfig as DashScopeRerankConfig, + ) from .llms.moonshot.chat.transformation import ( MoonshotChatConfig as MoonshotChatConfig, ) diff --git a/litellm/llms/dashscope/common_utils.py b/litellm/llms/dashscope/common_utils.py new file mode 100644 index 00000000000..b3b89cbbebf --- /dev/null +++ b/litellm/llms/dashscope/common_utils.py @@ -0,0 +1,28 @@ +""" +Common utilities for the DashScope LLM provider. +""" + +from typing import Optional + +import httpx + +from litellm.llms.base_llm.chat.transformation import BaseLLMException + + +class DashScopeError(BaseLLMException): + """Exception class for DashScope provider errors.""" + + def __init__( + self, + status_code: int, + message: str, + headers: Optional[httpx.Headers] = None, + ): + self.status_code = status_code + self.message = message + self.headers = headers or httpx.Headers() + super().__init__( + status_code=status_code, + message=message, + headers=dict(self.headers), + ) diff --git a/litellm/llms/dashscope/embed/__init__.py b/litellm/llms/dashscope/embed/__init__.py new file mode 100644 index 00000000000..4962b1f3251 --- /dev/null +++ b/litellm/llms/dashscope/embed/__init__.py @@ -0,0 +1,7 @@ +""" +DashScope Embedding Module +""" + +from .transformation import DashScopeEmbeddingConfig + +__all__ = ["DashScopeEmbeddingConfig"] diff --git a/litellm/llms/dashscope/embed/transformation.py b/litellm/llms/dashscope/embed/transformation.py new file mode 100644 index 00000000000..5bc0e5ca817 --- /dev/null +++ b/litellm/llms/dashscope/embed/transformation.py @@ -0,0 +1,191 @@ +""" +Transformation logic from OpenAI /v1/embeddings format to DashScope's /v1/embeddings format. + +Supports +- text-embedding-v4 +- text-embedding-v3 + +Endpoint +- https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings + +Docs - https://help.aliyun.com/zh/model-studio/text-embedding-synchronous-api +""" + +from typing import List, Optional, Union + +import httpx + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues +from litellm.types.utils import EmbeddingResponse, Usage + +from ..common_utils import DashScopeError + +DEFAULT_API_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1" + + +class DashScopeEmbeddingConfig(BaseEmbeddingConfig): + """ + Reference: https://help.aliyun.com/zh/model-studio/text-embedding-synchronous-api + + DashScope exposes an OpenAI-compatible /v1/embeddings endpoint, so the + request and response shapes are nearly identical to OpenAI's. + """ + + def __init__(self) -> None: + pass + + def get_supported_openai_params(self, model: str) -> List[str]: + # DashScope's compatible-mode embeddings API accepts the same params as OpenAI. + # `dimensions` / `encoding_format` are only honored by text-embedding-v3 / v4; + # earlier versions silently ignore them server-side. + return ["dimensions", "encoding_format", "user"] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool = False, + ) -> dict: + supported = self.get_supported_openai_params(model) + for k, v in non_default_params.items(): + if v is None: + continue + if k in supported: + optional_params[k] = v + # unsupported params are dropped when drop_params=True; + # the upstream _check_valid_arg already raised UnsupportedParamsError + # for drop_params=False before this method is called. + return optional_params + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + if api_key is None: + api_key = get_secret_str("DASHSCOPE_API_KEY") + if api_key is None: + raise ValueError( + "DashScope API key is required. Set 'DASHSCOPE_API_KEY' env var or pass api_key explicitly." + ) + default_headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {api_key}", + } + return {**default_headers, **headers} + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + base = api_base or get_secret_str("DASHSCOPE_API_BASE") or DEFAULT_API_BASE + base = base.rstrip("/") + if base.endswith("/embeddings"): + return base + return f"{base}/embeddings" + + def transform_embedding_request( + self, + model: str, + input: AllEmbeddingInputValues, + optional_params: dict, + headers: dict, + ) -> dict: + data: dict = { + "model": model, + "input": input, + } + for key in ("dimensions", "encoding_format", "user"): + value = optional_params.get(key) + if value is not None: + data[key] = value + return data + + def transform_embedding_response( + self, + model: str, + raw_response: httpx.Response, + model_response: EmbeddingResponse, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str], + request_data: dict, + optional_params: dict, + litellm_params: dict, + ) -> EmbeddingResponse: + try: + response_json = raw_response.json() + except Exception as e: + raise DashScopeError( + status_code=raw_response.status_code, + message=f"Failed to parse DashScope response as JSON: {str(e)}", + ) + + logging_obj.post_call( + input=request_data.get("input"), + api_key=api_key, + additional_args={"complete_input_dict": request_data}, + original_response=response_json, + ) + + if "error" in response_json: + error = response_json["error"] + message = ( + error.get("message", str(error)) + if isinstance(error, dict) + else str(error) + ) + raise DashScopeError( + status_code=raw_response.status_code, + message=message, + ) + + model_response.object = "list" + model_response.data = response_json.get("data", []) + model_response.model = response_json.get("model", model) + + usage = response_json.get("usage") or {} + prompt_tokens = usage.get("prompt_tokens", 0) + total_tokens = usage.get("total_tokens", prompt_tokens) + setattr( + model_response, + "usage", + Usage( + prompt_tokens=prompt_tokens, + completion_tokens=0, + total_tokens=total_tokens, + ), + ) + + if "id" in response_json: + setattr(model_response, "id", response_json["id"]) + + return model_response + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: Union[dict, httpx.Headers], + ) -> BaseLLMException: + if isinstance(headers, dict): + headers = httpx.Headers(headers) + return DashScopeError( + status_code=status_code, + message=error_message, + headers=headers, + ) diff --git a/litellm/llms/dashscope/rerank/__init__.py b/litellm/llms/dashscope/rerank/__init__.py new file mode 100644 index 00000000000..2a1401f6dc0 --- /dev/null +++ b/litellm/llms/dashscope/rerank/__init__.py @@ -0,0 +1,7 @@ +""" +DashScope Rerank Module +""" + +from .transformation import DashScopeRerankConfig + +__all__ = ["DashScopeRerankConfig"] diff --git a/litellm/llms/dashscope/rerank/transformation.py b/litellm/llms/dashscope/rerank/transformation.py new file mode 100644 index 00000000000..faa1688b5a1 --- /dev/null +++ b/litellm/llms/dashscope/rerank/transformation.py @@ -0,0 +1,238 @@ +""" +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.) + +Endpoint +- https://dashscope.aliyuncs.com/compatible-api/v1/reranks + +Note: chat/embed live under `/compatible-mode/v1/`, but DashScope's rerank +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. + +Empirically, qwen3-rerank accepts `return_documents=true` and echoes +`results[].document.text` back, even though the public docs list the flag +as supported only for gte-rerank-v2 / qwen3-vl-rerank. + +Docs - https://help.aliyun.com/zh/model-studio/text-rerank-api +""" + +from typing import Any, Dict, List, Optional, Union + +import httpx + +from litellm._uuid import uuid +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +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, + RerankTokens, +) + +from ..common_utils import DashScopeError + +DEFAULT_RERANK_URL = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks" + + +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, + results[].relevance_score, optionally results[].document.text (when + return_documents=true), plus a top-level usage.total_tokens counter. + """ + + def __init__(self) -> None: + pass + + def get_complete_url( + self, + api_base: Optional[str], + model: str, + optional_params: Optional[dict] = None, + ) -> str: + if api_base is None: + api_base = get_secret_str("DASHSCOPE_API_BASE_RERANK") or DEFAULT_RERANK_URL + + if api_base == DEFAULT_RERANK_URL: + return DEFAULT_RERANK_URL + + cleaned = api_base.rstrip("/") + if cleaned.endswith("/reranks") or cleaned.endswith("/rerank"): + return cleaned + + if cleaned.endswith("/v1"): + return f"{cleaned}/reranks" + + # Unknown base: append /reranks rather than silently ignoring the caller's api_base. + return f"{cleaned}/reranks" + + def validate_environment( + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + optional_params: Optional[dict] = None, + ) -> dict: + if api_key is None: + api_key = get_secret_str("DASHSCOPE_API_KEY") + if api_key is None: + raise ValueError( + "DashScope API key is required. Set 'DASHSCOPE_API_KEY' env var or pass api_key explicitly." + ) + + default_headers = { + "Authorization": f"Bearer {api_key}", + "accept": "application/json", + "content-type": "application/json", + } + return {**default_headers, **headers} + + def get_supported_cohere_rerank_params(self, model: str) -> list: + return ["query", "documents", "top_n", "return_documents"] + + def map_cohere_rerank_params( + self, + non_default_params: Optional[dict], + model: str, + drop_params: bool, + query: str, + documents: List[Union[str, Dict[str, Any]]], + custom_llm_provider: Optional[str] = None, + top_n: Optional[int] = None, + rank_fields: Optional[List[str]] = None, + return_documents: Optional[bool] = True, + max_chunks_per_doc: Optional[int] = None, + max_tokens_per_doc: Optional[int] = None, + ) -> Dict: + # qwen3-rerank accepts query/documents/top_n/return_documents. The + # rest (rank_fields, max_*_per_doc) are silently dropped. + params: OptionalRerankParams = OptionalRerankParams( + query=query, + documents=documents, + ) + 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) + + def transform_rerank_request( + self, + model: str, + optional_rerank_params: Dict, + headers: dict, + litellm_params: Optional[dict] = None, + ) -> dict: + if "query" not in optional_rerank_params: + raise ValueError("query is required for DashScope rerank") + if "documents" not in optional_rerank_params: + raise ValueError("documents is required for DashScope rerank") + + request: Dict[str, Any] = { + "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"] + 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, + model: str, + raw_response: httpx.Response, + model_response: RerankResponse, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str] = None, + request_data: dict = {}, + optional_params: dict = {}, + litellm_params: dict = {}, + ) -> RerankResponse: + try: + response_json = raw_response.json() + except Exception: + raise DashScopeError( + status_code=raw_response.status_code, + message=raw_response.text, + ) + + logging_obj.post_call( + input=request_data.get("query"), + api_key=api_key, + additional_args={"complete_input_dict": request_data}, + original_response=response_json, + ) + + # DashScope error envelope: {"code": "...", "message": "...", "request_id": "..."} + 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)), + ) + + results = 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}", + ) + + # 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: List[dict] = [] + for r in results: + item: Dict[str, Any] = { + "index": r["index"], + "relevance_score": r["relevance_score"], + } + doc = r.get("document") + if isinstance(doc, dict): + item["document"] = doc + elif isinstance(doc, str): + # Defensive: spec says dict, but normalize string-shaped echoes. + item["document"] = {"text": doc} + transformed_results.append(item) + + usage = response_json.get("usage") or {} + total_tokens = usage.get("total_tokens") + billed_units = RerankBilledUnits(total_tokens=total_tokens) + tokens = RerankTokens(input_tokens=total_tokens) + meta = RerankResponseMeta(billed_units=billed_units, tokens=tokens) + + return RerankResponse( + id=response_json.get("id") or str(uuid.uuid4()), + results=transformed_results, # type: ignore + meta=meta, + ) + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: Union[dict, httpx.Headers], + ) -> BaseLLMException: + if isinstance(headers, dict): + headers = httpx.Headers(headers) + return DashScopeError( + status_code=status_code, + message=error_message, + headers=headers, + ) diff --git a/litellm/main.py b/litellm/main.py index c3d1c2e05b0..b5364f8ba17 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -5720,6 +5720,33 @@ def embedding( # noqa: PLR0915 aembedding=aembedding, headers=headers, ) + elif custom_llm_provider == "dashscope": + dashscope_key = ( + api_key or litellm.api_key or get_secret_str("DASHSCOPE_API_KEY") + ) + if dashscope_key is None: + raise ValueError( + "Missing API key for DashScope. Set DASHSCOPE_API_KEY environment variable or pass api_key parameter." + ) + if extra_headers is not None and isinstance(extra_headers, dict): + headers = extra_headers + else: + headers = {} + response = base_llm_http_handler.embedding( + model=model, + input=input, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + logging_obj=logging, + api_base=api_base, + optional_params=optional_params, + litellm_params={}, + model_response=EmbeddingResponse(), + api_key=dashscope_key, + client=client, + aembedding=aembedding, + headers=headers, + ) elif custom_llm_provider == "ovhcloud": api_key = api_key or litellm.api_key or get_secret_str("OVHCLOUD_API_KEY") api_base = ( diff --git a/litellm/utils.py b/litellm/utils.py index da80e4ae164..891b4ee9fa3 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8374,6 +8374,12 @@ class ProviderConfigManager: ) return VolcEngineEmbeddingConfig() + elif litellm.LlmProviders.DASHSCOPE == provider: + from litellm.llms.dashscope.embed.transformation import ( + DashScopeEmbeddingConfig, + ) + + return DashScopeEmbeddingConfig() elif litellm.LlmProviders.OVHCLOUD == provider: return litellm.OVHCloudEmbeddingConfig() elif litellm.LlmProviders.SNOWFLAKE == provider: @@ -8453,6 +8459,12 @@ 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() return litellm.CohereRerankConfig() @staticmethod diff --git a/tests/test_litellm/llms/dashscope/test_dashscope_embedding_transformation.py b/tests/test_litellm/llms/dashscope/test_dashscope_embedding_transformation.py new file mode 100644 index 00000000000..5e4d0177e8d --- /dev/null +++ b/tests/test_litellm/llms/dashscope/test_dashscope_embedding_transformation.py @@ -0,0 +1,141 @@ +""" +Unit tests for DashScope embedding transformation. +""" + +import json +import os +import sys +from unittest.mock import MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.dashscope.common_utils import DashScopeError +from litellm.llms.dashscope.embed.transformation import ( + DEFAULT_API_BASE, + DashScopeEmbeddingConfig, +) +from litellm.types.utils import EmbeddingResponse + + +def test_validate_environment_and_url(): + config = DashScopeEmbeddingConfig() + headers = config.validate_environment( + headers={}, + model="text-embedding-v4", + messages=[], + optional_params={}, + litellm_params={}, + api_key="sk-test", + ) + assert headers["Authorization"] == "Bearer sk-test" + + url = config.get_complete_url( + api_base=None, + api_key="sk-test", + model="text-embedding-v4", + optional_params={}, + litellm_params={}, + ) + assert url == f"{DEFAULT_API_BASE}/embeddings" + + +def test_transform_embedding_request(): + config = DashScopeEmbeddingConfig() + data = config.transform_embedding_request( + model="text-embedding-v4", + input=["风急天高猿啸哀"], + optional_params={"dimensions": 1024, "encoding_format": "float"}, + headers={}, + ) + assert data == { + "model": "text-embedding-v4", + "input": ["风急天高猿啸哀"], + "dimensions": 1024, + "encoding_format": "float", + } + + +def test_transform_embedding_response_success(): + config = DashScopeEmbeddingConfig() + payload = { + "data": [ + {"embedding": [0.1, 0.2], "index": 0, "object": "embedding"}, + ], + "model": "text-embedding-v4", + "object": "list", + "usage": {"prompt_tokens": 5, "total_tokens": 5}, + "id": "73591b79-xxxx", + } + raw = httpx.Response( + status_code=200, + content=json.dumps(payload).encode("utf-8"), + request=httpx.Request("POST", "https://example.com"), + ) + result = config.transform_embedding_response( + model="text-embedding-v4", + raw_response=raw, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="sk-x", + request_data={"input": ["a"]}, + optional_params={}, + litellm_params={}, + ) + assert result.model == "text-embedding-v4" + assert len(result.data) == 1 + assert result.usage.prompt_tokens == 5 + + +def test_transform_embedding_request_user_param(): + config = DashScopeEmbeddingConfig() + data = config.transform_embedding_request( + model="text-embedding-v4", + input=["hello"], + optional_params={"user": "user-123"}, + headers={}, + ) + assert data["user"] == "user-123" + + +def test_map_openai_params_drops_unsupported_with_drop_params(): + config = DashScopeEmbeddingConfig() + result = config.map_openai_params( + non_default_params={"dimensions": 512, "unknown_param": "value"}, + optional_params={}, + model="text-embedding-v4", + drop_params=True, + ) + assert result == {"dimensions": 512} + assert "unknown_param" not in result + + +def test_transform_embedding_response_error(): + config = DashScopeEmbeddingConfig() + payload = { + "error": { + "message": "Incorrect API key provided.", + "type": "invalid_request_error", + "code": "invalid_api_key", + } + } + raw = httpx.Response( + status_code=401, + content=json.dumps(payload).encode("utf-8"), + request=httpx.Request("POST", "https://example.com"), + ) + with pytest.raises(DashScopeError) as exc: + config.transform_embedding_response( + model="text-embedding-v4", + raw_response=raw, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="sk-bad", + request_data={"input": ["a"]}, + optional_params={}, + litellm_params={}, + ) + assert exc.value.status_code == 401 + assert "Incorrect API key" in exc.value.message diff --git a/tests/test_litellm/llms/dashscope/test_dashscope_rerank_transformation.py b/tests/test_litellm/llms/dashscope/test_dashscope_rerank_transformation.py new file mode 100644 index 00000000000..26e3881f83c --- /dev/null +++ b/tests/test_litellm/llms/dashscope/test_dashscope_rerank_transformation.py @@ -0,0 +1,322 @@ +""" +Unit tests for DashScope rerank transformation. +""" + +import json +import os +import sys +from unittest.mock import MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.dashscope.common_utils import DashScopeError +from litellm.llms.dashscope.rerank.transformation import ( + DEFAULT_RERANK_URL, + DashScopeRerankConfig, +) +from litellm.types.rerank import RerankResponse + + +class TestDashScopeRerankURL: + def setup_method(self): + self.config = DashScopeRerankConfig() + + def test_default_url(self): + url = self.config.get_complete_url(api_base=None, model="qwen3-rerank") + assert url == DEFAULT_RERANK_URL + + def test_explicit_v1_base_appends_reranks(self): + url = self.config.get_complete_url( + api_base="https://dashscope.aliyuncs.com/compatible-mode/v1", + model="qwen3-rerank", + ) + assert url == "https://dashscope.aliyuncs.com/compatible-mode/v1/reranks" + + def test_intl_v1_base_appends_reranks(self): + url = self.config.get_complete_url( + api_base="https://dashscope-intl.aliyuncs.com/compatible-mode/v1", + model="qwen3-rerank", + ) + assert url == "https://dashscope-intl.aliyuncs.com/compatible-mode/v1/reranks" + + def test_already_complete_url_passthrough(self): + full = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks" + assert self.config.get_complete_url(api_base=full, model="qwen3-rerank") == full + + def test_trailing_slash_stripped(self): + full = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks/" + assert self.config.get_complete_url( + api_base=full, model="qwen3-rerank" + ) == full.rstrip("/") + + def test_custom_v1_base_appends_reranks(self): + url = self.config.get_complete_url( + api_base="https://my-proxy.example.com/v1", model="qwen3-rerank" + ) + assert url == "https://my-proxy.example.com/v1/reranks" + + +class TestDashScopeRerankRequest: + def setup_method(self): + self.config = DashScopeRerankConfig() + + def test_validate_environment_with_explicit_key(self): + headers = self.config.validate_environment( + headers={}, model="qwen3-rerank", api_key="sk-test" + ) + assert headers["Authorization"] == "Bearer sk-test" + assert headers["content-type"] == "application/json" + + def test_validate_environment_missing_key(self, monkeypatch): + monkeypatch.delenv("DASHSCOPE_API_KEY", raising=False) + with pytest.raises(ValueError, match="DASHSCOPE_API_KEY"): + self.config.validate_environment( + headers={}, model="qwen3-rerank", api_key=None + ) + + def test_validate_environment_falls_back_to_env(self, monkeypatch): + monkeypatch.setenv("DASHSCOPE_API_KEY", "env-key") + headers = self.config.validate_environment( + headers={}, model="qwen3-rerank", api_key=None + ) + assert headers["Authorization"] == "Bearer env-key" + + def test_supported_params(self): + assert self.config.get_supported_cohere_rerank_params("qwen3-rerank") == [ + "query", + "documents", + "top_n", + "return_documents", + ] + + def test_map_params_drops_unsupported(self): + # qwen3-rerank accepts query/documents/top_n/return_documents. + # rank_fields and max_*_per_doc are silently dropped. + params = self.config.map_cohere_rerank_params( + non_default_params={}, + model="qwen3-rerank", + drop_params=False, + query="什么是文本排序模型", + documents=["d1", "d2"], + top_n=2, + rank_fields=["title"], + return_documents=True, + max_chunks_per_doc=5, + max_tokens_per_doc=100, + ) + assert params == { + "query": "什么是文本排序模型", + "documents": ["d1", "d2"], + "top_n": 2, + "return_documents": True, + } + + def test_transform_request_full(self): + body = self.config.transform_rerank_request( + model="qwen3-rerank", + optional_rerank_params={ + "query": "如何制作美味的苹果派?", + "documents": ["a", "b"], + "top_n": 5, + "return_documents": True, + }, + headers={}, + ) + assert body == { + "model": "qwen3-rerank", + "query": "如何制作美味的苹果派?", + "documents": ["a", "b"], + "top_n": 5, + "return_documents": True, + } + + def test_transform_request_omits_unset_optional(self): + body = self.config.transform_rerank_request( + model="qwen3-rerank", + optional_rerank_params={"query": "q", "documents": ["a"]}, + headers={}, + ) + assert "top_n" not in body + assert "return_documents" not in body + + def test_transform_request_requires_query(self): + with pytest.raises(ValueError, match="query"): + self.config.transform_rerank_request( + model="qwen3-rerank", + optional_rerank_params={"documents": ["a"]}, + headers={}, + ) + + def test_transform_request_requires_documents(self): + with pytest.raises(ValueError, match="documents"): + self.config.transform_rerank_request( + model="qwen3-rerank", + optional_rerank_params={"query": "q"}, + headers={}, + ) + + +class TestDashScopeRerankResponse: + def setup_method(self): + self.config = DashScopeRerankConfig() + self.logging = MagicMock() + + def _resp(self, body, status_code=200): + return httpx.Response( + status_code=status_code, content=json.dumps(body).encode() + ) + + def test_success_response(self): + body = { + "object": "list", + "results": [ + {"index": 0, "relevance_score": 0.93}, + {"index": 2, "relevance_score": 0.34}, + ], + "model": "qwen3-rerank", + "id": "85ba5752", + "usage": {"total_tokens": 79}, + } + out = self.config.transform_rerank_response( + model="qwen3-rerank", + raw_response=self._resp(body), + model_response=RerankResponse(), + logging_obj=self.logging, + api_key="sk", + request_data={"query": "q"}, + ) + assert out.id == "85ba5752" + assert out.results == [ + {"index": 0, "relevance_score": 0.93}, + {"index": 2, "relevance_score": 0.34}, + ] + assert out.meta == { + "billed_units": {"total_tokens": 79}, + "tokens": {"input_tokens": 79}, + } + + def test_response_with_return_documents_real_payload(self): + # Verbatim sample from a real qwen3-rerank call with return_documents=true. + body = { + "object": "list", + "results": [ + { + "document": { + "text": "苹果派的制作步骤包括准备面团、切苹果、调制馅料、组装和烘烤。" + }, + "index": 1, + "relevance_score": 0.8304247466067356, + }, + { + "document": { + "text": "制作苹果派时,预先煮软苹果可以缩短烘烤时间。" + }, + "index": 3, + "relevance_score": 0.7142660211908354, + }, + ], + "model": "qwen3-rerank", + "id": "e191b077-97c4-9929-b121-c2fbd2c7b0af", + "usage": {"total_tokens": 192}, + } + out = self.config.transform_rerank_response( + model="qwen3-rerank", + raw_response=self._resp(body), + model_response=RerankResponse(), + logging_obj=self.logging, + request_data={"query": "如何制作美味的苹果派?"}, + ) + assert out.id == "e191b077-97c4-9929-b121-c2fbd2c7b0af" + assert out.results == [ + { + "index": 1, + "relevance_score": 0.8304247466067356, + "document": { + "text": "苹果派的制作步骤包括准备面团、切苹果、调制馅料、组装和烘烤。" + }, + }, + { + "index": 3, + "relevance_score": 0.7142660211908354, + "document": {"text": "制作苹果派时,预先煮软苹果可以缩短烘烤时间。"}, + }, + ] + assert out.meta == { + "billed_units": {"total_tokens": 192}, + "tokens": {"input_tokens": 192}, + } + + def test_response_string_document_normalized(self): + # Defensive path: if a future API revision returns a bare string, + # normalize to {"text": ...} so downstream code stays consistent. + body = { + "results": [{"index": 0, "relevance_score": 0.9, "document": "hello"}], + "model": "qwen3-rerank", + "usage": {"total_tokens": 5}, + } + out = self.config.transform_rerank_response( + model="qwen3-rerank", + raw_response=self._resp(body), + model_response=RerankResponse(), + logging_obj=self.logging, + ) + assert out.results[0]["document"] == {"text": "hello"} + + def test_missing_id_generates_uuid(self): + body = {"results": [{"index": 0, "relevance_score": 0.5}], "usage": {}} + out = self.config.transform_rerank_response( + model="qwen3-rerank", + raw_response=self._resp(body), + model_response=RerankResponse(), + logging_obj=self.logging, + ) + assert out.id is not None and len(out.id) > 0 + + def test_error_envelope_raises(self): + 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", + 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): + bad = httpx.Response(status_code=500, content=b"bad gateway") + with pytest.raises(DashScopeError): + self.config.transform_rerank_response( + model="qwen3-rerank", + raw_response=bad, + model_response=RerankResponse(), + logging_obj=self.logging, + ) + + def test_get_error_class(self): + err = self.config.get_error_class( + error_message="boom", status_code=500, headers={} + ) + assert isinstance(err, DashScopeError) + assert err.status_code == 500 + + +class TestProviderConfigManagerDispatch: + def test_dashscope_returns_rerank_config(self): + import litellm + from litellm.utils import ProviderConfigManager + + cfg = ProviderConfigManager.get_provider_rerank_config( + model="qwen3-rerank", + provider=litellm.LlmProviders.DASHSCOPE, + api_base=None, + present_version_params=[], + ) + assert isinstance(cfg, DashScopeRerankConfig)