diff --git a/litellm/litellm_core_utils/get_model_cost_map.py b/litellm/litellm_core_utils/get_model_cost_map.py index b4a721d8e97..de873d2afaa 100644 --- a/litellm/litellm_core_utils/get_model_cost_map.py +++ b/litellm/litellm_core_utils/get_model_cost_map.py @@ -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 diff --git a/tests/llm_translation/test_model_cost_map_resilience.py b/tests/llm_translation/test_model_cost_map_resilience.py index 2c9f32028c8..61e375eabeb 100644 --- a/tests/llm_translation/test_model_cost_map_resilience.py +++ b/tests/llm_translation/test_model_cost_map_resilience.py @@ -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 )