mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge 1238550d39 into 22a349ee70
This commit is contained in:
commit
1eea574eb4
6 changed files with 420 additions and 42 deletions
|
|
@ -1,9 +1,10 @@
|
|||
import asyncio
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
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 ProxyException, 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,22 +74,54 @@ 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 _validated_block_fallbacks(raw: object) -> Sequence[Mapping[str, Sequence[str]] | str] | None:
|
||||
try:
|
||||
return _BLOCK_GATE_FALLBACKS_ADAPTER.validate_python(raw)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _reachable_block_fallbacks(
|
||||
llm_router: LitellmRouter, data: Mapping[str, object], route_type: str
|
||||
) -> Sequence[Mapping[str, Sequence[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
|
||||
if "fallbacks" in data:
|
||||
return _validated_block_fallbacks(data.get("fallbacks"))
|
||||
# Router.fallbacks is an untyped list attribute; the adapter validates it into a typed view.
|
||||
return _validated_block_fallbacks(llm_router.fallbacks) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped attr, validated by adapter
|
||||
|
||||
|
||||
def _raise_if_model_fully_blocked(
|
||||
llm_router: LitellmRouter,
|
||||
model_name: Any,
|
||||
team_id: str | None,
|
||||
reachable_fallbacks: Sequence[Mapping[str, Sequence[str]] | str] | 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._is_blocked_without_reachable_fallback(
|
||||
model_name=model_name,
|
||||
reachable_fallbacks=reachable_fallbacks,
|
||||
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 = {
|
||||
|
|
@ -535,22 +586,15 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr
|
|||
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,
|
||||
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:
|
||||
|
|
|
|||
|
|
@ -10604,6 +10604,55 @@ 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,
|
||||
fallbacks: Sequence[Mapping[str, Sequence[str]] | str],
|
||||
team_id: str | None = None,
|
||||
visited: frozenset[str] = frozenset(),
|
||||
) -> bool:
|
||||
"""
|
||||
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. Traversal stops on a repeat group and after
|
||||
`max_fallbacks` hops, matching the runtime limit and bounding recursion so a long
|
||||
client-supplied chain cannot exhaust the stack.
|
||||
"""
|
||||
if model_name in visited or len(visited) >= self.max_fallbacks:
|
||||
return False
|
||||
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, fallbacks, team_id, next_visited)
|
||||
for group in fallback_model_group
|
||||
)
|
||||
|
||||
def _is_blocked_without_reachable_fallback(
|
||||
self,
|
||||
model_name: str,
|
||||
reachable_fallbacks: Sequence[Mapping[str, Sequence[str]] | str] | None,
|
||||
team_id: str | None,
|
||||
) -> bool:
|
||||
"""
|
||||
True when every deployment of `model_name` is blocked and no reachable fallback
|
||||
can serve the request. `reachable_fallbacks` is the already-resolved chain the
|
||||
caller would attempt, or None when no fallback can run on this path.
|
||||
"""
|
||||
deployments: Final = self.get_model_list(model_name=model_name, team_id=team_id) or []
|
||||
if not self._are_all_deployments_blocked(deployments):
|
||||
return False
|
||||
return reachable_fallbacks is None or not self._has_reachable_fallback(
|
||||
model_name=model_name, fallbacks=reachable_fallbacks, team_id=team_id
|
||||
)
|
||||
|
||||
async def async_get_fully_unhealthy_model_names(self) -> set[str]:
|
||||
"""
|
||||
Returns the set of model names where every backing deployment is currently
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import hashlib
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
|
@ -214,7 +214,7 @@ def _check_stripped_model_group(model_group: str, fallback_key: str) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
def get_fallback_model_group(fallbacks: list[Any], model_group: str) -> tuple[list[str] | None, int | None]:
|
||||
def get_fallback_model_group(fallbacks: Sequence[Any], model_group: str) -> tuple[list[str] | None, int | None]:
|
||||
"""
|
||||
Returns:
|
||||
- fallback_model_group: List[str] of fallback model groups. example: ["gpt-4", "gpt-3.5-turbo"]
|
||||
|
|
|
|||
|
|
@ -57,6 +57,7 @@ IGNORE_FUNCTIONS = [
|
|||
"_filter_argument_value", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by blocking the tool call at the cap.
|
||||
"_redact_scanned_content", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by returning "[REDACTED]" at the cap.
|
||||
"_iter_fallback_targets", # max depth set (2 * ROUTER_MAX_FALLBACKS); fails closed by raising ValueError at the cap.
|
||||
"_has_reachable_fallback", # bounded by the visited set: every model group is expanded at most once, so a cycle terminates.
|
||||
"json_string_leaves", # max depth set (MAX_STRUCTURED_CONTENT_SCAN_DEPTH); fails closed by raising at the cap so nothing goes unscanned.
|
||||
"with_json_string_leaves", # transitively bounded: only runs on a tree json_string_leaves already walked under the cap.
|
||||
"json_unrewritable_labels", # max depth set (MAX_STRUCTURED_CONTENT_SCAN_DEPTH); returns the None sentinel at the cap so the caller blocks.
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -40,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)
|
||||
|
|
@ -75,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(
|
||||
|
|
@ -95,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
|
||||
|
|
@ -106,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(
|
||||
|
|
@ -130,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,
|
||||
|
|
@ -205,6 +200,158 @@ async def test_route_request_returns_403_when_model_is_fully_blocked(monkeypatch
|
|||
assert "Model is blocked" in exc_info.value.message
|
||||
|
||||
|
||||
def _blocked_primary_healthy_fallback_router(fallback_blocked: bool = False) -> "litellm.Router":
|
||||
return 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": 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", "messages": [{"role": "user", "content": "hi"}]},
|
||||
llm_router=router,
|
||||
user_model=None,
|
||||
route_type="acompletion",
|
||||
)
|
||||
if asyncio.iscoroutine(result):
|
||||
result.close()
|
||||
mock_acompletion.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_blocks_eval_route_despite_healthy_fallback(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),
|
||||
)
|
||||
|
||||
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_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;
|
||||
|
|
|
|||
|
|
@ -55,3 +55,140 @@ 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),
|
||||
]
|
||||
)
|
||||
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),
|
||||
]
|
||||
)
|
||||
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", fallbacks=[]) 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),
|
||||
]
|
||||
)
|
||||
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),
|
||||
]
|
||||
)
|
||||
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),
|
||||
]
|
||||
)
|
||||
assert router._has_reachable_fallback("primary", fallbacks=[{"*": ["fallback"]}]) is True
|
||||
|
||||
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", fallbacks=["fallback"]) is True
|
||||
|
||||
def test_chain_longer_than_max_fallbacks_fails_closed(self):
|
||||
hops = 6
|
||||
router = Router(
|
||||
model_list=[_deployment("m0", "m0", blocked=True)]
|
||||
+ [_deployment(f"m{i}", f"m{i}", blocked=True) for i in range(1, hops)]
|
||||
+ [_deployment("healthy", "h0", blocked=False)],
|
||||
max_fallbacks=hops - 2,
|
||||
)
|
||||
chain = [{f"m{i}": [f"m{i + 1}"]} for i in range(hops - 1)] + [{f"m{hops - 1}": ["healthy"]}]
|
||||
assert router._has_reachable_fallback("m0", fallbacks=chain) is False
|
||||
|
||||
|
||||
class TestIsBlockedWithoutReachableFallback:
|
||||
def test_blocked_and_no_fallback_returns_true(self):
|
||||
router = Router(model_list=[_deployment("primary", "p0", blocked=True)])
|
||||
assert router._is_blocked_without_reachable_fallback("primary", reachable_fallbacks=None, team_id=None) is True
|
||||
|
||||
def test_not_all_blocked_returns_false(self):
|
||||
router = Router(model_list=[_deployment("primary", "p0", blocked=False)])
|
||||
assert router._is_blocked_without_reachable_fallback("primary", reachable_fallbacks=None, team_id=None) is False
|
||||
|
||||
def test_blocked_with_reachable_fallback_returns_false(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
_deployment("primary", "p0", blocked=True),
|
||||
_deployment("fallback", "f0", blocked=False),
|
||||
]
|
||||
)
|
||||
reachable = [{"primary": ["fallback"]}]
|
||||
assert (
|
||||
router._is_blocked_without_reachable_fallback("primary", reachable_fallbacks=reachable, team_id=None)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_blocked_with_fully_blocked_fallback_returns_true(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
_deployment("primary", "p0", blocked=True),
|
||||
_deployment("fallback", "f0", blocked=True),
|
||||
]
|
||||
)
|
||||
reachable = [{"primary": ["fallback"]}]
|
||||
assert (
|
||||
router._is_blocked_without_reachable_fallback("primary", reachable_fallbacks=reachable, team_id=None)
|
||||
is True
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue