mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(router): move retry-policy retries off the refusing deployment on every router entrypoint (#40306)
Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
0f886d9006
commit
8ce4c05019
2 changed files with 250 additions and 1 deletions
|
|
@ -3656,6 +3656,13 @@ class Router:
|
|||
effective_model_info: Final = kwargs.get("model_info") or deployment.get("model_info") or MappingProxyType({})
|
||||
self._set_failed_deployment_id_on_exception(exception, MappingProxyType({"model_info": effective_model_info}))
|
||||
|
||||
@staticmethod
|
||||
def _stamp_retry_skip_deployment_id(exception: Exception, kwargs: Mapping[str, object]) -> None:
|
||||
effective_model_info: Final = kwargs.get("model_info")
|
||||
deployment_id: Final = effective_model_info.get("id") if isinstance(effective_model_info, Mapping) else None
|
||||
if isinstance(deployment_id, str) and deployment_id:
|
||||
exception.retry_skip_deployment_id = deployment_id # pyright: ignore[reportAttributeAccessIssue] # dynamic stamp, read by _deployment_ids_to_skip_on_retry
|
||||
|
||||
def _update_kwargs_with_default_litellm_params(
|
||||
self, kwargs: dict, metadata_variable_name: str | None = "metadata"
|
||||
) -> None:
|
||||
|
|
@ -4358,6 +4365,7 @@ class Router:
|
|||
model=model,
|
||||
messages=[{"role": "user", "content": "prompt"}],
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
|
||||
data: Final = deployment["litellm_params"].copy()
|
||||
|
|
@ -4388,6 +4396,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.image_generation(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aimage_generation(self, prompt: str, model: str, **kwargs):
|
||||
|
|
@ -4472,6 +4481,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.aimage_generation(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def atranscription(self, file: FileTypes, model: str, **kwargs):
|
||||
|
|
@ -4576,6 +4586,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.atranscription(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aspeech(self, model: str, input: str, voice: str | None = None, **kwargs):
|
||||
|
|
@ -4690,6 +4701,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.aspeech(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def arerank(self, model: str, **kwargs):
|
||||
|
|
@ -4748,6 +4760,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.arerank(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
def text_completion(
|
||||
|
|
@ -4882,6 +4895,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.atext_completion(model=%s)\x1b[31m Exception %s\x1b[0m", model, e)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aadapter_completion(
|
||||
|
|
@ -4972,6 +4986,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.aadapter_completion(model=%s)\x1b[31m Exception %s\x1b[0m", model, e)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def _asearch_with_fallbacks(self, original_function: Callable, **kwargs):
|
||||
|
|
@ -5754,6 +5769,7 @@ class Router:
|
|||
model=model,
|
||||
input=input,
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
|
||||
data: Final = deployment["litellm_params"].copy()
|
||||
|
|
@ -5792,6 +5808,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.embedding(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aembedding(
|
||||
|
|
@ -5879,6 +5896,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.aembedding(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
#### FILES API ####
|
||||
|
|
@ -6252,6 +6270,7 @@ class Router:
|
|||
)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aretrieve_batch(
|
||||
|
|
@ -6472,6 +6491,7 @@ class Router:
|
|||
)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def alist_batches(
|
||||
|
|
@ -7589,7 +7609,9 @@ class Router:
|
|||
|
||||
@staticmethod
|
||||
def _deployment_ids_to_skip_on_retry(exception: Exception, already_skipped: object) -> tuple[str, ...]:
|
||||
failed_deployment_id: Final[str | None] = getattr(exception, "failed_deployment_id", None)
|
||||
failed_deployment_id: Final[str | None] = getattr(exception, "retry_skip_deployment_id", None) or getattr(
|
||||
exception, "failed_deployment_id", None
|
||||
)
|
||||
status_code: Final = getattr(exception, "status_code", None)
|
||||
if not failed_deployment_id or not isinstance(status_code, int):
|
||||
return ()
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import logging
|
|||
import os
|
||||
import threading
|
||||
from datetime import datetime
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -27,6 +28,7 @@ from litellm.llms.bedrock.common_utils import BedrockError
|
|||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES,
|
||||
)
|
||||
from litellm.types.llms.openai import ChatCompletionRequest
|
||||
from litellm.router import (
|
||||
MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS,
|
||||
FallbackAwareAnthropicMessagesStream,
|
||||
|
|
@ -13972,6 +13974,31 @@ def test_router_deployment_ids_to_skip_on_retry(status_code, failed_deployment_i
|
|||
assert litellm.Router._deployment_ids_to_skip_on_retry(exception, already_skipped) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"kwargs,failed_deployment_id,expected",
|
||||
[
|
||||
({"model_info": {"id": "rejecting"}}, None, ("rejecting",)),
|
||||
({"model_info": {"id": "rejecting"}}, "cooldown-target", ("rejecting",)),
|
||||
({"model_info": {"id": ""}}, None, ()),
|
||||
({"model_info": {"id": 7}}, None, ()),
|
||||
({"model_info": "rejecting"}, None, ()),
|
||||
({}, None, ()),
|
||||
({}, "cooldown-target", ("cooldown-target",)),
|
||||
],
|
||||
)
|
||||
def test_router_retry_skip_stamp_feeds_deployment_ids_to_skip_on_retry(
|
||||
kwargs: Mapping[str, object], failed_deployment_id: str | None, expected: tuple[str, ...]
|
||||
):
|
||||
exception = Exception("upstream refused this request")
|
||||
exception.status_code = 400
|
||||
exception.failed_deployment_id = failed_deployment_id
|
||||
|
||||
litellm.Router._stamp_retry_skip_deployment_id(exception, kwargs)
|
||||
|
||||
assert litellm.Router._deployment_ids_to_skip_on_retry(exception, None) == expected
|
||||
assert exception.failed_deployment_id == failed_deployment_id
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"value,expected",
|
||||
[
|
||||
|
|
@ -14165,6 +14192,206 @@ async def test_router_retry_policy_400_never_returns_to_a_deployment_that_alread
|
|||
assert response.choices[0].message.content == "hi back"
|
||||
|
||||
|
||||
_LIT_7114_CHAT_OK = {
|
||||
"id": "chatcmpl-lit-7114",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-5.6",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hi back"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3},
|
||||
}
|
||||
_LIT_7114_EMBEDDING_OK = {
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1]}],
|
||||
"model": "text-embedding-3-large",
|
||||
"usage": {"prompt_tokens": 1, "total_tokens": 1},
|
||||
}
|
||||
_LIT_7114_IMAGE_OK = {"created": 1, "data": [{"b64_json": "aGk="}]}
|
||||
_LIT_7114_BATCH_OK = {
|
||||
"id": "batch_lit_7114",
|
||||
"object": "batch",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"input_file_id": "file-lit-7114",
|
||||
"completion_window": "24h",
|
||||
"status": "validating",
|
||||
"created_at": 1,
|
||||
}
|
||||
|
||||
|
||||
class _PassthroughAdapter(CustomLogger):
|
||||
def translate_completion_input_params(self, kwargs: ChatCompletionRequest) -> ChatCompletionRequest:
|
||||
return ChatCompletionRequest(**kwargs)
|
||||
|
||||
def translate_completion_output_params(self, response: litellm.ModelResponse) -> litellm.ModelResponse:
|
||||
return response
|
||||
|
||||
|
||||
def _lit_7114_router(litellm_model: str) -> litellm.Router:
|
||||
api_base_suffix: Final = "" if litellm_model.startswith("cohere/") else "/v1"
|
||||
return litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-5.6",
|
||||
"litellm_params": {
|
||||
"model": litellm_model,
|
||||
"api_key": "sk-fake",
|
||||
"api_base": f"https://{host}.local{api_base_suffix}",
|
||||
"weight": weight,
|
||||
},
|
||||
"model_info": {"id": host},
|
||||
}
|
||||
for host, weight in (("rejecting", 1), ("accepting", 0))
|
||||
],
|
||||
num_retries=2,
|
||||
retry_policy={"BadRequestErrorRetries": 2},
|
||||
disable_cooldowns=True,
|
||||
)
|
||||
|
||||
|
||||
def _lit_7114_mock_upstreams(
|
||||
respx_mock: respx.MockRouter, path: str, refusal_status: int, success_body: Mapping[str, object] | bytes
|
||||
) -> tuple[respx.Route, respx.Route]:
|
||||
ok_response: Final = (
|
||||
httpx.Response(200, content=success_body)
|
||||
if isinstance(success_body, bytes)
|
||||
else httpx.Response(200, json=success_body)
|
||||
)
|
||||
rejecting: Final = respx_mock.post(f"https://rejecting.local{path}").mock(
|
||||
return_value=httpx.Response(
|
||||
refusal_status, json={"error": _UPSTREAM_400, "message": "upstream refused this request"}
|
||||
)
|
||||
)
|
||||
accepting: Final = respx_mock.post(f"https://accepting.local{path}").mock(return_value=ok_response)
|
||||
return rejecting, accepting
|
||||
|
||||
|
||||
_LIT_7114_ASYNC_ENTRYPOINTS: Final[
|
||||
Mapping[str, tuple[str, str, int, Mapping[str, object] | bytes, Callable[[litellm.Router], Awaitable[object]]]]
|
||||
] = {
|
||||
"aembedding": (
|
||||
"openai/text-embedding-3-large",
|
||||
"/v1/embeddings",
|
||||
400,
|
||||
_LIT_7114_EMBEDDING_OK,
|
||||
lambda router: router.aembedding(model="gpt-5.6", input="hi"),
|
||||
),
|
||||
"aimage_generation": (
|
||||
"openai/gpt-image-1",
|
||||
"/v1/images/generations",
|
||||
400,
|
||||
_LIT_7114_IMAGE_OK,
|
||||
lambda router: router.aimage_generation(model="gpt-5.6", prompt="a cat"),
|
||||
),
|
||||
"atext_completion": (
|
||||
"text-completion-openai/gpt-3.5-turbo-instruct",
|
||||
"/v1/completions",
|
||||
400,
|
||||
{"id": "c", "object": "text_completion", "created": 1, "model": "i", "choices": [{"text": "hi", "index": 0}]},
|
||||
lambda router: router.atext_completion(model="gpt-5.6", prompt="hi"),
|
||||
),
|
||||
"aspeech": (
|
||||
"openai/gpt-4o-mini-tts",
|
||||
"/v1/audio/speech",
|
||||
400,
|
||||
b"RIFF",
|
||||
lambda router: router.aspeech(model="gpt-5.6", input="hi", voice="alloy"),
|
||||
),
|
||||
"atranscription": (
|
||||
"openai/gpt-4o-transcribe",
|
||||
"/v1/audio/transcriptions",
|
||||
400,
|
||||
{"text": "hi"},
|
||||
lambda router: router.atranscription(model="gpt-5.6", file=("hi.wav", b"RIFF", "audio/wav")),
|
||||
),
|
||||
"arerank": (
|
||||
"cohere/rerank-v3.5",
|
||||
"/v2/rerank",
|
||||
400,
|
||||
{"id": "r", "results": [{"index": 0, "relevance_score": 0.9}], "meta": {}},
|
||||
lambda router: router.arerank(model="gpt-5.6", query="hi", documents=["hi"]),
|
||||
),
|
||||
"aadapter_completion": (
|
||||
"openai/gpt-5.6",
|
||||
"/v1/chat/completions",
|
||||
400,
|
||||
_LIT_7114_CHAT_OK,
|
||||
lambda router: router.aadapter_completion(
|
||||
adapter_id="lit-7114", model="gpt-5.6", messages=[{"role": "user", "content": "hi"}]
|
||||
),
|
||||
),
|
||||
"acreate_batch": (
|
||||
"openai/gpt-5.6",
|
||||
"/v1/batches",
|
||||
401,
|
||||
_LIT_7114_BATCH_OK,
|
||||
lambda router: router.acreate_batch(
|
||||
model="gpt-5.6", completion_window="24h", endpoint="/v1/chat/completions", input_file_id="file-lit-7114"
|
||||
),
|
||||
),
|
||||
"acancel_batch": (
|
||||
"openai/gpt-5.6",
|
||||
"/v1/batches/batch_lit_7114/cancel",
|
||||
401,
|
||||
{**_LIT_7114_BATCH_OK, "status": "cancelling"},
|
||||
lambda router: router.acancel_batch(model="gpt-5.6", batch_id="batch_lit_7114"),
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entrypoint", sorted(_LIT_7114_ASYNC_ENTRYPOINTS))
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_retry_moves_off_the_refusing_deployment_on_every_async_entrypoint(
|
||||
monkeypatch: pytest.MonkeyPatch, entrypoint: str
|
||||
):
|
||||
litellm_model, path, refusal_status, success_body, call = _LIT_7114_ASYNC_ENTRYPOINTS[entrypoint]
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "adapters", [{"id": "lit-7114", "adapter": _PassthroughAdapter()}])
|
||||
router: Final = _lit_7114_router(litellm_model)
|
||||
|
||||
with respx.mock as respx_mock:
|
||||
rejecting, accepting = _lit_7114_mock_upstreams(respx_mock, path, refusal_status, success_body)
|
||||
response: Final = await call(router)
|
||||
|
||||
assert response is not None
|
||||
assert rejecting.call_count == 1
|
||||
assert accepting.call_count == 1
|
||||
|
||||
|
||||
_LIT_7114_SYNC_ENTRYPOINTS: Final[
|
||||
Mapping[str, tuple[str, str, Mapping[str, object], Callable[[litellm.Router], object]]]
|
||||
] = {
|
||||
"embedding": (
|
||||
"openai/text-embedding-3-large",
|
||||
"/v1/embeddings",
|
||||
_LIT_7114_EMBEDDING_OK,
|
||||
lambda router: router.embedding(model="gpt-5.6", input="hi"),
|
||||
),
|
||||
"image_generation": (
|
||||
"openai/gpt-image-1",
|
||||
"/v1/images/generations",
|
||||
_LIT_7114_IMAGE_OK,
|
||||
lambda router: router.image_generation(model="gpt-5.6", prompt="a cat"),
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entrypoint", sorted(_LIT_7114_SYNC_ENTRYPOINTS))
|
||||
def test_router_retry_policy_400_moves_off_the_refusing_deployment_on_every_sync_entrypoint(
|
||||
monkeypatch: pytest.MonkeyPatch, entrypoint: str
|
||||
):
|
||||
litellm_model, path, success_body, call = _LIT_7114_SYNC_ENTRYPOINTS[entrypoint]
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
router: Final = _lit_7114_router(litellm_model)
|
||||
|
||||
with respx.mock as respx_mock:
|
||||
rejecting, accepting = _lit_7114_mock_upstreams(respx_mock, path, 400, success_body)
|
||||
response: Final = call(router)
|
||||
|
||||
assert response is not None
|
||||
assert rejecting.call_count == 1
|
||||
assert accepting.call_count == 1
|
||||
|
||||
|
||||
def _make_failure_logging_obj():
|
||||
return LiteLLMLogging(
|
||||
model="gpt-5.6",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue