From 5d207d85ee547bc323d704bc6aaab901aa0be8b6 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 8 Oct 2026 22:38:38 +0000 Subject: [PATCH] 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__ 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> --- .../fallback_management_endpoints.py | 93 +++-- litellm/proxy/utils.py | 10 +- .../test_fallback_management_public_names.py | 379 ++++++++++++++++++ .../test_fallback_management_endpoints.py | 273 +++++++++++++ tests/unit/proxy/test_proxy_utils.py | 21 +- .../test_config_param_cache.py | 18 +- 6 files changed, 743 insertions(+), 51 deletions(-) create mode 100644 tests/integration/management/test_fallback_management_public_names.py diff --git a/litellm/proxy/management_endpoints/fallback_management_endpoints.py b/litellm/proxy/management_endpoints/fallback_management_endpoints.py index 543538f5edd..460b282a9f6 100644 --- a/litellm/proxy/management_endpoints/fallback_management_endpoints.py +++ b/litellm/proxy/management_endpoints/fallback_management_endpoints.py @@ -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) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 4ff23624c1b..4743970a526 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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: diff --git a/tests/integration/management/test_fallback_management_public_names.py b/tests/integration/management/test_fallback_management_public_names.py new file mode 100644 index 00000000000..b54465c1fd0 --- /dev/null +++ b/tests/integration/management/test_fallback_management_public_names.py @@ -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 diff --git a/tests/unit/proxy/test_fallback_management_endpoints.py b/tests/unit/proxy/test_fallback_management_endpoints.py index 054dafbf5a7..3ee1b64580d 100644 --- a/tests/unit/proxy/test_fallback_management_endpoints.py +++ b/tests/unit/proxy/test_fallback_management_endpoints.py @@ -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 = [] diff --git a/tests/unit/proxy/test_proxy_utils.py b/tests/unit/proxy/test_proxy_utils.py index ba61e393281..92bb163d6ec 100644 --- a/tests/unit/proxy/test_proxy_utils.py +++ b/tests/unit/proxy/test_proxy_utils.py @@ -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 diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py b/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py index 0d408de9ec6..cdaf5c1ee3b 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py @@ -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