mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
5133485009
commit
5d207d85ee
6 changed files with 743 additions and 51 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue