From adc9b64076748b19912cad71bc1872eab684c8c3 Mon Sep 17 00:00:00 2001 From: mateo Date: Mon, 10 Aug 2026 16:55:20 +0000 Subject: [PATCH] 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> --- .../get_fallback_generalizations.py | 91 ++++++++++-------- tests/local_testing/test_get_llm_provider.py | 2 +- .../test_get_fallback_generalizations.py | 92 ++++++++++--------- .../test_get_model_cost_map.py | 8 +- ...azure_anthropic_messages_transformation.py | 2 +- .../test_anthropic_claude3_transformation.py | 2 +- ...artner_models_anthropic_messages_config.py | 2 +- 7 files changed, 108 insertions(+), 91 deletions(-) diff --git a/litellm/litellm_core_utils/get_fallback_generalizations.py b/litellm/litellm_core_utils/get_fallback_generalizations.py index 739067f3fd8..ed0fd511ae7 100644 --- a/litellm/litellm_core_utils/get_fallback_generalizations.py +++ b/litellm/litellm_core_utils/get_fallback_generalizations.py @@ -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 diff --git a/tests/local_testing/test_get_llm_provider.py b/tests/local_testing/test_get_llm_provider.py index 7ee003753b5..4a72cd0e586 100644 --- a/tests/local_testing/test_get_llm_provider.py +++ b/tests/local_testing/test_get_llm_provider.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_get_fallback_generalizations.py b/tests/test_litellm/litellm_core_utils/test_get_fallback_generalizations.py index 66f50c15a07..b4902160dd3 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_fallback_generalizations.py +++ b/tests/test_litellm/litellm_core_utils/test_get_fallback_generalizations.py @@ -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" diff --git a/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py b/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py index 1a02ae2166a..3d0435129cb 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py +++ b/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py @@ -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"] diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py index ba3223ef587..398cb79d3d3 100644 --- a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py +++ b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py @@ -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, diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index 754e72a7536..a4d6245bc5f 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -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, diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py index 57b37ee8227..317c80380c7 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py @@ -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,