diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 8a35c398cbc..1c16f6c0caa 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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: diff --git a/litellm/router.py b/litellm/router.py index 1b3bbad8e74..ca0e015af37 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 ) diff --git a/tests/test_litellm/test_model_block_unblock.py b/tests/test_litellm/test_model_block_unblock.py index 0a74996c8fe..e1ed6f2453d 100644 --- a/tests/test_litellm/test_model_block_unblock.py +++ b/tests/test_litellm/test_model_block_unblock.py @@ -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; diff --git a/tests/test_litellm/test_router_block_helpers.py b/tests/test_litellm/test_router_block_helpers.py index aaae7cd3437..4abfe368654 100644 --- a/tests/test_litellm/test_router_block_helpers.py +++ b/tests/test_litellm/test_router_block_helpers.py @@ -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