mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
1533b7b813
commit
0220609ab6
4 changed files with 825 additions and 23 deletions
47
.github/workflows/test-model-map.yaml
vendored
47
.github/workflows/test-model-map.yaml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
317
scripts/validate_model_cost_map.py
Normal file
317
scripts/validate_model_cost_map.py
Normal 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()
|
||||
295
tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py
Normal file
295
tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py
Normal 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"
|
||||
Loading…
Add table
Reference in a new issue