mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
validate_model_cost_map
This commit is contained in:
parent
b80fefc268
commit
a0a6b7196a
2 changed files with 55 additions and 45 deletions
|
|
@ -25,10 +25,14 @@ class GetModelCostMap:
|
|||
"""
|
||||
Handles fetching, validating, and loading the model cost map.
|
||||
|
||||
All methods are static — no instance state is needed. This class groups
|
||||
the helpers that support `get_model_cost_map()` into a single namespace.
|
||||
The backup model count is cached on first access so that validation
|
||||
never needs to parse the full backup JSON — only the count is needed
|
||||
for the shrinkage check. The full backup is only loaded when it must
|
||||
be *returned* as a fallback.
|
||||
"""
|
||||
|
||||
_backup_model_count: int = -1 # -1 = not yet loaded
|
||||
|
||||
@staticmethod
|
||||
def load_local_model_cost_map() -> dict:
|
||||
"""Load the local backup model cost map bundled with the package."""
|
||||
|
|
@ -39,6 +43,14 @@ class GetModelCostMap:
|
|||
)
|
||||
return content
|
||||
|
||||
@classmethod
|
||||
def _get_backup_model_count(cls) -> int:
|
||||
"""Return the number of models in the local backup (cached)."""
|
||||
if cls._backup_model_count < 0:
|
||||
backup = cls.load_local_model_cost_map()
|
||||
cls._backup_model_count = len(backup)
|
||||
return cls._backup_model_count
|
||||
|
||||
@staticmethod
|
||||
def _check_is_valid_dict(fetched_map: dict) -> bool:
|
||||
"""Check 1: fetched map is a non-empty dict."""
|
||||
|
|
@ -59,10 +71,11 @@ class GetModelCostMap:
|
|||
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
@classmethod
|
||||
def _check_model_count_not_reduced(
|
||||
cls,
|
||||
fetched_map: dict,
|
||||
backup_map: dict,
|
||||
backup_model_count: int,
|
||||
min_model_count: int = MODEL_COST_MAP_MIN_MODEL_COUNT,
|
||||
max_shrink_pct: float = MODEL_COST_MAP_MAX_SHRINK_PERCENT,
|
||||
) -> bool:
|
||||
|
|
@ -79,25 +92,25 @@ class GetModelCostMap:
|
|||
)
|
||||
return False
|
||||
|
||||
backup_count = len(backup_map) if isinstance(backup_map, dict) else 0
|
||||
if backup_count > 0 and fetched_count < backup_count * max_shrink_pct:
|
||||
if backup_model_count > 0 and fetched_count < backup_model_count * max_shrink_pct:
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Fetched model cost map shrank significantly "
|
||||
"(fetched=%d, backup=%d, threshold=%.0f%%). "
|
||||
"This may indicate a corrupted upstream file. "
|
||||
"Falling back to local backup.",
|
||||
fetched_count,
|
||||
backup_count,
|
||||
backup_model_count,
|
||||
max_shrink_pct * 100,
|
||||
)
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
@classmethod
|
||||
def validate_model_cost_map(
|
||||
cls,
|
||||
fetched_map: dict,
|
||||
backup_map: dict,
|
||||
backup_model_count: int,
|
||||
min_model_count: int = MODEL_COST_MAP_MIN_MODEL_COUNT,
|
||||
max_shrink_pct: float = MODEL_COST_MAP_MAX_SHRINK_PERCENT,
|
||||
) -> bool:
|
||||
|
|
@ -113,12 +126,12 @@ class GetModelCostMap:
|
|||
|
||||
Returns True if all checks pass, False otherwise.
|
||||
"""
|
||||
if not GetModelCostMap._check_is_valid_dict(fetched_map):
|
||||
if not cls._check_is_valid_dict(fetched_map):
|
||||
return False
|
||||
|
||||
if not GetModelCostMap._check_model_count_not_reduced(
|
||||
if not cls._check_model_count_not_reduced(
|
||||
fetched_map=fetched_map,
|
||||
backup_map=backup_map,
|
||||
backup_model_count=backup_model_count,
|
||||
min_model_count=min_model_count,
|
||||
max_shrink_pct=max_shrink_pct,
|
||||
):
|
||||
|
|
@ -146,6 +159,10 @@ def get_model_cost_map(url: str) -> dict:
|
|||
1. If ``LITELLM_LOCAL_MODEL_COST_MAP`` is set, uses the local backup only.
|
||||
2. Otherwise fetches from ``url``, validates integrity, and falls back
|
||||
to the local backup on any failure.
|
||||
|
||||
The backup model count is cached in ``GetModelCostMap`` so validation
|
||||
only costs a cheap integer comparison — the full backup JSON is only
|
||||
parsed when it needs to be *returned* as a fallback.
|
||||
"""
|
||||
if (
|
||||
os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", False)
|
||||
|
|
@ -153,23 +170,8 @@ def get_model_cost_map(url: str) -> dict:
|
|||
):
|
||||
return GetModelCostMap.load_local_model_cost_map()
|
||||
|
||||
backup_map = GetModelCostMap.load_local_model_cost_map()
|
||||
|
||||
try:
|
||||
content = GetModelCostMap.fetch_remote_model_cost_map(url)
|
||||
|
||||
# Validate fetched JSON integrity before using it
|
||||
if not GetModelCostMap.validate_model_cost_map(
|
||||
fetched_map=content, backup_map=backup_map
|
||||
):
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Fetched model cost map failed integrity check. "
|
||||
"Using local backup instead. url=%s",
|
||||
url,
|
||||
)
|
||||
return backup_map
|
||||
|
||||
return content
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Failed to fetch remote model cost map from %s: %s. "
|
||||
|
|
@ -177,4 +179,18 @@ def get_model_cost_map(url: str) -> dict:
|
|||
url,
|
||||
str(e),
|
||||
)
|
||||
return backup_map
|
||||
return GetModelCostMap.load_local_model_cost_map()
|
||||
|
||||
# Validate fetched JSON integrity — uses cached backup count, no file I/O
|
||||
if not GetModelCostMap.validate_model_cost_map(
|
||||
fetched_map=content,
|
||||
backup_model_count=GetModelCostMap._get_backup_model_count(),
|
||||
):
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Fetched model cost map failed integrity check. "
|
||||
"Using local backup instead. url=%s",
|
||||
url,
|
||||
)
|
||||
return GetModelCostMap.load_local_model_cost_map()
|
||||
|
||||
return content
|
||||
|
|
|
|||
|
|
@ -59,40 +59,37 @@ class TestCheckModelCountNotReduced:
|
|||
small_map = {f"model-{i}": {} for i in range(5)}
|
||||
assert (
|
||||
GetModelCostMap._check_model_count_not_reduced(
|
||||
fetched_map=small_map, backup_map={}, min_model_count=10
|
||||
fetched_map=small_map, backup_model_count=0, min_model_count=10
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_should_reject_significant_shrinkage(self):
|
||||
"""Fetched map that shrunk >50% vs backup should fail."""
|
||||
backup = {f"model-{i}": {} for i in range(100)}
|
||||
fetched = {f"model-{i}": {} for i in range(40)} # 40% of backup
|
||||
fetched = {f"model-{i}": {} for i in range(40)} # 40% of 100
|
||||
assert (
|
||||
GetModelCostMap._check_model_count_not_reduced(
|
||||
fetched_map=fetched, backup_map=backup, min_model_count=10
|
||||
fetched_map=fetched, backup_model_count=100, min_model_count=10
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_should_accept_when_above_threshold(self):
|
||||
"""Fetched map at 60% of backup (above 50% threshold) should pass."""
|
||||
backup = {f"model-{i}": {} for i in range(100)}
|
||||
fetched = {f"model-{i}": {} for i in range(60)}
|
||||
assert (
|
||||
GetModelCostMap._check_model_count_not_reduced(
|
||||
fetched_map=fetched, backup_map=backup, min_model_count=10
|
||||
fetched_map=fetched, backup_model_count=100, min_model_count=10
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_should_accept_growth(self):
|
||||
"""Fetched map larger than backup should pass."""
|
||||
backup = {f"model-{i}": {} for i in range(100)}
|
||||
fetched = {f"model-{i}": {} for i in range(120)}
|
||||
assert (
|
||||
GetModelCostMap._check_model_count_not_reduced(
|
||||
fetched_map=fetched, backup_map=backup, min_model_count=10
|
||||
fetched_map=fetched, backup_model_count=100, min_model_count=10
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
|
@ -102,7 +99,7 @@ class TestCheckModelCountNotReduced:
|
|||
fetched = {f"model-{i}": {} for i in range(15)}
|
||||
assert (
|
||||
GetModelCostMap._check_model_count_not_reduced(
|
||||
fetched_map=fetched, backup_map={}, min_model_count=10
|
||||
fetched_map=fetched, backup_model_count=0, min_model_count=10
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
|
@ -113,41 +110,38 @@ class TestValidateModelCostMap:
|
|||
|
||||
def test_should_reject_non_dict(self):
|
||||
"""Non-dict should fail at check 1."""
|
||||
assert GetModelCostMap.validate_model_cost_map(fetched_map="not a dict", backup_map={}) is False
|
||||
assert GetModelCostMap.validate_model_cost_map(fetched_map="not a dict", backup_model_count=0) is False
|
||||
|
||||
def test_should_reject_empty_map(self):
|
||||
"""Empty dict should fail at check 1."""
|
||||
assert GetModelCostMap.validate_model_cost_map(fetched_map={}, backup_map={}) is False
|
||||
assert GetModelCostMap.validate_model_cost_map(fetched_map={}, backup_model_count=0) is False
|
||||
|
||||
def test_should_reject_significant_shrinkage(self):
|
||||
"""Should fail at check 2 (shrinkage)."""
|
||||
backup = {f"model-{i}": {} for i in range(100)}
|
||||
fetched = {f"model-{i}": {} for i in range(40)}
|
||||
assert (
|
||||
GetModelCostMap.validate_model_cost_map(
|
||||
fetched_map=fetched, backup_map=backup, min_model_count=10
|
||||
fetched_map=fetched, backup_model_count=100, min_model_count=10
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_should_accept_valid_map(self):
|
||||
"""Should pass both checks."""
|
||||
backup = {f"model-{i}": {} for i in range(100)}
|
||||
fetched = {f"model-{i}": {} for i in range(120)}
|
||||
assert (
|
||||
GetModelCostMap.validate_model_cost_map(
|
||||
fetched_map=fetched, backup_map=backup, min_model_count=10
|
||||
fetched_map=fetched, backup_model_count=100, min_model_count=10
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_should_accept_equal_size_map(self):
|
||||
"""Equal size should pass both checks."""
|
||||
backup = {f"model-{i}": {} for i in range(100)}
|
||||
fetched = {f"model-{i}": {} for i in range(100)}
|
||||
assert (
|
||||
GetModelCostMap.validate_model_cost_map(
|
||||
fetched_map=fetched, backup_map=backup, min_model_count=10
|
||||
fetched_map=fetched, backup_model_count=100, min_model_count=10
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue