mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
fix(proxy): let a fully blocked model fall back to a healthy group
A model group with every deployment blocked returned 403 at the proxy gate before routing, so a configured fallback to a healthy group never ran. Only raise when no per-request or router-level fallback reaches a group with an unblocked deployment. Fixes #36665
This commit is contained in:
parent
09889e1986
commit
e22d5af139
4 changed files with 222 additions and 12 deletions
|
|
@ -55,22 +55,34 @@ def _is_a2a_agent_model(model_name: Any) -> bool:
|
|||
return isinstance(model_name, str) and model_name.startswith("a2a/")
|
||||
|
||||
|
||||
def _raise_if_model_fully_blocked(llm_router: LitellmRouter, model_name: Any, team_id: str | None) -> None:
|
||||
def _raise_if_model_fully_blocked(
|
||||
llm_router: LitellmRouter,
|
||||
model_name: Any,
|
||||
team_id: str | None,
|
||||
request_fallbacks: list[dict[str, list[str]]] | None = None,
|
||||
) -> None:
|
||||
if not isinstance(model_name, str) or not model_name:
|
||||
return
|
||||
if not isinstance(llm_router, litellm.Router):
|
||||
return
|
||||
deployments: Final = llm_router.get_model_list(model_name=model_name, team_id=team_id) or []
|
||||
if llm_router._are_all_deployments_blocked(deployments):
|
||||
raise litellm.PermissionDeniedError(
|
||||
message="Model is blocked",
|
||||
model=model_name,
|
||||
llm_provider="",
|
||||
response=httpx.Response(
|
||||
status_code=403,
|
||||
request=httpx.Request(method="POST", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
if not llm_router._are_all_deployments_blocked(deployments):
|
||||
return
|
||||
if llm_router._has_reachable_fallback(
|
||||
model_name=model_name,
|
||||
request_fallbacks=request_fallbacks if isinstance(request_fallbacks, list) else None,
|
||||
team_id=team_id,
|
||||
):
|
||||
return
|
||||
raise litellm.PermissionDeniedError(
|
||||
message="Model is blocked",
|
||||
model=model_name,
|
||||
llm_provider="",
|
||||
response=httpx.Response(
|
||||
status_code=403,
|
||||
request=httpx.Request(method="POST", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
ROUTE_ENDPOINT_MAPPING: Final = {
|
||||
|
|
@ -489,7 +501,12 @@ async def route_request(
|
|||
else:
|
||||
return getattr(litellm, f"{route_type}")(**data)
|
||||
elif llm_router is not None:
|
||||
_raise_if_model_fully_blocked(llm_router=llm_router, model_name=data.get("model"), team_id=team_id)
|
||||
_raise_if_model_fully_blocked(
|
||||
llm_router=llm_router,
|
||||
model_name=data.get("model"),
|
||||
team_id=team_id,
|
||||
request_fallbacks=data.get("fallbacks"),
|
||||
)
|
||||
# Evals API: always route to litellm directly (not through router)
|
||||
# But extract model credentials if a model is provided
|
||||
if route_type in [
|
||||
|
|
|
|||
|
|
@ -9976,6 +9976,35 @@ class Router:
|
|||
deployments: Final = self.get_model_list(model_name=model) or []
|
||||
return self._are_all_deployments_blocked(deployments=deployments)
|
||||
|
||||
def _model_group_has_unblocked_deployment(self, model_name: str, team_id: str | None) -> bool:
|
||||
deployments: Final = self.get_model_list(model_name=model_name, team_id=team_id) or []
|
||||
return any((deployment.get("model_info") or {}).get("blocked") is not True for deployment in deployments)
|
||||
|
||||
def _has_reachable_fallback(
|
||||
self,
|
||||
model_name: str,
|
||||
request_fallbacks: list[dict[str, list[str]]] | None = None,
|
||||
team_id: str | None = None,
|
||||
visited: frozenset[str] = frozenset(),
|
||||
) -> bool:
|
||||
"""
|
||||
True when `model_name`'s configured fallback chain reaches a model group with at
|
||||
least one unblocked deployment. Follows per-request and router-level fallbacks and
|
||||
skips already-visited groups so a self-referential chain terminates.
|
||||
"""
|
||||
if model_name in visited:
|
||||
return False
|
||||
combined_fallbacks: Final[list[dict[str, list[str]]]] = [*(request_fallbacks or []), *(self.fallbacks or [])]
|
||||
fallback_model_group, _ = get_fallback_model_group(fallbacks=combined_fallbacks, model_group=model_name)
|
||||
if not fallback_model_group:
|
||||
return False
|
||||
next_visited: Final = visited | {model_name}
|
||||
return any(
|
||||
self._model_group_has_unblocked_deployment(group, team_id)
|
||||
or self._has_reachable_fallback(group, request_fallbacks, team_id, next_visited)
|
||||
for group in fallback_model_group
|
||||
)
|
||||
|
||||
async def async_get_fully_unhealthy_model_names(self) -> set[str]:
|
||||
"""
|
||||
Returns the set of model names where every backing deployment is currently
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -205,6 +206,76 @@ async def test_route_request_returns_403_when_model_is_fully_blocked(monkeypatch
|
|||
assert "Model is blocked" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_allows_fallback_when_primary_fully_blocked(monkeypatch):
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "primary",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "x"},
|
||||
"model_info": {"id": "p0", "blocked": True},
|
||||
},
|
||||
{
|
||||
"model_name": "fallback",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "x"},
|
||||
"model_info": {"id": "f0", "blocked": False},
|
||||
},
|
||||
],
|
||||
fallbacks=[{"primary": ["fallback"]}],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.route_llm_request.add_shared_session_to_data",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
|
||||
result = await route_request(
|
||||
data={"model": "primary"},
|
||||
llm_router=router,
|
||||
user_model=None,
|
||||
route_type="acreate_eval",
|
||||
)
|
||||
if asyncio.iscoroutine(result):
|
||||
result.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_blocks_when_primary_and_fallback_fully_blocked(monkeypatch):
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "primary",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "x"},
|
||||
"model_info": {"id": "p0", "blocked": True},
|
||||
},
|
||||
{
|
||||
"model_name": "fallback",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "x"},
|
||||
"model_info": {"id": "f0", "blocked": True},
|
||||
},
|
||||
],
|
||||
fallbacks=[{"primary": ["fallback"]}],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.route_llm_request.add_shared_session_to_data",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.PermissionDeniedError) as exc_info:
|
||||
await route_request(
|
||||
data={"model": "primary"},
|
||||
llm_router=router,
|
||||
user_model=None,
|
||||
route_type="acreate_eval",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Model is blocked" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_block_surfaces_wholesale_reload_failure(monkeypatch):
|
||||
"""The write endpoints owe the caller an error when the pod failed to reload at all;
|
||||
|
|
|
|||
|
|
@ -55,3 +55,96 @@ class TestIsModelFullyBlocked:
|
|||
def test_unblocked_deployment_returns_false(self):
|
||||
router = _make_router("gpt-4o", blocked=False)
|
||||
assert router._is_model_fully_blocked("gpt-4o") is False
|
||||
|
||||
|
||||
def _deployment(model_name: str, dep_id: str, blocked: bool) -> dict:
|
||||
return {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"},
|
||||
"model_info": {"id": dep_id, "blocked": blocked},
|
||||
}
|
||||
|
||||
|
||||
class TestModelGroupHasUnblockedDeployment:
|
||||
def test_one_unblocked_returns_true(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
_deployment("group", "d0", blocked=True),
|
||||
_deployment("group", "d1", blocked=False),
|
||||
]
|
||||
)
|
||||
assert router._model_group_has_unblocked_deployment("group", team_id=None) is True
|
||||
|
||||
def test_all_blocked_returns_false(self):
|
||||
router = Router(model_list=[_deployment("group", "d0", blocked=True)])
|
||||
assert router._model_group_has_unblocked_deployment("group", team_id=None) is False
|
||||
|
||||
def test_unknown_group_returns_false(self):
|
||||
router = _make_router("group", blocked=False)
|
||||
assert router._model_group_has_unblocked_deployment("missing", team_id=None) is False
|
||||
|
||||
|
||||
class TestHasReachableFallback:
|
||||
def test_reachable_when_direct_fallback_has_unblocked_deployment(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
_deployment("primary", "p0", blocked=True),
|
||||
_deployment("fallback", "f0", blocked=False),
|
||||
],
|
||||
fallbacks=[{"primary": ["fallback"]}],
|
||||
)
|
||||
assert router._has_reachable_fallback("primary") is True
|
||||
|
||||
def test_not_reachable_when_fallback_also_fully_blocked(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
_deployment("primary", "p0", blocked=True),
|
||||
_deployment("fallback", "f0", blocked=True),
|
||||
],
|
||||
fallbacks=[{"primary": ["fallback"]}],
|
||||
)
|
||||
assert router._has_reachable_fallback("primary") is False
|
||||
|
||||
def test_not_reachable_without_fallbacks(self):
|
||||
router = Router(model_list=[_deployment("primary", "p0", blocked=True)])
|
||||
assert router._has_reachable_fallback("primary") is False
|
||||
|
||||
def test_reachable_through_multi_level_chain(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
_deployment("primary", "p0", blocked=True),
|
||||
_deployment("mid", "m0", blocked=True),
|
||||
_deployment("healthy", "h0", blocked=False),
|
||||
],
|
||||
fallbacks=[{"primary": ["mid"]}, {"mid": ["healthy"]}],
|
||||
)
|
||||
assert router._has_reachable_fallback("primary") is True
|
||||
|
||||
def test_self_referential_chain_terminates(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
_deployment("primary", "p0", blocked=True),
|
||||
_deployment("fallback", "f0", blocked=True),
|
||||
],
|
||||
fallbacks=[{"primary": ["fallback"]}, {"fallback": ["primary"]}],
|
||||
)
|
||||
assert router._has_reachable_fallback("primary") is False
|
||||
|
||||
def test_generic_star_fallback_is_honored(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
_deployment("primary", "p0", blocked=True),
|
||||
_deployment("fallback", "f0", blocked=False),
|
||||
],
|
||||
fallbacks=[{"*": ["fallback"]}],
|
||||
)
|
||||
assert router._has_reachable_fallback("primary") is True
|
||||
|
||||
def test_request_level_fallbacks_are_honored(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
_deployment("primary", "p0", blocked=True),
|
||||
_deployment("fallback", "f0", blocked=False),
|
||||
]
|
||||
)
|
||||
assert router._has_reachable_fallback("primary", request_fallbacks=[{"primary": ["fallback"]}]) is True
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue