diff --git a/litellm/__init__.py b/litellm/__init__.py index 9e52fdaaf35..b0761da7f7c 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1647,6 +1647,9 @@ if TYPE_CHECKING: from .llms.jina_ai.rerank.transformation import ( JinaAIRerankConfig as JinaAIRerankConfig, ) + from .llms.scaleway.rerank.transformation import ( + ScalewayRerankConfig as ScalewayRerankConfig, + ) from .llms.deepinfra.rerank.transformation import ( DeepinfraRerankConfig as DeepinfraRerankConfig, ) diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index ffcbffb05b6..fcd2eed5387 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -151,6 +151,7 @@ LLM_CONFIG_NAMES: Final = ( "AzureAIRerankConfig", "InfinityRerankConfig", "JinaAIRerankConfig", + "ScalewayRerankConfig", "DeepinfraRerankConfig", "HostedVLLMRerankConfig", "NvidiaNimRerankConfig", @@ -688,6 +689,7 @@ _LLM_CONFIGS_IMPORT_MAP: Final = { "InfinityRerankConfig", ), "JinaAIRerankConfig": (".llms.jina_ai.rerank.transformation", "JinaAIRerankConfig"), + "ScalewayRerankConfig": (".llms.scaleway.rerank.transformation", "ScalewayRerankConfig"), "DeepinfraRerankConfig": ( ".llms.deepinfra.rerank.transformation", "DeepinfraRerankConfig", diff --git a/litellm/llms/scaleway/rerank/transformation.py b/litellm/llms/scaleway/rerank/transformation.py new file mode 100644 index 00000000000..31f273bdc43 --- /dev/null +++ b/litellm/llms/scaleway/rerank/transformation.py @@ -0,0 +1,51 @@ +""" +Support for Scaleway's `/v1/rerank` endpoint. + +The request and response match Jina AI's, so this reuses that config. + +API reference: https://www.scaleway.com/en/developers/api/generative-apis/#path-rerank-create-a-reranking +""" + +from collections.abc import Mapping +from typing import Final + +from litellm.llms.jina_ai.rerank.transformation import JinaAIRerankConfig +from litellm.secret_managers.main import get_secret_str + +SCALEWAY_API_BASE: Final = "https://api.scaleway.ai/v1" + + +class ScalewayRerankConfig(JinaAIRerankConfig): + def get_supported_cohere_rerank_params(self, model: str) -> list[str]: # mutable-ok: BaseRerankConfig contract + return ["query", "top_n", "documents"] + + def get_complete_url( + self, + api_base: str | None, + model: str, + optional_params: Mapping[str, object] | None = None, + ) -> str: + base: Final = SCALEWAY_API_BASE if api_base is None else api_base.rstrip("/") + return f"{base}/rerank" + + def validate_environment( + self, + headers: Mapping[str, str], + model: str, + api_key: str | None = None, + optional_params: Mapping[str, object] | None = None, + litellm_params: Mapping[str, object] | None = None, + ) -> dict[str, str]: # mutable-ok: BaseRerankConfig contract + key: Final = api_key or get_secret_str("SCW_SECRET_KEY") + if not key: + raise ValueError( + "Scaleway API key not found. Pass `api_key=...` or set the SCW_SECRET_KEY environment variable." + ) + provider_headers: Final = { + "accept": "application/json", + "content-type": "application/json", + "authorization": f"Bearer {key}", + } + # Header names are case-insensitive, so match on the lowercase name. + caller_headers: Final = {name: value for name, value in headers.items() if name.lower() not in provider_headers} + return {**caller_headers, **provider_headers} diff --git a/litellm/utils.py b/litellm/utils.py index fed8459bfac..26f412d3b75 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8861,8 +8861,13 @@ class ProviderConfigManager: return litellm.AzureAIRerankConfig() elif litellm.LlmProviders.INFINITY == provider: return litellm.InfinityRerankConfig() - elif litellm.LlmProviders.JINA_AI == provider: - return litellm.JinaAIRerankConfig() + elif provider in (litellm.LlmProviders.JINA_AI, litellm.LlmProviders.SCALEWAY): + # Scaleway's rerank API matches Jina's, so its config extends Jina's. + return ( + litellm.ScalewayRerankConfig() + if provider == litellm.LlmProviders.SCALEWAY + else litellm.JinaAIRerankConfig() + ) elif litellm.LlmProviders.HOSTED_VLLM == provider: return litellm.HostedVLLMRerankConfig() elif litellm.LlmProviders.HUGGINGFACE == provider: diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 223711a92b6..eb27d3fe810 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2397,7 +2397,7 @@ "audio_speech": false, "moderations": false, "batches": false, - "rerank": false, + "rerank": true, "a2a": true, "interactions": true } diff --git a/tests/unit/llms/scaleway/test_scaleway_rerank_transformation.py b/tests/unit/llms/scaleway/test_scaleway_rerank_transformation.py new file mode 100644 index 00000000000..dd448a048c6 --- /dev/null +++ b/tests/unit/llms/scaleway/test_scaleway_rerank_transformation.py @@ -0,0 +1,136 @@ +import json +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest +import respx + +import litellm +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + +SCALEWAY_RERANK_BODY = { + "id": "rerank-a89e6d7b8b97492ea81569c65fbfff49", + "model": "qwen3-embedding-8b", + "usage": {"total_tokens": 99}, + "results": [ + { + "index": 1, + "document": {"text": "Oceans can be sorted by size: Pacific, Atlantic, Indian", "multi_modal": None}, + "relevance_score": 0.6456239223480225, + }, + { + "index": 0, + "document": {"text": "The Pacific is approximately 165 million km²", "multi_modal": None}, + "relevance_score": 0.6059925556182861, + }, + ], +} + +DOCUMENTS = ["The Pacific is approximately 165 million km²", "Oceans can be sorted by size: Pacific, Atlantic, Indian"] + + +def test_scaleway_rerank_posts_to_the_documented_endpoint(respx_mock: respx.MockRouter, monkeypatch): + monkeypatch.delenv("SCALEWAY_API_BASE", raising=False) + route = respx_mock.post("https://api.scaleway.ai/v1/rerank") + route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY) + + response = litellm.rerank( + model="scaleway/qwen3-embedding-8b", + query="What is the biggest area of water on earth ?", + documents=DOCUMENTS, + top_n=2, + api_key="scw-key", + ) + + request = route.calls[0].request + assert request.headers["authorization"] == "Bearer scw-key" + assert json.loads(request.content) == { + "model": "qwen3-embedding-8b", + "query": "What is the biggest area of water on earth ?", + "documents": DOCUMENTS, + "top_n": 2, + } + assert [r["index"] for r in response.results] == [1, 0] + assert response.results[0]["relevance_score"] == pytest.approx(0.6456239223480225) + assert response.results[0]["document"]["text"].startswith("Oceans") + assert response.id == SCALEWAY_RERANK_BODY["id"] + assert response.meta["billed_units"]["total_tokens"] == 99 + + +def test_scaleway_rerank_reads_the_key_from_scw_secret_key(respx_mock: respx.MockRouter, monkeypatch): + monkeypatch.setenv("SCW_SECRET_KEY", "env-scw-key") + route = respx_mock.post("https://api.scaleway.ai/v1/rerank") + route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY) + + litellm.rerank(model="scaleway/qwen3-embedding-8b", query="q", documents=DOCUMENTS) + + assert route.calls[0].request.headers["authorization"] == "Bearer env-scw-key" + + +def test_scaleway_rerank_honors_api_base(respx_mock: respx.MockRouter): + route = respx_mock.post("https://scw.example/v1/rerank") + route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY) + + litellm.rerank( + model="scaleway/qwen3-embedding-8b", + query="q", + documents=DOCUMENTS, + api_key="scw-key", + api_base="https://scw.example/v1/", + ) + + assert route.called + + +def test_scaleway_rerank_does_not_send_return_documents(respx_mock: respx.MockRouter): + """The Scaleway API has no such field, so it must not reach the request body.""" + route = respx_mock.post("https://api.scaleway.ai/v1/rerank") + route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY) + + litellm.rerank( + model="scaleway/qwen3-embedding-8b", + query="q", + documents=DOCUMENTS, + return_documents=True, + api_key="scw-key", + ) + + assert "return_documents" not in json.loads(route.calls[0].request.content) + + +def test_scaleway_rerank_without_a_key_names_the_env_var(monkeypatch): + monkeypatch.delenv("SCW_SECRET_KEY", raising=False) + + with pytest.raises(litellm.APIConnectionError, match="SCW_SECRET_KEY"): + litellm.rerank(model="scaleway/qwen3-embedding-8b", query="q", documents=DOCUMENTS) + + +def test_scaleway_rerank_caller_headers_cannot_replace_the_provider_key(respx_mock: respx.MockRouter): + route = respx_mock.post("https://api.scaleway.ai/v1/rerank") + route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY) + + litellm.rerank( + model="scaleway/qwen3-embedding-8b", + query="q", + documents=DOCUMENTS, + api_key="scw-key", + headers={"Authorization": "Bearer caller-key", "x-trace": "abc"}, + ) + + request = route.calls[0].request + assert request.headers["authorization"] == "Bearer scw-key" + assert request.headers["x-trace"] == "abc" + + +@pytest.mark.asyncio +async def test_scaleway_arerank_posts_to_the_documented_endpoint(): + client = MagicMock(spec=AsyncHTTPHandler) + client.post = AsyncMock(return_value=httpx.Response(200, json=SCALEWAY_RERANK_BODY)) + + response = await litellm.arerank( + model="scaleway/qwen3-embedding-8b", query="q", documents=DOCUMENTS, api_key="scw-key", client=client + ) + + assert client.post.await_args.kwargs["url"] == "https://api.scaleway.ai/v1/rerank" + assert client.post.await_args.kwargs["headers"]["authorization"] == "Bearer scw-key" + assert [r["index"] for r in response.results] == [1, 0]