mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(scaleway): add rerank support (#44160)
* feat(scaleway): add rerank support Fixes #43856 Signed-off-by: Ankit Jha <ankit.jha@tradomate.one> * fix(scaleway): caller headers cannot replace the provider key; inject the async client in tests Signed-off-by: Ankit Jha <ankit.jha@tradomate.one> * refactor(utils): fold Scaleway into the Jina rerank branch to stay inside the complexity budget Signed-off-by: Ankit Jha <ankit.jha@tradomate.one> --------- Signed-off-by: Ankit Jha <ankit.jha@tradomate.one>
This commit is contained in:
parent
119179942e
commit
0d5ea45c86
6 changed files with 200 additions and 3 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
51
litellm/llms/scaleway/rerank/transformation.py
Normal file
51
litellm/llms/scaleway/rerank/transformation.py
Normal file
|
|
@ -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}
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -2397,7 +2397,7 @@
|
|||
"audio_speech": false,
|
||||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": false,
|
||||
"rerank": true,
|
||||
"a2a": true,
|
||||
"interactions": true
|
||||
}
|
||||
|
|
|
|||
136
tests/unit/llms/scaleway/test_scaleway_rerank_transformation.py
Normal file
136
tests/unit/llms/scaleway/test_scaleway_rerank_transformation.py
Normal file
|
|
@ -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]
|
||||
Loading…
Add table
Reference in a new issue