mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
2946c9a183
commit
adc9b64076
7 changed files with 108 additions and 91 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue