fix(mistral): strip metadata field from messages to prevent extra_forbidden error

This commit is contained in:
Rish Dass 2026-06-20 17:19:54 -05:00
parent 15aa40b36e
commit 1c1638ee8c
310 changed files with 3415 additions and 2081 deletions

View file

@ -53,8 +53,6 @@ from agent_shin_shared import ( # noqa: E402 -- sys.path adjusted above
ALLOWLIST_LOGINS,
GRACE_COMMENT_MARKER,
GRACE_PERIOD_SECONDS,
GREPTILE_BOT_LOGINS,
SCORE_PATTERN,
extract_greptile_score,
gh,
list_open_items,

View file

@ -53,7 +53,6 @@ from agent_shin_shared import ( # noqa: E402
)
from triage_with_llm import ( # noqa: E402
DEFAULT_MODEL,
call_llm_judge,
fetch_issue,
fetch_pr,
gh,

View file

@ -53,8 +53,6 @@ from agent_shin_shared import ( # noqa: E402 -- sys.path adjusted above
ALLOWLIST_LOGINS,
GRACE_COMMENT_MARKER,
GRACE_PERIOD_SECONDS,
GREPTILE_BOT_LOGINS,
SCORE_PATTERN,
extract_greptile_score,
gh,
parse_iso8601,

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

@ -12,68 +12,76 @@ import subprocess
import xml.etree.ElementTree as ET
from collections import defaultdict
from datetime import datetime
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 +97,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 +454,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)

View file

@ -9,7 +9,6 @@ from pathlib import Path
import testing.postgresql
DESTRUCTIVE_PATTERN = re.compile(r"\bDROP\s+(COLUMN|TABLE|INDEX)\b", re.IGNORECASE)
DEFAULT_BASE_BRANCH = "litellm_internal_staging"
@ -45,7 +44,7 @@ def _print_freshness_failure(
print("", file=out)
print("Options:", file=out)
print(
f" - Fix the above and re-run, OR pass --base-branch <name> if your", file=out
" - Fix the above and re-run, OR pass --base-branch <name> if your", file=out
)
print(
f" base branch is not '{base_branch}', OR pass --skip-freshness-check",

View file

@ -257,7 +257,6 @@
],
"source": [
"from litellm import acompletion\n",
"import asyncio\n",
"\n",
"async def test_get_response():\n",
" user_message = \"Hello, how are you?\"\n",
@ -313,8 +312,8 @@
}
],
"source": [
"from litellm import acompletion\n",
"import asyncio, os, traceback\n",
"import os\n",
"import traceback\n",
"\n",
"async def completion_call():\n",
" try:\n",
@ -384,7 +383,6 @@
}
],
"source": [
"import litellm\n",
"\n",
"\n",
"async def main():\n",
@ -421,9 +419,7 @@
}
],
"source": [
"import asyncio\n",
"\n",
"import litellm\n",
"\n",
"\n",
"async def main():\n",

View file

@ -6,7 +6,6 @@ from tabulate import tabulate
from termcolor import colored
import os
# Define the list of models to benchmark
# select any LLM listed here: https://docs.litellm.ai/docs/providers
models = ["gpt-3.5-turbo", "claude-2"]

View file

@ -74,7 +74,6 @@
"source": [
"# Setup environment and API keys\n",
"import os\n",
"import asyncio\n",
"from google.adk.agents import Agent\n",
"from google.adk.models.lite_llm import LiteLlm # For multi-model support\n",
"from google.adk.sessions import InMemorySessionService\n",

View file

@ -15,14 +15,13 @@ Usage:
import json
import os
from typing import Any, Dict, List, Optional
from typing import Any, Dict, Optional
import httpx
from fastapi import FastAPI, HTTPException, Header, Query
from fastapi.responses import JSONResponse
import uvicorn
app = FastAPI(
title="Braintrust Prompt Wrapper",
description="Wrapper server for Braintrust prompts to work with LiteLLM",
@ -264,7 +263,7 @@ def main():
print(f"🚀 Starting Braintrust Prompt Wrapper Server on {host}:{port}")
print(f"📚 API Documentation available at http://{host}:{port}/docs")
print(
f"🔑 Make sure to set BRAINTRUST_API_KEY environment variable or pass token in Authorization header"
"🔑 Make sure to set BRAINTRUST_API_KEY environment variable or pass token in Authorization header"
)
uvicorn.run(app, host=host, port=port)

View file

@ -6,7 +6,6 @@ This example shows how to use the CLI authentication token
in your Python scripts after running `litellm-proxy login`.
"""
from textwrap import indent
import litellm
LITELLM_BASE_URL = "http://localhost:4000/"

View file

@ -2,7 +2,7 @@
Simple xAI Voice Agent using LiveKit SDK with LiteLLM Gateway
This example shows how to use LiveKit's xAI realtime plugin through LiteLLM proxy.
LiteLLM acts as a unified interface, allowing you to switch between xAI, OpenAI,
LiteLLM acts as a unified interface, allowing you to switch between xAI, OpenAI,
and Azure realtime APIs without changing your agent code.
"""
@ -28,7 +28,7 @@ async def run_voice_agent():
url = f"ws://{PROXY_URL.replace('http://', '').replace('https://', '')}/v1/realtime?model={MODEL}"
headers = {"Authorization": f"Bearer {API_KEY}"}
print(f"🎙️ Connecting to voice agent...")
print("🎙️ Connecting to voice agent...")
print(f" Model: {MODEL}")
print(f" Proxy: {PROXY_URL}")
print()
@ -114,7 +114,7 @@ def main():
except Exception as e:
print(f"\n❌ Error: {e}")
print("\nMake sure LiteLLM proxy is running:")
print(f" litellm --config config.yaml --port 4000")
print(" litellm --config config.yaml --port 4000")
if __name__ == "__main__":

View file

@ -1,14 +1,14 @@
"""
LiteLLM Migration Script!
Takes a config.yaml and calls /model/new
Takes a config.yaml and calls /model/new
Inputs:
- File path to config.yaml
- Proxy base url to your hosted proxy
Step 1: Reads your config.yaml
Step 2: reads `model_list` and loops through all models
Step 2: reads `model_list` and loops through all models
Step 3: calls `<proxy-base-url>/model/new` for each model
"""

View file

@ -515,11 +515,10 @@ if __name__ == "__main__":
print("=" * 80)
print(f"Server starting on: http://{host}:{port}")
print(f"Bearer Token: {bearer_token}")
print(f"Endpoint: POST /guardrail/{{id}}/version/{{version}}/apply")
print("Endpoint: POST /guardrail/{id}/version/{version}/apply")
print("=" * 80)
print("\nExample curl command:")
print(
f"""
print(f"""
curl -X POST "http://{host}:{port}/guardrail/test-guardrail/version/1/apply" \\
-H "Authorization: Bearer {bearer_token}" \\
-H "Content-Type: application/json" \\
@ -533,8 +532,7 @@ curl -X POST "http://{host}:{port}/guardrail/test-guardrail/version/1/apply" \\
}}
]
}}'
"""
)
""")
print("=" * 80)
uvicorn.run(app, host=host, port=port)

View file

@ -14,12 +14,9 @@ Test the endpoint:
curl "http://localhost:8080/beta/litellm_prompt_management?prompt_id=hello-world-prompt"
"""
import os
import json
from typing import Any, Dict, List, Optional
from fastapi import FastAPI, HTTPException, Header, Query, status
from fastapi.responses import JSONResponse
from pydantic import BaseModel, Field
# ============================================================================
@ -364,7 +361,7 @@ if __name__ == "__main__":
print("=" * 70)
print("Mock Prompt Management API Server")
print("=" * 70)
print(f"\nStarting server on http://localhost:8080")
print("\nStarting server on http://localhost:8080")
print(f"\nAvailable prompts: {len(PROMPTS_DB)}")
for prompt_id in PROMPTS_DB.keys():
print(f" - {prompt_id}")

View file

@ -157,7 +157,7 @@ class VeoVideoGenerator:
Returns:
True if download successful, False otherwise
"""
print(f"⬇️ Downloading video...")
print("⬇️ Downloading video...")
print(f"Original URI: {video_uri}")
# Convert Google URI to LiteLLM proxy URI
@ -198,7 +198,7 @@ class VeoVideoGenerator:
if os.path.exists(output_filename):
file_size = os.path.getsize(output_filename)
if file_size > 0:
print(f"✅ Video downloaded successfully!")
print("✅ Video downloaded successfully!")
print(f"📁 Saved as: {output_filename}")
print(f"📏 File size: {file_size / (1024*1024):.2f} MB")
return True

View file

@ -3,7 +3,7 @@ import csv
import json
import asyncio
from datetime import datetime
from typing import Optional, List, Dict, Any
from typing import Any
import os
@ -172,7 +172,7 @@ async def migrate_verification_tokens():
)
continue
print(f"\nMigration completed!")
print("\nMigration completed!")
print(f"Successfully processed: {processed_count} records")
print(f"Errors encountered: {error_count} records")

View file

@ -15,11 +15,10 @@ sys.path.insert(
) # Adds the parent directory to the system path
import json
import sys
from typing import Any, List, Literal, Optional
from typing import Any, List, Optional
from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.logging_utils import (

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

@ -15,7 +15,7 @@ sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import sys
from typing import Literal, Optional
from typing import Optional
from fastapi import HTTPException

View file

@ -7,7 +7,7 @@
# Thank you users! We ❤️ you! - Krrish & Ishaan
## This provides an LLM Guard Integration for content moderation on the proxy
from typing import Literal, Optional
from typing import Optional
import aiohttp
from fastapi import HTTPException

View file

@ -349,8 +349,10 @@ class BaseEmailLogger(CustomLogger):
)
# Calculate percentage and alert threshold
percentage = threshold_pct if threshold_pct is not None else int(
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE * 100
percentage = (
threshold_pct
if threshold_pct is not None
else int(EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE * 100)
)
threshold_fraction = percentage / 100.0
alert_threshold_str = (
@ -609,9 +611,7 @@ class BaseEmailLogger(CustomLogger):
continue
_id = user_info.token or user_info.user_id or "default_id"
_cache_key = (
f"email_budget_alerts:max_budget_alert:{threshold_pct}:{_id}"
)
_cache_key = f"email_budget_alerts:max_budget_alert:{threshold_pct}:{_id}"
result = await _cache.async_get_cache(key=_cache_key)
if result is not None:
@ -630,7 +630,9 @@ class BaseEmailLogger(CustomLogger):
continue
recipient_emails = list(set(emails))
event_message = f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
event_message = (
f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
)
webhook_event = WebhookEvent(
event="max_budget_alert",
event_message=event_message,

View file

@ -15,7 +15,6 @@ from litellm.llms.custom_httpx.http_handler import (
from .base_email import BaseEmailLogger
SENDGRID_API_ENDPOINT = "https://api.sendgrid.com/v3/mail/send"
@ -79,4 +78,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

@ -7,7 +7,7 @@ GET - /audit/{id} - Get audit log by id
GET - /audit - Get all audit logs
"""
from typing import Any, Dict, List, Optional
from typing import Any, Dict, Optional
#### AUDIT LOGGING ####
from fastapi import APIRouter, Depends, HTTPException, Query
@ -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

@ -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,21 +328,32 @@ class CheckBatchCost:
# CheckBatchCost bypasses async_post_call_success_hook, so convert raw
# output/error file IDs to managed base64 IDs before the DB write here.
managed_files_hook = self.proxy_logging_obj.get_proxy_hook("managed_files")
managed_files_hook = self.proxy_logging_obj.get_proxy_hook(
"managed_files"
)
if managed_files_hook is not None:
from litellm.proxy._types import UserAPIKeyAuth
_minimal_auth = UserAPIKeyAuth(
user_id=job.created_by or "default-user-id",
team_id=getattr(job, "team_id", None),
)
for _file_attr in ["output_file_id", "error_file_id"]:
_raw_file_id = getattr(response, _file_attr, None)
if _raw_file_id and not _is_base64_encoded_unified_file_id(_raw_file_id):
if _raw_file_id and not _is_base64_encoded_unified_file_id(
_raw_file_id
):
try:
_unified_file_id = managed_files_hook.get_unified_output_file_id(
output_file_id=_raw_file_id,
model_id=model_id,
model_name=str(model_name) if model_name else deployment_info.model_name or None,
_unified_file_id = (
managed_files_hook.get_unified_output_file_id(
output_file_id=_raw_file_id,
model_id=model_id,
model_name=(
str(model_name)
if model_name
else deployment_info.model_name or None
),
)
)
await managed_files_hook.store_unified_file_id(
file_id=_unified_file_id,
@ -338,7 +375,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,
@ -385,7 +426,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,
@ -394,7 +437,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

@ -6,7 +6,6 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
from fastapi import HTTPException
import litellm
from litellm import Router, verbose_logger
from litellm._uuid import uuid
from litellm.integrations.custom_logger import CustomLogger
@ -41,7 +40,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 +76,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 +108,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 +136,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 +193,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 +209,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 +235,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 +260,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 +286,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 +302,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 +325,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 +366,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 +402,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 +431,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

@ -8,7 +8,6 @@ All /vector_store management endpoints
/vector_store/list
"""
import copy
import json
from typing import List, Optional
@ -147,12 +146,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

@ -106,7 +106,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
# Health & ops
"/health",
"/metrics",
"/watsonx"
"/watsonx",
)
GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(

View file

@ -266,6 +266,9 @@ class MistralConfig(OpenAIGPTConfig):
for m in messages:
m = MistralConfig._handle_name_in_message(m)
m = MistralConfig._handle_tool_call_message(m)
m.pop("metadata", None)
if MistralConfig._is_empty_assistant_message(m):
continue
m = strip_none_values_from_message(m) # prevents 'extra_forbidden' error

View file

@ -35,7 +35,7 @@ import httpx
class EvalCase:
category: str
prompt: str
ideal: str # criteria the judge checks the response against
ideal: str # criteria the judge checks the response against
EVAL_CASES: List[EvalCase] = [
@ -177,14 +177,19 @@ async def evaluate(
async with httpx.AsyncClient() as client:
for i, case in enumerate(EVAL_CASES, 1):
print(f"\n[{i}/{len(EVAL_CASES)}] category={case.category}")
print(f" prompt : {case.prompt[:80]}{'…' if len(case.prompt) > 80 else ''}")
print(
f" prompt : {case.prompt[:80]}{'…' if len(case.prompt) > 80 else ''}"
)
session_id = f"eval-{uuid.uuid4()}"
# Round 1: single-turn real request — get the actual LLM response to judge.
try:
response, chosen = await _chat(
client, proxy_url, api_key, router,
client,
proxy_url,
api_key,
router,
[{"role": "user", "content": case.prompt}],
session_id=session_id,
)
@ -194,16 +199,25 @@ async def evaluate(
continue
print(f" model : {chosen or router}")
print(f" response : {response[:120].replace(chr(10), ' ')}{'…' if len(response) > 120 else ''}")
print(
f" response : {response[:120].replace(chr(10), ' ')}{'…' if len(response) > 120 else ''}"
)
# Judge the real response.
judge_msgs = [
{"role": "system", "content": JUDGE_SYSTEM},
{"role": "user", "content": _judge_user(case.prompt, case.ideal, response)},
{
"role": "user",
"content": _judge_user(case.prompt, case.ideal, response),
},
]
try:
verdict, _ = await _chat(
client, proxy_url, api_key, judge_model, judge_msgs,
client,
proxy_url,
api_key,
judge_model,
judge_msgs,
)
except Exception as exc: # noqa: BLE001
print(f" ERROR calling judge: {exc}", file=sys.stderr)
@ -227,15 +241,19 @@ async def evaluate(
# On PASS → satisfaction follow-up (+alpha). On FAIL → neutral (no signal).
follow_up = SATISFY_FOLLOWUP if is_pass else NEUTRAL_FOLLOWUP
bandit_msgs = [
{"role": "user", "content": case.prompt},
{"role": "user", "content": case.prompt},
{"role": "assistant", "content": response},
{"role": "user", "content": "ok continue"},
{"role": "user", "content": "ok continue"},
{"role": "assistant", "content": FAB_ASSISTANT},
{"role": "user", "content": follow_up},
{"role": "user", "content": follow_up},
]
try:
await _chat(
client, proxy_url, api_key, router, bandit_msgs,
client,
proxy_url,
api_key,
router,
bandit_msgs,
session_id=session_id,
)
except Exception as exc: # noqa: BLE001
@ -257,11 +275,17 @@ async def evaluate(
# Entry point
# ---------------------------------------------------------------------------
def main() -> None:
ap = argparse.ArgumentParser(description="Evaluate the adaptive router with LLM-as-judge.")
ap.add_argument("--proxy-url", default="http://localhost:4000")
ap.add_argument("--api-key", required=True, help="proxy API key")
ap.add_argument("--router", default="smart-cheap-router", help="adaptive router model name")
ap.add_argument("--judge-model", default="smart", help="model name for the judge (via proxy)")
ap = argparse.ArgumentParser(
description="Evaluate the adaptive router with LLM-as-judge."
)
ap.add_argument("--proxy-url", default="http://localhost:4000")
ap.add_argument("--api-key", required=True, help="proxy API key")
ap.add_argument(
"--router", default="smart-cheap-router", help="adaptive router model name"
)
ap.add_argument(
"--judge-model", default="smart", help="model name for the judge (via proxy)"
)
args = ap.parse_args()
asyncio.run(evaluate(args.proxy_url, args.api_key, args.router, args.judge_model))

View file

@ -72,8 +72,8 @@ PROMPTS: Dict[str, List[str]] = {
# so that signals attribute to the right (type, model) bandit cell.
SATISFY: Dict[str, str] = {
"code_generation": "thanks, that works! now write me a python function that does the inverse",
"factual_lookup": "perfect, thanks! who is the current prime minister?",
"writing": "great, thanks! now write a follow-up email confirming attendance",
"factual_lookup": "perfect, thanks! who is the current prime minister?",
"writing": "great, thanks! now write a follow-up email confirming attendance",
}
# Neutral follow-up — does not match any signal regex, does not move the bandit.
@ -83,8 +83,8 @@ NEUTRAL_FOLLOWUP = "ok, noted"
# Defaults: smart dominates code/writing; both are fine for factual_lookup.
ORACLE: Dict[str, Dict[str, float]] = {
"code_generation": {"smart": 0.92, "fast": 0.35},
"factual_lookup": {"smart": 0.90, "fast": 0.85},
"writing": {"smart": 0.85, "fast": 0.55},
"factual_lookup": {"smart": 0.90, "fast": 0.85},
"writing": {"smart": 0.85, "fast": 0.55},
}
# Fabricated assistant turn — content doesn't matter for the hook, only the role.
@ -94,11 +94,11 @@ FAB_ASSISTANT = "Got it. Working on that now."
def _build_messages(prompt: str, last_user: str) -> List[Dict[str, str]]:
"""5-message conversation that passes the SIGNAL_GATE_MIN_MESSAGES=4 gate."""
return [
{"role": "user", "content": prompt},
{"role": "user", "content": prompt},
{"role": "assistant", "content": FAB_ASSISTANT},
{"role": "user", "content": "ok continue"},
{"role": "user", "content": "ok continue"},
{"role": "assistant", "content": FAB_ASSISTANT},
{"role": "user", "content": last_user},
{"role": "user", "content": last_user},
]
@ -155,7 +155,11 @@ async def _drive_one_session(
#
# Round 1: neutral follow-up → no signal fires, but we learn the pick.
ok, chosen = await _send(
client, proxy_url, api_key, router, session_id,
client,
proxy_url,
api_key,
router,
session_id,
_build_messages(prompt, NEUTRAL_FOLLOWUP),
mock_response=FAB_ASSISTANT,
)
@ -171,10 +175,15 @@ async def _drive_one_session(
# follow-up matches satisfaction → +alpha for (request_type, chosen).
history = _build_messages(prompt, NEUTRAL_FOLLOWUP) + [
{"role": "assistant", "content": FAB_ASSISTANT},
{"role": "user", "content": follow_up},
{"role": "user", "content": follow_up},
]
await _send(
client, proxy_url, api_key, router, session_id, history,
client,
proxy_url,
api_key,
router,
session_id,
history,
mock_response=FAB_ASSISTANT,
)
return chosen
@ -183,13 +192,22 @@ async def _drive_one_session(
async def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--proxy-url", default="http://localhost:4000")
ap.add_argument("--api-key", required=True, help="proxy key with /v1/chat/completions perms")
ap.add_argument("--router", default="smart-cheap-router")
ap.add_argument("--rounds", type=int, default=100)
ap.add_argument("--rate", type=float, default=0.5,
help="seconds between sessions; lower = faster")
ap.add_argument("--types", default="code_generation,factual_lookup,writing",
help="comma-separated subset of request types to drive")
ap.add_argument(
"--api-key", required=True, help="proxy key with /v1/chat/completions perms"
)
ap.add_argument("--router", default="smart-cheap-router")
ap.add_argument("--rounds", type=int, default=100)
ap.add_argument(
"--rate",
type=float,
default=0.5,
help="seconds between sessions; lower = faster",
)
ap.add_argument(
"--types",
default="code_generation,factual_lookup,writing",
help="comma-separated subset of request types to drive",
)
args = ap.parse_args()
types = [t.strip() for t in args.types.split(",") if t.strip() in PROMPTS]
@ -207,7 +225,12 @@ async def main() -> None:
rt = random.choice(types)
prompt = random.choice(PROMPTS[rt])
chosen = await _drive_one_session(
client, args.proxy_url, args.api_key, args.router, rt, prompt,
client,
args.proxy_url,
args.api_key,
args.router,
rt,
prompt,
)
if chosen:
counts[(rt, chosen)] = counts.get((rt, chosen), 0) + 1

View file

@ -24,7 +24,6 @@ import shlex
import signal
import statistics
import subprocess
import sys
import tempfile
import time
from dataclasses import dataclass
@ -34,7 +33,6 @@ from typing import Any, Optional
import aiohttp
from aiohttp import web
DEFAULT_MODEL = "perf-test-model"
DEFAULT_API_KEY = "sk-1234"

View file

@ -8,7 +8,6 @@ import statistics
import aiohttp
REQUEST_BODY = {
"model": "db-openai-endpoint",
"messages": [{"role": "user", "content": "hi"}],
@ -150,7 +149,7 @@ def print_aggregate(results):
cov_throughput = (
statistics.stdev(run_throughputs) / statistics.mean(run_throughputs) * 100
)
print(f"\n Run-to-run variance:")
print("\n Run-to-run variance:")
print(f" Latency CoV: {cov_latency:.1f}%")
print(f" Throughput CoV: {cov_throughput:.1f}%")

View file

@ -132,7 +132,7 @@ def run_scenario(
else:
runner = lambda: bench_with_chunk(
wrapper, spec["chunk_factory"], iterations
) # noqa: E731
)
for _ in range(warmup):
runner()

View file

@ -11,7 +11,7 @@ USAGE EXAMPLES:
export PROVIDER_URL='https://api.openai.com/v1/chat/completions'
export LITELLM_PROXY_API_KEY='sk-1234'
export PROVIDER_API_KEY='sk-openai-key'
# Run from scripts directory
cd scripts
python benchmark_proxy_vs_provider.py
@ -266,7 +266,7 @@ async def benchmark_endpoint(
print(f"\nStarting benchmark for {url}")
if warmup:
print(f" Warming up with 5 requests...")
print(" Warming up with 5 requests...")
await warmup_endpoint(
url, headers, payload, num_warmup=5, timeout_seconds=timeout_seconds
)
@ -359,7 +359,7 @@ def print_results(name: str, results: BenchmarkResults):
if "latency_stats" in stats:
latency = stats["latency_stats"]
print(f"\nLatency Statistics (seconds):")
print("\nLatency Statistics (seconds):")
print(f" Mean: {latency['mean']:.4f}s")
print(f" Median (p50): {latency['median']:.4f}s")
print(f" Min: {latency['min']:.4f}s")
@ -369,12 +369,12 @@ def print_results(name: str, results: BenchmarkResults):
print(f" p99: {latency['p99']:.4f}s")
if stats["status_codes"]:
print(f"\nStatus Codes:")
print("\nStatus Codes:")
for code, count in sorted(stats["status_codes"].items()):
print(f" {code}: {count}")
if results.errors:
print(f"\nErrors (showing first 5 unique):")
print("\nErrors (showing first 5 unique):")
unique_errors = list(set(results.errors))[:5]
for error in unique_errors:
count = results.errors.count(error)
@ -439,7 +439,7 @@ def print_run_variance(name: str, results_list: List[BenchmarkResults]):
throughputs.append(stats["requests_per_second"])
if mean_latencies:
print(f"\nMean Latency Variance:")
print("\nMean Latency Variance:")
print(f" Runs: {len(mean_latencies)}")
print(f" Mean: {mean(mean_latencies):.4f}s")
print(f" Min: {min(mean_latencies):.4f}s")
@ -456,7 +456,7 @@ def print_run_variance(name: str, results_list: List[BenchmarkResults]):
)
if throughputs:
print(f"\nThroughput Variance:")
print("\nThroughput Variance:")
print(f" Mean: {mean(throughputs):.2f} req/s")
print(f" Min: {min(throughputs):.2f} req/s")
print(f" Max: {max(throughputs):.2f} req/s")
@ -475,18 +475,18 @@ def compare_results(
provider_stats = provider_results.calculate_stats()
print(f"\n{'='*60}")
print(f"Comparison: LiteLLM Proxy vs Direct Provider")
print("Comparison: LiteLLM Proxy vs Direct Provider")
print(f"{'='*60}")
# Success Rate Comparison
print(f"\nSuccess Rate:")
print("\nSuccess Rate:")
print(f" Proxy: {proxy_stats['success_rate']:.2f}%")
print(f" Provider: {provider_stats['success_rate']:.2f}%")
diff = proxy_stats["success_rate"] - provider_stats["success_rate"]
print(f" Difference: {diff:+.2f}%")
# Throughput Comparison
print(f"\nThroughput (requests/second):")
print("\nThroughput (requests/second):")
print(f" Proxy: {proxy_stats['requests_per_second']:.2f}")
print(f" Provider: {provider_stats['requests_per_second']:.2f}")
diff = proxy_stats["requests_per_second"] - provider_stats["requests_per_second"]
@ -494,7 +494,7 @@ def compare_results(
# Latency Comparison
if "latency_stats" in proxy_stats and "latency_stats" in provider_stats:
print(f"\nLatency Comparison (seconds):")
print("\nLatency Comparison (seconds):")
proxy_latency = proxy_stats["latency_stats"]
provider_latency = provider_stats["latency_stats"]
@ -509,7 +509,7 @@ def compare_results(
)
# Total Time Comparison
print(f"\nTotal Time:")
print("\nTotal Time:")
print(f" Proxy: {proxy_stats['total_time']:.2f}s")
print(f" Provider: {provider_stats['total_time']:.2f}s")
diff = proxy_stats["total_time"] - provider_stats["total_time"]
@ -662,7 +662,7 @@ Examples:
print("=" * 60)
print("LiteLLM Proxy vs Provider Benchmark")
print("=" * 60)
print(f"Configuration (from environment variables):")
print("Configuration (from environment variables):")
print(f" Proxy URL: {LITELLM_PROXY_URL}")
print(f" Provider URL: {PROVIDER_URL}")
print(
@ -685,15 +685,15 @@ Examples:
)
if not args.max_concurrent:
print(f"\nTip: Use --max-concurrent 100 for more realistic load testing")
print(f" (prevents overwhelming the server with all requests at once)")
print("\nTip: Use --max-concurrent 100 for more realistic load testing")
print(" (prevents overwhelming the server with all requests at once)")
if args.parallel:
print(f"\nWARNING: Running benchmarks in parallel may affect results due to:")
print(f" - Shared network bandwidth")
print(f" - Provider endpoint receiving double load (via proxy + direct)")
print(f" - Potential rate limiting issues")
print(f" - Resource contention")
print("\nWARNING: Running benchmarks in parallel may affect results due to:")
print(" - Shared network bandwidth")
print(" - Provider endpoint receiving double load (via proxy + direct)")
print(" - Potential rate limiting issues")
print(" - Resource contention")
# Run benchmarks multiple times if requested
all_proxy_results = []
@ -703,7 +703,7 @@ Examples:
if args.runs > 1:
print(f"\nRunning {args.runs} benchmark runs for statistical accuracy...")
print(f" Results will be averaged across all runs.\n")
print(" Results will be averaged across all runs.\n")
overall_start_time = time.perf_counter()
@ -718,7 +718,7 @@ Examples:
print(f"{'='*60}")
if args.parallel:
print(f"\nRunning both benchmarks in parallel...")
print("\nRunning both benchmarks in parallel...")
proxy_results, provider_results = await asyncio.gather(
benchmark_endpoint(
LITELLM_PROXY_URL,
@ -740,9 +740,9 @@ Examples:
),
)
else:
print(f"\nRunning benchmarks sequentially (proxy first, then provider)...")
print("\nRunning benchmarks sequentially (proxy first, then provider)...")
if run_num == 1:
print(f" This ensures accurate results without interference.\n")
print(" This ensures accurate results without interference.\n")
proxy_results = await benchmark_endpoint(
LITELLM_PROXY_URL,
@ -755,7 +755,7 @@ Examples:
)
if run_num < args.runs or args.runs == 1:
print(f"\nWaiting 3 seconds before starting provider benchmark...")
print("\nWaiting 3 seconds before starting provider benchmark...")
await asyncio.sleep(3) # Longer pause to ensure clean separation
provider_results = await benchmark_endpoint(
@ -773,7 +773,7 @@ Examples:
# Brief pause between runs
if run_num < args.runs:
print(f"\nWaiting 5 seconds before next run...")
print("\nWaiting 5 seconds before next run...")
await asyncio.sleep(5)
overall_benchmark_time = time.perf_counter() - overall_start_time
@ -790,7 +790,7 @@ Examples:
raise RuntimeError("Benchmark results not initialized")
final_proxy_results = proxy_results
final_provider_results = provider_results
print(f"\nResults:")
print("\nResults:")
# Print individual results
print_results("LiteLLM Proxy", final_proxy_results)

View file

@ -223,6 +223,7 @@ async def drive_async(
# Repeat × take-min runner
# ---------------------------------------------------------------------------
@dataclass
class Result:
label: str

View file

@ -66,7 +66,12 @@ def _load_head(rel: str) -> dict | None:
def _ref_is_commit(ref: str) -> bool:
return _run(["git", "rev-parse", "--verify", "--quiet", f"{ref}^{{commit}}"]).returncode == 0
return (
_run(
["git", "rev-parse", "--verify", "--quiet", f"{ref}^{{commit}}"]
).returncode
== 0
)
def _load_base(rel: str, ref: str) -> dict | None:
@ -102,9 +107,11 @@ def regressions_for(rel: str, base: dict | None, head: dict | None) -> list[Regr
Regression(
rel,
rule,
f"rule dropped (ceiling {base_cap} -> removed)"
if rule not in head_caps
else f"ceiling raised {base_cap} -> {head_caps[rule]}",
(
f"rule dropped (ceiling {base_cap} -> removed)"
if rule not in head_caps
else f"ceiling raised {base_cap} -> {head_caps[rule]}"
),
)
for rule, base_cap in sorted(base_caps.items())
if rule not in head_caps or head_caps[rule] > base_cap
@ -141,7 +148,9 @@ def main() -> int:
regressions.extend(regressions_for(rel, base, head))
if regressions:
print(f"FAIL: budget ceiling(s) loosened vs base {args.base} (merge-base {ref[:12]}):")
print(
f"FAIL: budget ceiling(s) loosened vs base {args.base} (merge-base {ref[:12]}):"
)
for reg in regressions:
print(f" {reg.budget} {reg.rule}: {reg.detail}")
print(

View file

@ -1,6 +1,6 @@
#!/usr/bin/env python3
"""Type-discipline checker: the rules ruff can't enforce.
Rules
-----
LIT001 Mutable collection in a type annotation, anywhere it appears: function
@ -41,16 +41,16 @@ LIT008 `**kwargs` parameter. The keyword contract is erased and everything it c
LIT000 Setup failure: a target file could not be read, or contains a syntax error.
Reported as a violation rather than crashing the run.
Usage
-----
python check_type_discipline.py litellm/ tests/
Exit code 1 if any violation is found. Stdlib only.
"""
from __future__ import annotations
import ast
import io
import re
@ -60,27 +60,48 @@ from dataclasses import dataclass
from pathlib import Path
from collections.abc import Iterable, Iterator, Sequence
from typing import NamedTuple
# Mutable collection types, banned in *every* annotation. Name-based, so `dict`,
# `typing.Dict`, `collections.deque`, and `collections.abc.MutableMapping` all match
# however they were imported. The read-only interfaces (Mapping, Sequence, the
# immutable AbstractSet / `abc.Set`, Collection) and the immutable concretes (tuple,
# frozenset) are the escape hatch and are deliberately absent -- as is the bare name
# `Set`, which collides with the read-only `collections.abc.Set`.
MUTABLE_COLLECTIONS = frozenset((
"dict", "list", "set",
"Dict", "List", "DefaultDict", "OrderedDict", "Counter", "Deque", "ChainMap",
"deque", "defaultdict",
"MutableMapping", "MutableSequence", "MutableSet",
))
MUTABLE_COLLECTIONS = frozenset(
(
"dict",
"list",
"set",
"Dict",
"List",
"DefaultDict",
"OrderedDict",
"Counter",
"Deque",
"ChainMap",
"deque",
"defaultdict",
"MutableMapping",
"MutableSequence",
"MutableSet",
)
)
# Callables whose result is a fresh *mutable* collection (LIT002). `tuple` and
# `frozenset` are deliberately absent -- they are the wrappers you reach for, and
# a generator expression fed to them is the blessed one-shot build.
MUTABLE_CONSTRUCTORS = frozenset((
"dict", "list", "set",
"deque", "defaultdict", "OrderedDict", "Counter", "ChainMap",
))
MUTABLE_CONSTRUCTORS = frozenset(
(
"dict",
"list",
"set",
"deque",
"defaultdict",
"OrderedDict",
"Counter",
"ChainMap",
)
)
# A *qualified* call (`x.deque()`) counts as construction only for names that are rarely
# method names; `dict`/`list`/`set` are dropped here because `.dict()` / `.set()` / `.list()`
# are common methods (e.g. pydantic's `model.dict()`), not collection construction. A
@ -88,7 +109,7 @@ MUTABLE_CONSTRUCTORS = frozenset((
QUALIFIED_CONSTRUCTORS = MUTABLE_CONSTRUCTORS - frozenset(("dict", "list", "set"))
UNSAFE_GUARDS = frozenset(("TypeGuard", "TypeIs"))
MIN_REASON_LEN = 3
NOQA_RE = re.compile(
r"#\s*noqa"
r"(?P<colon>:\s*(?P<codes>[A-Z]+[0-9]+(?:\s*,\s*[A-Z]+[0-9]+)*))?"
@ -110,18 +131,18 @@ OK_SUPPRESSIONS: tuple[tuple[str, re.Pattern[str]], ...] = (
("guard-ok", GUARD_OK_RE),
("kwargs-ok", KWARGS_OK_RE),
)
class Violation(NamedTuple):
path: Path
line: int
code: str
message: str
def render(self) -> str:
return f"{self.path}:{self.line}: {self.code} {self.message}"
@dataclass(frozen=True, slots=True)
class Comments:
"""The lines carrying each valid `*-ok` suppression."""
@ -130,17 +151,17 @@ class Comments:
cast_ok_lines: frozenset[int]
guard_ok_lines: frozenset[int]
kwargs_ok_lines: frozenset[int]
# --------------------------------------------------------------------------- #
# Comment scanning (LIT003 / LIT004 / LIT005)
# --------------------------------------------------------------------------- #
def _reason_of(rest: str) -> str:
return rest.strip().lstrip("#-").strip()
def _valid_ok(regex: re.Pattern[str], text: str) -> bool:
"""True iff `text` carries this suppression with a reason of usable length."""
m = regex.search(text)
@ -152,30 +173,55 @@ def _comment_violations(path: Path, line_no: int, text: str) -> Iterator[Violati
for token, regex in OK_SUPPRESSIONS:
m = regex.search(text)
if m and len((m.group("reason") or "").strip()) < MIN_REASON_LEN:
yield Violation(path, line_no, "LIT005", f"{token} requires a reason: `# {token}: <reason>`")
yield Violation(
path,
line_no,
"LIT005",
f"{token} requires a reason: `# {token}: <reason>`",
)
m = NOQA_RE.search(text)
if m:
if not m.group("codes"):
yield Violation(path, line_no, "LIT003", "noqa requires rule codes: `# noqa: XXX123 # <reason>`")
yield Violation(
path,
line_no,
"LIT003",
"noqa requires rule codes: `# noqa: XXX123 # <reason>`",
)
elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN:
yield Violation(path, line_no, "LIT003", "noqa requires a reason: `# noqa: XXX123 # <reason>`")
yield Violation(
path,
line_no,
"LIT003",
"noqa requires a reason: `# noqa: XXX123 # <reason>`",
)
m = IGNORE_RE.search(text)
if m:
codes = m.group("codes")
if not codes or codes == "[]":
yield Violation(path, line_no, "LIT004",
"ignore requires codes: `# pyright: ignore[ruleName] # <reason>`")
yield Violation(
path,
line_no,
"LIT004",
"ignore requires codes: `# pyright: ignore[ruleName] # <reason>`",
)
elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN:
yield Violation(path, line_no, "LIT004",
"ignore requires a reason: `# pyright: ignore[ruleName] # <reason>`")
yield Violation(
path,
line_no,
"LIT004",
"ignore requires a reason: `# pyright: ignore[ruleName] # <reason>`",
)
def scan_comments(path: Path, source: str) -> tuple[Comments, tuple[Violation, ...]]:
try:
tokens = tokenize.generate_tokens(io.StringIO(source).readline)
comment_toks = tuple((t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT)
comment_toks = tuple(
(t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT
)
except (tokenize.TokenError, SyntaxError):
# tokenize raises TokenError (EOF mid-construct) or a SyntaxError subclass
# (IndentationError / TabError) on malformed source; defer to ast.parse below,
@ -192,13 +238,17 @@ def scan_comments(path: Path, source: str) -> tuple[Comments, tuple[Violation, .
guard_ok_lines=_lines_with(GUARD_OK_RE),
kwargs_ok_lines=_lines_with(KWARGS_OK_RE),
),
tuple(v for line, text in comment_toks for v in _comment_violations(path, line, text)),
tuple(
v
for line, text in comment_toks
for v in _comment_violations(path, line, text)
),
)
# --------------------------------------------------------------------------- #
def mutable_names_in(annotation: ast.expr) -> Iterator[str]:
"""Yield mutable-collection names anywhere inside an annotation expression.
@ -219,11 +269,13 @@ def mutable_names_in(annotation: ast.expr) -> Iterator[str]:
except SyntaxError:
continue
yield from mutable_names_in(inner)
def _mutable_ann(path: Path, line: int, name: str, where: str) -> Violation:
return Violation(
path, line, "LIT001",
path,
line,
"LIT001",
f"mutable `{name}` in {where}: a mutable collection can be grown or rewritten "
f"by whoever holds it. Annotate a read-only view -- Mapping[...], Sequence[...], "
f"AbstractSet[...], tuple[X, ...], frozenset[X], or a frozen dataclass / "
@ -233,13 +285,19 @@ def _mutable_ann(path: Path, line: int, name: str, where: str) -> Violation:
def _annotation_violations(
path: Path, annotation: ast.expr | None, line: int, where: str, ok_lines: frozenset[int]
path: Path,
annotation: ast.expr | None,
line: int,
where: str,
ok_lines: frozenset[int],
) -> Iterator[Violation]:
if annotation is None or line in ok_lines:
return
yield from (_mutable_ann(path, line, name, where) for name in mutable_names_in(annotation))
yield from (
_mutable_ann(path, line, name, where) for name in mutable_names_in(annotation)
)
def _function_violations(
path: Path, node: ast.FunctionDef | ast.AsyncFunctionDef, comments: Comments
) -> Iterator[Violation]:
@ -247,14 +305,22 @@ def _function_violations(
args = node.args
for arg in (*args.posonlyargs, *args.args, *args.kwonlyargs):
yield from _annotation_violations(
path, arg.annotation, arg.lineno, f"parameter `{arg.arg}` of `{node.name}`", mutable_ok
path,
arg.annotation,
arg.lineno,
f"parameter `{arg.arg}` of `{node.name}`",
mutable_ok,
)
# *args is allowed when typed (it's just a tuple); ruff ANN002 forces the
# annotation, so here we only add the LIT001 mutable-collection check on the element type.
if args.vararg is not None:
yield from _annotation_violations(
path, args.vararg.annotation, args.vararg.lineno, f"`*args` of `{node.name}`", mutable_ok
path,
args.vararg.annotation,
args.vararg.lineno,
f"`*args` of `{node.name}`",
mutable_ok,
)
# **kwargs is banned outright (LIT008): it erases the keyword contract and forces
@ -262,7 +328,9 @@ def _function_violations(
# cannot ban the syntax, so this rule does.
if args.kwarg is not None and args.kwarg.lineno not in comments.kwargs_ok_lines:
yield Violation(
path, args.kwarg.lineno, "LIT008",
path,
args.kwarg.lineno,
"LIT008",
f"`**{args.kwarg.arg}` is banned: it erases the keyword contract and forces "
f"Any-typing; declare explicit keyword parameters, or accept one frozen payload "
f"(frozen dataclass / NamedTuple / ReadOnly TypedDict) "
@ -271,11 +339,17 @@ def _function_violations(
if node.returns is not None:
yield from _annotation_violations(
path, node.returns, node.returns.lineno, f"return type of `{node.name}`", mutable_ok
path,
node.returns,
node.returns.lineno,
f"return type of `{node.name}`",
mutable_ok,
)
def iter_annotation_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]:
def iter_annotation_violations(
path: Path, tree: ast.AST, comments: Comments
) -> Iterator[Violation]:
# Every annotation is in scope: signatures (params / *args / return) plus every
# `x: T` -- class attribute, local, or module global. The latter three are all
# ast.AnnAssign, so one walk covers them; only the signature annotations (which
@ -286,8 +360,11 @@ def iter_annotation_violations(path: Path, tree: ast.AST, comments: Comments) ->
elif isinstance(node, ast.AnnAssign):
target = node.target.id if isinstance(node.target, ast.Name) else "<target>"
yield from _annotation_violations(
path, node.annotation, node.lineno,
f"the type of `{target}`", comments.mutable_ok_lines,
path,
node.annotation,
node.lineno,
f"the type of `{target}`",
comments.mutable_ok_lines,
)
@ -308,39 +385,54 @@ def _is_cast_call(node: ast.Call) -> bool:
)
def iter_cast_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]:
def iter_cast_violations(
path: Path, tree: ast.AST, comments: Comments
) -> Iterator[Violation]:
for node in ast.walk(tree):
if isinstance(node, ast.Call) and _is_cast_call(node) and node.lineno not in comments.cast_ok_lines:
if (
isinstance(node, ast.Call)
and _is_cast_call(node)
and node.lineno not in comments.cast_ok_lines
):
yield Violation(
path, node.lineno, "LIT006",
path,
node.lineno,
"LIT006",
"cast() is an unchecked assertion (the type checker takes it on faith); "
"validate into a frozen dataclass/NamedTuple/ReadOnly TypedDict at the "
"boundary instead (suppress: `# cast-ok: <reason>`)",
)
def iter_guard_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]:
def iter_guard_violations(
path: Path, tree: ast.AST, comments: Comments
) -> Iterator[Violation]:
# TypeGuard/TypeIs are legal only as a function's return annotation (`-> TypeGuard[int]`),
# so the walk is confined to `node.returns`; a runtime name that merely happens to read
# `TypeGuard` is not a narrowing predicate. ruff bans the import; this flags the use.
for node in ast.walk(tree):
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) or node.returns is None:
if (
not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
or node.returns is None
):
continue
for sub in ast.walk(node.returns):
name = (
sub.id if isinstance(sub, ast.Name)
else sub.attr if isinstance(sub, ast.Attribute)
else None
sub.id
if isinstance(sub, ast.Name)
else sub.attr if isinstance(sub, ast.Attribute) else None
)
if name in UNSAFE_GUARDS and sub.lineno not in comments.guard_ok_lines:
yield Violation(
path, sub.lineno, "LIT007",
path,
sub.lineno,
"LIT007",
f"`{name}` narrowing predicate: the checker never verifies the body, so a "
f"wrong guard silently corrupts types; parse into a concrete type instead "
f"(suppress: `# guard-ok: <reason>`)",
)
# --------------------------------------------------------------------------- #
# Mutable-collection construction (LIT002)
# --------------------------------------------------------------------------- #
@ -395,7 +487,9 @@ def _construction_kind(node: ast.expr) -> str | None:
return None
def iter_construction_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]:
def iter_construction_violations(
path: Path, tree: ast.AST, comments: Comments
) -> Iterator[Violation]:
in_annotation = _annotation_node_ids(tree)
for node in ast.walk(tree):
if not isinstance(node, ast.expr) or id(node) in in_annotation:
@ -404,32 +498,37 @@ def iter_construction_violations(path: Path, tree: ast.AST, comments: Comments)
if kind is None or node.lineno in comments.mutable_ok_lines:
continue
yield Violation(
path, node.lineno, "LIT002",
path,
node.lineno,
"LIT002",
f"mutable {kind}: this builds a collection that can be grown or rewritten. "
f"Build it in one shot and freeze it -- a tuple/frozenset wrapping a generator "
f"(`tuple(f(x) for x in xs)`), a tuple literal, or a frozen dataclass / NamedTuple "
f"/ ReadOnly TypedDict (suppress: `# mutable-ok: <reason>`)",
)
# --------------------------------------------------------------------------- #
# Driver
# --------------------------------------------------------------------------- #
def check_file(path: Path) -> tuple[Violation, ...]:
try:
source = path.read_text(encoding="utf-8")
except (OSError, UnicodeDecodeError) as exc:
return (Violation(path, 0, "LIT000", f"could not read file: {exc}"),)
comments, violations = scan_comments(path, source)
try:
tree = ast.parse(source, filename=str(path))
except SyntaxError as exc:
return (*violations, Violation(path, exc.lineno or 0, "LIT000", f"syntax error: {exc.msg}"))
return (
*violations,
Violation(path, exc.lineno or 0, "LIT000", f"syntax error: {exc.msg}"),
)
return (
*violations,
*iter_annotation_violations(path, tree, comments),
@ -437,8 +536,8 @@ def check_file(path: Path) -> tuple[Violation, ...]:
*iter_guard_violations(path, tree, comments),
*iter_construction_violations(path, tree, comments),
)
def collect_paths(raw: Iterable[str]) -> Iterator[Path]:
for item in raw:
p = Path(item)
@ -446,24 +545,23 @@ def collect_paths(raw: Iterable[str]) -> Iterator[Path]:
yield from sorted(p.rglob("*.py"))
elif p.suffix == ".py":
yield p
def main(argv: Sequence[str]) -> int:
paths = tuple(a for a in argv if not a.startswith("-"))
if not paths:
print("usage: check_type_discipline.py <files-or-dirs>...", file=sys.stderr)
return 2
violations = sorted(v for path in collect_paths(paths) for v in check_file(path))
for v in violations:
print(v.render())
if violations:
print(f"\n{len(violations)} violation(s).", file=sys.stderr)
return 1
return 0
if __name__ == "__main__":
raise SystemExit(main(sys.argv[1:]))

View file

@ -42,8 +42,7 @@ from litellm.types.utils import CallTypes
PROBLEMS = [
{
"id": "has_close_elements",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List
def has_close_elements(numbers: List[float], threshold: float) -> bool:
@ -54,10 +53,8 @@ PROBLEMS = [
>>> has_close_elements([1.0, 2.8, 3.0, 4.0, 5.0, 2.0], 0.3)
True
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert has_close_elements([1.0, 2.0, 3.9, 4.0, 5.0, 2.2], 0.3) == True
assert has_close_elements([1.0, 2.0, 3.9, 4.0, 5.0, 2.2], 0.05) == False
assert has_close_elements([1.0, 2.0, 5.9, 4.0, 5.0], 0.95) == True
@ -65,13 +62,11 @@ PROBLEMS = [
assert has_close_elements([1.0, 2.0, 3.0, 4.0, 5.0], 2.0) == True
assert has_close_elements([], 0.5) == False
print("PASSED")
"""
),
"""),
},
{
"id": "separate_paren_groups",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List
def separate_paren_groups(paren_string: str) -> List[str]:
@ -82,22 +77,18 @@ PROBLEMS = [
>>> separate_paren_groups('( ) (( )) (( )( ))')
['()', '(())', '(()())']
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert separate_paren_groups('(()()) ((())) () ((())()())') == ['(()())', '((()))', '()', '((())()())']
assert separate_paren_groups('() (()) ((())) (((())))') == ['()', '(())', '((()))', '(((())))']
assert separate_paren_groups('(()(()))') == ['(()(()))']
assert separate_paren_groups('( ) (( )) (( )( ))') == ['()', '(())', '(()())']
print("PASSED")
"""
),
"""),
},
{
"id": "truncate_number",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
def truncate_number(number: float) -> float:
\"\"\"Given a positive floating point number, it can be decomposed into
an integer part (largest integer smaller than given number) and decimals
@ -106,21 +97,17 @@ PROBLEMS = [
>>> truncate_number(3.5)
0.5
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert truncate_number(3.5) == 0.5
assert abs(truncate_number(1.33) - 0.33) < 1e-6
assert abs(truncate_number(123.456) - 0.456) < 1e-6
print("PASSED")
"""
),
"""),
},
{
"id": "below_zero",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List
def below_zero(operations: List[int]) -> bool:
@ -132,10 +119,8 @@ PROBLEMS = [
>>> below_zero([1, 2, -4, 5])
True
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert below_zero([]) == False
assert below_zero([1, 2, -3, 1, 2, -3]) == False
assert below_zero([1, 2, -4, 5, 6]) == True
@ -143,13 +128,11 @@ PROBLEMS = [
assert below_zero([1, -1, 2, -2, 5, -5, 4, -5]) == True
assert below_zero([1, -2]) == True
print("PASSED")
"""
),
"""),
},
{
"id": "mean_absolute_deviation",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List
def mean_absolute_deviation(numbers: List[float]) -> float:
@ -161,21 +144,17 @@ PROBLEMS = [
>>> mean_absolute_deviation([1.0, 2.0, 3.0, 4.0])
1.0
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert abs(mean_absolute_deviation([1.0, 2.0, 3.0, 4.0]) - 1.0) < 1e-6
assert abs(mean_absolute_deviation([1.0, 2.0, 3.0, 4.0, 5.0]) - 1.2) < 1e-6
assert abs(mean_absolute_deviation([1.0, 1.0, 1.0, 1.0]) - 0.0) < 1e-6
print("PASSED")
"""
),
"""),
},
{
"id": "intersperse",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List
def intersperse(numbers: List[int], delimiter: int) -> List[int]:
@ -185,21 +164,17 @@ PROBLEMS = [
>>> intersperse([1, 2, 3], 4)
[1, 4, 2, 4, 3]
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert intersperse([], 7) == []
assert intersperse([5, 6, 3, 2], 8) == [5, 8, 6, 8, 3, 8, 2]
assert intersperse([2, 2, 2], 2) == [2, 2, 2, 2, 2]
print("PASSED")
"""
),
"""),
},
{
"id": "parse_nested_parens",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List
def parse_nested_parens(paren_string: str) -> List[int]:
@ -209,21 +184,17 @@ PROBLEMS = [
>>> parse_nested_parens('(()()) ((())) () ((())())')
[2, 3, 1, 3]
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert parse_nested_parens('(()()) ((())) () ((())())') == [2, 3, 1, 3]
assert parse_nested_parens('() (()) ((())) (((())))') == [1, 2, 3, 4]
assert parse_nested_parens('(()(())((())))') == [4]
print("PASSED")
"""
),
"""),
},
{
"id": "filter_by_substring",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List
def filter_by_substring(strings: List[str], substring: str) -> List[str]:
@ -233,22 +204,18 @@ PROBLEMS = [
>>> filter_by_substring(['abc', 'bacd', 'cde', 'array'], 'a')
['abc', 'bacd', 'array']
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert filter_by_substring([], 'john') == []
assert filter_by_substring(['xxx', 'asd', 'xxy', 'john doe', 'xxxuj', 'xxx'], 'xxx') == ['xxx', 'xxxuj', 'xxx']
assert filter_by_substring(['xxx', 'asd', 'aaber', 'john doe', 'xxxuj', 'xxx'], 'xx') == ['xxx', 'xxxuj', 'xxx']
assert filter_by_substring(['grunt', 'hierarchial', 'abc', 'hierarchial'], 'hi') == ['hierarchial', 'hierarchial']
print("PASSED")
"""
),
"""),
},
{
"id": "sum_product",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List, Tuple
def sum_product(numbers: List[int]) -> Tuple[int, int]:
@ -259,23 +226,19 @@ PROBLEMS = [
>>> sum_product([1, 2, 3, 4])
(10, 24)
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert sum_product([]) == (0, 1)
assert sum_product([1, 1, 1]) == (3, 1)
assert sum_product([100, 0]) == (100, 0)
assert sum_product([3, 5, 7]) == (15, 105)
assert sum_product([10]) == (10, 10)
print("PASSED")
"""
),
"""),
},
{
"id": "max_element",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List
def max_element(l: List[int]) -> int:
@ -285,21 +248,17 @@ PROBLEMS = [
>>> max_element([5, 3, -5, 2, -3, 3, 9, 0, 123, 1, -10])
123
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert max_element([1, 2, 3]) == 3
assert max_element([5, 3, -5, 2, -3, 3, 9, 0, 124, 1, -10]) == 124
assert max_element([-1, -2, -3]) == -1
print("PASSED")
"""
),
"""),
},
{
"id": "fizz_buzz",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
def fizz_buzz(n: int) -> int:
\"\"\"Return the number of times the digit 7 appears in integers less than n which are divisible by 11 or 13.
>>> fizz_buzz(50)
@ -309,10 +268,8 @@ PROBLEMS = [
>>> fizz_buzz(79)
3
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert fizz_buzz(50) == 0
assert fizz_buzz(78) == 2
assert fizz_buzz(79) == 3
@ -320,13 +277,11 @@ PROBLEMS = [
assert fizz_buzz(200) == 6
assert fizz_buzz(4000) == 192
print("PASSED")
"""
),
"""),
},
{
"id": "sort_by_binary_len",
"prompt": textwrap.dedent(
"""\
"prompt": textwrap.dedent("""\
from typing import List
def sort_array(arr: List[int]) -> List[int]:
@ -339,10 +294,8 @@ PROBLEMS = [
>>> sort_array([1, 0, 2, 3, 4])
[0, 1, 2, 4, 3]
\"\"\"
"""
),
"tests": textwrap.dedent(
"""\
"""),
"tests": textwrap.dedent("""\
assert sort_array([1, 5, 2, 3, 4]) == [1, 2, 4, 3, 5]
assert sort_array([-2, -3, -4, -5, -6]) == [-6, -5, -4, -3, -2]
assert sort_array([1, 0, 2, 3, 4]) == [0, 1, 2, 4, 3]
@ -350,8 +303,7 @@ PROBLEMS = [
assert sort_array([2, 5, 77, 4, 5, 3, 5, 7, 2, 3, 4]) == [2, 2, 4, 4, 3, 3, 5, 5, 5, 7, 77]
assert sort_array([3, 6, 44, 12, 32, 5]) == [32, 3, 5, 6, 12, 44]
print("PASSED")
"""
),
"""),
},
]
@ -360,8 +312,7 @@ PROBLEMS = [
# compressor to identify and drop them.
DISTRACTOR_SNIPPETS = [
# distractor 0 — database connection pool
textwrap.dedent(
"""\
textwrap.dedent("""\
# db_pool.py
import threading
from contextlib import contextmanager
@ -408,11 +359,9 @@ DISTRACTOR_SNIPPETS = [
for conn in self._pool:
conn.close()
self._pool.clear()
"""
),
"""),
# distractor 1 — HTTP retry logic
textwrap.dedent(
"""\
textwrap.dedent("""\
# http_retry.py
import time
import random
@ -456,11 +405,9 @@ DISTRACTOR_SNIPPETS = [
resp = requests.get(url, params=params, timeout=30)
resp.raise_for_status()
return resp.json()
"""
),
"""),
# distractor 2 — LRU cache implementation
textwrap.dedent(
"""\
textwrap.dedent("""\
# lru_cache.py
from collections import OrderedDict
from threading import RLock
@ -509,11 +456,9 @@ DISTRACTOR_SNIPPETS = [
def __contains__(self, key):
return key in self._cache
"""
),
"""),
# distractor 3 — CSV report generator
textwrap.dedent(
"""\
textwrap.dedent("""\
# report_gen.py
import csv
import io
@ -566,11 +511,9 @@ DISTRACTOR_SNIPPETS = [
except (ValueError, KeyError):
return False
return self.filter_rows(in_range)
"""
),
"""),
# distractor 4 — async task queue
textwrap.dedent(
"""\
textwrap.dedent("""\
# task_queue.py
import asyncio
import logging
@ -640,11 +583,9 @@ DISTRACTOR_SNIPPETS = [
async def shutdown(self):
for w in self._workers:
w.cancel()
"""
),
"""),
# distractor 5 — config parser with env var interpolation
textwrap.dedent(
"""\
textwrap.dedent("""\
# config_parser.py
import os
import re
@ -707,8 +648,7 @@ DISTRACTOR_SNIPPETS = [
if val is None:
raise ConfigError(f"Required config key missing: {key}")
return val
"""
),
"""),
]
@ -1036,7 +976,7 @@ def run_benchmark(
print(f" Avg total tokens: {base_agg['avg_total_tokens']}")
print(f" Avg latency: {base_agg['avg_latency_ms']}ms")
print(f"\n Compressed (litellm.compress → then call model):")
print("\n Compressed (litellm.compress → then call model):")
print(
f" Pass rate: {comp_agg['pass_rate']}% ({comp_agg['passed']}/{comp_agg['total']})"
)
@ -1054,7 +994,7 @@ def run_benchmark(
latency_diff = base_agg["avg_latency_ms"] - comp_agg["avg_latency_ms"]
pass_diff = comp_agg["pass_rate"] - base_agg["pass_rate"]
print(f"\n Delta (compressed vs baseline):")
print("\n Delta (compressed vs baseline):")
print(f" Token savings: {token_savings} tokens ({token_pct}%)")
print(f" Latency delta: {latency_diff:+.1f}ms")
print(f" Pass rate delta: {pass_diff:+.1f}%")

View file

@ -1,6 +1,6 @@
#!/usr/bin/env python3
"""
Bench LiteLLM_HealthCheckTable + PrismaClient
Bench LiteLLM_HealthCheckTable + PrismaClient
- set DATABASE_URL to your Postgres
- Run ```prisma generate``` to install prisma client before running test )
- This test writes to the default "public" database. Make sure to run cleanup after testing

View file

@ -360,7 +360,7 @@ class LiteLLMHealthCheckClient:
# Print detailed results for each model (matching Go output format)
print(f"\n{'='*60}", file=sys.stderr)
print(f"Starting health check queries\n", file=sys.stderr)
print("Starting health check queries\n", file=sys.stderr)
for model_id, result in results.items():
if result.get("healthy"):
@ -383,7 +383,7 @@ class LiteLLMHealthCheckClient:
print(f"---- {model_id} ----\n❌ ERROR: {error}\n\n", file=sys.stderr)
print(f"{'='*60}", file=sys.stderr)
print(f"Health Check Summary", file=sys.stderr)
print("Health Check Summary", file=sys.stderr)
print(f"{'='*60}", file=sys.stderr)
print(f"Total models: {len(results)}", file=sys.stderr)
print(f"Healthy: {healthy_count}", file=sys.stderr)

View file

@ -44,6 +44,7 @@ Notes
- `converse-stream` is still JSON-only placeholder (different inner event shapes).
- Use real (or any non-empty) AWS creds in the environment of the **proxy**; signing still runs.
"""
from __future__ import annotations
import argparse

View file

@ -38,7 +38,7 @@ class SlowHandler(BaseHTTPRequestHandler):
return None
return self.rfile.read(length)
def do_POST(self) -> None: # noqa: N802
def do_POST(self) -> None:
if self.path != "/cygnal/monitor":
self.send_error(404, "Not Found")
return

View file

@ -11,6 +11,7 @@ reader to write tests that kill the survivors.
Run after `mutmut run` and `mutmut export-cicd-stats`. Expects mutmut to be
invokable as `uv run --no-sync --with mutmut==<version> mutmut <subcommand>`.
"""
from __future__ import annotations
import ast
@ -373,9 +374,7 @@ def render(config: dict, survivors: list[str], stats: dict | None) -> str:
out.append("## Task")
out.append("")
out.append(
dedent(
"""\
out.append(dedent("""\
For each surviving mutant listed above, write a new test in the
existing test file (matching its conventions, fixtures, and naming
style) that:
@ -388,9 +387,7 @@ def render(config: dict, survivors: list[str], stats: dict | None) -> str:
which mutant numbers in the test name or docstring.
Do not modify the source file. Only add tests.
"""
).strip()
)
""").strip())
out.append("")
return "\n".join(out)

View file

@ -98,7 +98,9 @@ def base_counts(ref: str) -> dict:
# the body (or the `worktree add` itself) failed. rmtree is already best-effort.
subprocess.run(
["git", "worktree", "remove", "--force", str(worktree)],
cwd=REPO_ROOT, capture_output=True, text=True,
cwd=REPO_ROOT,
capture_output=True,
text=True,
)
shutil.rmtree(parent, ignore_errors=True)
@ -110,7 +112,8 @@ def over_ceiling(head: dict, budget: dict) -> frozenset:
comparison cannot change the verdict and the base worktree scan can be skipped.
"""
return frozenset(
rule for rule, spec in budget.items()
rule
for rule, spec in budget.items()
if head.get(rule, 0) > spec["baseline"] + spec["slack"]
)

View file

@ -25,7 +25,12 @@ from typing import AsyncIterator, Final
import uvicorn
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import JSONResponse, PlainTextResponse, Response, StreamingResponse
from starlette.responses import (
JSONResponse,
PlainTextResponse,
Response,
StreamingResponse,
)
from starlette.routing import Route
_CANNED_CONTENT: Final = "Hello! This is a mock response from the fake OpenAI endpoint."
@ -159,7 +164,14 @@ async def _text_completion_stream(model: str, with_usage: bool) -> AsyncIterator
"object": "text_completion",
"created": created,
"model": model,
"choices": [{"text": text, "index": 0, "logprobs": None, "finish_reason": finish_reason}],
"choices": [
{
"text": text,
"index": 0,
"logprobs": None,
"finish_reason": finish_reason,
}
],
}
yield f"data: {json.dumps(chunk(_CANNED_CONTENT, None))}\n\n"
@ -190,7 +202,10 @@ async def embeddings(request: Request) -> Response:
return JSONResponse(
{
"object": "list",
"data": [{"object": "embedding", "index": i, "embedding": [0.0] * 1536} for i in range(max(count, 1))],
"data": [
{"object": "embedding", "index": i, "embedding": [0.0] * 1536}
for i in range(max(count, 1))
],
"model": _requested_model(body),
"usage": {"prompt_tokens": 5, "total_tokens": 5},
}

View file

@ -19,7 +19,6 @@ from litellm.batches.batch_utils import (
from litellm.cost_calculator import batch_cost_calculator
from litellm.types.utils import Usage
# --- helpers ---

View file

@ -21,7 +21,6 @@ from unittest.mock import patch, MagicMock
import httpx
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
_BEDROCK_TEST_AWS_ENV = {
"AWS_ACCESS_KEY_ID": "test-access-key",
"AWS_SECRET_ACCESS_KEY": "test-secret-key",
@ -99,7 +98,9 @@ class _CaptureAsyncHTTPHandler(AsyncHTTPHandler):
raw = json if json is not None else (data if data is not None else content)
payload = raw if isinstance(raw, dict) else json_module.loads(raw)
job_name = payload["jobName"]
job_arn = f"arn:aws:bedrock:us-west-2:941277531214:model-invocation-job/{job_name}"
job_arn = (
f"arn:aws:bedrock:us-west-2:941277531214:model-invocation-job/{job_name}"
)
self.batch_jobs[job_arn] = {
"jobArn": job_arn,
"jobName": job_name,

View file

@ -20,7 +20,6 @@ sys.path.insert(0, os.path.abspath("../.."))
import litellm
SERVER_URL = "https://exampleopenaiendpoint-production-0ee2.up.railway.app/v1"

View file

@ -12,7 +12,6 @@ import litellm
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.token_counter import token_counter
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------

View file

@ -2,7 +2,6 @@ import ast
import os
from typing import List, Dict, Any
ALLOWED_FILE = os.path.normpath("litellm/_uuid.py")

View file

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

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

@ -801,7 +801,9 @@ async def test_list_projects_returns_timestamps():
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock, patch
from litellm_enterprise.proxy.management_endpoints.project_endpoints import list_projects
from litellm_enterprise.proxy.management_endpoints.project_endpoints import (
list_projects,
)
from litellm.proxy._types import LiteLLM_ProjectTable
now = datetime(2024, 1, 15, 12, 0, 0, tzinfo=timezone.utc)

View file

@ -14,7 +14,6 @@ from litellm.proxy.guardrails.guardrail_registry import (
)
from litellm.proxy.guardrails.guardrail_hooks.akto.akto import AktoGuardrail
# ---------------------------------------------------------------------------
# Registry tests
# ---------------------------------------------------------------------------

View file

@ -6,7 +6,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio

View file

@ -21,7 +21,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
ContentFilterCategoryConfig,
)
# Test cases: (sentence, expected_result, reason)
TEST_CASES = [
# ALWAYS BLOCK - Explicit prohibited practices (1-10)

View file

@ -23,7 +23,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
ContentFilterCategoryConfig,
)
# ── helpers ──────────────────────────────────────────────────────────────
POLICY_DIR = os.path.abspath(

View file

@ -28,7 +28,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
ContentFilterCategoryConfig,
)
# ── helpers ──────────────────────────────────────────────────────────────
POLICY_DIR = os.path.abspath(

View file

@ -7,7 +7,6 @@ import sys
import traceback
from unittest.mock import AsyncMock, MagicMock, patch
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path

View file

@ -6,7 +6,6 @@ import os
import sys
import traceback
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path

View file

@ -278,6 +278,7 @@ def test_embedding_via_proxy(proxy_url: str) -> None:
assert len(embedding) >= 64
assert all(isinstance(x, (int, float)) for x in embedding)
def test_model_list_advertises_oci_models(proxy_url: str) -> None:
"""The /v1/models registry advertises every OCI alias from the config."""
r = httpx.get(
@ -288,7 +289,9 @@ def test_model_list_advertises_oci_models(proxy_url: str) -> None:
assert r.status_code == 200, r.text
advertised = {row["id"] for row in r.json()["data"]}
for expected in CHAT_MODELS + ["oci-embed"]:
assert expected in advertised, f"{expected} missing from /v1/models: {advertised}"
assert (
expected in advertised
), f"{expected} missing from /v1/models: {advertised}"
def test_chat_completion_no_drop_params(proxy_url_no_drop_params: str) -> None:

View file

@ -11,7 +11,6 @@ import pytest
from litellm.types.utils import Embedding
from litellm.main import bedrock_embedding, embedding, EmbeddingResponse, Usage
_mock_model_id = (
"arn:aws:bedrock:us-east-1:123412341234:application-inference-profile/abc123123"
)

View file

@ -11,7 +11,6 @@ from litellm.llms.bedrock.common_utils import (
)
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
NOVA_ARN = "arn:aws:bedrock:us-east-1:123456789012:custom-model-deployment/a1b2c3d4e5f6"
NOVA_MODEL = f"bedrock/nova/{NOVA_ARN}"
NOVA2_MODEL = f"bedrock/nova-2/{NOVA_ARN}"

View file

@ -11,7 +11,6 @@ import pytest
from litellm import get_model_info
MODEL_NAME = "nvidia.nemotron-super-3-120b"

View file

@ -12,7 +12,6 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
BedrockConverseMessagesProcessor,
)
MODEL = "anthropic.claude-v2"
PROVIDER = "bedrock_converse"
@ -546,9 +545,9 @@ def test_bedrock_converse_sorts_text_before_tooluse_sync():
tool_indices = [i for i, b in enumerate(content) if "toolUse" in b]
# All text blocks must come before all toolUse blocks
assert max(text_indices) < min(tool_indices), (
f"text blocks at {text_indices} should all precede toolUse blocks at {tool_indices}"
)
assert max(text_indices) < min(
tool_indices
), f"text blocks at {text_indices} should all precede toolUse blocks at {tool_indices}"
@pytest.mark.asyncio
@ -567,9 +566,9 @@ async def test_bedrock_converse_sorts_text_before_tooluse_async():
text_indices = [i for i, b in enumerate(content) if "text" in b]
tool_indices = [i for i, b in enumerate(content) if "toolUse" in b]
assert max(text_indices) < min(tool_indices), (
f"text blocks at {text_indices} should all precede toolUse blocks at {tool_indices}"
)
assert max(text_indices) < min(
tool_indices
), f"text blocks at {text_indices} should all precede toolUse blocks at {tool_indices}"
@pytest.mark.asyncio
@ -577,7 +576,9 @@ async def test_bedrock_converse_content_ordering_sync_async_parity():
"""Sync and async paths should produce identical content block ordering."""
messages = _make_tooluse_before_text_messages()
sync_result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
messages, MODEL, PROVIDER
async_result = (
await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
messages, MODEL, PROVIDER
)
)
assert sync_result == async_result

View file

@ -30,6 +30,7 @@ def _attrify(d: dict):
None)` (et al), which returns None for plain dicts — that would silently
skip the row.
"""
class _AttrDict(dict):
def __getattr__(self, k):
try:
@ -120,7 +121,9 @@ async def test_reset_budget_keys_partial_failure():
key1, key2, key3, key4, key5, key6 = (
_attrify(k) for k in [key1, key2, key3, key4, key5, key6]
)
prisma_client.get_data = AsyncMock(return_value=[key1, key2, key3, key4, key5, key6])
prisma_client.get_data = AsyncMock(
return_value=[key1, key2, key3, key4, key5, key6]
)
async def fake_reset_key(key, current_time):
if key["id"] == "key1":
@ -207,7 +210,9 @@ async def test_reset_budget_users_partial_failure():
user1, user2, user3, user4, user5, user6 = (
_attrify(u) for u in [user1, user2, user3, user4, user5, user6]
)
prisma_client.get_data = AsyncMock(return_value=[user1, user2, user3, user4, user5, user6])
prisma_client.get_data = AsyncMock(
return_value=[user1, user2, user3, user4, user5, user6]
)
async def fake_reset_user(user, current_time):
if user["id"] == "user1":

View file

@ -2490,9 +2490,7 @@ def test_get_base_model_from_metadata():
"litellm_params": {"base_model": "azure/gpt-5-mini"}
}
result = _get_base_model_from_metadata(model_call_details_with_direct_base_model)
assert (
result == "azure/gpt-5-mini"
), f"Expected 'azure/gpt-5-mini', got {result}"
assert result == "azure/gpt-5-mini", f"Expected 'azure/gpt-5-mini', got {result}"
# Test 4: metadata takes precedence over litellm_metadata
model_call_details_with_both = {

View file

@ -12,7 +12,6 @@ from litellm.responses.litellm_completion_transformation.transformation import (
)
from litellm.types.utils import ModelResponse
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.integrations.custom_logger import CustomLogger

View file

@ -2,7 +2,7 @@
Test to reproduce and verify fix for Anthropic tool_result issue with empty call_id.
This test reproduces the exact error:
"messages.0.content.0: unexpected `tool_use_id` found in `tool_result` blocks: tool_use_id.
"messages.0.content.0: unexpected `tool_use_id` found in `tool_result` blocks: tool_use_id.
Each `tool_result` block must have a corresponding `tool_use` block in the previous message."
The issue occurs when:

View file

@ -2,12 +2,12 @@
Unit tests for BaseResponsesAPIStreamingIterator
Tests core functionality including:
1. Processing chunks and handling ResponseCompletedEvent
1. Processing chunks and handling ResponseCompletedEvent
2. Ensuring _update_responses_api_response_id_with_model_id is called for final chunk
3. Verifying ID update is NOT called for non-final chunks (delta events)
4. Edge case handling for invalid JSON, empty chunks, and [DONE] markers
These tests ensure the streaming iterator correctly processes response chunks
These tests ensure the streaming iterator correctly processes response chunks
and applies model ID updates only to completed responses, as required for proper
response tracking and logging.
"""

View file

@ -387,9 +387,7 @@ def test_process_chunk_completed_response_updates_id_and_usage_cost(monkeypatch)
# Chunk must include a top-level "response" key so BaseResponsesAPIStreamingIterator
# runs _update_responses_api_response_id_with_model_id (see streaming_iterator.py).
event = iterator._process_chunk(
json.dumps(
{"type": "response.completed", "response": {"id": "resp_live"}}
)
json.dumps({"type": "response.completed", "response": {"id": "resp_live"}})
)
finally:
litellm.include_cost_in_streaming_usage = original_include_cost

View file

@ -22,10 +22,8 @@ sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm import completion
# Large document for caching tests (needs 1024+ tokens for Claude models)
LARGE_DOCUMENT_FOR_CACHING = (
"""
LARGE_DOCUMENT_FOR_CACHING = """
This is a comprehensive legal agreement between Party A and Party B.
ARTICLE 1: DEFINITIONS
@ -77,9 +75,7 @@ ARTICLE 9: GENERAL PROVISIONS
9.5 Waiver of any provision shall not constitute ongoing waiver.
IN WITNESS WHEREOF, the parties have executed this Agreement.
"""
* 8
) # Repeat to ensure we have enough tokens (need 1024+ for Claude models)
""" * 8 # Repeat to ensure we have enough tokens (need 1024+ for Claude models)
class TestBedrockAnthropicPromptCachingRegression:

View file

@ -1,6 +1,7 @@
"""
Tests for Crusoe provider integration
"""
import os
from unittest import mock
@ -96,9 +97,9 @@ def test_crusoe_models_configuration():
for model in crusoe_models:
model_info = get_model_info(model)
assert model_info is not None, f"Model info not found for {model}"
assert model_info.get("litellm_provider") == "crusoe", (
f"{model} should have crusoe as provider"
)
assert (
model_info.get("litellm_provider") == "crusoe"
), f"{model} should have crusoe as provider"
assert model_info.get("mode") == "chat", f"{model} should be in chat mode"
finally:
litellm.model_cost = original_model_cost

View file

@ -191,7 +191,11 @@ def test_deepseek_fill_reasoning_content_multiturn():
# Case 1: assistant message already has reasoning_content — should be left as-is
messages_with_rc = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi", "reasoning_content": "I thought about it"},
{
"role": "assistant",
"content": "Hi",
"reasoning_content": "I thought about it",
},
{"role": "user", "content": "Follow up"},
]
result = config._fill_reasoning_content(messages_with_rc)
@ -259,9 +263,9 @@ def test_deepseek_fill_reasoning_content_guard_in_transform_request():
litellm_params={},
headers={},
)
assert result["messages"][1].get("reasoning_content") == " ", (
"reasoning_content should be injected when thinking is enabled"
)
assert (
result["messages"][1].get("reasoning_content") == " "
), "reasoning_content should be injected when thinking is enabled"
# Case 2: reasoning model + thinking NOT in optional_params -> no injection
result = config.transform_request(
@ -271,9 +275,9 @@ def test_deepseek_fill_reasoning_content_guard_in_transform_request():
litellm_params={},
headers={},
)
assert "reasoning_content" not in result["messages"][1], (
"reasoning_content should not be injected when thinking is not enabled"
)
assert (
"reasoning_content" not in result["messages"][1]
), "reasoning_content should not be injected when thinking is not enabled"
# Case 3: non-reasoning model + thinking enabled -> no injection
result = config.transform_request(
@ -283,6 +287,6 @@ def test_deepseek_fill_reasoning_content_guard_in_transform_request():
litellm_params={},
headers={},
)
assert "reasoning_content" not in result["messages"][1], (
"reasoning_content should not be injected for non-reasoning models"
)
assert (
"reasoning_content" not in result["messages"][1]
), "reasoning_content should not be injected for non-reasoning models"

View file

@ -16,7 +16,6 @@ import litellm
from litellm import completion
import json
GEMINI_3_IMAGE_SIZE_MAPPINGS = [
("512x512", "1:1", "512"),
("1024x1024", "1:1", "1K"),
@ -557,9 +556,10 @@ def test_gemini_image_generation_openai_size_auto_uses_google_defaults(size: str
map_openai_size_to_gemini_image_config,
)
assert map_openai_size_to_gemini_image_config(
size, "gemini-3-pro-image-preview"
) is None
assert (
map_openai_size_to_gemini_image_config(size, "gemini-3-pro-image-preview")
is None
)
def test_gemini_imagen_models_use_predict_endpoint():

View file

@ -1,7 +1,7 @@
"""
Test for Gemini image generation usage metadata extraction.
This test verifies the fix for issue #18323 where image_generation()
This test verifies the fix for issue #18323 where image_generation()
was returning usage=0 while completion() returned proper token usage.
"""

View file

@ -1683,7 +1683,9 @@ class TestMissingChoicesGuard:
assert "no 'choices'" in exc_info.value.message
def test_convert_to_model_response_object_stream_true_no_choices_raises_api_error(self):
def test_convert_to_model_response_object_stream_true_no_choices_raises_api_error(
self,
):
"""Missing choices via stream=True path raises APIError when generator is consumed."""
from litellm.exceptions import APIError
@ -2471,6 +2473,13 @@ class TestConvertToModelResponseObjectCompletion:
def test_model_response_none_raises(self):
with pytest.raises(Exception):
convert_to_model_response_object(
response_object={"choices": [{"message": {"content": "hi", "role": "assistant"}, "finish_reason": "stop"}]},
response_object={
"choices": [
{
"message": {"content": "hi", "role": "assistant"},
"finish_reason": "stop",
}
]
},
model_response_object=None,
)

View file

@ -2,7 +2,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio

View file

@ -38,7 +38,6 @@ from litellm.llms.vertex_ai.gemini.transformation import (
)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
litellm.num_retries = 3
litellm.cache = None
user_message = "Write a short poem about the sky"
@ -1104,9 +1103,7 @@ def vertex_httpx_mock_post_valid_response(*args, **kwargs):
{
"content": {
"role": "model",
"parts": [
{
"text": """{
"parts": [{"text": """{
"recipes": [
{"recipe_name": "Chocolate Chip Cookies"},
{"recipe_name": "Oatmeal Raisin Cookies"},
@ -1114,9 +1111,7 @@ def vertex_httpx_mock_post_valid_response(*args, **kwargs):
{"recipe_name": "Sugar Cookies"},
{"recipe_name": "Snickerdoodles"}
]
}"""
}
],
}"""}],
},
"finishReason": "STOP",
"safetyRatings": [

View file

@ -4,7 +4,7 @@ Test for preset_cache_key multiple values bug fix.
This test verifies that get_cache_key doesn't raise TypeError when kwargs
already contains preset_cache_key.
Issue: When get_cache_key(**kwargs) is called with kwargs containing
Issue: When get_cache_key(**kwargs) is called with kwargs containing
preset_cache_key, the call to _set_preset_cache_key_in_kwargs() would fail with:
TypeError: got multiple values for keyword argument 'preset_cache_key'
"""

View file

@ -450,12 +450,9 @@ def c():
litellm.enable_caching_on_provider_specific_optional_params = False
embedding_large_text = (
"""
embedding_large_text = """
small text
"""
* 5
)
""" * 5
# # test_caching_with_models()

View file

@ -1666,15 +1666,13 @@ def custom_callback(
#################################################
print(
f"""
print(f"""
Model: {model},
Messages: {messages},
User: {user},
Seed: {kwargs["seed"]},
temperature: {kwargs["temperature"]},
"""
)
""")
assert kwargs["user"] == "ishaans app"
assert kwargs["model"] == "gpt-3.5-turbo-1106"

View file

@ -93,8 +93,7 @@ class testCustomCallbackProxy(CustomLogger):
print("\n\n in custom callback vars my custom logger, ", vars(my_custom_logger))
print(
f"""
print(f"""
Model: {model},
Messages: {messages},
User: {user},
@ -102,8 +101,7 @@ class testCustomCallbackProxy(CustomLogger):
Cost: {cost},
Response: {response}
Proxy Metadata: {metadata}
"""
)
""")
return
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):

View file

@ -267,7 +267,6 @@ def test_aaparallel_function_call_with_anthropic_thinking(model):
from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
_PARALLEL_TOOL_HISTORY_MESSAGES = [
{
"role": "user",

View file

@ -21,7 +21,6 @@ verbose_logger.setLevel(logging.DEBUG)
litellm.set_verbose = True
import time
# test_langsmith_logging()

View file

@ -16,7 +16,6 @@ import litellm
from litellm import Router
from litellm.caching import RedisCache, RedisClusterCache
## Scenarios
## 1. 2 models - openai + azure - 1 model group "gpt-3.5-turbo",
## 2. 2 models - openai, azure - 2 diff model groups, 1 caching group

View file

@ -2,7 +2,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio

View file

@ -2,7 +2,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio

View file

@ -6,7 +6,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio

View file

@ -2,7 +2,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio

View file

@ -2,7 +2,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio

View file

@ -2,7 +2,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio

Some files were not shown because too many files have changed in this diff Show more