style: run black formatter on files from staging merge

This commit is contained in:
Ishaan Jaffer 2026-04-18 10:38:29 -07:00
parent 94fbcac6f4
commit af24a3f720
No known key found for this signature in database
43 changed files with 1920 additions and 1216 deletions

View file

@ -2,6 +2,7 @@ import asyncio
import aiohttp
import json
# Asynchronously fetch data from a given URL
async def fetch_data(url):
try:
@ -15,22 +16,24 @@ async def fetch_data(url):
resp_json = await resp.json()
print("Fetch the data from URL.")
# Return the 'data' field from the JSON response
return resp_json['data']
return resp_json["data"]
except Exception as e:
# Print an error message if fetching data fails
print("Error fetching data from URL:", e)
return None
# Synchronize local data with remote data
def sync_local_data_with_remote(local_data, remote_data):
# Update existing keys in local_data with values from remote_data
for key in (set(local_data) & set(remote_data)):
for key in set(local_data) & set(remote_data):
local_data[key].update(remote_data[key])
# Add new keys from remote_data to local_data
for key in (set(remote_data) - set(local_data)):
for key in set(remote_data) - set(local_data):
local_data[key] = remote_data[key]
# Write data to the json file
def write_to_file(file_path, data):
try:
@ -43,6 +46,7 @@ def write_to_file(file_path, data):
# Print an error message if writing to file fails
print("Error updating JSON file:", e)
# Update the existing models and add the missing models for OpenRouter
def transform_openrouter_data(data):
transformed = {}
@ -54,33 +58,41 @@ def transform_openrouter_data(data):
}
# Add 'max_output_tokens' as a field if it is not None
if "top_provider" in row and "max_completion_tokens" in row["top_provider"] and row["top_provider"]["max_completion_tokens"] is not None:
obj['max_output_tokens'] = int(row["top_provider"]["max_completion_tokens"])
if (
"top_provider" in row
and "max_completion_tokens" in row["top_provider"]
and row["top_provider"]["max_completion_tokens"] is not None
):
obj["max_output_tokens"] = int(row["top_provider"]["max_completion_tokens"])
# Add the field 'output_cost_per_token'
obj.update({
"output_cost_per_token": float(row["pricing"]["completion"]),
})
obj.update(
{
"output_cost_per_token": float(row["pricing"]["completion"]),
}
)
# Add field 'input_cost_per_image' if it exists and is non-zero
if "pricing" in row and "image" in row["pricing"] and float(row["pricing"]["image"]) != 0.0:
obj['input_cost_per_image'] = float(row["pricing"]["image"])
if (
"pricing" in row
and "image" in row["pricing"]
and float(row["pricing"]["image"]) != 0.0
):
obj["input_cost_per_image"] = float(row["pricing"]["image"])
# Add the fields 'litellm_provider' and 'mode'
obj.update({
"litellm_provider": "openrouter",
"mode": "chat"
})
obj.update({"litellm_provider": "openrouter", "mode": "chat"})
# Add the 'supports_vision' field if the modality is 'multimodal'
if row.get('architecture', {}).get('modality') == 'multimodal':
obj['supports_vision'] = True
if row.get("architecture", {}).get("modality") == "multimodal":
obj["supports_vision"] = True
# Use a composite key to store the transformed object
transformed[f'openrouter/{row["id"]}'] = obj
return transformed
# Update the existing models and add the missing models for Vercel AI Gateway
def transform_vercel_ai_gateway_data(data):
transformed = {}
@ -89,20 +101,30 @@ def transform_vercel_ai_gateway_data(data):
"max_tokens": row["context_window"],
"input_cost_per_token": float(row["pricing"]["input"]),
"output_cost_per_token": float(row["pricing"]["output"]),
'max_output_tokens': row['max_tokens'],
'max_input_tokens': row["context_window"],
"max_output_tokens": row["max_tokens"],
"max_input_tokens": row["context_window"],
}
# Handle cache pricing if available
if "pricing" in row:
if "input_cache_read" in row["pricing"] and row["pricing"]["input_cache_read"] is not None:
obj['cache_read_input_token_cost'] = float(f"{float(row['pricing']['input_cache_read']):e}")
if "input_cache_write" in row["pricing"] and row["pricing"]["input_cache_write"] is not None:
obj['cache_creation_input_token_cost'] = float(f"{float(row['pricing']['input_cache_write']):e}")
if (
"input_cache_read" in row["pricing"]
and row["pricing"]["input_cache_read"] is not None
):
obj["cache_read_input_token_cost"] = float(
f"{float(row['pricing']['input_cache_read']):e}"
)
if (
"input_cache_write" in row["pricing"]
and row["pricing"]["input_cache_write"] is not None
):
obj["cache_creation_input_token_cost"] = float(
f"{float(row['pricing']['input_cache_write']):e}"
)
mode = "embedding" if "embedding" in row["id"].lower() else "chat"
obj.update({"litellm_provider": "vercel_ai_gateway", "mode": mode})
transformed[f'vercel_ai_gateway/{row["id"]}'] = obj
@ -126,24 +148,31 @@ def load_local_data(file_path):
print("Error decoding JSON:", e)
return None
def main():
local_file_path = "model_prices_and_context_window.json" # Path to the local data file
openrouter_url = "https://openrouter.ai/api/v1/models" # URL to fetch OpenRouter data
vercel_ai_gateway_url = "https://ai-gateway.vercel.sh/v1/models" # URL to fetch Vercel AI Gateway data
local_file_path = (
"model_prices_and_context_window.json" # Path to the local data file
)
openrouter_url = (
"https://openrouter.ai/api/v1/models" # URL to fetch OpenRouter data
)
vercel_ai_gateway_url = (
"https://ai-gateway.vercel.sh/v1/models" # URL to fetch Vercel AI Gateway data
)
# Load local data from file
local_data = load_local_data(local_file_path)
# Fetch OpenRouter data
openrouter_data = asyncio.run(fetch_data(openrouter_url))
# Transform the fetched OpenRouter data
openrouter_data = transform_openrouter_data(openrouter_data)
# Fetch Vercel AI Gateway data
vercel_data = asyncio.run(fetch_data(vercel_ai_gateway_url))
# Transform the fetched Vercel AI Gateway data
vercel_data = transform_vercel_ai_gateway_data(vercel_data)
# Combine both datasets
all_remote_data = {**openrouter_data, **vercel_data}
@ -154,6 +183,7 @@ def main():
else:
print("Failed to fetch model data from either local file or URL.")
# Entry point of the script
if __name__ == "__main__":
main()

View file

@ -16,64 +16,75 @@ from pathlib import Path
import json
from typing import Dict, List, Tuple, Optional
# ANSI color codes for terminal output
class Colors:
GREEN = '\033[92m'
RED = '\033[91m'
YELLOW = '\033[93m'
BLUE = '\033[94m'
PURPLE = '\033[95m'
CYAN = '\033[96m'
RESET = '\033[0m'
BOLD = '\033[1m'
GREEN = "\033[92m"
RED = "\033[91m"
YELLOW = "\033[93m"
BLUE = "\033[94m"
PURPLE = "\033[95m"
CYAN = "\033[96m"
RESET = "\033[0m"
BOLD = "\033[1m"
def print_colored(message: str, color: str = Colors.RESET):
"""Print colored message to terminal"""
print(f"{color}{message}{Colors.RESET}")
def get_provider_from_test_file(test_file: str) -> str:
"""Map test file names to provider names"""
provider_mapping = {
'test_anthropic': 'Anthropic',
'test_azure': 'Azure',
'test_bedrock': 'AWS Bedrock',
'test_openai': 'OpenAI',
'test_vertex': 'Google Vertex AI',
'test_gemini': 'Google Vertex AI',
'test_cohere': 'Cohere',
'test_databricks': 'Databricks',
'test_groq': 'Groq',
'test_together': 'Together AI',
'test_mistral': 'Mistral',
'test_deepseek': 'DeepSeek',
'test_replicate': 'Replicate',
'test_huggingface': 'HuggingFace',
'test_fireworks': 'Fireworks AI',
'test_perplexity': 'Perplexity',
'test_cloudflare': 'Cloudflare',
'test_voyage': 'Voyage AI',
'test_xai': 'xAI',
'test_nvidia': 'NVIDIA',
'test_watsonx': 'IBM watsonx',
'test_azure_ai': 'Azure AI',
'test_snowflake': 'Snowflake',
'test_infinity': 'Infinity',
'test_jina': 'Jina AI',
'test_deepgram': 'Deepgram',
'test_clarifai': 'Clarifai',
'test_triton': 'Triton',
"test_anthropic": "Anthropic",
"test_azure": "Azure",
"test_bedrock": "AWS Bedrock",
"test_openai": "OpenAI",
"test_vertex": "Google Vertex AI",
"test_gemini": "Google Vertex AI",
"test_cohere": "Cohere",
"test_databricks": "Databricks",
"test_groq": "Groq",
"test_together": "Together AI",
"test_mistral": "Mistral",
"test_deepseek": "DeepSeek",
"test_replicate": "Replicate",
"test_huggingface": "HuggingFace",
"test_fireworks": "Fireworks AI",
"test_perplexity": "Perplexity",
"test_cloudflare": "Cloudflare",
"test_voyage": "Voyage AI",
"test_xai": "xAI",
"test_nvidia": "NVIDIA",
"test_watsonx": "IBM watsonx",
"test_azure_ai": "Azure AI",
"test_snowflake": "Snowflake",
"test_infinity": "Infinity",
"test_jina": "Jina AI",
"test_deepgram": "Deepgram",
"test_clarifai": "Clarifai",
"test_triton": "Triton",
}
for key, provider in provider_mapping.items():
if key in test_file:
return provider
# For cross-provider test files
if any(name in test_file for name in ['test_optional_params', 'test_prompt_factory',
'test_router', 'test_text_completion']):
return f'Cross-Provider Tests ({test_file})'
return 'Other Tests'
if any(
name in test_file
for name in [
"test_optional_params",
"test_prompt_factory",
"test_router",
"test_text_completion",
]
):
return f"Cross-Provider Tests ({test_file})"
return "Other Tests"
def format_duration(seconds: float) -> str:
"""Format duration in human-readable format"""
@ -89,290 +100,355 @@ def format_duration(seconds: float) -> str:
return f"{hours}h {minutes}m"
def generate_markdown_report(junit_xml_path: str, output_path: str, tag: str = None, commit: str = None):
def generate_markdown_report(
junit_xml_path: str, output_path: str, tag: str = None, commit: str = None
):
"""Generate a beautiful markdown report from JUnit XML"""
try:
tree = ET.parse(junit_xml_path)
root = tree.getroot()
# Handle both testsuite and testsuites root
if root.tag == 'testsuites':
suites = root.findall('testsuite')
if root.tag == "testsuites":
suites = root.findall("testsuite")
else:
suites = [root]
# Overall statistics
total_tests = 0
total_failures = 0
total_errors = 0
total_skipped = 0
total_time = 0.0
# Provider breakdown
provider_stats = defaultdict(lambda: {'passed': 0, 'failed': 0, 'skipped': 0, 'errors': 0, 'time': 0.0})
provider_stats = defaultdict(
lambda: {"passed": 0, "failed": 0, "skipped": 0, "errors": 0, "time": 0.0}
)
provider_tests = defaultdict(list)
for suite in suites:
total_tests += int(suite.get('tests', 0))
total_failures += int(suite.get('failures', 0))
total_errors += int(suite.get('errors', 0))
total_skipped += int(suite.get('skipped', 0))
total_time += float(suite.get('time', 0))
for testcase in suite.findall('testcase'):
classname = testcase.get('classname', '')
test_name = testcase.get('name', '')
test_time = float(testcase.get('time', 0))
total_tests += int(suite.get("tests", 0))
total_failures += int(suite.get("failures", 0))
total_errors += int(suite.get("errors", 0))
total_skipped += int(suite.get("skipped", 0))
total_time += float(suite.get("time", 0))
for testcase in suite.findall("testcase"):
classname = testcase.get("classname", "")
test_name = testcase.get("name", "")
test_time = float(testcase.get("time", 0))
# Extract test file name from classname
if '.' in classname:
parts = classname.split('.')
test_file = parts[-2] if len(parts) > 1 else 'unknown'
if "." in classname:
parts = classname.split(".")
test_file = parts[-2] if len(parts) > 1 else "unknown"
else:
test_file = 'unknown'
test_file = "unknown"
provider = get_provider_from_test_file(test_file)
provider_stats[provider]['time'] += test_time
provider_stats[provider]["time"] += test_time
# Check test status
if testcase.find('failure') is not None:
provider_stats[provider]['failed'] += 1
failure = testcase.find('failure')
failure_msg = failure.get('message', '') if failure is not None else ''
provider_tests[provider].append({
'name': test_name,
'status': 'FAILED',
'time': test_time,
'message': failure_msg
})
elif testcase.find('error') is not None:
provider_stats[provider]['errors'] += 1
error = testcase.find('error')
error_msg = error.get('message', '') if error is not None else ''
provider_tests[provider].append({
'name': test_name,
'status': 'ERROR',
'time': test_time,
'message': error_msg
})
elif testcase.find('skipped') is not None:
provider_stats[provider]['skipped'] += 1
skip = testcase.find('skipped')
skip_msg = skip.get('message', '') if skip is not None else ''
provider_tests[provider].append({
'name': test_name,
'status': 'SKIPPED',
'time': test_time,
'message': skip_msg
})
if testcase.find("failure") is not None:
provider_stats[provider]["failed"] += 1
failure = testcase.find("failure")
failure_msg = (
failure.get("message", "") if failure is not None else ""
)
provider_tests[provider].append(
{
"name": test_name,
"status": "FAILED",
"time": test_time,
"message": failure_msg,
}
)
elif testcase.find("error") is not None:
provider_stats[provider]["errors"] += 1
error = testcase.find("error")
error_msg = error.get("message", "") if error is not None else ""
provider_tests[provider].append(
{
"name": test_name,
"status": "ERROR",
"time": test_time,
"message": error_msg,
}
)
elif testcase.find("skipped") is not None:
provider_stats[provider]["skipped"] += 1
skip = testcase.find("skipped")
skip_msg = skip.get("message", "") if skip is not None else ""
provider_tests[provider].append(
{
"name": test_name,
"status": "SKIPPED",
"time": test_time,
"message": skip_msg,
}
)
else:
provider_stats[provider]['passed'] += 1
provider_tests[provider].append({
'name': test_name,
'status': 'PASSED',
'time': test_time,
'message': ''
})
provider_stats[provider]["passed"] += 1
provider_tests[provider].append(
{
"name": test_name,
"status": "PASSED",
"time": test_time,
"message": "",
}
)
passed = total_tests - total_failures - total_errors - total_skipped
# Generate the markdown report
with open(output_path, 'w') as f:
with open(output_path, "w") as f:
# Header
f.write("# LLM Translation Test Results\n\n")
# Metadata table
f.write("## Test Run Information\n\n")
f.write("| Field | Value |\n")
f.write("|-------|-------|\n")
f.write(f"| **Tag** | `{tag or 'N/A'}` |\n")
f.write(f"| **Date** | {datetime.utcnow().strftime('%Y-%m-%d %H:%M:%S UTC')} |\n")
f.write(
f"| **Date** | {datetime.utcnow().strftime('%Y-%m-%d %H:%M:%S UTC')} |\n"
)
f.write(f"| **Commit** | `{commit or 'N/A'}` |\n")
f.write(f"| **Duration** | {format_duration(total_time)} |\n")
f.write("\n")
# Overall statistics with visual elements
f.write("## Overall Statistics\n\n")
# Summary box
f.write("```\n")
f.write(f"Total Tests: {total_tests}\n")
f.write(f"├── Passed: {passed:>4} ({(passed/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
f.write(f"├── Failed: {total_failures:>4} ({(total_failures/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
f.write(f"├── Errors: {total_errors:>4} ({(total_errors/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
f.write(f"└── Skipped: {total_skipped:>4} ({(total_skipped/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
f.write(
f"├── Passed: {passed:>4} ({(passed/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
)
f.write(
f"├── Failed: {total_failures:>4} ({(total_failures/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
)
f.write(
f"├── Errors: {total_errors:>4} ({(total_errors/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
)
f.write(
f"└── Skipped: {total_skipped:>4} ({(total_skipped/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
)
f.write("```\n\n")
# Provider summary table
f.write("## Results by Provider\n\n")
f.write("| Provider | Total | Pass | Fail | Error | Skip | Pass Rate | Duration |\n")
f.write("|----------|-------|------|------|-------|------|-----------|----------|")
f.write(
"| Provider | Total | Pass | Fail | Error | Skip | Pass Rate | Duration |\n"
)
f.write(
"|----------|-------|------|------|-------|------|-----------|----------|"
)
# Sort providers: specific providers first, then cross-provider tests
sorted_providers = []
cross_provider = []
for p in sorted(provider_stats.keys()):
if 'Cross-Provider' in p or p == 'Other Tests':
if "Cross-Provider" in p or p == "Other Tests":
cross_provider.append(p)
else:
sorted_providers.append(p)
all_providers = sorted_providers + cross_provider
for provider in all_providers:
stats = provider_stats[provider]
total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped']
pass_rate = (stats['passed'] / total * 100) if total > 0 else 0
f.write(f"\n| {provider} | {total} | {stats['passed']} | {stats['failed']} | ")
total = (
stats["passed"]
+ stats["failed"]
+ stats["errors"]
+ stats["skipped"]
)
pass_rate = (stats["passed"] / total * 100) if total > 0 else 0
f.write(
f"\n| {provider} | {total} | {stats['passed']} | {stats['failed']} | "
)
f.write(f"{stats['errors']} | {stats['skipped']} | {pass_rate:.1f}% | ")
f.write(f"{format_duration(stats['time'])} |")
# Detailed test results by provider
f.write("\n\n## Detailed Test Results\n\n")
for provider in sorted_providers:
if provider_tests[provider]:
stats = provider_stats[provider]
total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped']
total = (
stats["passed"]
+ stats["failed"]
+ stats["errors"]
+ stats["skipped"]
)
f.write(f"### {provider}\n\n")
f.write(f"**Summary:** {stats['passed']}/{total} passed ")
f.write(f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%) ")
f.write(
f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%) "
)
f.write(f"in {format_duration(stats['time'])}\n\n")
# Group tests by status
tests_by_status = defaultdict(list)
for test in provider_tests[provider]:
tests_by_status[test['status']].append(test)
tests_by_status[test["status"]].append(test)
# Show failed tests first (if any)
if tests_by_status['FAILED']:
if tests_by_status["FAILED"]:
f.write("<details>\n<summary>Failed Tests</summary>\n\n")
for test in tests_by_status['FAILED']:
for test in tests_by_status["FAILED"]:
f.write(f"- `{test['name']}` ({test['time']:.2f}s)\n")
if test['message']:
if test["message"]:
# Truncate long error messages
msg = test['message'][:200] + '...' if len(test['message']) > 200 else test['message']
msg = (
test["message"][:200] + "..."
if len(test["message"]) > 200
else test["message"]
)
f.write(f" > {msg}\n")
f.write("\n</details>\n\n")
# Show errors (if any)
if tests_by_status['ERROR']:
if tests_by_status["ERROR"]:
f.write("<details>\n<summary>Error Tests</summary>\n\n")
for test in tests_by_status['ERROR']:
for test in tests_by_status["ERROR"]:
f.write(f"- `{test['name']}` ({test['time']:.2f}s)\n")
f.write("\n</details>\n\n")
# Show passed tests in collapsible section
if tests_by_status['PASSED']:
if tests_by_status["PASSED"]:
f.write("<details>\n<summary>Passed Tests</summary>\n\n")
for test in tests_by_status['PASSED']:
for test in tests_by_status["PASSED"]:
f.write(f"- `{test['name']}` ({test['time']:.2f}s)\n")
f.write("\n</details>\n\n")
# Show skipped tests (if any)
if tests_by_status['SKIPPED']:
if tests_by_status["SKIPPED"]:
f.write("<details>\n<summary>Skipped Tests</summary>\n\n")
for test in tests_by_status['SKIPPED']:
for test in tests_by_status["SKIPPED"]:
f.write(f"- `{test['name']}`\n")
f.write("\n</details>\n\n")
# Cross-provider tests in a separate section
if cross_provider:
f.write("### Cross-Provider Tests\n\n")
for provider in cross_provider:
if provider_tests[provider]:
stats = provider_stats[provider]
total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped']
total = (
stats["passed"]
+ stats["failed"]
+ stats["errors"]
+ stats["skipped"]
)
f.write(f"#### {provider}\n\n")
f.write(f"**Summary:** {stats['passed']}/{total} passed ")
f.write(f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%)\n\n")
f.write(
f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%)\n\n"
)
# For cross-provider tests, just show counts
f.write(f"- Passed: {stats['passed']}\n")
if stats['failed'] > 0:
if stats["failed"] > 0:
f.write(f"- Failed: {stats['failed']}\n")
if stats['errors'] > 0:
if stats["errors"] > 0:
f.write(f"- Errors: {stats['errors']}\n")
if stats['skipped'] > 0:
if stats["skipped"] > 0:
f.write(f"- Skipped: {stats['skipped']}\n")
f.write("\n")
print_colored(f"Report generated: {output_path}", Colors.GREEN)
except Exception as e:
print_colored(f"Error generating report: {e}", Colors.RED)
raise
def run_tests(test_path: str = "tests/llm_translation/",
junit_xml: str = "test-results/junit.xml",
report_path: str = "test-results/llm_translation_report.md",
tag: str = None,
commit: str = None) -> int:
def run_tests(
test_path: str = "tests/llm_translation/",
junit_xml: str = "test-results/junit.xml",
report_path: str = "test-results/llm_translation_report.md",
tag: str = None,
commit: str = None,
) -> int:
"""Run the LLM translation tests and generate report"""
# Create test results directory
os.makedirs(os.path.dirname(junit_xml), exist_ok=True)
print_colored("Starting LLM Translation Tests", Colors.BOLD + Colors.BLUE)
print_colored(f"Test directory: {test_path}", Colors.CYAN)
print_colored(f"Output: {junit_xml}", Colors.CYAN)
print()
# Run pytest
cmd = [
"uv", "run", "--no-sync", "pytest", test_path,
"uv",
"run",
"--no-sync",
"pytest",
test_path,
f"--junitxml={junit_xml}",
"-v",
"--tb=short",
"--maxfail=500",
"-n", "auto"
"-n",
"auto",
]
# Add timeout if pytest-timeout is installed
try:
subprocess.run(["uv", "run", "--no-sync", "python", "-c", "import pytest_timeout"],
capture_output=True, check=True)
subprocess.run(
["uv", "run", "--no-sync", "python", "-c", "import pytest_timeout"],
capture_output=True,
check=True,
)
cmd.extend(["--timeout=300"])
except:
print_colored("Warning: pytest-timeout not installed, skipping timeout option", Colors.YELLOW)
print_colored(
"Warning: pytest-timeout not installed, skipping timeout option",
Colors.YELLOW,
)
print_colored("Running pytest with command:", Colors.YELLOW)
print(f" {' '.join(cmd)}")
print()
# Run the tests
result = subprocess.run(cmd, capture_output=False)
# Generate the report regardless of test outcome
if os.path.exists(junit_xml):
print()
print_colored("Generating test report...", Colors.BLUE)
generate_markdown_report(junit_xml, report_path, tag, commit)
# Print summary to console
print()
print_colored("Test Summary:", Colors.BOLD + Colors.PURPLE)
# Parse XML for quick summary
tree = ET.parse(junit_xml)
root = tree.getroot()
if root.tag == 'testsuites':
suites = root.findall('testsuite')
if root.tag == "testsuites":
suites = root.findall("testsuite")
else:
suites = [root]
total = sum(int(s.get('tests', 0)) for s in suites)
failures = sum(int(s.get('failures', 0)) for s in suites)
errors = sum(int(s.get('errors', 0)) for s in suites)
skipped = sum(int(s.get('skipped', 0)) for s in suites)
total = sum(int(s.get("tests", 0)) for s in suites)
failures = sum(int(s.get("failures", 0)) for s in suites)
errors = sum(int(s.get("errors", 0)) for s in suites)
skipped = sum(int(s.get("skipped", 0)) for s in suites)
passed = total - failures - errors - skipped
print(f" Total: {total}")
print_colored(f" Passed: {passed}", Colors.GREEN)
if failures > 0:
@ -381,59 +457,75 @@ def run_tests(test_path: str = "tests/llm_translation/",
print_colored(f" Errors: {errors}", Colors.RED)
if skipped > 0:
print_colored(f" Skipped: {skipped}", Colors.YELLOW)
if total > 0:
pass_rate = (passed / total) * 100
color = Colors.GREEN if pass_rate >= 80 else Colors.YELLOW if pass_rate >= 60 else Colors.RED
color = (
Colors.GREEN
if pass_rate >= 80
else Colors.YELLOW if pass_rate >= 60 else Colors.RED
)
print_colored(f" Pass Rate: {pass_rate:.1f}%", color)
else:
print_colored("No test results found!", Colors.RED)
print()
print_colored("Test run complete!", Colors.BOLD + Colors.GREEN)
return result.returncode
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="Run LLM Translation Tests")
parser.add_argument("--test-path", default="tests/llm_translation/",
help="Path to test directory")
parser.add_argument("--junit-xml", default="test-results/junit.xml",
help="Path for JUnit XML output")
parser.add_argument("--report", default="test-results/llm_translation_report.md",
help="Path for markdown report")
parser.add_argument(
"--test-path", default="tests/llm_translation/", help="Path to test directory"
)
parser.add_argument(
"--junit-xml",
default="test-results/junit.xml",
help="Path for JUnit XML output",
)
parser.add_argument(
"--report",
default="test-results/llm_translation_report.md",
help="Path for markdown report",
)
parser.add_argument("--tag", help="Git tag or version")
parser.add_argument("--commit", help="Git commit SHA")
args = parser.parse_args()
# Get git info if not provided
if not args.commit:
try:
result = subprocess.run(["git", "rev-parse", "HEAD"],
capture_output=True, text=True)
result = subprocess.run(
["git", "rev-parse", "HEAD"], capture_output=True, text=True
)
if result.returncode == 0:
args.commit = result.stdout.strip()
except:
pass
if not args.tag:
try:
result = subprocess.run(["git", "describe", "--tags", "--abbrev=0"],
capture_output=True, text=True)
result = subprocess.run(
["git", "describe", "--tags", "--abbrev=0"],
capture_output=True,
text=True,
)
if result.returncode == 0:
args.tag = result.stdout.strip()
except:
pass
exit_code = run_tests(
test_path=args.test_path,
junit_xml=args.junit_xml,
report_path=args.report,
tag=args.tag,
commit=args.commit
commit=args.commit,
)
sys.exit(exit_code)

12
cache_demo_config.yaml Normal file
View file

@ -0,0 +1,12 @@
model_list:
- model_name: bedrock-claude-haiku
litellm_params:
model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0
aws_region_name: us-east-1
litellm_settings:
success_callback: []
failure_callback: []
general_settings:
store_model_in_db: false

75
cache_demo_request.py Normal file
View file

@ -0,0 +1,75 @@
"""
Demo script: back-to-back Bedrock streaming requests with prompt caching.
Request 1: populates the cache (cache_creation_input_tokens)
Request 2: reads from cache (cache_read_input_tokens)
"""
import json
import time
import httpx
PROXY_URL = "http://localhost:4001"
API_KEY = "sk-1234"
# ~5000 token system prompt (above claude-haiku-4-5's 2048-token min for caching on Bedrock)
LARGE_SYSTEM_PROMPT = (
"AWS Bedrock provides managed ML infrastructure for enterprise workloads. "
"Anthropic Claude models support prompt caching for cost optimization. "
) * 200
def make_streaming_request(req_num: int, label: str) -> None:
print(f"\n{'='*60}")
print(f"Request {req_num}: {label}")
print(f"{'='*60}")
payload = {
"model": "bedrock-claude-haiku",
"messages": [
{
"role": "system",
"content": [
{
"type": "text",
"text": LARGE_SYSTEM_PROMPT,
"cache_control": {"type": "ephemeral"},
}
],
},
{"role": "user", "content": f"Say only: 'Request {req_num} done'"},
],
"stream": True,
"max_tokens": 20,
}
full_response = ""
with httpx.Client(timeout=60) as client:
with client.stream(
"POST",
f"{PROXY_URL}/v1/chat/completions",
json=payload,
headers={"Authorization": f"Bearer {API_KEY}"},
) as r:
r.raise_for_status()
for line in r.iter_lines():
if line.startswith("data: ") and line != "data: [DONE]":
chunk = json.loads(line[6:])
delta = chunk.get("choices", [{}])[0].get("delta", {})
if content := delta.get("content"):
full_response += content
print(f"Response: {full_response!r}")
if __name__ == "__main__":
print("Sending request 1 (cache write)...")
make_streaming_request(1, "cache WRITE (populates cache)")
print("\nWaiting 2s between requests...")
time.sleep(2)
print("Sending request 2 (cache read)...")
make_streaming_request(2, "cache READ (hits cache)")
print("\n\nDone. Check SpendLogs in the DB or the UI at http://localhost:4001")

View file

@ -14,53 +14,74 @@ from litellm.types.utils import StandardCallbackDynamicParams
class EnterpriseCallbackControls:
@staticmethod
def is_callback_disabled_dynamically(
callback: litellm.CALLBACK_TYPES,
litellm_params: dict,
standard_callback_dynamic_params: StandardCallbackDynamicParams
) -> bool:
"""
Check if a callback is disabled via the x-litellm-disable-callbacks header or via `litellm_disabled_callbacks` in standard_callback_dynamic_params.
Args:
callback: The callback to check (can be string, CustomLogger instance, or callable)
litellm_params: Parameters containing proxy server request info
Returns:
bool: True if the callback should be disabled, False otherwise
"""
from litellm.litellm_core_utils.custom_logger_registry import (
CustomLoggerRegistry,
)
callback: litellm.CALLBACK_TYPES,
litellm_params: dict,
standard_callback_dynamic_params: StandardCallbackDynamicParams,
) -> bool:
"""
Check if a callback is disabled via the x-litellm-disable-callbacks header or via `litellm_disabled_callbacks` in standard_callback_dynamic_params.
Args:
callback: The callback to check (can be string, CustomLogger instance, or callable)
litellm_params: Parameters containing proxy server request info
Returns:
bool: True if the callback should be disabled, False otherwise
"""
from litellm.litellm_core_utils.custom_logger_registry import (
CustomLoggerRegistry,
)
try:
disabled_callbacks = EnterpriseCallbackControls.get_disabled_callbacks(
litellm_params, standard_callback_dynamic_params
)
verbose_logger.debug(
f"Dynamically disabled callbacks from {X_LITELLM_DISABLE_CALLBACKS}: {disabled_callbacks}"
)
verbose_logger.debug(
f"Checking if {callback} is disabled via headers. Disable callbacks from headers: {disabled_callbacks}"
)
if disabled_callbacks is not None:
#########################################################
# premium user check
#########################################################
if (
not EnterpriseCallbackControls._should_allow_dynamic_callback_disabling()
):
return False
#########################################################
if isinstance(callback, str):
if callback.lower() in disabled_callbacks:
verbose_logger.debug(
f"Not logging to {callback} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}"
)
return True
elif isinstance(callback, CustomLogger):
# get the string name of the callback
callback_str = (
CustomLoggerRegistry.get_callback_str_from_class_type(
callback.__class__
)
)
if (
callback_str is not None
and callback_str.lower() in disabled_callbacks
):
verbose_logger.debug(
f"Not logging to {callback_str} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}"
)
return True
return False
except Exception as e:
verbose_logger.debug(f"Error checking disabled callbacks header: {str(e)}")
return False
try:
disabled_callbacks = EnterpriseCallbackControls.get_disabled_callbacks(litellm_params, standard_callback_dynamic_params)
verbose_logger.debug(f"Dynamically disabled callbacks from {X_LITELLM_DISABLE_CALLBACKS}: {disabled_callbacks}")
verbose_logger.debug(f"Checking if {callback} is disabled via headers. Disable callbacks from headers: {disabled_callbacks}")
if disabled_callbacks is not None:
#########################################################
# premium user check
#########################################################
if not EnterpriseCallbackControls._should_allow_dynamic_callback_disabling():
return False
#########################################################
if isinstance(callback, str):
if callback.lower() in disabled_callbacks:
verbose_logger.debug(f"Not logging to {callback} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}")
return True
elif isinstance(callback, CustomLogger):
# get the string name of the callback
callback_str = CustomLoggerRegistry.get_callback_str_from_class_type(callback.__class__)
if callback_str is not None and callback_str.lower() in disabled_callbacks:
verbose_logger.debug(f"Not logging to {callback_str} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}")
return True
return False
except Exception as e:
verbose_logger.debug(
f"Error checking disabled callbacks header: {str(e)}"
)
return False
@staticmethod
def get_disabled_callbacks(litellm_params: dict, standard_callback_dynamic_params: StandardCallbackDynamicParams) -> Optional[List[str]]:
def get_disabled_callbacks(
litellm_params: dict,
standard_callback_dynamic_params: StandardCallbackDynamicParams,
) -> Optional[List[str]]:
"""
Get the disabled callbacks from the standard callback dynamic params.
"""
@ -71,18 +92,24 @@ class EnterpriseCallbackControls:
request_headers = get_proxy_server_request_headers(litellm_params)
disabled_callbacks = request_headers.get(X_LITELLM_DISABLE_CALLBACKS, None)
if disabled_callbacks is not None:
disabled_callbacks = set([cb.strip().lower() for cb in disabled_callbacks.split(",")])
disabled_callbacks = set(
[cb.strip().lower() for cb in disabled_callbacks.split(",")]
)
return list(disabled_callbacks)
#########################################################
# check if disabled via request body
#########################################################
if standard_callback_dynamic_params.get("litellm_disabled_callbacks", None) is not None:
return standard_callback_dynamic_params.get("litellm_disabled_callbacks", None)
if (
standard_callback_dynamic_params.get("litellm_disabled_callbacks", None)
is not None
):
return standard_callback_dynamic_params.get(
"litellm_disabled_callbacks", None
)
return None
@staticmethod
def _should_allow_dynamic_callback_disabling():
import litellm
@ -90,10 +117,14 @@ class EnterpriseCallbackControls:
# Check if admin has disabled this feature
if litellm.allow_dynamic_callback_disabling is not True:
verbose_logger.debug("Dynamic callback disabling is disabled by admin via litellm.allow_dynamic_callback_disabling")
verbose_logger.debug(
"Dynamic callback disabling is disabled by admin via litellm.allow_dynamic_callback_disabling"
)
return False
if premium_user:
return True
verbose_logger.warning(f"Disabling callbacks using request headers is an enterprise feature. {CommonProxyErrors.not_premium_user.value}")
return False
verbose_logger.warning(
f"Disabling callbacks using request headers is an enterprise feature. {CommonProxyErrors.not_premium_user.value}"
)
return False

View file

@ -79,4 +79,4 @@ class SendGridEmailLogger(BaseEmailLogger):
verbose_logger.debug(
f"SendGrid response status={response.status_code}, body={response.text}"
)
return
return

View file

@ -1,6 +1,7 @@
"""
This is the litellm SMTP email integration
"""
import asyncio
from typing import List

View file

@ -1,6 +1,7 @@
"""
Enterprise specific logging utils
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingMetadata

View file

@ -153,11 +153,11 @@ async def get_audit_logs(
# Return paginated response
return PaginatedAuditLogResponse(
audit_logs=[
AuditLogResponse(**audit_log.model_dump()) for audit_log in audit_logs
]
if audit_logs
else [],
audit_logs=(
[AuditLogResponse(**audit_log.model_dump()) for audit_log in audit_logs]
if audit_logs
else []
),
total=total_count,
page=page,
page_size=page_size,

View file

@ -7,4 +7,4 @@ including custom SSO handlers and advanced authentication features.
from .custom_sso_handler import EnterpriseCustomSSOHandler
__all__ = ["EnterpriseCustomSSOHandler"]
__all__ = ["EnterpriseCustomSSOHandler"]

View file

@ -26,12 +26,12 @@ from litellm.proxy.management_endpoints.types import CustomOpenID
class EnterpriseCustomSSOHandler:
"""
Enterprise Custom SSO Handler for LiteLLM Proxy
This class provides methods for handling custom SSO authentication flows
where users can implement their own authentication logic by processing
request headers and returning user information in OpenID format.
"""
@staticmethod
async def handle_custom_ui_sso_sign_in(
request: Request,
@ -40,16 +40,16 @@ class EnterpriseCustomSSOHandler:
Allow a user to execute their custom code to parse incoming request headers and return a OpenID object
Use this when you have an OAuth proxy in front of LiteLLM (where the OAuth proxy has already authenticated the user)
Args:
request: The FastAPI request object containing headers and other request data
Returns:
RedirectResponse: Redirect response that sends the user to the LiteLLM UI with authentication token
Raises:
ValueError: If custom_ui_sso_sign_in_handler is not configured
Example:
This method is typically called when a user has already been authenticated by an
external OAuth proxy and the proxy has added custom headers containing user information.
@ -63,24 +63,31 @@ class EnterpriseCustomSSOHandler:
premium_user,
user_custom_ui_sso_sign_in_handler,
)
if premium_user is not True:
raise ValueError(CommonProxyErrors.not_premium_user.value)
if user_custom_ui_sso_sign_in_handler is None:
raise ValueError("custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings.")
custom_sso_login_handler = cast(CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler)
openid_response: OpenID = await custom_sso_login_handler.handle_custom_ui_sso_sign_in(
request=request,
raise ValueError(
"custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings."
)
custom_sso_login_handler = cast(
CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler
)
openid_response: OpenID = (
await custom_sso_login_handler.handle_custom_ui_sso_sign_in(
request=request,
)
)
# Import here to avoid circular imports
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
return await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=openid_response,
request=request,
received_response=None,
generic_client_id=None,
ui_access_mode=None,
)
)

View file

@ -53,7 +53,9 @@ class CheckBatchCost:
"user_api_key_alias": getattr(user_row, "user_alias", None),
}
except Exception as e:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}")
verbose_proxy_logger.error(
f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}"
)
return {}
async def _cleanup_stale_managed_objects(self) -> None:
@ -62,11 +64,22 @@ class CheckBatchCost:
in non-terminal states as 'stale_expired'. These will never complete and
should not be polled.
"""
cutoff = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
cutoff = datetime.now(timezone.utc) - timedelta(
days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS
)
result = await self.prisma_client.db.litellm_managedobjecttable.update_many(
where={
"file_purpose": "batch",
"status": {"not_in": ["completed", "complete", "failed", "expired", "cancelled", "stale_expired"]},
"status": {
"not_in": [
"completed",
"complete",
"failed",
"expired",
"cancelled",
"stale_expired",
]
},
"created_at": {"lt": cutoff},
},
data={"status": "stale_expired"},
@ -120,9 +133,12 @@ class CheckBatchCost:
try:
from litellm.integrations.prometheus import PrometheusLogger
prom_logger = PrometheusLogger.get_instance()
except Exception as e:
verbose_proxy_logger.error(f"CheckBatchCost: could not get Prometheus logger: {e}")
verbose_proxy_logger.error(
f"CheckBatchCost: could not get Prometheus logger: {e}"
)
prom_logger = None
processed_models: List[Tuple[Optional[str], Optional[str]]] = []
@ -161,7 +177,11 @@ class CheckBatchCost:
order={"created_at": "asc"},
)
except Exception as query_err:
if "batch_processed" not in str(query_err).lower() and "unknown column" not in str(query_err).lower() and "does not exist" not in str(query_err).lower():
if (
"batch_processed" not in str(query_err).lower()
and "unknown column" not in str(query_err).lower()
and "does not exist" not in str(query_err).lower()
):
raise
# Permanent schema gap — cache the result so future cycles skip straight to fallback
self._has_batch_processed_column = False
@ -216,14 +236,13 @@ class CheckBatchCost:
f"Skipping job {unified_object_id} because of error querying model ID: {model_id} for cost and usage of batch ID: {batch_id}: {e}"
)
if prom_logger:
prom_logger.record_check_batch_cost_error("provider_retrieval_error")
prom_logger.record_check_batch_cost_error(
"provider_retrieval_error"
)
continue
## RETRIEVE THE BATCH JOB OUTPUT FILE
if (
response.status == "completed"
and response.output_file_id is not None
):
if response.status == "completed" and response.output_file_id is not None:
verbose_proxy_logger.info(
f"Batch ID: {batch_id} is complete, tracking cost and usage"
)
@ -250,20 +269,25 @@ class CheckBatchCost:
decoded = _is_base64_encoded_unified_file_id(raw_output_file_id)
if decoded:
try:
raw_output_file_id = decoded.split("llm_output_file_id,")[1].split(";")[0]
raw_output_file_id = decoded.split("llm_output_file_id,")[
1
].split(";")[0]
except (IndexError, AttributeError):
pass
credentials = self.llm_router.get_deployment_credentials_with_provider(model_id) or {}
credentials = (
self.llm_router.get_deployment_credentials_with_provider(model_id)
or {}
)
_file_content = await afile_content(
file_id=raw_output_file_id,
**credentials,
)
# Access content - handle both direct attribute and method call
if hasattr(_file_content, 'content'):
if hasattr(_file_content, "content"):
content_bytes = _file_content.content # type: ignore[union-attr]
elif hasattr(_file_content, 'read'):
elif hasattr(_file_content, "read"):
content_bytes = await _file_content.read() # type: ignore[misc]
else:
content_bytes = _file_content # type: ignore[assignment]
@ -290,7 +314,9 @@ class CheckBatchCost:
f"Skipping job {unified_object_id} because it is not a valid deployment info"
)
if prom_logger:
prom_logger.record_check_batch_cost_error("deployment_not_found")
prom_logger.record_check_batch_cost_error(
"deployment_not_found"
)
continue
custom_llm_provider = deployment_info.litellm_params.custom_llm_provider
litellm_model_name = deployment_info.litellm_params.model
@ -302,7 +328,11 @@ class CheckBatchCost:
# Pass deployment model_info so custom batch pricing
# (input_cost_per_token_batches etc.) is used for cost calc
deployment_model_info = deployment_info.model_info.model_dump() if deployment_info.model_info else {}
deployment_model_info = (
deployment_info.model_info.model_dump()
if deployment_info.model_info
else {}
)
batch_cost, batch_usage, batch_models = (
await calculate_batch_cost_and_usage(
file_content_dictionary=file_content_as_dict,
@ -349,7 +379,9 @@ class CheckBatchCost:
# Record batch duration (completed_at - created_at)
if prom_logger and response.completed_at and response.created_at:
duration_seconds = float(response.completed_at - response.created_at)
duration_seconds = float(
response.completed_at - response.created_at
)
if duration_seconds >= 0:
prom_logger.record_managed_batch_duration(
duration_seconds=duration_seconds,
@ -358,7 +390,9 @@ class CheckBatchCost:
)
# Track this job for the final metrics summary
processed_models.append((model_name, str(llm_provider) if llm_provider else None))
processed_models.append(
(model_name, str(llm_provider) if llm_provider else None)
)
# mark the job as complete
try:

View file

@ -33,9 +33,7 @@ class CheckResponsesCost:
self.prisma_client: PrismaClient = prisma_client
self.llm_router: Router = llm_router
async def _expire_stale_rows(
self, cutoff: datetime, batch_size: int
) -> int:
async def _expire_stale_rows(self, cutoff: datetime, batch_size: int) -> int:
"""Execute the bounded UPDATE that marks stale rows as 'stale_expired'.
Isolated so it can be swapped / mocked in tests without touching the
@ -74,7 +72,9 @@ class CheckResponsesCost:
rows per invocation to avoid overwhelming the DB when there is a large
backlog.
"""
cutoff = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
cutoff = datetime.now(timezone.utc) - timedelta(
days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS
)
result = await self._expire_stale_rows(cutoff, STALE_OBJECT_CLEANUP_BATCH_SIZE)
if result > 0:
verbose_proxy_logger.warning(
@ -105,7 +105,7 @@ class CheckResponsesCost:
take=MAX_OBJECTS_PER_POLL_CYCLE,
order={"created_at": "asc"},
)
verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check")
completed_jobs = []
@ -120,29 +120,33 @@ class CheckResponsesCost:
# Get the stored response object to extract model information
stored_response = job.file_object
model_name = stored_response.get("model", None)
# Decrypt the response ID
responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(unified_object_id)
responses_id_security, _, _ = (
ResponsesIDSecurity()._decrypt_response_id(unified_object_id)
)
# Prepare metadata with model information for cost tracking
litellm_metadata = {
"user_api_key_user_id": job.created_by or "default-user-id",
}
# Add model information if available
if model_name:
litellm_metadata["model"] = model_name
litellm_metadata["model_group"] = model_name # Use same value for model_group
litellm_metadata["model_group"] = (
model_name # Use same value for model_group
)
response = await litellm.aget_responses(
response_id=responses_id_security,
litellm_metadata=litellm_metadata,
)
verbose_proxy_logger.debug(
f"Response {unified_object_id} status: {response.status}, model: {model_name}"
)
except Exception as e:
verbose_proxy_logger.info(
f"Skipping job {unified_object_id} due to error: {e}"
@ -155,7 +159,7 @@ class CheckResponsesCost:
f"Response {unified_object_id} is complete. Cost automatically tracked by aget_responses."
)
completed_jobs.append(job)
elif response.status in ["failed", "cancelled"]:
verbose_proxy_logger.info(
f"Response {unified_object_id} has status {response.status}, marking as complete"
@ -171,4 +175,3 @@ class CheckResponsesCost:
verbose_proxy_logger.info(
f"Marked {len(completed_jobs)} response jobs as completed"
)

View file

@ -125,7 +125,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
db_data["storage_backend"] = hidden_params["storage_backend"]
if "storage_url" in hidden_params:
db_data["storage_url"] = hidden_params["storage_url"]
verbose_logger.debug(
f"Storage metadata: storage_backend={db_data.get('storage_backend')}, "
f"storage_url={db_data.get('storage_url')}"
@ -285,28 +285,28 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
raise Exception(
"Filtering by 'target_model_names' is not supported when using managed batches."
)
where_clause: Dict[str, Any] = {"file_purpose": "batch"}
# Filter by user who created the batch
if user_api_key_dict.user_id:
where_clause["created_by"] = user_api_key_dict.user_id
if after:
where_clause["id"] = {"gt": after}
# Fetch more than needed to allow for post-fetch filtering
fetch_limit = limit or 20
if target_model_names:
# Fetch extra to account for filtering
fetch_limit = max(fetch_limit * 3, 100)
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where=where_clause,
take=fetch_limit,
order={"created_at": "desc"},
)
batch_objects: List[LiteLLMBatch] = []
for batch in batches:
try:
@ -314,7 +314,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
if len(batch_objects) >= (limit or 20):
break
batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object
batch_data = (
json.loads(batch.file_object)
if isinstance(batch.file_object, str)
else batch.file_object
)
batch_obj = LiteLLMBatch(**batch_data)
batch_obj.id = batch.unified_object_id
batch_objects.append(batch_obj)
@ -324,7 +328,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
f"Failed to parse batch object {batch.unified_object_id}: {e}"
)
continue
return {
"object": "list",
"data": batch_objects,
@ -377,11 +381,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"""
Check if the user has access to a list of file IDs.
Only checks managed (unified) file IDs.
Args:
file_ids: List of file IDs to check access for
user_api_key_dict: User API key authentication details
Raises:
HTTPException: If user doesn't have access to any of the files
"""
@ -419,10 +423,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
### HANDLE TRANSFORMATIONS ###
# Check both completion and acompletion call types
is_completion_call = (
call_type == CallTypes.completion.value
call_type == CallTypes.completion.value
or call_type == CallTypes.acompletion.value
)
if is_completion_call:
messages = data.get("messages")
model = data.get("model", "")
@ -431,22 +435,27 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
if file_ids:
# Check user has access to all managed files
await self.check_file_ids_access(file_ids, user_api_key_dict)
# Check if any files are stored in storage backends and need base64 conversion
# This is needed for Vertex AI/Gemini which requires base64 content
is_vertex_ai = model and ("vertex_ai" in model or "gemini" in model.lower())
is_vertex_ai = model and (
"vertex_ai" in model or "gemini" in model.lower()
)
if is_vertex_ai:
await self._convert_storage_files_to_base64(
messages=messages,
file_ids=file_ids,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
)
model_file_id_mapping = await self.get_model_file_id_mapping(
file_ids, user_api_key_dict.parent_otel_span
)
data["model_file_id_mapping"] = model_file_id_mapping
elif call_type == CallTypes.aresponses.value or call_type == CallTypes.responses.value:
elif (
call_type == CallTypes.aresponses.value
or call_type == CallTypes.responses.value
):
# Handle managed files in responses API input and tools
file_ids = []
@ -611,7 +620,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
if model_id is None:
model_id = cast(
Optional[str],
kwargs.get("litellm_metadata", {}).get("model_info", {}).get("id", None),
kwargs.get("litellm_metadata", {})
.get("model_info", {})
.get("id", None),
)
mapped_file_id: Optional[str] = None
if input_file_id and model_file_id_mapping and model_id:
@ -648,7 +659,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
) -> List[str]:
"""
Gets file ids from responses API input.
The input can be:
- A string (no files)
- A list of input items, where each item can have:
@ -656,32 +667,35 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
- content: a list that can contain items with type: "input_file" and file_id
"""
file_ids: List[str] = []
if isinstance(input, str):
return file_ids
if not isinstance(input, list):
return file_ids
for item in input:
if not isinstance(item, dict):
continue
# Check for direct input_file type
if item.get("type") == "input_file":
file_id = item.get("file_id")
if file_id:
file_ids.append(file_id)
# Check for input_file in content array
content = item.get("content")
if isinstance(content, list):
for content_item in content:
if isinstance(content_item, dict) and content_item.get("type") == "input_file":
if (
isinstance(content_item, dict)
and content_item.get("type") == "input_file"
):
file_id = content_item.get("file_id")
if file_id:
file_ids.append(file_id)
return file_ids
def get_file_ids_from_responses_tools(
@ -689,7 +703,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
) -> List[str]:
"""
Gets file ids from responses API tools parameter.
The tools can contain code_interpreter with container.file_ids:
[
{
@ -699,14 +713,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
]
"""
file_ids: List[str] = []
if not isinstance(tools, list):
return file_ids
for tool in tools:
if not isinstance(tool, dict):
continue
# Check for code_interpreter with container file_ids
if tool.get("type") == "code_interpreter":
container = tool.get("container")
@ -716,7 +730,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
for file_id in container_file_ids:
if isinstance(file_id, str):
file_ids.append(file_id)
return file_ids
def get_vector_store_ids_from_file_search_tools(
@ -916,10 +930,17 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Emit Prometheus metrics for managed file creation
prom_logger = self._get_prometheus_logger()
if prom_logger:
first_model = target_model_names_list[0] if target_model_names_list else None
first_model = (
target_model_names_list[0] if target_model_names_list else None
)
first_provider = ""
if responses:
first_provider = getattr(responses[0], "_hidden_params", {}).get("custom_llm_provider") or ""
first_provider = (
getattr(responses[0], "_hidden_params", {}).get(
"custom_llm_provider"
)
or ""
)
prom_logger.record_managed_file_created(
model=first_model or "",
api_provider=first_provider,
@ -1073,16 +1094,24 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_name=resolved_model_name,
)
setattr(response, file_attr, unified_file_id)
# Use llm_router credentials when available. Without credentials,
# Azure and other auth-required providers return 500/401.
file_object = None
try:
# Import module and use getattr for better testability with mocks
import litellm.proxy.proxy_server as proxy_server_module
_llm_router = getattr(proxy_server_module, 'llm_router', None)
_llm_router = getattr(
proxy_server_module, "llm_router", None
)
if _llm_router is not None and model_id:
_creds = _llm_router.get_deployment_credentials_with_provider(model_id) or {}
_creds = (
_llm_router.get_deployment_credentials_with_provider(
model_id
)
or {}
)
file_object = await litellm.afile_retrieve(
file_id=original_file_id,
**_creds,
@ -1099,7 +1128,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
verbose_logger.warning(
f"Failed to retrieve file object for {file_attr}={original_file_id}: {str(e)}. Storing with None and will fetch on-demand."
)
await self.store_unified_file_id(
file_id=unified_file_id,
file_object=file_object,
@ -1128,6 +1157,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
from litellm.litellm_core_utils.get_llm_provider_logic import (
get_llm_provider,
)
_, batch_provider, _, _ = get_llm_provider(model=model_name)
except Exception:
if "/" in model_name:
@ -1199,7 +1229,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Case 1 : This is not a managed file
if not stored_file_object:
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
# Case 2: Managed file and the file object exists in the database
# The stored file_object has the raw provider ID. Replace with the unified ID
# so callers see a consistent ID (matching Case 3 which does response.id = file_id).
@ -1217,13 +1247,21 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
try:
model_id, model_file_id = next(iter(stored_file_object.model_mappings.items()))
credentials = llm_router.get_deployment_credentials_with_provider(model_id) or {}
response = await litellm.afile_retrieve(file_id=model_file_id, **credentials)
model_id, model_file_id = next(
iter(stored_file_object.model_mappings.items())
)
credentials = (
llm_router.get_deployment_credentials_with_provider(model_id) or {}
)
response = await litellm.afile_retrieve(
file_id=model_file_id, **credentials
)
response.id = file_id # Replace with unified ID
return response
except Exception as e:
raise Exception(f"Failed to retrieve file {file_id} from provider: {str(e)}") from e
raise Exception(
f"Failed to retrieve file {file_id} from provider: {str(e)}"
) from e
async def afile_list(
self,
@ -1245,19 +1283,19 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
import litellm.proxy.proxy_server as proxy_server_module
# Check if the scheduler has the batch cost checking job registered
scheduler = getattr(proxy_server_module, 'scheduler', None)
scheduler = getattr(proxy_server_module, "scheduler", None)
if scheduler is None:
return False
# Check if the check_batch_cost_job exists in the scheduler
try:
job = scheduler.get_job('check_batch_cost_job')
job = scheduler.get_job("check_batch_cost_job")
if job is not None:
return True
except Exception:
# Job not found or scheduler doesn't support get_job
pass
return False
except Exception as e:
verbose_logger.warning(
@ -1265,28 +1303,26 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
return False
async def _get_batches_referencing_file(
self, file_id: str
) -> List[Dict[str, Any]]:
async def _get_batches_referencing_file(self, file_id: str) -> List[Dict[str, Any]]:
"""
Find batches that reference this file and still need cost tracking.
Find batches that are in non-terminal state and have not yet been processed by CheckBatchCost.
Args:
file_id: The unified file ID to check
Returns:
List of batch objects referencing this file in non-terminal state
(max 10 for error message display)
"""
# Prepare list of file IDs to check (both unified and provider IDs)
file_ids_to_check = [file_id]
# Get model-specific file IDs for this unified file ID if it's a managed file
try:
model_file_id_mapping = await self.get_model_file_id_mapping(
[file_id], litellm_parent_otel_span=None
)
if model_file_id_mapping and file_id in model_file_id_mapping:
# Add all provider file IDs for this unified file
provider_file_ids = list(model_file_id_mapping[file_id].values())
@ -1296,59 +1332,67 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
f"Could not get model file ID mapping for {file_id}: {e}. "
f"Will only check unified file ID."
)
MAX_MATCHES_TO_RETURN = 10
MAX_MATCHES_TO_RETURN = 10
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where={
"file_purpose": "batch",
"batch_processed": False,
"status": {"not_in": ["failed", "expired", "cancelled"]}
"status": {"not_in": ["failed", "expired", "cancelled"]},
},
take=MAX_MATCHES_TO_RETURN,
order={"created_at": "desc"},
)
referencing_batches = []
for batch in batches:
try:
# Parse the batch file_object to check for file references
batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object
batch_data = (
json.loads(batch.file_object)
if isinstance(batch.file_object, str)
else batch.file_object
)
# Extract file IDs from batch
# Batches typically reference the unified file ID in input_file_id
# Output and error files are generated by the provider
input_file_id = batch_data.get("input_file_id")
output_file_id = batch_data.get("output_file_id")
error_file_id = batch_data.get("error_file_id")
referenced_file_ids = [fid for fid in [input_file_id, output_file_id, error_file_id] if fid]
referenced_file_ids = [
fid for fid in [input_file_id, output_file_id, error_file_id] if fid
]
# Check if any referenced file ID matches the file we're trying to delete
if any(ref_id in file_ids_to_check for ref_id in referenced_file_ids):
referencing_batches.append({
"batch_id": batch.unified_object_id,
"status": batch.status,
"created_at": batch.created_at,
})
referencing_batches.append(
{
"batch_id": batch.unified_object_id,
"status": batch.status,
"created_at": batch.created_at,
}
)
except Exception as e:
verbose_logger.warning(
f"Error parsing batch object {batch.unified_object_id}: {e}"
)
continue
return referencing_batches
async def _check_file_deletion_allowed(self, file_id: str) -> None:
"""
Check if file deletion should be blocked due to batch references.
Blocks deletion if:
1. File is referenced by any batch in non-terminal state, AND
2. Batch polling is configured (user wants cost tracking)
Args:
file_id: The unified file ID to check
Raises:
HTTPException: If file deletion should be blocked
"""
@ -1356,39 +1400,45 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
if not self._is_batch_polling_enabled():
# Batch polling not configured, allow deletion
return
# Check if file is referenced by any non-terminal batches
referencing_batches = await self._get_batches_referencing_file(file_id)
if referencing_batches:
# File is referenced by non-terminal batches and polling is enabled
MAX_BATCHES_IN_ERROR = 5 # Limit batches shown in error message for readability
MAX_BATCHES_IN_ERROR = (
5 # Limit batches shown in error message for readability
)
# Show up to MAX_BATCHES_IN_ERROR in the error message
batches_to_show = referencing_batches[:MAX_BATCHES_IN_ERROR]
batch_statuses = [f"{b['batch_id']}: {b['status']}" for b in batches_to_show]
batch_statuses = [
f"{b['batch_id']}: {b['status']}" for b in batches_to_show
]
# Determine the count message
count_message = f"{len(referencing_batches)}"
if len(referencing_batches) >= 10: # MAX_MATCHES_TO_RETURN from _get_batches_referencing_file
if (
len(referencing_batches) >= 10
): # MAX_MATCHES_TO_RETURN from _get_batches_referencing_file
count_message = "10+"
error_message = (
f"Cannot delete file {file_id}. "
f"The file is referenced by {count_message} batch(es) in non-terminal state"
)
# Add specific batch details if not too many
if len(referencing_batches) <= MAX_BATCHES_IN_ERROR:
error_message += f": {', '.join(batch_statuses)}. "
else:
error_message += f" (showing {MAX_BATCHES_IN_ERROR} most recent): {', '.join(batch_statuses)}. "
error_message += (
f"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. "
f"Alternatively, wait for all batches to complete and for cost to be computed (batch_processed=true)."
)
# Record blocked deletion metric
prom_logger = self._get_prometheus_logger()
if prom_logger:
@ -1419,7 +1469,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
specific_model_file_id_mapping = model_file_id_mapping.get(file_id)
if specific_model_file_id_mapping:
# Remove conflicting keys from data to avoid duplicate keyword arguments
filtered_data = {k: v for k, v in data.items() if k not in ("model", "file_id")}
filtered_data = {
k: v for k, v in data.items() if k not in ("model", "file_id")
}
for model_id, model_file_id in specific_model_file_id_mapping.items():
delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **filtered_data) # type: ignore
@ -1480,7 +1532,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
) -> None:
"""
Convert files stored in storage backends to base64 format for Vertex AI/Gemini.
This method checks if any managed files are stored in storage backends,
downloads them, and converts them to base64 format in the messages.
"""
@ -1488,29 +1540,29 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
for file_id in file_ids:
# Check if this is a base64 encoded unified file ID
decoded_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
if not decoded_unified_file_id:
continue
# Check database for storage backend info
# IMPORTANT: The database stores the base64 encoded unified_file_id (not the decoded version)
# So we query with the original file_id (which is base64 encoded)
db_file = await self.prisma_client.db.litellm_managedfiletable.find_first(
where={"unified_file_id": file_id}
)
if not db_file or not db_file.storage_backend or not db_file.storage_url:
continue
# File is stored in a storage backend, download and convert to base64
try:
from litellm.llms.base_llm.files.storage_backend_factory import (
get_storage_backend,
)
storage_backend_name = db_file.storage_backend
storage_url = db_file.storage_url
# Get storage backend (uses same env vars as callback)
try:
storage_backend = get_storage_backend(storage_backend_name)
@ -1519,18 +1571,22 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
f"Storage backend '{storage_backend_name}' error for file {file_id}: {str(e)}"
)
continue
file_content = await storage_backend.download_file(storage_url)
# Determine content type from file object
content_type = self._get_content_type_from_file_object(db_file.file_object)
content_type = self._get_content_type_from_file_object(
db_file.file_object
)
# Convert to base64
base64_data = base64.b64encode(file_content).decode("utf-8")
base64_data_uri = f"data:{content_type};base64,{base64_data}"
# Update messages to use base64 instead of file_id
self._update_messages_with_base64_data(messages, file_id, base64_data_uri, content_type)
self._update_messages_with_base64_data(
messages, file_id, base64_data_uri, content_type
)
except Exception as e:
verbose_logger.exception(
f"Error converting file {file_id} from storage backend to base64: {str(e)}"
@ -1541,21 +1597,21 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
def _get_content_type_from_file_object(self, file_object: Optional[Any]) -> str:
"""
Determine content type from file object.
Uses the MIME type utility for consistent detection and normalization.
Args:
file_object: The file object from the database (can be dict, JSON string, or None)
Returns:
str: MIME type (defaults to "application/octet-stream" if cannot be determined)
"""
# Use utility function for detection
content_type = get_content_type_from_file_object(file_object)
# Normalize for Gemini/Vertex AI (requires image/jpeg, not image/jpg)
content_type = normalize_mime_type_for_provider(content_type, provider="gemini")
return content_type
def _update_messages_with_base64_data(
@ -1567,7 +1623,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
) -> None:
"""
Update messages to replace file_id with base64 data URI.
Args:
messages: List of messages to update
file_id: The file ID to replace
@ -1582,7 +1638,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
if element.get("type") == "file":
file_element = cast(ChatCompletionFileObject, element)
file_element_file = file_element.get("file", {})
if file_element_file.get("file_id") == file_id:
# Replace file_id with base64 data
file_element_file["file_data"] = base64_data_uri
@ -1590,7 +1646,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
file_element_file["format"] = content_type
# Remove file_id to ensure only file_data is used
file_element_file.pop("file_id", None)
verbose_logger.debug(
f"Converted file {file_id} from storage backend to base64 with format {content_type}"
)

View file

@ -41,7 +41,7 @@ class _PROXY_LiteLLMManagedVectorStores(
):
"""
Managed vector stores with target_model_names support.
This class provides functionality to:
- Create vector stores across multiple models
- Retrieve vector stores by unified ID
@ -77,14 +77,14 @@ class _PROXY_LiteLLMManagedVectorStores(
) -> str:
"""
Generate the format string for the unified vector store ID.
Format:
litellm_proxy:vector_store;unified_id,<uuid>;target_model_names,<models>;resource_id,<vs_id>;model_id,<model_id>
"""
# VectorStoreCreateResponse is a TypedDict, so resource_object is a dictionary
# Extract provider resource ID from the response
provider_resource_id = resource_object.get("id", "")
# Model ID is stored in hidden params if the response object supports it
# For TypedDict responses, we need to check if _hidden_params was added
hidden_params: Dict[str, Any] = {}
@ -109,20 +109,18 @@ class _PROXY_LiteLLMManagedVectorStores(
) -> VectorStoreCreateResponse:
"""
Create a vector store for a specific model.
Args:
llm_router: LiteLLM router instance
model: Model name to create vector store for
request_data: Request data for vector store creation
litellm_parent_otel_span: OpenTelemetry span for tracing
Returns:
VectorStoreCreateResponse from the provider
"""
# Use the router to create the vector store
response = await llm_router.avector_store_create(
model=model, **request_data
)
response = await llm_router.avector_store_create(model=model, **request_data)
return response
# ============================================================================
@ -139,14 +137,14 @@ class _PROXY_LiteLLMManagedVectorStores(
) -> VectorStoreCreateResponse:
"""
Create a vector store across multiple models.
Args:
create_request: Vector store creation request parameters
llm_router: LiteLLM router instance
target_model_names_list: List of target model names
litellm_parent_otel_span: OpenTelemetry span for tracing
user_api_key_dict: User API key authentication details
Returns:
VectorStoreCreateResponse with unified ID
"""
@ -196,7 +194,7 @@ class _PROXY_LiteLLMManagedVectorStores(
# VectorStoreCreateResponse is a TypedDict, so we need to create a new dict with the unified ID
response = responses[0].copy()
response["id"] = unified_id
verbose_logger.info(
f"Successfully created managed vector store with unified ID: {unified_id}"
)
@ -212,13 +210,13 @@ class _PROXY_LiteLLMManagedVectorStores(
) -> Dict[str, Any]:
"""
List vector stores created by a user.
Args:
user_api_key_dict: User API key authentication details
limit: Maximum number of vector stores to return
after: Cursor for pagination
order: Sort order ('asc' or 'desc')
Returns:
Dictionary with list of vector stores and pagination info
"""
@ -238,23 +236,23 @@ class _PROXY_LiteLLMManagedVectorStores(
) -> bool:
"""
Check if user has access to a vector store.
Args:
vector_store_id: The unified vector store ID
user_api_key_dict: User API key authentication details
Returns:
True if user has access, False otherwise
"""
is_unified_id = is_base64_encoded_unified_id(vector_store_id)
if is_unified_id:
# Check access for managed vector store
return await self.can_user_access_unified_resource_id(
vector_store_id,
user_api_key_dict,
)
# Not a managed vector store, allow access
return True
@ -263,24 +261,22 @@ class _PROXY_LiteLLMManagedVectorStores(
) -> bool:
"""
Check if user has access to a managed vector store in request data.
Args:
data: Request data containing vector_store_id
user_api_key_dict: User API key authentication details
Returns:
True if this is a managed vector store and user has access
Raises:
HTTPException: If user doesn't have access
"""
vector_store_id = cast(Optional[str], data.get("vector_store_id"))
is_unified_id = (
is_base64_encoded_unified_id(vector_store_id)
if vector_store_id
else False
is_base64_encoded_unified_id(vector_store_id) if vector_store_id else False
)
if is_unified_id and vector_store_id:
if await self.can_user_access_unified_resource_id(
vector_store_id, user_api_key_dict
@ -291,7 +287,7 @@ class _PROXY_LiteLLMManagedVectorStores(
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}",
)
return False
# ============================================================================
@ -307,18 +303,18 @@ class _PROXY_LiteLLMManagedVectorStores(
) -> Union[Exception, str, Dict, None]:
"""
Pre-call hook to handle vector store operations.
This hook intercepts vector store requests and:
- Validates access for managed vector stores
- Transforms unified IDs to provider-specific IDs
- Adds model routing information
Args:
user_api_key_dict: User API key authentication details
cache: Cache instance
data: Request data
call_type: Type of call being made
Returns:
Modified request data or None
"""
@ -330,40 +326,40 @@ class _PROXY_LiteLLMManagedVectorStores(
# Handle vector store search operations
if call_type == "avector_store_search":
vector_store_id = data.get("vector_store_id")
if vector_store_id:
# Check if it's a managed vector store ID
decoded_id = is_base64_encoded_unified_id(vector_store_id)
if decoded_id:
verbose_logger.debug(
f"Processing managed vector store search: {vector_store_id}"
)
# Check access
has_access = await self.can_user_access_unified_resource_id(
vector_store_id, user_api_key_dict
)
if not has_access:
raise HTTPException(
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}",
)
# Parse the unified ID to extract components
parsed_id = parse_unified_id(vector_store_id)
if parsed_id:
# Extract the model ID and provider resource ID
model_id = parsed_id.get("model_id")
provider_resource_id = parsed_id.get("provider_resource_id")
target_model_names = parsed_id.get("target_model_names", [])
verbose_logger.debug(
f"Decoded vector store - model_id: {model_id}, provider_resource_id: {provider_resource_id}, target_model_names: {target_model_names}"
)
# Determine which model to use for routing
# Priority: model_id (deployment ID) > first target_model_name
routing_model = None
@ -371,28 +367,28 @@ class _PROXY_LiteLLMManagedVectorStores(
routing_model = model_id
elif target_model_names and len(target_model_names) > 0:
routing_model = target_model_names[0]
# Set the model for routing
if routing_model:
data["model"] = routing_model
verbose_logger.info(
f"Routing vector store search to model: {routing_model}"
)
# Replace the unified ID with the provider-specific ID
if provider_resource_id:
data["vector_store_id"] = provider_resource_id
verbose_logger.debug(
f"Replaced unified ID with provider resource ID: {provider_resource_id}"
)
# Handle vector store retrieve/delete operations
elif call_type in ("avector_store_retrieve", "avector_store_delete"):
await self.check_managed_vector_store_access(data, user_api_key_dict)
# If it's a managed vector store, we'll handle it in the endpoint
# No need to transform here as the endpoint will route to the hook
return data
# ============================================================================
@ -407,15 +403,15 @@ class _PROXY_LiteLLMManagedVectorStores(
) -> Any:
"""
Post-call hook to transform responses.
This hook can be used to transform responses if needed.
For now, it just passes through the response.
Args:
data: Request data
user_api_key_dict: User API key authentication details
response: Response from the provider
Returns:
Potentially modified response
"""
@ -436,21 +432,21 @@ class _PROXY_LiteLLMManagedVectorStores(
) -> List[Dict]:
"""
Filter deployments based on vector store availability.
This is used by the router to select only deployments that have
the vector store available.
Note: This method signature is a compromise between CustomLogger and BaseManagedResource
parent classes which have incompatible signatures. The type: ignore[override] is necessary
due to this multiple inheritance conflict.
Args:
model: Model name
healthy_deployments: List of healthy deployments
messages: Messages (unused for vector stores, required by CustomLogger interface)
request_kwargs: Request kwargs containing vector_store_id and mappings
parent_otel_span: OpenTelemetry span for tracing
Returns:
Filtered list of deployments
"""

View file

@ -2,7 +2,6 @@
Enterprise internal user management endpoints
"""
from fastapi import APIRouter, Depends, HTTPException
from litellm.proxy._types import UserAPIKeyAuth

View file

@ -147,12 +147,12 @@ async def list_vector_stores(
vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db(
prisma_client=prisma_client
)
# Also clean up in-memory registry to remove any deleted vector stores
if litellm.vector_store_registry is not None:
db_vector_store_ids = {
vs.get("vector_store_id")
for vs in vector_stores_from_db
vs.get("vector_store_id")
for vs in vector_stores_from_db
if vs.get("vector_store_id")
}
# Remove any in-memory vector stores that no longer exist in database

View file

@ -39,15 +39,23 @@ class EmailEvent(str, enum.Enum):
soft_budget_crossed = "Soft Budget Crossed"
max_budget_alert = "Max Budget Alert"
class EmailEventSettings(BaseModel):
event: EmailEvent
enabled: bool
class EmailEventSettingsUpdateRequest(BaseModel):
settings: List[EmailEventSettings]
class EmailEventSettingsResponse(BaseModel):
settings: List[EmailEventSettings]
class DefaultEmailSettings(BaseModel):
"""Default settings for email events"""
settings: Dict[EmailEvent, bool] = Field(
default_factory=lambda: {
EmailEvent.virtual_key_created: True, # On by default
@ -57,10 +65,12 @@ class DefaultEmailSettings(BaseModel):
EmailEvent.max_budget_alert: True, # On by default
}
)
def to_dict(self) -> Dict[str, bool]:
"""Convert to dictionary with string keys for storage"""
return {event.value: enabled for event, enabled in self.settings.items()}
@classmethod
def get_defaults(cls) -> Dict[str, bool]:
"""Get the default settings as a dictionary with string keys"""
return cls().to_dict()
return cls().to_dict()

View file

@ -51,8 +51,7 @@ class PrometheusLabelFactoryContext:
self.enum_values = enum_values
enum_dict = enum_values.model_dump()
self._sanitized_enum: Dict[str, Optional[str]] = {
k: _sanitize_prometheus_label_value(v)
for k, v in enum_dict.items()
k: _sanitize_prometheus_label_value(v) for k, v in enum_dict.items()
}
self._custom_by_sanitized_key: Dict[str, Optional[str]] = {}
if enum_values.custom_metadata_labels is not None:

View file

@ -294,9 +294,7 @@ class Authenticator:
access_token_url = os.getenv(
"GITHUB_COPILOT_ACCESS_TOKEN_URL", DEFAULT_GITHUB_ACCESS_TOKEN_URL
)
client_id = os.getenv(
"GITHUB_COPILOT_CLIENT_ID", DEFAULT_GITHUB_CLIENT_ID
)
client_id = os.getenv("GITHUB_COPILOT_CLIENT_ID", DEFAULT_GITHUB_CLIENT_ID)
for attempt in range(max_attempts):
try:

View file

@ -79,7 +79,9 @@ class BasePassthroughUtils:
for header_name, header_value in request_headers.items():
if header_name.lower().startswith(PASS_THROUGH_HEADER_PREFIX):
# Strip the 'x-pass-' prefix and normalize to lowercase
actual_header_name = header_name[len(PASS_THROUGH_HEADER_PREFIX) :].lower()
actual_header_name = header_name[
len(PASS_THROUGH_HEADER_PREFIX) :
].lower()
if actual_header_name in _PASS_THROUGH_PROTECTED_HEADERS or any(
actual_header_name.startswith(p)
for p in _PASS_THROUGH_PROTECTED_HEADER_PREFIXES

View file

@ -784,7 +784,7 @@ class UserAPIKeyLabelValues:
org_id: Optional[str] = None
org_alias: Optional[str] = None
#Added for test compatibility.
# Added for test compatibility.
def __init__(self, **kwargs: Any) -> None:
"""
Match former Pydantic behavior: unknown keys are ignored; ``api_key_hash`` maps to

View file

@ -23,8 +23,6 @@ def event_loop():
loop.close()
@pytest.fixture(scope="function", autouse=True)
def setup_and_teardown():
"""

View file

@ -289,11 +289,10 @@ async def test_increment_remaining_budget_metrics(prometheus_logger):
future_reset_time_team = datetime.now() + timedelta(hours=10)
future_reset_time_key = datetime.now() + timedelta(hours=12)
# Mock the get_team_object and get_key_object functions to return objects with budget reset times
with patch(
"litellm.proxy.auth.auth_checks.get_team_object"
) as mock_get_team, patch(
"litellm.proxy.auth.auth_checks.get_key_object"
) as mock_get_key:
with (
patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team,
patch("litellm.proxy.auth.auth_checks.get_key_object") as mock_get_key,
):
mock_get_team.return_value = MagicMock(budget_reset_at=future_reset_time_team)
mock_get_key.return_value = MagicMock(budget_reset_at=future_reset_time_key)
@ -1518,9 +1517,12 @@ async def test_initialize_remaining_budget_metrics(prometheus_logger):
"""
litellm.prometheus_initialize_budget_metrics = True
# Mock the prisma client and get_paginated_teams function
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams:
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams,
):
# Create mock team data with proper datetime objects for budget_reset_at
future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now
mock_teams = [
@ -1613,11 +1615,15 @@ async def test_initialize_remaining_budget_metrics_exception_handling(
"""
litellm.prometheus_initialize_budget_metrics = True
# Mock the prisma client and get_paginated_teams function to raise an exception
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams, patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys:
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams,
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys,
):
# Make get_paginated_teams raise an exception
mock_get_teams.side_effect = Exception("Database error")
mock_list_keys.side_effect = Exception("Key listing error")
@ -1636,9 +1642,7 @@ async def test_initialize_remaining_budget_metrics_exception_handling(
# Mock litellm_organizationtable to raise an exception for org budget metrics
mock_orgtable = MagicMock()
mock_orgtable.find_many = MagicMock(
side_effect=Exception("Org database error")
)
mock_orgtable.find_many = MagicMock(side_effect=Exception("Org database error"))
mock_orgtable.count = MagicMock(side_effect=Exception("Org count error"))
mock_db = MagicMock()
@ -1699,9 +1703,12 @@ async def test_initialize_api_key_budget_metrics(prometheus_logger):
"""
litellm.prometheus_initialize_budget_metrics = True
# Mock the prisma client and _list_key_helper function
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys:
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys,
):
# Create mock key data with proper datetime objects for budget_reset_at
future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now
key1 = UserAPIKeyAuth(

View file

@ -944,8 +944,8 @@ def test_callback_failure_metric_different_callbacks(prometheus_logger):
async def test_langfuse_callback_failure_metric(prometheus_logger):
"""
Test that Langfuse callback failures are properly tracked in Prometheus metrics.
This test verifies that when Langfuse logging fails, the
This test verifies that when Langfuse logging fails, the
litellm_callback_logging_failures_metric is incremented with callback_name="langfuse".
"""
from unittest.mock import MagicMock, patch
@ -957,16 +957,20 @@ async def test_langfuse_callback_failure_metric(prometheus_logger):
# Get initial value
initial_value = 0
try:
initial_value = prometheus_logger.litellm_callback_logging_failures_metric.labels(
callback_name="langfuse"
)._value.get()
initial_value = (
prometheus_logger.litellm_callback_logging_failures_metric.labels(
callback_name="langfuse"
)._value.get()
)
except Exception:
initial_value = 0
# Create Langfuse logger with mocked initialization
with patch("litellm.integrations.langfuse.langfuse_prompt_management.langfuse_client_init"):
with patch(
"litellm.integrations.langfuse.langfuse_prompt_management.langfuse_client_init"
):
langfuse_logger = LangfusePromptManagement()
# Mock the log_event_on_langfuse to raise an exception
with patch(
"litellm.integrations.langfuse.langfuse_prompt_management.LangFuseHandler.get_langfuse_logger_for_request"
@ -974,14 +978,16 @@ async def test_langfuse_callback_failure_metric(prometheus_logger):
mock_logger = MagicMock()
mock_logger.log_event_on_langfuse.side_effect = Exception("Langfuse API error")
mock_get_logger.return_value = mock_logger
# Mock handle_callback_failure to track calls
with patch.object(prometheus_logger, "increment_callback_logging_failure") as mock_increment:
with patch.object(
prometheus_logger, "increment_callback_logging_failure"
) as mock_increment:
# Inject prometheus logger into the langfuse logger
langfuse_logger.handle_callback_failure = lambda callback_name: mock_increment(
callback_name=callback_name
langfuse_logger.handle_callback_failure = (
lambda callback_name: mock_increment(callback_name=callback_name)
)
# Call async_log_success_event - should catch exception and increment metric
await langfuse_logger.async_log_success_event(
kwargs={},
@ -989,10 +995,10 @@ async def test_langfuse_callback_failure_metric(prometheus_logger):
start_time=None,
end_time=None,
)
# Verify that increment was called with correct callback name
mock_increment.assert_called_once_with(callback_name="langfuse")
print("✓ Langfuse callback failure metric test passed")
@ -1000,8 +1006,8 @@ async def test_langfuse_callback_failure_metric(prometheus_logger):
async def test_langfuse_otel_callback_failure_metric(prometheus_logger):
"""
Test that Langfuse OTEL callback failures are properly tracked in Prometheus metrics.
This test verifies that when Langfuse OTEL logging fails, the
This test verifies that when Langfuse OTEL logging fails, the
litellm_callback_logging_failures_metric is incremented with callback_name="langfuse_otel".
"""
from unittest.mock import MagicMock, patch
@ -1011,50 +1017,58 @@ async def test_langfuse_otel_callback_failure_metric(prometheus_logger):
# Get initial value
initial_value = 0
try:
initial_value = prometheus_logger.litellm_callback_logging_failures_metric.labels(
callback_name="langfuse_otel"
)._value.get()
initial_value = (
prometheus_logger.litellm_callback_logging_failures_metric.labels(
callback_name="langfuse_otel"
)._value.get()
)
except Exception:
initial_value = 0
# Create Langfuse OTEL logger with mocked initialization
with patch("litellm.integrations.opentelemetry.OpenTelemetry.__init__", return_value=None):
with patch(
"litellm.integrations.opentelemetry.OpenTelemetry.__init__", return_value=None
):
langfuse_otel_logger = LangfuseOtelLogger(callback_name="langfuse_otel")
langfuse_otel_logger.callback_name = "langfuse_otel"
# Mock handle_callback_failure to track calls
with patch.object(prometheus_logger, "increment_callback_logging_failure") as mock_increment:
with patch.object(
prometheus_logger, "increment_callback_logging_failure"
) as mock_increment:
# Inject prometheus logger into the langfuse otel logger
langfuse_otel_logger.handle_callback_failure = lambda callback_name: mock_increment(
callback_name=callback_name
langfuse_otel_logger.handle_callback_failure = (
lambda callback_name: mock_increment(callback_name=callback_name)
)
# Test that the OpenTelemetry base class set_attributes exception handler works
# This is where langfuse_otel failures are caught and tracked
with patch.object(langfuse_otel_logger, "set_attributes") as mock_set_attributes:
with patch.object(
langfuse_otel_logger, "set_attributes"
) as mock_set_attributes:
# Simulate the exception handling in set_attributes
def set_attributes_with_error(*args, **kwargs):
# This simulates what happens in the real set_attributes method
try:
raise Exception("Attribute error")
except Exception as e:
langfuse_otel_logger.handle_callback_failure(callback_name=langfuse_otel_logger.callback_name)
langfuse_otel_logger.handle_callback_failure(
callback_name=langfuse_otel_logger.callback_name
)
mock_set_attributes.side_effect = set_attributes_with_error
# Call set_attributes
try:
langfuse_otel_logger.set_attributes(
span=MagicMock(),
kwargs={},
response_obj={}
span=MagicMock(), kwargs={}, response_obj={}
)
except Exception:
pass
# Verify that increment was called with correct callback name
mock_increment.assert_called_with(callback_name="langfuse_otel")
print("✓ Langfuse OTEL callback failure metric test passed")

View file

@ -49,10 +49,13 @@ async def test_enterprise_custom_auth_returns_string():
mock_user_auth = AsyncMock(return_value="sk-test-key")
request = MagicMock(spec=Request)
with patch(
"litellm.proxy.auth.user_api_key_auth.enterprise_custom_auth", mock_user_auth
), patch("litellm.proxy.proxy_server.master_key", "sk-1234"), patch(
"litellm.proxy.proxy_server.prisma_client", MagicMock()
with (
patch(
"litellm.proxy.auth.user_api_key_auth.enterprise_custom_auth",
mock_user_auth,
),
patch("litellm.proxy.proxy_server.master_key", "sk-1234"),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
):
# Verify the key is correctly handled in _user_api_key_auth_builder
with patch(

View file

@ -105,7 +105,9 @@ async def test_apply_guardrail_endpoint_with_presidio_guardrail():
mock_guardrail = Mock(spec=CustomGuardrail)
# Simulate masking PII entities - returns GenericGuardrailAPIInputs (dict with texts key)
mock_guardrail.apply_guardrail = AsyncMock(
return_value={"texts": ["My name is [PERSON] and my email is [EMAIL_ADDRESS]"]}
return_value={
"texts": ["My name is [PERSON] and my email is [EMAIL_ADDRESS]"]
}
)
# Configure the registry to return our mock guardrail

View file

@ -31,12 +31,16 @@ class TestAvailableEnterpriseUsers:
self, client, mock_user_api_key_auth
):
"""Test when max_users is set and user count is within limit"""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.proxy_server.premium_user",
True,
), patch(
"litellm.proxy.proxy_server.premium_user_data",
{"max_users": 10},
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.proxy_server.premium_user",
True,
),
patch(
"litellm.proxy.proxy_server.premium_user_data",
{"max_users": 10},
),
):
# Mock database count
mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=5)
@ -66,12 +70,16 @@ class TestAvailableEnterpriseUsers:
self, client, mock_user_api_key_auth
):
"""Test when max_users is not set (premium_user_data is None or doesn't contain max_users)"""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.proxy_server.premium_user",
True,
), patch(
"litellm.proxy.proxy_server.premium_user_data",
None,
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.proxy_server.premium_user",
True,
),
patch(
"litellm.proxy.proxy_server.premium_user_data",
None,
),
):
# Mock database count
mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=3)
@ -99,12 +107,16 @@ class TestAvailableEnterpriseUsers:
self, client, mock_user_api_key_auth
):
"""Test the current bug where total_users_remaining can be negative"""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.proxy_server.premium_user",
True,
), patch(
"litellm.proxy.proxy_server.premium_user_data",
{"key": "value"},
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.proxy_server.premium_user",
True,
),
patch(
"litellm.proxy.proxy_server.premium_user_data",
{"key": "value"},
),
):
# Mock database count higher than max_users to trigger the bug
mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=8)
@ -140,12 +152,15 @@ class TestAvailableEnterpriseUsers:
"""Test when prisma_client is None (no database connection)"""
from litellm.proxy._types import CommonProxyErrors
with patch(
"litellm.proxy.proxy_server.prisma_client",
None,
), patch(
"litellm.proxy.proxy_server.premium_user",
True,
with (
patch(
"litellm.proxy.proxy_server.prisma_client",
None,
),
patch(
"litellm.proxy.proxy_server.premium_user",
True,
),
):
# Override the dependency
client.app.dependency_overrides[mock_user_api_key_auth] = lambda: {

View file

@ -110,9 +110,7 @@ class TestAzureDocumentIntelligencePagesParam:
model="azure_ai/doc-intelligence/prebuilt-layout",
optional_params={"pages": "1-3,5"},
)
assert (
f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url
), url
assert f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url, url
assert "pages=1-3,5" in url, url
assert "/documentintelligence/documentModels/prebuilt-layout:analyze" in url
@ -168,4 +166,3 @@ class TestAzureDocumentIntelligencePagesParam:
assert "pages=3,4,5,6,7,8,9" in url
assert req.data == {"urlSource": "https://example.com/x.pdf"}

View file

@ -30,6 +30,7 @@ def no_invitation_wait(monkeypatch):
monkeypatch.setattr(BaseEmailLogger, "_wait_for_invitation_creation", _noop)
@pytest.fixture
def base_email_logger():
return BaseEmailLogger()
@ -283,7 +284,10 @@ async def test_send_key_created_email_without_key(
mock_send_email.assert_called_once()
call_args = mock_send_email.call_args[1]
assert "sk-secret-key-456" not in call_args["html_body"]
assert "[Key hidden for security - retrieve from dashboard]" in call_args["html_body"]
assert (
"[Key hidden for security - retrieve from dashboard]"
in call_args["html_body"]
)
@pytest.mark.asyncio
@ -317,7 +321,10 @@ async def test_send_key_rotated_email_without_key(
mock_send_email.assert_called_once()
call_args = mock_send_email.call_args[1]
assert "sk-secret-rotated-789" not in call_args["html_body"]
assert "[Key hidden for security - retrieve from dashboard]" in call_args["html_body"]
assert (
"[Key hidden for security - retrieve from dashboard]"
in call_args["html_body"]
)
@pytest.mark.asyncio
@ -371,52 +378,52 @@ async def test_get_invitation_link_creates_new_when_none_exist(base_email_logger
"""Test that _get_invitation_link creates a new invitation when none exist"""
# Mock prisma client with no existing invitation rows
mock_prisma = mock.MagicMock()
# Mock find_many to return empty list (no existing invitations)
async def mock_find_many_empty(*args, **kwargs):
return []
mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_empty
# Mock the create_invitation_for_user function
mock_created_invitation = mock.MagicMock()
mock_created_invitation.id = "new-invitation-id"
with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
with mock.patch(
"litellm.proxy.management_helpers.user_invitation.create_invitation_for_user",
return_value=mock_created_invitation
return_value=mock_created_invitation,
) as mock_create_invitation:
# Execute
result = await base_email_logger._get_invitation_link(
user_id="test-user", base_url="http://test.com"
)
# Verify that create_invitation_for_user was called
mock_create_invitation.assert_called_once()
call_args = mock_create_invitation.call_args[1]
assert call_args["data"].user_id == "test-user"
assert call_args["user_api_key_dict"].user_id == "test-user"
# Verify the returned link uses the new invitation ID
assert result == "http://test.com/ui?invitation_id=new-invitation-id"
@pytest.mark.asyncio
@pytest.mark.asyncio
async def test_get_invitation_link_uses_existing_when_available(base_email_logger):
"""Test that _get_invitation_link uses existing invitation when available"""
# Mock prisma client with existing invitation row
mock_invitation_row = mock.MagicMock()
mock_invitation_row.id = "existing-invitation-id"
mock_prisma = mock.MagicMock()
# Mock find_many to return existing invitation
async def mock_find_many_existing(*args, **kwargs):
return [mock_invitation_row]
mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_existing
with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
with mock.patch(
"litellm.proxy.management_helpers.user_invitation.create_invitation_for_user"
@ -425,10 +432,10 @@ async def test_get_invitation_link_uses_existing_when_available(base_email_logge
result = await base_email_logger._get_invitation_link(
user_id="test-user", base_url="http://test.com"
)
# Verify that create_invitation_for_user was NOT called
mock_create_invitation.assert_not_called()
# Verify the returned link uses the existing invitation ID
assert result == "http://test.com/ui?invitation_id=existing-invitation-id"
@ -438,33 +445,33 @@ async def test_get_invitation_link_creates_new_when_list_is_none(base_email_logg
"""Test that _get_invitation_link creates a new invitation when invitation_rows is None"""
# Mock prisma client to return None
mock_prisma = mock.MagicMock()
# Mock find_many to return None
async def mock_find_many_none(*args, **kwargs):
return None
mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_none
# Mock the create_invitation_for_user function
mock_created_invitation = mock.MagicMock()
mock_created_invitation.id = "new-invitation-from-none"
with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
with mock.patch(
"litellm.proxy.management_helpers.user_invitation.create_invitation_for_user",
return_value=mock_created_invitation
return_value=mock_created_invitation,
) as mock_create_invitation:
# Execute
result = await base_email_logger._get_invitation_link(
user_id="test-user", base_url="http://test.com"
)
# Verify that create_invitation_for_user was called
mock_create_invitation.assert_called_once()
call_args = mock_create_invitation.call_args[1]
assert call_args["data"].user_id == "test-user"
assert call_args["user_api_key_dict"].user_id == "test-user"
# Verify the returned link uses the new invitation ID
assert result == "http://test.com/ui?invitation_id=new-invitation-from-none"
@ -495,13 +502,15 @@ async def test_get_email_params_user_invitation(
user_email="test@example.com",
)
assert result.logo_url == "https://litellm-listing.s3.amazonaws.com/litellm_logo.png"
assert (
result.logo_url
== "https://litellm-listing.s3.amazonaws.com/litellm_logo.png"
)
assert result.support_contact == "support@berri.ai"
assert result.base_url == "http://test.com/ui?invitation_id=test-id"
assert result.recipient_email == "test@example.com"
@pytest.fixture
def mock_env_vars(monkeypatch):
"""Set up test environment variables"""
@ -513,69 +522,74 @@ def mock_env_vars(monkeypatch):
monkeypatch.setenv("PROXY_BASE_URL", "http://test.com")
monkeypatch.setenv("PROXY_API_URL", "https://test.com")
@pytest.mark.asyncio
async def test_get_email_params_custom_templates_premium_user(mock_env_vars):
"""Test that _get_email_params returns correct values with custom templates for premium users"""
# Mock premium_user as True
with patch("litellm.proxy.proxy_server.premium_user", True):
email_logger = BaseEmailLogger()
# Test invitation email params
invitation_params = await email_logger._get_email_params(
email_event=EmailEvent.new_user_invitation,
user_id="testid",
user_email="test@example.com",
event_message="New User Invitation"
event_message="New User Invitation",
)
assert invitation_params.subject == "Welcome to Test Company!"
assert invitation_params.signature == "Best regards,\nTest Company Team"
assert invitation_params.logo_url == "https://test-company.com/logo.png"
assert invitation_params.support_contact == "support@test-company.com"
assert invitation_params.base_url == "http://test.com"
# Test key created email params
key_params = await email_logger._get_email_params(
email_event=EmailEvent.virtual_key_created,
user_id="testid",
user_email="test@example.com",
event_message="API Key Created"
event_message="API Key Created",
)
assert key_params.subject == "Your Test Company API Key"
assert key_params.signature == "Best regards,\nTest Company Team"
@pytest.mark.asyncio
async def test_get_email_params_non_premium_user(mock_env_vars):
"""Test that non-premium users get default templates even when custom ones are provided"""
# Mock premium_user as False
with patch("litellm.proxy.proxy_server.premium_user", False):
email_logger = BaseEmailLogger()
# Test invitation email params
email_params = await email_logger._get_email_params(
email_event=EmailEvent.new_user_invitation,
user_email="test@example.com",
event_message="New User Invitation"
event_message="New User Invitation",
)
# Should use default values even though custom values are set in env
assert email_params.subject == "LiteLLM: New User Invitation"
assert email_params.signature == EMAIL_FOOTER
assert email_params.logo_url == "https://litellm-listing.s3.amazonaws.com/litellm_logo.png"
assert (
email_params.logo_url
== "https://litellm-listing.s3.amazonaws.com/litellm_logo.png"
)
assert email_params.support_contact == "support@berri.ai"
# Test key created email params
key_params = await email_logger._get_email_params(
email_event=EmailEvent.virtual_key_created,
user_email="test@example.com",
event_message="API Key Created"
event_message="API Key Created",
)
assert key_params.subject == "LiteLLM: API Key Created"
assert key_params.signature == EMAIL_FOOTER
@pytest.mark.asyncio
async def test_get_email_params_default_templates(monkeypatch):
"""Test that _get_email_params uses default templates when custom ones aren't provided"""
@ -583,28 +597,28 @@ async def test_get_email_params_default_templates(monkeypatch):
monkeypatch.delenv("EMAIL_SUBJECT_INVITATION", raising=False)
monkeypatch.delenv("EMAIL_SUBJECT_KEY_CREATED", raising=False)
monkeypatch.delenv("EMAIL_SIGNATURE", raising=False)
# Mock premium_user as True (shouldn't matter since no custom values are set)
with patch("litellm.proxy.proxy_server.premium_user", True):
email_logger = BaseEmailLogger()
# Test invitation email params with default template
invitation_params = await email_logger._get_email_params(
email_event=EmailEvent.new_user_invitation,
user_email="test@example.com",
event_message="New User Invitation"
event_message="New User Invitation",
)
assert invitation_params.subject == "LiteLLM: New User Invitation"
assert invitation_params.signature == EMAIL_FOOTER
# Test key created email params with default template
key_params = await email_logger._get_email_params(
email_event=EmailEvent.virtual_key_created,
user_email="test@example.com",
event_message="API Key Created"
event_message="API Key Created",
)
assert key_params.subject == "LiteLLM: API Key Created"
assert key_params.signature == EMAIL_FOOTER
@ -639,7 +653,10 @@ async def test_send_soft_budget_alert_email(
call_args = mock_send_email.call_args[1]
assert call_args["from_email"] == BaseEmailLogger.DEFAULT_LITELLM_EMAIL
assert call_args["to_email"] == ["test@example.com"]
assert call_args["subject"] == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
assert (
call_args["subject"]
== "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
)
assert "$100.0" in call_args["html_body"] # soft_budget
assert "$105.0" in call_args["html_body"] # spend
assert "$200.0" in call_args["html_body"] # max_budget
@ -673,13 +690,13 @@ async def test_send_soft_budget_alert_email_no_max_budget(
call_args = mock_send_email.call_args[1]
assert "$100.0" in call_args["html_body"] # soft_budget
assert "$105.0" in call_args["html_body"] # spend
assert "Maximum Budget" not in call_args["html_body"] # max_budget should not be shown
assert (
"Maximum Budget" not in call_args["html_body"]
) # max_budget should not be shown
@pytest.mark.asyncio
async def test_budget_alerts_soft_budget_crossed(
base_email_logger, mock_send_email
):
async def test_budget_alerts_soft_budget_crossed(base_email_logger, mock_send_email):
"""Test that budget_alerts sends email when soft budget is crossed"""
user_info = CallInfo(
user_id="test_user",
@ -708,11 +725,14 @@ async def test_budget_alerts_soft_budget_crossed(
mock_send_email.assert_called_once()
call_args = mock_send_email.call_args[1]
assert call_args["to_email"] == ["test@example.com"]
# Verify cache was set to prevent duplicate alerts
mock_cache.async_set_cache.assert_called_once()
cache_call_args = mock_cache.async_set_cache.call_args[1]
assert cache_call_args["key"] == "email_budget_alerts:soft_budget_crossed:test_user"
assert (
cache_call_args["key"]
== "email_budget_alerts:soft_budget_crossed:test_user"
)
assert cache_call_args["value"] == "SENT"
assert cache_call_args["ttl"] == EMAIL_BUDGET_ALERT_TTL
@ -766,9 +786,7 @@ async def test_budget_alerts_soft_budget_duplicate_prevention(
@pytest.mark.asyncio
async def test_budget_alerts_no_budgets(
base_email_logger, mock_send_email
):
async def test_budget_alerts_no_budgets(base_email_logger, mock_send_email):
"""Test that budget_alerts returns early when no budgets are set"""
user_info = CallInfo(
user_id="test_user",
@ -817,7 +835,10 @@ async def test_budget_alerts_uses_token_for_cache_key(
# Verify cache key uses token instead of user_id
mock_cache.async_set_cache.assert_called_once()
cache_call_args = mock_cache.async_set_cache.call_args[1]
assert cache_call_args["key"] == "email_budget_alerts:soft_budget_crossed:hashed_token_123"
assert (
cache_call_args["key"]
== "email_budget_alerts:soft_budget_crossed:hashed_token_123"
)
@pytest.mark.asyncio
@ -838,7 +859,9 @@ async def test_get_email_params_soft_budget_crossed(
)
# Should use default subject template for soft_budget_crossed
assert result.subject == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
assert (
result.subject == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
)
assert result.recipient_email == "test@example.com"
assert result.base_url == "http://test.com"
@ -867,15 +890,19 @@ async def test_budget_alerts_max_budget_alert_crossed(
"PROXY_BASE_URL": "http://test.com",
},
):
await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info)
await base_email_logger.budget_alerts(
type="max_budget_alert", user_info=user_info
)
mock_send_email.assert_called_once()
call_args = mock_send_email.call_args[1]
assert call_args["to_email"] == ["test@example.com"]
assert "Max Budget Alert" in call_args["subject"]
mock_cache.async_set_cache.assert_called_once()
cache_call_args = mock_cache.async_set_cache.call_args[1]
assert cache_call_args["key"] == "email_budget_alerts:max_budget_alert:test_user"
assert (
cache_call_args["key"] == "email_budget_alerts:max_budget_alert:test_user"
)
assert cache_call_args["value"] == "SENT"
assert cache_call_args["ttl"] == EMAIL_BUDGET_ALERT_TTL
assert cache_call_args["ttl"] == EMAIL_BUDGET_ALERT_TTL

View file

@ -87,7 +87,7 @@ async def test_send_email_success(mock_env_vars):
async def test_send_email_missing_api_key():
# Remove the API key from environment before initializing logger
original_key = os.environ.pop("RESEND_API_KEY", None)
try:
# Initialize the logger after removing the API key
logger = ResendEmailLogger()
@ -104,16 +104,19 @@ async def test_send_email_missing_api_key():
mock_response.raise_for_status.return_value = None
mock_response.status_code = 200
mock_response.json.return_value = {"id": "test_email_id"}
mock_async_client = mock.AsyncMock()
mock_async_client.post.return_value = mock_response
# Directly inject the mock client to bypass any caching
logger.async_httpx_client = mock_async_client
# Send email
await logger.send_email(
from_email=from_email, to_email=to_email, subject=subject, html_body=html_body
from_email=from_email,
to_email=to_email,
subject=subject,
html_body=html_body,
)
# Verify the HTTP client was called with None as the API key

View file

@ -32,12 +32,12 @@ def mock_env_vars():
# Store original values
original_api_key = os.environ.get("SENDGRID_API_KEY")
original_sender_email = os.environ.get("SENDGRID_SENDER_EMAIL")
# Set test API key and remove SENDGRID_SENDER_EMAIL to ensure isolation
os.environ["SENDGRID_API_KEY"] = "test_api_key"
if "SENDGRID_SENDER_EMAIL" in os.environ:
del os.environ["SENDGRID_SENDER_EMAIL"]
try:
yield
finally:
@ -46,7 +46,7 @@ def mock_env_vars():
os.environ["SENDGRID_API_KEY"] = original_api_key
elif "SENDGRID_API_KEY" in os.environ:
del os.environ["SENDGRID_API_KEY"]
if original_sender_email is not None:
os.environ["SENDGRID_SENDER_EMAIL"] = original_sender_email

View file

@ -18,168 +18,282 @@ from litellm.types.utils import StandardCallbackDynamicParams
class TestEnterpriseCallbackControls:
@pytest.fixture
def mock_premium_user(self):
"""Fixture to mock premium user check as True"""
with patch.object(EnterpriseCallbackControls, '_should_allow_dynamic_callback_disabling', return_value=True):
with patch.object(
EnterpriseCallbackControls,
"_should_allow_dynamic_callback_disabling",
return_value=True,
):
yield
@pytest.fixture
@pytest.fixture
def mock_non_premium_user(self):
"""Fixture to mock premium user check as False"""
with patch.object(EnterpriseCallbackControls, '_should_allow_dynamic_callback_disabling', return_value=False):
with patch.object(
EnterpriseCallbackControls,
"_should_allow_dynamic_callback_disabling",
return_value=False,
):
yield
@pytest.fixture
def mock_request_headers(self):
"""Fixture to mock get_proxy_server_request_headers"""
with patch('enterprise.litellm_enterprise.enterprise_callbacks.callback_controls.get_proxy_server_request_headers') as mock_headers:
with patch(
"enterprise.litellm_enterprise.enterprise_callbacks.callback_controls.get_proxy_server_request_headers"
) as mock_headers:
yield mock_headers
def test_callback_disabled_langfuse_string(self, mock_premium_user, mock_request_headers):
def test_callback_disabled_langfuse_string(
self, mock_premium_user, mock_request_headers
):
"""Test that 'langfuse' string callback is disabled when specified in headers"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
assert result is True
def test_callback_disabled_langfuse_customlogger(self, mock_premium_user, mock_request_headers):
def test_callback_disabled_langfuse_customlogger(
self, mock_premium_user, mock_request_headers
):
"""Test that LangfusePromptManagement CustomLogger instance is disabled when 'langfuse' specified in headers"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
langfuse_logger = LangfusePromptManagement()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(langfuse_logger, litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
langfuse_logger, litellm_params, standard_callback_dynamic_params
)
assert result is True
def test_callback_disabled_s3_v2_string(self, mock_premium_user, mock_request_headers):
def test_callback_disabled_s3_v2_string(
self, mock_premium_user, mock_request_headers
):
"""Test that 's3_v2' string callback is disabled when specified in headers"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "s3_v2"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("s3_v2", litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"s3_v2", litellm_params, standard_callback_dynamic_params
)
assert result is True
def test_callback_disabled_s3_v2_customlogger(self, mock_premium_user, mock_request_headers):
def test_callback_disabled_s3_v2_customlogger(
self, mock_premium_user, mock_request_headers
):
"""Test that S3Logger CustomLogger instance is disabled when 's3_v2' specified in headers"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "s3_v2"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
# Mock S3Logger to avoid async initialization issues
with patch('litellm.integrations.s3_v2.S3Logger.__init__', return_value=None):
with patch("litellm.integrations.s3_v2.S3Logger.__init__", return_value=None):
s3_logger = S3Logger()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(s3_logger, litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
s3_logger, litellm_params, standard_callback_dynamic_params
)
assert result is True
def test_callback_disabled_datadog_string(self, mock_premium_user, mock_request_headers):
def test_callback_disabled_datadog_string(
self, mock_premium_user, mock_request_headers
):
"""Test that 'datadog' string callback is disabled when specified in headers"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "datadog"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"datadog", litellm_params, standard_callback_dynamic_params
)
assert result is True
def test_callback_disabled_datadog_customlogger(self, mock_premium_user, mock_request_headers):
def test_callback_disabled_datadog_customlogger(
self, mock_premium_user, mock_request_headers
):
"""Test that DataDogLogger CustomLogger instance is disabled when 'datadog' specified in headers"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "datadog"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
# Mock DataDogLogger to avoid async initialization issues
with patch('litellm.integrations.datadog.datadog.DataDogLogger.__init__', return_value=None):
with patch(
"litellm.integrations.datadog.datadog.DataDogLogger.__init__",
return_value=None,
):
datadog_logger = DataDogLogger()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(datadog_logger, litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
datadog_logger, litellm_params, standard_callback_dynamic_params
)
assert result is True
def test_multiple_callbacks_disabled(self, mock_premium_user, mock_request_headers):
"""Test that multiple callbacks can be disabled with comma-separated list"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse,datadog,s3_v2"}
mock_request_headers.return_value = {
X_LITELLM_DISABLE_CALLBACKS: "langfuse,datadog,s3_v2"
}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
# Test each callback is disabled
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) is True
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) is True
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("s3_v2", litellm_params, standard_callback_dynamic_params) is True
# Test non-disabled callback is not disabled
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("prometheus", litellm_params, standard_callback_dynamic_params) is False
def test_callback_not_disabled_when_not_in_list(self, mock_premium_user, mock_request_headers):
# Test each callback is disabled
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
is True
)
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"datadog", litellm_params, standard_callback_dynamic_params
)
is True
)
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"s3_v2", litellm_params, standard_callback_dynamic_params
)
is True
)
# Test non-disabled callback is not disabled
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"prometheus", litellm_params, standard_callback_dynamic_params
)
is False
)
def test_callback_not_disabled_when_not_in_list(
self, mock_premium_user, mock_request_headers
):
"""Test that callbacks not in the disabled list are not disabled"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"datadog", litellm_params, standard_callback_dynamic_params
)
assert result is False
def test_callback_not_disabled_when_no_header(self, mock_premium_user, mock_request_headers):
def test_callback_not_disabled_when_no_header(
self, mock_premium_user, mock_request_headers
):
"""Test that callbacks are not disabled when the header is not present"""
mock_request_headers.return_value = {}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
assert result is False
def test_callback_not_disabled_when_header_none(self, mock_premium_user, mock_request_headers):
def test_callback_not_disabled_when_header_none(
self, mock_premium_user, mock_request_headers
):
"""Test that callbacks are not disabled when the header value is None"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: None}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
assert result is False
def test_non_premium_user_cannot_disable_callbacks(self, mock_non_premium_user, mock_request_headers):
def test_non_premium_user_cannot_disable_callbacks(
self, mock_non_premium_user, mock_request_headers
):
"""Test that non-premium users cannot disable callbacks even with the header"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
assert result is False
def test_case_insensitive_callback_matching(self, mock_premium_user, mock_request_headers):
def test_case_insensitive_callback_matching(
self, mock_premium_user, mock_request_headers
):
"""Test that callback matching is case insensitive"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "LANGFUSE,DataDog"}
mock_request_headers.return_value = {
X_LITELLM_DISABLE_CALLBACKS: "LANGFUSE,DataDog"
}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
# Test lowercase callbacks are disabled
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) is True
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) is True
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
is True
)
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"datadog", litellm_params, standard_callback_dynamic_params
)
is True
)
def test_whitespace_handling_in_disabled_callbacks(self, mock_premium_user, mock_request_headers):
def test_whitespace_handling_in_disabled_callbacks(
self, mock_premium_user, mock_request_headers
):
"""Test that whitespace around callback names is handled correctly"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: " langfuse , datadog , s3_v2 "}
mock_request_headers.return_value = {
X_LITELLM_DISABLE_CALLBACKS: " langfuse , datadog , s3_v2 "
}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) is True
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) is True
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("s3_v2", litellm_params, standard_callback_dynamic_params) is True
def test_custom_logger_not_in_registry(self, mock_premium_user, mock_request_headers):
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
is True
)
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"datadog", litellm_params, standard_callback_dynamic_params
)
is True
)
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"s3_v2", litellm_params, standard_callback_dynamic_params
)
is True
)
def test_custom_logger_not_in_registry(
self, mock_premium_user, mock_request_headers
):
"""Test that CustomLogger not in registry is not disabled"""
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "unknown_logger"}
mock_request_headers.return_value = {
X_LITELLM_DISABLE_CALLBACKS: "unknown_logger"
}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
# Create a mock CustomLogger that's not in the registry
class UnknownLogger(CustomLogger):
pass
unknown_logger = UnknownLogger()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(unknown_logger, litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
unknown_logger, litellm_params, standard_callback_dynamic_params
)
assert result is False
def test_exception_handling(self, mock_premium_user, mock_request_headers):
@ -188,32 +302,64 @@ class TestEnterpriseCallbackControls:
mock_request_headers.side_effect = Exception("Test exception")
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
assert result is False
def test_callback_disabled_via_request_body_langfuse(self, mock_premium_user, mock_request_headers):
def test_callback_disabled_via_request_body_langfuse(
self, mock_premium_user, mock_request_headers
):
"""Test that callbacks can be disabled via request body litellm_disabled_callbacks"""
mock_request_headers.return_value = {} # No headers
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams(litellm_disabled_callbacks=["langfuse"])
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params)
standard_callback_dynamic_params = StandardCallbackDynamicParams(
litellm_disabled_callbacks=["langfuse"]
)
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
assert result is True
def test_callback_disabled_via_request_body_multiple(self, mock_premium_user, mock_request_headers):
def test_callback_disabled_via_request_body_multiple(
self, mock_premium_user, mock_request_headers
):
"""Test that multiple callbacks can be disabled via request body"""
mock_request_headers.return_value = {} # No headers
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams(litellm_disabled_callbacks=["langfuse", "datadog", "s3_v2"])
standard_callback_dynamic_params = StandardCallbackDynamicParams(
litellm_disabled_callbacks=["langfuse", "datadog", "s3_v2"]
)
# Test each callback is disabled
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) is True
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) is True
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("s3_v2", litellm_params, standard_callback_dynamic_params) is True
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
is True
)
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"datadog", litellm_params, standard_callback_dynamic_params
)
is True
)
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"s3_v2", litellm_params, standard_callback_dynamic_params
)
is True
)
# Test non-disabled callback is not disabled
assert EnterpriseCallbackControls.is_callback_disabled_dynamically("prometheus", litellm_params, standard_callback_dynamic_params) is False
assert (
EnterpriseCallbackControls.is_callback_disabled_dynamically(
"prometheus", litellm_params, standard_callback_dynamic_params
)
is False
)
def test_admin_can_disable_dynamic_callback_disabling(self, mock_request_headers):
"""
@ -223,11 +369,13 @@ class TestEnterpriseCallbackControls:
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
# Mock litellm.allow_dynamic_callback_disabling set to False
with patch('litellm.allow_dynamic_callback_disabling', False):
with patch('litellm.proxy.proxy_server.premium_user', True):
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params)
with patch("litellm.allow_dynamic_callback_disabling", False):
with patch("litellm.proxy.proxy_server.premium_user", True):
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
assert result is False
def test_admin_can_enable_dynamic_callback_disabling(self, mock_request_headers):
@ -238,14 +386,18 @@ class TestEnterpriseCallbackControls:
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
# Mock litellm.allow_dynamic_callback_disabling set to True
with patch('litellm.allow_dynamic_callback_disabling', True):
with patch('litellm.proxy.proxy_server.premium_user', True):
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params)
with patch("litellm.allow_dynamic_callback_disabling", True):
with patch("litellm.proxy.proxy_server.premium_user", True):
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
assert result is True
def test_default_admin_setting_allows_dynamic_callback_disabling(self, mock_request_headers):
def test_default_admin_setting_allows_dynamic_callback_disabling(
self, mock_request_headers
):
"""
Test that when allow_dynamic_callback_disabling is not set,
it defaults to True and allows dynamic callback disabling for premium users
@ -253,8 +405,10 @@ class TestEnterpriseCallbackControls:
mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"}
litellm_params = {"proxy_server_request": {"url": "test"}}
standard_callback_dynamic_params = StandardCallbackDynamicParams()
# litellm.allow_dynamic_callback_disabling should default to True
with patch('litellm.proxy.proxy_server.premium_user', True):
result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params)
with patch("litellm.proxy.proxy_server.premium_user", True):
result = EnterpriseCallbackControls.is_callback_disabled_dynamically(
"langfuse", litellm_params, standard_callback_dynamic_params
)
assert result is True

View file

@ -18,11 +18,17 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
DECODED_UNIFIED_INPUT_FILE_ID = "litellm_proxy:application/octet-stream;unified_id,test-uuid;target_model_names,azure-gpt-4"
B64_UNIFIED_INPUT_FILE_ID = base64.urlsafe_b64encode(DECODED_UNIFIED_INPUT_FILE_ID.encode()).decode().rstrip("=")
B64_UNIFIED_INPUT_FILE_ID = (
base64.urlsafe_b64encode(DECODED_UNIFIED_INPUT_FILE_ID.encode())
.decode()
.rstrip("=")
)
RAW_INPUT_FILE_ID = "file-raw-provider-abc123"
DECODED_UNIFIED_BATCH_ID = "litellm_proxy;model_id:model-xyz;llm_batch_id:batch-123"
B64_UNIFIED_BATCH_ID = base64.urlsafe_b64encode(DECODED_UNIFIED_BATCH_ID.encode()).decode().rstrip("=")
B64_UNIFIED_BATCH_ID = (
base64.urlsafe_b64encode(DECODED_UNIFIED_BATCH_ID.encode()).decode().rstrip("=")
)
@pytest.mark.asyncio
@ -55,10 +61,16 @@ async def test_should_resolve_raw_input_file_id_to_unified():
mock_managed_file.unified_file_id = B64_UNIFIED_INPUT_FILE_ID
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=mock_db_object)
mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=mock_managed_file)
mock_prisma.db.litellm_managedobjecttable.find_first = AsyncMock(
return_value=mock_db_object
)
mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(
return_value=mock_managed_file
)
from litellm.proxy.openai_files_endpoints.common_utils import get_batch_from_database
from litellm.proxy.openai_files_endpoints.common_utils import (
get_batch_from_database,
)
_, response = await get_batch_from_database(
batch_id=B64_UNIFIED_BATCH_ID,

View file

@ -93,7 +93,9 @@ async def test_should_preserve_already_managed_input_file_id():
unified_batch_id = "bGl0ZWxsbV9wcm94eTpiYXRjaF9pZA"
decoded_unified = "litellm_proxy:application/octet-stream;unified_id,test-123"
base64_input_file_id = base64.urlsafe_b64encode(decoded_unified.encode()).decode().rstrip("=")
base64_input_file_id = (
base64.urlsafe_b64encode(decoded_unified.encode()).decode().rstrip("=")
)
batch_data = {
"id": "batch-raw-123",

View file

@ -14,63 +14,77 @@ import pytest
def test_enterprise_routes_all_imports_exist():
"""
Validate that all relative imports in enterprise_routes.py exist in the filesystem.
This catches any import errors from moved/deleted modules without hardcoding
specific module names. Works by checking that imported files actually exist.
"""
# Path to the enterprise_routes.py source file
enterprise_routes_path = os.path.join(
os.path.dirname(__file__),
"..", "..", "..", "..",
"enterprise", "litellm_enterprise", "proxy", "enterprise_routes.py"
"..",
"..",
"..",
"..",
"enterprise",
"litellm_enterprise",
"proxy",
"enterprise_routes.py",
)
enterprise_routes_path = os.path.normpath(enterprise_routes_path)
enterprise_proxy_dir = os.path.dirname(enterprise_routes_path)
if not os.path.exists(enterprise_routes_path):
pytest.skip(f"Enterprise routes file not found at {enterprise_routes_path}")
# Read and parse the source file
with open(enterprise_routes_path, "r") as f:
source_code = f.read()
try:
tree = ast.parse(source_code)
except SyntaxError as e:
pytest.fail(f"Syntax error in enterprise_routes.py: {e}")
# Check all relative imports
missing_imports = []
for node in ast.walk(tree):
if isinstance(node, ast.ImportFrom):
# level > 0 means it's a relative import (. or .. etc)
if node.level and node.level > 0:
module = node.module or ""
# Convert relative import to file path
# e.g., "audit_logging_endpoints" -> "audit_logging_endpoints.py"
# e.g., "vector_stores.endpoints" -> "vector_stores/endpoints.py"
module_path = module.replace(".", os.sep) if module else ""
# Check both .py file and package directory
file_path = os.path.join(enterprise_proxy_dir, module_path + ".py") if module_path else None
package_path = os.path.join(enterprise_proxy_dir, module_path, "__init__.py") if module_path else None
file_path = (
os.path.join(enterprise_proxy_dir, module_path + ".py")
if module_path
else None
)
package_path = (
os.path.join(enterprise_proxy_dir, module_path, "__init__.py")
if module_path
else None
)
# If module is empty (e.g., "from . import something"), skip check
if not module:
continue
file_exists = file_path and os.path.exists(file_path)
package_exists = package_path and os.path.exists(package_path)
if not file_exists and not package_exists:
missing_imports.append(
f"Line {node.lineno}: Cannot find '.{module}' "
f"(checked: {file_path} and {package_path})"
)
if missing_imports:
error_msg = "Found imports in enterprise_routes.py that don't exist:\n"
error_msg += "\n".join(missing_imports)

View file

@ -61,7 +61,7 @@ def _make_managed_files_instance_with_batches(
):
"""
Create a _PROXY_LiteLLMManagedFiles instance with mocked DB and batches.
Args:
file_id: The unified file ID
batches: List of batch records to return from DB
@ -79,7 +79,7 @@ def _make_managed_files_instance_with_batches(
# Mock prisma
mock_prisma = MagicMock()
# Mock file table queries
mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(
return_value=mock_file_record
@ -87,7 +87,7 @@ def _make_managed_files_instance_with_batches(
mock_prisma.db.litellm_managedfiletable.delete = AsyncMock(
return_value=mock_file_record
)
# Mock batch/object table queries
mock_prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=batches
@ -95,11 +95,13 @@ def _make_managed_files_instance_with_batches(
# Mock cache
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value={
"unified_file_id": file_id,
"model_mappings": {"model-123": "provider-file-abc"},
"flat_model_file_ids": ["provider-file-abc"],
})
mock_cache.async_get_cache = AsyncMock(
return_value={
"unified_file_id": file_id,
"model_mappings": {"model-123": "provider-file-abc"},
"flat_model_file_ids": ["provider-file-abc"],
}
)
mock_cache.async_set_cache = AsyncMock()
instance = _PROXY_LiteLLMManagedFiles(
@ -117,17 +119,17 @@ def test_is_batch_polling_enabled_when_job_registered():
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
instance = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=MagicMock(),
)
# Mock scheduler with registered job
mock_scheduler = MagicMock()
mock_job = MagicMock()
mock_scheduler.get_job.return_value = mock_job
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
assert instance._is_batch_polling_enabled() is True
@ -137,16 +139,16 @@ def test_is_batch_polling_disabled_when_job_not_registered():
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
instance = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=MagicMock(),
)
# Mock scheduler without registered job
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = None
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
assert instance._is_batch_polling_enabled() is False
@ -156,12 +158,12 @@ def test_is_batch_polling_disabled_when_no_scheduler():
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
instance = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=MagicMock(),
)
with patch("litellm.proxy.proxy_server.scheduler", None):
assert instance._is_batch_polling_enabled() is False
@ -174,26 +176,28 @@ async def test_get_batches_referencing_file_finds_batch_with_input_file():
"""Test finding a batch that references the file as input_file_id."""
unified_file_id = _make_unified_file_id("file-input-123")
unified_batch_id = _make_unified_batch_id("batch-123")
batch_file_object = {
"id": "batch-123",
"input_file_id": unified_file_id, # Batch references this file
"status": "validating",
}
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="validating",
file_object=batch_file_object,
)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
referencing_batches = await managed_files._get_batches_referencing_file(
unified_file_id
)
assert len(referencing_batches) == 1
assert referencing_batches[0]["batch_id"] == unified_batch_id
assert referencing_batches[0]["status"] == "validating"
@ -204,27 +208,29 @@ async def test_get_batches_referencing_file_finds_batch_with_output_file():
"""Test finding a batch that references the file as output_file_id."""
unified_file_id = _make_unified_file_id("file-output-456")
unified_batch_id = _make_unified_batch_id("batch-456")
batch_file_object = {
"id": "batch-456",
"input_file_id": "file-input-different",
"output_file_id": unified_file_id, # Batch references this file
"status": "in_progress",
}
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="in_progress",
file_object=batch_file_object,
)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
referencing_batches = await managed_files._get_batches_referencing_file(
unified_file_id
)
assert len(referencing_batches) == 1
assert referencing_batches[0]["status"] == "in_progress"
@ -234,27 +240,29 @@ async def test_get_batches_referencing_file_ignores_terminal_batches():
"""Test that batches in terminal states are not returned."""
unified_file_id = _make_unified_file_id("file-123")
unified_batch_id = _make_unified_batch_id("batch-completed")
batch_file_object = {
"id": "batch-completed",
"input_file_id": unified_file_id,
"status": "completed",
}
# Batch is in terminal state in DB
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="completed", # Terminal state
file_object=batch_file_object,
)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[], # Query returns no batches (terminal states filtered out)
)
referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
referencing_batches = await managed_files._get_batches_referencing_file(
unified_file_id
)
assert len(referencing_batches) == 0
@ -262,26 +270,36 @@ async def test_get_batches_referencing_file_ignores_terminal_batches():
async def test_get_batches_referencing_file_finds_multiple_batches():
"""Test finding multiple batches referencing the same file."""
unified_file_id = _make_unified_file_id("file-shared")
batch1 = _make_batch_db_record(
unified_object_id=_make_unified_batch_id("batch-1"),
status="validating",
file_object={"id": "batch-1", "input_file_id": unified_file_id, "status": "validating"},
file_object={
"id": "batch-1",
"input_file_id": unified_file_id,
"status": "validating",
},
)
batch2 = _make_batch_db_record(
unified_object_id=_make_unified_batch_id("batch-2"),
status="in_progress",
file_object={"id": "batch-2", "input_file_id": unified_file_id, "status": "in_progress"},
file_object={
"id": "batch-2",
"input_file_id": unified_file_id,
"status": "in_progress",
},
)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch1, batch2],
)
referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
referencing_batches = await managed_files._get_batches_referencing_file(
unified_file_id
)
assert len(referencing_batches) == 2
statuses = [b["status"] for b in referencing_batches]
assert "validating" in statuses
@ -300,32 +318,32 @@ async def test_file_deletion_blocked_when_batch_polling_enabled_and_batch_refere
"""
unified_file_id = _make_unified_file_id("file-to-delete")
unified_batch_id = _make_unified_batch_id("batch-active")
batch_file_object = {
"id": "batch-active",
"input_file_id": unified_file_id,
"status": "validating",
}
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="validating",
file_object=batch_file_object,
)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
# Mock scheduler with registered batch cost job
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock() # Job exists
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
with pytest.raises(HTTPException) as exc_info:
await managed_files._check_file_deletion_allowed(unified_file_id)
assert exc_info.value.status_code == 400
error_detail = exc_info.value.detail
assert "Cannot delete file" in error_detail
@ -342,28 +360,28 @@ async def test_file_deletion_allowed_when_batch_polling_disabled():
"""
unified_file_id = _make_unified_file_id("file-to-delete")
unified_batch_id = _make_unified_batch_id("batch-active")
batch_file_object = {
"id": "batch-active",
"input_file_id": unified_file_id,
"status": "validating",
}
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="validating",
file_object=batch_file_object,
)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
# Mock scheduler without registered job (batch cost tracking disabled)
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = None
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
# Should not raise an exception
await managed_files._check_file_deletion_allowed(unified_file_id)
@ -376,16 +394,16 @@ async def test_file_deletion_allowed_when_no_batches_reference_file():
even when batch cost tracking is enabled.
"""
unified_file_id = _make_unified_file_id("file-to-delete")
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[], # No batches reference this file
)
# Mock scheduler with registered job (batch cost tracking enabled)
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock()
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
# Should not raise an exception
await managed_files._check_file_deletion_allowed(unified_file_id)
@ -398,32 +416,32 @@ async def test_afile_delete_calls_check_deletion_allowed():
"""
unified_file_id = _make_unified_file_id("file-to-delete")
unified_batch_id = _make_unified_batch_id("batch-active")
batch_file_object = {
"id": "batch-active",
"input_file_id": unified_file_id,
"status": "in_progress",
}
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="in_progress",
file_object=batch_file_object,
)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
# Mock llm_router
mock_router = MagicMock()
mock_router.afile_delete = AsyncMock()
# Mock scheduler with registered job
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock()
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
with pytest.raises(HTTPException) as exc_info:
await managed_files.afile_delete(
@ -431,7 +449,7 @@ async def test_afile_delete_calls_check_deletion_allowed():
litellm_parent_otel_span=None,
llm_router=mock_router,
)
# Should raise error before calling router delete
assert exc_info.value.status_code == 400
mock_router.afile_delete.assert_not_called()
@ -444,7 +462,7 @@ async def test_database_limit_respected():
This is a performance optimization - we only fetch what we need.
"""
unified_file_id = _make_unified_file_id("file-shared")
# Create exactly 10 batches (what DB will return with take=10)
ten_batches = []
for i in range(10):
@ -454,30 +472,32 @@ async def test_database_limit_respected():
file_object={
"id": f"batch-{i}",
"input_file_id": unified_file_id,
"status": "validating"
"status": "validating",
},
)
ten_batches.append(batch)
# Mock will return only 10 batches (as DB would with take=10)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=ten_batches,
)
referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
referencing_batches = await managed_files._get_batches_referencing_file(
unified_file_id
)
# Should return all 10 that reference the file
assert len(referencing_batches) == 10
# Verify error message handles "10+" case (since we got exactly 10, might be more in DB)
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock()
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
with pytest.raises(HTTPException) as exc_info:
await managed_files._check_file_deletion_allowed(unified_file_id)
error_detail = exc_info.value.detail
# When we get exactly 10 matches, show "10+" to indicate there might be more
assert "10+ batch(es)" in error_detail
@ -491,32 +511,40 @@ async def test_error_message_includes_batch_details():
unified_file_id = _make_unified_file_id("file-to-delete")
batch1_id = _make_unified_batch_id("batch-1")
batch2_id = _make_unified_batch_id("batch-2")
batch1 = _make_batch_db_record(
unified_object_id=batch1_id,
status="validating",
file_object={"id": "batch-1", "input_file_id": unified_file_id, "status": "validating"},
file_object={
"id": "batch-1",
"input_file_id": unified_file_id,
"status": "validating",
},
)
batch2 = _make_batch_db_record(
unified_object_id=batch2_id,
status="in_progress",
file_object={"id": "batch-2", "output_file_id": unified_file_id, "status": "in_progress"},
file_object={
"id": "batch-2",
"output_file_id": unified_file_id,
"status": "in_progress",
},
)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch1, batch2],
)
# Mock scheduler with registered job
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock()
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
with pytest.raises(HTTPException) as exc_info:
await managed_files._check_file_deletion_allowed(unified_file_id)
error_detail = exc_info.value.detail
assert "2 batch(es)" in error_detail
assert "validating" in error_detail

View file

@ -144,6 +144,7 @@ async def test_check_batch_cost_should_call_afile_content_directly_with_credenti
# Mock the batch response (completed, with output file)
from litellm.types.utils import LiteLLMBatch
batch_response = LiteLLMBatch(
id="batch-123",
completion_window="24h",
@ -201,9 +202,11 @@ async def test_check_batch_cost_should_call_afile_content_directly_with_credenti
# Verify the DB update writes batch_processed, status, and file_object
mock_prisma.db.litellm_managedobjecttable.update.assert_called_once()
update_call_kwargs = mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs
update_call_kwargs = (
mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs
)
assert update_call_kwargs["data"]["batch_processed"] is True
assert update_call_kwargs["data"]["status"] == "complete"
assert "file_object" in update_call_kwargs["data"], (
"file_object must be written to DB so list_batches reads updated status"
)
assert (
"file_object" in update_call_kwargs["data"]
), "file_object must be written to DB so list_batches reads updated status"

View file

@ -110,10 +110,9 @@ async def test_should_pass_credentials_to_afile_retrieve():
mock_afile_retrieve = AsyncMock(return_value=_make_file_object("file-output-abc"))
with patch(
"litellm.afile_retrieve", mock_afile_retrieve
), patch(
"litellm.proxy.proxy_server.llm_router", mock_router
with (
patch("litellm.afile_retrieve", mock_afile_retrieve),
patch("litellm.proxy.proxy_server.llm_router", mock_router),
):
await managed_files.async_post_call_success_hook(
data={},
@ -128,7 +127,9 @@ async def test_should_pass_credentials_to_afile_retrieve():
f"afile_retrieve must receive api_key from router credentials. "
f"Got kwargs: {call_kwargs.kwargs}"
)
assert call_kwargs.kwargs.get("api_base") == "https://my-azure.openai.azure.com/", (
assert (
call_kwargs.kwargs.get("api_base") == "https://my-azure.openai.azure.com/"
), (
f"afile_retrieve must receive api_base from router credentials. "
f"Got kwargs: {call_kwargs.kwargs}"
)
@ -150,10 +151,9 @@ async def test_should_fallback_when_no_router():
mock_afile_retrieve = AsyncMock(return_value=_make_file_object("file-output-abc"))
with patch(
"litellm.afile_retrieve", mock_afile_retrieve
), patch(
"litellm.proxy.proxy_server.llm_router", None
with (
patch("litellm.afile_retrieve", mock_afile_retrieve),
patch("litellm.proxy.proxy_server.llm_router", None),
):
await managed_files.async_post_call_success_hook(
data={},

View file

@ -255,12 +255,19 @@ class TestGitHubCopilotAuthenticator:
"user_code": "UC",
"verification_uri": "https://example.com",
}
with patch.dict(os.environ, {"GITHUB_COPILOT_DEVICE_CODE_URL": custom_url}), \
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client):
with (
patch.dict(os.environ, {"GITHUB_COPILOT_DEVICE_CODE_URL": custom_url}),
patch(
"litellm.llms.github_copilot.authenticator._get_httpx_client",
return_value=mock_client,
),
):
authenticator._get_device_code()
assert mock_client.post.call_args[0][0] == custom_url
def test_get_device_code_with_custom_client_id(self, authenticator, mock_http_client):
def test_get_device_code_with_custom_client_id(
self, authenticator, mock_http_client
):
"""GITHUB_COPILOT_CLIENT_ID env var must appear as client_id in the device-code request body."""
mock_client, mock_response = mock_http_client
custom_id = "custom_client_id"
@ -269,30 +276,49 @@ class TestGitHubCopilotAuthenticator:
"user_code": "UC",
"verification_uri": "https://example.com",
}
with patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), \
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client):
with (
patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}),
patch(
"litellm.llms.github_copilot.authenticator._get_httpx_client",
return_value=mock_client,
),
):
authenticator._get_device_code()
assert mock_client.post.call_args[1]["json"]["client_id"] == custom_id
def test_poll_for_access_token_with_custom_url(self, authenticator, mock_http_client):
def test_poll_for_access_token_with_custom_url(
self, authenticator, mock_http_client
):
"""GITHUB_COPILOT_ACCESS_TOKEN_URL env var must be used by _poll_for_access_token at call time."""
mock_client, mock_response = mock_http_client
custom_url = "https://custom.example.com/token"
mock_response.json.return_value = {"access_token": "tok"}
with patch.dict(os.environ, {"GITHUB_COPILOT_ACCESS_TOKEN_URL": custom_url}), \
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \
patch("time.sleep"):
with (
patch.dict(os.environ, {"GITHUB_COPILOT_ACCESS_TOKEN_URL": custom_url}),
patch(
"litellm.llms.github_copilot.authenticator._get_httpx_client",
return_value=mock_client,
),
patch("time.sleep"),
):
authenticator._poll_for_access_token("dc")
assert mock_client.post.call_args[0][0] == custom_url
def test_poll_for_access_token_with_custom_client_id(self, authenticator, mock_http_client):
def test_poll_for_access_token_with_custom_client_id(
self, authenticator, mock_http_client
):
"""GITHUB_COPILOT_CLIENT_ID env var must appear as client_id in the polling request body."""
mock_client, mock_response = mock_http_client
custom_id = "custom_client_id"
mock_response.json.return_value = {"access_token": "tok"}
with patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), \
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \
patch("time.sleep"):
with (
patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}),
patch(
"litellm.llms.github_copilot.authenticator._get_httpx_client",
return_value=mock_client,
),
patch("time.sleep"),
):
authenticator._poll_for_access_token("dc")
assert mock_client.post.call_args[1]["json"]["client_id"] == custom_id
@ -301,9 +327,13 @@ class TestGitHubCopilotAuthenticator:
mock_client, mock_response = mock_http_client
custom_url = "https://custom.example.com/api-key"
mock_response.json.return_value = {"token": "api-tok", "expires_at": 9999999999}
with patch.dict(os.environ, {"GITHUB_COPILOT_API_KEY_URL": custom_url}), \
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \
patch.object(authenticator, "get_access_token", return_value="access-tok"):
with (
patch.dict(os.environ, {"GITHUB_COPILOT_API_KEY_URL": custom_url}),
patch(
"litellm.llms.github_copilot.authenticator._get_httpx_client",
return_value=mock_client,
),
patch.object(authenticator, "get_access_token", return_value="access-tok"),
):
authenticator._refresh_api_key()
assert mock_client.get.call_args[0][0] == custom_url

View file

@ -166,9 +166,9 @@ def test_anthropic_provider_fields_support_byok():
"Anthropic api_key must be optional so admins can configure BYOK models "
"without entering a key. See BYOK tutorial."
)
assert fields_by_key["api_key"].get("tooltip"), (
"Anthropic api_key must have a tooltip explaining the BYOK use case."
)
assert fields_by_key["api_key"].get(
"tooltip"
), "Anthropic api_key must have a tooltip explaining the BYOK use case."
assert "api_base" in fields_by_key, (
"Anthropic provider form must expose api_base so cloud customers "
"can override the upstream URL without env var access."
@ -176,16 +176,16 @@ def test_anthropic_provider_fields_support_byok():
api_base_field = fields_by_key["api_base"]
assert api_base_field["required"] is False
assert api_base_field["field_type"] == "text"
assert api_base_field.get("tooltip"), (
"api_base should have a tooltip explaining it is optional."
)
assert api_base_field.get(
"tooltip"
), "api_base should have a tooltip explaining it is optional."
# UI forms render fields in credential_fields order; api_base should come first
# so an admin sees the URL override before the key field.
field_order = [f["key"] for f in anthropic["credential_fields"]]
assert field_order.index("api_base") < field_order.index("api_key"), (
"api_base must appear before api_key in credential_fields (matches AI21 and ANTHROPIC_TEXT convention)."
)
assert field_order.index("api_base") < field_order.index(
"api_key"
), "api_base must appear before api_key in credential_fields (matches AI21 and ANTHROPIC_TEXT convention)."
def test_public_model_hub_with_healthy_model():

View file

@ -22,7 +22,7 @@ from litellm.types.integrations.prometheus import (
# Escapes per Prometheus text format
('he said "hi"', 'he said \\"hi\\"'),
(r"path\to\file", r"path\\to\\file"),
(r'quote\"slash\\', r'quote\\\"slash\\\\'),
(r"quote\"slash\\", r"quote\\\"slash\\\\"),
# Non-string inputs get coerced to str first
(123, "123"),
(True, "True"),
@ -31,4 +31,3 @@ from litellm.types.integrations.prometheus import (
)
def test_sanitize_prometheus_label_value_expected_outputs(value, expected):
assert _sanitize_prometheus_label_value(value) == expected