refactor(model_prices): load fallback_generalizations from its own file and URL

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-08-10 16:55:20 +00:00
parent 2946c9a183
commit adc9b64076
7 changed files with 108 additions and 91 deletions

View file

@ -1,20 +1,21 @@
"""
Loads the fallback generalization rules that ``fallback_generalizations`` compiles.
The rules live in their own ``litellm/fallback_generalizations.json`` file, bundled with
the package and served from the repository, so adding a model to the cost map never
shuffles the rule block around and the two can be published independently.
The rules live in their own ``litellm/fallback_generalizations.json``, bundled with the
package and served from the repository, so adding a model to the cost map never shuffles
the rule block around and the two files can be published independently.
Resolution order mirrors the model cost map: remote first, bundled copy on any failure.
Set ``LITELLM_LOCAL_FALLBACK_GENERALIZATIONS=True`` (or ``LITELLM_LOCAL_MODEL_COST_MAP=True``,
Resolution mirrors the model cost map: remote first, bundled copy on any failure. Set
``LITELLM_LOCAL_FALLBACK_GENERALIZATIONS=True`` (or ``LITELLM_LOCAL_MODEL_COST_MAP=True``,
which already means "no registry network calls") to skip the fetch, and
``LITELLM_FALLBACK_GENERALIZATIONS_URL`` to point at a different remote file.
"""
import json
import os
from collections.abc import Callable, Mapping
from importlib.resources import files
from typing import Final
from typing import Final, TypeAlias
import httpx
@ -23,57 +24,72 @@ from litellm.litellm_core_utils.fallback_generalizations import (
set_fallback_generalizations,
)
FallbackRules: TypeAlias = tuple[Mapping[str, object], ...]
JsonFetcher: TypeAlias = Callable[[str], object]
RULES_FIELD: Final = "rules"
LOCAL_FILE_NAME: Final = "fallback_generalizations.json"
LOCAL_ONLY_ENV_VARS: Final = ("LITELLM_LOCAL_FALLBACK_GENERALIZATIONS", "LITELLM_LOCAL_MODEL_COST_MAP")
FETCH_TIMEOUT_SECONDS: Final = 5
def _is_local_only() -> bool:
return any(
os.getenv(name, "").lower() == "true"
for name in ("LITELLM_LOCAL_FALLBACK_GENERALIZATIONS", "LITELLM_LOCAL_MODEL_COST_MAP")
)
return any(os.getenv(name, "").lower() == "true" for name in LOCAL_ONLY_ENV_VARS)
def load_local_fallback_generalizations() -> list:
"""Return the rules from the copy bundled with the package."""
content: Final = json.loads(files("litellm").joinpath(LOCAL_FILE_NAME).read_text(encoding="utf-8"))
return _extract_rules(content, source=LOCAL_FILE_NAME)
def _extract_rules(content: object, source: str) -> list:
if not isinstance(content, dict):
def _extract_rules(content: object, source: str) -> FallbackRules:
if not isinstance(content, Mapping):
verbose_logger.warning(
"LiteLLM: fallback generalizations from %s are not a dict (type=%s); ignoring them.",
"LiteLLM: fallback generalizations from %s are not an object (type=%s); ignoring them.",
source,
type(content).__name__,
)
return []
rules: Final = content.get(RULES_FIELD)
if not isinstance(rules, list):
return ()
raw: Final = content.get(RULES_FIELD)
if not isinstance(raw, list):
verbose_logger.warning(
"LiteLLM: fallback generalizations from %s have no '%s' list; ignoring them.",
source,
RULES_FIELD,
)
return []
return ()
rules: Final = tuple(rule for rule in raw if isinstance(rule, Mapping))
if len(rules) != len(raw):
verbose_logger.warning(
"LiteLLM: %s fallback generalization rule(s) from %s are not objects; skipping them.",
len(raw) - len(rules),
source,
)
return rules
def fetch_remote_fallback_generalizations(url: str, timeout: int = 5) -> list:
"""Fetch the rules from ``url``, raising on network/parse errors for the caller to handle."""
response: Final = httpx.get(url, timeout=timeout)
def load_local_fallback_generalizations() -> FallbackRules:
"""Return the rules from the copy bundled with the package."""
content: Final = json.loads(files("litellm").joinpath(LOCAL_FILE_NAME).read_text(encoding="utf-8"))
return _extract_rules(content, source=LOCAL_FILE_NAME)
def fetch_remote_json(url: str) -> object:
"""Fetch and parse ``url``, raising on any network, status, or parse failure."""
response: Final = httpx.get(url, timeout=FETCH_TIMEOUT_SECONDS)
response.raise_for_status()
return _extract_rules(response.json(), source=url)
return response.json()
def get_fallback_generalizations(url: str) -> list:
"""Return the active rule list, falling back to the bundled copy when the fetch is unusable."""
def get_fallback_generalizations(url: str, fetch: JsonFetcher = fetch_remote_json) -> FallbackRules:
"""Return the rules to install, falling back to the bundled copy whenever the fetch is unusable.
Routing and capability fallbacks must survive an unreachable, 404, or truncated remote
file, so anything that does not yield at least one rule resolves to the bundled copy.
"""
if _is_local_only():
return load_local_fallback_generalizations()
try:
remote: Final = fetch_remote_fallback_generalizations(url)
except Exception as e:
remote: Final = _extract_rules(fetch(url), source=url)
except (httpx.HTTPError, ValueError) as e:
verbose_logger.warning(
"LiteLLM: Failed to fetch fallback generalizations from %s: %s. Falling back to local backup.",
url,
@ -81,16 +97,13 @@ def get_fallback_generalizations(url: str) -> list:
)
return load_local_fallback_generalizations()
if not remote:
return load_local_fallback_generalizations()
return remote
return remote or load_local_fallback_generalizations()
def install_fallback_generalizations() -> list:
def install_fallback_generalizations(url: str | None = None, fetch: JsonFetcher = fetch_remote_json) -> FallbackRules:
"""Load the rules and install them into the generalizations registry."""
from litellm import fallback_generalizations_url
import litellm
rules: Final = get_fallback_generalizations(url=fallback_generalizations_url)
set_fallback_generalizations(rules)
rules: Final = get_fallback_generalizations(url=url or litellm.fallback_generalizations_url, fetch=fetch)
set_fallback_generalizations(list(rules))
return rules

View file

@ -500,7 +500,7 @@ def shipped_generalizations():
)
previous = list(get_fallback_generalization_rules())
rules = load_local_fallback_generalizations()
rules = list(load_local_fallback_generalizations())
set_fallback_generalizations(rules)
try:
yield rules

View file

@ -1,6 +1,6 @@
"""
Tests for loading the standalone fallback generalizations file: remote wins when it is
usable, the bundled copy covers every failure, and the local-only env flags skip HTTP.
Tests for loading the standalone fallback generalizations file: the remote copy wins when
it is usable, the bundled copy covers every failure, and the local-only env flags skip HTTP.
"""
import json
@ -12,7 +12,6 @@ import pytest
sys.path.insert(0, os.path.abspath("../../.."))
import litellm
from litellm.litellm_core_utils.fallback_generalizations import (
get_fallback_generalization_rules,
match_routing_generalization,
@ -26,13 +25,13 @@ from litellm.litellm_core_utils.get_fallback_generalizations import (
_URL = "https://example.invalid/fallback_generalizations.json"
_REMOTE_RULES = [
_REMOTE_RULES = (
{
"name": "remote-only",
"pattern": r"^widget-",
"model_info": {"litellm_provider": "openai"},
}
]
},
)
@pytest.fixture(autouse=True)
@ -51,48 +50,59 @@ def restore_generalizations():
set_fallback_generalizations(previous)
def _serve(monkeypatch, response: httpx.Response | Exception) -> dict:
"""Point httpx.get at a canned response and count the calls made to it."""
calls = {"count": 0}
def _serving(payload: object):
"""A fetcher returning ``payload``, or raising it when it is an exception."""
def fake_get(url, timeout=None):
calls["count"] += 1
if isinstance(response, Exception):
raise response
return response
def fetch(url: str) -> object:
if isinstance(payload, Exception):
raise payload
return payload
monkeypatch.setattr(
"litellm.litellm_core_utils.get_fallback_generalizations.httpx.get", fake_get
)
return calls
return fetch
def test_remote_rules_are_used_when_the_fetch_succeeds(monkeypatch):
_serve(monkeypatch, httpx.Response(200, json={"rules": _REMOTE_RULES}, request=httpx.Request("GET", _URL)))
assert get_fallback_generalizations(url=_URL) == _REMOTE_RULES
def _never_called(url: str) -> object:
raise AssertionError(f"fetched {url} while in local-only mode")
def test_network_failure_falls_back_to_the_bundled_file(monkeypatch):
def test_remote_rules_are_used_when_the_fetch_succeeds():
rules = get_fallback_generalizations(url=_URL, fetch=_serving({"rules": list(_REMOTE_RULES)}))
assert rules == _REMOTE_RULES
def test_network_failure_falls_back_to_the_bundled_file():
"""A rules fetch must never take routing down: an unreachable host keeps the shipped rules."""
_serve(monkeypatch, httpx.ConnectError("connection refused"))
assert get_fallback_generalizations(url=_URL) == load_local_fallback_generalizations()
rules = get_fallback_generalizations(url=_URL, fetch=_serving(httpx.ConnectError("connection refused")))
assert rules == load_local_fallback_generalizations()
def test_http_error_falls_back_to_the_bundled_file(monkeypatch):
"""The file is 404 on older branches until this lands, so a 404 must not wipe the rules."""
_serve(monkeypatch, httpx.Response(404, request=httpx.Request("GET", _URL)))
assert get_fallback_generalizations(url=_URL) == load_local_fallback_generalizations()
def test_http_error_falls_back_to_the_bundled_file():
"""The URL 404s on every branch that predates this file, so a 404 must not wipe the rules."""
error = httpx.HTTPStatusError(
"404", request=httpx.Request("GET", _URL), response=httpx.Response(404)
)
rules = get_fallback_generalizations(url=_URL, fetch=_serving(error))
assert rules == load_local_fallback_generalizations()
def test_unparseable_remote_body_falls_back_to_the_bundled_file():
rules = get_fallback_generalizations(url=_URL, fetch=_serving(json.JSONDecodeError("boom", "", 0)))
assert rules == load_local_fallback_generalizations()
@pytest.mark.parametrize(
"payload",
[{"rules": {}}, {}, []],
ids=["rules_not_a_list", "no_rules_key", "not_an_object"],
[{"rules": {}}, {"rules": []}, {}, [], "rules"],
ids=["rules_not_a_list", "rules_empty", "no_rules_key", "not_an_object", "not_json_object"],
)
def test_malformed_remote_payload_falls_back_to_the_bundled_file(monkeypatch, payload):
def test_malformed_remote_payload_falls_back_to_the_bundled_file(payload):
"""A truncated or reshaped remote file must not silently disable every rule."""
_serve(monkeypatch, httpx.Response(200, json=payload, request=httpx.Request("GET", _URL)))
assert get_fallback_generalizations(url=_URL) == load_local_fallback_generalizations()
assert get_fallback_generalizations(url=_URL, fetch=_serving(payload)) == load_local_fallback_generalizations()
def test_non_object_rules_are_skipped_without_dropping_the_valid_ones():
payload = {"rules": ["nonsense", *_REMOTE_RULES]}
assert get_fallback_generalizations(url=_URL, fetch=_serving(payload)) == _REMOTE_RULES
@pytest.mark.parametrize(
@ -100,21 +110,15 @@ def test_malformed_remote_payload_falls_back_to_the_bundled_file(monkeypatch, pa
["LITELLM_LOCAL_FALLBACK_GENERALIZATIONS", "LITELLM_LOCAL_MODEL_COST_MAP"],
)
def test_local_only_env_skips_the_fetch(monkeypatch, env_name):
"""Both flags mean "no registry network calls", so neither may issue an HTTP request."""
"""Both flags mean "no registry network calls", so neither may issue a request."""
monkeypatch.setenv(env_name, "True")
calls = _serve(monkeypatch, httpx.Response(200, json={"rules": _REMOTE_RULES}, request=httpx.Request("GET", _URL)))
assert get_fallback_generalizations(url=_URL) == load_local_fallback_generalizations()
assert calls["count"] == 0
assert get_fallback_generalizations(url=_URL, fetch=_never_called) == load_local_fallback_generalizations()
def test_install_uses_the_configured_url_and_compiles_the_rules(
monkeypatch, restore_generalizations
):
monkeypatch.setattr(litellm, "fallback_generalizations_url", _URL)
_serve(monkeypatch, httpx.Response(200, json={"rules": _REMOTE_RULES}, request=httpx.Request("GET", _URL)))
def test_install_compiles_the_loaded_rules(restore_generalizations):
installed = install_fallback_generalizations(url=_URL, fetch=_serving({"rules": list(_REMOTE_RULES)}))
assert install_fallback_generalizations() == _REMOTE_RULES
assert installed == _REMOTE_RULES
assert match_routing_generalization("widget-9") == "openai"

View file

@ -118,7 +118,7 @@ def test_finalize_drops_a_stale_block_and_installs_the_standalone_rules(monkeypa
assert FALLBACK_GENERALIZATIONS_KEY not in finalized
assert match_routing_generalization("widget-9") is None
assert get_fallback_generalization_rules() == load_local_fallback_generalizations()
assert tuple(get_fallback_generalization_rules()) == load_local_fallback_generalizations()
finally:
set_fallback_generalizations(previous)
@ -142,7 +142,7 @@ def test_shipped_backup_carries_the_claude_routing_rules():
"""The bundled rules file must ship the Claude routing rules so a fresh install
(or an offline fallback) routes unknown Claude models without code changes.
Bedrock-syntax ids must hit the bedrock rule before the bare-id Anthropic rule."""
rules = load_local_fallback_generalizations()
rules = list(load_local_fallback_generalizations())
names = [r.get("name") for r in rules]
assert names.index("bedrock-claude-ids") < names.index("anthropic-claude-ids")
@ -163,7 +163,7 @@ def test_shipped_routing_rules_never_match_through_an_unrecognized_namespace():
``bedrockz/anthropic.claude-...`` resolve to bedrock and slip through a
``bedrock/*`` key, so every shipped routing rule must anchor to the start of
the name and never match an id carrying an unrecognized namespace prefix."""
rules = load_local_fallback_generalizations()
rules = list(load_local_fallback_generalizations())
routing_rules = [r for r in rules if "litellm_provider" in r["model_info"]]
assert routing_rules
@ -202,7 +202,7 @@ def test_shipped_backup_marks_claude_4_6_plus_adaptive_not_4_0():
baseline block is never duplicated across rules and no rule needs ``extends``."""
backup = GetModelCostMap.load_local_model_cost_map()
rules = load_local_fallback_generalizations()
rules = list(load_local_fallback_generalizations())
baseline_rule = next(r for r in rules if r.get("name") == "claude-family-baseline")
adaptive_rule = next(r for r in rules if r.get("name") == "claude-adaptive-thinking")
assert "supports_adaptive_thinking" not in baseline_rule["model_info"]

View file

@ -482,7 +482,7 @@ def test_azure_claude_4_8_plus_cost_map_entries_carry_mid_conversation_system_fl
load_local_fallback_generalizations,
)
rules = load_local_fallback_generalizations()
rules = list(load_local_fallback_generalizations())
rule_pattern = next(
(r["pattern"] for r in rules if r["name"] == "claude-mid-conversation-system"),
None,

View file

@ -2174,7 +2174,7 @@ def test_bedrock_claude_4_8_plus_cost_map_entries_carry_mid_conversation_system_
load_local_fallback_generalizations,
)
rules = load_local_fallback_generalizations()
rules = list(load_local_fallback_generalizations())
pattern = re.compile(
next(r["pattern"] for r in rules if r["name"] == "claude-mid-conversation-system"),
re.IGNORECASE,

View file

@ -661,7 +661,7 @@ def test_vertex_claude_4_8_plus_cost_map_entries_carry_mid_conversation_system_f
load_local_fallback_generalizations,
)
rules = load_local_fallback_generalizations()
rules = list(load_local_fallback_generalizations())
rule_pattern = next(
(r["pattern"] for r in rules if r["name"] == "claude-mid-conversation-system"),
None,