fix(proxy): accept team-scoped models by their public name on POST /fallback (#45455)

* fix(proxy): accept team-scoped models by their public name on POST /fallback

create_fallback validated the primary and fallback models against the
router's stored model names only, so a model created through POST
/model/new with model_info.team_id was rejected with a 404 unless the
caller used the generated model_name_<team_id>_<uuid> name. The
endpoint now also accepts the team public model names, which request
time fallback matching already keys on, and lists both kinds of names
in the 404's available_models.

* fix(proxy): read fallback rules fresh and clear the config cache after a fallback write

A second POST or DELETE /fallback within the 60 s config cache TTL started from
a cached copy of router_settings and dropped every rule stored since that copy
was taken, by any instance. Both endpoints now evict the cached row before the
read and invalidate it after the upsert

* test(proxy): type the fallback endpoint tests and prove a team request fails over by public name

* fix(proxy): keep fallback writes working when the Redis config cache is down and type the stored settings read

* fix(proxy): keep the non-standard fallback shapes the router accepts on writes and resolve the config cache at call time

* test(proxy): mark the config cache outage test's result as Final

* fix(proxy): replace a same-key fallback rule in place and read a null rule list as empty

* test(proxy): let the stored router settings fixture carry a null rule list

* test(proxy): audit cells for fallback rules by team public name

* test(proxy): delete the fallback rules the audit cells save

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-08 22:38:38 +00:00 • committed by GitHub
parent 5133485009
commit 5d207d85ee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 743 additions and 51 deletions

View file

@ -11,15 +11,21 @@ DELETE /fallback/{model} - Delete fallbacks for a specific model
# pyright: reportMissingImports=false
import json
from typing import TYPE_CHECKING, Final, Literal
from collections.abc import Mapping
from typing import TYPE_CHECKING, Annotated, Final, Literal
from pydantic import Field, TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.model_checks import get_all_fallbacks
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.utils import PrismaClient, evict_config_param, invalidate_config_param
if TYPE_CHECKING:
from fastapi import APIRouter, Depends, HTTPException, status
from litellm.proxy.proxy_server import ProxyConfig
else:
try:
from fastapi import APIRouter, Depends, HTTPException, status
@ -37,6 +43,33 @@ from litellm.types.management_endpoints.router_settings_endpoints import (
router: Final = APIRouter()
ROUTER_SETTINGS_PARAM: Final = "router_settings"
FallbackRule = dict[str, list[str]]
StoredFallback = Annotated[FallbackRule | dict[str, object] | str, Field(union_mode="left_to_right")]
_STORED_FALLBACKS: Final = TypeAdapter(list[StoredFallback])
def _rule_covers(entry: StoredFallback, model: str) -> bool:
return isinstance(entry, dict) and model in entry
async def _router_settings_fresh_from_db(proxy_config: "ProxyConfig") -> dict[str, object]:
await evict_config_param(ROUTER_SETTINGS_PARAM)
config: Final = await proxy_config.get_config()
return config.get(ROUTER_SETTINGS_PARAM, {})
async def _persist_router_settings(prisma_client: PrismaClient, router_settings: Mapping[str, object]) -> None:
router_settings_json: Final = json.dumps(router_settings)
await ConfigRepository(prisma_client).table.upsert(
where={"param_name": ROUTER_SETTINGS_PARAM},
data={
"create": {"param_name": ROUTER_SETTINGS_PARAM, "param_value": router_settings_json},
"update": {"param_value": router_settings_json},
},
)
await invalidate_config_param(ROUTER_SETTINGS_PARAM)
@router.post(
"/fallback",
@ -85,24 +118,24 @@ async def create_fallback(
)
# Validate that the model exists in the router
model_names: Final = llm_router.model_names
if data.model not in model_names:
known_model_names: Final = frozenset(llm_router.model_names) | llm_router.team_public_model_names
if data.model not in known_model_names:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail={
"error": f"Model '{data.model}' not found in router",
"available_models": list(model_names),
"available_models": sorted(known_model_names),
},
)
# Validate that all fallback models exist in the router
invalid_fallback_models: Final = [m for m in data.fallback_models if m not in model_names]
invalid_fallback_models: Final = [m for m in data.fallback_models if m not in known_model_names]
if invalid_fallback_models:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": f"Invalid fallback models: {invalid_fallback_models}",
"available_models": list(model_names),
"available_models": sorted(known_model_names),
},
)
@ -122,9 +155,7 @@ async def create_fallback(
},
)
# Load existing config
config: Final = await proxy_config.get_config()
router_settings: Final = config.get("router_settings", {})
router_settings: Final = await _router_settings_fresh_from_db(proxy_config)
# Get the appropriate fallback list based on type
fallback_key = "fallbacks"
@ -134,12 +165,12 @@ async def create_fallback(
fallback_key = "content_policy_fallbacks"
# Get existing fallbacks
existing_fallbacks: Final[list[dict[str, list[str]]]] = router_settings.get(fallback_key, [])
existing_fallbacks: Final = _STORED_FALLBACKS.validate_python(router_settings.get(fallback_key) or [])
# Update or add the fallback configuration
fallback_updated = False
for i, fallback_dict in enumerate(existing_fallbacks):
if data.model in fallback_dict:
for i, rule in enumerate(existing_fallbacks):
if _rule_covers(rule, data.model):
# Update existing fallback
existing_fallbacks[i] = {data.model: data.fallback_models}
fallback_updated = True
@ -152,18 +183,7 @@ async def create_fallback(
# Update router settings
router_settings[fallback_key] = existing_fallbacks
# Save to database - convert router_settings to JSON string
router_settings_json: Final = json.dumps(router_settings)
await ConfigRepository(prisma_client).table.upsert(
where={"param_name": "router_settings"},
data={
"create": {
"param_name": "router_settings",
"param_value": router_settings_json,
},
"update": {"param_value": router_settings_json},
},
)
await _persist_router_settings(prisma_client, router_settings)
# Update the in-memory router configuration
setattr(llm_router, fallback_key, existing_fallbacks)
@ -291,9 +311,7 @@ async def delete_fallback(
},
)
# Load existing config
config: Final = await proxy_config.get_config()
router_settings: Final = config.get("router_settings", {})
router_settings: Final = await _router_settings_fresh_from_db(proxy_config)
# Get the appropriate fallback list based on type
fallback_key = "fallbacks"
@ -303,14 +321,14 @@ async def delete_fallback(
fallback_key = "content_policy_fallbacks"
# Get existing fallbacks
existing_fallbacks: Final[list[dict[str, list[str]]]] = router_settings.get(fallback_key, [])
existing_fallbacks: Final = _STORED_FALLBACKS.validate_python(router_settings.get(fallback_key) or [])
# Find and remove the fallback configuration
fallback_found = False
updated_fallbacks: Final = []
for fallback_dict in existing_fallbacks:
if model not in fallback_dict:
updated_fallbacks.append(fallback_dict)
for rule in existing_fallbacks:
if not _rule_covers(rule, model):
updated_fallbacks.append(rule)
else:
fallback_found = True
@ -323,18 +341,7 @@ async def delete_fallback(
# Update router settings
router_settings[fallback_key] = updated_fallbacks
# Save to database - convert router_settings to JSON string
router_settings_json: Final = json.dumps(router_settings)
await ConfigRepository(prisma_client).table.upsert(
where={"param_name": "router_settings"},
data={
"create": {
"param_name": "router_settings",
"param_value": router_settings_json,
},
"update": {"param_value": router_settings_json},
},
)
await _persist_router_settings(prisma_client, router_settings)
# Update the in-memory router configuration
setattr(llm_router, fallback_key, updated_fallbacks)

View file

@ -4396,9 +4396,13 @@ async def get_config_param(prisma_client: "PrismaClient", param_name: str) -> An
return row
async def evict_config_param(param_name: str) -> None:
with service_target(CONFIG_PARAMS_TARGET):
await litellm_config_cache.async_delete_cache(_config_cache_key(param_name))
async def evict_config_param(param_name: str, cache: DualCache | None = None) -> None:
target: Final = cache if cache is not None else litellm_config_cache
try:
with service_target(CONFIG_PARAMS_TARGET):
await target.async_delete_cache(_config_cache_key(param_name))
except Exception as e: # noqa: BLE001 # best-effort eviction; config writes must never fail on redis errors
verbose_proxy_logger.warning("config cache eviction of %s failed: %s", param_name, e)
async def invalidate_config_param(param_name: str) -> None:

View file

@ -0,0 +1,379 @@
import uuid
from collections.abc import Iterator
from concurrent.futures import ThreadPoolExecutor
from functools import partial
from itertools import chain
from pathlib import Path
from typing import Final
import httpx
import psutil
import pytest
from pydantic import JsonValue
from tests.integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.process import OwnedProxy, owned_proxy_process
from tests.integration._support.redis_process import owned_redis
pytestmark: Final = pytest.mark.timeout(300)
PROVIDER_KEY: Final = "integration-provider-key"
EVICTION_WARNING: Final = "config cache eviction of router_settings failed"
REFUSED: Final = frozenset({401, 403})
@pytest.fixture
def upstream(gateway: Gateway) -> Iterator[httpx.Client]:
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as client:
client.get("/__observations").raise_for_status()
yield client
def _observed_requests(upstream: httpx.Client) -> list[JsonValue]:
observed: Final = upstream.get("/__observations")
observed.raise_for_status()
requests: Final = object_value(observed.json())["requests"]
assert isinstance(requests, list)
return requests
def _calls_to(observed: list[JsonValue], provider_model: str) -> int:
return sum(object_value(object_value(request)["body"]).get("model") == provider_model for request in observed)
def _fallback_body(model: str, fallback_models: list[str], fallback_type: str = "general") -> dict[str, JsonValue]:
return {"model": model, "fallback_models": list(fallback_models), "fallback_type": fallback_type}
def _forget_fallback(gateway: Gateway, model: str, fallback_type: str) -> None:
gateway.request("DELETE", f"/fallback/{model}", params={"fallback_type": fallback_type})
def _create_fallback(
gateway: Gateway, scenario: Scenario, model: str, fallback_models: list[str], fallback_type: str = "general"
) -> httpx.Response:
scenario.cleanups.callback(_forget_fallback, gateway, model, fallback_type)
return eventually(
lambda: gateway.request("POST", "/fallback", _fallback_body(model, fallback_models, fallback_type)),
lambda response: response.status_code == 200,
seconds=30,
return_last_on_timeout=True,
)
def _fallback_models(response: httpx.Response) -> list[str]:
body: Final = response.json()
models: Final = body.get("fallback_models") if isinstance(body, dict) else None
return [string_value(entry) for entry in models] if isinstance(models, list) else []
def _available_models(response: httpx.Response) -> list[str]:
body: Final = response.json()
if not isinstance(body, dict):
return []
detail: Final = body.get("detail")
source: Final = detail if isinstance(detail, dict) else body
models: Final = source.get("available_models")
return [string_value(entry) for entry in models] if isinstance(models, list) else []
def _stored_fallbacks(fallback_key: str = "fallbacks") -> list[JsonValue]:
rows: Final = read_rows('SELECT param_value FROM "LiteLLM_Config" WHERE param_name = %s', ("router_settings",))
if not rows:
return []
settings: Final = object_value(rows[0]["param_value"])
entries: Final = settings.get(fallback_key) or []
assert isinstance(entries, list), entries
return entries
def _covered_models(entries: list[JsonValue]) -> frozenset[str]:
dict_entries: Final = (object_value(entry) for entry in entries if isinstance(entry, dict))
return frozenset(chain.from_iterable(dict_entries))
def _models_over_a_fresh_connection(gateway: Gateway, _: int) -> frozenset[str]:
with httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False) as client:
listed: Final = client.get("/v1/models", headers={"Authorization": f"Bearer {gateway.key}"})
assert listed.status_code == 200, listed.text
data: Final = object_value(listed.json())["data"]
assert isinstance(data, list), listed.text
return frozenset(string_value(object_value(entry)["id"]) for entry in data)
def _every_worker_serves(gateway: Gateway, model: str) -> bool:
with ThreadPoolExecutor(max_workers=16) as pool:
rounds: Final = tuple(
tuple(pool.map(partial(_models_over_a_fresh_connection, gateway), range(16))) for _ in range(2)
)
return all(model in seen for seen in chain.from_iterable(rounds))
def _wait_until_served(gateway: Gateway, model: str) -> None:
eventually(lambda: _every_worker_serves(gateway, model), lambda served: served, seconds=90)
def _team_model(gateway: Gateway, scenario: Scenario, team: str, provider_model: str) -> tuple[str, str]:
public: Final = f"ipub-{uuid.uuid4().hex}"
created: Final = gateway.post(
"/model/new",
{
"model_name": public,
"litellm_params": {
"model": f"openai/{provider_model}",
"api_key": PROVIDER_KEY,
"api_base": f"{gateway.upstream_url}/v1",
"num_retries": 0,
},
"model_info": {"team_id": team},
},
)
info: Final = object_value(created["model_info"])
scenario.cleanups.callback(scenario.delete_model, string_value(info["id"]))
return public, string_value(created["model_name"])
def _workers(owned: OwnedProxy) -> tuple[psutil.Process, ...]:
return tuple(child for child in psutil.Process(owned.process.pid).children() if _is_worker(child))
def _is_worker(child: psutil.Process) -> bool:
try:
return "spawn_main" in " ".join(child.cmdline()) and child.status() != psutil.STATUS_ZOMBIE
except psutil.Error:
return False
def test_post_by_team_public_name_creates_reads_back_and_peer_converges(gateway: Gateway, peer: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
primary: Final = scenario.model(model=f"openai/n1p-{uuid.uuid4().hex}", model_info={"team_id": team})
fallback: Final = scenario.model(model=f"openai/n1f-{uuid.uuid4().hex}", model_info={"team_id": team})
created: Final = _create_fallback(gateway, scenario, primary, [fallback])
assert created.status_code == 200, created.text
here: Final = eventually(
lambda: gateway.request("GET", f"/fallback/{primary}", params={"fallback_type": "general"}),
lambda response: response.status_code == 200 and fallback in _fallback_models(response),
seconds=30,
return_last_on_timeout=True,
)
assert here.status_code == 200 and fallback in _fallback_models(here), here.text
there: Final = eventually(
lambda: peer.request("GET", f"/fallback/{primary}", params={"fallback_type": "general"}),
lambda response: response.status_code == 200 and fallback in _fallback_models(response),
seconds=30,
return_last_on_timeout=True,
)
assert there.status_code == 200 and fallback in _fallback_models(there), there.text
def test_rule_by_team_public_name_fires_on_chat_completions(gateway: Gateway, upstream: httpx.Client) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
primary_provider: Final = f"n2prim-{uuid.uuid4().hex}"
fallback_provider: Final = f"n2fbk-{uuid.uuid4().hex}"
primary: Final = scenario.model(model=f"openai/{primary_provider}", num_retries=0, model_info={"team_id": team})
fallback: Final = scenario.model(model=f"openai/{fallback_provider}", model_info={"team_id": team})
team_key: Final = scenario.key(team_id=team)
created: Final = _create_fallback(gateway, scenario, primary, [fallback])
assert created.status_code == 200, created.text
upstream.post(f"/__scripts/{primary_provider}", json={"statuses": [500]}).raise_for_status()
upstream.get("/__observations").raise_for_status()
answered: Final = eventually(
lambda: gateway.request(
"POST",
"/v1/chat/completions",
{"model": primary, "messages": [{"role": "user", "content": f"fire {uuid.uuid4().hex}"}]},
key=team_key,
),
lambda response: response.status_code == 200,
seconds=60,
return_last_on_timeout=True,
)
assert answered.status_code == 200, answered.text
observed: Final = _observed_requests(upstream)
assert _calls_to(observed, primary_provider) >= 1, observed
assert _calls_to(observed, fallback_provider) >= 1, observed
def test_consecutive_creates_within_the_cache_window_keep_every_rule(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
target: Final = scenario.model(model=f"openai/n5t-{uuid.uuid4().hex}")
primaries: Final = tuple(scenario.model(model=f"openai/n5{tag}-{uuid.uuid4().hex}") for tag in "abc")
for model in (target, *primaries):
_wait_until_served(gateway, model)
created: Final = tuple(_create_fallback(gateway, scenario, primary, [target]) for primary in primaries)
assert all(response.status_code == 200 for response in created), [response.text for response in created]
covered: Final = _covered_models(_stored_fallbacks())
assert frozenset(primaries) <= covered, (primaries, covered)
def test_unknown_model_404_lists_team_public_and_internal_names(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
public, internal = _team_model(gateway, scenario, team, f"n6-{uuid.uuid4().hex}")
unknown: Final = f"n6-unknown-{uuid.uuid4().hex}"
refused: Final = eventually(
lambda: gateway.request("POST", "/fallback", _fallback_body(unknown, [public])),
lambda response: response.status_code == 404 and internal in _available_models(response),
seconds=30,
return_last_on_timeout=True,
)
assert refused.status_code == 404, refused.text
available: Final = _available_models(refused)
assert internal in available, (internal, available)
assert public in available, (public, available)
def test_create_by_generated_internal_name_still_works(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
_, primary_internal = _team_model(gateway, scenario, team, f"c1p-{uuid.uuid4().hex}")
_, fallback_internal = _team_model(gateway, scenario, team, f"c1f-{uuid.uuid4().hex}")
created: Final = _create_fallback(gateway, scenario, primary_internal, [fallback_internal])
assert created.status_code == 200, created.text
here: Final = eventually(
lambda: gateway.request("GET", f"/fallback/{primary_internal}", params={"fallback_type": "general"}),
lambda response: response.status_code == 200 and fallback_internal in _fallback_models(response),
seconds=30,
return_last_on_timeout=True,
)
assert here.status_code == 200 and fallback_internal in _fallback_models(here), here.text
def test_delete_reads_fresh_rules_and_keeps_the_others(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
target: Final = scenario.model(model=f"openai/c2t-{uuid.uuid4().hex}")
keep: Final = scenario.model(model=f"openai/c2k-{uuid.uuid4().hex}")
drop: Final = scenario.model(model=f"openai/c2d-{uuid.uuid4().hex}")
_wait_until_served(gateway, target)
_wait_until_served(gateway, keep)
_wait_until_served(gateway, drop)
assert _create_fallback(gateway, scenario, keep, [target]).status_code == 200
assert _create_fallback(gateway, scenario, drop, [target]).status_code == 200
removed: Final = gateway.request("DELETE", f"/fallback/{drop}", params={"fallback_type": "general"})
assert removed.status_code == 200, removed.text
covered: Final = _covered_models(_stored_fallbacks())
assert keep in covered, covered
assert drop not in covered, covered
gone: Final = gateway.request("GET", f"/fallback/{drop}", params={"fallback_type": "general"})
assert gone.status_code == 404, gone.text
def test_self_fallback_and_unknown_fallback_models_are_rejected(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(model=f"openai/c3-{uuid.uuid4().hex}")
_wait_until_served(gateway, model)
itself: Final = gateway.request("POST", "/fallback", _fallback_body(model, [model]))
assert itself.status_code == 400, itself.text
unknown_target: Final = gateway.request(
"POST", "/fallback", _fallback_body(model, [f"c3-missing-{uuid.uuid4().hex}"])
)
assert unknown_target.status_code == 400, unknown_target.text
def test_malformed_requests_are_rejected_and_the_proxy_stays_up(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(model=f"openai/c4-{uuid.uuid4().hex}")
fallback: Final = scenario.model(model=f"openai/c4f-{uuid.uuid4().hex}")
_wait_until_served(gateway, model)
_wait_until_served(gateway, fallback)
malformed: Final = (
{"model": model, "fallback_models": [], "fallback_type": "general"},
{"model": model, "fallback_type": "general"},
{"model": model, "fallback_models": [fallback], "fallback_type": "nonsense"},
{"model": 123, "fallback_models": [fallback], "fallback_type": "general"},
)
for body in malformed:
assert gateway.request("POST", "/fallback", body).status_code in (400, 422), body
healthy: Final = _create_fallback(gateway, scenario, model, [fallback])
assert healthy.status_code == 200, healthy.text
listed: Final = gateway.request("GET", "/v1/models")
assert listed.status_code == 200, listed.text
def test_team_key_is_forbidden_on_create_and_delete(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
model: Final = scenario.model(model=f"openai/c5-{uuid.uuid4().hex}", model_info={"team_id": team})
fallback: Final = scenario.model(model=f"openai/c5f-{uuid.uuid4().hex}", model_info={"team_id": team})
team_key: Final = scenario.key(team_id=team)
creating: Final = gateway.request("POST", "/fallback", _fallback_body(model, [fallback]), key=team_key)
assert creating.status_code in REFUSED, creating.text
assert model not in _covered_models(_stored_fallbacks())
admitted: Final = _create_fallback(gateway, scenario, model, [fallback])
assert admitted.status_code == 200, admitted.text
deleting: Final = gateway.request(
"DELETE", f"/fallback/{model}", params={"fallback_type": "general"}, key=team_key
)
assert deleting.status_code in REFUSED, deleting.text
assert model in _covered_models(_stored_fallbacks())
def test_context_window_and_content_policy_types_persist(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
target: Final = scenario.model(model=f"openai/c6t-{uuid.uuid4().hex}")
windowed: Final = scenario.model(model=f"openai/c6w-{uuid.uuid4().hex}")
policy: Final = scenario.model(model=f"openai/c6p-{uuid.uuid4().hex}")
_wait_until_served(gateway, target)
_wait_until_served(gateway, windowed)
_wait_until_served(gateway, policy)
window_create: Final = _create_fallback(gateway, scenario, windowed, [target], "context_window")
assert window_create.status_code == 200, window_create.text
window_read: Final = eventually(
lambda: gateway.request("GET", f"/fallback/{windowed}", params={"fallback_type": "context_window"}),
lambda response: response.status_code == 200 and target in _fallback_models(response),
seconds=30,
return_last_on_timeout=True,
)
assert window_read.status_code == 200 and target in _fallback_models(window_read), window_read.text
policy_create: Final = _create_fallback(gateway, scenario, policy, [target], "content_policy")
assert policy_create.status_code == 200, policy_create.text
covered: Final = _covered_models(_stored_fallbacks("content_policy_fallbacks"))
assert policy in covered, covered
def test_fallback_write_survives_a_redis_outage(gateway: Gateway, tmp_path: Path) -> None:
with owned_redis(tmp_path) as cache:
overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)}
with owned_proxy_process(gateway, tmp_path, overrides, workers=2, database_setup=()) as owned:
with owned.gateway.scenario() as scenario:
target: Final = scenario.model(model=f"openai/x1t-{uuid.uuid4().hex}")
first: Final = scenario.model(model=f"openai/x1a-{uuid.uuid4().hex}")
second: Final = scenario.model(model=f"openai/x1b-{uuid.uuid4().hex}")
third: Final = scenario.model(model=f"openai/x1c-{uuid.uuid4().hex}")
for model in (target, first, second, third):
_wait_until_served(owned.gateway, model)
assert _create_fallback(owned.gateway, scenario, first, [target]).status_code == 200
cache.stop()
degraded: Final = tuple(
_create_fallback(owned.gateway, scenario, model, [target]) for model in (second, third)
)
assert all(response.status_code == 200 for response in degraded), [r.text for r in degraded]
covered: Final = _covered_models(_stored_fallbacks())
assert {first, second, third} <= covered, covered
assert EVICTION_WARNING in owned.log.read_text(), owned.log.read_text()[-3000:]
cache.start()
def test_fallback_write_survives_a_worker_kill(gateway: Gateway, tmp_path: Path) -> None:
with owned_proxy_process(gateway, tmp_path, {}, workers=2, database_setup=()) as owned:
with owned.gateway.scenario() as scenario:
target: Final = scenario.model(model=f"openai/x2t-{uuid.uuid4().hex}")
before: Final = scenario.model(model=f"openai/x2a-{uuid.uuid4().hex}")
after: Final = scenario.model(model=f"openai/x2b-{uuid.uuid4().hex}")
last: Final = scenario.model(model=f"openai/x2c-{uuid.uuid4().hex}")
for model in (target, before, after, last):
_wait_until_served(owned.gateway, model)
assert _create_fallback(owned.gateway, scenario, before, [target]).status_code == 200
victim: Final = eventually(lambda: _workers(owned), lambda workers: len(workers) == 2)[0]
victim.kill()
fresh: Final = scenario.cleanups.enter_context(
httpx.Client(base_url=owned.gateway.client.base_url, timeout=15, trust_env=False)
)
survivor: Final = Gateway(fresh, owned.gateway.key, owned.gateway.upstream_url)
degraded: Final = tuple(_create_fallback(survivor, scenario, model, [target]) for model in (after, last))
assert all(response.status_code == 200 for response in degraded), [r.text for r in degraded]
covered: Final = _covered_models(_stored_fallbacks())
assert {before, after, last} <= covered, covered

View file

@ -8,17 +8,27 @@ Tests:
4. Validation tests (invalid models, duplicate fallbacks, etc.)
"""
import copy
import json
from collections.abc import AsyncIterator
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm import Router
from litellm.proxy.utils import evict_config_param, get_config_param
from litellm.proxy.management_endpoints.fallback_management_endpoints import (
FallbackCreateRequest,
FallbackDeleteResponse,
FallbackResponse,
create_fallback,
delete_fallback,
get_fallback,
)
from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict
class TestFallbackCreateRequest:
@ -104,6 +114,268 @@ class TestFallbackCreateRequest:
assert request.fallback_type == "content_policy"
TEAM_ID: Final = "team-a"
PRIMARY_INTERNAL_NAME: Final = f"model_name_{TEAM_ID}_1a6437cb-4cab-432c-8099-1d7411731a8b"
FALLBACK_INTERNAL_NAME: Final = f"model_name_{TEAM_ID}_83151607-5556-4bbf-ac65-c474dbc64eba"
PRIMARY_DEPLOYMENT_ID: Final = "team-primary-id"
FALLBACK_DEPLOYMENT_ID: Final = "team-fallback-id"
FallbackRules = list[dict[str, list[str]] | dict[str, object] | str]
RouterSettings = dict[str, FallbackRules | None]
def _litellm_params(mock_response: str | None) -> LiteLLMParamsTypedDict:
if mock_response is None:
return {"model": "openai/gpt-5.4-mini", "api_key": "fake"}
return {"model": "openai/gpt-5.4-mini", "api_key": "fake", "mock_response": mock_response}
def _team_scoped_deployment(
internal_name: str, public_name: str, deployment_id: str, mock_response: str | None = None
) -> DeploymentTypedDict:
return {
"model_name": internal_name,
"litellm_params": _litellm_params(mock_response),
"model_info": {"id": deployment_id, "team_id": TEAM_ID, "team_public_model_name": public_name},
}
def _team_router(primary_mock_response: str | None = None, fallback_mock_response: str | None = None) -> Router:
return Router(
model_list=[
{"model_name": "gpt-5.4-mini", "litellm_params": _litellm_params(None)},
_team_scoped_deployment(
PRIMARY_INTERNAL_NAME, "team-primary", PRIMARY_DEPLOYMENT_ID, primary_mock_response
),
_team_scoped_deployment(
FALLBACK_INTERNAL_NAME, "team-fallback", FALLBACK_DEPLOYMENT_ID, fallback_mock_response
),
],
num_retries=0,
)
class TestCreateFallbackForTeamScopedModels:
"""POST /fallback takes the public name a caller invokes a team-scoped model by, not only the stored internal one"""
@pytest.fixture
def router(self) -> Router:
return _team_router()
@pytest.fixture
def prisma_client(self) -> MagicMock:
client: Final = MagicMock()
client.db.litellm_config.upsert = AsyncMock()
return client
@pytest.fixture
def proxy_config(self) -> MagicMock:
config: Final = MagicMock()
config.get_config = AsyncMock(return_value={"router_settings": {}})
return config
async def _create(
self, request: FallbackCreateRequest, router: Router, prisma_client: MagicMock, proxy_config: MagicMock
) -> FallbackResponse:
with (
patch("litellm.proxy.proxy_server.llm_router", router),
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.proxy_config", proxy_config),
patch("litellm.proxy.proxy_server.store_model_in_db", True),
):
return await create_fallback(request, MagicMock())
async def test_public_names_create_a_rule_keyed_on_the_public_name(
self, router: Router, prisma_client: MagicMock, proxy_config: MagicMock
) -> None:
request: Final = FallbackCreateRequest(model="team-primary", fallback_models=["team-fallback"])
response: Final = await self._create(request, router, prisma_client, proxy_config)
assert response.model == "team-primary"
assert response.fallback_models == ["team-fallback"]
assert router.fallbacks == [{"team-primary": ["team-fallback"]}]
persisted: Final = json.loads(
prisma_client.db.litellm_config.upsert.call_args.kwargs["data"]["create"]["param_value"]
)
assert persisted["fallbacks"] == [{"team-primary": ["team-fallback"]}]
async def test_internal_names_keep_working(
self, router: Router, prisma_client: MagicMock, proxy_config: MagicMock
) -> None:
request: Final = FallbackCreateRequest(model=PRIMARY_INTERNAL_NAME, fallback_models=[FALLBACK_INTERNAL_NAME])
response: Final = await self._create(request, router, prisma_client, proxy_config)
assert router.fallbacks == [{PRIMARY_INTERNAL_NAME: [FALLBACK_INTERNAL_NAME]}]
assert response.model == PRIMARY_INTERNAL_NAME
async def test_public_fallback_target_behind_a_gateway_primary(
self, router: Router, prisma_client: MagicMock, proxy_config: MagicMock
) -> None:
request: Final = FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"])
await self._create(request, router, prisma_client, proxy_config)
assert router.fallbacks == [{"gpt-5.4-mini": ["team-fallback"]}]
async def test_unknown_name_is_rejected_and_the_error_names_the_public_names(
self, router: Router, prisma_client: MagicMock, proxy_config: MagicMock
) -> None:
request: Final = FallbackCreateRequest(model="team-missing", fallback_models=["team-fallback"])
with pytest.raises(HTTPException) as exc_info:
await self._create(request, router, prisma_client, proxy_config)
assert exc_info.value.status_code == 404
assert {"team-primary", "team-fallback", "gpt-5.4-mini"} <= set(exc_info.value.detail["available_models"])
async def test_a_team_request_fails_over_to_the_public_name_fallback(
self, prisma_client: MagicMock, proxy_config: MagicMock
) -> None:
router: Final = _team_router(primary_mock_response="litellm.RateLimitError", fallback_mock_response="pong")
request: Final = FallbackCreateRequest(model="team-primary", fallback_models=["team-fallback"])
await self._create(request, router, prisma_client, proxy_config)
response: Final = await router.acompletion(
model="team-primary",
messages=[{"role": "user", "content": "ping"}],
metadata={"user_api_key_team_id": TEAM_ID},
)
assert response.choices[0].message.content == "pong"
assert response._hidden_params["model_id"] == FALLBACK_DEPLOYMENT_ID
def _config_row(router_settings: RouterSettings) -> SimpleNamespace:
return SimpleNamespace(param_name="router_settings", param_value=router_settings)
class _StoredRouterSettings:
"""The LiteLLM_Config router_settings row, with the proxy objects that read it through the config cache"""
def __init__(self, router_settings: RouterSettings | None) -> None:
self.row: SimpleNamespace | None = None if router_settings is None else _config_row(router_settings)
self.prisma_client: Final = MagicMock()
self.prisma_client.get_generic_data = AsyncMock(side_effect=lambda **_: self.row)
self.prisma_client.db.litellm_config.upsert = AsyncMock(side_effect=self._upsert)
self.proxy_config: Final = MagicMock()
self.proxy_config.get_config = AsyncMock(side_effect=self._get_config)
async def _upsert(self, where: dict[str, str], data: dict[str, dict[str, str]]) -> None:
self.row = _config_row(json.loads(data["update"]["param_value"]))
async def _get_config(self) -> dict[str, RouterSettings]:
row: Final = await get_config_param(self.prisma_client, "router_settings")
return {"router_settings": copy.deepcopy(row.param_value) if row is not None else {}}
def written_by_another_instance(self, router_settings: RouterSettings) -> None:
self.row = _config_row(router_settings)
def stored_fallbacks(self) -> FallbackRules:
assert self.row is not None
return self.row.param_value["fallbacks"]
async def cached_fallbacks(self) -> FallbackRules:
row: Final = await get_config_param(self.prisma_client, "router_settings")
return row.param_value["fallbacks"]
TEAM_RULE: Final = {"team-primary": ["team-fallback"]}
GATEWAY_RULE: Final = {"gpt-5.4-mini": ["team-fallback"]}
MALFORMED_GATEWAY_RULE: Final = {"gpt-5.4-mini": "team-fallback"}
NON_STANDARD_RULES: Final = [
"claude-3-haiku",
{"model": "gpt-5.4-mini", "messages": [{"role": "user", "content": "retry"}]},
]
class TestFallbackWritesSeeTheLatestStoredRules:
"""A write reads the rules the database holds now and leaves no stale copy in the config cache behind"""
@pytest.fixture(autouse=True)
async def clean_config_cache(self) -> AsyncIterator[None]:
await evict_config_param("router_settings")
yield
await evict_config_param("router_settings")
async def _create(self, request: FallbackCreateRequest, stored: _StoredRouterSettings) -> FallbackResponse:
with (
patch("litellm.proxy.proxy_server.llm_router", _team_router()),
patch("litellm.proxy.proxy_server.prisma_client", stored.prisma_client),
patch("litellm.proxy.proxy_server.proxy_config", stored.proxy_config),
patch("litellm.proxy.proxy_server.store_model_in_db", True),
):
return await create_fallback(request, MagicMock())
async def _delete(self, model: str, stored: _StoredRouterSettings) -> FallbackDeleteResponse:
with (
patch("litellm.proxy.proxy_server.llm_router", _team_router()),
patch("litellm.proxy.proxy_server.prisma_client", stored.prisma_client),
patch("litellm.proxy.proxy_server.proxy_config", stored.proxy_config),
patch("litellm.proxy.proxy_server.store_model_in_db", True),
):
return await delete_fallback(model, "general", MagicMock())
async def test_a_second_create_keeps_the_first_rule(self) -> None:
stored: Final = _StoredRouterSettings(None)
await self._create(FallbackCreateRequest(model="team-primary", fallback_models=["team-fallback"]), stored)
await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored)
assert stored.stored_fallbacks() == [TEAM_RULE, GATEWAY_RULE]
assert await stored.cached_fallbacks() == [TEAM_RULE, GATEWAY_RULE]
async def test_a_create_keeps_a_rule_another_instance_stored_since_this_one_last_read(self) -> None:
stored: Final = _StoredRouterSettings({})
await get_config_param(stored.prisma_client, "router_settings")
stored.written_by_another_instance({"fallbacks": [TEAM_RULE]})
await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored)
assert stored.stored_fallbacks() == [TEAM_RULE, GATEWAY_RULE]
async def test_a_delete_keeps_a_rule_another_instance_stored_since_this_one_last_read(self) -> None:
stored: Final = _StoredRouterSettings({"fallbacks": [TEAM_RULE]})
await get_config_param(stored.prisma_client, "router_settings")
stored.written_by_another_instance({"fallbacks": [TEAM_RULE, GATEWAY_RULE]})
await self._delete("team-primary", stored)
assert stored.stored_fallbacks() == [GATEWAY_RULE]
assert await stored.cached_fallbacks() == [GATEWAY_RULE]
async def test_writes_keep_the_non_standard_rules_the_router_accepts(self) -> None:
stored: Final = _StoredRouterSettings({"fallbacks": [*NON_STANDARD_RULES, TEAM_RULE]})
await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored)
assert stored.stored_fallbacks() == [*NON_STANDARD_RULES, TEAM_RULE, GATEWAY_RULE]
await self._delete("team-primary", stored)
assert stored.stored_fallbacks() == [*NON_STANDARD_RULES, GATEWAY_RULE]
async def test_a_create_replaces_a_same_key_rule_whatever_its_targets_shape(self) -> None:
stored: Final = _StoredRouterSettings({"fallbacks": [MALFORMED_GATEWAY_RULE, TEAM_RULE]})
await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored)
assert stored.stored_fallbacks() == [GATEWAY_RULE, TEAM_RULE]
async def test_a_delete_removes_a_same_key_rule_whatever_its_targets_shape(self) -> None:
stored: Final = _StoredRouterSettings({"fallbacks": [MALFORMED_GATEWAY_RULE, TEAM_RULE]})
await self._delete("gpt-5.4-mini", stored)
assert stored.stored_fallbacks() == [TEAM_RULE]
async def test_a_create_treats_a_null_rule_list_as_empty(self) -> None:
stored: Final = _StoredRouterSettings({"fallbacks": None})
await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored)
assert stored.stored_fallbacks() == [GATEWAY_RULE]
@pytest.mark.asyncio
class TestCreateFallback:
"""Test the create_fallback endpoint"""
@ -113,6 +385,7 @@ class TestCreateFallback:
"""Create a mock router"""
router = MagicMock()
router.model_names = {"gpt-3.5-turbo", "gpt-4", "claude-3-haiku"}
router.team_public_model_names = frozenset()
router.fallbacks = []
router.context_window_fallbacks = []
router.content_policy_fallbacks = []

View file

@ -2,7 +2,7 @@ import asyncio
import json
import os
from datetime import datetime
from typing import Any, Dict, List, Optional, Union
from typing import Any, Dict, Final, List, Optional, Union
from unittest.mock import Mock
import pytest
@ -3329,3 +3329,22 @@ def test_handle_exception_on_proxy_preserves_auth_error_status_code():
result = handle_exception_on_proxy(auth_error)
assert int(result.code) == 401, f"Expected 401, got {result.code}"
class _RedisDown:
async def async_delete_cache(self, key: str) -> None:
raise ConnectionError("redis is down")
async def test_evict_config_param_clears_the_local_layer_and_survives_a_redis_outage() -> None:
from litellm.caching.caching import DualCache
from litellm.proxy.utils import _config_cache_key, evict_config_param
cache: Final = DualCache(redis_cache=_RedisDown())
await cache.in_memory_cache.async_set_cache(
_config_cache_key("router_settings"), {"param_name": "router_settings", "param_value": {"fallbacks": []}}
)
await evict_config_param("router_settings", cache=cache)
assert await cache.in_memory_cache.async_get_cache(_config_cache_key("router_settings")) is None

View file

@ -12,8 +12,9 @@ Symbols pinned here:
from __future__ import annotations
import logging
from types import SimpleNamespace
from typing import Any, List
from typing import Any, Final, List
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -214,14 +215,23 @@ async def test_invalidate_config_param_evicts_from_cache(
@pytest.mark.asyncio
async def test_invalidate_config_param_propagates_cache_error(
_swap_config_cache: Any,
async def test_invalidate_config_param_survives_a_cache_error(
_swap_config_cache: Any, caplog: pytest.LogCaptureFixture
) -> None:
_swap_config_cache.async_delete_cache = AsyncMock(
side_effect=ConnectionError("redis down")
)
with pytest.raises(ConnectionError):
caplog.set_level(logging.WARNING, logger="LiteLLM Proxy")
utils_mod.verbose_proxy_logger.addHandler(caplog.handler)
try:
await invalidate_config_param("p5")
finally:
utils_mod.verbose_proxy_logger.removeHandler(caplog.handler)
actual: Final = {
"delete_calls": _swap_config_cache.async_delete_cache.await_count,
"warned": "config cache eviction of p5 failed: redis down" in caplog.text,
}
assert actual == {"delete_calls": 1, "warned": True}
@pytest.mark.asyncio