diff --git a/.github/workflows/test-model-map.yaml b/.github/workflows/test-model-map.yaml index ae5ac402e23..f7d18291c53 100644 --- a/.github/workflows/test-model-map.yaml +++ b/.github/workflows/test-model-map.yaml @@ -3,6 +3,12 @@ name: Validate model_prices_and_context_window.json on: pull_request: branches: [ main ] + push: + branches: [ main ] + paths: + - 'model_prices_and_context_window.json' + - 'scripts/validate_model_cost_map.py' + - '.github/workflows/test-model-map.yaml' jobs: validate-model-prices-json: @@ -10,6 +16,45 @@ jobs: steps: - uses: actions/checkout@v4 - - name: Validate model_prices_and_context_window.json + - name: Set up Python + uses: actions/setup-python@v4 + with: + python-version: '3.11' + + # Step 1: Basic JSON syntax validation (fast fail) + - name: Validate JSON syntax run: | + echo "Checking JSON syntax..." jq empty model_prices_and_context_window.json + echo "✅ JSON syntax is valid" + + # Step 2: Comprehensive schema/semantic validation + - name: Validate model cost map schema + run: | + echo "Running comprehensive validation..." + python scripts/validate_model_cost_map.py model_prices_and_context_window.json + + # Step 3: Verify backup file is also valid + - name: Validate backup model cost map + run: | + echo "Validating backup file..." + python scripts/validate_model_cost_map.py litellm/model_prices_and_context_window_backup.json + + # Step 4: Ensure backup stays in sync (optional check) + - name: Check backup freshness + run: | + echo "Checking if backup matches main file..." + # Compare entry counts - backup should be reasonably close to main + MAIN_COUNT=$(jq 'keys | length' model_prices_and_context_window.json) + BACKUP_COUNT=$(jq 'keys | length' litellm/model_prices_and_context_window_backup.json) + echo "Main file entries: $MAIN_COUNT" + echo "Backup file entries: $BACKUP_COUNT" + + # Warn if backup is significantly out of date (more than 100 entries behind) + DIFF=$((MAIN_COUNT - BACKUP_COUNT)) + if [ $DIFF -gt 100 ]; then + echo "⚠️ WARNING: Backup file is $DIFF entries behind main file" + echo "Consider updating the backup during release" + else + echo "✅ Backup file is reasonably up to date" + fi diff --git a/litellm/litellm_core_utils/get_model_cost_map.py b/litellm/litellm_core_utils/get_model_cost_map.py index 9b86f4ca2f0..4a1ccde1fe7 100644 --- a/litellm/litellm_core_utils/get_model_cost_map.py +++ b/litellm/litellm_core_utils/get_model_cost_map.py @@ -9,39 +9,184 @@ export LITELLM_LOCAL_MODEL_COST_MAP=True """ import os +import sys +from typing import Optional, Tuple import httpx +# Minimum number of model entries expected in a valid model cost map +# This prevents accepting an empty or truncated response +MIN_MODEL_ENTRIES = 100 + +# Sample of required fields that should be present in model entries +# At minimum, every model entry should have a litellm_provider +REQUIRED_MODEL_FIELDS = {"litellm_provider"} + +# Valid modes for model entries +VALID_MODES = { + "audio_speech", + "audio_transcription", + "chat", + "completion", + "embedding", + "image_edit", + "image_generation", + "moderation", + "ocr", + "rerank", + "responses", + "search", + "vector_store", + "video_generation", +} + + +def _get_logger(): + """Get the verbose logger if available, otherwise return a no-op logger.""" + try: + from litellm._logging import verbose_logger + return verbose_logger + except ImportError: + import logging + return logging.getLogger(__name__) + + +def validate_model_cost_map(data: dict) -> Tuple[bool, Optional[str]]: + """ + Validate that fetched data is a valid model cost map. + + Returns: + Tuple of (is_valid, error_message) + - (True, None) if valid + - (False, "reason") if invalid + """ + if not isinstance(data, dict): + return False, f"Expected dict, got {type(data).__name__}" + + # Check minimum number of entries (should have hundreds of models) + # Exclude sample_spec from count + model_count = len([k for k in data.keys() if k != "sample_spec"]) + if model_count < MIN_MODEL_ENTRIES: + return False, f"Too few model entries: {model_count} (expected at least {MIN_MODEL_ENTRIES})" + + # Validate structure of model entries (sample check for performance) + errors = [] + entries_checked = 0 + max_entries_to_check = 50 # Check a sample for performance + + for key, value in data.items(): + if key == "sample_spec": + continue + + entries_checked += 1 + if entries_checked > max_entries_to_check: + break + + if not isinstance(value, dict): + errors.append(f"Entry '{key}' is not a dict") + continue + + # Check required fields + if "litellm_provider" not in value: + errors.append(f"Entry '{key}' missing required field 'litellm_provider'") + + # Validate mode if present + if "mode" in value and value["mode"] not in VALID_MODES: + errors.append(f"Entry '{key}' has invalid mode: '{value['mode']}'") + + # Validate cost fields are numeric if present + for cost_field in ["input_cost_per_token", "output_cost_per_token"]: + if cost_field in value and not isinstance(value[cost_field], (int, float)): + errors.append(f"Entry '{key}' has non-numeric {cost_field}: {type(value[cost_field]).__name__}") + + if errors: + # Return first few errors to avoid huge error messages + error_sample = errors[:5] + if len(errors) > 5: + error_sample.append(f"... and {len(errors) - 5} more errors") + return False, "; ".join(error_sample) + + return True, None + + +def _load_local_model_cost_map() -> dict: + """Load the local backup model cost map.""" + from importlib.resources import files + import json + + content = json.loads( + files("litellm") + .joinpath("model_prices_and_context_window_backup.json") + .read_text(encoding="utf-8") + ) + return content + + def get_model_cost_map(url: str) -> dict: - if ( - os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", False) - or os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", False) == "True" - ): - from importlib.resources import files - import json - - content = json.loads( - files("litellm") - .joinpath("model_prices_and_context_window_backup.json") - .read_text(encoding="utf-8") - ) - return content + """ + Get the model cost map, either from remote URL or local backup. + + Priority: + 1. If LITELLM_LOCAL_MODEL_COST_MAP=True, use local backup + 2. Try to fetch from remote URL + 3. Validate the remote response before accepting + 4. Fall back to local backup if remote fails or is invalid + + Args: + url: The URL to fetch the model cost map from + + Returns: + dict: The model cost map + """ + logger = _get_logger() + + # Check if local-only mode is enabled + local_mode = os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", "").lower() in ("true", "1", "yes") + + if local_mode: + logger.debug("LITELLM_LOCAL_MODEL_COST_MAP is set, using local model cost map") + return _load_local_model_cost_map() + # Try to fetch from remote try: response = httpx.get( url, timeout=5 ) # set a 5 second timeout for the get request response.raise_for_status() # Raise an exception if the request is unsuccessful content = response.json() + + # Validate the response before accepting + is_valid, error_message = validate_model_cost_map(content) + + if not is_valid: + logger.warning( + f"Remote model cost map failed validation: {error_message}. " + f"Falling back to local backup. " + f"Set LITELLM_LOCAL_MODEL_COST_MAP=True to always use local backup." + ) + return _load_local_model_cost_map() + + logger.debug(f"Successfully loaded remote model cost map with {len(content)} entries") return content - except Exception: - from importlib.resources import files - import json - - content = json.loads( - files("litellm") - .joinpath("model_prices_and_context_window_backup.json") - .read_text(encoding="utf-8") + + except httpx.TimeoutException: + logger.warning( + "Timeout fetching remote model cost map, falling back to local backup. " + "Set LITELLM_LOCAL_MODEL_COST_MAP=True to always use local backup." ) - return content + return _load_local_model_cost_map() + except httpx.HTTPStatusError as e: + logger.warning( + f"HTTP error fetching remote model cost map (status {e.response.status_code}), " + f"falling back to local backup. " + f"Set LITELLM_LOCAL_MODEL_COST_MAP=True to always use local backup." + ) + return _load_local_model_cost_map() + except Exception as e: + logger.warning( + f"Error fetching remote model cost map: {type(e).__name__}: {e}. " + f"Falling back to local backup. " + f"Set LITELLM_LOCAL_MODEL_COST_MAP=True to always use local backup." + ) + return _load_local_model_cost_map() diff --git a/scripts/validate_model_cost_map.py b/scripts/validate_model_cost_map.py new file mode 100644 index 00000000000..072bbfeb763 --- /dev/null +++ b/scripts/validate_model_cost_map.py @@ -0,0 +1,317 @@ +#!/usr/bin/env python3 +""" +Comprehensive validation script for model_prices_and_context_window.json + +This script validates the model cost map to prevent malformed entries from +breaking LiteLLM deployments. Run this in CI before merging any changes to +the model cost map. + +Usage: + python scripts/validate_model_cost_map.py [path_to_json] + +If no path is provided, defaults to model_prices_and_context_window.json +in the repository root. + +Exit codes: + 0 - Validation passed + 1 - Validation failed +""" + +import json +import sys +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple + + +# Validation configuration +MIN_MODEL_ENTRIES = 100 +MAX_COST_PER_TOKEN = 1.0 # $1 per token would be absurdly high + +VALID_MODES = { + "audio_speech", + "audio_transcription", + "chat", + "completion", + "embedding", + "image_edit", + "image_generation", + "moderation", + "ocr", + "rerank", + "responses", + "search", + "vector_store", + "video_generation", +} + +# Fields that should be numeric if present +NUMERIC_FIELDS = { + "input_cost_per_token", + "output_cost_per_token", + "input_cost_per_character", + "output_cost_per_character", + "input_cost_per_image", + "output_cost_per_image", + "input_cost_per_audio_token", + "output_cost_per_audio_token", + "input_cost_per_second", + "output_cost_per_second", + "max_tokens", + "max_input_tokens", + "max_output_tokens", +} + +# Fields that should be boolean if present +BOOLEAN_FIELDS = { + "supports_function_calling", + "supports_parallel_function_calling", + "supports_vision", + "supports_audio_input", + "supports_audio_output", + "supports_prompt_caching", + "supports_response_schema", + "supports_system_messages", + "supports_reasoning", + "supports_web_search", +} + + +class ValidationError: + def __init__(self, model: str, field: str, message: str, severity: str = "error"): + self.model = model + self.field = field + self.message = message + self.severity = severity # "error" or "warning" + + def __str__(self): + return f"[{self.severity.upper()}] {self.model}: {self.field} - {self.message}" + + +def validate_entry(model_name: str, entry: Any) -> List[ValidationError]: + """Validate a single model entry.""" + errors = [] + + # Skip sample_spec - it's documentation + if model_name == "sample_spec": + return errors + + # Entry must be a dict + if not isinstance(entry, dict): + errors.append(ValidationError( + model_name, "entry", f"Expected dict, got {type(entry).__name__}" + )) + return errors + + # Required field: litellm_provider + if "litellm_provider" not in entry: + errors.append(ValidationError( + model_name, "litellm_provider", "Missing required field" + )) + elif not isinstance(entry["litellm_provider"], str): + errors.append(ValidationError( + model_name, "litellm_provider", + f"Expected string, got {type(entry['litellm_provider']).__name__}" + )) + elif not entry["litellm_provider"].strip(): + errors.append(ValidationError( + model_name, "litellm_provider", "Cannot be empty string" + )) + + # Validate mode if present + if "mode" in entry: + if not isinstance(entry["mode"], str): + errors.append(ValidationError( + model_name, "mode", + f"Expected string, got {type(entry['mode']).__name__}" + )) + elif entry["mode"] not in VALID_MODES: + errors.append(ValidationError( + model_name, "mode", + f"Invalid mode '{entry['mode']}'. Valid modes: {sorted(VALID_MODES)}" + )) + + # Validate numeric fields + for field in NUMERIC_FIELDS: + if field in entry: + value = entry[field] + if isinstance(value, str): + # Some fields in sample_spec have string descriptions + if model_name != "sample_spec": + errors.append(ValidationError( + model_name, field, + f"Expected numeric, got string: '{value}'" + )) + elif not isinstance(value, (int, float)): + errors.append(ValidationError( + model_name, field, + f"Expected numeric, got {type(value).__name__}" + )) + elif isinstance(value, (int, float)) and value < 0: + errors.append(ValidationError( + model_name, field, + f"Cannot be negative: {value}", + severity="warning" + )) + + # Validate boolean fields + for field in BOOLEAN_FIELDS: + if field in entry: + value = entry[field] + if not isinstance(value, bool): + errors.append(ValidationError( + model_name, field, + f"Expected boolean, got {type(value).__name__}: {value}" + )) + + # Validate cost sanity (catch obviously wrong values) + for cost_field in ["input_cost_per_token", "output_cost_per_token"]: + if cost_field in entry: + value = entry[cost_field] + if isinstance(value, (int, float)) and value > MAX_COST_PER_TOKEN: + errors.append(ValidationError( + model_name, cost_field, + f"Suspiciously high cost: ${value}/token (max expected: ${MAX_COST_PER_TOKEN})", + severity="warning" + )) + + # Validate max_tokens relationships + max_tokens = entry.get("max_tokens") + max_input = entry.get("max_input_tokens") + max_output = entry.get("max_output_tokens") + + if max_input is not None and max_output is not None: + if isinstance(max_input, (int, float)) and isinstance(max_output, (int, float)): + if max_input < 0 or max_output < 0: + pass # Already caught above + elif max_input == 0 and max_output == 0: + errors.append(ValidationError( + model_name, "max_tokens", + "Both max_input_tokens and max_output_tokens are 0", + severity="warning" + )) + + return errors + + +def validate_model_cost_map(data: Dict[str, Any]) -> Tuple[List[ValidationError], Dict[str, int]]: + """ + Validate the entire model cost map. + + Returns: + Tuple of (errors, stats) + """ + all_errors = [] + stats = { + "total_entries": 0, + "valid_entries": 0, + "entries_with_errors": 0, + "entries_with_warnings": 0, + "total_errors": 0, + "total_warnings": 0, + } + + if not isinstance(data, dict): + all_errors.append(ValidationError( + "ROOT", "type", f"Expected dict at root, got {type(data).__name__}" + )) + return all_errors, stats + + # Count entries (excluding sample_spec) + model_entries = {k: v for k, v in data.items() if k != "sample_spec"} + stats["total_entries"] = len(model_entries) + + # Check minimum entries + if stats["total_entries"] < MIN_MODEL_ENTRIES: + all_errors.append(ValidationError( + "ROOT", "count", + f"Too few model entries: {stats['total_entries']} (minimum: {MIN_MODEL_ENTRIES})" + )) + + # Validate each entry + for model_name, entry in data.items(): + errors = validate_entry(model_name, entry) + + entry_errors = [e for e in errors if e.severity == "error"] + entry_warnings = [e for e in errors if e.severity == "warning"] + + if entry_errors: + stats["entries_with_errors"] += 1 + elif entry_warnings: + stats["entries_with_warnings"] += 1 + else: + stats["valid_entries"] += 1 + + stats["total_errors"] += len(entry_errors) + stats["total_warnings"] += len(entry_warnings) + + all_errors.extend(errors) + + return all_errors, stats + + +def main(): + # Determine file path + if len(sys.argv) > 1: + json_path = Path(sys.argv[1]) + else: + # Default to repository root + script_dir = Path(__file__).parent + json_path = script_dir.parent / "model_prices_and_context_window.json" + + if not json_path.exists(): + print(f"ERROR: File not found: {json_path}") + sys.exit(1) + + print(f"Validating: {json_path}") + print("-" * 60) + + # Load and parse JSON + try: + with open(json_path, "r", encoding="utf-8") as f: + data = json.load(f) + except json.JSONDecodeError as e: + print(f"ERROR: Invalid JSON: {e}") + sys.exit(1) + + # Validate + errors, stats = validate_model_cost_map(data) + + # Print statistics + print(f"Total model entries: {stats['total_entries']}") + print(f"Valid entries: {stats['valid_entries']}") + print(f"Entries with errors: {stats['entries_with_errors']}") + print(f"Entries with warnings: {stats['entries_with_warnings']}") + print("-" * 60) + + # Print errors (errors first, then warnings) + error_list = [e for e in errors if e.severity == "error"] + warning_list = [e for e in errors if e.severity == "warning"] + + if error_list: + print(f"\nERRORS ({len(error_list)}):") + for error in error_list[:50]: # Limit output + print(f" {error}") + if len(error_list) > 50: + print(f" ... and {len(error_list) - 50} more errors") + + if warning_list: + print(f"\nWARNINGS ({len(warning_list)}):") + for warning in warning_list[:20]: # Limit output + print(f" {warning}") + if len(warning_list) > 20: + print(f" ... and {len(warning_list) - 20} more warnings") + + # Exit with appropriate code + if error_list: + print(f"\n❌ VALIDATION FAILED: {len(error_list)} error(s) found") + sys.exit(1) + elif warning_list: + print(f"\n⚠️ VALIDATION PASSED with {len(warning_list)} warning(s)") + sys.exit(0) + else: + print("\n✅ VALIDATION PASSED") + sys.exit(0) + + +if __name__ == "__main__": + main() 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 new file mode 100644 index 00000000000..e6a3bd9bc55 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py @@ -0,0 +1,295 @@ +""" +Unit tests for get_model_cost_map.py validation logic. + +These tests ensure that malformed model cost maps are rejected before +they can break LLM calls. +""" + +import json +import os +import pytest +from unittest.mock import patch, MagicMock + +# Set local mode for tests by default to avoid network calls +os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + +from litellm.litellm_core_utils.get_model_cost_map import ( + validate_model_cost_map, + get_model_cost_map, + _load_local_model_cost_map, + MIN_MODEL_ENTRIES, + VALID_MODES, +) + + +class TestValidateModelCostMap: + """Tests for the validate_model_cost_map function.""" + + def test_valid_model_cost_map(self): + """Test that a valid model cost map passes validation.""" + # Create a minimal valid map with enough entries + valid_map = { + "sample_spec": {"litellm_provider": "example"}, + } + # Add enough model entries to pass minimum check + for i in range(MIN_MODEL_ENTRIES + 10): + valid_map[f"model-{i}"] = { + "litellm_provider": "openai", + "mode": "chat", + "input_cost_per_token": 0.00001, + "output_cost_per_token": 0.00002, + } + + is_valid, error = validate_model_cost_map(valid_map) + assert is_valid is True + assert error is None + + def test_rejects_non_dict(self): + """Test that non-dict input is rejected.""" + is_valid, error = validate_model_cost_map([]) + assert is_valid is False + assert "Expected dict" in error + + is_valid, error = validate_model_cost_map("string") + assert is_valid is False + assert "Expected dict" in error + + def test_rejects_too_few_entries(self): + """Test that maps with too few entries are rejected.""" + small_map = { + "model-1": {"litellm_provider": "openai"}, + "model-2": {"litellm_provider": "openai"}, + } + + is_valid, error = validate_model_cost_map(small_map) + assert is_valid is False + assert "Too few model entries" in error + + def test_rejects_missing_litellm_provider(self): + """Test that entries without litellm_provider are flagged.""" + invalid_map = {"sample_spec": {}} + for i in range(MIN_MODEL_ENTRIES + 10): + invalid_map[f"model-{i}"] = { + "mode": "chat", # Missing litellm_provider + } + + is_valid, error = validate_model_cost_map(invalid_map) + assert is_valid is False + assert "litellm_provider" in error + + def test_rejects_invalid_mode(self): + """Test that invalid mode values are flagged.""" + invalid_map = {"sample_spec": {}} + for i in range(MIN_MODEL_ENTRIES + 10): + invalid_map[f"model-{i}"] = { + "litellm_provider": "openai", + "mode": "invalid_mode_xyz", # Invalid mode + } + + is_valid, error = validate_model_cost_map(invalid_map) + assert is_valid is False + assert "invalid mode" in error.lower() + + def test_rejects_non_numeric_cost(self): + """Test that non-numeric cost values are flagged.""" + invalid_map = {"sample_spec": {}} + for i in range(MIN_MODEL_ENTRIES + 10): + invalid_map[f"model-{i}"] = { + "litellm_provider": "openai", + "input_cost_per_token": "not_a_number", # Should be numeric + } + + is_valid, error = validate_model_cost_map(invalid_map) + assert is_valid is False + assert "non-numeric" in error.lower() + + def test_accepts_all_valid_modes(self): + """Test that all valid modes are accepted.""" + for mode in VALID_MODES: + valid_map = {"sample_spec": {}} + for i in range(MIN_MODEL_ENTRIES + 10): + valid_map[f"model-{i}"] = { + "litellm_provider": "openai", + "mode": mode, + } + + is_valid, error = validate_model_cost_map(valid_map) + assert is_valid is True, f"Mode '{mode}' should be valid but got error: {error}" + + def test_sample_spec_is_ignored(self): + """Test that sample_spec entry doesn't count toward minimum.""" + # Only sample_spec - should fail due to too few entries + only_sample = {"sample_spec": {"litellm_provider": "example"}} + + is_valid, error = validate_model_cost_map(only_sample) + assert is_valid is False + assert "Too few" in error + + +class TestGetModelCostMap: + """Tests for the get_model_cost_map function.""" + + def test_local_mode_uses_backup(self): + """Test that local mode uses the backup file.""" + with patch.dict(os.environ, {"LITELLM_LOCAL_MODEL_COST_MAP": "True"}): + result = get_model_cost_map("https://fake-url.com/model.json") + + # Should return a dict with many entries + assert isinstance(result, dict) + assert len(result) > MIN_MODEL_ENTRIES + + def test_local_mode_variations(self): + """Test that various truthy values enable local mode.""" + for value in ["True", "true", "TRUE", "1", "yes", "YES"]: + with patch.dict(os.environ, {"LITELLM_LOCAL_MODEL_COST_MAP": value}): + # Should not make network request + with patch("httpx.get") as mock_get: + result = get_model_cost_map("https://fake-url.com/model.json") + mock_get.assert_not_called() + assert isinstance(result, dict) + + @patch("httpx.get") + def test_fallback_on_network_error(self, mock_get): + """Test fallback to local on network errors.""" + with patch.dict(os.environ, {"LITELLM_LOCAL_MODEL_COST_MAP": ""}): + mock_get.side_effect = Exception("Network error") + + result = get_model_cost_map("https://fake-url.com/model.json") + + # Should fall back to local backup + assert isinstance(result, dict) + assert len(result) > MIN_MODEL_ENTRIES + + @patch("httpx.get") + def test_fallback_on_invalid_response(self, mock_get): + """Test fallback when remote returns invalid data.""" + with patch.dict(os.environ, {"LITELLM_LOCAL_MODEL_COST_MAP": ""}): + # Mock a response with too few entries + mock_response = MagicMock() + mock_response.json.return_value = {"only": "one entry"} + mock_response.raise_for_status.return_value = None + mock_get.return_value = mock_response + + result = get_model_cost_map("https://fake-url.com/model.json") + + # Should fall back to local backup due to validation failure + assert isinstance(result, dict) + assert len(result) > MIN_MODEL_ENTRIES + + @patch("httpx.get") + def test_fallback_on_malformed_entry(self, mock_get): + """Test fallback when remote has malformed entries.""" + with patch.dict(os.environ, {"LITELLM_LOCAL_MODEL_COST_MAP": ""}): + # Create invalid map (missing litellm_provider) + invalid_map = {} + for i in range(MIN_MODEL_ENTRIES + 10): + invalid_map[f"model-{i}"] = {"mode": "chat"} # Missing litellm_provider + + mock_response = MagicMock() + mock_response.json.return_value = invalid_map + mock_response.raise_for_status.return_value = None + mock_get.return_value = mock_response + + result = get_model_cost_map("https://fake-url.com/model.json") + + # Should fall back to local backup due to validation failure + assert isinstance(result, dict) + # Local backup should have litellm_provider in entries + for key, value in list(result.items())[:5]: + if key != "sample_spec": + assert "litellm_provider" in value + + +class TestLocalBackup: + """Tests for the local backup file.""" + + def test_local_backup_is_valid(self): + """Test that the bundled local backup passes validation.""" + backup = _load_local_model_cost_map() + + is_valid, error = validate_model_cost_map(backup) + assert is_valid is True, f"Local backup failed validation: {error}" + + def test_local_backup_has_required_models(self): + """Test that local backup has common models.""" + backup = _load_local_model_cost_map() + + # Check for some common models that should always be present + common_models = [ + "gpt-4", + "gpt-3.5-turbo", + "claude-3-opus-20240229", + ] + + for model in common_models: + assert model in backup, f"Common model '{model}' missing from backup" + + def test_local_backup_entries_have_provider(self): + """Test that all entries in local backup have litellm_provider.""" + backup = _load_local_model_cost_map() + + for key, value in backup.items(): + if key == "sample_spec": + continue + assert "litellm_provider" in value, f"Entry '{key}' missing litellm_provider" + + +class TestRealWorldScenarios: + """Tests simulating real-world failure scenarios.""" + + @patch("httpx.get") + def test_scenario_empty_json_response(self, mock_get): + """Simulate GitHub returning empty JSON object.""" + with patch.dict(os.environ, {"LITELLM_LOCAL_MODEL_COST_MAP": ""}): + mock_response = MagicMock() + mock_response.json.return_value = {} + mock_response.raise_for_status.return_value = None + mock_get.return_value = mock_response + + result = get_model_cost_map("https://fake-url.com/model.json") + + # Should fall back, not return empty dict + assert len(result) > MIN_MODEL_ENTRIES + + @patch("httpx.get") + def test_scenario_truncated_response(self, mock_get): + """Simulate GitHub returning truncated response.""" + with patch.dict(os.environ, {"LITELLM_LOCAL_MODEL_COST_MAP": ""}): + # Only 5 entries - way below minimum + truncated = {f"model-{i}": {"litellm_provider": "openai"} for i in range(5)} + + mock_response = MagicMock() + mock_response.json.return_value = truncated + mock_response.raise_for_status.return_value = None + mock_get.return_value = mock_response + + result = get_model_cost_map("https://fake-url.com/model.json") + + # Should fall back due to too few entries + assert len(result) > MIN_MODEL_ENTRIES + + @patch("httpx.get") + def test_scenario_corrupted_entry(self, mock_get): + """Simulate one corrupted entry in otherwise valid response.""" + with patch.dict(os.environ, {"LITELLM_LOCAL_MODEL_COST_MAP": ""}): + # Most entries valid, but some corrupted + corrupted_map = {} + for i in range(MIN_MODEL_ENTRIES + 10): + if i == 5: + # Corrupted entry - not a dict + corrupted_map[f"model-{i}"] = "corrupted" + else: + corrupted_map[f"model-{i}"] = {"litellm_provider": "openai"} + + mock_response = MagicMock() + mock_response.json.return_value = corrupted_map + mock_response.raise_for_status.return_value = None + mock_get.return_value = mock_response + + result = get_model_cost_map("https://fake-url.com/model.json") + + # Should fall back due to corrupted entry + assert len(result) > MIN_MODEL_ENTRIES + # All entries in result should be dicts + for key, value in result.items(): + assert isinstance(value, dict), f"Entry '{key}' is not a dict"