Add validation and fallback protection for model cost map

This prevents upstream GitHub changes from breaking LLM calls by:

1. Runtime validation in get_model_cost_map.py:
   - Validates JSON structure before accepting remote data
   - Checks for required fields (litellm_provider)
   - Validates mode values against allowed list
   - Validates numeric cost fields
   - Falls back to local backup if validation fails
   - Logs warnings when falling back

2. Enhanced CI validation (.github/workflows/test-model-map.yaml):
   - JSON syntax validation (fast fail)
   - Comprehensive schema validation via Python script
   - Validates both main file and backup file
   - Checks backup freshness

3. New validation script (scripts/validate_model_cost_map.py):
   - Validates all entries have required fields
   - Checks mode values against allowed modes
   - Validates numeric fields are actually numeric
   - Sanity checks for suspiciously high costs
   - Detailed error reporting

4. Unit tests for validation logic

The key improvement: if GitHub returns valid JSON but semantically
broken data (missing fields, wrong types), we now detect this and
fall back to the local backup instead of using the bad data.

Customers can also set LITELLM_LOCAL_MODEL_COST_MAP=True to completely
disable remote fetching for maximum stability.

Co-authored-by: ishaan <ishaan@berri.ai>
This commit is contained in:
Cursor Agent 2026-01-30 21:26:58 +00:00
parent 1533b7b813
commit 0220609ab6
4 changed files with 825 additions and 23 deletions

View file

@ -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

View file

@ -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()

View file

@ -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()

View file

@ -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"