fix(proxy): scope block-gate fallback exemption to paths that run fallbacks

The exemption now mirrors what the router actually does: eval and run routes
call litellm directly so they keep the hard 403, a request `fallbacks` list
replaces the router-level chain instead of merging, and `disable_fallbacks`
suppresses the exemption. Reachability resolution moved to the caller so the
router helper takes the already-resolved list.
This commit is contained in:
Awshesh12 2026-08-13 08:56:21 +00:00
parent e22d5af139
commit 37edd13e5b
4 changed files with 177 additions and 79 deletions

View file

@ -4,6 +4,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
from fastapi import HTTPException, status
from pydantic import TypeAdapter, ValidationError
import litellm
from litellm.proxy._types import UserAPIKeyAuth
@ -27,6 +28,24 @@ GATED_MOCK_PARAM_NAMES: Final[tuple[str, ...]] = (
MOCK_TESTING_CONFIG_KEY: Final = "dangerously_allow_mock_testing_request_params"
# Eval and run routes call litellm directly instead of going through the router, so a
# fully blocked model on these paths never reaches a fallback chain.
EVAL_ROUTE_TYPES: Final[tuple[str, ...]] = (
"acreate_eval",
"alist_evals",
"aget_eval",
"aupdate_eval",
"adelete_eval",
"acancel_eval",
"acreate_run",
"alist_runs",
"aget_run",
"acancel_run",
"adelete_run",
)
_BLOCK_GATE_FALLBACKS_ADAPTER: Final = TypeAdapter(list[dict[str, list[str]] | str])
if TYPE_CHECKING:
from litellm.router import Router as _Router
@ -55,11 +74,29 @@ def _is_a2a_agent_model(model_name: Any) -> bool:
return isinstance(model_name, str) and model_name.startswith("a2a/")
def _reachable_block_fallbacks(
llm_router: LitellmRouter, data: dict, route_type: str
) -> list[dict[str, list[str]] | str] | None:
"""Fallbacks the router would actually attempt for a blocked primary, or None when
none can run: eval routes bypass the router, disabled fallbacks skip the chain, and a
request-supplied list replaces the router-level one, matching what the router does."""
if route_type in EVAL_ROUTE_TYPES:
return None
if data.get("disable_fallbacks") is True:
return None
request_fallbacks: Final[object] = data.get("fallbacks")
raw_fallbacks: Final[object] = request_fallbacks if isinstance(request_fallbacks, list) else llm_router.fallbacks
try:
return _BLOCK_GATE_FALLBACKS_ADAPTER.validate_python(raw_fallbacks)
except ValidationError:
return 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,
reachable_fallbacks: list[dict[str, list[str]] | str] | None,
) -> None:
if not isinstance(model_name, str) or not model_name:
return
@ -68,9 +105,9 @@ def _raise_if_model_fully_blocked(
deployments: Final = llm_router.get_model_list(model_name=model_name, team_id=team_id) or []
if not llm_router._are_all_deployments_blocked(deployments):
return
if llm_router._has_reachable_fallback(
if reachable_fallbacks is not None and llm_router._has_reachable_fallback(
model_name=model_name,
request_fallbacks=request_fallbacks if isinstance(request_fallbacks, list) else None,
fallbacks=reachable_fallbacks,
team_id=team_id,
):
return
@ -505,23 +542,11 @@ async def route_request(
llm_router=llm_router,
model_name=data.get("model"),
team_id=team_id,
request_fallbacks=data.get("fallbacks"),
reachable_fallbacks=_reachable_block_fallbacks(llm_router=llm_router, data=data, route_type=route_type),
)
# Evals API: always route to litellm directly (not through router)
# But extract model credentials if a model is provided
if route_type in [
"acreate_eval",
"alist_evals",
"aget_eval",
"aupdate_eval",
"adelete_eval",
"acancel_eval",
"acreate_run",
"alist_runs",
"aget_run",
"acancel_run",
"adelete_run",
]:
if route_type in EVAL_ROUTE_TYPES:
# If a model is provided, get its credentials from the router
model: Final = data.get("model")
if model and llm_router:

View file

@ -9983,25 +9983,25 @@ class Router:
def _has_reachable_fallback(
self,
model_name: str,
request_fallbacks: list[dict[str, list[str]]] | None = None,
fallbacks: list[dict[str, list[str]] | str],
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.
True when `fallbacks` routes `model_name` to a model group with at least one
unblocked deployment. `fallbacks` must already reflect the precedence the router
applies at call time, so a request-supplied list replaces the router-level chain
rather than merging with it. Visited groups are skipped so a cycle 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)
fallback_model_group, _ = get_fallback_model_group(fallbacks=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)
or self._has_reachable_fallback(group, fallbacks, team_id, next_visited)
for group in fallback_model_group
)

View file

@ -41,9 +41,7 @@ def _setup_model_block_mocks(monkeypatch, *, updated_blocked: bool):
# No reconcile ran in these tests, so both fields are None and the verdict falls
# back to reading the router live -- which is what the get_model_ids side_effects
# below drive.
mock_clear_cache = AsyncMock(
return_value=ReconcileOutcome(still_desired=None, live_after=None)
)
mock_clear_cache = AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None))
mock_audit_log = AsyncMock(return_value=None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
@ -76,8 +74,8 @@ async def test_model_block_endpoint_sets_blocked_true(monkeypatch):
block_model,
)
model_id, model_table, updated_row, mock_clear_cache, mock_audit_log = (
_setup_model_block_mocks(monkeypatch, updated_blocked=True)
model_id, model_table, updated_row, mock_clear_cache, mock_audit_log = _setup_model_block_mocks(
monkeypatch, updated_blocked=True
)
result = await block_model(
@ -96,9 +94,7 @@ async def test_model_block_endpoint_sets_blocked_true(monkeypatch):
assert "updated_at" in update_kwargs["data"]
mock_clear_cache.assert_awaited_once_with()
assert mock_audit_log.call_args.kwargs["action"] == "blocked"
assert (
mock_audit_log.call_args.kwargs["litellm_changed_by"] == "operator@example.com"
)
assert mock_audit_log.call_args.kwargs["litellm_changed_by"] == "operator@example.com"
@pytest.mark.asyncio
@ -107,8 +103,8 @@ async def test_model_unblock_endpoint_sets_blocked_false(monkeypatch):
unblock_model,
)
model_id, model_table, updated_row, mock_clear_cache, mock_audit_log = (
_setup_model_block_mocks(monkeypatch, updated_blocked=False)
model_id, model_table, updated_row, mock_clear_cache, mock_audit_log = _setup_model_block_mocks(
monkeypatch, updated_blocked=False
)
result = await unblock_model(
@ -131,9 +127,7 @@ async def test_model_block_endpoint_requires_proxy_admin(monkeypatch):
block_model,
)
model_id, model_table, _, _, _ = _setup_model_block_mocks(
monkeypatch, updated_blocked=True
)
model_id, model_table, _, _, _ = _setup_model_block_mocks(monkeypatch, updated_blocked=True)
non_admin = UserAPIKeyAuth(
user_id="internal-user",
user_role=LitellmUserRoles.INTERNAL_USER,
@ -206,11 +200,8 @@ 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(
def _blocked_primary_healthy_fallback_router(fallback_blocked: bool = False) -> "litellm.Router":
return litellm.Router(
model_list=[
{
"model_name": "primary",
@ -220,45 +211,108 @@ async def test_route_request_allows_fallback_when_primary_fully_blocked(monkeypa
{
"model_name": "fallback",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "x"},
"model_info": {"id": "f0", "blocked": False},
"model_info": {"id": "f0", "blocked": fallback_blocked},
},
],
fallbacks=[{"primary": ["fallback"]}],
)
def test_block_gate_allows_request_when_reachable_fallback_supplied():
from litellm.proxy.route_llm_request import _raise_if_model_fully_blocked
router = _blocked_primary_healthy_fallback_router()
_raise_if_model_fully_blocked(
llm_router=router,
model_name="primary",
team_id=None,
reachable_fallbacks=[{"primary": ["fallback"]}],
)
def test_block_gate_raises_when_no_fallback_reaches_healthy_group():
from litellm.proxy.route_llm_request import _raise_if_model_fully_blocked
router = _blocked_primary_healthy_fallback_router()
with pytest.raises(litellm.PermissionDeniedError) as exc_info:
_raise_if_model_fully_blocked(
llm_router=router,
model_name="primary",
team_id=None,
reachable_fallbacks=None,
)
assert exc_info.value.status_code == 403
def test_reachable_block_fallbacks_none_for_eval_route():
from litellm.proxy.route_llm_request import _reachable_block_fallbacks
router = _blocked_primary_healthy_fallback_router()
result = _reachable_block_fallbacks(llm_router=router, data={"model": "primary"}, route_type="acreate_eval")
assert result is None
def test_reachable_block_fallbacks_none_when_disabled():
from litellm.proxy.route_llm_request import _reachable_block_fallbacks
router = _blocked_primary_healthy_fallback_router()
result = _reachable_block_fallbacks(
llm_router=router,
data={"model": "primary", "disable_fallbacks": True},
route_type="acompletion",
)
assert result is None
def test_reachable_block_fallbacks_request_list_replaces_router_chain():
from litellm.proxy.route_llm_request import _reachable_block_fallbacks
router = _blocked_primary_healthy_fallback_router()
result = _reachable_block_fallbacks(
llm_router=router,
data={"model": "primary", "fallbacks": [{"other": ["x"]}]},
route_type="acompletion",
)
assert result == [{"other": ["x"]}]
def test_reachable_block_fallbacks_uses_router_chain_without_request_list():
from litellm.proxy.route_llm_request import _reachable_block_fallbacks
router = _blocked_primary_healthy_fallback_router()
result = _reachable_block_fallbacks(llm_router=router, data={"model": "primary"}, route_type="acompletion")
assert result == [{"primary": ["fallback"]}]
@pytest.mark.asyncio
async def test_route_request_allows_completion_fallback_when_primary_fully_blocked(monkeypatch):
from litellm.proxy.route_llm_request import route_request
router = _blocked_primary_healthy_fallback_router()
monkeypatch.setattr(
"litellm.proxy.route_llm_request.add_shared_session_to_data",
AsyncMock(return_value=None),
)
mock_acompletion = AsyncMock(return_value="ok")
monkeypatch.setattr(router, "acompletion", mock_acompletion)
result = await route_request(
data={"model": "primary"},
data={"model": "primary", "messages": [{"role": "user", "content": "hi"}]},
llm_router=router,
user_model=None,
route_type="acreate_eval",
route_type="acompletion",
)
if asyncio.iscoroutine(result):
result.close()
mock_acompletion.assert_called_once()
@pytest.mark.asyncio
async def test_route_request_blocks_when_primary_and_fallback_fully_blocked(monkeypatch):
async def test_route_request_blocks_eval_route_despite_healthy_fallback(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"]}],
)
router = _blocked_primary_healthy_fallback_router()
monkeypatch.setattr(
"litellm.proxy.route_llm_request.add_shared_session_to_data",
AsyncMock(return_value=None),
@ -276,6 +330,28 @@ async def test_route_request_blocks_when_primary_and_fallback_fully_blocked(monk
assert "Model is blocked" in exc_info.value.message
@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 = _blocked_primary_healthy_fallback_router(fallback_blocked=True)
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", "messages": [{"role": "user", "content": "hi"}]},
llm_router=router,
user_model=None,
route_type="acompletion",
)
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;

View file

@ -90,24 +90,22 @@ class TestHasReachableFallback:
model_list=[
_deployment("primary", "p0", blocked=True),
_deployment("fallback", "f0", blocked=False),
],
fallbacks=[{"primary": ["fallback"]}],
]
)
assert router._has_reachable_fallback("primary") is True
assert router._has_reachable_fallback("primary", fallbacks=[{"primary": ["fallback"]}]) 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
assert router._has_reachable_fallback("primary", fallbacks=[{"primary": ["fallback"]}]) 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
assert router._has_reachable_fallback("primary", fallbacks=[]) is False
def test_reachable_through_multi_level_chain(self):
router = Router(
@ -115,36 +113,35 @@ class TestHasReachableFallback:
_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
chain = [{"primary": ["mid"]}, {"mid": ["healthy"]}]
assert router._has_reachable_fallback("primary", fallbacks=chain) 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
chain = [{"primary": ["fallback"]}, {"fallback": ["primary"]}]
assert router._has_reachable_fallback("primary", fallbacks=chain) 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
assert router._has_reachable_fallback("primary", fallbacks=[{"*": ["fallback"]}]) is True
def test_request_level_fallbacks_are_honored(self):
def test_string_form_fallback_is_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
assert router._has_reachable_fallback("primary", fallbacks=["fallback"]) is True