fix(rerank): map provider errors with the resolved provider on sync and async paths

This commit is contained in:
mateo-berri 2026-09-01 12:46:04 -07:00
parent 695d943745
commit d3dab8e294
2 changed files with 80 additions and 3 deletions

View file

@ -43,10 +43,17 @@ async def arerank(
"""
Async: Reranks a list of documents based on their relevance to the query
"""
_custom_llm_provider: str | None = None # rebind-ok: set by the get_llm_provider unpack; read in the except
try:
loop: Final = asyncio.get_event_loop()
kwargs["arerank"] = True
_, _custom_llm_provider, _, _ = litellm.get_llm_provider( # rebind-ok: see pre-declaration above
model=model,
custom_llm_provider=custom_llm_provider,
api_base=kwargs.get("api_base", None),
)
func: Final = partial(
rerank,
model,
@ -70,7 +77,11 @@ async def arerank(
response = init_response
return response
except Exception as e:
raise e
raise exception_type(
model=model,
custom_llm_provider=_custom_llm_provider or custom_llm_provider,
original_exception=e,
)
@client
@ -115,6 +126,7 @@ def rerank(
model_info: Final = kwargs.get("model_info", None)
user: Final = kwargs.get("user", None)
client: Final = kwargs.get("client", None)
_custom_llm_provider: str | None = None # rebind-ok: set by the get_llm_provider unpack; read in the except
try:
_is_async: Final = kwargs.pop("arerank", False) is True
optional_params: Final = GenericLiteLLMParams(**kwargs)
@ -127,7 +139,7 @@ def rerank(
(
model,
_custom_llm_provider,
_custom_llm_provider, # rebind-ok: see pre-declaration above
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
@ -538,4 +550,8 @@ def rerank(
return response
except Exception as e:
verbose_logger.error("Error in rerank: %s", e)
raise exception_type(model=model, custom_llm_provider=custom_llm_provider, original_exception=e)
raise exception_type(
model=model,
custom_llm_provider=_custom_llm_provider or custom_llm_provider,
original_exception=e,
)

View file

@ -111,6 +111,67 @@ def test_together_rerank_honors_api_base(respx_mock: respx.MockRouter):
assert mock_route.calls[0].request.headers["authorization"] == "Bearer fake-together-key"
DASHSCOPE_404_BODY = {
"error": {
"message": "The model `does-not-exist` does not exist or you do not have access to it.",
"type": "invalid_request_error",
"param": None,
"code": "model_not_found",
},
"request_id": "mock-request-id",
}
def test_rerank_error_names_provider_and_keeps_body(respx_mock: respx.MockRouter, monkeypatch):
"""Regression for the rerank error path mapping with the unresolved provider param:
a provider 404 surfaced as 'None - ' instead of naming the provider and its error body."""
monkeypatch.delenv("DASHSCOPE_API_BASE", raising=False)
monkeypatch.delenv("DASHSCOPE_API_BASE_RERANK", raising=False)
mock_route = respx_mock.post("https://dashscope.example/v1/reranks")
mock_route.return_value = httpx.Response(404, json=DASHSCOPE_404_BODY)
with pytest.raises(litellm.NotFoundError) as exc_info:
litellm.rerank(
model="dashscope/does-not-exist",
query=MARKER_QUERY,
documents=[MARKER_DOC],
api_key="fake-dashscope-key",
api_base="https://dashscope.example/v1",
)
assert mock_route.called
assert "DashscopeException" in str(exc_info.value)
assert "does not exist or you do not have access to it" in str(exc_info.value)
assert "None - " not in str(exc_info.value)
@pytest.mark.asyncio
async def test_arerank_error_is_mapped_to_litellm_exception(respx_mock: respx.MockRouter, monkeypatch):
"""Regression for arerank's bare re-raise: provider errors escaped as raw
provider exception classes instead of the mapped litellm exception contract."""
monkeypatch.delenv("DASHSCOPE_API_BASE", raising=False)
monkeypatch.delenv("DASHSCOPE_API_BASE_RERANK", raising=False)
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
mock_route = respx_mock.post("https://dashscope.example/v1/reranks")
mock_route.return_value = httpx.Response(404, json=DASHSCOPE_404_BODY)
with pytest.raises(litellm.NotFoundError) as exc_info:
await litellm.arerank(
model="dashscope/does-not-exist",
query=MARKER_QUERY,
documents=[MARKER_DOC],
api_key="fake-dashscope-key",
api_base="https://dashscope.example/v1",
)
assert mock_route.called
assert "DashscopeException" in str(exc_info.value)
assert "does not exist or you do not have access to it" in str(exc_info.value)
assert "None - " not in str(exc_info.value)
@pytest.mark.asyncio
async def test_together_rerank_async_honors_env_api_base(respx_mock: respx.MockRouter, monkeypatch):
"""Regression: TOGETHER_AI_API_BASE was honored by chat but ignored by rerank."""