diff --git a/.github/scripts/close_low_quality_prs.py b/.github/scripts/close_low_quality_prs.py
index 7b9bbb579e3..c6d275b378b 100644
--- a/.github/scripts/close_low_quality_prs.py
+++ b/.github/scripts/close_low_quality_prs.py
@@ -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,
diff --git a/.github/scripts/triage_rollout_heads_up.py b/.github/scripts/triage_rollout_heads_up.py
index a5dedb1c9e7..7d8117c22eb 100644
--- a/.github/scripts/triage_rollout_heads_up.py
+++ b/.github/scripts/triage_rollout_heads_up.py
@@ -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,
diff --git a/.github/scripts/triage_with_llm.py b/.github/scripts/triage_with_llm.py
index d2536058e01..47400213219 100644
--- a/.github/scripts/triage_with_llm.py
+++ b/.github/scripts/triage_with_llm.py
@@ -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,
diff --git a/.github/workflows/auto_update_price_and_context_window_file.py b/.github/workflows/auto_update_price_and_context_window_file.py
index 461d8d347d9..b92a0568e37 100644
--- a/.github/workflows/auto_update_price_and_context_window_file.py
+++ b/.github/workflows/auto_update_price_and_context_window_file.py
@@ -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()
diff --git a/.github/workflows/run_llm_translation_tests.py b/.github/workflows/run_llm_translation_tests.py
index 3f3a70efe92..ff8252393a0 100644
--- a/.github/workflows/run_llm_translation_tests.py
+++ b/.github/workflows/run_llm_translation_tests.py
@@ -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("\nFailed Tests
\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 \n\n")
-
+
# Show errors (if any)
- if tests_by_status['ERROR']:
+ if tests_by_status["ERROR"]:
f.write("\nError Tests
\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 \n\n")
-
+
# Show passed tests in collapsible section
- if tests_by_status['PASSED']:
+ if tests_by_status["PASSED"]:
f.write("\nPassed Tests
\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 \n\n")
-
+
# Show skipped tests (if any)
- if tests_by_status['SKIPPED']:
+ if tests_by_status["SKIPPED"]:
f.write("\nSkipped Tests
\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 \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)
diff --git a/ci_cd/run_migration.py b/ci_cd/run_migration.py
index feec4046ee1..1eccb19cd05 100644
--- a/ci_cd/run_migration.py
+++ b/ci_cd/run_migration.py
@@ -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 if your", file=out
+ " - Fix the above and re-run, OR pass --base-branch if your", file=out
)
print(
f" base branch is not '{base_branch}', OR pass --skip-freshness-check",
diff --git a/cookbook/LiteLLM_CometAPI.ipynb b/cookbook/LiteLLM_CometAPI.ipynb
index 0a7ab581ae3..f1ee397c28b 100644
--- a/cookbook/LiteLLM_CometAPI.ipynb
+++ b/cookbook/LiteLLM_CometAPI.ipynb
@@ -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",
diff --git a/cookbook/benchmark/benchmark.py b/cookbook/benchmark/benchmark.py
index b38d185a166..2b3ddacad80 100644
--- a/cookbook/benchmark/benchmark.py
+++ b/cookbook/benchmark/benchmark.py
@@ -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"]
diff --git a/cookbook/google_adk_litellm_tutorial.ipynb b/cookbook/google_adk_litellm_tutorial.ipynb
index 27914edbba8..448098d9c08 100644
--- a/cookbook/google_adk_litellm_tutorial.ipynb
+++ b/cookbook/google_adk_litellm_tutorial.ipynb
@@ -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",
diff --git a/cookbook/litellm_proxy_server/braintrust_prompt_wrapper_server.py b/cookbook/litellm_proxy_server/braintrust_prompt_wrapper_server.py
index 6379314c5b6..af13842db8c 100644
--- a/cookbook/litellm_proxy_server/braintrust_prompt_wrapper_server.py
+++ b/cookbook/litellm_proxy_server/braintrust_prompt_wrapper_server.py
@@ -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)
diff --git a/cookbook/litellm_proxy_server/cli_token_usage.py b/cookbook/litellm_proxy_server/cli_token_usage.py
index 6306970cdde..f3ce878b645 100644
--- a/cookbook/litellm_proxy_server/cli_token_usage.py
+++ b/cookbook/litellm_proxy_server/cli_token_usage.py
@@ -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/"
diff --git a/cookbook/livekit_agent_sdk/main.py b/cookbook/livekit_agent_sdk/main.py
index c68e5534ea8..f0f625a03a5 100644
--- a/cookbook/livekit_agent_sdk/main.py
+++ b/cookbook/livekit_agent_sdk/main.py
@@ -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__":
diff --git a/cookbook/misc/migrate_proxy_config.py b/cookbook/misc/migrate_proxy_config.py
index 31c3f32c08a..dc039070a5d 100644
--- a/cookbook/misc/migrate_proxy_config.py
+++ b/cookbook/misc/migrate_proxy_config.py
@@ -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 `/model/new` for each model
"""
diff --git a/cookbook/mock_guardrail_server/mock_bedrock_guardrail_server.py b/cookbook/mock_guardrail_server/mock_bedrock_guardrail_server.py
index 7bf9cc32484..327f70b16b7 100644
--- a/cookbook/mock_guardrail_server/mock_bedrock_guardrail_server.py
+++ b/cookbook/mock_guardrail_server/mock_bedrock_guardrail_server.py
@@ -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)
diff --git a/cookbook/mock_prompt_management_server/mock_prompt_management_server.py b/cookbook/mock_prompt_management_server/mock_prompt_management_server.py
index 295a96e12a9..f390d40c007 100644
--- a/cookbook/mock_prompt_management_server/mock_prompt_management_server.py
+++ b/cookbook/mock_prompt_management_server/mock_prompt_management_server.py
@@ -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}")
diff --git a/cookbook/veo_video_generation.py b/cookbook/veo_video_generation.py
index 4df2d946a01..2e447fdb6a0 100644
--- a/cookbook/veo_video_generation.py
+++ b/cookbook/veo_video_generation.py
@@ -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
diff --git a/db_scripts/migrate_keys.py b/db_scripts/migrate_keys.py
index 5c940e069b3..573d81a0365 100644
--- a/db_scripts/migrate_keys.py
+++ b/db_scripts/migrate_keys.py
@@ -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")
diff --git a/enterprise/enterprise_hooks/aporia_ai.py b/enterprise/enterprise_hooks/aporia_ai.py
index 28b49bfce21..6f3fa8c0c3e 100644
--- a/enterprise/enterprise_hooks/aporia_ai.py
+++ b/enterprise/enterprise_hooks/aporia_ai.py
@@ -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 (
diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py b/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py
index 8824f4c02de..7353b995d2a 100644
--- a/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py
+++ b/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py
@@ -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
\ No newline at end of file
+ verbose_logger.warning(
+ f"Disabling callbacks using request headers is an enterprise feature. {CommonProxyErrors.not_premium_user.value}"
+ )
+ return False
diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/llama_guard.py b/enterprise/litellm_enterprise/enterprise_callbacks/llama_guard.py
index 5e1aebdbdfb..b659b628a75 100644
--- a/enterprise/litellm_enterprise/enterprise_callbacks/llama_guard.py
+++ b/enterprise/litellm_enterprise/enterprise_callbacks/llama_guard.py
@@ -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
diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py b/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py
index ad8aabf77b6..d5384ea5a63 100644
--- a/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py
+++ b/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py
@@ -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
diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py
index 89c3b854686..4fb6679a6eb 100644
--- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py
+++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py
@@ -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,
diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py
index 8fc2d66d531..a1e8def2bb2 100644
--- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py
+++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py
@@ -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
\ No newline at end of file
+ return
diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/smtp_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/smtp_email.py
index 8efdaf231b7..8e4dbde437b 100644
--- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/smtp_email.py
+++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/smtp_email.py
@@ -1,6 +1,7 @@
"""
This is the litellm SMTP email integration
"""
+
import asyncio
from typing import List
diff --git a/enterprise/litellm_enterprise/litellm_core_utils/litellm_logging.py b/enterprise/litellm_enterprise/litellm_core_utils/litellm_logging.py
index 44ba0063ffe..24941e90ab8 100644
--- a/enterprise/litellm_enterprise/litellm_core_utils/litellm_logging.py
+++ b/enterprise/litellm_enterprise/litellm_core_utils/litellm_logging.py
@@ -1,6 +1,7 @@
"""
Enterprise specific logging utils
"""
+
from litellm.litellm_core_utils.litellm_logging import StandardLoggingMetadata
diff --git a/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py b/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py
index 18ac29b9781..d100e38c1c9 100644
--- a/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py
+++ b/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py
@@ -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,
diff --git a/enterprise/litellm_enterprise/proxy/auth/__init__.py b/enterprise/litellm_enterprise/proxy/auth/__init__.py
index f67826ca7fa..dc70b57ab55 100644
--- a/enterprise/litellm_enterprise/proxy/auth/__init__.py
+++ b/enterprise/litellm_enterprise/proxy/auth/__init__.py
@@ -7,4 +7,4 @@ including custom SSO handlers and advanced authentication features.
from .custom_sso_handler import EnterpriseCustomSSOHandler
-__all__ = ["EnterpriseCustomSSOHandler"]
\ No newline at end of file
+__all__ = ["EnterpriseCustomSSOHandler"]
diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
index ee7745d0add..705a682f298 100644
--- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
+++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
@@ -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:
diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py
index dc0168683c8..8b2d15c1578 100644
--- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py
+++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py
@@ -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"
)
-
diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py
index 254d816039c..d50b77010ad 100644
--- a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py
+++ b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py
@@ -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,;target_model_names,;resource_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
"""
diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py
index 2f53f9e9281..48b6dd76348 100644
--- a/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py
+++ b/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py
@@ -2,7 +2,6 @@
Enterprise internal user management endpoints
"""
-
from fastapi import APIRouter, Depends, HTTPException
from litellm.proxy._types import UserAPIKeyAuth
diff --git a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py
index 5e799599862..19dbcb39cae 100644
--- a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py
+++ b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py
@@ -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
diff --git a/enterprise/litellm_enterprise/types/enterprise_callbacks/send_emails.py b/enterprise/litellm_enterprise/types/enterprise_callbacks/send_emails.py
index 380b0a6facb..d9d5a989abb 100644
--- a/enterprise/litellm_enterprise/types/enterprise_callbacks/send_emails.py
+++ b/enterprise/litellm_enterprise/types/enterprise_callbacks/send_emails.py
@@ -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()
\ No newline at end of file
+ return cls().to_dict()
diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py
index 144bb4c473f..951088e76c1 100644
--- a/gateway/routes/allowlist.py
+++ b/gateway/routes/allowlist.py
@@ -106,7 +106,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
# Health & ops
"/health",
"/metrics",
- "/watsonx"
+ "/watsonx",
)
GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(
diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py
index f1ad3708236..ad090ed7ff5 100644
--- a/litellm/llms/mistral/chat/transformation.py
+++ b/litellm/llms/mistral/chat/transformation.py
@@ -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
diff --git a/scripts/adaptive_router_demo/eval.py b/scripts/adaptive_router_demo/eval.py
index b02e4a37d31..66b91c71932 100644
--- a/scripts/adaptive_router_demo/eval.py
+++ b/scripts/adaptive_router_demo/eval.py
@@ -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))
diff --git a/scripts/adaptive_router_demo/traffic.py b/scripts/adaptive_router_demo/traffic.py
index eae5506eaee..dd64d95c90c 100644
--- a/scripts/adaptive_router_demo/traffic.py
+++ b/scripts/adaptive_router_demo/traffic.py
@@ -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
diff --git a/scripts/benchmark_chat_completions_perf.py b/scripts/benchmark_chat_completions_perf.py
index 2c211f674fe..2baca734dcb 100644
--- a/scripts/benchmark_chat_completions_perf.py
+++ b/scripts/benchmark_chat_completions_perf.py
@@ -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"
diff --git a/scripts/benchmark_mock.py b/scripts/benchmark_mock.py
index 55dbb1d4134..27f076e8485 100644
--- a/scripts/benchmark_mock.py
+++ b/scripts/benchmark_mock.py
@@ -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}%")
diff --git a/scripts/benchmark_model_response_creator.py b/scripts/benchmark_model_response_creator.py
index 881870d3854..8c595516fa4 100644
--- a/scripts/benchmark_model_response_creator.py
+++ b/scripts/benchmark_model_response_creator.py
@@ -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()
diff --git a/scripts/benchmark_proxy_vs_provider.py b/scripts/benchmark_proxy_vs_provider.py
index 6196580b230..d8c3a0a8b95 100755
--- a/scripts/benchmark_proxy_vs_provider.py
+++ b/scripts/benchmark_proxy_vs_provider.py
@@ -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)
diff --git a/scripts/benchmark_streaming_chunk_overhead.py b/scripts/benchmark_streaming_chunk_overhead.py
index 11fbea6a6a3..ab192a0a3e5 100644
--- a/scripts/benchmark_streaming_chunk_overhead.py
+++ b/scripts/benchmark_streaming_chunk_overhead.py
@@ -223,6 +223,7 @@ async def drive_async(
# Repeat × take-min runner
# ---------------------------------------------------------------------------
+
@dataclass
class Result:
label: str
diff --git a/scripts/budget_ratchet_check.py b/scripts/budget_ratchet_check.py
index 861d65489e8..a1e67fcba13 100644
--- a/scripts/budget_ratchet_check.py
+++ b/scripts/budget_ratchet_check.py
@@ -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(
diff --git a/scripts/check_type_discipline.py b/scripts/check_type_discipline.py
index 6e152541863..a028bdfdd1c 100644
--- a/scripts/check_type_discipline.py
+++ b/scripts/check_type_discipline.py
@@ -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:\s*(?P[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}: `")
-
+ yield Violation(
+ path,
+ line_no,
+ "LIT005",
+ f"{token} requires a reason: `# {token}: `",
+ )
+
m = NOQA_RE.search(text)
if m:
if not m.group("codes"):
- yield Violation(path, line_no, "LIT003", "noqa requires rule codes: `# noqa: XXX123 # `")
+ yield Violation(
+ path,
+ line_no,
+ "LIT003",
+ "noqa requires rule codes: `# noqa: XXX123 # `",
+ )
elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN:
- yield Violation(path, line_no, "LIT003", "noqa requires a reason: `# noqa: XXX123 # `")
-
+ yield Violation(
+ path,
+ line_no,
+ "LIT003",
+ "noqa requires a reason: `# noqa: XXX123 # `",
+ )
+
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] # `")
+ yield Violation(
+ path,
+ line_no,
+ "LIT004",
+ "ignore requires codes: `# pyright: ignore[ruleName] # `",
+ )
elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN:
- yield Violation(path, line_no, "LIT004",
- "ignore requires a reason: `# pyright: ignore[ruleName] # `")
-
-
+ yield Violation(
+ path,
+ line_no,
+ "LIT004",
+ "ignore requires a reason: `# pyright: ignore[ruleName] # `",
+ )
+
+
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 ""
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: `)",
)
-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: `)",
)
-
-
+
+
# --------------------------------------------------------------------------- #
# 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: `)",
)
-
-
+
+
# --------------------------------------------------------------------------- #
# 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 ...", 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:]))
-
\ No newline at end of file
diff --git a/scripts/eval_compression.py b/scripts/eval_compression.py
index a169cc02d74..4ba82002ee4 100644
--- a/scripts/eval_compression.py
+++ b/scripts/eval_compression.py
@@ -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}%")
diff --git a/scripts/health_check/benchmark_get_all_latest_health_checks.py b/scripts/health_check/benchmark_get_all_latest_health_checks.py
index 45845554c86..5d0957d4f7d 100644
--- a/scripts/health_check/benchmark_get_all_latest_health_checks.py
+++ b/scripts/health_check/benchmark_get_all_latest_health_checks.py
@@ -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
diff --git a/scripts/health_check/health_check_client.py b/scripts/health_check/health_check_client.py
index 9ef8b934961..a04143f8ddc 100644
--- a/scripts/health_check/health_check_client.py
+++ b/scripts/health_check/health_check_client.py
@@ -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)
diff --git a/scripts/mock_bedrock_passthrough_target.py b/scripts/mock_bedrock_passthrough_target.py
index e993cd99bde..73462aa6dfc 100644
--- a/scripts/mock_bedrock_passthrough_target.py
+++ b/scripts/mock_bedrock_passthrough_target.py
@@ -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
diff --git a/scripts/mock_grayswan_timeout_server.py b/scripts/mock_grayswan_timeout_server.py
index bc40e13dfa5..b4945937a11 100644
--- a/scripts/mock_grayswan_timeout_server.py
+++ b/scripts/mock_grayswan_timeout_server.py
@@ -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
diff --git a/scripts/mutation_report.py b/scripts/mutation_report.py
index a606e3f71cf..22641d88499 100644
--- a/scripts/mutation_report.py
+++ b/scripts/mutation_report.py
@@ -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== mutmut `.
"""
+
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)
diff --git a/scripts/type_discipline_gate.py b/scripts/type_discipline_gate.py
index c111486e56a..5ad9b96a340 100644
--- a/scripts/type_discipline_gate.py
+++ b/scripts/type_discipline_gate.py
@@ -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"]
)
diff --git a/tests/_fake_openai_endpoint_server.py b/tests/_fake_openai_endpoint_server.py
index 409f569070b..1b2429b127d 100644
--- a/tests/_fake_openai_endpoint_server.py
+++ b/tests/_fake_openai_endpoint_server.py
@@ -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},
}
diff --git a/tests/batches_tests/test_batch_custom_pricing.py b/tests/batches_tests/test_batch_custom_pricing.py
index cb2ca385ffc..3d3590105c9 100644
--- a/tests/batches_tests/test_batch_custom_pricing.py
+++ b/tests/batches_tests/test_batch_custom_pricing.py
@@ -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 ---
diff --git a/tests/batches_tests/test_bedrock_files_and_batches.py b/tests/batches_tests/test_bedrock_files_and_batches.py
index 431d5a2a60c..e7f7ef89565 100644
--- a/tests/batches_tests/test_bedrock_files_and_batches.py
+++ b/tests/batches_tests/test_bedrock_files_and_batches.py
@@ -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,
diff --git a/tests/batches_tests/test_hosted_vllm_batches_and_files.py b/tests/batches_tests/test_hosted_vllm_batches_and_files.py
index c7a25c71c53..b7f21f6cab2 100644
--- a/tests/batches_tests/test_hosted_vllm_batches_and_files.py
+++ b/tests/batches_tests/test_hosted_vllm_batches_and_files.py
@@ -20,7 +20,6 @@ sys.path.insert(0, os.path.abspath("../.."))
import litellm
-
SERVER_URL = "https://exampleopenaiendpoint-production-0ee2.up.railway.app/v1"
diff --git a/tests/benchmarks/test_benchmarks.py b/tests/benchmarks/test_benchmarks.py
index 123dad93e11..d5ade58980a 100644
--- a/tests/benchmarks/test_benchmarks.py
+++ b/tests/benchmarks/test_benchmarks.py
@@ -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
# ---------------------------------------------------------------------------
diff --git a/tests/code_coverage_tests/check_fastuuid_usage.py b/tests/code_coverage_tests/check_fastuuid_usage.py
index e0433454371..bb3176b57a3 100644
--- a/tests/code_coverage_tests/check_fastuuid_usage.py
+++ b/tests/code_coverage_tests/check_fastuuid_usage.py
@@ -2,7 +2,6 @@ import ast
import os
from typing import List, Dict, Any
-
ALLOWED_FILE = os.path.normpath("litellm/_uuid.py")
diff --git a/tests/enterprise/conftest.py b/tests/enterprise/conftest.py
index 0365bbbcfa0..524ab85b938 100644
--- a/tests/enterprise/conftest.py
+++ b/tests/enterprise/conftest.py
@@ -23,8 +23,6 @@ def event_loop():
loop.close()
-
-
@pytest.fixture(scope="function", autouse=True)
def setup_and_teardown():
"""
diff --git a/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py b/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py
index ebea96e2152..23d9d8db972 100644
--- a/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py
+++ b/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py
@@ -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")
diff --git a/tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py b/tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py
index a45df5df008..ab5c5576625 100644
--- a/tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py
+++ b/tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py
@@ -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(
diff --git a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py
index c55b66b402b..c65e6359100 100644
--- a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py
+++ b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py
@@ -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)
diff --git a/tests/guardrails_tests/test_akto_guardrails.py b/tests/guardrails_tests/test_akto_guardrails.py
index 901cdd3b95e..1b6e7ccfda7 100644
--- a/tests/guardrails_tests/test_akto_guardrails.py
+++ b/tests/guardrails_tests/test_akto_guardrails.py
@@ -14,7 +14,6 @@ from litellm.proxy.guardrails.guardrail_registry import (
)
from litellm.proxy.guardrails.guardrail_hooks.akto.akto import AktoGuardrail
-
# ---------------------------------------------------------------------------
# Registry tests
# ---------------------------------------------------------------------------
diff --git a/tests/guardrails_tests/test_custom_guardrail.py b/tests/guardrails_tests/test_custom_guardrail.py
index af1270756f2..95ac82b8d0f 100644
--- a/tests/guardrails_tests/test_custom_guardrail.py
+++ b/tests/guardrails_tests/test_custom_guardrail.py
@@ -6,7 +6,6 @@ import io
import os
import sys
-
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
diff --git a/tests/guardrails_tests/test_eu_ai_act_article5.py b/tests/guardrails_tests/test_eu_ai_act_article5.py
index bda7bf6f517..d602b206b53 100644
--- a/tests/guardrails_tests/test_eu_ai_act_article5.py
+++ b/tests/guardrails_tests/test_eu_ai_act_article5.py
@@ -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)
diff --git a/tests/guardrails_tests/test_sg_mas_ai_guardrails.py b/tests/guardrails_tests/test_sg_mas_ai_guardrails.py
index 668ee704692..47bfd1b8a10 100644
--- a/tests/guardrails_tests/test_sg_mas_ai_guardrails.py
+++ b/tests/guardrails_tests/test_sg_mas_ai_guardrails.py
@@ -23,7 +23,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
ContentFilterCategoryConfig,
)
-
# ── helpers ──────────────────────────────────────────────────────────────
POLICY_DIR = os.path.abspath(
diff --git a/tests/guardrails_tests/test_sg_pdpa_guardrails.py b/tests/guardrails_tests/test_sg_pdpa_guardrails.py
index fd7133bc745..03f55777f21 100644
--- a/tests/guardrails_tests/test_sg_pdpa_guardrails.py
+++ b/tests/guardrails_tests/test_sg_pdpa_guardrails.py
@@ -28,7 +28,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
ContentFilterCategoryConfig,
)
-
# ── helpers ──────────────────────────────────────────────────────────────
POLICY_DIR = os.path.abspath(
diff --git a/tests/image_gen_tests/test_image_generation.py b/tests/image_gen_tests/test_image_generation.py
index 873777189c9..2e6e08276fe 100644
--- a/tests/image_gen_tests/test_image_generation.py
+++ b/tests/image_gen_tests/test_image_generation.py
@@ -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
diff --git a/tests/image_gen_tests/test_image_variation.py b/tests/image_gen_tests/test_image_variation.py
index 301835057a7..ef528c391e1 100644
--- a/tests/image_gen_tests/test_image_variation.py
+++ b/tests/image_gen_tests/test_image_variation.py
@@ -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
diff --git a/tests/integration/test_oci_proxy_integration.py b/tests/integration/test_oci_proxy_integration.py
index f41e4826a04..4c1241ba891 100644
--- a/tests/integration/test_oci_proxy_integration.py
+++ b/tests/integration/test_oci_proxy_integration.py
@@ -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:
diff --git a/tests/litellm/llms/bedrock/embed/test_embedding.py b/tests/litellm/llms/bedrock/embed/test_embedding.py
index 261448842f4..163c143d825 100644
--- a/tests/litellm/llms/bedrock/embed/test_embedding.py
+++ b/tests/litellm/llms/bedrock/embed/test_embedding.py
@@ -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"
)
diff --git a/tests/litellm/llms/bedrock/test_nova_imported_models.py b/tests/litellm/llms/bedrock/test_nova_imported_models.py
index e3677aaf9e6..b5fd2630cbe 100644
--- a/tests/litellm/llms/bedrock/test_nova_imported_models.py
+++ b/tests/litellm/llms/bedrock/test_nova_imported_models.py
@@ -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}"
diff --git a/tests/litellm/test_bedrock_nemotron_super.py b/tests/litellm/test_bedrock_nemotron_super.py
index 8b081f10d1d..7b4efd86b28 100644
--- a/tests/litellm/test_bedrock_nemotron_super.py
+++ b/tests/litellm/test_bedrock_nemotron_super.py
@@ -11,7 +11,6 @@ import pytest
from litellm import get_model_info
-
MODEL_NAME = "nvidia.nemotron-super-3-120b"
diff --git a/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py b/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py
index c32917efe87..87af6a35f4b 100644
--- a/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py
+++ b/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py
@@ -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
diff --git a/tests/litellm_utils_tests/test_proxy_budget_reset.py b/tests/litellm_utils_tests/test_proxy_budget_reset.py
index 5c96eb619bf..468f02b1d40 100644
--- a/tests/litellm_utils_tests/test_proxy_budget_reset.py
+++ b/tests/litellm_utils_tests/test_proxy_budget_reset.py
@@ -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":
diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py
index d64633413a0..7c5927b4cb2 100644
--- a/tests/litellm_utils_tests/test_utils.py
+++ b/tests/litellm_utils_tests/test_utils.py
@@ -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 = {
diff --git a/tests/llm_responses_api_testing/test_anthropic_responses_api.py b/tests/llm_responses_api_testing/test_anthropic_responses_api.py
index 68ff22e8938..3399df148f6 100644
--- a/tests/llm_responses_api_testing/test_anthropic_responses_api.py
+++ b/tests/llm_responses_api_testing/test_anthropic_responses_api.py
@@ -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
diff --git a/tests/llm_responses_api_testing/test_anthropic_tool_result_empty_call_id.py b/tests/llm_responses_api_testing/test_anthropic_tool_result_empty_call_id.py
index 08b1c1784e7..3d256794cb8 100644
--- a/tests/llm_responses_api_testing/test_anthropic_tool_result_empty_call_id.py
+++ b/tests/llm_responses_api_testing/test_anthropic_tool_result_empty_call_id.py
@@ -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:
diff --git a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py
index 37fcc602d37..1bcc17ccbaf 100644
--- a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py
+++ b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py
@@ -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.
"""
diff --git a/tests/llm_responses_api_testing/test_responses_hooks.py b/tests/llm_responses_api_testing/test_responses_hooks.py
index 3799a0b9121..09ac9c7a9e3 100644
--- a/tests/llm_responses_api_testing/test_responses_hooks.py
+++ b/tests/llm_responses_api_testing/test_responses_hooks.py
@@ -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
diff --git a/tests/llm_translation/test_bedrock_anthropic_regression.py b/tests/llm_translation/test_bedrock_anthropic_regression.py
index 8b8ce0a6cc8..d81e76cd4b7 100644
--- a/tests/llm_translation/test_bedrock_anthropic_regression.py
+++ b/tests/llm_translation/test_bedrock_anthropic_regression.py
@@ -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:
diff --git a/tests/llm_translation/test_crusoe.py b/tests/llm_translation/test_crusoe.py
index 56aa4e4cd42..158152d8aeb 100644
--- a/tests/llm_translation/test_crusoe.py
+++ b/tests/llm_translation/test_crusoe.py
@@ -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
diff --git a/tests/llm_translation/test_deepseek_completion.py b/tests/llm_translation/test_deepseek_completion.py
index 2ede5d3f3f8..2da2b9be848 100644
--- a/tests/llm_translation/test_deepseek_completion.py
+++ b/tests/llm_translation/test_deepseek_completion.py
@@ -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"
diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py
index 0a5aebdf91b..b9b1a4753bb 100644
--- a/tests/llm_translation/test_gemini.py
+++ b/tests/llm_translation/test_gemini.py
@@ -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():
diff --git a/tests/llm_translation/test_gemini_image_usage.py b/tests/llm_translation/test_gemini_image_usage.py
index 096f9c4796c..75ce8c81469 100644
--- a/tests/llm_translation/test_gemini_image_usage.py
+++ b/tests/llm_translation/test_gemini_image_usage.py
@@ -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.
"""
diff --git a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py
index 9a69f513069..e2a43711f74 100644
--- a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py
+++ b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py
@@ -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,
)
diff --git a/tests/local_testing/create_mock_standard_logging_payload.py b/tests/local_testing/create_mock_standard_logging_payload.py
index 106328e95e2..775efda96e1 100644
--- a/tests/local_testing/create_mock_standard_logging_payload.py
+++ b/tests/local_testing/create_mock_standard_logging_payload.py
@@ -2,7 +2,6 @@ import io
import os
import sys
-
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
diff --git a/tests/local_testing/test_amazing_vertex_completion.py b/tests/local_testing/test_amazing_vertex_completion.py
index 2382b8a5197..8f7e6ad557f 100644
--- a/tests/local_testing/test_amazing_vertex_completion.py
+++ b/tests/local_testing/test_amazing_vertex_completion.py
@@ -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": [
diff --git a/tests/local_testing/test_cache_preset_key.py b/tests/local_testing/test_cache_preset_key.py
index d6518c5a073..cf496fe4b86 100644
--- a/tests/local_testing/test_cache_preset_key.py
+++ b/tests/local_testing/test_cache_preset_key.py
@@ -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'
"""
diff --git a/tests/local_testing/test_caching.py b/tests/local_testing/test_caching.py
index 0c7c0157651..a68e8c91584 100644
--- a/tests/local_testing/test_caching.py
+++ b/tests/local_testing/test_caching.py
@@ -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()
diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py
index cce6d33e799..e5c838c7e5d 100644
--- a/tests/local_testing/test_completion.py
+++ b/tests/local_testing/test_completion.py
@@ -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"
diff --git a/tests/local_testing/test_configs/custom_callbacks.py b/tests/local_testing/test_configs/custom_callbacks.py
index 42f88b5d19b..7ef6c1aaad5 100644
--- a/tests/local_testing/test_configs/custom_callbacks.py
+++ b/tests/local_testing/test_configs/custom_callbacks.py
@@ -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):
diff --git a/tests/local_testing/test_function_calling.py b/tests/local_testing/test_function_calling.py
index 3c7e004b62e..dab3585833b 100644
--- a/tests/local_testing/test_function_calling.py
+++ b/tests/local_testing/test_function_calling.py
@@ -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",
diff --git a/tests/local_testing/test_langsmith.py b/tests/local_testing/test_langsmith.py
index af7ac46a1cf..6f09b2701b0 100644
--- a/tests/local_testing/test_langsmith.py
+++ b/tests/local_testing/test_langsmith.py
@@ -21,7 +21,6 @@ verbose_logger.setLevel(logging.DEBUG)
litellm.set_verbose = True
import time
-
# test_langsmith_logging()
diff --git a/tests/local_testing/test_router_caching.py b/tests/local_testing/test_router_caching.py
index cb223b661b4..ad358699a48 100644
--- a/tests/local_testing/test_router_caching.py
+++ b/tests/local_testing/test_router_caching.py
@@ -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
diff --git a/tests/logging_callback_tests/create_mock_standard_logging_payload.py b/tests/logging_callback_tests/create_mock_standard_logging_payload.py
index 106328e95e2..775efda96e1 100644
--- a/tests/logging_callback_tests/create_mock_standard_logging_payload.py
+++ b/tests/logging_callback_tests/create_mock_standard_logging_payload.py
@@ -2,7 +2,6 @@ import io
import os
import sys
-
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
diff --git a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py
index d6d0652ed77..615148cd853 100644
--- a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py
+++ b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py
@@ -2,7 +2,6 @@ import io
import os
import sys
-
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
diff --git a/tests/logging_callback_tests/test_datadog_llm_obs.py b/tests/logging_callback_tests/test_datadog_llm_obs.py
index 56aae7aa8bf..b19963f8f73 100644
--- a/tests/logging_callback_tests/test_datadog_llm_obs.py
+++ b/tests/logging_callback_tests/test_datadog_llm_obs.py
@@ -6,7 +6,6 @@ import io
import os
import sys
-
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
diff --git a/tests/logging_callback_tests/test_gcs_pub_sub.py b/tests/logging_callback_tests/test_gcs_pub_sub.py
index f4cc9735177..6c99e724afe 100644
--- a/tests/logging_callback_tests/test_gcs_pub_sub.py
+++ b/tests/logging_callback_tests/test_gcs_pub_sub.py
@@ -2,7 +2,6 @@ import io
import os
import sys
-
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
diff --git a/tests/logging_callback_tests/test_generic_api_callback.py b/tests/logging_callback_tests/test_generic_api_callback.py
index fbe74d017a6..639e6dec626 100644
--- a/tests/logging_callback_tests/test_generic_api_callback.py
+++ b/tests/logging_callback_tests/test_generic_api_callback.py
@@ -2,7 +2,6 @@ import io
import os
import sys
-
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
diff --git a/tests/logging_callback_tests/test_langsmith_unit_test.py b/tests/logging_callback_tests/test_langsmith_unit_test.py
index 9cc1acd1ee4..2fc19827f49 100644
--- a/tests/logging_callback_tests/test_langsmith_unit_test.py
+++ b/tests/logging_callback_tests/test_langsmith_unit_test.py
@@ -2,7 +2,6 @@ import io
import os
import sys
-
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
diff --git a/tests/logging_callback_tests/test_view_request_resp_logs.py b/tests/logging_callback_tests/test_view_request_resp_logs.py
index ea778a44e67..66463e315c3 100644
--- a/tests/logging_callback_tests/test_view_request_resp_logs.py
+++ b/tests/logging_callback_tests/test_view_request_resp_logs.py
@@ -25,7 +25,6 @@ from litellm.integrations.gcs_bucket.gcs_bucket import (
)
from litellm.types.utils import StandardCallbackDynamicParams
-
# This is the response payload that GCS would return.
mock_response_data = {
"id": "chatcmpl-9870a859d6df402795f75dc5fca5b2e0",
diff --git a/tests/mcp_tests/test_mcp_logging.py b/tests/mcp_tests/test_mcp_logging.py
index 55b49aa0d29..73de99ac894 100644
--- a/tests/mcp_tests/test_mcp_logging.py
+++ b/tests/mcp_tests/test_mcp_logging.py
@@ -5,7 +5,6 @@ import asyncio
from typing import Optional
from unittest.mock import AsyncMock, patch
-
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py
index eea2f2721ab..827f5528a82 100644
--- a/tests/mcp_tests/test_mcp_server.py
+++ b/tests/mcp_tests/test_mcp_server.py
@@ -18,7 +18,6 @@ from litellm.proxy._types import LiteLLM_ObjectPermissionTable
from mcp.types import Tool as MCPTool, CallToolResult, ListToolsResult
from mcp.types import TextContent
-
mcp_server_manager = MCPServerManager()
diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py
index 2dd57e13d3b..62fc7ad3ff8 100644
--- a/tests/mcp_tests/test_proxy_mcp_e2e.py
+++ b/tests/mcp_tests/test_proxy_mcp_e2e.py
@@ -20,7 +20,6 @@ from litellm.proxy.proxy_server import (
initialize,
)
-
CONFIG_TEMPLATE_PATH = Path("tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml")
MCP_SERVER_SCRIPT = Path("tests/mcp_tests/mcp_server.py")
PROJECT_ROOT = Path(__file__).resolve().parents[2]
diff --git a/tests/ocr_tests/base_ocr_unit_tests.py b/tests/ocr_tests/base_ocr_unit_tests.py
index ae65efd952d..a09158ce18c 100644
--- a/tests/ocr_tests/base_ocr_unit_tests.py
+++ b/tests/ocr_tests/base_ocr_unit_tests.py
@@ -9,7 +9,6 @@ import litellm
import os
from abc import ABC, abstractmethod
-
# Test resources
TEST_IMAGE_PATH = "test_image_edit.png"
# Tiny in-repo PDF served via jsdelivr (sha-pinned, immutable). The arxiv
diff --git a/tests/ocr_tests/test_ocr_azure_document_intelligence.py b/tests/ocr_tests/test_ocr_azure_document_intelligence.py
index 7269890b7b6..27428f96a21 100644
--- a/tests/ocr_tests/test_ocr_azure_document_intelligence.py
+++ b/tests/ocr_tests/test_ocr_azure_document_intelligence.py
@@ -110,9 +110,7 @@ class TestAzureDocumentIntelligencePagesParam:
model="azure_ai/doc-intelligence/prebuilt-layout",
optional_params={"pages": "1-3,5"},
)
- assert (
- f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url
- ), url
+ assert f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url, url
assert "pages=1-3,5" in url, url
assert "/documentintelligence/documentModels/prebuilt-layout:analyze" in url
@@ -168,4 +166,3 @@ class TestAzureDocumentIntelligencePagesParam:
assert "pages=3,4,5,6,7,8,9" in url
assert req.data == {"urlSource": "https://example.com/x.pdf"}
-
diff --git a/tests/old_proxy_tests/tests/bursty_load_test_completion.py b/tests/old_proxy_tests/tests/bursty_load_test_completion.py
index 642bd9f14d4..41944c03aa2 100644
--- a/tests/old_proxy_tests/tests/bursty_load_test_completion.py
+++ b/tests/old_proxy_tests/tests/bursty_load_test_completion.py
@@ -3,7 +3,6 @@ from openai import AsyncOpenAI
from litellm._uuid import uuid
import traceback
-
litellm_client = AsyncOpenAI(api_key="test", base_url="http://0.0.0.0:8000")
diff --git a/tests/old_proxy_tests/tests/load_test_embedding_100.py b/tests/old_proxy_tests/tests/load_test_embedding_100.py
index 8cd4d250249..bfb4e137e80 100644
--- a/tests/old_proxy_tests/tests/load_test_embedding_100.py
+++ b/tests/old_proxy_tests/tests/load_test_embedding_100.py
@@ -3,7 +3,6 @@ from openai import AsyncOpenAI
from litellm._uuid import uuid
import traceback
-
litellm_client = AsyncOpenAI(api_key="test", base_url="http://0.0.0.0:8000")
diff --git a/tests/old_proxy_tests/tests/test_openai_request_with_traceparent.py b/tests/old_proxy_tests/tests/test_openai_request_with_traceparent.py
index 2f8455dcbe9..cde68002a75 100644
--- a/tests/old_proxy_tests/tests/test_openai_request_with_traceparent.py
+++ b/tests/old_proxy_tests/tests/test_openai_request_with_traceparent.py
@@ -8,7 +8,6 @@ from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator
-
trace.set_tracer_provider(TracerProvider())
memory_exporter = InMemorySpanExporter()
span_processor = SimpleSpanProcessor(memory_exporter)
diff --git a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py
index c6f4128f2c5..a2c6338ede6 100644
--- a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py
+++ b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py
@@ -11,7 +11,6 @@ import sys
import time
from unittest.mock import patch, MagicMock, AsyncMock
-
BASE_URL = "http://localhost:4000" # Replace with your actual base URL
API_KEY = "sk-1234" # Replace with your actual API key
@@ -234,7 +233,9 @@ async def test_list_batches_with_target_model_names():
# Test data
target_model_names = "gpt-5.5,gpt-5-mini"
- expected_model = "gpt-5.5" # Should use the first model from the comma-separated list
+ expected_model = (
+ "gpt-5.5" # Should use the first model from the comma-separated list
+ )
# Mock response for list_batches
mock_batch_response = {
diff --git a/tests/openai_endpoints_tests/test_openai_files_endpoints.py b/tests/openai_endpoints_tests/test_openai_files_endpoints.py
index 6be692b278c..547bb554835 100644
--- a/tests/openai_endpoints_tests/test_openai_files_endpoints.py
+++ b/tests/openai_endpoints_tests/test_openai_files_endpoints.py
@@ -6,7 +6,6 @@ import aiohttp, openai
from openai import OpenAI, AsyncOpenAI
from typing import Optional, List, Union
-
BASE_URL = "http://localhost:4000" # Replace with your actual base URL
API_KEY = "sk-1234" # Replace with your actual API key
diff --git a/tests/otel_tests/test_e2e_model_access.py b/tests/otel_tests/test_e2e_model_access.py
index 5b5f2a89c8d..d1afcc72800 100644
--- a/tests/otel_tests/test_e2e_model_access.py
+++ b/tests/otel_tests/test_e2e_model_access.py
@@ -5,7 +5,6 @@ import json
from httpx import AsyncClient
from typing import Any, Optional, List, Literal
-
# The proxy strips client-supplied `mock_response` unless the calling key or
# team has this admin-metadata flag set. See `_UNTRUSTED_ROOT_CONTROL_FIELDS`
# in litellm/proxy/litellm_pre_call_utils.py.
@@ -152,9 +151,7 @@ async def test_model_access_update():
async with aiohttp.ClientSession() as session:
# Both models should now work
await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5")
- await mock_chat_completion(
- session=session, key=key, model="openai/gpt-5-mini"
- )
+ await mock_chat_completion(session=session, key=key, model="openai/gpt-5-mini")
# Non-OpenAI model should still fail
with pytest.raises(Exception) as exc_info:
@@ -274,9 +271,7 @@ async def test_team_model_access_update():
async with aiohttp.ClientSession() as session:
# Both models should now work
await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5")
- await mock_chat_completion(
- session=session, key=key, model="openai/gpt-5-mini"
- )
+ await mock_chat_completion(session=session, key=key, model="openai/gpt-5-mini")
# Non-OpenAI model should still fail
with pytest.raises(Exception) as exc_info:
diff --git a/tests/otel_tests/test_team_member_permissions.py b/tests/otel_tests/test_team_member_permissions.py
index ddb8b741c45..36926b9f525 100644
--- a/tests/otel_tests/test_team_member_permissions.py
+++ b/tests/otel_tests/test_team_member_permissions.py
@@ -4,13 +4,13 @@
Invalid Permissions:
- - User tries creating a key with team_id = team_id -> expect to fail. Invalid Permissions
- - User tries editing a key with team_id = team_id -> expect to fail. Invalid Permissions
- - User tries deleting a key with team_id = team_id -> expect to fail. Invalid Permissions
- - User tries regenerating a key with team_id = team_id -> expect to fail. Invalid Permissions
+ - User tries creating a key with team_id = team_id -> expect to fail. Invalid Permissions
+ - User tries editing a key with team_id = team_id -> expect to fail. Invalid Permissions
+ - User tries deleting a key with team_id = team_id -> expect to fail. Invalid Permissions
+ - User tries regenerating a key with team_id = team_id -> expect to fail. Invalid Permissions
Valid Permissions:
- - User tries calling /key/info with team_id, expect to get valid response
+ - User tries calling /key/info with team_id, expect to get valid response
@@ -26,7 +26,7 @@
Invalid Permissions:
- User tries creating a key with team_id = team_id -> expect to fail. Invalid Permissions
- - User tries calling /key/info with team_id, expect to get valid response
+ - User tries calling /key/info with team_id, expect to get valid response
diff --git a/tests/pass_through_tests/test_openai_assistants_passthrough.py b/tests/pass_through_tests/test_openai_assistants_passthrough.py
index 28568005fd6..c8e9a9ef0f0 100644
--- a/tests/pass_through_tests/test_openai_assistants_passthrough.py
+++ b/tests/pass_through_tests/test_openai_assistants_passthrough.py
@@ -6,7 +6,6 @@ import tempfile
from typing_extensions import override
from openai import AssistantEventHandler
-
client = openai.OpenAI(base_url="http://0.0.0.0:4000/openai", api_key="sk-1234")
diff --git a/tests/proxy_admin_ui_tests/test_route_check_unit_tests.py b/tests/proxy_admin_ui_tests/test_route_check_unit_tests.py
index f0cc6985e66..3b1db1c327e 100644
--- a/tests/proxy_admin_ui_tests/test_route_check_unit_tests.py
+++ b/tests/proxy_admin_ui_tests/test_route_check_unit_tests.py
@@ -14,7 +14,6 @@ import io
import os
import time
-
# this file is to test litellm/proxy
sys.path.insert(
diff --git a/tests/proxy_admin_ui_tests/test_usage_endpoints.py b/tests/proxy_admin_ui_tests/test_usage_endpoints.py
index 54ad136f082..d8ef87f9333 100644
--- a/tests/proxy_admin_ui_tests/test_usage_endpoints.py
+++ b/tests/proxy_admin_ui_tests/test_usage_endpoints.py
@@ -1,5 +1,5 @@
"""
-Tests the following endpoints used by the UI
+Tests the following endpoints used by the UI
/global/spend/logs
/global/spend/keys
@@ -9,7 +9,7 @@ Tests the following endpoints used by the UI
For all tests - test the following:
-- Response is valid
+- Response is valid
- Response for Admin User is different from response from Internal User
"""
diff --git a/tests/proxy_e2e_anthropic_messages_tests/test_claude_agent_sdk.py b/tests/proxy_e2e_anthropic_messages_tests/test_claude_agent_sdk.py
index 48eb7d85ec1..7f5fcd38c87 100644
--- a/tests/proxy_e2e_anthropic_messages_tests/test_claude_agent_sdk.py
+++ b/tests/proxy_e2e_anthropic_messages_tests/test_claude_agent_sdk.py
@@ -12,7 +12,6 @@ import pytest
import asyncio
from claude_agent_sdk import ClaudeSDKClient, ClaudeAgentOptions
-
# Test models from test_config.yaml
# Note: bedrock-converse-claude-sonnet-4.5 removed temporarily as the Bedrock Converse API
# for Claude Sonnet 4.5 may not be available in all regions/accounts
diff --git a/tests/proxy_unit_tests/conftest.py b/tests/proxy_unit_tests/conftest.py
index a0326f64ed7..6ff9ffe84c4 100644
--- a/tests/proxy_unit_tests/conftest.py
+++ b/tests/proxy_unit_tests/conftest.py
@@ -16,7 +16,6 @@ sys.path.insert(
import litellm
import litellm.proxy.proxy_server
-
# Top-level assignments of these types are the ones importlib.reload(litellm)
# would have effectively reset. We snapshot them at conftest import time and
# deep-copy the snapshot back before every test.
diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py
index e7136ecb195..429681907e8 100644
--- a/tests/proxy_unit_tests/test_auth_checks.py
+++ b/tests/proxy_unit_tests/test_auth_checks.py
@@ -39,9 +39,9 @@ from litellm.proxy.utils import CallInfo
async def test_get_end_user_object(customer_spend, customer_budget):
"""
Scenario 1: normal - get_end_user_object returns the cached user
- Scenario 2: user over budget - NOTE: budget enforcement now happens in
+ Scenario 2: user over budget - NOTE: budget enforcement now happens in
common_checks() via _check_end_user_budget(), not in get_end_user_object()
-
+
This test verifies that get_end_user_object correctly retrieves the end user
from cache. Budget enforcement is tested separately in test_check_end_user_budget().
"""
@@ -81,12 +81,12 @@ async def test_check_end_user_budget(customer_spend, customer_budget):
Test _check_end_user_budget enforcement:
- Scenario 1: customer_spend=0, customer_budget=10 - should pass (under budget)
- Scenario 2: customer_spend=10, customer_budget=0 - should fail (over budget)
-
- Note: Budget enforcement for end users happens in common_checks() via
+
+ Note: Budget enforcement for end users happens in common_checks() via
_check_end_user_budget(), not in get_end_user_object().
"""
from litellm.proxy.auth.auth_checks import _check_end_user_budget
-
+
_budget = LiteLLM_BudgetTable(max_budget=customer_budget)
end_user_obj = LiteLLM_EndUserTable(
user_id="my-test-customer",
@@ -94,9 +94,9 @@ async def test_check_end_user_budget(customer_spend, customer_budget):
litellm_budget_table=_budget,
blocked=False,
)
-
+
should_exceed = customer_spend > customer_budget
-
+
try:
await _check_end_user_budget(
end_user_obj=end_user_obj,
diff --git a/tests/proxy_unit_tests/test_configs/custom_callbacks.py b/tests/proxy_unit_tests/test_configs/custom_callbacks.py
index c7d66c068e9..ac705cd28a9 100644
--- a/tests/proxy_unit_tests/test_configs/custom_callbacks.py
+++ b/tests/proxy_unit_tests/test_configs/custom_callbacks.py
@@ -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):
diff --git a/tests/proxy_unit_tests/test_default_end_user_budget_simple.py b/tests/proxy_unit_tests/test_default_end_user_budget_simple.py
index 6170b0a972e..9010647da41 100644
--- a/tests/proxy_unit_tests/test_default_end_user_budget_simple.py
+++ b/tests/proxy_unit_tests/test_default_end_user_budget_simple.py
@@ -136,12 +136,12 @@ async def test_budget_enforcement_blocks_over_budget_users():
"""
Core scenario: Budget limits are actually enforced via _check_end_user_budget.
Users who exceed their budget should be blocked.
-
+
Note: Budget enforcement happens in common_checks() via _check_end_user_budget(),
not in get_end_user_object(). get_end_user_object only fetches the user data.
"""
from litellm.proxy.auth.auth_checks import _check_end_user_budget
-
+
end_user_id = f"test_user_{uuid.uuid4().hex}"
default_budget_id = str(uuid.uuid4())
litellm.max_end_user_budget_id = default_budget_id
@@ -182,7 +182,7 @@ async def test_budget_enforcement_blocks_over_budget_users():
user_api_key_cache=mock_cache,
route="/chat/completions",
)
-
+
# Verify user was fetched with default budget applied
assert result is not None
assert result.litellm_budget_table is not None
diff --git a/tests/proxy_unit_tests/test_deprecated_key_grace_period.py b/tests/proxy_unit_tests/test_deprecated_key_grace_period.py
index a91ecf95f32..643b9775578 100644
--- a/tests/proxy_unit_tests/test_deprecated_key_grace_period.py
+++ b/tests/proxy_unit_tests/test_deprecated_key_grace_period.py
@@ -24,7 +24,6 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
-
# ── helpers ───────────────────────────────────────────────────────────────────
HASHED_TOKEN = "165efe575c98fe7e65d98cb2de71b68842049e286afd33a92d3491c340216880"
diff --git a/tests/proxy_unit_tests/test_gemini_agents_endpoints.py b/tests/proxy_unit_tests/test_gemini_agents_endpoints.py
index bdac9348f71..8a42f8d53db 100644
--- a/tests/proxy_unit_tests/test_gemini_agents_endpoints.py
+++ b/tests/proxy_unit_tests/test_gemini_agents_endpoints.py
@@ -23,7 +23,6 @@ from litellm.proxy.google_endpoints.agents_endpoints import (
_merge_query_params_into_data,
)
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py
index 921fbfa320f..0455acb23c7 100644
--- a/tests/proxy_unit_tests/test_proxy_server.py
+++ b/tests/proxy_unit_tests/test_proxy_server.py
@@ -2119,14 +2119,11 @@ async def test_model_info_alias_without_prisma(hidden):
models = resp["data"]
- alias_found = any(
- m["model_name"] == model_alias
- for m in models
- )
+ alias_found = any(m["model_name"] == model_alias for m in models)
assert alias_found is (not hidden)
-
+
@pytest.mark.parametrize("hidden", [True, False])
@pytest.mark.asyncio
@pytest.mark.skip(reason="Requires reliable external DB connection (prisma).")
diff --git a/tests/proxy_unit_tests/test_search_api_logging.py b/tests/proxy_unit_tests/test_search_api_logging.py
index 71bbe5351a2..3298389083b 100644
--- a/tests/proxy_unit_tests/test_search_api_logging.py
+++ b/tests/proxy_unit_tests/test_search_api_logging.py
@@ -2,7 +2,7 @@
Test search API logging and cost tracking in proxy.
Tests that search API requests are properly logged to LiteLLM_SpendLogs
-with correct fields populated (call_type, model, custom_llm_provider,
+with correct fields populated (call_type, model, custom_llm_provider,
model_group, spend, etc.)
"""
diff --git a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py
index 0daa5b17ffa..e6885d72f39 100644
--- a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py
+++ b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py
@@ -423,7 +423,9 @@ async def test_async_log_success_event_pushes_redis_increments_when_redis_config
_push_in_memory_increments_to_redis when Redis is wired so other workers see spend.
"""
dual_cache = DualCache()
- dual_cache.redis_cache = object() # truthy placeholder; push only checks is not None
+ dual_cache.redis_cache = (
+ object()
+ ) # truthy placeholder; push only checks is not None
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
model = "gpt-4"
kwargs = {
@@ -471,7 +473,9 @@ async def test_async_log_success_event_skips_redis_push_without_redis(budget_lim
},
},
}
- with patch.object(budget_limiter, "_increment_spend_for_key", new_callable=AsyncMock):
+ with patch.object(
+ budget_limiter, "_increment_spend_for_key", new_callable=AsyncMock
+ ):
with patch.object(
budget_limiter,
"_push_in_memory_increments_to_redis",
diff --git a/tests/router_unit_tests/create_mock_standard_logging_payload.py b/tests/router_unit_tests/create_mock_standard_logging_payload.py
index 106328e95e2..775efda96e1 100644
--- a/tests/router_unit_tests/create_mock_standard_logging_payload.py
+++ b/tests/router_unit_tests/create_mock_standard_logging_payload.py
@@ -2,7 +2,6 @@ import io
import os
import sys
-
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
diff --git a/tests/router_unit_tests/test_get_model_list_alias_optimization.py b/tests/router_unit_tests/test_get_model_list_alias_optimization.py
index 145c7e8092e..6d7ad3ebc9b 100644
--- a/tests/router_unit_tests/test_get_model_list_alias_optimization.py
+++ b/tests/router_unit_tests/test_get_model_list_alias_optimization.py
@@ -20,9 +20,7 @@ def test_get_model_list_from_model_alias_should_not_iterate_for_non_alias_lookup
{f"alias-{idx}": "gpt-5.5" for idx in range(200)}
)
- model_alias_list = router.get_model_list_from_model_alias(
- model_name="gpt-5-mini"
- )
+ model_alias_list = router.get_model_list_from_model_alias(model_name="gpt-5-mini")
assert model_alias_list == []
diff --git a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py
index 25bf79cd575..6db6283e760 100644
--- a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py
+++ b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py
@@ -250,15 +250,18 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
async def fake_original(**_kwargs):
return streaming_iter
- with patch.object(
- router,
- "_ageneric_api_call_with_fallbacks",
- new=AsyncMock(return_value=streaming_iter),
- ), patch.object(
- router,
- "_aresponses_streaming_iterator",
- new=AsyncMock(return_value=wrapped),
- ) as mock_wrap:
+ with (
+ patch.object(
+ router,
+ "_ageneric_api_call_with_fallbacks",
+ new=AsyncMock(return_value=streaming_iter),
+ ),
+ patch.object(
+ router,
+ "_aresponses_streaming_iterator",
+ new=AsyncMock(return_value=wrapped),
+ ) as mock_wrap,
+ ):
out = await router._aresponses_with_streaming_fallbacks(
original_function=fake_original,
model="primary",
diff --git a/tests/router_unit_tests/test_router_cooldown_utils.py b/tests/router_unit_tests/test_router_cooldown_utils.py
index ea0cd74d877..a8e3df260eb 100644
--- a/tests/router_unit_tests/test_router_cooldown_utils.py
+++ b/tests/router_unit_tests/test_router_cooldown_utils.py
@@ -112,9 +112,7 @@ def test_should_cooldown_deployment_rate_limit_error(testing_litellm_router):
Test the _should_cooldown_deployment function when a rate limit error occurs
"""
# Test 429 error (rate limit) -> always cooldown a deployment returning 429s
- _exception = litellm.exceptions.RateLimitError(
- "Rate limit", "openai", "gpt-5-mini"
- )
+ _exception = litellm.exceptions.RateLimitError("Rate limit", "openai", "gpt-5-mini")
assert (
_should_cooldown_deployment(
testing_litellm_router, "test_deployment", 429, _exception
@@ -150,9 +148,7 @@ async def test_should_cooldown_deployment(testing_litellm_router):
verbose_router_logger.setLevel(logging.DEBUG)
# Test 429 error (rate limit) -> always cooldown a deployment returning 429s
- _exception = litellm.exceptions.RateLimitError(
- "Rate limit", "openai", "gpt-5-mini"
- )
+ _exception = litellm.exceptions.RateLimitError("Rate limit", "openai", "gpt-5-mini")
assert (
_should_cooldown_deployment(
testing_litellm_router, "test_deployment", 429, _exception
diff --git a/tests/router_unit_tests/test_router_handle_error.py b/tests/router_unit_tests/test_router_handle_error.py
index a84c90ccb78..e19c12f5e7a 100644
--- a/tests/router_unit_tests/test_router_handle_error.py
+++ b/tests/router_unit_tests/test_router_handle_error.py
@@ -15,7 +15,6 @@ from collections import defaultdict
from dotenv import load_dotenv
from unittest.mock import AsyncMock, MagicMock
-
load_dotenv()
diff --git a/tests/router_unit_tests/test_router_index_management.py b/tests/router_unit_tests/test_router_index_management.py
index 983fc0c4c3b..9dbb43a2e1b 100644
--- a/tests/router_unit_tests/test_router_index_management.py
+++ b/tests/router_unit_tests/test_router_index_management.py
@@ -142,7 +142,10 @@ class TestRouterIndexManagement:
router._update_team_model_index(model, 0)
assert router.team_model_to_deployment_indices[("team-abc", "gpt-5.5")] == [0]
router._update_team_model_index(model, 2)
- assert router.team_model_to_deployment_indices[("team-abc", "gpt-5.5")] == [0, 2]
+ assert router.team_model_to_deployment_indices[("team-abc", "gpt-5.5")] == [
+ 0,
+ 2,
+ ]
router._update_team_model_index(
{"model_name": "x", "model_info": {"id": "dep-2"}}, 5
diff --git a/tests/search_tests/test_google_pse_search.py b/tests/search_tests/test_google_pse_search.py
index 21d58a95491..9ce55b428da 100644
--- a/tests/search_tests/test_google_pse_search.py
+++ b/tests/search_tests/test_google_pse_search.py
@@ -10,7 +10,6 @@ sys.path.insert(0, os.path.abspath("../.."))
from tests.search_tests.base_search_unit_tests import BaseSearchTest
-
# class TestGooglePSESearch(BaseSearchTest):
# """
# Tests for Google PSE Search functionality.
diff --git a/tests/store_model_in_db_tests/test_callbacks_in_db.py b/tests/store_model_in_db_tests/test_callbacks_in_db.py
index 4a851251a3e..51545333c65 100644
--- a/tests/store_model_in_db_tests/test_callbacks_in_db.py
+++ b/tests/store_model_in_db_tests/test_callbacks_in_db.py
@@ -1,9 +1,9 @@
"""
PROD TEST - DO NOT Delete this Test
-e2e test for langfuse callback in DB
+e2e test for langfuse callback in DB
- Add langfuse callback to DB - with /config/update
-- wait 20 seconds for the callback to be loaded into the instance
+- wait 20 seconds for the callback to be loaded into the instance
- Make a /chat/completions request to the proxy
- Check if the request is logged in Langfuse
"""
diff --git a/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py b/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py
index 5503a5668bf..886465a8c92 100644
--- a/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py
+++ b/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py
@@ -14,7 +14,6 @@ import json
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
-
SAMPLE_ARN = "arn:aws:bedrock-agentcore:us-west-2:123456789:runtime/my_agent"
SAMPLE_MODEL = f"bedrock/agentcore/{SAMPLE_ARN}"
SAMPLE_PARAMS = {
diff --git a/tests/test_litellm/a2a_protocol/test_send_message_response.py b/tests/test_litellm/a2a_protocol/test_send_message_response.py
index 832aa288c7a..098a7d41f4e 100644
--- a/tests/test_litellm/a2a_protocol/test_send_message_response.py
+++ b/tests/test_litellm/a2a_protocol/test_send_message_response.py
@@ -9,9 +9,7 @@ def test_from_dict_backfills_id_on_agent_error_response():
"error": {"code": -32054, "message": "Session not found"},
}
- response = LiteLLMSendMessageResponse.from_dict(
- agent_error, request_id="r1"
- )
+ response = LiteLLMSendMessageResponse.from_dict(agent_error, request_id="r1")
assert response.id == "r1"
assert response.error == {"code": -32054, "message": "Session not found"}
@@ -25,9 +23,7 @@ def test_from_dict_preserves_existing_id():
"error": {"code": -32001, "message": "Task not found"},
}
- response = LiteLLMSendMessageResponse.from_dict(
- payload, request_id="r1"
- )
+ response = LiteLLMSendMessageResponse.from_dict(payload, request_id="r1")
assert response.id == "upstream-id"
diff --git a/tests/test_litellm/caching/test_check_and_fix_namespace_none_guard.py b/tests/test_litellm/caching/test_check_and_fix_namespace_none_guard.py
index c049c3157f4..99a3ea82492 100644
--- a/tests/test_litellm/caching/test_check_and_fix_namespace_none_guard.py
+++ b/tests/test_litellm/caching/test_check_and_fix_namespace_none_guard.py
@@ -3,6 +3,7 @@ Test that check_and_fix_namespace handles None key gracefully.
Regression test for https://github.com/BerriAI/litellm/issues/30424
"""
+
from unittest.mock import MagicMock
from litellm.caching.redis_cache import RedisCache
diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py
index 6a1de0586dd..429bec04546 100644
--- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py
+++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py
@@ -887,9 +887,7 @@ def test_transform_request_system_only_message_maps_to_system_input_item():
{
"type": "message",
"role": "system",
- "content": [
- {"type": "input_text", "text": "You are a helpful assistant."}
- ],
+ "content": [{"type": "input_text", "text": "You are a helpful assistant."}],
}
]
# System content lives in input only; not duplicated into instructions.
diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py
index 40439a78a49..d247b02074d 100644
--- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py
+++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py
@@ -32,12 +32,12 @@ def mock_env_vars():
# Store original values
original_api_key = os.environ.get("SENDGRID_API_KEY")
original_sender_email = os.environ.get("SENDGRID_SENDER_EMAIL")
-
+
# Set test API key and remove SENDGRID_SENDER_EMAIL to ensure isolation
os.environ["SENDGRID_API_KEY"] = "test_api_key"
if "SENDGRID_SENDER_EMAIL" in os.environ:
del os.environ["SENDGRID_SENDER_EMAIL"]
-
+
try:
yield
finally:
@@ -46,7 +46,7 @@ def mock_env_vars():
os.environ["SENDGRID_API_KEY"] = original_api_key
elif "SENDGRID_API_KEY" in os.environ:
del os.environ["SENDGRID_API_KEY"]
-
+
if original_sender_email is not None:
os.environ["SENDGRID_SENDER_EMAIL"] = original_sender_email
diff --git a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py
index 6e9c3c0354b..180041b5f71 100644
--- a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py
+++ b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py
@@ -16,13 +16,18 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
-
DECODED_UNIFIED_INPUT_FILE_ID = "litellm_proxy:application/octet-stream;unified_id,test-uuid;target_model_names,azure-gpt-4"
-B64_UNIFIED_INPUT_FILE_ID = base64.urlsafe_b64encode(DECODED_UNIFIED_INPUT_FILE_ID.encode()).decode().rstrip("=")
+B64_UNIFIED_INPUT_FILE_ID = (
+ base64.urlsafe_b64encode(DECODED_UNIFIED_INPUT_FILE_ID.encode())
+ .decode()
+ .rstrip("=")
+)
RAW_INPUT_FILE_ID = "file-raw-provider-abc123"
DECODED_UNIFIED_BATCH_ID = "litellm_proxy;model_id:model-xyz;llm_batch_id:batch-123"
-B64_UNIFIED_BATCH_ID = base64.urlsafe_b64encode(DECODED_UNIFIED_BATCH_ID.encode()).decode().rstrip("=")
+B64_UNIFIED_BATCH_ID = (
+ base64.urlsafe_b64encode(DECODED_UNIFIED_BATCH_ID.encode()).decode().rstrip("=")
+)
@pytest.mark.asyncio
@@ -55,10 +60,16 @@ async def test_should_resolve_raw_input_file_id_to_unified():
mock_managed_file.unified_file_id = B64_UNIFIED_INPUT_FILE_ID
mock_prisma = MagicMock()
- mock_prisma.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=mock_db_object)
- mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=mock_managed_file)
+ mock_prisma.db.litellm_managedobjecttable.find_first = AsyncMock(
+ return_value=mock_db_object
+ )
+ mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(
+ return_value=mock_managed_file
+ )
- from litellm.proxy.openai_files_endpoints.common_utils import get_batch_from_database
+ from litellm.proxy.openai_files_endpoints.common_utils import (
+ get_batch_from_database,
+ )
_, response = await get_batch_from_database(
batch_id=B64_UNIFIED_BATCH_ID,
diff --git a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py
index 420f5f9789c..c80b8a848ca 100644
--- a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py
+++ b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py
@@ -93,7 +93,9 @@ async def test_should_preserve_already_managed_input_file_id():
unified_batch_id = "bGl0ZWxsbV9wcm94eTpiYXRjaF9pZA"
decoded_unified = "litellm_proxy:application/octet-stream;unified_id,test-123"
- base64_input_file_id = base64.urlsafe_b64encode(decoded_unified.encode()).decode().rstrip("=")
+ base64_input_file_id = (
+ base64.urlsafe_b64encode(decoded_unified.encode()).decode().rstrip("=")
+ )
batch_data = {
"id": "batch-raw-123",
diff --git a/tests/test_litellm/enterprise/proxy/test_enterprise_routes.py b/tests/test_litellm/enterprise/proxy/test_enterprise_routes.py
index a9bf33a21ac..34c8c9c5c5d 100644
--- a/tests/test_litellm/enterprise/proxy/test_enterprise_routes.py
+++ b/tests/test_litellm/enterprise/proxy/test_enterprise_routes.py
@@ -14,63 +14,77 @@ import pytest
def test_enterprise_routes_all_imports_exist():
"""
Validate that all relative imports in enterprise_routes.py exist in the filesystem.
-
+
This catches any import errors from moved/deleted modules without hardcoding
specific module names. Works by checking that imported files actually exist.
"""
# Path to the enterprise_routes.py source file
enterprise_routes_path = os.path.join(
os.path.dirname(__file__),
- "..", "..", "..", "..",
- "enterprise", "litellm_enterprise", "proxy", "enterprise_routes.py"
+ "..",
+ "..",
+ "..",
+ "..",
+ "enterprise",
+ "litellm_enterprise",
+ "proxy",
+ "enterprise_routes.py",
)
-
+
enterprise_routes_path = os.path.normpath(enterprise_routes_path)
enterprise_proxy_dir = os.path.dirname(enterprise_routes_path)
-
+
if not os.path.exists(enterprise_routes_path):
pytest.skip(f"Enterprise routes file not found at {enterprise_routes_path}")
-
+
# Read and parse the source file
with open(enterprise_routes_path, "r") as f:
source_code = f.read()
-
+
try:
tree = ast.parse(source_code)
except SyntaxError as e:
pytest.fail(f"Syntax error in enterprise_routes.py: {e}")
-
+
# Check all relative imports
missing_imports = []
-
+
for node in ast.walk(tree):
if isinstance(node, ast.ImportFrom):
# level > 0 means it's a relative import (. or .. etc)
if node.level and node.level > 0:
module = node.module or ""
-
+
# Convert relative import to file path
# e.g., "audit_logging_endpoints" -> "audit_logging_endpoints.py"
# e.g., "vector_stores.endpoints" -> "vector_stores/endpoints.py"
module_path = module.replace(".", os.sep) if module else ""
-
+
# Check both .py file and package directory
- file_path = os.path.join(enterprise_proxy_dir, module_path + ".py") if module_path else None
- package_path = os.path.join(enterprise_proxy_dir, module_path, "__init__.py") if module_path else None
-
+ file_path = (
+ os.path.join(enterprise_proxy_dir, module_path + ".py")
+ if module_path
+ else None
+ )
+ package_path = (
+ os.path.join(enterprise_proxy_dir, module_path, "__init__.py")
+ if module_path
+ else None
+ )
+
# If module is empty (e.g., "from . import something"), skip check
if not module:
continue
-
+
file_exists = file_path and os.path.exists(file_path)
package_exists = package_path and os.path.exists(package_path)
-
+
if not file_exists and not package_exists:
missing_imports.append(
f"Line {node.lineno}: Cannot find '.{module}' "
f"(checked: {file_path} and {package_path})"
)
-
+
if missing_imports:
error_msg = "Found imports in enterprise_routes.py that don't exist:\n"
error_msg += "\n".join(missing_imports)
diff --git a/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py b/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py
index 852077dcf0c..3c7aace7d31 100644
--- a/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py
+++ b/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py
@@ -61,7 +61,7 @@ def _make_managed_files_instance_with_batches(
):
"""
Create a _PROXY_LiteLLMManagedFiles instance with mocked DB and batches.
-
+
Args:
file_id: The unified file ID
batches: List of batch records to return from DB
@@ -79,7 +79,7 @@ def _make_managed_files_instance_with_batches(
# Mock prisma
mock_prisma = MagicMock()
-
+
# Mock file table queries
mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(
return_value=mock_file_record
@@ -87,7 +87,7 @@ def _make_managed_files_instance_with_batches(
mock_prisma.db.litellm_managedfiletable.delete = AsyncMock(
return_value=mock_file_record
)
-
+
# Mock batch/object table queries
mock_prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=batches
@@ -95,11 +95,13 @@ def _make_managed_files_instance_with_batches(
# Mock cache
mock_cache = MagicMock()
- mock_cache.async_get_cache = AsyncMock(return_value={
- "unified_file_id": file_id,
- "model_mappings": {"model-123": "provider-file-abc"},
- "flat_model_file_ids": ["provider-file-abc"],
- })
+ mock_cache.async_get_cache = AsyncMock(
+ return_value={
+ "unified_file_id": file_id,
+ "model_mappings": {"model-123": "provider-file-abc"},
+ "flat_model_file_ids": ["provider-file-abc"],
+ }
+ )
mock_cache.async_set_cache = AsyncMock()
instance = _PROXY_LiteLLMManagedFiles(
@@ -117,17 +119,17 @@ def test_is_batch_polling_enabled_when_job_registered():
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
-
+
instance = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=MagicMock(),
)
-
+
# Mock scheduler with registered job
mock_scheduler = MagicMock()
mock_job = MagicMock()
mock_scheduler.get_job.return_value = mock_job
-
+
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
assert instance._is_batch_polling_enabled() is True
@@ -137,16 +139,16 @@ def test_is_batch_polling_disabled_when_job_not_registered():
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
-
+
instance = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=MagicMock(),
)
-
+
# Mock scheduler without registered job
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = None
-
+
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
assert instance._is_batch_polling_enabled() is False
@@ -156,12 +158,12 @@ def test_is_batch_polling_disabled_when_no_scheduler():
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
-
+
instance = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=MagicMock(),
prisma_client=MagicMock(),
)
-
+
with patch("litellm.proxy.proxy_server.scheduler", None):
assert instance._is_batch_polling_enabled() is False
@@ -174,26 +176,28 @@ async def test_get_batches_referencing_file_finds_batch_with_input_file():
"""Test finding a batch that references the file as input_file_id."""
unified_file_id = _make_unified_file_id("file-input-123")
unified_batch_id = _make_unified_batch_id("batch-123")
-
+
batch_file_object = {
"id": "batch-123",
"input_file_id": unified_file_id, # Batch references this file
"status": "validating",
}
-
+
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="validating",
file_object=batch_file_object,
)
-
+
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
-
- referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
-
+
+ referencing_batches = await managed_files._get_batches_referencing_file(
+ unified_file_id
+ )
+
assert len(referencing_batches) == 1
assert referencing_batches[0]["batch_id"] == unified_batch_id
assert referencing_batches[0]["status"] == "validating"
@@ -204,27 +208,29 @@ async def test_get_batches_referencing_file_finds_batch_with_output_file():
"""Test finding a batch that references the file as output_file_id."""
unified_file_id = _make_unified_file_id("file-output-456")
unified_batch_id = _make_unified_batch_id("batch-456")
-
+
batch_file_object = {
"id": "batch-456",
"input_file_id": "file-input-different",
"output_file_id": unified_file_id, # Batch references this file
"status": "in_progress",
}
-
+
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="in_progress",
file_object=batch_file_object,
)
-
+
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
-
- referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
-
+
+ referencing_batches = await managed_files._get_batches_referencing_file(
+ unified_file_id
+ )
+
assert len(referencing_batches) == 1
assert referencing_batches[0]["status"] == "in_progress"
@@ -234,27 +240,29 @@ async def test_get_batches_referencing_file_ignores_terminal_batches():
"""Test that batches in terminal states are not returned."""
unified_file_id = _make_unified_file_id("file-123")
unified_batch_id = _make_unified_batch_id("batch-completed")
-
+
batch_file_object = {
"id": "batch-completed",
"input_file_id": unified_file_id,
"status": "completed",
}
-
+
# Batch is in terminal state in DB
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="completed", # Terminal state
file_object=batch_file_object,
)
-
+
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[], # Query returns no batches (terminal states filtered out)
)
-
- referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
-
+
+ referencing_batches = await managed_files._get_batches_referencing_file(
+ unified_file_id
+ )
+
assert len(referencing_batches) == 0
@@ -262,26 +270,36 @@ async def test_get_batches_referencing_file_ignores_terminal_batches():
async def test_get_batches_referencing_file_finds_multiple_batches():
"""Test finding multiple batches referencing the same file."""
unified_file_id = _make_unified_file_id("file-shared")
-
+
batch1 = _make_batch_db_record(
unified_object_id=_make_unified_batch_id("batch-1"),
status="validating",
- file_object={"id": "batch-1", "input_file_id": unified_file_id, "status": "validating"},
+ file_object={
+ "id": "batch-1",
+ "input_file_id": unified_file_id,
+ "status": "validating",
+ },
)
-
+
batch2 = _make_batch_db_record(
unified_object_id=_make_unified_batch_id("batch-2"),
status="in_progress",
- file_object={"id": "batch-2", "input_file_id": unified_file_id, "status": "in_progress"},
+ file_object={
+ "id": "batch-2",
+ "input_file_id": unified_file_id,
+ "status": "in_progress",
+ },
)
-
+
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch1, batch2],
)
-
- referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
-
+
+ referencing_batches = await managed_files._get_batches_referencing_file(
+ unified_file_id
+ )
+
assert len(referencing_batches) == 2
statuses = [b["status"] for b in referencing_batches]
assert "validating" in statuses
@@ -300,32 +318,32 @@ async def test_file_deletion_blocked_when_batch_polling_enabled_and_batch_refere
"""
unified_file_id = _make_unified_file_id("file-to-delete")
unified_batch_id = _make_unified_batch_id("batch-active")
-
+
batch_file_object = {
"id": "batch-active",
"input_file_id": unified_file_id,
"status": "validating",
}
-
+
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="validating",
file_object=batch_file_object,
)
-
+
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
-
+
# Mock scheduler with registered batch cost job
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock() # Job exists
-
+
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
with pytest.raises(HTTPException) as exc_info:
await managed_files._check_file_deletion_allowed(unified_file_id)
-
+
assert exc_info.value.status_code == 400
error_detail = exc_info.value.detail
assert "Cannot delete file" in error_detail
@@ -342,28 +360,28 @@ async def test_file_deletion_allowed_when_batch_polling_disabled():
"""
unified_file_id = _make_unified_file_id("file-to-delete")
unified_batch_id = _make_unified_batch_id("batch-active")
-
+
batch_file_object = {
"id": "batch-active",
"input_file_id": unified_file_id,
"status": "validating",
}
-
+
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="validating",
file_object=batch_file_object,
)
-
+
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
-
+
# Mock scheduler without registered job (batch cost tracking disabled)
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = None
-
+
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
# Should not raise an exception
await managed_files._check_file_deletion_allowed(unified_file_id)
@@ -376,16 +394,16 @@ async def test_file_deletion_allowed_when_no_batches_reference_file():
even when batch cost tracking is enabled.
"""
unified_file_id = _make_unified_file_id("file-to-delete")
-
+
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[], # No batches reference this file
)
-
+
# Mock scheduler with registered job (batch cost tracking enabled)
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock()
-
+
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
# Should not raise an exception
await managed_files._check_file_deletion_allowed(unified_file_id)
@@ -398,32 +416,32 @@ async def test_afile_delete_calls_check_deletion_allowed():
"""
unified_file_id = _make_unified_file_id("file-to-delete")
unified_batch_id = _make_unified_batch_id("batch-active")
-
+
batch_file_object = {
"id": "batch-active",
"input_file_id": unified_file_id,
"status": "in_progress",
}
-
+
batch_record = _make_batch_db_record(
unified_object_id=unified_batch_id,
status="in_progress",
file_object=batch_file_object,
)
-
+
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch_record],
)
-
+
# Mock llm_router
mock_router = MagicMock()
mock_router.afile_delete = AsyncMock()
-
+
# Mock scheduler with registered job
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock()
-
+
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
with pytest.raises(HTTPException) as exc_info:
await managed_files.afile_delete(
@@ -431,7 +449,7 @@ async def test_afile_delete_calls_check_deletion_allowed():
litellm_parent_otel_span=None,
llm_router=mock_router,
)
-
+
# Should raise error before calling router delete
assert exc_info.value.status_code == 400
mock_router.afile_delete.assert_not_called()
@@ -444,7 +462,7 @@ async def test_database_limit_respected():
This is a performance optimization - we only fetch what we need.
"""
unified_file_id = _make_unified_file_id("file-shared")
-
+
# Create exactly 10 batches (what DB will return with take=10)
ten_batches = []
for i in range(10):
@@ -454,30 +472,32 @@ async def test_database_limit_respected():
file_object={
"id": f"batch-{i}",
"input_file_id": unified_file_id,
- "status": "validating"
+ "status": "validating",
},
)
ten_batches.append(batch)
-
+
# Mock will return only 10 batches (as DB would with take=10)
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=ten_batches,
)
-
- referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id)
-
+
+ referencing_batches = await managed_files._get_batches_referencing_file(
+ unified_file_id
+ )
+
# Should return all 10 that reference the file
assert len(referencing_batches) == 10
-
+
# Verify error message handles "10+" case (since we got exactly 10, might be more in DB)
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock()
-
+
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
with pytest.raises(HTTPException) as exc_info:
await managed_files._check_file_deletion_allowed(unified_file_id)
-
+
error_detail = exc_info.value.detail
# When we get exactly 10 matches, show "10+" to indicate there might be more
assert "10+ batch(es)" in error_detail
@@ -491,32 +511,40 @@ async def test_error_message_includes_batch_details():
unified_file_id = _make_unified_file_id("file-to-delete")
batch1_id = _make_unified_batch_id("batch-1")
batch2_id = _make_unified_batch_id("batch-2")
-
+
batch1 = _make_batch_db_record(
unified_object_id=batch1_id,
status="validating",
- file_object={"id": "batch-1", "input_file_id": unified_file_id, "status": "validating"},
+ file_object={
+ "id": "batch-1",
+ "input_file_id": unified_file_id,
+ "status": "validating",
+ },
)
-
+
batch2 = _make_batch_db_record(
unified_object_id=batch2_id,
status="in_progress",
- file_object={"id": "batch-2", "output_file_id": unified_file_id, "status": "in_progress"},
+ file_object={
+ "id": "batch-2",
+ "output_file_id": unified_file_id,
+ "status": "in_progress",
+ },
)
-
+
managed_files = _make_managed_files_instance_with_batches(
file_id=unified_file_id,
batches=[batch1, batch2],
)
-
+
# Mock scheduler with registered job
mock_scheduler = MagicMock()
mock_scheduler.get_job.return_value = MagicMock()
-
+
with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler):
with pytest.raises(HTTPException) as exc_info:
await managed_files._check_file_deletion_allowed(unified_file_id)
-
+
error_detail = exc_info.value.detail
assert "2 batch(es)" in error_detail
assert "validating" in error_detail
diff --git a/tests/test_litellm/google_genai/test_google_genai_adapter.py b/tests/test_litellm/google_genai/test_google_genai_adapter.py
index f21564546a8..90ef89699ea 100644
--- a/tests/test_litellm/google_genai/test_google_genai_adapter.py
+++ b/tests/test_litellm/google_genai/test_google_genai_adapter.py
@@ -2,6 +2,7 @@
"""
Test to verify the Google GenAI generate_content adapter functionality
"""
+
import json
import os
import sys
diff --git a/tests/test_litellm/google_genai/test_google_genai_adapter_fixes.py b/tests/test_litellm/google_genai/test_google_genai_adapter_fixes.py
index da56b094d95..8fc3eca9adf 100644
--- a/tests/test_litellm/google_genai/test_google_genai_adapter_fixes.py
+++ b/tests/test_litellm/google_genai/test_google_genai_adapter_fixes.py
@@ -2,6 +2,7 @@
"""
Test to verify the Google GenAI adapter fixes
"""
+
import json
import os
import sys
diff --git a/tests/test_litellm/google_genai/test_google_genai_handler.py b/tests/test_litellm/google_genai/test_google_genai_handler.py
index 0dc218d297b..da1c0bb611b 100644
--- a/tests/test_litellm/google_genai/test_google_genai_handler.py
+++ b/tests/test_litellm/google_genai/test_google_genai_handler.py
@@ -2,6 +2,7 @@
"""
Test to verify the Google GenAI generate_content handler functionality
"""
+
import json
import os
import sys
diff --git a/tests/test_litellm/google_genai/test_google_genai_main.py b/tests/test_litellm/google_genai/test_google_genai_main.py
index 5854e4b55af..1801ee10f30 100644
--- a/tests/test_litellm/google_genai/test_google_genai_main.py
+++ b/tests/test_litellm/google_genai/test_google_genai_main.py
@@ -2,6 +2,7 @@
"""
Test to verify the Google GenAI generate_content adapter functionality
"""
+
import json
import os
import sys
diff --git a/tests/test_litellm/google_genai/test_google_genai_transformation.py b/tests/test_litellm/google_genai/test_google_genai_transformation.py
index 6b0cd500a82..db5de7641d4 100644
--- a/tests/test_litellm/google_genai/test_google_genai_transformation.py
+++ b/tests/test_litellm/google_genai/test_google_genai_transformation.py
@@ -2,6 +2,7 @@
"""
Test to verify the Google GenAI transformation logic for generateContent parameters
"""
+
import os
import sys
diff --git a/tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py b/tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py
index fdf56aedbc9..0767b83e8eb 100644
--- a/tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py
+++ b/tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py
@@ -20,7 +20,6 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanE
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py
index d849582b3c4..78f5259fc58 100644
--- a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py
+++ b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py
@@ -127,16 +127,14 @@ def test_input_validation():
# Create a temporary directory with a test prompt
with tempfile.TemporaryDirectory() as temp_dir:
prompt_file = Path(temp_dir) / "test_validation.prompt"
- prompt_file.write_text(
- """---
+ prompt_file.write_text("""---
input:
schema:
name: string
age: integer
active: boolean
---
-Hello {{name}}, you are {{age}} years old and {'active' if active else 'inactive'}."""
- )
+Hello {{name}}, you are {{age}} years old and {'active' if active else 'inactive'}.""")
manager = PromptManager(prompt_directory=str(temp_dir))
@@ -248,16 +246,14 @@ def test_frontmatter_parsing():
with tempfile.TemporaryDirectory() as temp_dir:
# Test with frontmatter
prompt_with_frontmatter = Path(temp_dir) / "with_frontmatter.prompt"
- prompt_with_frontmatter.write_text(
- """---
+ prompt_with_frontmatter.write_text("""---
model: gpt-4
temperature: 0.8
input:
schema:
topic: string
---
-Write about {{topic}}."""
- )
+Write about {{topic}}.""")
# Test without frontmatter
prompt_without_frontmatter = Path(temp_dir) / "without_frontmatter.prompt"
@@ -370,8 +366,7 @@ def test_prompt_file_to_json_conversion():
# Create a temporary prompt file with frontmatter
with tempfile.TemporaryDirectory() as temp_dir:
prompt_file = Path(temp_dir) / "test_conversion.prompt"
- prompt_file.write_text(
- """---
+ prompt_file.write_text("""---
model: gpt-4
temperature: 0.7
max_tokens: 200
@@ -384,8 +379,7 @@ output:
---
You are an AI assistant. Given the context: {{context}}
-Please respond to: {{user_input}}"""
- )
+Please respond to: {{user_input}}""")
manager = PromptManager()
json_data = manager.prompt_file_to_json(prompt_file)
diff --git a/tests/test_litellm/integrations/rubrik_test_helpers.py b/tests/test_litellm/integrations/rubrik_test_helpers.py
index 1bdb8cb247b..f33d4413ddd 100644
--- a/tests/test_litellm/integrations/rubrik_test_helpers.py
+++ b/tests/test_litellm/integrations/rubrik_test_helpers.py
@@ -5,9 +5,7 @@ from typing import Any, Dict
from litellm.types.utils import GenericGuardrailAPIInputs
-def make_tool_call_dict(
- tc_id: str, name: str, arguments: str = "{}"
-) -> Dict[str, Any]:
+def make_tool_call_dict(tc_id: str, name: str, arguments: str = "{}") -> Dict[str, Any]:
"""Create a tool call dict matching the ChatCompletionMessageToolCall schema."""
return {
"id": tc_id,
diff --git a/tests/test_litellm/integrations/test_openmeter.py b/tests/test_litellm/integrations/test_openmeter.py
index 248b9b34909..6c93c47865e 100644
--- a/tests/test_litellm/integrations/test_openmeter.py
+++ b/tests/test_litellm/integrations/test_openmeter.py
@@ -400,9 +400,7 @@ class TestOpenMeterIntegration:
"model": "gpt-4",
"response_cost": 0.002,
"litellm_call_id": "test-call-id",
- "litellm_params": {
- "metadata": {"user_api_key_user_id": "real-tenant-id"}
- },
+ "litellm_params": {"metadata": {"user_api_key_user_id": "real-tenant-id"}},
}
response_obj = {
@@ -445,9 +443,7 @@ class TestOpenMeterIntegration:
"model": "gpt-4",
"response_cost": 0.002,
"litellm_call_id": "test-call-id",
- "litellm_params": {
- "metadata": {"user_api_key_user_id": "key-user"}
- },
+ "litellm_params": {"metadata": {"user_api_key_user_id": "key-user"}},
}
response_obj = {
diff --git a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py b/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py
index ace9399cf53..80b52651b7b 100644
--- a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py
+++ b/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py
@@ -43,7 +43,6 @@ from litellm.integrations.opentelemetry import (
)
from litellm.proxy._types import UserAPIKeyAuth
-
GUARDRAIL_SPAN_NAME = "guardrail"
PROXY_SPAN_NAME = "Received Proxy Server Request"
diff --git a/tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py b/tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py
index bb035c4c3ee..d7614c970d1 100644
--- a/tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py
+++ b/tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py
@@ -31,7 +31,6 @@ from litellm.types.integrations.prometheus import (
UserAPIKeyLabelValues,
)
-
# ---------------------------------------------------------------------------
# Label / enum wiring
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py b/tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py
index 72a4e80717b..d99f5b62e15 100644
--- a/tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py
+++ b/tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py
@@ -21,7 +21,6 @@ from litellm.types.integrations.prometheus import (
UserAPIKeyLabelValues,
)
-
TOKEN_DETAIL_METRICS = [
"litellm_input_cached_tokens_metric",
"litellm_input_cache_creation_tokens_metric",
diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py
index b2d5225070c..97011edd1d1 100644
--- a/tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py
+++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py
@@ -18,7 +18,6 @@ from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/interactions/test_agents_http_handler.py b/tests/test_litellm/interactions/test_agents_http_handler.py
index 6947503e0bb..2c7efb1bef6 100644
--- a/tests/test_litellm/interactions/test_agents_http_handler.py
+++ b/tests/test_litellm/interactions/test_agents_http_handler.py
@@ -32,7 +32,6 @@ from litellm.types.agents import (
)
from litellm.types.router import GenericLiteLLMParams
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/interactions/test_agents_main_and_utils.py b/tests/test_litellm/interactions/test_agents_main_and_utils.py
index 7c0183d20c6..e2899d78f34 100644
--- a/tests/test_litellm/interactions/test_agents_main_and_utils.py
+++ b/tests/test_litellm/interactions/test_agents_main_and_utils.py
@@ -38,7 +38,6 @@ from litellm.interactions.agents.utils import get_provider_agents_api_config
from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig
from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig
-
_HANDLER_PATH = "litellm.interactions.agents.main.agents_http_handler"
diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_azure_assistant_cost_tracking.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_azure_assistant_cost_tracking.py
index e8bf54f7ffc..3b99167c034 100644
--- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_azure_assistant_cost_tracking.py
+++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_azure_assistant_cost_tracking.py
@@ -3,7 +3,7 @@ Test Azure OpenAI Assistant Features Cost Tracking
Tests cost calculation for Azure's new assistant features:
- File Search (storage-based pricing)
-- Code Interpreter (session-based pricing)
+- Code Interpreter (session-based pricing)
- Computer Use (token-based pricing)
- Vector Store (storage-based pricing)
"""
diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py
index b67ea91bb0b..fc3308d6e36 100644
--- a/tests/test_litellm/litellm_core_utils/test_core_helpers.py
+++ b/tests/test_litellm/litellm_core_utils/test_core_helpers.py
@@ -176,7 +176,11 @@ class TestRedactNestedMatchAndRegexKeys:
{
"sensitiveInformationPolicy": {
"piiEntities": [
- {"type": "NAME", "match": "secret-name", "action": "BLOCKED"}
+ {
+ "type": "NAME",
+ "match": "secret-name",
+ "action": "BLOCKED",
+ }
]
},
"wordPolicy": {
@@ -187,16 +191,22 @@ class TestRedactNestedMatchAndRegexKeys:
"regex": "should-redact-key-named-regex",
}
out = redact_nested_match_and_regex_keys(payload)
- assert out["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][0][
- "match"
- ] == "[REDACTED]"
+ assert (
+ out["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][0][
+ "match"
+ ]
+ == "[REDACTED]"
+ )
assert out["assessments"][0]["wordPolicy"]["customWords"][0]["match"] == (
"[REDACTED]"
)
assert out["regex"] == "[REDACTED]"
- assert payload["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][
- 0
- ]["match"] == "secret-name"
+ assert (
+ payload["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][0][
+ "match"
+ ]
+ == "secret-name"
+ )
def test_passes_through_none_and_str(self):
assert redact_nested_match_and_regex_keys(None) is None
diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py
index b5eb7af88b3..67f87daa172 100644
--- a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py
+++ b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py
@@ -325,6 +325,7 @@ def test_cache_read_input_tokens_retained():
assert usage.cache_read_input_tokens == 11775
assert usage.prompt_tokens_details.cached_tokens == 11775
+
def test_cache_read_input_tokens_retained_genericstreamingchunk():
chunk1 = GenericStreamingChunk(
text="Test1",
@@ -362,6 +363,7 @@ def test_cache_read_input_tokens_retained_genericstreamingchunk():
assert usage.prompt_tokens_details.cached_tokens == 543
+
def test_stream_chunk_builder_litellm_usage_chunks():
"""
Validate ChunkProcessor.calculate_usage uses provided usage fields from streaming chunks
diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
index e95cd656cc4..1409b4fb507 100644
--- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
+++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
@@ -2306,7 +2306,9 @@ def test_chunk_creator_tool_calls_not_dropped_on_finish(
tool_calls=[
ChatCompletionDeltaToolCall(
id="call_abc",
- function=Function(name="get_weather", arguments='{"city":"NYC"}'),
+ function=Function(
+ name="get_weather", arguments='{"city":"NYC"}'
+ ),
type="function",
index=0,
)
diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py
index 60e5a797627..19690da4551 100644
--- a/tests/test_litellm/litellm_core_utils/test_token_counter.py
+++ b/tests/test_litellm/litellm_core_utils/test_token_counter.py
@@ -203,9 +203,7 @@ def test_tokenizers():
try:
llama3_tokenizer = create_pretrained_tokenizer("Xenova/llama-3-tokenizer")
except Exception as e:
- pytest.skip(
- f"custom tokenizer download failed (HF hub unreachable): {e}"
- )
+ pytest.skip(f"custom tokenizer download failed (HF hub unreachable): {e}")
llama3_tokens_2 = token_counter(
custom_tokenizer=llama3_tokenizer, text=sample_text
)
diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py
index b9bda07336f..d88260629ee 100644
--- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py
+++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py
@@ -22,7 +22,6 @@ from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming
_parse_sse_events,
)
-
# ---------------------------------------------------------------------------
# Helpers to build SSE byte payloads
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py
index e44413cf837..65320446524 100644
--- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py
+++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py
@@ -2,7 +2,6 @@ import os
import sys
from typing import List
-
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import (
diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_auto_summary_messages.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_auto_summary_messages.py
index 07c0012b04d..dd5d1341713 100644
--- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_auto_summary_messages.py
+++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_auto_summary_messages.py
@@ -31,13 +31,16 @@ def _call_handler_and_capture_optional_params(thinking=None, **extra_kwargs):
"""
captured = {}
- with patch(
- "litellm.llms.anthropic.experimental_pass_through.messages.handler."
- "base_llm_http_handler"
- ) as mock_handler, patch(
- "litellm.llms.anthropic.experimental_pass_through.messages.handler."
- "ProviderConfigManager"
- ) as mock_pcm:
+ with (
+ patch(
+ "litellm.llms.anthropic.experimental_pass_through.messages.handler."
+ "base_llm_http_handler"
+ ) as mock_handler,
+ patch(
+ "litellm.llms.anthropic.experimental_pass_through.messages.handler."
+ "ProviderConfigManager"
+ ) as mock_pcm,
+ ):
# Make get_provider_anthropic_messages_config return a non-None config
# so the handler takes the native Anthropic path
mock_pcm.get_provider_anthropic_messages_config.return_value = MagicMock()
@@ -119,8 +122,9 @@ class TestReasoningAutoSummaryMessages:
def test_env_var_enables_auto_summary(self):
"""LITELLM_REASONING_AUTO_SUMMARY=true env var enables the feature."""
- with patch.object(litellm, "reasoning_auto_summary", False), patch.dict(
- os.environ, {"LITELLM_REASONING_AUTO_SUMMARY": "true"}
+ with (
+ patch.object(litellm, "reasoning_auto_summary", False),
+ patch.dict(os.environ, {"LITELLM_REASONING_AUTO_SUMMARY": "true"}),
):
params = _call_handler_and_capture_optional_params(
thinking={"type": "adaptive", "budget_tokens": 5000}
diff --git a/tests/test_litellm/llms/anthropic/test_message_sanitization.py b/tests/test_litellm/llms/anthropic/test_message_sanitization.py
index 79ed321d0ee..37e82a838af 100644
--- a/tests/test_litellm/llms/anthropic/test_message_sanitization.py
+++ b/tests/test_litellm/llms/anthropic/test_message_sanitization.py
@@ -363,7 +363,9 @@ class TestMessageSanitization:
assert len(result) == 1
assert result[0]["role"] == "user"
text_blocks = [
- b for b in result[0]["content"] if isinstance(b, dict) and b.get("type") == "text"
+ b
+ for b in result[0]["content"]
+ if isinstance(b, dict) and b.get("type") == "text"
]
assert len(text_blocks) == 3
# No text block may be empty — that's the contract Anthropic enforces.
@@ -398,7 +400,9 @@ class TestMessageSanitization:
assert len(result) == 1
text_blocks = [
- b for b in result[0]["content"] if isinstance(b, dict) and b.get("type") == "text"
+ b
+ for b in result[0]["content"]
+ if isinstance(b, dict) and b.get("type") == "text"
]
assert len(text_blocks) == 3
assert text_blocks[0]["text"] == "real content"
diff --git a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py
index 4638bc4df0f..f3488f29a75 100644
--- a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py
+++ b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py
@@ -189,8 +189,7 @@ async def test_construct_url_forwards_transcription_intent_ga_without_model_quer
)
assert url == (
- "wss://my-endpoint.openai.azure.com/openai/v1/realtime"
- "?intent=transcription"
+ "wss://my-endpoint.openai.azure.com/openai/v1/realtime" "?intent=transcription"
)
diff --git a/tests/test_litellm/llms/azure/test_azure_cost_calculation.py b/tests/test_litellm/llms/azure/test_azure_cost_calculation.py
index 53c91032b34..89fff013360 100644
--- a/tests/test_litellm/llms/azure/test_azure_cost_calculation.py
+++ b/tests/test_litellm/llms/azure/test_azure_cost_calculation.py
@@ -8,7 +8,6 @@ import litellm
from litellm.llms.azure.cost_calculation import cost_per_token
from litellm.types.utils import Usage
-
# Register a test model with tier-specific pricing
TEST_MODEL = "test-azure-gpt-4.1"
TEST_MODEL_COST = {
diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py b/tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py
index 20260c744f8..b9316fb9af1 100644
--- a/tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py
+++ b/tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py
@@ -459,18 +459,21 @@ class TestAzureAIServiceTierCostCalculation:
@pytest.fixture(autouse=True)
def register_test_model(self):
import litellm
- litellm.register_model(model_cost={
- "test-azure-ai-model": {
- "input_cost_per_token": 0.001,
- "output_cost_per_token": 0.002,
- "input_cost_per_token_priority": 0.01,
- "output_cost_per_token_priority": 0.02,
- "input_cost_per_token_flex": 0.0005,
- "output_cost_per_token_flex": 0.001,
- "litellm_provider": "azure_ai",
- "max_tokens": 8192,
+
+ litellm.register_model(
+ model_cost={
+ "test-azure-ai-model": {
+ "input_cost_per_token": 0.001,
+ "output_cost_per_token": 0.002,
+ "input_cost_per_token_priority": 0.01,
+ "output_cost_per_token_priority": 0.02,
+ "input_cost_per_token_flex": 0.0005,
+ "output_cost_per_token_flex": 0.001,
+ "litellm_provider": "azure_ai",
+ "max_tokens": 8192,
+ }
}
- })
+ )
def test_service_tier_priority_higher_cost(self):
"""Priority tier should cost more than standard for azure_ai."""
diff --git a/tests/test_litellm/llms/base_llm/test_managed_resource_isolation.py b/tests/test_litellm/llms/base_llm/test_managed_resource_isolation.py
index b5fcd9d8219..1efd10e921b 100644
--- a/tests/test_litellm/llms/base_llm/test_managed_resource_isolation.py
+++ b/tests/test_litellm/llms/base_llm/test_managed_resource_isolation.py
@@ -10,7 +10,6 @@ from litellm.llms.base_llm.managed_resources.isolation import (
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
-
# ---------------------------------------------------------------------------
# build_owner_filter
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/llms/bedrock/batches/test_batch_metadata_sanitization.py b/tests/test_litellm/llms/bedrock/batches/test_batch_metadata_sanitization.py
index 8de47331614..e5030d84552 100644
--- a/tests/test_litellm/llms/bedrock/batches/test_batch_metadata_sanitization.py
+++ b/tests/test_litellm/llms/bedrock/batches/test_batch_metadata_sanitization.py
@@ -60,7 +60,9 @@ class TestGetOpenaiCompatibleBatchMetadata:
def test_non_dict_input_returns_empty_dict(self):
assert BedrockBatchesConfig._get_openai_compatible_batch_metadata(None) == {}
- assert BedrockBatchesConfig._get_openai_compatible_batch_metadata("string") == {}
+ assert (
+ BedrockBatchesConfig._get_openai_compatible_batch_metadata("string") == {}
+ )
assert BedrockBatchesConfig._get_openai_compatible_batch_metadata(123) == {}
def test_empty_dict_returns_empty_dict(self):
@@ -82,7 +84,9 @@ class TestGetOpenaiCompatibleBatchMetadata:
# All values must be strings
for key, value in result.items():
- assert isinstance(value, str), f"metadata[{key!r}] is {type(value)}, not str"
+ assert isinstance(
+ value, str
+ ), f"metadata[{key!r}] is {type(value)}, not str"
# Excluded keys
assert "standard_logging_guardrail_information" not in result
diff --git a/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py b/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py
index 744ed50dbcb..11f9e76a5d1 100644
--- a/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py
+++ b/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py
@@ -75,7 +75,9 @@ class TestExtractConverseTexts:
"system": [{"text": "sys text"}],
"messages": [{"role": "user", "content": [{"text": "user text"}]}],
}
- texts, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False)
+ texts, holders = _extract_converse_texts(
+ body, skip_system=False, skip_tool=False
+ )
assert texts == ["sys text", "user text"]
assert holders[0] == (body["system"][0], "text")
assert holders[1] == (body["messages"][0]["content"][0], "text")
@@ -85,7 +87,9 @@ class TestExtractConverseTexts:
"system": [{"text": "sys text"}],
"messages": [{"role": "user", "content": [{"text": "user text"}]}],
}
- texts, holders = _extract_converse_texts(body, skip_system=True, skip_tool=False)
+ texts, holders = _extract_converse_texts(
+ body, skip_system=True, skip_tool=False
+ )
assert texts == ["user text"]
assert holders == [(body["messages"][0]["content"][0], "text")]
@@ -130,7 +134,9 @@ class TestExtractConverseTexts:
}
]
}
- texts, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False)
+ texts, holders = _extract_converse_texts(
+ body, skip_system=False, skip_tool=False
+ )
assert texts == ["hello", "blocked tool text", "blocked json value"]
tool_content = body["messages"][0]["content"][1]["toolResult"]["content"]
assert holders[1] == (tool_content[0], "text")
@@ -153,7 +159,9 @@ class TestExtractConverseTexts:
}
]
}
- texts, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False)
+ texts, holders = _extract_converse_texts(
+ body, skip_system=False, skip_tool=False
+ )
assert texts == ["blocked input value"]
tool_use_input = body["messages"][0]["content"][0]["toolUse"]["input"]
assert holders[0] == (tool_use_input, "query")
@@ -253,7 +261,10 @@ class TestWriteBackTexts:
}
_, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False)
_write_back_texts(["masked"], holders)
- assert body["messages"][0]["content"][0]["toolResult"]["content"][0]["text"] == "masked"
+ assert (
+ body["messages"][0]["content"][0]["toolResult"]["content"][0]["text"]
+ == "masked"
+ )
def test_extra_non_text_fields_untouched(self):
body = {
@@ -278,14 +289,14 @@ class TestWriteBackTexts:
_, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False)
_write_back_texts(["replaced"], holders)
assert body["messages"][0]["content"][0]["text"] == "replaced"
- assert body["messages"][0]["content"][1] == original["messages"][0]["content"][1]
+ assert (
+ body["messages"][0]["content"][1] == original["messages"][0]["content"][1]
+ )
assert body["inferenceConfig"] == original["inferenceConfig"]
def test_fewer_guardrailed_texts_logs_warning(self, monkeypatch):
body = {
- "messages": [
- {"role": "user", "content": [{"text": "a"}, {"text": "b"}]}
- ]
+ "messages": [{"role": "user", "content": [{"text": "a"}, {"text": "b"}]}]
}
_, holders = _extract_converse_texts(body, skip_system=False, skip_tool=False)
assert len(holders) == 2
@@ -298,7 +309,9 @@ class TestWriteBackTexts:
_write_back_texts(["masked"], holders)
- assert warnings, "mismatched guardrail output count must not be silently dropped"
+ assert (
+ warnings
+ ), "mismatched guardrail output count must not be silently dropped"
assert body["messages"][0]["content"][0]["text"] == "masked"
assert body["messages"][0]["content"][1]["text"] == "b"
@@ -326,7 +339,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
data = _converse_data()
guardrail = _make_guardrail({"texts": ["[REDACTED]", "[REDACTED]"]})
- result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ result = await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
body = result["data"]
assert body["system"][0]["text"] == "[REDACTED]"
@@ -347,7 +362,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlocked("Blocked"))
with pytest.raises(GuardrailBlocked):
- await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
@pytest.mark.asyncio
async def test_tool_result_text_scanned_and_masked(self):
@@ -365,7 +382,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
{"texts": ["You are helpful.", "Hello world", "[REDACTED]"]}
)
- result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ result = await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
assert "My SSN is 123-45-6789" in sent_texts
@@ -390,7 +409,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
{"texts": ["You are helpful.", "Hello world", "[REDACTED]"]}
)
- result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ result = await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
assert "SSN 123-45-6789" in sent_texts
@@ -409,7 +430,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
{"texts": ["You are helpful.", "Hello world", "[REDACTED]"]}
)
- result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ result = await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
assert "email john@example.com" in sent_texts
@@ -431,7 +454,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlocked("Blocked"))
with pytest.raises(GuardrailBlocked):
- await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
assert "blocked content" in sent_texts
@@ -454,10 +479,20 @@ class TestBedrockPassthroughGuardrailHandlerInput:
]
}
guardrail = _make_guardrail(
- {"texts": ["You are helpful.", "Hello world", "lookup", "[REDACTED]", "object"]}
+ {
+ "texts": [
+ "You are helpful.",
+ "Hello world",
+ "lookup",
+ "[REDACTED]",
+ "object",
+ ]
+ }
)
- result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ result = await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
assert "email john@example.com" in sent_texts
@@ -479,7 +514,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlocked("Blocked"))
with pytest.raises(GuardrailBlocked):
- await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
assert "blocked content" in sent_texts
@@ -495,7 +532,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
{"texts": ["You are helpful.", "Hello world", "[REDACTED]"]}
)
- result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ result = await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
assert "ssn 123-45-6789" in sent_texts
@@ -533,7 +572,9 @@ class TestBedrockPassthroughGuardrailHandlerInput:
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlocked("Blocked"))
with pytest.raises(GuardrailBlocked):
- await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
@pytest.mark.asyncio
async def test_missing_messages_field_skips(self):
@@ -580,7 +621,9 @@ class TestBedrockPassthroughGuardrailHandlerOutput:
response = self._converse_response("Model reply")
guardrail = _make_guardrail({"texts": ["Model reply"]})
- await handler.process_output_response(response=response, guardrail_to_apply=guardrail)
+ await handler.process_output_response(
+ response=response, guardrail_to_apply=guardrail
+ )
call_args = guardrail.apply_guardrail.call_args
assert call_args.kwargs["input_type"] == "response"
@@ -592,13 +635,17 @@ class TestBedrockPassthroughGuardrailHandlerOutput:
response = self._converse_response("Bad content")
guardrail = _make_guardrail({"texts": ["[MASKED]"]})
- result = await handler.process_output_response(response=response, guardrail_to_apply=guardrail)
+ result = await handler.process_output_response(
+ response=response, guardrail_to_apply=guardrail
+ )
assert result["output"]["message"]["content"][0]["text"] == "[MASKED]"
assert result["stopReason"] == "end_turn"
@pytest.mark.asyncio
- async def test_response_guardrail_returning_no_texts_preserves_output(self, monkeypatch):
+ async def test_response_guardrail_returning_no_texts_preserves_output(
+ self, monkeypatch
+ ):
"""A guardrail that returns no texts must leave the response untouched and
not warn, mirroring the request path's empty-result guard."""
handler = BedrockPassthroughGuardrailHandler()
@@ -611,7 +658,9 @@ class TestBedrockPassthroughGuardrailHandlerOutput:
lambda *args, **kwargs: warnings.append(args),
)
- result = await handler.process_output_response(response=response, guardrail_to_apply=guardrail)
+ result = await handler.process_output_response(
+ response=response, guardrail_to_apply=guardrail
+ )
assert not warnings
assert result["output"]["message"]["content"][0]["text"] == "Model reply"
@@ -648,9 +697,7 @@ class TestBedrockPassthroughGuardrailHandlerOutput:
},
"stopReason": "end_turn",
}
- guardrail = _make_guardrail(
- {"texts": ["[V]", "[REASON]", "[INPUT]"]}
- )
+ guardrail = _make_guardrail({"texts": ["[V]", "[REASON]", "[INPUT]"]})
result = await handler.process_output_response(
response=response, guardrail_to_apply=guardrail
@@ -680,7 +727,10 @@ class TestBedrockPassthroughGuardrailHandlerOutput:
"citationsContent": {
"content": [{"text": "Contact john@example.com"}],
"citations": [
- {"source": "https://example.com", "title": "Example"}
+ {
+ "source": "https://example.com",
+ "title": "Example",
+ }
],
}
}
@@ -707,7 +757,9 @@ class TestBedrockPassthroughGuardrailHandlerOutput:
handler = BedrockPassthroughGuardrailHandler()
guardrail = _make_guardrail({"texts": []})
- result = await handler.process_output_response(response="raw string", guardrail_to_apply=guardrail)
+ result = await handler.process_output_response(
+ response="raw string", guardrail_to_apply=guardrail
+ )
assert result == "raw string"
guardrail.apply_guardrail.assert_not_called()
@@ -718,7 +770,9 @@ class TestBedrockPassthroughGuardrailHandlerOutput:
response = {"stopReason": "end_turn"}
guardrail = _make_guardrail({"texts": []})
- result = await handler.process_output_response(response=response, guardrail_to_apply=guardrail)
+ result = await handler.process_output_response(
+ response=response, guardrail_to_apply=guardrail
+ )
assert result == {"stopReason": "end_turn"}
guardrail.apply_guardrail.assert_not_called()
@@ -753,7 +807,11 @@ def _build_event_stream_frame(event_type: str, payload: dict) -> bytes:
name_b = name.encode()
value_b = value.encode()
return (
- struct.pack("!B", len(name_b)) + name_b + struct.pack("!B", 7) + struct.pack("!H", len(value_b)) + value_b
+ struct.pack("!B", len(name_b))
+ + name_b
+ + struct.pack("!B", 7)
+ + struct.pack("!H", len(value_b))
+ + value_b
)
headers_bytes = (
@@ -788,7 +846,10 @@ class TestDeAnonymizeConverseStream:
_build_event_stream_frame("messageStart", {"role": "assistant"})
+ _build_event_stream_frame(
"contentBlockDelta",
- {"contentBlockIndex": 0, "delta": {"text": " works at "}},
+ {
+ "contentBlockIndex": 0,
+ "delta": {"text": " works at "},
+ },
)
+ _build_event_stream_frame("contentBlockStop", {"contentBlockIndex": 0})
+ _build_event_stream_frame("messageStop", {"stopReason": "end_turn"})
@@ -847,7 +908,10 @@ class TestDeAnonymizeConverseStream:
}
async def mock_hook(data, user_api_key_dict, response):
- assert response["output"]["message"]["content"][0]["text"] == " called."
+ assert (
+ response["output"]["message"]["content"][0]["text"]
+ == " called."
+ )
return de_anon_response
result = await BedrockPassthroughGuardrailHandler.de_anonymize_event_stream(
@@ -1061,7 +1125,10 @@ class TestDeAnonymizeConverseStream:
"""toolUse.input deltas carry model-generated tool arguments and must be guardrailed instead of being forwarded raw."""
stream_bytes = _build_event_stream_frame(
"contentBlockDelta",
- {"contentBlockIndex": 0, "delta": {"toolUse": {"input": '{"q":""}'}}},
+ {
+ "contentBlockIndex": 0,
+ "delta": {"toolUse": {"input": '{"q":""}'}},
+ },
)
result = await BedrockPassthroughGuardrailHandler.de_anonymize_event_stream(
@@ -1086,7 +1153,9 @@ class TestDeAnonymizeConverseStream:
"delta": {
"citationsContent": {
"content": [{"text": "Contact "}],
- "citations": [{"source": "https://example.com", "title": "Example"}],
+ "citations": [
+ {"source": "https://example.com", "title": "Example"}
+ ],
}
},
},
@@ -1127,7 +1196,10 @@ class TestDeAnonymizeConverseStream:
{"contentBlockIndex": 0, "delta": {"text": "Hi "}},
) + _build_event_stream_frame(
"contentBlockDelta",
- {"contentBlockIndex": 1, "delta": {"reasoningContent": {"text": "works at "}}},
+ {
+ "contentBlockIndex": 1,
+ "delta": {"reasoningContent": {"text": "works at "}},
+ },
)
result = await BedrockPassthroughGuardrailHandler.de_anonymize_event_stream(
@@ -1148,7 +1220,10 @@ class TestDeAnonymizeConverseStream:
"""A reasoning delta carrying only a signature has no guardrailable text; it must be forwarded untouched and the guardrail must not run."""
stream_bytes = _build_event_stream_frame(
"contentBlockDelta",
- {"contentBlockIndex": 0, "delta": {"reasoningContent": {"signature": "sig"}}},
+ {
+ "contentBlockIndex": 0,
+ "delta": {"reasoningContent": {"signature": "sig"}},
+ },
)
hook_spy = AsyncMock()
diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py
index 3f91f6ac26e..d709bd54979 100644
--- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py
+++ b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py
@@ -880,7 +880,11 @@ def test_different_roles_without_session_names_should_not_share_cache():
},
),
],
- ids=["no_region_or_endpoint", "bedrock_region_ignored_for_sts", "explicit_sts_endpoint"],
+ ids=[
+ "no_region_or_endpoint",
+ "bedrock_region_ignored_for_sts",
+ "explicit_sts_endpoint",
+ ],
)
def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs):
"""
@@ -1272,7 +1276,11 @@ def test_sts_endpoint_region_matches_bedrock_region_param():
},
),
],
- ids=["no_region_or_endpoint", "bedrock_region_ignored_for_sts", "explicit_sts_endpoint"],
+ ids=[
+ "no_region_or_endpoint",
+ "bedrock_region_ignored_for_sts",
+ "explicit_sts_endpoint",
+ ],
)
def test_explicit_credentials_used_when_provided(role_kwargs, expected_client_kwargs):
"""
diff --git a/tests/test_litellm/llms/crusoe/test_crusoe.py b/tests/test_litellm/llms/crusoe/test_crusoe.py
index 0a05126919a..32a2f174f3d 100644
--- a/tests/test_litellm/llms/crusoe/test_crusoe.py
+++ b/tests/test_litellm/llms/crusoe/test_crusoe.py
@@ -39,7 +39,10 @@ def test_crusoe_dynamic_config_env_vars():
with patch.dict(
os.environ,
- {"CRUSOE_API_KEY": "test-key", "CRUSOE_API_BASE": "https://custom.crusoe.com/v1"},
+ {
+ "CRUSOE_API_KEY": "test-key",
+ "CRUSOE_API_BASE": "https://custom.crusoe.com/v1",
+ },
):
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
@@ -69,7 +72,9 @@ def test_crusoe_supported_params():
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
config = create_config_class(JSONProviderRegistry.get("crusoe"))()
- params = config.get_supported_openai_params(model="meta-llama/Llama-3.3-70B-Instruct")
+ params = config.get_supported_openai_params(
+ model="meta-llama/Llama-3.3-70B-Instruct"
+ )
assert isinstance(params, list)
assert len(params) > 0
@@ -91,7 +96,9 @@ def test_crusoe_param_mapping_max_completion_tokens():
drop_params=False,
)
- assert "max_tokens" in optional_params, "max_completion_tokens should be mapped to max_tokens"
+ assert (
+ "max_tokens" in optional_params
+ ), "max_completion_tokens should be mapped to max_tokens"
assert optional_params["max_tokens"] == 1024
assert "max_completion_tokens" not in optional_params
diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py
index 5515e6ce815..660ab88c2b2 100644
--- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py
+++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py
@@ -158,4 +158,7 @@ def test_socket_factory_uses_tcp_keepalive_when_keepidle_unavailable(monkeypatch
assert (
setsockopt_calls[(socket.IPPROTO_TCP, fake_socket_module.TCP_KEEPALIVE)] == 60
)
- assert (socket.IPPROTO_TCP, getattr(socket, "TCP_KEEPIDLE", -1)) not in setsockopt_calls
+ assert (
+ socket.IPPROTO_TCP,
+ getattr(socket, "TCP_KEEPIDLE", -1),
+ ) not in setsockopt_calls
diff --git a/tests/test_litellm/llms/custom_httpx/test_mock_transport.py b/tests/test_litellm/llms/custom_httpx/test_mock_transport.py
index c2d4e146428..b3b5f82f02f 100644
--- a/tests/test_litellm/llms/custom_httpx/test_mock_transport.py
+++ b/tests/test_litellm/llms/custom_httpx/test_mock_transport.py
@@ -10,7 +10,6 @@ import pytest
from litellm.llms.custom_httpx.mock_transport import MockOpenAITransport
-
# ---------------------------------------------------------------------------
# Non-streaming
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/llms/databricks/test_databricks_e2e.py b/tests/test_litellm/llms/databricks/test_databricks_e2e.py
index 669f9e94639..90846f28920 100644
--- a/tests/test_litellm/llms/databricks/test_databricks_e2e.py
+++ b/tests/test_litellm/llms/databricks/test_databricks_e2e.py
@@ -17,13 +17,13 @@ Purpose:
LiteLLM Integration Tests:
This test file includes tests for different ways of calling Databricks via LiteLLM:
-
+
1. LiteLLM SDK Direct - Using litellm.completion() with user_agent parameter
2. LangChain + LiteLLM - Using ChatLiteLLM wrapper (requires langchain-community)
3. LiteLLM Async - Using litellm.acompletion() async API
4. LiteLLM Streaming - Using litellm.completion() with stream=True
5. LiteLLM Embedding - Using litellm.embedding() with user_agent parameter
-
+
All tests use the CUSTOM_USER_AGENT value from the config file and call
Databricks endpoints through LiteLLM's unified interface.
@@ -31,7 +31,7 @@ Prerequisites:
- Valid Databricks workspace access
- Configured credentials (OAuth Service Principal, PAT, or Databricks CLI)
- Access to serving endpoints (e.g., databricks-gpt-oss-120b)
-
+
Optional Dependencies (for LiteLLM integration tests):
- pip install langchain-litellm # For LangChain tests (recommended)
diff --git a/tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py b/tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py
index d59ab975ef2..a950b4a5f2c 100644
--- a/tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py
+++ b/tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py
@@ -53,7 +53,7 @@ def test_file():
)
def test_audio_file_handling(fixture_name, request):
handler = DeepgramAudioTranscriptionConfig()
- (audio_file, expected_output) = request.getfixturevalue(fixture_name)
+ audio_file, expected_output = request.getfixturevalue(fixture_name)
result = handler.transform_audio_transcription_request(
model="deepseek-audio-transcription",
audio_file=audio_file,
diff --git a/tests/test_litellm/llms/fastcrw/search/test_transformation.py b/tests/test_litellm/llms/fastcrw/search/test_transformation.py
index adf8fec087c..9b71e881ae1 100644
--- a/tests/test_litellm/llms/fastcrw/search/test_transformation.py
+++ b/tests/test_litellm/llms/fastcrw/search/test_transformation.py
@@ -79,7 +79,9 @@ def test_validate_environment_missing_key_raises():
def test_get_complete_url_default_base():
with patch.dict(os.environ, {}, clear=True):
- assert _config().get_complete_url(None, {}) == "https://fastcrw.com/api/v1/search"
+ assert (
+ _config().get_complete_url(None, {}) == "https://fastcrw.com/api/v1/search"
+ )
def test_get_complete_url_appends_search():
@@ -100,7 +102,9 @@ def test_get_complete_url_reads_env_base():
with patch.dict(
os.environ, {"CRW_API_BASE": "https://env-base.local/v1"}, clear=True
):
- assert _config().get_complete_url(None, {}) == "https://env-base.local/v1/search"
+ assert (
+ _config().get_complete_url(None, {}) == "https://env-base.local/v1/search"
+ )
def test_transform_search_request_basic():
diff --git a/tests/test_litellm/llms/gemini/test_gemini_tts.py b/tests/test_litellm/llms/gemini/test_gemini_tts.py
index 98f3ac0f4e5..53b433ae06a 100644
--- a/tests/test_litellm/llms/gemini/test_gemini_tts.py
+++ b/tests/test_litellm/llms/gemini/test_gemini_tts.py
@@ -367,7 +367,6 @@ class TestGeminiTTSSpeechConfigInRequestBody:
assert "responseModalities" in generation_config
assert "AUDIO" in generation_config["responseModalities"]
-
@pytest.mark.parametrize(
"model,custom_llm_provider",
[
diff --git a/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py b/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py
index 6c846a90c71..7c787f20b8b 100644
--- a/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py
+++ b/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py
@@ -255,12 +255,19 @@ class TestGitHubCopilotAuthenticator:
"user_code": "UC",
"verification_uri": "https://example.com",
}
- with patch.dict(os.environ, {"GITHUB_COPILOT_DEVICE_CODE_URL": custom_url}), \
- patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client):
+ with (
+ patch.dict(os.environ, {"GITHUB_COPILOT_DEVICE_CODE_URL": custom_url}),
+ patch(
+ "litellm.llms.github_copilot.authenticator._get_httpx_client",
+ return_value=mock_client,
+ ),
+ ):
authenticator._get_device_code()
assert mock_client.post.call_args[0][0] == custom_url
- def test_get_device_code_with_custom_client_id(self, authenticator, mock_http_client):
+ def test_get_device_code_with_custom_client_id(
+ self, authenticator, mock_http_client
+ ):
"""GITHUB_COPILOT_CLIENT_ID env var must appear as client_id in the device-code request body."""
mock_client, mock_response = mock_http_client
custom_id = "custom_client_id"
@@ -269,30 +276,49 @@ class TestGitHubCopilotAuthenticator:
"user_code": "UC",
"verification_uri": "https://example.com",
}
- with patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), \
- patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client):
+ with (
+ patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}),
+ patch(
+ "litellm.llms.github_copilot.authenticator._get_httpx_client",
+ return_value=mock_client,
+ ),
+ ):
authenticator._get_device_code()
assert mock_client.post.call_args[1]["json"]["client_id"] == custom_id
- def test_poll_for_access_token_with_custom_url(self, authenticator, mock_http_client):
+ def test_poll_for_access_token_with_custom_url(
+ self, authenticator, mock_http_client
+ ):
"""GITHUB_COPILOT_ACCESS_TOKEN_URL env var must be used by _poll_for_access_token at call time."""
mock_client, mock_response = mock_http_client
custom_url = "https://custom.example.com/token"
mock_response.json.return_value = {"access_token": "tok"}
- with patch.dict(os.environ, {"GITHUB_COPILOT_ACCESS_TOKEN_URL": custom_url}), \
- patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \
- patch("time.sleep"):
+ with (
+ patch.dict(os.environ, {"GITHUB_COPILOT_ACCESS_TOKEN_URL": custom_url}),
+ patch(
+ "litellm.llms.github_copilot.authenticator._get_httpx_client",
+ return_value=mock_client,
+ ),
+ patch("time.sleep"),
+ ):
authenticator._poll_for_access_token("dc")
assert mock_client.post.call_args[0][0] == custom_url
- def test_poll_for_access_token_with_custom_client_id(self, authenticator, mock_http_client):
+ def test_poll_for_access_token_with_custom_client_id(
+ self, authenticator, mock_http_client
+ ):
"""GITHUB_COPILOT_CLIENT_ID env var must appear as client_id in the polling request body."""
mock_client, mock_response = mock_http_client
custom_id = "custom_client_id"
mock_response.json.return_value = {"access_token": "tok"}
- with patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), \
- patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \
- patch("time.sleep"):
+ with (
+ patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}),
+ patch(
+ "litellm.llms.github_copilot.authenticator._get_httpx_client",
+ return_value=mock_client,
+ ),
+ patch("time.sleep"),
+ ):
authenticator._poll_for_access_token("dc")
assert mock_client.post.call_args[1]["json"]["client_id"] == custom_id
@@ -301,9 +327,13 @@ class TestGitHubCopilotAuthenticator:
mock_client, mock_response = mock_http_client
custom_url = "https://custom.example.com/api-key"
mock_response.json.return_value = {"token": "api-tok", "expires_at": 9999999999}
- with patch.dict(os.environ, {"GITHUB_COPILOT_API_KEY_URL": custom_url}), \
- patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \
- patch.object(authenticator, "get_access_token", return_value="access-tok"):
+ with (
+ patch.dict(os.environ, {"GITHUB_COPILOT_API_KEY_URL": custom_url}),
+ patch(
+ "litellm.llms.github_copilot.authenticator._get_httpx_client",
+ return_value=mock_client,
+ ),
+ patch.object(authenticator, "get_access_token", return_value="access-tok"),
+ ):
authenticator._refresh_api_key()
assert mock_client.get.call_args[0][0] == custom_url
-
diff --git a/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py b/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py
index 8a072fa5097..9e1f0b8a344 100644
--- a/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py
+++ b/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py
@@ -11,7 +11,6 @@ sys.path.insert(
import litellm
import pytest
-
MOCK_EMBEDDING_RESPONSE = [[0.1, 0.2, 0.3, 0.4, 0.5]]
diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py
index 3ee53bb46cd..2f9b9b42631 100644
--- a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py
+++ b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py
@@ -44,6 +44,27 @@ async def test_mistral_chat_transformation():
class TestMistralReasoningSupport:
"""Test suite for Mistral Magistral reasoning functionality."""
+ def test_mistral_strips_metadata_from_messages(self):
+ from litellm.llms.mistral.chat.transformation import MistralConfig
+
+ config = MistralConfig()
+ messages = [
+ {"role": "user", "content": "hello"},
+ {
+ "role": "assistant",
+ "content": "hi",
+ "metadata": {
+ "tool_outputs_trimmed": True,
+ "trimmed_by": "async_context_compression",
+ },
+ },
+ ]
+ result = config._transform_messages(messages, model="mistral-large-latest")
+ for msg in result:
+ assert (
+ "metadata" not in msg
+ ), "metadata should be stripped before sending to Mistral"
+
def test_get_supported_openai_params_magistral_model(self):
"""Test that magistral models support reasoning parameters."""
mistral_config = MistralConfig()
diff --git a/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py b/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py
index 2767deae176..c3d12f371f3 100644
--- a/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py
+++ b/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py
@@ -122,7 +122,10 @@ class TestModelScopeConfig:
"role": "user",
"content": [
{"type": "text", "text": "What is this?"},
- {"type": "image_url", "image_url": {"url": "https://example.com/img.png"}},
+ {
+ "type": "image_url",
+ "image_url": {"url": "https://example.com/img.png"},
+ },
],
}
]
@@ -173,7 +176,10 @@ class TestModelScopeConfig:
"role": "user",
"content": [
{"type": "text", "text": "Describe this image"},
- {"type": "image_url", "image_url": {"url": "https://example.com/photo.jpg"}},
+ {
+ "type": "image_url",
+ "image_url": {"url": "https://example.com/photo.jpg"},
+ },
],
},
]
@@ -303,11 +309,17 @@ class TestModelScopeConfig:
"finish_reason": "stop",
}
],
- "usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6},
+ "usage": {
+ "prompt_tokens": 5,
+ "completion_tokens": 1,
+ "total_tokens": 6,
+ },
},
)
- respx_mock.post(f"{api_base}/chat/completions").mock(side_effect=capture_request)
+ respx_mock.post(f"{api_base}/chat/completions").mock(
+ side_effect=capture_request
+ )
response = completion(
model=f"modelscope/{DEFAULT_MODEL}",
@@ -358,11 +370,17 @@ class TestModelScopeConfig:
"finish_reason": "stop",
}
],
- "usage": {"prompt_tokens": 100, "completion_tokens": 8, "total_tokens": 108},
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 8,
+ "total_tokens": 108,
+ },
},
)
- respx_mock.post(f"{api_base}/chat/completions").mock(side_effect=capture_request)
+ respx_mock.post(f"{api_base}/chat/completions").mock(
+ side_effect=capture_request
+ )
response = completion(
model=f"modelscope/{DEFAULT_MODEL}",
diff --git a/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py b/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py
index 95ade4290e9..536a1f2d2a8 100644
--- a/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py
+++ b/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py
@@ -9,7 +9,9 @@ import os
import sys
from unittest.mock import patch
-sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path
+sys.path.insert(
+ 0, os.path.abspath("../../../../..")
+) # Adds the parent directory to the system path
import pytest
@@ -265,7 +267,10 @@ class TestMoonshotConfig:
assert result["messages"][0]["role"] == "user"
assert result["messages"][0]["content"] == "What's the weather like?"
assert result["messages"][1]["role"] == "user"
- assert result["messages"][1]["content"] == "Please select a tool to handle the current issue."
+ assert (
+ result["messages"][1]["content"]
+ == "Please select a tool to handle the current issue."
+ )
# Check that tool_choice was removed but tools are preserved
assert "tool_choice" not in result
@@ -303,7 +308,10 @@ class TestMoonshotConfig:
# Check that the message was added
assert len(result["messages"]) == 2
- assert result["messages"][1]["content"] == "Please select a tool to handle the current issue."
+ assert (
+ result["messages"][1]["content"]
+ == "Please select a tool to handle the current issue."
+ )
def test_tool_choice_non_required_preserved(self):
"""Test that non-'required' tool_choice values are preserved"""
@@ -528,7 +536,9 @@ class TestMoonshotConfig:
assert result[0].get("reasoning_content") == "stored thinking"
# The promoted key must be removed from provider_specific_fields to
# avoid sending the value twice in the serialised request body
- assert "reasoning_content" not in (result[0].get("provider_specific_fields") or {})
+ assert "reasoning_content" not in (
+ result[0].get("provider_specific_fields") or {}
+ )
def test_reasoning_model_fill_called_from_transform_request(self):
"""transform_request injects reasoning_content end-to-end for reasoning models."""
@@ -628,7 +638,10 @@ class TestMoonshotConfig:
result = config.fill_reasoning_content(messages)
# reasoning_content should be preserved, not replaced with placeholder
- assert result[0].get("reasoning_content") == "User wants weather"
+ assert (
+ result[0].get("reasoning_content")
+ == "User wants weather"
+ )
def test_reasoning_content_preserved_in_multi_turn_flow(self):
"""reasoning_content is preserved through multi-turn conversation flow.
@@ -672,7 +685,10 @@ class TestMoonshotConfig:
result = config.fill_reasoning_content(messages)
# reasoning_content should be preserved in the assistant message
- assert result[1].get("reasoning_content") == "Planning to call weather tool"
+ assert (
+ result[1].get("reasoning_content")
+ == "Planning to call weather tool"
+ )
class TestKimiK26ModelRegistry:
@@ -685,7 +701,9 @@ class TestKimiK26ModelRegistry:
def test_kimi_k26_in_model_cost_map(self, model_cost_map):
"""kimi-k2.6 should be present in the model cost map."""
- assert "moonshot/kimi-k2.6" in model_cost_map, "moonshot/kimi-k2.6 not found in model_cost"
+ assert (
+ "moonshot/kimi-k2.6" in model_cost_map
+ ), "moonshot/kimi-k2.6 not found in model_cost"
def test_kimi_k26_pricing(self, model_cost_map):
"""kimi-k2.6 pricing should match official Kimi API rates."""
@@ -741,6 +759,10 @@ class TestMoonshotResponseSchemaSupport:
def test_live_model_supports_response_schema(self, model, model_cost_map):
assert model_cost_map[model].get("supports_response_schema") is True
- def test_supports_response_schema_utility_reports_true(self, model_cost_map, monkeypatch):
+ def test_supports_response_schema_utility_reports_true(
+ self, model_cost_map, monkeypatch
+ ):
monkeypatch.setattr(litellm, "model_cost", model_cost_map)
- assert litellm.utils.supports_response_schema(model="moonshot/kimi-k2.5") is True
+ assert (
+ litellm.utils.supports_response_schema(model="moonshot/kimi-k2.5") is True
+ )
diff --git a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py
index 53c9e4b207c..e24e4956ab6 100644
--- a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py
+++ b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py
@@ -1478,6 +1478,7 @@ class TestOCIChatConfigErrorPaths:
drop_params=False,
)
assert "max_retries" not in result
+
def test_map_openai_params_cohere_n_default_dropped(self):
"""Cohere has no numGenerations field, but n=1 (and None) is the OpenAI
default single-generation request. It must be dropped silently rather
diff --git a/tests/test_litellm/llms/ollama/test_ollama_model_info.py b/tests/test_litellm/llms/ollama/test_ollama_model_info.py
index 8d46151ecce..3491c843d85 100644
--- a/tests/test_litellm/llms/ollama/test_ollama_model_info.py
+++ b/tests/test_litellm/llms/ollama/test_ollama_model_info.py
@@ -146,9 +146,7 @@ class TestOllamaModelInfo:
assert models == []
assert call_headers[0] == {"Authorization": "Bearer explicit-api-key"}
- def test_get_models_empty_key_does_not_leak_to_provided_api_base(
- self, monkeypatch
- ):
+ def test_get_models_empty_key_does_not_leak_to_provided_api_base(self, monkeypatch):
"""An empty explicit key must not fall back to server-side creds for a custom base."""
call_headers = []
diff --git a/tests/test_litellm/llms/openai_like/test_tensormesh_provider.py b/tests/test_litellm/llms/openai_like/test_tensormesh_provider.py
index c94b2cbfa80..1fbf69c8170 100644
--- a/tests/test_litellm/llms/openai_like/test_tensormesh_provider.py
+++ b/tests/test_litellm/llms/openai_like/test_tensormesh_provider.py
@@ -152,7 +152,9 @@ class TestTensormeshCostMap:
"tensormesh/google/gemma-4-31B-it",
}
for model in TENSORMESH_MODELS:
- assert litellm.supports_reasoning(model) is (model in reasoning_models), model
+ assert litellm.supports_reasoning(model) is (
+ model in reasoning_models
+ ), model
def test_cost_is_wired_and_cache_reads_are_free(self):
prompt_cost, completion_cost = litellm.cost_per_token(
diff --git a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py
index c8751fb2d95..69cb8a111c9 100644
--- a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py
+++ b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py
@@ -56,7 +56,6 @@ def test_ovhcloud_audio_transcription_config_installed():
assert isinstance(config, BaseAudioTranscriptionConfig)
-
class TestOVHCloudDurationFieldMigration:
"""Tests for OVHCloud duration -> seconds field migration."""
@@ -98,8 +97,6 @@ class TestOVHCloudDurationFieldMigration:
assert result.text == "Hello world"
assert result._hidden_params["duration"] == 2.71
-
-
def test_seconds_zero_mapped_to_duration(self):
"""seconds=0.0 must not be treated as falsy and lost."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
@@ -111,4 +108,4 @@ class TestOVHCloudDurationFieldMigration:
mock_response = MagicMock()
mock_response.json.return_value = {"text": "silence", "seconds": 0.0}
result = config.transform_audio_transcription_response(mock_response)
- assert result._hidden_params["duration"] == 0.0
\ No newline at end of file
+ assert result._hidden_params["duration"] == 0.0
diff --git a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py
index 40d57c76d02..831a9a16096 100644
--- a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py
+++ b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py
@@ -301,7 +301,7 @@ class TestOVHCloudReasoningFieldMigration:
"""New `reasoning` field should be mapped to `reasoning_content`."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
- sync_stream=True,
+ sync_stream=True,
)
chunk = {
"id": "test-id",
@@ -324,7 +324,7 @@ class TestOVHCloudReasoningFieldMigration:
"""Legacy `reasoning_content` field should pass through untouched."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
- sync_stream=True,
+ sync_stream=True,
)
chunk = {
"id": "test-id",
@@ -341,13 +341,15 @@ class TestOVHCloudReasoningFieldMigration:
],
}
result = handler.chunk_parser(chunk)
- assert result.choices[0]["delta"]["reasoning_content"] == "Already correct field."
+ assert (
+ result.choices[0]["delta"]["reasoning_content"] == "Already correct field."
+ )
def test_streaming_both_fields_legacy_wins(self):
"""When both fields present, existing `reasoning_content` is not overwritten."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
- sync_stream=True,
+ sync_stream=True,
)
chunk = {
"id": "test-id",
@@ -365,5 +367,3 @@ class TestOVHCloudReasoningFieldMigration:
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "legacy field"
-
-
diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py
index e59fbc9f272..d219be7105f 100644
--- a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py
+++ b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py
@@ -1,7 +1,7 @@
"""
Integration tests for Perplexity cost calculation and transformation.
-Tests the end-to-end functionality of Perplexity cost calculation
+Tests the end-to-end functionality of Perplexity cost calculation
including integration with the main LiteLLM cost calculator.
"""
@@ -105,9 +105,7 @@ class TestPerplexityIntegration:
citation_tokens = citation_chars // 4
expected_prompt_cost = (100 * 2e-6) + (citation_tokens * 2e-6)
- expected_completion_cost = (
- ((50 - 10) * 8e-6) + (10 * 3e-6) + (2 / 1000 * 0.005)
- )
+ expected_completion_cost = ((50 - 10) * 8e-6) + (10 * 3e-6) + (2 / 1000 * 0.005)
expected_total = expected_prompt_cost + expected_completion_cost
assert math.isclose(total_cost, expected_total, rel_tol=1e-6)
diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py
index 7e13459bca1..ebfb0292f98 100644
--- a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py
+++ b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py
@@ -10,7 +10,6 @@ sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.sagemaker.common_utils import AWSEventStreamDecoder
from litellm.llms.sagemaker.completion.transformation import SagemakerConfig
-
# --------------------------------------------------------------------------- #
# get_sagemaker_response_stream_shape lazy-load tests #
# --------------------------------------------------------------------------- #
diff --git a/tests/test_litellm/llms/scaleway/test_scaleway_audio_transcription_transformation.py b/tests/test_litellm/llms/scaleway/test_scaleway_audio_transcription_transformation.py
index 407e1d19fb3..c237e28081e 100644
--- a/tests/test_litellm/llms/scaleway/test_scaleway_audio_transcription_transformation.py
+++ b/tests/test_litellm/llms/scaleway/test_scaleway_audio_transcription_transformation.py
@@ -10,7 +10,6 @@ from litellm.llms.scaleway.audio_transcription.transformation import (
)
from litellm.types.utils import TranscriptionResponse
-
# ---------------------------------------------------------------------------
# get_complete_url
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/llms/snowflake/chat/test_snowflake_chat_transformation.py b/tests/test_litellm/llms/snowflake/chat/test_snowflake_chat_transformation.py
index a182656e4a8..d0affbdead7 100644
--- a/tests/test_litellm/llms/snowflake/chat/test_snowflake_chat_transformation.py
+++ b/tests/test_litellm/llms/snowflake/chat/test_snowflake_chat_transformation.py
@@ -142,7 +142,9 @@ class TestSnowflakeToolTransformation:
"type": "function",
"function": {
"name": "get_weather",
- "arguments": json.dumps({"location": "Paris, France", "unit": "celsius"}),
+ "arguments": json.dumps(
+ {"location": "Paris, France", "unit": "celsius"}
+ ),
},
}
],
@@ -215,7 +217,9 @@ class TestSnowflakeToolTransformation:
"type": "function",
"function": {
"name": "get_weather",
- "arguments": json.dumps({"location": "Tokyo, Japan"}),
+ "arguments": json.dumps(
+ {"location": "Tokyo, Japan"}
+ ),
},
}
],
diff --git a/tests/test_litellm/llms/snowflake/test_snowflake_native_endpoints.py b/tests/test_litellm/llms/snowflake/test_snowflake_native_endpoints.py
index fb21e2e6f6b..91d5a1aba36 100644
--- a/tests/test_litellm/llms/snowflake/test_snowflake_native_endpoints.py
+++ b/tests/test_litellm/llms/snowflake/test_snowflake_native_endpoints.py
@@ -22,7 +22,6 @@ from litellm.llms.snowflake.chat.transformation import (
)
from litellm.types.utils import ModelResponse
-
# ─── Fixtures ──────────────────────────────────────────────────────────────
ACCOUNT_ID = "myaccount"
@@ -69,6 +68,7 @@ def _make_anthropic_response(content: str = "Hello!") -> httpx.Response:
# ─── SnowflakeConfig (OpenAI-compatible) ───────────────────────────────────
+
class TestSnowflakeConfigURL:
def setup_method(self):
self.cfg = SnowflakeConfig()
@@ -82,7 +82,10 @@ class TestSnowflakeConfigURL:
optional_params=optional_params,
litellm_params={},
)
- assert url == f"https://{ACCOUNT_ID}.snowflakecomputing.com/api/v2/cortex/v1/chat/completions"
+ assert (
+ url
+ == f"https://{ACCOUNT_ID}.snowflakecomputing.com/api/v2/cortex/v1/chat/completions"
+ )
def test_url_with_explicit_api_base(self):
url = self.cfg.get_complete_url(
@@ -140,7 +143,10 @@ class TestSnowflakeConfigAuth:
litellm_params={},
api_key=PAT_TOKEN,
)
- assert headers["X-Snowflake-Authorization-Token-Type"] == "PROGRAMMATIC_ACCESS_TOKEN"
+ assert (
+ headers["X-Snowflake-Authorization-Token-Type"]
+ == "PROGRAMMATIC_ACCESS_TOKEN"
+ )
assert headers["Authorization"] == "Bearer my-secret-pat-token"
def test_jwt_auth_sets_keypair_header(self):
@@ -179,7 +185,10 @@ class TestSnowflakeConfigRequest:
"function": {
"name": "get_weather",
"description": "Get weather",
- "parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
+ "parameters": {
+ "type": "object",
+ "properties": {"city": {"type": "string"}},
+ },
},
}
]
@@ -266,6 +275,7 @@ class TestSnowflakeConfigResponse:
# ─── SnowflakeConfig ────────────────────────────────────────
+
class TestAnthropicConfigURL:
def setup_method(self):
self.cfg = SnowflakeConfig()
@@ -290,7 +300,10 @@ class TestAnthropicConfigURL:
optional_params={"account_id": ACCOUNT_ID},
litellm_params={},
)
- assert f"https://{ACCOUNT_ID}.snowflakecomputing.com/api/v2/cortex/v1/messages" == url
+ assert (
+ f"https://{ACCOUNT_ID}.snowflakecomputing.com/api/v2/cortex/v1/messages"
+ == url
+ )
class TestAnthropicConfigAuth:
@@ -317,7 +330,10 @@ class TestAnthropicConfigAuth:
litellm_params={},
api_key=PAT_TOKEN,
)
- assert headers["X-Snowflake-Authorization-Token-Type"] == "PROGRAMMATIC_ACCESS_TOKEN"
+ assert (
+ headers["X-Snowflake-Authorization-Token-Type"]
+ == "PROGRAMMATIC_ACCESS_TOKEN"
+ )
assert headers["anthropic-version"] == "2023-06-01"
assert "Bearer" in headers["Authorization"]
@@ -475,6 +491,7 @@ class TestAnthropicConfigResponse:
# ─── Model detection helper ────────────────────────────────────────────────
+
class TestIsClaudeModel:
def test_claude_model_detected(self):
assert _is_claude_model("snowflake/claude-sonnet-4-5") is True
@@ -490,6 +507,7 @@ class TestIsClaudeModel:
# ─── Anthropic Tool Transformation Tests ──────────────────────────────────
+
class TestAnthropicToolTransformation:
def setup_method(self):
self.cfg = SnowflakeConfig()
@@ -528,7 +546,9 @@ class TestAnthropicToolTransformation:
def test_tools_already_in_anthropic_format_pass_through(self):
messages = [{"role": "user", "content": "hi"}]
- tools = [{"name": "my_tool", "input_schema": {"type": "object", "properties": {}}}]
+ tools = [
+ {"name": "my_tool", "input_schema": {"type": "object", "properties": {}}}
+ ]
body = self.cfg.transform_request(
model="snowflake/claude-sonnet-4-5",
messages=messages,
@@ -619,7 +639,10 @@ class TestAnthropicMultiTurnToolMessages:
headers={},
)
assistant_msg = body["messages"][1]
- assert assistant_msg["content"][0] == {"type": "text", "text": "Let me check that for you."}
+ assert assistant_msg["content"][0] == {
+ "type": "text",
+ "text": "Let me check that for you.",
+ }
assert assistant_msg["content"][1]["type"] == "tool_use"
assert assistant_msg["content"][1]["name"] == "get_weather"
@@ -629,7 +652,13 @@ class TestAnthropicMultiTurnToolMessages:
{
"role": "assistant",
"content": None,
- "tool_calls": [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": "{}"}}],
+ "tool_calls": [
+ {
+ "id": "c1",
+ "type": "function",
+ "function": {"name": "f", "arguments": "{}"},
+ }
+ ],
},
{"role": "tool", "tool_call_id": "c1", "content": "result"},
]
@@ -653,7 +682,10 @@ class TestAnthropicMultiTurnToolMessages:
{
"id": "call_bad",
"type": "function",
- "function": {"name": "broken_tool", "arguments": "not valid json{{{"},
+ "function": {
+ "name": "broken_tool",
+ "arguments": "not valid json{{{",
+ },
}
],
},
@@ -681,7 +713,10 @@ class TestAnthropicMultiTurnToolMessages:
{
"id": "call_dict",
"type": "function",
- "function": {"name": "dict_tool", "arguments": {"already": "parsed"}},
+ "function": {
+ "name": "dict_tool",
+ "arguments": {"already": "parsed"},
+ },
}
],
},
@@ -702,9 +737,19 @@ class TestAnthropicMultiTurnToolMessages:
{
"role": "assistant",
"content": None,
- "tool_calls": [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": "{}"}}],
+ "tool_calls": [
+ {
+ "id": "c1",
+ "type": "function",
+ "function": {"name": "f", "arguments": "{}"},
+ }
+ ],
+ },
+ {
+ "role": "tool",
+ "tool_call_id": "c1",
+ "content": {"result_key": "result_value"},
},
- {"role": "tool", "tool_call_id": "c1", "content": {"result_key": "result_value"}},
]
body = self.cfg.transform_request(
model="snowflake/claude-sonnet-4-5",
diff --git a/tests/test_litellm/llms/test_file_content_block.py b/tests/test_litellm/llms/test_file_content_block.py
index 5552c1a4d68..051ef335b69 100644
--- a/tests/test_litellm/llms/test_file_content_block.py
+++ b/tests/test_litellm/llms/test_file_content_block.py
@@ -104,7 +104,9 @@ def _explicit_null_file_in_content() -> List[AllMessageValues]:
def test_gemini_convert_messages_malformed_file_raises_bad_request():
"""_gemini_convert_messages_with_history should raise BadRequestError (not KeyError)
when a content block has type='file' but no 'file' sub-field."""
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
_gemini_convert_messages_with_history(
messages=_malformed(),
model="gemini-2.0-flash",
@@ -113,7 +115,9 @@ def test_gemini_convert_messages_malformed_file_raises_bad_request():
def test_gemini_convert_messages_explicit_null_file_field_raises_bad_request():
"""Explicit JSON null for `file` must be rejected like a missing `file` key."""
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
_gemini_convert_messages_with_history(
messages=_explicit_null_file_in_content(),
model="gemini-2.0-flash",
@@ -129,20 +133,26 @@ def test_google_ai_studio_transform_messages_malformed_file_raises_bad_request()
"""GoogleAIStudioGeminiConfig._transform_messages should raise BadRequestError
when a content block has type='file' but no 'file' sub-field."""
config = GoogleAIStudioGeminiConfig()
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
config._transform_messages(messages=_malformed(), model="gemini-2.0-flash")
def test_google_ai_studio_transform_messages_explicit_null_file_field_raises_bad_request():
"""Explicit JSON null for `file` must be rejected like a missing `file` key."""
config = GoogleAIStudioGeminiConfig()
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
config._transform_messages(
messages=_explicit_null_file_in_content(), model="gemini-2.0-flash"
)
-def test_google_ai_studio_transform_messages_http_file_id_converts_to_base64(monkeypatch):
+def test_google_ai_studio_transform_messages_http_file_id_converts_to_base64(
+ monkeypatch,
+):
"""Google AI Studio rejects raw HTTP(S) file URLs; _transform_messages should
fetch and replace them with base64 `file_data` before conversion."""
# Data URL shape so downstream Gemini media parsing accepts the inlined bytes
@@ -179,7 +189,9 @@ def test_google_ai_studio_transform_messages_http_file_id_converts_to_base64(mon
config._transform_messages(messages=messages, model="gemini-2.0-flash")
content = messages[0].get("content")
assert isinstance(content, list)
- file_block = next(c for c in content if isinstance(c, dict) and c.get("type") == "file")
+ file_block = next(
+ c for c in content if isinstance(c, dict) and c.get("type") == "file"
+ )
file_field = file_block.get("file")
assert isinstance(file_field, dict)
assert file_field.get("file_data") == fake_file_data
@@ -222,7 +234,9 @@ def test_google_ai_studio_transform_messages_http_file_id_convert_failure_leaves
config._transform_messages(messages=messages, model="gemini-2.0-flash")
content = messages[0].get("content")
assert isinstance(content, list)
- file_block = next(c for c in content if isinstance(c, dict) and c.get("type") == "file")
+ file_block = next(
+ c for c in content if isinstance(c, dict) and c.get("type") == "file"
+ )
file_field = file_block.get("file")
assert isinstance(file_field, dict)
assert file_field.get("file_id") == https_id
@@ -247,7 +261,9 @@ def test_update_messages_with_model_file_ids_malformed_skips_non_openai_file_blo
assert result == messages
content = result[0].get("content")
assert isinstance(content, list)
- file_block = next(c for c in content if isinstance(c, dict) and c.get("type") == "file")
+ file_block = next(
+ c for c in content if isinstance(c, dict) and c.get("type") == "file"
+ )
assert "file" not in file_block
@@ -284,7 +300,10 @@ def test_get_file_ids_from_messages_well_formed_returns_ids():
"role": "user",
"content": [
{"type": "text", "text": "hello"},
- {"type": "file", "file": {"file_id": "file-abc123", "format": "pdf"}},
+ {
+ "type": "file",
+ "file": {"file_id": "file-abc123", "format": "pdf"},
+ },
],
}
],
@@ -301,13 +320,19 @@ def test_get_file_ids_from_messages_well_formed_returns_ids():
def test_bedrock_process_file_message_malformed_raises_bad_request():
"""_process_file_message should raise BadRequestError (not KeyError)
when the file object is missing the 'file' sub-field."""
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
BedrockConverseMessagesProcessor._process_file_message(MALFORMED_FILE_OBJECT)
def test_bedrock_process_file_message_explicit_null_file_field_raises_bad_request():
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
- BedrockConverseMessagesProcessor._process_file_message(EXPLICIT_NULL_FILE_OBJECT)
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
+ BedrockConverseMessagesProcessor._process_file_message(
+ EXPLICIT_NULL_FILE_OBJECT
+ )
def test_bedrock_async_process_file_message_malformed_raises_bad_request():
@@ -349,7 +374,9 @@ def test_openai_apply_common_transform_malformed_file_raises_bad_request():
malformed_block: OpenAIMessageContentListBlock = cast(
OpenAIMessageContentListBlock, {"type": "file"}
)
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
config._apply_common_transform_content_item(malformed_block)
@@ -359,7 +386,9 @@ def test_openai_apply_common_transform_explicit_null_file_field_raises_bad_reque
OpenAIMessageContentListBlock,
{"type": "file", "file": None},
)
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
config._apply_common_transform_content_item(explicit_null_block)
@@ -384,12 +413,16 @@ def test_openai_apply_common_transform_well_formed_file_does_not_raise():
def test_anthropic_process_openai_file_message_malformed_raises_bad_request():
"""anthropic_process_openai_file_message should raise BadRequestError (not KeyError)
when the file object is missing the 'file' sub-field."""
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
anthropic_process_openai_file_message(MALFORMED_FILE_OBJECT)
def test_anthropic_process_openai_file_message_explicit_null_file_field_raises_bad_request():
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
anthropic_process_openai_file_message(EXPLICIT_NULL_FILE_OBJECT)
@@ -411,12 +444,16 @@ def test_anthropic_process_openai_file_message_well_formed_file_id_does_not_rais
def test_migrate_file_to_image_url_malformed_raises_bad_request():
"""migrate_file_to_image_url should raise BadRequestError (not KeyError)
when the file object is missing the 'file' sub-field."""
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
migrate_file_to_image_url(MALFORMED_FILE_OBJECT)
def test_migrate_file_to_image_url_explicit_null_file_field_raises_bad_request():
- with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"):
+ with pytest.raises(
+ litellm.BadRequestError, match="missing the required 'file' field"
+ ):
migrate_file_to_image_url(EXPLICIT_NULL_FILE_OBJECT)
diff --git a/tests/test_litellm/llms/test_polling_url_origin_match.py b/tests/test_litellm/llms/test_polling_url_origin_match.py
index f1f910bc731..b6b932ad937 100644
--- a/tests/test_litellm/llms/test_polling_url_origin_match.py
+++ b/tests/test_litellm/llms/test_polling_url_origin_match.py
@@ -15,7 +15,6 @@ from unittest.mock import MagicMock, patch
import httpx
import pytest
-
# Azure DALL-E sync + async paths route through ``assert_same_origin``
# the same way as the cases below. The helper itself is unit-tested in
# ``tests/test_litellm/litellm_core_utils/test_url_utils.py``; the
diff --git a/tests/test_litellm/llms/test_predibase_transformation.py b/tests/test_litellm/llms/test_predibase_transformation.py
index 1600878a586..7d67996906a 100644
--- a/tests/test_litellm/llms/test_predibase_transformation.py
+++ b/tests/test_litellm/llms/test_predibase_transformation.py
@@ -233,7 +233,9 @@ def test_predibase_transform_response_non_dict_payload():
raw_response.headers = {}
raw_response.json.return_value = []
- with pytest.raises(PredibaseError, match="'completion_response' is not a dictionary"):
+ with pytest.raises(
+ PredibaseError, match="'completion_response' is not a dictionary"
+ ):
config.transform_response(
model="predibase-model",
raw_response=raw_response,
@@ -375,7 +377,9 @@ def test_predibase_transform_response_best_of_invalid_value_falls_back(monkeypat
assert result.choices[0].message.content == "primary-output"
-def test_predibase_transform_response_empty_output_sets_completion_tokens_zero(monkeypatch):
+def test_predibase_transform_response_empty_output_sets_completion_tokens_zero(
+ monkeypatch,
+):
config = PredibaseConfig()
logging_obj = Mock()
encoding = Mock()
@@ -429,7 +433,10 @@ def test_predibase_transform_response_usage_fallbacks(monkeypatch):
raw_response = httpx.Response(
status_code=200,
- json={"generated_text": "ok", "details": {"tokens": [], "finish_reason": "stop"}},
+ json={
+ "generated_text": "ok",
+ "details": {"tokens": [], "finish_reason": "stop"},
+ },
)
result = config.transform_response(
@@ -539,13 +546,17 @@ def test_predibase_completion_sync_returns_transform_response(monkeypatch):
def fake_transform_response(self, **kwargs):
return expected
- monkeypatch.setattr(PredibaseConfig, "validate_environment", fake_validate_environment)
+ monkeypatch.setattr(
+ PredibaseConfig, "validate_environment", fake_validate_environment
+ )
monkeypatch.setattr(PredibaseConfig, "get_complete_url", fake_get_complete_url)
monkeypatch.setattr(PredibaseConfig, "transform_request", fake_transform_request)
monkeypatch.setattr(PredibaseConfig, "transform_response", fake_transform_response)
monkeypatch.setattr(
"litellm.module_level_client.post",
- lambda *args, **kwargs: httpx.Response(status_code=200, json={"generated_text": "ok"}),
+ lambda *args, **kwargs: httpx.Response(
+ status_code=200, json={"generated_text": "ok"}
+ ),
)
result = handler.completion(
@@ -586,7 +597,9 @@ def test_predibase_completion_passes_existing_config_to_async_completion(monkeyp
captured["async_kwargs"] = kwargs
return "async-result"
- monkeypatch.setattr(PredibaseConfig, "validate_environment", fake_validate_environment)
+ monkeypatch.setattr(
+ PredibaseConfig, "validate_environment", fake_validate_environment
+ )
monkeypatch.setattr(PredibaseConfig, "get_complete_url", fake_get_complete_url)
monkeypatch.setattr(PredibaseConfig, "transform_request", fake_transform_request)
monkeypatch.setattr(handler, "async_completion", fake_async_completion)
diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py
index 6d913ad5d1d..ef084dcc004 100644
--- a/tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py
+++ b/tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py
@@ -22,7 +22,6 @@ from litellm.llms.vertex_ai.gemini.transformation import (
)
from litellm.types.llms.vertex_ai import HttpxPartType
-
# --- Response extraction tests ---
diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py b/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py
index 10fc68ecaad..50d7ebf213a 100644
--- a/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py
+++ b/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py
@@ -23,7 +23,12 @@ def test_missing_url_inside_image_url_dict_raises_bad_request_error():
"""When image_url is a dict but 'url' key is absent, a BadRequestError is raised."""
messages = cast(
List[AllMessageValues],
- [{"role": "user", "content": [{"type": "image_url", "image_url": {"detail": "high"}}]}],
+ [
+ {
+ "role": "user",
+ "content": [{"type": "image_url", "image_url": {"detail": "high"}}],
+ }
+ ],
)
with pytest.raises(litellm.BadRequestError) as exc_info:
_gemini_convert_messages_with_history(messages, model="gemini-1.5-pro")
diff --git a/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py
index bb4e6c67e9e..4bd0b93027d 100644
--- a/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py
+++ b/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py
@@ -20,7 +20,6 @@ from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation
from litellm.types.llms.vertex_ai import VertexAIBatchEmbeddingsResponseObject
from litellm.types.utils import EmbeddingResponse
-
IMAGE_DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII"
GCS_URL = "gs://my-bucket/image.png"
@@ -72,7 +71,9 @@ class TestBuildPartForInput:
assert part["file_data"]["file_uri"] == GCS_URL
def test_file_reference_resolved(self):
- resolved = {"files/abc": {"mime_type": "image/jpeg", "uri": "https://example.com/abc"}}
+ resolved = {
+ "files/abc": {"mime_type": "image/jpeg", "uri": "https://example.com/abc"}
+ }
part = _build_part_for_input("files/abc", resolved_files=resolved)
assert part["file_data"] is not None
assert part["file_data"]["mime_type"] == "image/jpeg"
@@ -94,7 +95,9 @@ class TestTransformOpenaiInputGeminiContent:
def test_multiple_texts(self):
result = transform_openai_input_gemini_content(
- input=["hello", "world"], model="gemini-embedding-2-preview", optional_params={}
+ input=["hello", "world"],
+ model="gemini-embedding-2-preview",
+ optional_params={},
)
assert len(result["requests"]) == 2
assert result["requests"][0]["content"]["parts"][0]["text"] == "hello"
@@ -109,7 +112,10 @@ class TestTransformOpenaiInputGeminiContent:
)
assert len(result["requests"]) == 2
# First request is text
- assert result["requests"][0]["content"]["parts"][0]["text"] == "The food was delicious"
+ assert (
+ result["requests"][0]["content"]["parts"][0]["text"]
+ == "The food was delicious"
+ )
# Second request is image
assert result["requests"][1]["content"]["parts"][0]["inline_data"] is not None
diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex.py b/tests/test_litellm/llms/vertex_ai/test_vertex.py
index ec73e5e42be..97ddb4a7306 100644
--- a/tests/test_litellm/llms/vertex_ai/test_vertex.py
+++ b/tests/test_litellm/llms/vertex_ai/test_vertex.py
@@ -1282,7 +1282,6 @@ def test_process_gemini_media():
assert base64_result["inline_data"]["data"] == "/9j/4AAQSkZJRg..."
-
def test_get_image_mime_type_from_url():
"""Test the _get_image_mime_type_from_url function for different image URLs"""
from litellm.llms.vertex_ai.gemini.transformation import (
@@ -1519,8 +1518,8 @@ def test_vertex_parallel_tool_calls_true():
def test_vertex_parallel_tool_calls_false_multiple_tools_dropped():
"""
- parallel_tool_calls=False with multiple tools is dropped for Gemini
- (unsupported upstream). Request should succeed without the param.
+ parallel_tool_calls=False with multiple tools is dropped for Gemini
+ (unsupported upstream). Request should succeed without the param.
"""
tools = [
{"type": "function", "function": {"name": "get_weather"}},
diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py
index 034f85f5a0b..ae28cd403c2 100644
--- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py
+++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py
@@ -232,7 +232,9 @@ def test_datastore_search_request_rejects_target_selecting_fields(field):
def test_search_request_rejects_unsupported_extra_body_field():
- with pytest.raises(BadRequestError, match="Unsupported Vertex AI Search extra_body"):
+ with pytest.raises(
+ BadRequestError, match="Unsupported Vertex AI Search extra_body"
+ ):
_search_request(extra_body={"notARealField": True})
@@ -257,9 +259,7 @@ def test_search_request_forwards_supported_extra_body_fields():
def test_datastore_search_request_forwards_supported_extra_body_fields():
- _, body = _datastore_search_request(
- extra_body={"filter": 'category: ANY("docs")'}
- )
+ _, body = _datastore_search_request(extra_body={"filter": 'category: ANY("docs")'})
assert body["filter"] == 'category: ANY("docs")'
diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py b/tests/test_litellm/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py
index e0eccad80e2..a6e63fce078 100644
--- a/tests/test_litellm/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py
+++ b/tests/test_litellm/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py
@@ -2,6 +2,7 @@
Split from test_vertex.py to satisfy CI per-file size limits.
"""
+
import asyncio
import os
import sys
@@ -222,7 +223,9 @@ def test_get_gcs_object_content_type_http_error_explicit_vs_anonymous():
),
):
assert (
- gt._get_gcs_object_content_type(image_url="gs://public-bucket/public-object")
+ gt._get_gcs_object_content_type(
+ image_url="gs://public-bucket/public-object"
+ )
is None
)
mock_v2.get_access_token.assert_not_called()
@@ -247,7 +250,9 @@ def test_get_gcs_object_content_type_anonymous_success_no_auth_header():
),
):
assert (
- gt._get_gcs_object_content_type(image_url="gs://public-bucket/public-object")
+ gt._get_gcs_object_content_type(
+ image_url="gs://public-bucket/public-object"
+ )
== "image/jpeg"
)
mock_v.get_access_token.assert_not_called()
@@ -340,7 +345,9 @@ def test_async_transform_request_body_offloads_extensionless_gs_not_plain_text()
async def run_plain():
with patch(
"litellm.llms.vertex_ai.gemini.transformation.asyncify",
- side_effect=AssertionError("asyncify must not run without extensionless gs://"),
+ side_effect=AssertionError(
+ "asyncify must not run without extensionless gs://"
+ ),
):
return await gemini_transformation.async_transform_request_body(
gemini_api_key=None,
diff --git a/tests/test_litellm/llms/voyage/test_voyage_multimodal_embedding.py b/tests/test_litellm/llms/voyage/test_voyage_multimodal_embedding.py
index f283e7fe0df..70f94986d48 100644
--- a/tests/test_litellm/llms/voyage/test_voyage_multimodal_embedding.py
+++ b/tests/test_litellm/llms/voyage/test_voyage_multimodal_embedding.py
@@ -90,9 +90,7 @@ class TestVoyageMultimodalEmbeddings:
request = config.transform_embedding_request(
"voyage-multimodal-3.5", "hello", {}, {}
)
- assert request["inputs"] == [
- {"content": [{"type": "text", "text": "hello"}]}
- ]
+ assert request["inputs"] == [{"content": [{"type": "text", "text": "hello"}]}]
def test_multimodal_embedding_response_transformation(self):
from litellm.llms.voyage.embedding.transformation_multimodal import (
@@ -103,9 +101,7 @@ class TestVoyageMultimodalEmbeddings:
config = VoyageMultimodalEmbeddingConfig()
response_payload = {
"object": "list",
- "data": [
- {"object": "embedding", "embedding": [0.1, 0.2], "index": 0}
- ],
+ "data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}],
"model": "voyage-multimodal-3.5",
"usage": {
"text_tokens": 2,
@@ -156,9 +152,7 @@ class TestVoyageMultimodalEmbeddings:
{"dimensions": 512}, {}, "voyage-multimodal-3.5", False
)
assert optional_params == {"output_dimension": 512}
- assert (
- config.map_openai_params({}, {}, "voyage-multimodal-3.5", False) == {}
- )
+ assert config.map_openai_params({}, {}, "voyage-multimodal-3.5", False) == {}
def test_validate_environment_uses_api_key(self):
from litellm.llms.voyage.embedding.transformation_multimodal import (
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_callback_oauth_error_responses.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_callback_oauth_error_responses.py
index 11ef40b9961..445124c4832 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_callback_oauth_error_responses.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_callback_oauth_error_responses.py
@@ -75,9 +75,7 @@ class TestCallbackOAuthErrorResponses:
# Sanity: must not leak the Pydantic validation error.
assert "Field required" not in body
- def test_idp_error_html_escapes_user_controlled_fields(
- self, callback_test_client
- ):
+ def test_idp_error_html_escapes_user_controlled_fields(self, callback_test_client):
"""A malicious IdP must not be able to inject HTML/JS via error params."""
resp = callback_test_client.get(
"/callback",
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_is_tool_name_prefixed.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_is_tool_name_prefixed.py
index 8f09e2410c4..1084e64eab3 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_is_tool_name_prefixed.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_is_tool_name_prefixed.py
@@ -8,7 +8,6 @@ import pytest
from litellm.proxy._experimental.mcp_server.utils import is_tool_name_prefixed
-
# ---------------------------------------------------------------------------
# Legacy behaviour (no known_server_prefixes passed)
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_tool_conversion.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_tool_conversion.py
index b4b219e958c..447f586104f 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_tool_conversion.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_tool_conversion.py
@@ -15,7 +15,6 @@ from litellm.proxy._experimental.mcp_server.sampling_handler import (
_convert_single_content,
)
-
# ---------------------------------------------------------------------------
# Helpers — lightweight MCP type stand-ins
# ---------------------------------------------------------------------------
@@ -25,7 +24,9 @@ def _text(text: str) -> SimpleNamespace:
return SimpleNamespace(type="text", text=text)
-def _tool_use(*, name: str, tool_id: str, input_data: Dict[str, Any]) -> SimpleNamespace:
+def _tool_use(
+ *, name: str, tool_id: str, input_data: Dict[str, Any]
+) -> SimpleNamespace:
return SimpleNamespace(type="tool_use", name=name, id=tool_id, input=input_data)
@@ -53,7 +54,9 @@ class TestConvertSingleContentToolUse:
def test_should_produce_function_call_dict(self):
"""tool_use must produce a proper function-call dict, not a text stub."""
- tu = _tool_use(name="get_weather", tool_id="call_123", input_data={"city": "NYC"})
+ tu = _tool_use(
+ name="get_weather", tool_id="call_123", input_data={"city": "NYC"}
+ )
result = _convert_single_content(tu)
assert result["_marker_type"] == "tool_use"
@@ -129,9 +132,12 @@ class TestConvertMcpMessagesMultiTurnTools:
"""An assistant message with tool_use content should produce
a proper tool_calls array, not a text stub."""
messages = [
- _sampling_msg("assistant", _tool_use(
- name="search", tool_id="call_1", input_data={"query": "LiteLLM"}
- )),
+ _sampling_msg(
+ "assistant",
+ _tool_use(
+ name="search", tool_id="call_1", input_data={"query": "LiteLLM"}
+ ),
+ ),
]
result = _convert_mcp_messages_to_openai(messages)
@@ -148,10 +154,13 @@ class TestConvertMcpMessagesMultiTurnTools:
"""A user message with tool_result content should produce
a separate role='tool' message."""
messages = [
- _sampling_msg("user", _tool_result(
- tool_use_id="call_1",
- content=[_text("Found 42 results")],
- )),
+ _sampling_msg(
+ "user",
+ _tool_result(
+ tool_use_id="call_1",
+ content=[_text("Found 42 results")],
+ ),
+ ),
]
result = _convert_mcp_messages_to_openai(messages)
@@ -167,14 +176,21 @@ class TestConvertMcpMessagesMultiTurnTools:
"""
messages = [
_sampling_msg("user", _text("What's the weather in NYC?")),
- _sampling_msg("assistant", _tool_use(
- name="get_weather", tool_id="call_w1",
- input_data={"city": "NYC"},
- )),
- _sampling_msg("user", _tool_result(
- tool_use_id="call_w1",
- content=[_text("72°F, sunny")],
- )),
+ _sampling_msg(
+ "assistant",
+ _tool_use(
+ name="get_weather",
+ tool_id="call_w1",
+ input_data={"city": "NYC"},
+ ),
+ ),
+ _sampling_msg(
+ "user",
+ _tool_result(
+ tool_use_id="call_w1",
+ content=[_text("72°F, sunny")],
+ ),
+ ),
_sampling_msg("assistant", _text("It's 72°F and sunny in NYC!")),
]
result = _convert_mcp_messages_to_openai(messages)
@@ -200,10 +216,13 @@ class TestConvertMcpMessagesMultiTurnTools:
def test_should_handle_mixed_text_and_tool_use_in_assistant(self):
"""An assistant message with both text and tool_use content."""
messages = [
- _sampling_msg("assistant", [
- _text("Let me check that for you."),
- _tool_use(name="lookup", tool_id="call_lu1", input_data={"id": 42}),
- ]),
+ _sampling_msg(
+ "assistant",
+ [
+ _text("Let me check that for you."),
+ _tool_use(name="lookup", tool_id="call_lu1", input_data={"id": 42}),
+ ],
+ ),
]
result = _convert_mcp_messages_to_openai(messages)
@@ -218,10 +237,13 @@ class TestConvertMcpMessagesMultiTurnTools:
def test_should_handle_multiple_tool_uses_in_single_message(self):
"""Multiple tool_use items in a single assistant message → multiple tool_calls."""
messages = [
- _sampling_msg("assistant", [
- _tool_use(name="tool_a", tool_id="call_a", input_data={}),
- _tool_use(name="tool_b", tool_id="call_b", input_data={"x": 1}),
- ]),
+ _sampling_msg(
+ "assistant",
+ [
+ _tool_use(name="tool_a", tool_id="call_a", input_data={}),
+ _tool_use(name="tool_b", tool_id="call_b", input_data={"x": 1}),
+ ],
+ ),
]
result = _convert_mcp_messages_to_openai(messages)
@@ -234,10 +256,13 @@ class TestConvertMcpMessagesMultiTurnTools:
def test_should_handle_multiple_tool_results_in_single_message(self):
"""Multiple tool_result items in a single user message → multiple tool messages."""
messages = [
- _sampling_msg("user", [
- _tool_result(tool_use_id="call_a", content=[_text("Result A")]),
- _tool_result(tool_use_id="call_b", content=[_text("Result B")]),
- ]),
+ _sampling_msg(
+ "user",
+ [
+ _tool_result(tool_use_id="call_a", content=[_text("Result A")]),
+ _tool_result(tool_use_id="call_b", content=[_text("Result B")]),
+ ],
+ ),
]
result = _convert_mcp_messages_to_openai(messages)
@@ -270,9 +295,10 @@ class TestConvertMcpMessagesMarkerHoisting:
def test_should_hoist_tool_use_arriving_on_user_role(self):
messages = [
- _sampling_msg("user", _tool_use(
- name="search", tool_id="call_1", input_data={"q": "x"}
- )),
+ _sampling_msg(
+ "user",
+ _tool_use(name="search", tool_id="call_1", input_data={"q": "x"}),
+ ),
]
result = _convert_mcp_messages_to_openai(messages)
@@ -282,9 +308,9 @@ class TestConvertMcpMessagesMarkerHoisting:
def test_should_hoist_tool_result_arriving_on_assistant_role(self):
messages = [
- _sampling_msg("assistant", _tool_result(
- tool_use_id="call_1", content=[_text("done")]
- )),
+ _sampling_msg(
+ "assistant", _tool_result(tool_use_id="call_1", content=[_text("done")])
+ ),
]
result = _convert_mcp_messages_to_openai(messages)
@@ -295,10 +321,13 @@ class TestConvertMcpMessagesMarkerHoisting:
def test_should_keep_text_when_hoisting_tool_use_on_user_role(self):
messages = [
- _sampling_msg("user", [
- _text("here you go"),
- _tool_use(name="lookup", tool_id="call_2", input_data={}),
- ]),
+ _sampling_msg(
+ "user",
+ [
+ _text("here you go"),
+ _tool_use(name="lookup", tool_id="call_2", input_data={}),
+ ],
+ ),
]
result = _convert_mcp_messages_to_openai(messages)
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py
index d80d15e2140..b67d7cb4134 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py
@@ -722,9 +722,7 @@ async def test_handle_streamable_http_mcp_delegated_server_surfaces_upstream_cha
delegated_server.needs_user_oauth_token = True
delegated_server.server_id = "delegated-oauth-server"
- upstream_challenge = (
- 'Bearer resource_metadata="https://upstream.example.com/.well-known/oauth-protected-resource"'
- )
+ upstream_challenge = 'Bearer resource_metadata="https://upstream.example.com/.well-known/oauth-protected-resource"'
with (
patch(
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py
index 2558df8533b..2888077a23a 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py
@@ -489,9 +489,7 @@ class TestGetToolsByNames:
{"name": "send_email", "description": "send mail"},
]
- matched = filter_instance._get_tools_by_names(
- ["send_email"], available_tools
- )
+ matched = filter_instance._get_tools_by_names(["send_email"], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == "send_email"
@@ -503,9 +501,7 @@ class TestGetToolsByNames:
client_name = "litellm_" + canonical
available_tools = [{"name": client_name, "description": "scrape"}]
- matched = filter_instance._get_tools_by_names(
- [canonical], available_tools
- )
+ matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
# Must return the incoming tool unchanged so the client-facing
@@ -516,13 +512,9 @@ class TestGetToolsByNames:
"""Some clients use dash as alias separator; accept that too."""
filter_instance = self._make_filter()
canonical = "weather_svc-get_weather"
- available_tools = [
- {"name": "mcp-" + canonical, "description": "weather"}
- ]
+ available_tools = [{"name": "mcp-" + canonical, "description": "weather"}]
- matched = filter_instance._get_tools_by_names(
- [canonical], available_tools
- )
+ matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == "mcp-" + canonical
@@ -552,9 +544,7 @@ class TestGetToolsByNames:
{"name": "litellm_" + canonical, "description": "wrapped"},
]
- matched = filter_instance._get_tools_by_names(
- [canonical], available_tools
- )
+ matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == canonical
@@ -567,9 +557,7 @@ class TestGetToolsByNames:
separator-anchored suffixes of ``litellm_api-fs-read_file``.
"""
filter_instance = self._make_filter()
- available_tools = [
- {"name": "litellm_api-fs-read_file", "description": "read"}
- ]
+ available_tools = [{"name": "litellm_api-fs-read_file", "description": "read"}]
matched = filter_instance._get_tools_by_names(
["fs-read_file", "api-fs-read_file"], available_tools
@@ -590,9 +578,7 @@ class TestGetToolsByNames:
{"name": "my_" + canonical, "description": "plain search"},
]
- matched = filter_instance._get_tools_by_names(
- [canonical], available_tools
- )
+ matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == "my_" + canonical
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_header_isolation.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_header_isolation.py
index 554f98d7209..06da009134f 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_agent_header_isolation.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_header_isolation.py
@@ -16,7 +16,6 @@ import pytest
from litellm.constants import DEFAULT_A2A_AGENT_TIMEOUT
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py
index 5ec5d12784f..52634cc25fe 100644
--- a/tests/test_litellm/proxy/auth/test_auth_checks.py
+++ b/tests/test_litellm/proxy/auth/test_auth_checks.py
@@ -2365,7 +2365,9 @@ async def test_virtual_key_budget_check_reads_from_spend_counter():
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
proxy_logging_obj.budget_alerts = AsyncMock()
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:key:test-hashed-token":
return 1.5
return fallback_spend
@@ -2397,7 +2399,9 @@ async def test_virtual_key_budget_check_fallback_no_counter():
proxy_logging_obj.budget_alerts = AsyncMock()
# get_current_spend returns fallback_spend when no counter exists
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
return fallback_spend
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
@@ -2424,7 +2428,9 @@ async def test_team_budget_check_reads_from_spend_counter():
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
proxy_logging_obj.budget_alerts = AsyncMock()
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:team:test-team":
return 1.5
return fallback_spend
@@ -2449,7 +2455,9 @@ async def test_end_user_budget_check_reads_from_spend_counter():
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:end_user:customer-1":
return 1.5
return fallback_spend
@@ -2475,7 +2483,9 @@ async def test_tag_budget_check_reads_from_spend_counter():
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:tag:paid-tag":
return 1.5
return fallback_spend
@@ -2523,7 +2533,9 @@ async def test_team_member_budget_check_reads_from_spend_counter():
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:team_member:test-user:test-team":
return 1.5
return fallback_spend
@@ -2756,7 +2768,9 @@ async def test_team_member_budget_check_falls_back_to_team_default_budget_id():
return_value=fake_budget_row
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:team_member:test-user:test-team":
return 70.0
return fallback_spend
@@ -2853,7 +2867,9 @@ async def test_team_member_budget_check_per_member_override_wins_over_team_defau
mocked_spend = 70.0
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:team_member:test-user:test-team":
return mocked_spend
return fallback_spend
@@ -2943,7 +2959,9 @@ async def test_team_member_budget_check_null_clone_falls_back_to_team_default():
return_value=fake_default_row
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:team_member:test-user:test-team":
return 500.0
return fallback_spend
@@ -3010,7 +3028,9 @@ async def test_team_member_budget_check_null_clone_with_null_default_skips_enfor
return_value=fake_default_row
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:team_member:test-user:test-team":
return 1000.0
return fallback_spend
@@ -3077,7 +3097,9 @@ async def test_team_member_budget_check_zero_team_default_treated_as_no_cap():
return_value=fake_default_row
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:team_member:test-user:test-team":
return 0.0
return fallback_spend
@@ -3135,7 +3157,9 @@ async def test_team_member_budget_check_zero_per_member_row_still_blocks():
prisma_client = MagicMock()
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:team_member:test-user:test-team":
return 0.0
return fallback_spend
diff --git a/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py b/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py
index e49f025df2e..3203878a1e0 100644
--- a/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py
+++ b/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py
@@ -106,7 +106,9 @@ async def test_custom_auth_enforces_end_user_budget_when_common_checks_skipped()
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:end_user:customer-1":
return 5.0
return fallback_spend
diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py
index 63510086f95..b547ec877e2 100644
--- a/tests/test_litellm/proxy/auth/test_handle_jwt.py
+++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py
@@ -1206,7 +1206,13 @@ async def test_auth_builder_returns_team_membership_object():
JWTAuthManager,
"get_objects",
new_callable=AsyncMock,
- return_value=(user_object, None, None, mock_team_membership, user_object.user_id),
+ return_value=(
+ user_object,
+ None,
+ None,
+ mock_team_membership,
+ user_object.user_id,
+ ),
) as mock_get_objects,
patch.object(
JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock
@@ -3509,9 +3515,7 @@ def test_canonical_user_id_no_change_when_ids_match():
user_object = LiteLLM_UserTable(user_id=same, user_email=same)
assert (
- JWTAuthManager._canonical_user_id_from_db(
- user_id=same, user_object=user_object
- )
+ JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object)
== same
)
@@ -3802,12 +3806,15 @@ async def test_get_objects_team_membership_uses_rebound_user_id():
user_id_jwt_field="email", user_id_upsert=True
)
- with patch(
- "litellm.proxy.auth.handle_jwt.get_user_object",
- side_effect=fake_get_user_object,
- ), patch(
- "litellm.proxy.auth.handle_jwt.get_team_membership",
- side_effect=fake_get_team_membership,
+ with (
+ patch(
+ "litellm.proxy.auth.handle_jwt.get_user_object",
+ side_effect=fake_get_user_object,
+ ),
+ patch(
+ "litellm.proxy.auth.handle_jwt.get_team_membership",
+ side_effect=fake_get_team_membership,
+ ),
):
(
user_object,
diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py
index 8d686900ea6..02b1f698132 100644
--- a/tests/test_litellm/proxy/auth/test_model_checks.py
+++ b/tests/test_litellm/proxy/auth/test_model_checks.py
@@ -436,7 +436,9 @@ def test_wildcard_custom_prefix_keeps_org_segment_for_non_provider_first_segment
result = get_known_models_from_wildcard(
wildcard_model="my_hf/*",
- litellm_params=LiteLLM_Params(model="huggingface/*", custom_llm_provider="huggingface"),
+ litellm_params=LiteLLM_Params(
+ model="huggingface/*", custom_llm_provider="huggingface"
+ ),
)
assert result == ["my_hf/meta-llama/Llama-3-8B"]
diff --git a/tests/test_litellm/proxy/auth/test_onboarding.py b/tests/test_litellm/proxy/auth/test_onboarding.py
index c81f4cb7d66..d55a5472af1 100644
--- a/tests/test_litellm/proxy/auth/test_onboarding.py
+++ b/tests/test_litellm/proxy/auth/test_onboarding.py
@@ -18,7 +18,6 @@ from fastapi import HTTPException
import litellm
from litellm.proxy._types import InvitationClaim
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/common_utils/test_cache_codec.py b/tests/test_litellm/proxy/common_utils/test_cache_codec.py
index 044d4c2d1a7..8aa16ddbce1 100644
--- a/tests/test_litellm/proxy/common_utils/test_cache_codec.py
+++ b/tests/test_litellm/proxy/common_utils/test_cache_codec.py
@@ -54,7 +54,9 @@ class TestCacheCodecSerialize:
def test_with_model_type_already_correct_instance_skips_revalidation(self):
"""Fast-path: value is already model_type — model_validate must NOT be called."""
m = _SampleModel(name="fast", count=7)
- with patch.object(_SampleModel, "model_validate", wraps=_SampleModel.model_validate) as mock_validate:
+ with patch.object(
+ _SampleModel, "model_validate", wraps=_SampleModel.model_validate
+ ) as mock_validate:
out = CacheCodec.serialize(m, model_type=_SampleModel)
assert out == {"name": "fast", "count": 7}
mock_validate.assert_not_called()
@@ -62,7 +64,9 @@ class TestCacheCodecSerialize:
def test_with_model_type_subclass_instance_skips_revalidation(self):
"""Subclass is isinstance of base → should also take the fast path."""
sub = _SampleSubModel(name="sub", count=2)
- with patch.object(_SampleModel, "model_validate", wraps=_SampleModel.model_validate) as mock_validate:
+ with patch.object(
+ _SampleModel, "model_validate", wraps=_SampleModel.model_validate
+ ) as mock_validate:
out = CacheCodec.serialize(sub, model_type=_SampleModel)
assert out == {"name": "sub", "count": 2}
mock_validate.assert_not_called()
diff --git a/tests/test_litellm/proxy/common_utils/test_key_rotation_e2e.py b/tests/test_litellm/proxy/common_utils/test_key_rotation_e2e.py
index d6e1d22fdde..08432baab7c 100644
--- a/tests/test_litellm/proxy/common_utils/test_key_rotation_e2e.py
+++ b/tests/test_litellm/proxy/common_utils/test_key_rotation_e2e.py
@@ -579,7 +579,9 @@ class TestDeprecatedKeyLookupDbE2E:
"""
database_url = os.getenv("DATABASE_URL")
if not database_url:
- pytest.skip("DATABASE_URL not set; skipping DB-backed key-rotation E2E test.")
+ pytest.skip(
+ "DATABASE_URL not set; skipping DB-backed key-rotation E2E test."
+ )
db_url = cast(str, database_url)
proxy_logging_obj = MagicMock()
diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
index 0b683745369..ae50b72eeb2 100644
--- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
+++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
@@ -445,7 +445,9 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
("user", {"user_id": "uid-all-1"}),
("team", {"team_id": "tid-all-1"}),
]:
- writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == table_name]
+ writes = [
+ c for c in mock_prisma_client.db.batch_calls if c["table"] == table_name
+ ]
assert len(writes) == 1, f"expected 1 {table_name} write, got {len(writes)}"
assert writes[0]["where"] == where
assert writes[0]["data"]["spend"] == 0
@@ -1525,7 +1527,9 @@ def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch):
counter_cache.in_memory_cache.set_cache.assert_not_called()
-def test_reset_budget_for_keys_writes_only_spend_and_reset_at(reset_budget_job, mock_prisma_client):
+def test_reset_budget_for_keys_writes_only_spend_and_reset_at(
+ reset_budget_job, mock_prisma_client
+):
"""
Regression for #27730 (the trigger-half).
@@ -1543,8 +1547,8 @@ def test_reset_budget_for_keys_writes_only_spend_and_reset_at(reset_budget_job,
"budget_duration": "30d",
"budget_reset_at": now,
"token": "sk-problematic",
- "object_permission_id": "perm-abc", # would be rejected on update
- "budget_limits": [{"max_budget": 5}], # would be rejected on update
+ "object_permission_id": "perm-abc", # would be rejected on update
+ "budget_limits": [{"max_budget": 5}], # would be rejected on update
"metadata": {"some": "thing"},
},
)
diff --git a/tests/test_litellm/proxy/db/test_tool_registry_writer.py b/tests/test_litellm/proxy/db/test_tool_registry_writer.py
index 7bf1ffda4fe..3bad9916452 100644
--- a/tests/test_litellm/proxy/db/test_tool_registry_writer.py
+++ b/tests/test_litellm/proxy/db/test_tool_registry_writer.py
@@ -325,10 +325,7 @@ async def test_sync_tool_policy_from_db_retries_on_transport_error_first_read():
assert len(invocations) == 2
mock_prisma_client.attempt_db_reconnect.assert_awaited_once()
reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs
- assert (
- reconnect_kwargs["reason"]
- == "sync_tool_policy_from_db_tools_lookup_failure"
- )
+ assert reconnect_kwargs["reason"] == "sync_tool_policy_from_db_tools_lookup_failure"
assert registry.is_initialized()
@@ -361,7 +358,4 @@ async def test_sync_tool_policy_from_db_retries_on_transport_error_second_read()
assert len(perms_invocations) == 2
mock_prisma_client.attempt_db_reconnect.assert_awaited_once()
reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs
- assert (
- reconnect_kwargs["reason"]
- == "sync_tool_policy_from_db_perms_lookup_failure"
- )
+ assert reconnect_kwargs["reason"] == "sync_tool_policy_from_db_perms_lookup_failure"
diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py
index bb079ea6580..03f2d8e347b 100644
--- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py
@@ -314,15 +314,13 @@ class TestContentFilterGuardrail:
# Create a temporary blocked words file
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f:
- f.write(
- """blocked_words:
+ f.write("""blocked_words:
- keyword: "test_keyword"
action: "BLOCK"
description: "Test keyword"
- keyword: "another_word"
action: "MASK"
-"""
- )
+""")
temp_file = f.name
try:
diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py
index 428f2faf041..d2c7d7ba72c 100644
--- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py
@@ -49,7 +49,9 @@ def test_cato_guard_config_no_api_key(monkeypatch):
monkeypatch.delenv("CATO_API_KEY", raising=False)
litellm.set_verbose = True
litellm.guardrail_name_config_map = {}
- with pytest.raises(CatoNetworksGuardrailMissingSecrets, match="Couldn't get Cato Networks api key"):
+ with pytest.raises(
+ CatoNetworksGuardrailMissingSecrets, match="Couldn't get Cato Networks api key"
+ ):
init_guardrails_v2(
all_guardrails=[
{
@@ -82,7 +84,9 @@ async def test_block_callback(mode: str):
config_file_path="",
)
cato_guardrails = [
- callback for callback in litellm.callbacks if isinstance(callback, CatoNetworksGuardrail)
+ callback
+ for callback in litellm.callbacks
+ if isinstance(callback, CatoNetworksGuardrail)
]
assert len(cato_guardrails) == 1
cato_guardrail = cato_guardrails[0]
@@ -145,7 +149,9 @@ async def test_anonymize_callback__it_returns_redacted_content(mode: str):
config_file_path="",
)
cato_guardrails = [
- callback for callback in litellm.callbacks if isinstance(callback, CatoNetworksGuardrail)
+ callback
+ for callback in litellm.callbacks
+ if isinstance(callback, CatoNetworksGuardrail)
]
assert len(cato_guardrails) == 1
cato_guardrail = cato_guardrails[0]
@@ -192,7 +198,9 @@ async def test_post_call__with_anonymized_entities__it_doesnt_deanonymize_output
config_file_path="",
)
cato_guardrails = [
- callback for callback in litellm.callbacks if isinstance(callback, CatoNetworksGuardrail)
+ callback
+ for callback in litellm.callbacks
+ if isinstance(callback, CatoNetworksGuardrail)
]
assert len(cato_guardrails) == 1
cato_guardrail = cato_guardrails[0]
@@ -375,7 +383,9 @@ def test_init_uses_cato_api_base_env_var(monkeypatch):
def test_init_explicit_args_take_precedence_over_env(monkeypatch):
monkeypatch.setenv("CATO_API_KEY", "env-key")
monkeypatch.setenv("CATO_API_BASE", "https://env.example.com")
- guard = CatoNetworksGuardrail(api_key="explicit-key", api_base="https://explicit.example.com")
+ guard = CatoNetworksGuardrail(
+ api_key="explicit-key", api_base="https://explicit.example.com"
+ )
assert guard.api_key == "explicit-key"
assert guard.api_base == "https://explicit.example.com"
assert guard.ws_api_base == "wss://explicit.example.com"
@@ -386,10 +396,13 @@ def test_init_http_api_base_maps_to_ws():
assert guard.ws_api_base == "ws://insecure.example.com"
-@pytest.mark.parametrize("api_base", [
- "https://api.aisec.catonetworks.com/",
- "https://api.aisec.catonetworks.com",
-])
+@pytest.mark.parametrize(
+ "api_base",
+ [
+ "https://api.aisec.catonetworks.com/",
+ "https://api.aisec.catonetworks.com",
+ ],
+)
def test_base_url_trailing_slash(monkeypatch, api_base):
monkeypatch.setenv("CATO_API_KEY", "test-key")
guardrail = CatoNetworksGuardrail(api_base=api_base)
@@ -684,9 +697,7 @@ async def test_call_cato_guardrail_inspects_responses_api_input():
await guard.call_cato_guardrail(data, hook="pre_call", key_alias=None)
assert exc.value.status_code == 400
- assert any(
- "hunter2" in (m.get("content") or "") for m in captured["messages"]
- )
+ assert any("hunter2" in (m.get("content") or "") for m in captured["messages"])
@pytest.mark.asyncio
@@ -1642,7 +1653,10 @@ async def test_post_call_success_hook_anonymize_action_redacts_content():
anonymize_response = _make_response(
{
"analysis_result": {"policy_drill_down": {"PII": {}}},
- "required_action": {"action_type": "anonymize_action", "policy_name": "PII"},
+ "required_action": {
+ "action_type": "anonymize_action",
+ "policy_name": "PII",
+ },
"redacted_chat": {
"all_redacted_messages": [
{"role": "user", "content": "hi"},
@@ -1680,7 +1694,10 @@ async def test_post_call_success_hook_anonymize_action_applies_empty_redacted_ou
anonymize_response = _make_response(
{
"analysis_result": {"policy_drill_down": {"PII": {}}},
- "required_action": {"action_type": "anonymize_action", "policy_name": "PII"},
+ "required_action": {
+ "action_type": "anonymize_action",
+ "policy_name": "PII",
+ },
"redacted_chat": {
"all_redacted_messages": [
{"role": "user", "content": "hi"},
@@ -1718,7 +1735,10 @@ async def test_post_call_success_hook_anonymize_action_empty_redacted_messages_k
anonymize_response = _make_response(
{
"analysis_result": {"policy_drill_down": {"PII": {}}},
- "required_action": {"action_type": "anonymize_action", "policy_name": "PII"},
+ "required_action": {
+ "action_type": "anonymize_action",
+ "policy_name": "PII",
+ },
"redacted_chat": {"all_redacted_messages": []},
}
)
@@ -1751,7 +1771,10 @@ async def test_post_call_success_hook_anonymize_action_missing_content_key_keeps
anonymize_response = _make_response(
{
"analysis_result": {"policy_drill_down": {"PII": {}}},
- "required_action": {"action_type": "anonymize_action", "policy_name": "PII"},
+ "required_action": {
+ "action_type": "anonymize_action",
+ "policy_name": "PII",
+ },
"redacted_chat": {
"all_redacted_messages": [
{"role": "user", "content": "hi"},
@@ -1794,7 +1817,10 @@ async def test_post_call_success_hook_anonymize_action_partial_redacted_keeps_ou
anonymize_response = _make_response(
{
"analysis_result": {"policy_drill_down": {"PII": {}}},
- "required_action": {"action_type": "anonymize_action", "policy_name": "PII"},
+ "required_action": {
+ "action_type": "anonymize_action",
+ "policy_name": "PII",
+ },
"redacted_chat": {
"all_redacted_messages": [
{"role": "user", "content": "[REDACTED_INPUT_1]"},
@@ -1984,7 +2010,10 @@ async def test_post_call_success_hook_redacts_tool_call_arguments_keeps_none_con
anonymize_response = _make_response(
{
"analysis_result": {"policy_drill_down": {"PII": {}}},
- "required_action": {"action_type": "anonymize_action", "policy_name": "PII"},
+ "required_action": {
+ "action_type": "anonymize_action",
+ "policy_name": "PII",
+ },
"redacted_chat": {
"all_redacted_messages": [
{"role": "user", "content": "email my doctor"},
@@ -2165,7 +2194,10 @@ async def test_post_call_success_hook_redacts_responses_api_output_text():
anonymize_response = _make_response(
{
"analysis_result": {"policy_drill_down": {"PII": {}}},
- "required_action": {"action_type": "anonymize_action", "policy_name": "PII"},
+ "required_action": {
+ "action_type": "anonymize_action",
+ "policy_name": "PII",
+ },
"redacted_chat": {
"all_redacted_messages": [
{"role": "user", "content": "hi"},
@@ -2210,7 +2242,10 @@ async def test_post_call_success_hook_redacts_responses_api_function_call_argume
anonymize_response = _make_response(
{
"analysis_result": {"policy_drill_down": {"PII": {}}},
- "required_action": {"action_type": "anonymize_action", "policy_name": "PII"},
+ "required_action": {
+ "action_type": "anonymize_action",
+ "policy_name": "PII",
+ },
"redacted_chat": {
"all_redacted_messages": [
{"role": "user", "content": "email my doctor"},
diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cisco_ai_defense_mcp.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cisco_ai_defense_mcp.py
index 137b7d24023..1b0d77942de 100644
--- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cisco_ai_defense_mcp.py
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cisco_ai_defense_mcp.py
@@ -126,7 +126,11 @@ class TestCiscoAIDefenseMCPMode:
@pytest.mark.parametrize(
"verdict_extra",
[
- {"sanitized_payload": {"params": {"arguments": {"note": "ssn [REDACTED]"}}}},
+ {
+ "sanitized_payload": {
+ "params": {"arguments": {"note": "ssn [REDACTED]"}}
+ }
+ },
{"sanitized_text": "ssn [REDACTED]"},
],
ids=["structured_arguments", "sanitized_text_fallback"],
diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py
index f8fd9a0a185..679e04ccec6 100644
--- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py
@@ -533,12 +533,41 @@ async def test_apply_guardrail_no_metadata_skips_user_fields(
@pytest.mark.parametrize(
"litellm_metadata, metadata",
[
- (None, {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}),
- ({"trace_id": "t1"}, {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}),
- (["unexpected"], {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}),
- ({"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}, {"trace_id": "t1"}),
+ (
+ None,
+ {
+ "user_api_key_user_id": "uid-abc",
+ "user_api_key_user_email": "alice@example.com",
+ },
+ ),
+ (
+ {"trace_id": "t1"},
+ {
+ "user_api_key_user_id": "uid-abc",
+ "user_api_key_user_email": "alice@example.com",
+ },
+ ),
+ (
+ ["unexpected"],
+ {
+ "user_api_key_user_id": "uid-abc",
+ "user_api_key_user_email": "alice@example.com",
+ },
+ ),
+ (
+ {
+ "user_api_key_user_id": "uid-abc",
+ "user_api_key_user_email": "alice@example.com",
+ },
+ {"trace_id": "t1"},
+ ),
+ ],
+ ids=[
+ "identity_in_metadata_llm_none",
+ "identity_in_metadata_llm_user_dict",
+ "identity_in_metadata_llm_non_mapping",
+ "identity_in_litellm_metadata",
],
- ids=["identity_in_metadata_llm_none", "identity_in_metadata_llm_user_dict", "identity_in_metadata_llm_non_mapping", "identity_in_litellm_metadata"],
)
async def test_apply_guardrail_reads_identity_from_either_metadata_bag(
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
diff --git a/tests/test_litellm/proxy/guardrails/test_content_utils.py b/tests/test_litellm/proxy/guardrails/test_content_utils.py
index 099fca78a62..d3428dfbafa 100644
--- a/tests/test_litellm/proxy/guardrails/test_content_utils.py
+++ b/tests/test_litellm/proxy/guardrails/test_content_utils.py
@@ -8,7 +8,6 @@ from litellm.proxy.guardrails._content_utils import (
walk_user_text,
)
-
# ── iter_message_text ────────────────────────────────────────────────────────────
diff --git a/tests/test_litellm/proxy/guardrails/test_llm_as_a_judge.py b/tests/test_litellm/proxy/guardrails/test_llm_as_a_judge.py
index c9fde4ffbae..9178a66611e 100644
--- a/tests/test_litellm/proxy/guardrails/test_llm_as_a_judge.py
+++ b/tests/test_litellm/proxy/guardrails/test_llm_as_a_judge.py
@@ -13,7 +13,6 @@ from litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge import (
initialize_guardrail,
)
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@@ -39,8 +38,20 @@ def _make_guardrail(**overrides) -> LLMAsAJudgeGuardrail:
def _make_verdict_response(overall_score: float) -> dict:
return {
"verdicts": [
- {"criterion_name": "Accuracy", "score": overall_score, "reasoning": "ok", "passed": True, "weight": 60},
- {"criterion_name": "Safety", "score": overall_score, "reasoning": "ok", "passed": True, "weight": 40},
+ {
+ "criterion_name": "Accuracy",
+ "score": overall_score,
+ "reasoning": "ok",
+ "passed": True,
+ "weight": 60,
+ },
+ {
+ "criterion_name": "Safety",
+ "score": overall_score,
+ "reasoning": "ok",
+ "passed": True,
+ "weight": 40,
+ },
],
"overall_score": overall_score,
}
@@ -90,7 +101,15 @@ def test_build_judge_prompt_missing_name_and_weight():
def _make_litellm_params(**overrides):
params = MagicMock()
- for attr in ("guardrail_name", "judge_model", "criteria", "on_failure", "overall_threshold", "mode", "default_on"):
+ for attr in (
+ "guardrail_name",
+ "judge_model",
+ "criteria",
+ "on_failure",
+ "overall_threshold",
+ "mode",
+ "default_on",
+ ):
setattr(params, attr, None)
for k, v in overrides.items():
setattr(params, k, v)
@@ -98,12 +117,19 @@ def _make_litellm_params(**overrides):
def _make_guardrail_dict(name="g", **litellm_params_overrides):
- raw = {"judge_model": "gpt-4o-mini", "criteria": CRITERIA_100, "on_failure": "block", "overall_threshold": 80.0}
+ raw = {
+ "judge_model": "gpt-4o-mini",
+ "criteria": CRITERIA_100,
+ "on_failure": "block",
+ "overall_threshold": 80.0,
+ }
raw.update(litellm_params_overrides)
return {"guardrail_name": name, "litellm_params": raw}
-@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.logging_callback_manager")
+@patch(
+ "litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.logging_callback_manager"
+)
def test_initialize_guardrail_ok(mock_mgr):
lp = _make_litellm_params()
g = _make_guardrail_dict()
@@ -160,11 +186,18 @@ async def test_apply_guardrail_empty_response_passthrough():
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion")
async def test_apply_guardrail_passes_above_threshold(mock_completion):
mock_completion.return_value = MagicMock(
- choices=[MagicMock(message=MagicMock(content=json.dumps(_make_verdict_response(90.0))))]
+ choices=[
+ MagicMock(
+ message=MagicMock(content=json.dumps(_make_verdict_response(90.0)))
+ )
+ ]
)
guardrail = _make_guardrail(overall_threshold=80.0)
inputs = {"texts": ["good response"]}
- request_data: dict = {"messages": [{"role": "user", "content": "hi"}], "metadata": {}}
+ request_data: dict = {
+ "messages": [{"role": "user", "content": "hi"}],
+ "metadata": {},
+ }
result = await guardrail.apply_guardrail(inputs, request_data, "response")
assert result is inputs
assert request_data["metadata"]["eval_information"]["passed"] is True
@@ -174,7 +207,11 @@ async def test_apply_guardrail_passes_above_threshold(mock_completion):
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion")
async def test_apply_guardrail_blocks_below_threshold(mock_completion):
mock_completion.return_value = MagicMock(
- choices=[MagicMock(message=MagicMock(content=json.dumps(_make_verdict_response(50.0))))]
+ choices=[
+ MagicMock(
+ message=MagicMock(content=json.dumps(_make_verdict_response(50.0)))
+ )
+ ]
)
guardrail = _make_guardrail(overall_threshold=80.0, on_failure="block")
inputs = {"texts": ["bad response"]}
@@ -188,7 +225,11 @@ async def test_apply_guardrail_blocks_below_threshold(mock_completion):
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion")
async def test_apply_guardrail_log_mode_does_not_block(mock_completion):
mock_completion.return_value = MagicMock(
- choices=[MagicMock(message=MagicMock(content=json.dumps(_make_verdict_response(50.0))))]
+ choices=[
+ MagicMock(
+ message=MagicMock(content=json.dumps(_make_verdict_response(50.0)))
+ )
+ ]
)
guardrail = _make_guardrail(overall_threshold=80.0, on_failure="log")
inputs = {"texts": ["bad response"]}
diff --git a/tests/test_litellm/proxy/guardrails/test_qostodian_nexus_guardrail.py b/tests/test_litellm/proxy/guardrails/test_qostodian_nexus_guardrail.py
index 6daa3e1430d..f7b4eb762a6 100644
--- a/tests/test_litellm/proxy/guardrails/test_qostodian_nexus_guardrail.py
+++ b/tests/test_litellm/proxy/guardrails/test_qostodian_nexus_guardrail.py
@@ -232,9 +232,9 @@ def test_qostodian_nexus_builtin_extra_headers():
]
for header in expected_headers:
- assert header in instance.extra_headers, (
- f"Expected built-in header '{header}' to be in extra_headers"
- )
+ assert (
+ header in instance.extra_headers
+ ), f"Expected built-in header '{header}' to be in extra_headers"
def test_qostodian_nexus_extra_headers_merged():
diff --git a/tests/test_litellm/proxy/hooks/test_post_call_response_headers_hook.py b/tests/test_litellm/proxy/hooks/test_post_call_response_headers_hook.py
index 660b0b0162a..11d0196f5fb 100644
--- a/tests/test_litellm/proxy/hooks/test_post_call_response_headers_hook.py
+++ b/tests/test_litellm/proxy/hooks/test_post_call_response_headers_hook.py
@@ -439,6 +439,5 @@ async def test_litellm_call_info_hidden_params_takes_priority():
)
assert (
- inspector.received_call_info["custom_llm_provider"]
- == "hidden_params_value"
+ inspector.received_call_info["custom_llm_provider"] == "hidden_params_value"
)
diff --git a/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py
index 02b4e32db86..f0e640e36e9 100644
--- a/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py
+++ b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py
@@ -69,7 +69,6 @@ from litellm.proxy.hooks.rate_limiter_utils import (
from litellm.proxy.utils import InternalUsageCache
from litellm.types.agents import AgentResponse
-
# ---------------------------------------------------------------------------
# Helper class itself
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py
index f2ccfcd0155..edd592ac2c5 100644
--- a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py
+++ b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py
@@ -644,10 +644,7 @@ async def test_get_all_search_tools_from_db_retries_on_transport_error():
assert len(invocations) == 2
mock_prisma_client.attempt_db_reconnect.assert_awaited_once()
reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs
- assert (
- reconnect_kwargs["reason"]
- == "get_all_search_tools_from_db_lookup_failure"
- )
+ assert reconnect_kwargs["reason"] == "get_all_search_tools_from_db_lookup_failure"
@contextlib.contextmanager
diff --git a/tests/test_litellm/proxy/management_endpoints/test_activity_tenant_scoping.py b/tests/test_litellm/proxy/management_endpoints/test_activity_tenant_scoping.py
index 0855c201945..4686c5e0475 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_activity_tenant_scoping.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_activity_tenant_scoping.py
@@ -14,7 +14,6 @@ import pytest
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
-
# ---------------------------------------------------------------------------
# /team/daily/activity — per-team admin/permission requirement
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py
index 4bdef2e8f96..24bd529f70c 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py
@@ -289,10 +289,7 @@ class TestCacheSettingsManager:
assert len(invocations) == 2
mock_prisma_client.attempt_db_reconnect.assert_awaited_once()
reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs
- assert (
- reconnect_kwargs["reason"]
- == "init_cache_settings_in_db_lookup_failure"
- )
+ assert reconnect_kwargs["reason"] == "init_cache_settings_in_db_lookup_failure"
# ── Audit-log emission for /cache/settings ────────────────────────────────────
diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py
index 2d2c18bb46e..fa983601879 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py
@@ -758,6 +758,8 @@ class TestBuildAggregatedSqlQuery:
]
assert "model = $4" in sql
assert "api_key = $5" in sql
+
+
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_empty_result_set():
"""Regression test for the empty-range 500.
diff --git a/tests/test_litellm/proxy/management_endpoints/test_delete_verification_tokens_failed.py b/tests/test_litellm/proxy/management_endpoints/test_delete_verification_tokens_failed.py
index 63e584e49bc..6720002c2bf 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_delete_verification_tokens_failed.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_delete_verification_tokens_failed.py
@@ -26,7 +26,6 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
delete_verification_tokens,
)
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py b/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py
index a06d79306ab..b9072b2952d 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py
@@ -14,7 +14,6 @@ from fastapi import HTTPException
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
-
# ---------------------------------------------------------------------------
# /project/update — _check_user_permission_for_project
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py b/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py
index 443089b5f01..b0aae5ce393 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py
@@ -25,7 +25,6 @@ from litellm.proxy.management_endpoints.team_endpoints import (
)
from litellm.proxy.proxy_server import ProxyConfig
-
# ---------------------------------------------------------------------------
# _update_config_fields: default_team_params loaded from DB on startup
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
index a649bc7225e..cd6fccda244 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
@@ -4826,9 +4826,7 @@ async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed():
"team_id": "standalone-uncapped-123",
"organization_id": None,
"max_budget": None,
- "members_with_roles": [
- {"user_id": "uncapped-team-admin", "role": "admin"}
- ],
+ "members_with_roles": [{"user_id": "uncapped-team-admin", "role": "admin"}],
}
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(
return_value=mock_existing_team
diff --git a/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py
index a337ff6d888..bf55fd560e0 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py
@@ -17,7 +17,6 @@ sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.management_endpoints.workflow_management_endpoints import router
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py
index e8a74e41dae..da899cd62f4 100644
--- a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py
+++ b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py
@@ -17,7 +17,6 @@ from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import (
stream_usage_ai_chat,
)
-
SAMPLE_AGGREGATED_RESPONSE = {
"results": [
{
diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
index 8bb7b52af14..70adf4bc3b0 100644
--- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
+++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
@@ -1719,7 +1719,9 @@ class TestBedrockLLMProxyRoute:
}
with pytest.raises(HTTPException) as exc_info:
- await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
+ await handler.process_input_messages(
+ data=data, guardrail_to_apply=guardrail
+ )
assert exc_info.value.status_code == 400
assert "Blocked by guardrail" in str(exc_info.value.detail)
diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_invitation.py b/tests/test_litellm/proxy/proxy_server/test_routes_invitation.py
index 5b54a63d8a2..e41db4e26c7 100644
--- a/tests/test_litellm/proxy/proxy_server/test_routes_invitation.py
+++ b/tests/test_litellm/proxy/proxy_server/test_routes_invitation.py
@@ -17,7 +17,6 @@ import pytest
from .conftest import VOLATILE_KEYS, normalize
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@@ -309,9 +308,7 @@ def test_invitation_delete_admin_happy(client, auth_as, monkeypatch, mock_prisma
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
- response = client.post(
- "/invitation/delete", json={"invitation_id": "inv-del"}
- )
+ response = client.post("/invitation/delete", json={"invitation_id": "inv-del"})
assert response.status_code == 200
assert normalize(response.json()) == {
@@ -342,9 +339,7 @@ def test_invitation_delete_non_admin_forbidden(
monkeypatch.setattr(ps, "_user_has_admin_privileges", _no_privileges)
with auth_as(LitellmUserRoles.INTERNAL_USER):
- response = client.post(
- "/invitation/delete", json={"invitation_id": "inv-del"}
- )
+ response = client.post("/invitation/delete", json={"invitation_id": "inv-del"})
assert response.status_code == 400
err_text = str(response.json())
@@ -360,9 +355,7 @@ def test_invitation_delete_unknown_id_400(client, auth_as, monkeypatch, mock_pri
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
- response = client.post(
- "/invitation/delete", json={"invitation_id": "ghost"}
- )
+ response = client.post("/invitation/delete", json={"invitation_id": "ghost"})
assert response.status_code == 400
assert response.json() == {
@@ -378,9 +371,7 @@ def test_invitation_delete_db_not_connected_400(client, auth_as, monkeypatch):
monkeypatch.setattr(ps, "prisma_client", None)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
- response = client.post(
- "/invitation/delete", json={"invitation_id": "inv-del"}
- )
+ response = client.post("/invitation/delete", json={"invitation_id": "inv-del"})
assert response.status_code == 400
err_text = str(response.json())
diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_misc.py b/tests/test_litellm/proxy/proxy_server/test_routes_misc.py
index 0c45e31afd2..4003f002f21 100644
--- a/tests/test_litellm/proxy/proxy_server/test_routes_misc.py
+++ b/tests/test_litellm/proxy/proxy_server/test_routes_misc.py
@@ -17,7 +17,6 @@ import pytest
from .conftest import normalize
-
# ---------------------------------------------------------------------------
# GET /
# ---------------------------------------------------------------------------
@@ -75,9 +74,11 @@ def test_get_routes_invalid_method_405(client):
"""POST against the GET-only /routes endpoint is rejected (error path)."""
response = client.post("/routes")
assert response.status_code == 405
- body = response.json() if response.headers.get("content-type", "").startswith(
- "application/json"
- ) else {}
+ body = (
+ response.json()
+ if response.headers.get("content-type", "").startswith("application/json")
+ else {}
+ )
assert isinstance(body, dict)
@@ -139,7 +140,9 @@ def test_get_logo_url_returns_http_url_when_set(client, monkeypatch):
monkeypatch.setenv("UI_LOGO_PATH", "https://example.invalid/logo.png")
response = client.get("/get_logo_url")
assert response.status_code == 200
- assert normalize(response.json()) == {"logo_url": "https://example.invalid/logo.png"}
+ assert normalize(response.json()) == {
+ "logo_url": "https://example.invalid/logo.png"
+ }
def test_get_logo_url_blank_when_local_path(client, monkeypatch):
diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py b/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py
index 35ae9a3568e..2fdfd96bf0d 100644
--- a/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py
+++ b/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py
@@ -16,7 +16,6 @@ import pytest
from .conftest import normalize
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@@ -126,9 +125,7 @@ def test_onboarding_get_token_master_key_missing_500(client, monkeypatch, mock_p
assert "Master Key not set" in str(err_blob)
-def test_onboarding_get_token_invalid_invite_link_401(
- client, monkeypatch, mock_prisma
-):
+def test_onboarding_get_token_invalid_invite_link_401(client, monkeypatch, mock_prisma):
"""Unknown invite link → 401 with the not-in-db error message."""
from litellm.proxy import proxy_server as ps
@@ -161,7 +158,9 @@ def test_onboarding_get_token_expired_invite_401(client, monkeypatch, mock_prism
response = client.get("/onboarding/get_token", params={"invite_link": "inv-123"})
assert response.status_code == 401
- assert response.json().get("detail", {}).get("error") == "Invitation link has expired."
+ assert (
+ response.json().get("detail", {}).get("error") == "Invitation link has expired."
+ )
def test_onboarding_get_token_missing_query_param_422(client, monkeypatch, mock_prisma):
@@ -270,9 +269,7 @@ def test_claim_onboarding_link_invalid_invite_401(client, monkeypatch, mock_pris
}
-def test_claim_onboarding_link_user_id_mismatch_401(
- client, monkeypatch, mock_prisma
-):
+def test_claim_onboarding_link_user_id_mismatch_401(client, monkeypatch, mock_prisma):
"""Invitation belongs to a different user_id → 401 with mismatch error."""
from litellm.proxy import proxy_server as ps
@@ -317,9 +314,7 @@ def test_claim_onboarding_link_missing_field_422(client, monkeypatch, mock_prism
assert any("password" in str(item) for item in body["detail"])
-def test_claim_onboarding_link_bad_onboarding_jwt_401(
- client, monkeypatch, mock_prisma
-):
+def test_claim_onboarding_link_bad_onboarding_jwt_401(client, monkeypatch, mock_prisma):
"""Onboarding JWT decodes but token_type / invitation_link don't match → 401."""
from litellm.proxy import proxy_server as ps
diff --git a/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py b/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py
index 656e1406f07..ac43418e41f 100644
--- a/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py
+++ b/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py
@@ -181,6 +181,8 @@ def test_rag_ingest_blocks_clientside_credentials(client_internal_user, blocked_
assert blocked_field in str(
body
), f"Response should mention '{blocked_field}': {body}"
+
+
class TestRagIngestSSRFBlocked:
"""
aws_sts_endpoint and related credential-redirect fields must be rejected
@@ -222,7 +224,9 @@ class TestRagIngestSSRFBlocked:
error_text = (
detail.get("error", "") if isinstance(detail, dict) else str(detail)
)
- assert field in error_text, f"Error should name the offending field: {error_text}"
+ assert (
+ field in error_text
+ ), f"Error should name the offending field: {error_text}"
def test_clean_bedrock_ingest_options_not_rejected(self, client_internal_user):
with patch(
@@ -239,6 +243,6 @@ class TestRagIngestSSRFBlocked:
},
},
)
- assert response.status_code != 400, (
- f"Clean Bedrock ingest_options should not be rejected: {response.json()}"
- )
+ assert (
+ response.status_code != 400
+ ), f"Clean Bedrock ingest_options should not be rejected: {response.json()}"
diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py
index e0e51e7b966..4aedbd624cd 100644
--- a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py
+++ b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py
@@ -267,9 +267,7 @@ async def test_client_secrets_transcription_rejects_disallowed_nested_model(
"model": "gpt-4o-realtime-preview",
"audio": {
"input": {
- "transcription": {
- "model": "gpt-realtime-whisper"
- }
+ "transcription": {"model": "gpt-realtime-whisper"}
}
},
},
@@ -343,9 +341,7 @@ async def test_client_secrets_transcription_routes_on_nested_model(
"model": "gpt-4o-realtime-preview",
"audio": {
"input": {
- "transcription": {
- "model": "gpt-realtime-whisper"
- }
+ "transcription": {"model": "gpt-realtime-whisper"}
}
},
},
@@ -595,9 +591,7 @@ async def test_transcription_sessions_rejects_disallowed_resolved_model(
response = client.post(
"/v1/realtime/transcription_sessions",
headers={"Authorization": "Bearer sk-test-master-key"},
- json={
- "input_audio_transcription": {"model": "gpt-realtime-whisper"}
- },
+ json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}},
)
assert response.status_code == 403
@@ -641,9 +635,7 @@ async def test_transcription_sessions_rejects_disallowed_team_model_scope(
response = client.post(
"/v1/realtime/transcription_sessions",
headers={"Authorization": "Bearer sk-test-master-key"},
- json={
- "input_audio_transcription": {"model": "gpt-realtime-whisper"}
- },
+ json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}},
)
assert response.status_code == 403
@@ -686,9 +678,7 @@ async def test_transcription_sessions_rejects_disallowed_project_model_scope(
response = client.post(
"/v1/realtime/transcription_sessions",
headers={"Authorization": "Bearer sk-test-master-key"},
- json={
- "input_audio_transcription": {"model": "gpt-realtime-whisper"}
- },
+ json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}},
)
assert response.status_code == 403
@@ -741,9 +731,7 @@ async def test_transcription_sessions_rejects_disallowed_team_member_model_scope
response = client.post(
"/v1/realtime/transcription_sessions",
headers={"Authorization": "Bearer sk-test-master-key"},
- json={
- "input_audio_transcription": {"model": "gpt-realtime-whisper"}
- },
+ json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}},
)
assert response.status_code == 403
diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py
index 07d1a9d14f9..1022a48a552 100644
--- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py
+++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py
@@ -209,7 +209,9 @@ class TestManagedResponsesWSFirstMessage:
entering its receive loop. Regression for clients that connect without
?model= (e.g. Codex) and send model inside the first response.create event.
"""
- from litellm.responses.streaming_iterator import ManagedResponsesWebSocketHandler
+ from litellm.responses.streaming_iterator import (
+ ManagedResponsesWebSocketHandler,
+ )
first = json.dumps(
{
@@ -250,7 +252,9 @@ class TestManagedResponsesWSFirstMessage:
@pytest.mark.asyncio
async def test_no_first_message_falls_through_to_loop(self):
"""When first_message is None, run() goes straight to receive_text()."""
- from litellm.responses.streaming_iterator import ManagedResponsesWebSocketHandler
+ from litellm.responses.streaming_iterator import (
+ ManagedResponsesWebSocketHandler,
+ )
subsequent = json.dumps({"type": "response.create", "model": "gpt-4o-mini"})
@@ -285,7 +289,9 @@ class TestResponsesWSStreamingFirstMessage:
"""
from litellm.responses.streaming_iterator import ResponsesWebSocketStreaming
- first = json.dumps({"type": "response.create", "model": "gpt-4o-mini", "input": []})
+ first = json.dumps(
+ {"type": "response.create", "model": "gpt-4o-mini", "input": []}
+ )
ws = MagicMock()
ws.receive_text = AsyncMock(side_effect=Exception("disconnect"))
@@ -355,6 +361,7 @@ class TestWSModelExtraction:
from litellm.proxy.response_api_endpoints.endpoints import (
_extract_model_from_first_ws_event,
)
+
event = {"type": "response.create", "model": "gpt-4o", "input": "hello"}
assert _extract_model_from_first_ws_event(event) == "gpt-4o"
@@ -362,13 +369,18 @@ class TestWSModelExtraction:
from litellm.proxy.response_api_endpoints.endpoints import (
_extract_model_from_first_ws_event,
)
- event = {"type": "response.create", "response": {"model": "gpt-4o", "input": "hello"}}
+
+ event = {
+ "type": "response.create",
+ "response": {"model": "gpt-4o", "input": "hello"},
+ }
assert _extract_model_from_first_ws_event(event) == "gpt-4o"
def test_nested_format_takes_precedence_over_flat(self):
from litellm.proxy.response_api_endpoints.endpoints import (
_extract_model_from_first_ws_event,
)
+
event = {
"type": "response.create",
"model": "flat-model",
@@ -380,6 +392,7 @@ class TestWSModelExtraction:
from litellm.proxy.response_api_endpoints.endpoints import (
_extract_model_from_first_ws_event,
)
+
event = {"type": "response.create", "input": "hello"}
assert _extract_model_from_first_ws_event(event) is None
@@ -696,18 +709,14 @@ class TestManagedResponsesSameProvider:
assert "custom_llm_provider" not in call_kwargs
def test_unresolvable_connection_model_falls_back_to_custom_provider(self):
- handler = self._handler(
- "my-custom-deployment", custom_llm_provider="openai"
- )
+ handler = self._handler("my-custom-deployment", custom_llm_provider="openai")
assert handler._same_provider("gpt-4o-mini") is True
call_kwargs: dict = {}
handler._inject_credentials(call_kwargs, model="gpt-4o-mini")
assert call_kwargs["custom_llm_provider"] == "openai"
def test_unresolvable_connection_model_still_drops_cross_provider(self):
- handler = self._handler(
- "my-custom-deployment", custom_llm_provider="openai"
- )
+ handler = self._handler("my-custom-deployment", custom_llm_provider="openai")
call_kwargs: dict = {}
handler._inject_credentials(call_kwargs, model="vertex_ai/gemini-2.0-flash")
assert "custom_llm_provider" not in call_kwargs
diff --git a/tests/test_litellm/proxy/test_caching_routes.py b/tests/test_litellm/proxy/test_caching_routes.py
index 840ba054cc9..7f108c9419a 100644
--- a/tests/test_litellm/proxy/test_caching_routes.py
+++ b/tests/test_litellm/proxy/test_caching_routes.py
@@ -201,12 +201,12 @@ def test_cache_ping_no_cache_does_not_expose_internals():
raw_body = response.text
# No Python traceback or source-file paths must appear in the response
- assert "traceback" not in raw_body.lower(), (
- "CWE-209: Python traceback exposed in /cache/ping no-cache response"
- )
- assert 'File "' not in raw_body, (
- "CWE-209: Python stack frame paths exposed in /cache/ping no-cache response"
- )
+ assert (
+ "traceback" not in raw_body.lower()
+ ), "CWE-209: Python traceback exposed in /cache/ping no-cache response"
+ assert (
+ 'File "' not in raw_body
+ ), "CWE-209: Python stack frame paths exposed in /cache/ping no-cache response"
data = response.json()
# Response must use the ProxyException envelope
diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py
index 364861d6e31..2193f817ad2 100644
--- a/tests/test_litellm/proxy/test_common_request_processing.py
+++ b/tests/test_litellm/proxy/test_common_request_processing.py
@@ -75,7 +75,9 @@ class TestProxyBaseLLMRequestProcessing:
assert result.headers["x-litellm-version"] == "test-version"
@pytest.mark.asyncio
- async def test_base_passthrough_process_llm_request_returns_fastapi_response_from_guardrails(self, monkeypatch):
+ async def test_base_passthrough_process_llm_request_returns_fastapi_response_from_guardrails(
+ self, monkeypatch
+ ):
"""Post-call guardrails return a FastAPI Response; must not call httpx aread()."""
import json
@@ -258,7 +260,9 @@ class TestProxyBaseLLMRequestProcessing:
assert kwargs["request_headers"] == {"authorization": "Bearer sk-test"}
@pytest.mark.asyncio
- async def test_common_processing_pre_call_logic_pre_call_hook_receives_litellm_call_id(self, monkeypatch):
+ async def test_common_processing_pre_call_logic_pre_call_hook_receives_litellm_call_id(
+ self, monkeypatch
+ ):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
mock_request = MagicMock(spec=Request)
mock_request.headers = {}
@@ -266,12 +270,16 @@ class TestProxyBaseLLMRequestProcessing:
async def mock_add_litellm_data_to_request(*args, **kwargs):
return {}
- async def mock_common_processing_pre_call_logic(user_api_key_dict, data, call_type):
+ async def mock_common_processing_pre_call_logic(
+ user_api_key_dict, data, call_type
+ ):
data_copy = copy.deepcopy(data)
return data_copy
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
- mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=mock_common_processing_pre_call_logic)
+ mock_proxy_logging_obj.pre_call_hook = AsyncMock(
+ side_effect=mock_common_processing_pre_call_logic
+ )
monkeypatch.setattr(
litellm.proxy.common_request_processing,
"add_litellm_data_to_request",
@@ -307,7 +315,9 @@ class TestProxyBaseLLMRequestProcessing:
pytest.fail("litellm_call_id is not a valid UUID")
assert data_passed["litellm_call_id"] == returned_data["litellm_call_id"]
- def test_add_dd_apm_tags_for_litellm_call_id_uses_dd_tracing_helper(self, monkeypatch):
+ def test_add_dd_apm_tags_for_litellm_call_id_uses_dd_tracing_helper(
+ self, monkeypatch
+ ):
mock_set_active_span_tag = MagicMock(return_value=True)
import litellm.proxy.dd_span_tagger
@@ -319,10 +329,14 @@ class TestProxyBaseLLMRequestProcessing:
DDSpanTagger.tag_call_id("test-call-id")
- mock_set_active_span_tag.assert_called_once_with("litellm.call_id", "test-call-id")
+ mock_set_active_span_tag.assert_called_once_with(
+ "litellm.call_id", "test-call-id"
+ )
@pytest.mark.asyncio
- async def test_should_apply_hierarchical_router_settings_as_override(self, monkeypatch):
+ async def test_should_apply_hierarchical_router_settings_as_override(
+ self, monkeypatch
+ ):
"""
Test that hierarchical router settings are stored as router_settings_override
instead of creating a full user_config with model_list.
@@ -337,12 +351,16 @@ class TestProxyBaseLLMRequestProcessing:
async def mock_add_litellm_data_to_request(*args, **kwargs):
return {}
- async def mock_common_processing_pre_call_logic(user_api_key_dict, data, call_type):
+ async def mock_common_processing_pre_call_logic(
+ user_api_key_dict, data, call_type
+ ):
data_copy = copy.deepcopy(data)
return data_copy
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
- mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=mock_common_processing_pre_call_logic)
+ mock_proxy_logging_obj.pre_call_hook = AsyncMock(
+ side_effect=mock_common_processing_pre_call_logic
+ )
monkeypatch.setattr(
litellm.proxy.common_request_processing,
"add_litellm_data_to_request",
@@ -358,7 +376,9 @@ class TestProxyBaseLLMRequestProcessing:
"timeout": 30.0,
"num_retries": 3,
}
- mock_proxy_config._get_hierarchical_router_settings = AsyncMock(return_value=mock_router_settings)
+ mock_proxy_config._get_hierarchical_router_settings = AsyncMock(
+ return_value=mock_router_settings
+ )
mock_llm_router = MagicMock()
@@ -412,18 +432,24 @@ class TestProxyBaseLLMRequestProcessing:
# Test with stream timeout header
headers_with_timeout = {"x-litellm-stream-timeout": "30.5"}
- result = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers_with_timeout)
+ result = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(
+ headers_with_timeout
+ )
assert result == 30.5
# Test without stream timeout header
headers_without_timeout = {}
- result = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers_without_timeout)
+ result = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(
+ headers_without_timeout
+ )
assert result is None
# Test with invalid header value (should raise ValueError when converting to float)
headers_with_invalid = {"x-litellm-stream-timeout": "invalid"}
with pytest.raises(ValueError):
- LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers_with_invalid)
+ LiteLLMProxyRequestSetup._get_stream_timeout_from_request(
+ headers_with_invalid
+ )
@pytest.mark.asyncio
async def test_build_litellm_proxy_success_headers_from_llm_response(self):
@@ -518,7 +544,9 @@ class TestProxyBaseLLMRequestProcessing:
)
assert headers["x-litellm-model-id"] == "stream-model-id"
- assert headers["x-litellm-model-api-base"] == ("https://generativelanguage.googleapis.com/v1beta")
+ assert headers["x-litellm-model-api-base"] == (
+ "https://generativelanguage.googleapis.com/v1beta"
+ )
assert headers["llm_provider-x"] == "y"
@pytest.mark.asyncio
@@ -958,7 +986,9 @@ class TestProxyBaseLLMRequestProcessing:
assert "x-litellm-key-spend" in headers_1
expected_spend_1 = 0.001 + 0.0005 # Initial spend + current request cost
- assert float(headers_1["x-litellm-key-spend"]) == pytest.approx(expected_spend_1, abs=1e-10)
+ assert float(headers_1["x-litellm-key-spend"]) == pytest.approx(
+ expected_spend_1, abs=1e-10
+ )
assert float(headers_1["x-litellm-response-cost"]) == response_cost_1
# Test case 2: response_cost is provided as string
@@ -971,7 +1001,9 @@ class TestProxyBaseLLMRequestProcessing:
assert "x-litellm-key-spend" in headers_2
expected_spend_2 = 0.001 + 0.0003 # Initial spend + current request cost
- assert float(headers_2["x-litellm-key-spend"]) == pytest.approx(expected_spend_2, abs=1e-10)
+ assert float(headers_2["x-litellm-key-spend"]) == pytest.approx(
+ expected_spend_2, abs=1e-10
+ )
# Test case 3: response_cost is None (should use original spend)
headers_3 = ProxyBaseLLMRequestProcessing.get_custom_headers(
@@ -981,7 +1013,9 @@ class TestProxyBaseLLMRequestProcessing:
)
assert "x-litellm-key-spend" in headers_3
- assert float(headers_3["x-litellm-key-spend"]) == 0.001 # Should use original spend
+ assert (
+ float(headers_3["x-litellm-key-spend"]) == 0.001
+ ) # Should use original spend
# Test case 4: response_cost is 0 (should not change spend)
headers_4 = ProxyBaseLLMRequestProcessing.get_custom_headers(
@@ -991,7 +1025,9 @@ class TestProxyBaseLLMRequestProcessing:
)
assert "x-litellm-key-spend" in headers_4
- assert float(headers_4["x-litellm-key-spend"]) == 0.001 # Should remain unchanged for 0 cost
+ assert (
+ float(headers_4["x-litellm-key-spend"]) == 0.001
+ ) # Should remain unchanged for 0 cost
# Test case 5: user_api_key_dict.spend is None (should default to 0.0)
mock_user_api_key_dict.spend = None
@@ -1013,7 +1049,9 @@ class TestProxyBaseLLMRequestProcessing:
)
assert "x-litellm-key-spend" in headers_6
- assert float(headers_6["x-litellm-key-spend"]) == 0.001 # Should use original spend
+ assert (
+ float(headers_6["x-litellm-key-spend"]) == 0.001
+ ) # Should use original spend
# Test case 7: response_cost is invalid string (should fallback to original spend)
headers_7 = ProxyBaseLLMRequestProcessing.get_custom_headers(
@@ -1023,7 +1061,9 @@ class TestProxyBaseLLMRequestProcessing:
)
assert "x-litellm-key-spend" in headers_7
- assert float(headers_7["x-litellm-key-spend"]) == 0.001 # Should use original spend on error
+ assert (
+ float(headers_7["x-litellm-key-spend"]) == 0.001
+ ) # Should use original spend on error
@pytest.mark.asyncio
async def test_queue_time_seconds_is_set_in_metadata(self, monkeypatch):
@@ -1084,10 +1124,12 @@ class TestProxyBaseLLMRequestProcessing:
# Verify queue_time_seconds is set and non-negative
metadata = returned_data.get("metadata", {})
- assert "queue_time_seconds" in metadata, "queue_time_seconds should be set in metadata"
- assert metadata["queue_time_seconds"] >= 0.5, (
- f"queue_time_seconds should be at least 0.5, got {metadata['queue_time_seconds']}"
- )
+ assert (
+ "queue_time_seconds" in metadata
+ ), "queue_time_seconds should be set in metadata"
+ assert (
+ metadata["queue_time_seconds"] >= 0.5
+ ), f"queue_time_seconds should be at least 0.5, got {metadata['queue_time_seconds']}"
@pytest.mark.asyncio
@@ -1238,7 +1280,9 @@ class TestCommonRequestProcessingHelpers:
the original status code instead of hardcoding 500.
"""
mock_gen = AsyncMock()
- mock_gen.__anext__.side_effect = HTTPException(status_code=400, detail="Content blocked by guardrail")
+ mock_gen.__anext__.side_effect = HTTPException(
+ status_code=400, detail="Content blocked by guardrail"
+ )
response = await create_response(mock_gen, "text/event-stream", {})
assert response.status_code == 400
@@ -1332,8 +1376,14 @@ class TestCommonRequestProcessingHelpers:
response = await create_response(mock_gen, "text/event-stream", {})
content = await self.consume_stream(response)
payload = json.loads(content[0][len("data: ") :].strip())
- assert payload["error"]["message"] == "MCP request blocked: no rewritable argument field present"
- assert payload["error"]["provider_specific_fields"]["error"]["code"] == "panw_prisma_airs_blocked"
+ assert (
+ payload["error"]["message"]
+ == "MCP request blocked: no rewritable argument field present"
+ )
+ assert (
+ payload["error"]["provider_specific_fields"]["error"]["code"]
+ == "panw_prisma_airs_blocked"
+ )
async def test_serialize_http_exception_detail_helper(self):
"""Direct unit coverage for the L1 helper across all branches."""
@@ -1344,11 +1394,15 @@ class TestCommonRequestProcessingHelpers:
assert _serialize_http_exception_detail("plain") == ("plain", None)
- msg, fields = _serialize_http_exception_detail({"error": "Violated", "extra": "x"})
+ msg, fields = _serialize_http_exception_detail(
+ {"error": "Violated", "extra": "x"}
+ )
assert msg == "Violated"
assert fields == {"error": "Violated", "extra": "x"}
- msg, fields = _serialize_http_exception_detail({"error": {"message": "blocked", "code": "x"}})
+ msg, fields = _serialize_http_exception_detail(
+ {"error": {"message": "blocked", "code": "x"}}
+ )
assert msg == "blocked"
assert fields == {"error": {"message": "blocked", "code": "x"}}
@@ -1388,7 +1442,9 @@ class TestCommonRequestProcessingHelpers:
yield "data: [DONE]\n\n"
custom_headers = {"X-Custom-Header": "TestValue"}
- response = await create_response(mock_generator(), "text/event-stream", custom_headers)
+ response = await create_response(
+ mock_generator(), "text/event-stream", custom_headers
+ )
assert response.headers["x-custom-header"] == "TestValue"
async def test_create_streaming_response_disables_proxy_buffering(self):
@@ -1408,7 +1464,9 @@ class TestCommonRequestProcessingHelpers:
error_stream.__anext__.side_effect = ValueError("boom")
for generator in (normal_stream(), empty_stream(), error_stream):
- response = await create_response(generator, "text/event-stream", {"X-Custom-Header": "keep"})
+ response = await create_response(
+ generator, "text/event-stream", {"X-Custom-Header": "keep"}
+ )
assert isinstance(response, StreamingResponse)
assert response.headers["x-accel-buffering"] == "no"
assert response.headers["cache-control"] == "no-cache"
@@ -1507,9 +1565,9 @@ class TestCommonRequestProcessingHelpers:
for i, call in enumerate(actual_calls):
args, kwargs = call
- assert args[0] == "streaming.chunk.yield", (
- f"Call {i} should have operation name 'streaming.chunk.yield', got {args[0]}"
- )
+ assert (
+ args[0] == "streaming.chunk.yield"
+ ), f"Call {i} should have operation name 'streaming.chunk.yield', got {args[0]}"
async def test_create_streaming_response_skips_dd_trace_when_disabled(self):
"""When DD tracing is disabled (the default), the per-chunk span
@@ -1690,7 +1748,9 @@ class TestOverrideOpenAIResponseModel:
# _hidden_params is an attribute (not a dict key) accessed via getattr
response_obj = MagicMock()
response_obj.model = fallback_model
- response_obj._hidden_params = {"additional_headers": {"x-litellm-attempted-fallbacks": 1}}
+ response_obj._hidden_params = {
+ "additional_headers": {"x-litellm-attempted-fallbacks": 1}
+ }
# Call the function - should preserve fallback model
_override_openai_response_model(
@@ -1817,7 +1877,9 @@ class TestOverrideOpenAIResponseModel:
# Create a mock object response
response_obj = MagicMock()
response_obj.model = downstream_model
- response_obj._hidden_params = {"additional_headers": {"x-litellm-attempted-fallbacks": None}}
+ response_obj._hidden_params = {
+ "additional_headers": {"x-litellm-attempted-fallbacks": None}
+ }
# Call the function - should override to requested model
_override_openai_response_model(
@@ -1862,7 +1924,9 @@ class TestOverrideOpenAIResponseModel:
# Create a mock object response
response_obj = MagicMock()
response_obj.model = fallback_model
- response_obj._hidden_params = {"additional_headers": {"x-litellm-attempted-fallbacks": 1}}
+ response_obj._hidden_params = {
+ "additional_headers": {"x-litellm-attempted-fallbacks": 1}
+ }
# Call the function with None requested_model
_override_openai_response_model(
@@ -2028,7 +2092,10 @@ class TestIsAzureModelRouterRequest:
def test_detects_model_router_with_underscore(self):
assert _is_azure_model_router_request("azure_ai/model_router") is True
- assert _is_azure_model_router_request("azure_ai/model_router/my-deployment") is True
+ assert (
+ _is_azure_model_router_request("azure_ai/model_router/my-deployment")
+ is True
+ )
def test_detects_model_router_with_hyphen(self):
assert _is_azure_model_router_request("azure_ai/model-router") is True
@@ -2252,7 +2319,9 @@ class TestDDSpanTaggerTagRequest:
def test_tags_key_alias_and_model(self):
"""key_alias and requested_model are set on the span when present."""
- user_key = self._make_user_api_key_dict(key_alias="my-prod-key", token="hashed123")
+ user_key = self._make_user_api_key_dict(
+ key_alias="my-prod-key", token="hashed123"
+ )
with patch("litellm.proxy.dd_span_tagger.set_active_span_tag") as mock_set_tag:
DDSpanTagger.tag_request(
@@ -2286,7 +2355,9 @@ class TestDDSpanTaggerTagRequest:
requested_model="claude-3-5-sonnet",
)
- mock_set_tag.assert_called_once_with("litellm.requested_model", "claude-3-5-sonnet")
+ mock_set_tag.assert_called_once_with(
+ "litellm.requested_model", "claude-3-5-sonnet"
+ )
class TestHasAttributeErrorInChain:
@@ -2375,7 +2446,9 @@ class TestHandleLLMApiExceptionDictDetail:
)
proxy_exc = await self._invoke(exc)
assert proxy_exc.message == "Violated guardrail policy"
- assert proxy_exc.provider_specific_fields["guardrail_name"] == "bedrock-pii-guard"
+ assert (
+ proxy_exc.provider_specific_fields["guardrail_name"] == "bedrock-pii-guard"
+ )
# No Python repr leakage of the dict into the message field.
assert "{'error':" not in proxy_exc.message
@@ -2719,7 +2792,9 @@ class TestAsyncStreamingDataGeneratorFastPath:
proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock())
hook_spy = AsyncMock(side_effect=lambda **kw: kw["response"])
- monkeypatch.setattr(proxy_logging_obj, "async_post_call_streaming_hook", hook_spy)
+ monkeypatch.setattr(
+ proxy_logging_obj, "async_post_call_streaming_hook", hook_spy
+ )
chunks = [b"event: a\ndata: {}\n\n", b"event: b\ndata: {}\n\n"]
out = [
@@ -2752,7 +2827,9 @@ class TestAsyncStreamingDataGeneratorFastPath:
proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock())
hook_spy = AsyncMock(side_effect=lambda **kw: kw["response"])
- monkeypatch.setattr(proxy_logging_obj, "async_post_call_streaming_hook", hook_spy)
+ monkeypatch.setattr(
+ proxy_logging_obj, "async_post_call_streaming_hook", hook_spy
+ )
out = [
c
@@ -2794,7 +2871,9 @@ class TestDisconnectGatherCleanup:
import asyncio
import litellm.proxy.common_request_processing as cpr
- from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
+ from litellm.proxy.common_request_processing import (
+ ProxyBaseLLMRequestProcessing,
+ )
async def slow_llm():
await asyncio.sleep(9999)
@@ -2812,7 +2891,9 @@ class TestDisconnectGatherCleanup:
monkeypatch.setattr(cpr, "route_request", fake_route_request)
- processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"})
+ processing_obj = ProxyBaseLLMRequestProcessing(
+ data={"model": "gemini-2.0-flash"}
+ )
monkeypatch.setattr(
processing_obj,
"common_processing_pre_call_logic",
@@ -2844,7 +2925,9 @@ class TestDisconnectGatherCleanup:
import asyncio
import litellm.proxy.common_request_processing as cpr
- from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
+ from litellm.proxy.common_request_processing import (
+ ProxyBaseLLMRequestProcessing,
+ )
async def fake_gather(*_tasks, **_kwargs):
raise asyncio.CancelledError()
@@ -2859,7 +2942,9 @@ class TestDisconnectGatherCleanup:
monkeypatch.setattr(cpr.asyncio, "gather", fake_gather)
- processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"})
+ processing_obj = ProxyBaseLLMRequestProcessing(
+ data={"model": "gemini-2.0-flash"}
+ )
monkeypatch.setattr(
processing_obj,
"common_processing_pre_call_logic",
@@ -2894,7 +2979,9 @@ class TestDisconnectGatherCleanup:
import asyncio
import litellm.proxy.common_request_processing as cpr
- from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
+ from litellm.proxy.common_request_processing import (
+ ProxyBaseLLMRequestProcessing,
+ )
hook_cancelled = False
@@ -2922,7 +3009,9 @@ class TestDisconnectGatherCleanup:
monkeypatch.setattr(cpr, "route_request", fake_route_request)
- processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"})
+ processing_obj = ProxyBaseLLMRequestProcessing(
+ data={"model": "gemini-2.0-flash"}
+ )
monkeypatch.setattr(
processing_obj,
"common_processing_pre_call_logic",
@@ -2987,7 +3076,9 @@ class TestDisconnectGatherCleanup:
import asyncio
import litellm.proxy.common_request_processing as cpr
- from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
+ from litellm.proxy.common_request_processing import (
+ ProxyBaseLLMRequestProcessing,
+ )
async def failing_llm():
raise ValueError("llm api error")
@@ -3008,7 +3099,9 @@ class TestDisconnectGatherCleanup:
monkeypatch.setattr(cpr, "route_request", fake_route_request)
- processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"})
+ processing_obj = ProxyBaseLLMRequestProcessing(
+ data={"model": "gemini-2.0-flash"}
+ )
monkeypatch.setattr(
processing_obj,
"common_processing_pre_call_logic",
@@ -3059,9 +3152,7 @@ class TestStreamingClientDisconnectLogging:
assert recorded is True
assert request_data["metadata"]["client_disconnected"] is True
- assert (
- request_data["metadata"]["error_information"]["error_code"] == "499"
- )
+ assert request_data["metadata"]["error_information"]["error_code"] == "499"
assert (
mock_logging_obj.model_call_details["litellm_params"]["metadata"][
"error_information"
@@ -3107,7 +3198,9 @@ class TestStreamingClientDisconnectLogging:
request_data = {
"metadata": {},
"litellm_params": {"metadata": {}},
- "litellm_logging_obj": MagicMock(model_call_details={"metadata": {}, "litellm_params": {}}),
+ "litellm_logging_obj": MagicMock(
+ model_call_details={"metadata": {}, "litellm_params": {}}
+ ),
}
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
@@ -3204,6 +3297,8 @@ class TestStreamingClientDisconnectLogging:
assert request_data["metadata"]["error_information"]["error_code"] == "499"
ProxyLogging._callback_capabilities_cache.clear()
+
+
class TestCancelOnDisconnect:
"""
Coverage for the opt-in `general_settings.cancel_on_disconnect` flag:
@@ -3230,9 +3325,7 @@ class TestCancelOnDisconnect:
llm_call = asyncio.get_running_loop().create_future()
disconnect_event = asyncio.Event()
- await _cancel_llm_call_on_client_disconnect(
- request, llm_call, disconnect_event
- )
+ await _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event)
assert llm_call.cancelled()
assert disconnect_event.is_set()
@@ -3265,9 +3358,7 @@ class TestCancelOnDisconnect:
llm_call = asyncio.get_running_loop().create_future()
disconnect_event = asyncio.Event()
- await _cancel_llm_call_on_client_disconnect(
- request, llm_call, disconnect_event
- )
+ await _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event)
assert not llm_call.cancelled()
assert not disconnect_event.is_set()
@@ -3303,9 +3394,7 @@ class TestCancelOnDisconnect:
proxy_logging_obj.post_call_success_hook = AsyncMock(
side_effect=lambda data, user_api_key_dict, response: response
)
- proxy_logging_obj.post_call_response_headers_hook = AsyncMock(
- return_value=None
- )
+ proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value=None)
async def fake_route_request(**kwargs):
return llm_call()
@@ -3409,13 +3498,18 @@ class TestAllmPassthroughRoutePostCallGuardrails:
cb = MagicMock(spec=CustomGuardrail)
cb.guardrail_name = name
- cb.event_hook = [GuardrailEventHooks.pre_call.value, GuardrailEventHooks.post_call.value]
+ cb.event_hook = [
+ GuardrailEventHooks.pre_call.value,
+ GuardrailEventHooks.post_call.value,
+ ]
cb._event_hook_is_event_type = lambda et: et.value in cb.event_hook
cb.should_run_guardrail = MagicMock(return_value=True)
return cb
@pytest.mark.asyncio
- async def test_post_call_hook_receives_parsed_dict_not_httpx_response(self, monkeypatch):
+ async def test_post_call_hook_receives_parsed_dict_not_httpx_response(
+ self, monkeypatch
+ ):
"""
post_call_success_hook must be called with the parsed JSON dict when the
non-streaming allm_passthrough_route response is application/json.
@@ -3452,7 +3546,11 @@ class TestAllmPassthroughRoutePostCallGuardrails:
proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock())
monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", capture_hook)
- with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True):
+ with patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails_for_passthrough",
+ return_value=True,
+ ):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=httpx_response,
@@ -3463,9 +3561,9 @@ class TestAllmPassthroughRoutePostCallGuardrails:
)
assert len(received_responses) == 1
- assert isinstance(received_responses[0], dict), (
- "post_call_success_hook must receive parsed dict, not httpx.Response"
- )
+ assert isinstance(
+ received_responses[0], dict
+ ), "post_call_success_hook must receive parsed dict, not httpx.Response"
assert received_responses[0]["stopReason"] == "end_turn"
assert isinstance(result, Response)
body = json.loads(result.body)
@@ -3502,7 +3600,11 @@ class TestAllmPassthroughRoutePostCallGuardrails:
proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock())
monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", non_dict_hook)
- with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True):
+ with patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails_for_passthrough",
+ return_value=True,
+ ):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=httpx_response,
@@ -3540,7 +3642,11 @@ class TestAllmPassthroughRoutePostCallGuardrails:
hook_spy = AsyncMock()
monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", hook_spy)
- with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True):
+ with patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails_for_passthrough",
+ return_value=True,
+ ):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=httpx_response,
@@ -3581,7 +3687,11 @@ class TestAllmPassthroughRoutePostCallGuardrails:
hook_spy = AsyncMock()
monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", hook_spy)
- with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=False):
+ with patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails_for_passthrough",
+ return_value=False,
+ ):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=httpx_response,
@@ -3637,7 +3747,9 @@ class TestEventStreamAllmPassthroughRoute:
@pytest.mark.asyncio
async def test_bedrock_provider_dispatches_to_handler(self):
stream_bytes = _build_event_stream_frame("messageStart", {"role": "assistant"})
- expected_bytes = _build_event_stream_frame("messageStart", {"role": "assistant"}) + b"extra"
+ expected_bytes = (
+ _build_event_stream_frame("messageStart", {"role": "assistant"}) + b"extra"
+ )
proxy_logging_obj = MagicMock()
user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
@@ -3646,7 +3758,9 @@ class TestEventStreamAllmPassthroughRoute:
"litellm.llms.bedrock.passthrough.guardrail_translation.handler.BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
new=AsyncMock(return_value=expected_bytes),
) as mock_handler:
- processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"})
+ processing_obj = ProxyBaseLLMRequestProcessing(
+ data={"custom_llm_provider": "bedrock"}
+ )
result = await processing_obj._handle_event_stream_allm_passthrough_route(
body_bytes=stream_bytes,
proxy_logging_obj=proxy_logging_obj,
@@ -3661,7 +3775,9 @@ class TestEventStreamAllmPassthroughRoute:
stream_bytes = _build_event_stream_frame("messageStart", {"role": "assistant"})
proxy_logging_obj = MagicMock()
- processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "anthropic"})
+ processing_obj = ProxyBaseLLMRequestProcessing(
+ data={"custom_llm_provider": "anthropic"}
+ )
result = await processing_obj._handle_event_stream_allm_passthrough_route(
body_bytes=stream_bytes,
proxy_logging_obj=proxy_logging_obj,
@@ -3674,10 +3790,15 @@ class TestEventStreamAllmPassthroughRoute:
async def test_non_streaming_response_includes_custom_headers(self):
import json
- body = {"output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}}}
+ body = {
+ "output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}}
+ }
mock_response = MagicMock()
mock_response.status_code = 200
- mock_response.headers = {"content-type": "application/json", "content-length": "99"}
+ mock_response.headers = {
+ "content-type": "application/json",
+ "content-length": "99",
+ }
mock_response.aread = AsyncMock(return_value=json.dumps(body).encode())
async def mock_hook(data, user_api_key_dict, response):
@@ -3693,7 +3814,11 @@ class TestEventStreamAllmPassthroughRoute:
"content-length": "99",
}
- with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True):
+ with patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails_for_passthrough",
+ return_value=True,
+ ):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=mock_response,
@@ -3777,14 +3902,17 @@ class TestAllmPassthroughStreamingProviderGate:
processing_obj = self._build_processing_obj("anthropic")
chunks = [b"chunk-1", b"chunk-2"]
- with patch.object(
- ProxyBaseLLMRequestProcessing,
- "_has_post_call_guardrails",
- return_value=False,
- ), patch.object(
- ProxyBaseLLMRequestProcessing,
- "_has_post_call_guardrails_for_passthrough",
- return_value=True,
+ with (
+ patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails",
+ return_value=False,
+ ),
+ patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails_for_passthrough",
+ return_value=True,
+ ),
):
result = await self._run(processing_obj, monkeypatch, chunks)
@@ -3801,19 +3929,23 @@ class TestAllmPassthroughStreamingProviderGate:
)
chunks = [b"raw-1", b"raw-2"]
- with patch.object(
- ProxyBaseLLMRequestProcessing,
- "_has_post_call_guardrails",
- return_value=False,
- ), patch.object(
- ProxyBaseLLMRequestProcessing,
- "_has_post_call_guardrails_for_passthrough",
- return_value=True,
- ), patch(
- "litellm.llms.bedrock.passthrough.guardrail_translation.handler."
- "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
- new=AsyncMock(return_value=b"modified-body"),
- ) as mock_handler:
+ with (
+ patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails",
+ return_value=False,
+ ),
+ patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails_for_passthrough",
+ return_value=True,
+ ),
+ patch(
+ "litellm.llms.bedrock.passthrough.guardrail_translation.handler."
+ "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
+ new=AsyncMock(return_value=b"modified-body"),
+ ) as mock_handler,
+ ):
result = await self._run(processing_obj, monkeypatch, chunks)
assert isinstance(result, Response)
@@ -3829,19 +3961,23 @@ class TestAllmPassthroughStreamingProviderGate:
)
chunks = [b"raw-1", b"raw-2"]
- with patch.object(
- ProxyBaseLLMRequestProcessing,
- "_has_post_call_guardrails",
- return_value=False,
- ), patch.object(
- ProxyBaseLLMRequestProcessing,
- "_has_post_call_guardrails_for_passthrough",
- return_value=True,
- ), patch(
- "litellm.llms.bedrock.passthrough.guardrail_translation.handler."
- "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
- new=AsyncMock(return_value=b"modified-body"),
- ) as mock_handler:
+ with (
+ patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails",
+ return_value=False,
+ ),
+ patch.object(
+ ProxyBaseLLMRequestProcessing,
+ "_has_post_call_guardrails_for_passthrough",
+ return_value=True,
+ ),
+ patch(
+ "litellm.llms.bedrock.passthrough.guardrail_translation.handler."
+ "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
+ new=AsyncMock(return_value=b"modified-body"),
+ ) as mock_handler,
+ ):
result = await self._run(processing_obj, monkeypatch, chunks)
assert isinstance(result, StreamingResponse)
diff --git a/tests/test_litellm/proxy/test_component_allowlists.py b/tests/test_litellm/proxy/test_component_allowlists.py
index 926ce3bee66..9e27b6229f5 100644
--- a/tests/test_litellm/proxy/test_component_allowlists.py
+++ b/tests/test_litellm/proxy/test_component_allowlists.py
@@ -96,16 +96,19 @@ def test_gateway_plus_backend_covers_full_app():
def test_backend_mount_paths_defined():
"""BACKEND_MOUNT_PATHS constant must exist and be a frozenset."""
- assert isinstance(BACKEND_MOUNT_PATHS, frozenset), \
- f"BACKEND_MOUNT_PATHS must be a frozenset, got {type(BACKEND_MOUNT_PATHS)}"
- assert len(BACKEND_MOUNT_PATHS) > 0, \
- "BACKEND_MOUNT_PATHS must contain at least one Mount path"
+ assert isinstance(
+ BACKEND_MOUNT_PATHS, frozenset
+ ), f"BACKEND_MOUNT_PATHS must be a frozenset, got {type(BACKEND_MOUNT_PATHS)}"
+ assert (
+ len(BACKEND_MOUNT_PATHS) > 0
+ ), "BACKEND_MOUNT_PATHS must contain at least one Mount path"
def test_swagger_mount_in_backend_allowlist():
"""The /swagger Mount must be in BACKEND_MOUNT_PATHS."""
- assert "/swagger" in BACKEND_MOUNT_PATHS, \
- "/swagger Mount path must be in BACKEND_MOUNT_PATHS"
+ assert (
+ "/swagger" in BACKEND_MOUNT_PATHS
+ ), "/swagger Mount path must be in BACKEND_MOUNT_PATHS"
def test_backend_keeps_swagger_mount():
@@ -115,8 +118,9 @@ def test_backend_keeps_swagger_mount():
for r in app.router.routes
if isinstance(r, Mount) and getattr(r, "path", None) in BACKEND_MOUNT_PATHS
}
- assert "/swagger" in backend_mounts, \
- "/swagger Mount is expected on the proxy app and should be in BACKEND_MOUNT_PATHS"
+ assert (
+ "/swagger" in backend_mounts
+ ), "/swagger Mount is expected on the proxy app and should be in BACKEND_MOUNT_PATHS"
def test_backend_drops_non_allowlisted_mounts():
@@ -128,8 +132,10 @@ def test_backend_drops_non_allowlisted_mounts():
}
non_backend_mounts = all_mounts - BACKEND_MOUNT_PATHS
- assert len(non_backend_mounts) > 0, \
- "Expected at least one non-backend Mount (e.g., /ui, /_next) to verify filtering logic"
+ assert (
+ len(non_backend_mounts) > 0
+ ), "Expected at least one non-backend Mount (e.g., /ui, /_next) to verify filtering logic"
for mount_path in non_backend_mounts:
- assert mount_path not in BACKEND_MOUNT_PATHS, \
- f"Mount {mount_path} should not be in BACKEND_MOUNT_PATHS"
+ assert (
+ mount_path not in BACKEND_MOUNT_PATHS
+ ), f"Mount {mount_path} should not be in BACKEND_MOUNT_PATHS"
diff --git a/tests/test_litellm/proxy/test_dynamic_mcp_route.py b/tests/test_litellm/proxy/test_dynamic_mcp_route.py
index 592cebd957c..bf11ce7f69e 100644
--- a/tests/test_litellm/proxy/test_dynamic_mcp_route.py
+++ b/tests/test_litellm/proxy/test_dynamic_mcp_route.py
@@ -19,7 +19,6 @@ from unittest.mock import ANY, AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
-
# ---------------------------------------------------------------------------
# helpers
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/test_fastapi_offline_routes.py b/tests/test_litellm/proxy/test_fastapi_offline_routes.py
index f3fc3d3ea28..7788e54b0b7 100644
--- a/tests/test_litellm/proxy/test_fastapi_offline_routes.py
+++ b/tests/test_litellm/proxy/test_fastapi_offline_routes.py
@@ -1,7 +1,7 @@
"""
Unit test for testing /routes endpoint with FastAPIOffline app initialization.
-This test verifies that the /routes endpoint works correctly when the proxy
+This test verifies that the /routes endpoint works correctly when the proxy
server is initialized using FastAPIOffline instead of regular FastAPI.
"""
diff --git a/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py b/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py
index 2d8a9f30c1b..5d1767bfb02 100644
--- a/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py
+++ b/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py
@@ -233,4 +233,6 @@ async def test_filter_db_fallback_receives_resolved_model_names():
"gpt-4o",
"gpt-5",
}, f"DB query should receive resolved model names, got {queried_names}"
- assert "Group-A" not in queried_names, "Raw access group name should not be in DB query"
+ assert (
+ "Group-A" not in queried_names
+ ), "Raw access group name should not be in DB query"
diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py
index 6b692180559..0e921c341c6 100644
--- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py
+++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py
@@ -4150,7 +4150,9 @@ class TestApplyClientTagPolicyPreAuth:
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:tag:paid":
return 0.50
return fallback_spend
@@ -4207,7 +4209,9 @@ class TestApplyClientTagPolicyPreAuth:
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:tag:tenant:acme":
return 0.50
return fallback_spend
@@ -4362,7 +4366,9 @@ class TestApplyKeyTagsPreAuth:
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:tag:engineering":
return 0.50
return fallback_spend
@@ -4413,7 +4419,9 @@ class TestApplyKeyTagsPreAuth:
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
- async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
+ async def mock_get_current_spend(
+ counter_key, fallback_spend, max_budget=None, **kwargs
+ ):
if counter_key == "spend:tag:engineering":
return 0.05
return fallback_spend
diff --git a/tests/test_litellm/proxy/test_modify_response_streaming_passthrough.py b/tests/test_litellm/proxy/test_modify_response_streaming_passthrough.py
index da57d9c616e..9731abb9592 100644
--- a/tests/test_litellm/proxy/test_modify_response_streaming_passthrough.py
+++ b/tests/test_litellm/proxy/test_modify_response_streaming_passthrough.py
@@ -39,22 +39,24 @@ async def _run_streaming_block_and_get_wrapper(exception):
user_api_key_dict = UserAPIKeyAuth()
outer_body = {"model": "gpt-4o", "messages": [], "stream": True}
- with patch(
- "litellm.proxy.proxy_server._read_request_body",
- new_callable=AsyncMock,
- return_value=outer_body,
- ), patch(
- "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request",
- new_callable=AsyncMock,
- side_effect=exception,
- ), patch(
- "litellm.proxy.proxy_server.proxy_logging_obj"
- ) as mock_proxy_logging, patch(
- "litellm.proxy.proxy_server.select_data_generator",
- return_value=iter([]),
- ), patch(
- "litellm.CustomStreamWrapper"
- ) as mock_csw:
+ with (
+ patch(
+ "litellm.proxy.proxy_server._read_request_body",
+ new_callable=AsyncMock,
+ return_value=outer_body,
+ ),
+ patch(
+ "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request",
+ new_callable=AsyncMock,
+ side_effect=exception,
+ ),
+ patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
+ patch(
+ "litellm.proxy.proxy_server.select_data_generator",
+ return_value=iter([]),
+ ),
+ patch("litellm.CustomStreamWrapper") as mock_csw,
+ ):
mock_proxy_logging.post_call_failure_hook = AsyncMock()
await chat_completion(
diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py
index a909c510581..3d60eda04e9 100644
--- a/tests/test_litellm/proxy/test_proxy_utils.py
+++ b/tests/test_litellm/proxy/test_proxy_utils.py
@@ -604,6 +604,8 @@ def test_create_model_info_response_no_router_keeps_base_fields():
"created": response["created"],
"owned_by": "openai",
}
+
+
class TestPostCallFailureHookLLMExceptionAlerting:
"""The llm_exceptions alert is for infra / LLM-API failures, not user
errors (https://github.com/BerriAI/litellm/issues/3395). Already-normalized
diff --git a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py
index d0cb5ec5465..3cb2f5e434c 100644
--- a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py
+++ b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py
@@ -16,7 +16,6 @@ import litellm.proxy.proxy_server as ps
from litellm.caching.caching import RedisCache
from litellm.caching.dual_cache import DualCache
-
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py
index 47dc6e6d37d..ff5da38f84c 100644
--- a/tests/test_litellm/proxy/test_route_llm_request.py
+++ b/tests/test_litellm/proxy/test_route_llm_request.py
@@ -149,9 +149,9 @@ async def test_route_request_vector_store_routes_model_none_no_api_key_in_body()
mock_method.assert_called_once()
actual_kwargs = mock_method.call_args.kwargs
for key, value in data.items():
- assert actual_kwargs.get(key) == value, (
- f"{route_type}: expected {key}={value!r}, got {actual_kwargs.get(key)!r}"
- )
+ assert (
+ actual_kwargs.get(key) == value
+ ), f"{route_type}: expected {key}={value!r}, got {actual_kwargs.get(key)!r}"
llm_router.reset_mock()
diff --git a/tests/test_litellm/proxy/test_team_org_move.py b/tests/test_litellm/proxy/test_team_org_move.py
index 2dc961bec85..ad5f9c38594 100644
--- a/tests/test_litellm/proxy/test_team_org_move.py
+++ b/tests/test_litellm/proxy/test_team_org_move.py
@@ -6,6 +6,7 @@ Covers the SSO/Entra scenario where:
- Non-proxy-admins (team admins) must have all team members pre-added to the org,
preserving the original security model (no privilege escalation via team move).
"""
+
from unittest.mock import AsyncMock, MagicMock
import pytest
@@ -47,10 +48,10 @@ def _make_org(organization_id="org-1", members=None, models=None):
def _make_team(team_id="team-1", member_ids=None, organization_id=None):
- members = [
- Member(user_id=uid, role="user") for uid in (member_ids or [])
- ]
- members.append(Member(user_id=SpecialProxyStrings.default_user_id.value, role="admin"))
+ members = [Member(user_id=uid, role="user") for uid in (member_ids or [])]
+ members.append(
+ Member(user_id=SpecialProxyStrings.default_user_id.value, role="admin")
+ )
return LiteLLM_TeamTable(
team_id=team_id,
team_alias="test-team",
@@ -120,12 +121,18 @@ class TestValidateTeamOrgChange:
team = _make_team(member_ids=["u1"], organization_id="org-1")
org = _make_org(organization_id="org-1")
- assert validate_team_org_change(
- team=team, organization=org, llm_router=router, is_proxy_admin=False
- ) is True
- assert validate_team_org_change(
- team=team, organization=org, llm_router=router, is_proxy_admin=True
- ) is True
+ assert (
+ validate_team_org_change(
+ team=team, organization=org, llm_router=router, is_proxy_admin=False
+ )
+ is True
+ )
+ assert (
+ validate_team_org_change(
+ team=team, organization=org, llm_router=router, is_proxy_admin=True
+ )
+ is True
+ )
def test_default_user_excluded_from_membership_check(self):
"""default_user_id is never checked for org membership."""
@@ -149,6 +156,7 @@ class TestAutoAddTeamMembersToOrg:
mock_add = AsyncMock()
import litellm.proxy.management_endpoints.team_endpoints as te
+
original = te.add_member_to_organization
te.add_member_to_organization = mock_add
@@ -174,6 +182,7 @@ class TestAutoAddTeamMembersToOrg:
mock_add = AsyncMock()
import litellm.proxy.management_endpoints.team_endpoints as te
+
original = te.add_member_to_organization
te.add_member_to_organization = mock_add
@@ -197,6 +206,7 @@ class TestAutoAddTeamMembersToOrg:
mock_add = AsyncMock()
import litellm.proxy.management_endpoints.team_endpoints as te
+
original = te.add_member_to_organization
te.add_member_to_organization = mock_add
@@ -219,6 +229,7 @@ class TestAutoAddTeamMembersToOrg:
mock_add = AsyncMock(side_effect=Exception("duplicate key"))
import litellm.proxy.management_endpoints.team_endpoints as te
+
original = te.add_member_to_organization
te.add_member_to_organization = mock_add
diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py
index d1270b60b19..b67488ffa28 100644
--- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py
+++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py
@@ -19,9 +19,7 @@ from litellm.proxy.utils import _cache_user_row
async def test_cache_user_row_caches_on_miss(
mock_dual_cache: Any,
) -> None:
- user_row = SimpleNamespace(
- user_id="u1", spend=2.5, max_budget=10.0, name="Alice"
- )
+ user_row = SimpleNamespace(user_id="u1", spend=2.5, max_budget=10.0, name="Alice")
user_row.model_dump_json = MagicMock(
return_value='{"user_id":"u1","spend":2.5,"max_budget":10.0,"name":"Alice"}'
)
diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py
index 761835078f4..5e2b45a17a2 100644
--- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py
+++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py
@@ -31,9 +31,7 @@ from litellm.proxy.utils import (
@pytest.fixture(autouse=True)
-def _swap_config_cache(
- monkeypatch: pytest.MonkeyPatch, mock_dual_cache: Any
-) -> Any:
+def _swap_config_cache(monkeypatch: pytest.MonkeyPatch, mock_dual_cache: Any) -> Any:
"""Replace the module-level cache so tests see a clean store per run."""
monkeypatch.setattr(utils_mod, "litellm_config_cache", mock_dual_cache)
return mock_dual_cache
@@ -198,7 +196,9 @@ async def test_invalidate_config_param_evicts_from_cache(
_swap_config_cache: Any,
) -> None:
cache_key = _config_cache_key("p4")
- await _swap_config_cache.async_set_cache(cache_key, {"param_name": "p4", "param_value": 1})
+ await _swap_config_cache.async_set_cache(
+ cache_key, {"param_name": "p4", "param_value": 1}
+ )
await invalidate_config_param("p4")
actual = {
"store_empty": _swap_config_cache._store == {},
diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py
index 2fedd6bb134..f296123078d 100644
--- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py
+++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py
@@ -31,7 +31,6 @@ import pytest
from litellm.proxy.utils import PrismaClient
-
pytestmark = pytest.mark.skipif(
sys.platform == "win32", reason="engine watcher is Unix-only"
)
@@ -72,9 +71,7 @@ def test_is_engine_alive_false_when_process_lookup_fails(
prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch
) -> None:
prisma_client._engine_pid = 99999
- monkeypatch.setattr(
- "os.kill", MagicMock(side_effect=ProcessLookupError())
- )
+ monkeypatch.setattr("os.kill", MagicMock(side_effect=ProcessLookupError()))
assert prisma_client._is_engine_alive() is False
@@ -102,15 +99,18 @@ def test_reap_all_zombies_returns_set_of_reaped_pids(
"contains_111": 111 in reaped,
"contains_222": 222 in reaped,
}
- assert pinned == {"type": "set", "size": 2, "contains_111": True, "contains_222": True}
+ assert pinned == {
+ "type": "set",
+ "size": 2,
+ "contains_111": True,
+ "contains_222": True,
+ }
def test_reap_all_zombies_handles_no_children_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
- monkeypatch.setattr(
- "os.waitpid", MagicMock(side_effect=ChildProcessError())
- )
+ monkeypatch.setattr("os.waitpid", MagicMock(side_effect=ChildProcessError()))
assert PrismaClient._reap_all_zombies() == set()
@@ -153,9 +153,7 @@ async def test_try_waitpid_watch_starts_thread_for_live_child(
async def test_try_waitpid_watch_returns_false_for_non_child(
prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch
) -> None:
- monkeypatch.setattr(
- "os.waitpid", MagicMock(side_effect=ChildProcessError())
- )
+ monkeypatch.setattr("os.waitpid", MagicMock(side_effect=ChildProcessError()))
assert prisma_client._try_waitpid_watch(123) is False
@@ -245,7 +243,9 @@ async def test_on_engine_death_from_thread_schedules_reconnect(
"confirmed_dead": prisma_client._engine_confirmed_dead,
"cleanup_called": prisma_client._cleanup_engine_watcher.call_count,
"reconnect_called": prisma_client.attempt_db_reconnect.await_count,
- "reconnect_reason": prisma_client.attempt_db_reconnect.await_args.kwargs["reason"],
+ "reconnect_reason": prisma_client.attempt_db_reconnect.await_args.kwargs[
+ "reason"
+ ],
}
assert pinned == {
"confirmed_dead": True,
@@ -462,7 +462,9 @@ async def test_start_engine_watcher_picks_waitpid_when_available(
prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(prisma_client, "_get_engine_pid", MagicMock(return_value=12345))
- monkeypatch.setattr(prisma_client, "_try_waitpid_watch", MagicMock(return_value=True))
+ monkeypatch.setattr(
+ prisma_client, "_try_waitpid_watch", MagicMock(return_value=True)
+ )
pidfd_called = MagicMock(return_value=False)
monkeypatch.setattr(prisma_client, "_try_pidfd_watch", pidfd_called)
await prisma_client._start_engine_watcher()
@@ -495,8 +497,12 @@ async def test_start_engine_watcher_falls_back_to_polling_when_no_kernel_apis(
prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(prisma_client, "_get_engine_pid", MagicMock(return_value=4242))
- monkeypatch.setattr(prisma_client, "_try_waitpid_watch", MagicMock(return_value=False))
- monkeypatch.setattr(prisma_client, "_try_pidfd_watch", MagicMock(return_value=False))
+ monkeypatch.setattr(
+ prisma_client, "_try_waitpid_watch", MagicMock(return_value=False)
+ )
+ monkeypatch.setattr(
+ prisma_client, "_try_pidfd_watch", MagicMock(return_value=False)
+ )
monkeypatch.setattr(prisma_client, "_poll_engine_proc", AsyncMock())
await prisma_client._start_engine_watcher()
await asyncio.sleep(0)
@@ -516,7 +522,9 @@ def test_stop_engine_watcher_clears_dead_flag(
def test_stop_engine_watcher_error_in_cleanup_propagates(
prisma_client: PrismaClient,
) -> None:
- prisma_client._cleanup_engine_watcher = MagicMock(side_effect=RuntimeError("cleanup boom"))
+ prisma_client._cleanup_engine_watcher = MagicMock(
+ side_effect=RuntimeError("cleanup boom")
+ )
with pytest.raises(RuntimeError, match="cleanup boom"):
prisma_client._stop_engine_watcher()
diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py
index 220fff1a881..3568be3eead 100644
--- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py
+++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py
@@ -48,7 +48,9 @@ async def test_health_check_returns_query_raw_result(
async def test_health_check_raises_when_query_raw_fails(
prisma_client: PrismaClient,
) -> None:
- prisma_client.db.query_raw = AsyncMock(side_effect=RuntimeError("connection refused"))
+ prisma_client.db.query_raw = AsyncMock(
+ side_effect=RuntimeError("connection refused")
+ )
with pytest.raises(RuntimeError, match="connection refused"):
await prisma_client.health_check()
@@ -115,7 +117,9 @@ async def test_set_spend_logs_row_count_error_raises_through_backoff(
await prisma_client._set_spend_logs_row_count_in_proxy_state()
-def test_validate_response_time_passes_finite_value(prisma_client: PrismaClient) -> None:
+def test_validate_response_time_passes_finite_value(
+ prisma_client: PrismaClient,
+) -> None:
inputs = {
"ok": prisma_client._validate_response_time(123.45),
"none": prisma_client._validate_response_time(None),
diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py
index 4e547b81acc..d1fcb1d3215 100644
--- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py
+++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py
@@ -21,10 +21,15 @@ from litellm.proxy.utils import PrismaClient
@pytest.mark.asyncio
-async def test_insert_data_hashes_token_and_upserts(prisma_client: PrismaClient) -> None:
+async def test_insert_data_hashes_token_and_upserts(
+ prisma_client: PrismaClient,
+) -> None:
token = "sk-secret-1"
- response = SimpleNamespace(token=hashlib.sha256(token.encode()).hexdigest(),
- key_alias="alias", user_id="u1")
+ response = SimpleNamespace(
+ token=hashlib.sha256(token.encode()).hexdigest(),
+ key_alias="alias",
+ user_id="u1",
+ )
prisma_client.db.litellm_verificationtoken.upsert = AsyncMock(return_value=response)
data = {
"token": token,
@@ -56,14 +61,18 @@ async def test_insert_data_hashes_token_and_upserts(prisma_client: PrismaClient)
@pytest.mark.asyncio
-async def test_insert_data_strips_null_budget_limits(prisma_client: PrismaClient) -> None:
+async def test_insert_data_strips_null_budget_limits(
+ prisma_client: PrismaClient,
+) -> None:
prisma_client.db.litellm_verificationtoken.upsert = AsyncMock(return_value=None)
await prisma_client.insert_data(
data={"token": "sk-1", "budget_limits": None}, table_name="key"
)
- create_payload = prisma_client.db.litellm_verificationtoken.upsert.await_args.kwargs[
- "data"
- ]["create"]
+ create_payload = (
+ prisma_client.db.litellm_verificationtoken.upsert.await_args.kwargs["data"][
+ "create"
+ ]
+ )
assert "budget_limits" not in create_payload
@@ -78,11 +87,13 @@ async def test_insert_data_team_serializes_members(prisma_client: PrismaClient)
"members_with_roles": [{"role": "admin", "user_id": "u1"}],
}
result = await prisma_client.insert_data(data=data, table_name="team")
- create_payload = prisma_client.db.litellm_teamtable.upsert.await_args.kwargs["data"][
- "create"
- ]
+ create_payload = prisma_client.db.litellm_teamtable.upsert.await_args.kwargs[
+ "data"
+ ]["create"]
assert result.team_id == "t1"
- assert create_payload["members_with_roles"] == json.dumps(data["members_with_roles"])
+ assert create_payload["members_with_roles"] == json.dumps(
+ data["members_with_roles"]
+ )
assert create_payload["team_id"] == "t1"
diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py
index 739e942de52..dfb740cca21 100644
--- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py
+++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py
@@ -82,9 +82,7 @@ async def test_send_email_error_missing_sender_email(
) -> None:
monkeypatch.delenv("SMTP_SENDER_EMAIL", raising=False)
with pytest.raises(ValueError, match="SMTP_SENDER_EMAIL"):
- await send_email(
- receiver_email="x@y", subject="s", html="h
"
- )
+ await send_email(receiver_email="x@y", subject="s", html="h
")
@pytest.mark.asyncio
@@ -113,7 +111,5 @@ async def test_send_email_smtp_failure_is_swallowed(
does not raise so a failing email never blocks the proxy.
"""
in_memory_smtp.raise_on_send = RuntimeError("smtp boom")
- await send_email(
- receiver_email="to@invalid", subject="Hi", html="x
"
- )
+ await send_email(receiver_email="to@invalid", subject="Hi", html="x
")
assert in_memory_smtp.sent == []
diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py
index a0b3af54750..251a2f2b0ab 100644
--- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py
+++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py
@@ -31,7 +31,9 @@ async def test_update_spend_invokes_writer_and_skips_empty_queue(
) -> None:
proxy_logging = MagicMock()
proxy_logging.db_spend_update_writer = MagicMock()
- proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock()
+ proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = (
+ AsyncMock()
+ )
mock_prisma_client.spend_log_transactions = []
await update_spend(
@@ -62,7 +64,9 @@ async def test_update_spend_processes_logs_when_queue_nonempty(
) -> None:
proxy_logging = MagicMock()
proxy_logging.db_spend_update_writer = MagicMock()
- proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock()
+ proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = (
+ AsyncMock()
+ )
mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")]
import litellm.proxy.utils as utils_mod
@@ -84,8 +88,8 @@ async def test_update_spend_handler_failure_propagates(
) -> None:
proxy_logging = MagicMock()
proxy_logging.db_spend_update_writer = MagicMock()
- proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock(
- side_effect=RuntimeError("handler down")
+ proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = (
+ AsyncMock(side_effect=RuntimeError("handler down"))
)
with pytest.raises(RuntimeError, match="handler down"):
await update_spend(
@@ -221,7 +225,11 @@ async def test_update_spend_logs_job_processes_and_clears_queue(
"queue_after": mock_prisma_client.spend_log_transactions,
"first_data_request_id": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[
"data"
- ][0]["request_id"],
+ ][
+ 0
+ ][
+ "request_id"
+ ],
"skip_duplicates_set": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[
"skip_duplicates"
],
@@ -243,8 +251,12 @@ async def test_monitor_spend_logs_queue_invokes_job_when_queue_nonempty(
import litellm.proxy.utils as utils_mod
import litellm.constants as constants_mod
- monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False)
- monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_SIZE_THRESHOLD", 1, raising=False)
+ monkeypatch.setattr(
+ constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False
+ )
+ monkeypatch.setattr(
+ constants_mod, "SPEND_LOG_QUEUE_SIZE_THRESHOLD", 1, raising=False
+ )
proxy_logging = MagicMock()
mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")]
@@ -277,7 +289,9 @@ async def test_monitor_spend_logs_queue_swallows_errors_and_backs_off(
import litellm.proxy.utils as utils_mod
import litellm.constants as constants_mod
- monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False)
+ monkeypatch.setattr(
+ constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False
+ )
sleep_count = {"n": 0}
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/_harness_smoke_test.py b/tests/test_litellm/proxy/utils/proxy_logging/_harness_smoke_test.py
index 1ec01f8c563..98086512300 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/_harness_smoke_test.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/_harness_smoke_test.py
@@ -10,7 +10,11 @@ import pytest
def test_normalize_replaces_volatile_keys(normalize_fn):
raw = {"id": 7, "name": "x", "nested": {"created_at": 1, "value": 2}}
- expected = {"id": "", "name": "x", "nested": {"created_at": "", "value": 2}}
+ expected = {
+ "id": "",
+ "name": "x",
+ "nested": {"created_at": "", "value": 2},
+ }
assert normalize_fn(raw) == expected
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_alerting.py b/tests/test_litellm/proxy/utils/proxy_logging/test_alerting.py
index cede859cb38..18284337dc2 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_alerting.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_alerting.py
@@ -16,7 +16,6 @@ from fastapi import HTTPException
import litellm
from litellm.proxy._types import AlertType, CallInfo
-
# ---------------------------------------------------------------------------
# failed_tracking_alert
# ---------------------------------------------------------------------------
@@ -39,7 +38,9 @@ async def test_failed_tracking_alert_forwards_to_slack(proxy_logging):
captured.update(kwargs)
proxy_logging.slack_alerting_instance = MagicMock(failed_tracking_alert=fake_alert)
- await proxy_logging.failed_tracking_alert(error_message="db down", failing_model="gpt-4")
+ await proxy_logging.failed_tracking_alert(
+ error_message="db down", failing_model="gpt-4"
+ )
snapshot = {
"error_message": captured["error_message"],
"failing_model": captured["failing_model"],
@@ -112,11 +113,17 @@ async def test_budget_alerts_slack_when_slack_alerting(proxy_logging):
"user_info_is_callinfo": isinstance(captured["user_info"], CallInfo),
"user_id": captured["user_info"].user_id,
}
- assert snapshot == {"type": "user_budget", "user_info_is_callinfo": True, "user_id": "u1"}
+ assert snapshot == {
+ "type": "user_budget",
+ "user_info_is_callinfo": True,
+ "user_id": "u1",
+ }
@pytest.mark.asyncio
-async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_global(proxy_logging):
+async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_global(
+ proxy_logging,
+):
proxy_logging.alerting = None
proxy_logging.slack_alerting_instance = MagicMock(budget_alerts=AsyncMock())
proxy_logging.email_logging_instance = MagicMock(budget_alerts=AsyncMock())
@@ -146,7 +153,9 @@ async def test_budget_alerts_slack_failure_raises(proxy_logging):
async def test_alerting_handler_no_op_when_alerting_is_none(proxy_logging):
proxy_logging.alerting = None
proxy_logging.slack_alerting_instance = MagicMock(send_alert=AsyncMock())
- await proxy_logging.alerting_handler(message="x", level="High", alert_type=AlertType.db_exceptions)
+ await proxy_logging.alerting_handler(
+ message="x", level="High", alert_type=AlertType.db_exceptions
+ )
proxy_logging.slack_alerting_instance.send_alert.assert_not_called()
@@ -160,7 +169,10 @@ async def test_alerting_handler_sends_to_slack(proxy_logging):
proxy_logging.slack_alerting_instance = MagicMock(send_alert=fake_send)
await proxy_logging.alerting_handler(
- message="hi", level="High", alert_type=AlertType.db_exceptions, request_data={"metadata": {}}
+ message="hi",
+ level="High",
+ alert_type=AlertType.db_exceptions,
+ request_data={"metadata": {}},
)
snapshot = {
"message": captured["message"],
@@ -177,11 +189,15 @@ async def test_alerting_handler_sends_to_slack(proxy_logging):
@pytest.mark.asyncio
-async def test_alerting_handler_sentry_without_sdk_error_raises(proxy_logging, monkeypatch):
+async def test_alerting_handler_sentry_without_sdk_error_raises(
+ proxy_logging, monkeypatch
+):
proxy_logging.alerting = ["sentry"]
monkeypatch.setattr(litellm.utils, "sentry_sdk_instance", None)
with pytest.raises(Exception, match="SENTRY_DSN"):
- await proxy_logging.alerting_handler(message="x", level="Low", alert_type=AlertType.db_exceptions)
+ await proxy_logging.alerting_handler(
+ message="x", level="Low", alert_type=AlertType.db_exceptions
+ )
# ---------------------------------------------------------------------------
@@ -190,29 +206,45 @@ async def test_alerting_handler_sentry_without_sdk_error_raises(proxy_logging, m
@pytest.mark.asyncio
-async def test_failure_handler_skips_when_db_exceptions_not_in_alert_types(proxy_logging):
+async def test_failure_handler_skips_when_db_exceptions_not_in_alert_types(
+ proxy_logging,
+):
proxy_logging.alert_types = ["llm_too_slow"] # type: ignore[list-item]
proxy_logging.alerting_handler = AsyncMock()
- proxy_logging.service_logging_obj = MagicMock(async_service_failure_hook=AsyncMock())
- await proxy_logging.failure_handler(original_exception=Exception("x"), duration=1.0, call_type="db_read")
+ proxy_logging.service_logging_obj = MagicMock(
+ async_service_failure_hook=AsyncMock()
+ )
+ await proxy_logging.failure_handler(
+ original_exception=Exception("x"), duration=1.0, call_type="db_read"
+ )
proxy_logging.alerting_handler.assert_not_called()
proxy_logging.service_logging_obj.async_service_failure_hook.assert_not_called()
@pytest.mark.asyncio
-async def test_failure_handler_logs_db_error_and_calls_service_logging(proxy_logging, monkeypatch):
+async def test_failure_handler_logs_db_error_and_calls_service_logging(
+ proxy_logging, monkeypatch
+):
proxy_logging.alert_types = [AlertType.db_exceptions]
proxy_logging.alerting_handler = AsyncMock()
- proxy_logging.service_logging_obj = MagicMock(async_service_failure_hook=AsyncMock())
+ proxy_logging.service_logging_obj = MagicMock(
+ async_service_failure_hook=AsyncMock()
+ )
monkeypatch.setattr(litellm.utils, "capture_exception", None)
await proxy_logging.failure_handler(
original_exception=HTTPException(status_code=500, detail="boom"),
duration=1.5,
call_type="db_write",
)
- call_kwargs = proxy_logging.service_logging_obj.async_service_failure_hook.call_args.kwargs
+ call_kwargs = (
+ proxy_logging.service_logging_obj.async_service_failure_hook.call_args.kwargs
+ )
snapshot = {
- "service": call_kwargs["service"].value if hasattr(call_kwargs["service"], "value") else call_kwargs["service"],
+ "service": (
+ call_kwargs["service"].value
+ if hasattr(call_kwargs["service"], "value")
+ else call_kwargs["service"]
+ ),
"duration": call_kwargs["duration"],
"call_type": call_kwargs["call_type"],
}
@@ -224,10 +256,14 @@ async def test_failure_handler_logs_db_error_and_calls_service_logging(proxy_log
@pytest.mark.asyncio
-async def test_failure_handler_with_capture_exception_invoked(proxy_logging, monkeypatch):
+async def test_failure_handler_with_capture_exception_invoked(
+ proxy_logging, monkeypatch
+):
proxy_logging.alert_types = [AlertType.db_exceptions]
proxy_logging.alerting_handler = AsyncMock()
- proxy_logging.service_logging_obj = MagicMock(async_service_failure_hook=AsyncMock())
+ proxy_logging.service_logging_obj = MagicMock(
+ async_service_failure_hook=AsyncMock()
+ )
captured: Dict[str, Any] = {}
def fake_capture(error):
@@ -235,7 +271,9 @@ async def test_failure_handler_with_capture_exception_invoked(proxy_logging, mon
monkeypatch.setattr(litellm.utils, "capture_exception", fake_capture)
err = RuntimeError("real")
- await proxy_logging.failure_handler(original_exception=err, duration=1.0, call_type="db_read")
+ await proxy_logging.failure_handler(
+ original_exception=err, duration=1.0, call_type="db_read"
+ )
snapshot = {
"captured_is_input": captured["error"] is err,
"service_failure_called": proxy_logging.service_logging_obj.async_service_failure_hook.called,
@@ -249,7 +287,9 @@ async def test_failure_handler_with_capture_exception_invoked(proxy_logging, mon
@pytest.mark.asyncio
-async def test_failure_handler_propagates_service_logging_error_raises(proxy_logging, monkeypatch):
+async def test_failure_handler_propagates_service_logging_error_raises(
+ proxy_logging, monkeypatch
+):
proxy_logging.alert_types = [AlertType.db_exceptions]
proxy_logging.alerting_handler = AsyncMock()
proxy_logging.service_logging_obj = MagicMock(
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py
index 45b81acbce1..4b2222f879e 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py
@@ -49,7 +49,9 @@ def _clear_caps_cache():
ProxyLogging._callback_capabilities_cache.clear()
-def test_callback_capabilities_with_no_callbacks_returns_defaults(mock_callbacks_disabled):
+def test_callback_capabilities_with_no_callbacks_returns_defaults(
+ mock_callbacks_disabled,
+):
caps = ProxyLogging._callback_capabilities()
snapshot = {
"headers": caps.has_post_call_response_headers,
@@ -128,17 +130,23 @@ def test_callback_capabilities_callback_resolution_error_raises(monkeypatch):
# ---------------------------------------------------------------------------
-def test_has_post_call_response_headers_callbacks_truth_table(monkeypatch, mock_callbacks_disabled):
+def test_has_post_call_response_headers_callbacks_truth_table(
+ monkeypatch, mock_callbacks_disabled
+):
"""One snapshot covering true + false + cache invalidation."""
snapshot = {
"empty_returns_false": ProxyLogging.has_post_call_response_headers_callbacks(),
}
monkeypatch.setattr(litellm, "callbacks", [_OverridesResponseHeaders()])
ProxyLogging._callback_capabilities_cache.clear()
- snapshot["override_returns_true"] = ProxyLogging.has_post_call_response_headers_callbacks()
+ snapshot["override_returns_true"] = (
+ ProxyLogging.has_post_call_response_headers_callbacks()
+ )
monkeypatch.setattr(litellm, "callbacks", [_PlainLogger()])
ProxyLogging._callback_capabilities_cache.clear()
- snapshot["plain_logger_false"] = ProxyLogging.has_post_call_response_headers_callbacks()
+ snapshot["plain_logger_false"] = (
+ ProxyLogging.has_post_call_response_headers_callbacks()
+ )
assert snapshot == {
"empty_returns_false": False,
"override_returns_true": True,
@@ -185,13 +193,17 @@ def test_has_streaming_callbacks_error_when_resolution_fails(monkeypatch):
ProxyLogging.has_streaming_callbacks()
-def test_has_streaming_chunk_hook_overrides_truth_table(monkeypatch, mock_callbacks_disabled):
+def test_has_streaming_chunk_hook_overrides_truth_table(
+ monkeypatch, mock_callbacks_disabled
+):
snapshot = {
"empty_false": ProxyLogging.has_streaming_chunk_hook_overrides(),
}
monkeypatch.setattr(litellm, "callbacks", [_OverridesPerChunk()])
ProxyLogging._callback_capabilities_cache.clear()
- snapshot["per_chunk_override_true"] = ProxyLogging.has_streaming_chunk_hook_overrides()
+ snapshot["per_chunk_override_true"] = (
+ ProxyLogging.has_streaming_chunk_hook_overrides()
+ )
monkeypatch.setattr(litellm, "callbacks", [_OverridesIterator()])
ProxyLogging._callback_capabilities_cache.clear()
snapshot["only_iterator_false"] = ProxyLogging.has_streaming_chunk_hook_overrides()
@@ -213,7 +225,9 @@ def test_has_streaming_chunk_hook_overrides_error_raises(monkeypatch):
ProxyLogging.has_streaming_chunk_hook_overrides()
-def test_needs_iterator_wrap_truth_table(proxy_logging, monkeypatch, mock_callbacks_disabled):
+def test_needs_iterator_wrap_truth_table(
+ proxy_logging, monkeypatch, mock_callbacks_disabled
+):
snapshot = {
"empty_false": proxy_logging.needs_iterator_wrap(),
}
@@ -241,7 +255,9 @@ def test_needs_iterator_wrap_error_raises(proxy_logging, monkeypatch):
proxy_logging.needs_iterator_wrap()
-def test_needs_per_chunk_streaming_hook_truth_table(proxy_logging, monkeypatch, mock_callbacks_disabled):
+def test_needs_per_chunk_streaming_hook_truth_table(
+ proxy_logging, monkeypatch, mock_callbacks_disabled
+):
snapshot = {
"empty_false": proxy_logging.needs_per_chunk_streaming_hook(),
}
@@ -250,7 +266,9 @@ def test_needs_per_chunk_streaming_hook_truth_table(proxy_logging, monkeypatch,
snapshot["per_chunk_override_true"] = proxy_logging.needs_per_chunk_streaming_hook()
monkeypatch.setattr(litellm, "callbacks", [_OverridesIterator()])
ProxyLogging._callback_capabilities_cache.clear()
- snapshot["only_iter_override_false"] = proxy_logging.needs_per_chunk_streaming_hook()
+ snapshot["only_iter_override_false"] = (
+ proxy_logging.needs_per_chunk_streaming_hook()
+ )
assert snapshot == {
"empty_false": False,
"per_chunk_override_true": True,
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py
index 3c5d879c2dc..27c60daea62 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py
@@ -32,7 +32,9 @@ def _make_guardrail(name="g1", should_run=True, response=None):
@pytest.mark.asyncio
-async def test_during_call_hook_no_guardrail_fast_path_returns_data(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled):
+async def test_during_call_hook_no_guardrail_fast_path_returns_data(
+ proxy_logging, make_user_api_key_auth, mock_callbacks_disabled
+):
data = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1}
out = await proxy_logging.during_call_hook(
data=data,
@@ -43,7 +45,9 @@ async def test_during_call_hook_no_guardrail_fast_path_returns_data(proxy_loggin
@pytest.mark.asyncio
-async def test_during_call_hook_runs_guardrails_in_parallel(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_during_call_hook_runs_guardrails_in_parallel(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
g1 = _make_guardrail("a")
g2 = _make_guardrail("b")
monkeypatch.setattr(litellm, "callbacks", [g1, g2])
@@ -62,7 +66,9 @@ async def test_during_call_hook_runs_guardrails_in_parallel(proxy_logging, make_
@pytest.mark.asyncio
-async def test_during_call_hook_guardrail_skipped_when_should_not_run(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_during_call_hook_guardrail_skipped_when_should_not_run(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
g = _make_guardrail("g", should_run=False)
monkeypatch.setattr(litellm, "callbacks", [g])
await proxy_logging.during_call_hook(
@@ -74,7 +80,9 @@ async def test_during_call_hook_guardrail_skipped_when_should_not_run(proxy_logg
@pytest.mark.asyncio
-async def test_during_call_hook_guardrail_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_during_call_hook_guardrail_error_raises(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
g = _make_guardrail("bad")
g.async_moderation_hook = AsyncMock(side_effect=RuntimeError("blocked"))
monkeypatch.setattr(litellm, "callbacks", [g])
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py
index 1ff9fbf8d83..d6a914bc88b 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py
@@ -42,15 +42,21 @@ def test_should_use_guardrail_load_balancing_truth_table(proxy_logging):
router = MagicMock()
router.guardrail_list = [{"guardrail_name": "g1"}, {"guardrail_name": "g1"}]
with patch("litellm.proxy.proxy_server.llm_router", router):
- snapshot["multiple_deployments"] = proxy_logging._should_use_guardrail_load_balancing("g1")
+ snapshot["multiple_deployments"] = (
+ proxy_logging._should_use_guardrail_load_balancing("g1")
+ )
router.guardrail_list = [{"guardrail_name": "g1"}]
with patch("litellm.proxy.proxy_server.llm_router", router):
- snapshot["single_deployment"] = proxy_logging._should_use_guardrail_load_balancing("g1")
+ snapshot["single_deployment"] = (
+ proxy_logging._should_use_guardrail_load_balancing("g1")
+ )
with patch("litellm.proxy.proxy_server.llm_router", None):
snapshot["no_router"] = proxy_logging._should_use_guardrail_load_balancing("g1")
router.guardrail_list = [{"guardrail_name": "other"}, {"guardrail_name": "other"}]
with patch("litellm.proxy.proxy_server.llm_router", router):
- snapshot["unmatched_name"] = proxy_logging._should_use_guardrail_load_balancing("g1")
+ snapshot["unmatched_name"] = proxy_logging._should_use_guardrail_load_balancing(
+ "g1"
+ )
assert snapshot == {
"multiple_deployments": True,
"single_deployment": False,
@@ -98,7 +104,9 @@ async def test_execute_guardrail_hook_pre_call(proxy_logging, make_user_api_key_
@pytest.mark.asyncio
-async def test_execute_guardrail_hook_during_call(proxy_logging, make_user_api_key_auth):
+async def test_execute_guardrail_hook_during_call(
+ proxy_logging, make_user_api_key_auth
+):
cb = _make_guardrail()
out = await proxy_logging._execute_guardrail_hook(
callback=cb,
@@ -125,7 +133,9 @@ async def test_execute_guardrail_hook_post_call(proxy_logging, make_user_api_key
@pytest.mark.asyncio
-async def test_execute_guardrail_hook_unknown_hook_type_raises(proxy_logging, make_user_api_key_auth):
+async def test_execute_guardrail_hook_unknown_hook_type_raises(
+ proxy_logging, make_user_api_key_auth
+):
cb = _make_guardrail()
with pytest.raises(ValueError, match="Unknown hook_type"):
await proxy_logging._execute_guardrail_hook(
@@ -237,7 +247,9 @@ async def test_process_guardrail_callback_enriches_and_reraises_http_exception(
cb = _make_guardrail()
cb.should_run_guardrail = MagicMock(return_value=True)
detail = {"error": "blocked"}
- cb.async_pre_call_hook = AsyncMock(side_effect=HTTPException(status_code=400, detail=detail))
+ cb.async_pre_call_hook = AsyncMock(
+ side_effect=HTTPException(status_code=400, detail=detail)
+ )
cb.event_hook = "pre_call"
proxy_logging._should_use_guardrail_load_balancing = MagicMock(return_value=False)
@@ -265,7 +277,9 @@ def test_process_guardrail_metadata_calls_header_helper(proxy_logging, monkeypat
from litellm.proxy.common_utils import callback_utils
- monkeypatch.setattr(callback_utils, "add_guardrail_to_applied_guardrails_header", fake_add)
+ monkeypatch.setattr(
+ callback_utils, "add_guardrail_to_applied_guardrails_header", fake_add
+ )
data = {"metadata": {"guardrails": ["g1", "g2"]}}
proxy_logging._process_guardrail_metadata(data)
snapshot = {
@@ -290,7 +304,9 @@ def test_process_guardrail_metadata_skips_already_applied(proxy_logging, monkeyp
from litellm.proxy.common_utils import callback_utils
- monkeypatch.setattr(callback_utils, "add_guardrail_to_applied_guardrails_header", fake_add)
+ monkeypatch.setattr(
+ callback_utils, "add_guardrail_to_applied_guardrails_header", fake_add
+ )
data = {"metadata": {"guardrails": ["g1", "g2"], "applied_guardrails": ["g1"]}}
proxy_logging._process_guardrail_metadata(data)
assert calls == ["g2"]
@@ -318,7 +334,9 @@ def test_process_guardrail_metadata_invalid_data_raises(proxy_logging):
@pytest.mark.asyncio
-async def test_maybe_execute_pipelines_no_pipelines_returns_data(proxy_logging, make_user_api_key_auth):
+async def test_maybe_execute_pipelines_no_pipelines_returns_data(
+ proxy_logging, make_user_api_key_auth
+):
data = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1}
out = await proxy_logging._maybe_execute_pipelines(
data=data,
@@ -330,13 +348,20 @@ async def test_maybe_execute_pipelines_no_pipelines_returns_data(proxy_logging,
@pytest.mark.asyncio
-async def test_maybe_execute_pipelines_skips_pipelines_with_other_mode(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_maybe_execute_pipelines_skips_pipelines_with_other_mode(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
pipeline = MagicMock()
pipeline.mode = "post_call" # not pre_call
- data = {"metadata": {"_guardrail_pipelines": [("p1", pipeline)]}, "model": "m", "messages": []}
+ data = {
+ "metadata": {"_guardrail_pipelines": [("p1", pipeline)]},
+ "model": "m",
+ "messages": [],
+ }
executed = MagicMock()
monkeypatch.setattr(
- "litellm.proxy.policy_engine.pipeline_executor.PipelineExecutor.execute_steps", executed
+ "litellm.proxy.policy_engine.pipeline_executor.PipelineExecutor.execute_steps",
+ executed,
)
out = await proxy_logging._maybe_execute_pipelines(
data=data,
@@ -358,7 +383,11 @@ async def test_maybe_execute_pipelines_blocks_on_block_terminal_action_raises(
fake_result = MagicMock()
fake_result.terminal_action = "block"
fake_result.step_results = []
- data = {"metadata": {"_guardrail_pipelines": [("policy-1", pipeline)]}, "messages": [], "model": "m"}
+ data = {
+ "metadata": {"_guardrail_pipelines": [("policy-1", pipeline)]},
+ "messages": [],
+ "model": "m",
+ }
async def fake_execute_steps(**kwargs):
return fake_result
@@ -386,7 +415,9 @@ def test_handle_pipeline_result_allow_with_modifications():
result = MagicMock()
result.terminal_action = "allow"
result.modified_data = {"b": 2, "c": 3}
- out = ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p")
+ out = ProxyLogging._handle_pipeline_result(
+ result=result, data=data, policy_name="p"
+ )
assert out == {"a": 1, "b": 2, "c": 3}
@@ -395,7 +426,9 @@ def test_handle_pipeline_result_block_raises_http_exception():
result.terminal_action = "block"
result.step_results = []
with pytest.raises(HTTPException) as info:
- ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p")
+ ProxyLogging._handle_pipeline_result(
+ result=result, data={"model": "m"}, policy_name="p"
+ )
detail = info.value.detail
snapshot = {
"is_dict": isinstance(detail, dict),
@@ -414,14 +447,19 @@ def test_handle_pipeline_result_modify_response_raises_modify_exception():
result.terminal_action = "modify_response"
result.modify_response_message = "filtered"
with pytest.raises(ModifyResponseException):
- ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p")
+ ProxyLogging._handle_pipeline_result(
+ result=result, data={"model": "m"}, policy_name="p"
+ )
def test_handle_pipeline_result_unknown_action_returns_data():
data = {"a": 1, "b": 2, "c": 3}
result = MagicMock()
result.terminal_action = "something_else"
- assert ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p") is data
+ assert (
+ ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p")
+ is data
+ )
# ---------------------------------------------------------------------------
@@ -461,16 +499,26 @@ async def test_run_guardrail_task_with_enrichment_enriches_http_exception_raises
@pytest.mark.asyncio
-async def test_process_prompt_template_no_op_when_no_prompt_spec(proxy_logging, monkeypatch):
+async def test_process_prompt_template_no_op_when_no_prompt_spec(
+ proxy_logging, monkeypatch
+):
from litellm.proxy.prompts import prompt_registry
monkeypatch.setattr(
- prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_callback_by_id", lambda *a, **kw: None
+ prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
+ "get_prompt_callback_by_id",
+ lambda *a, **kw: None,
)
monkeypatch.setattr(
- prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: None
+ prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
+ "get_prompt_by_id",
+ lambda *a, **kw: None,
)
- data: Dict[str, Any] = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1}
+ data: Dict[str, Any] = {
+ "messages": [{"role": "user"}],
+ "model": "m",
+ "temperature": 0.1,
+ }
await proxy_logging._process_prompt_template(
data=data,
litellm_logging_obj=MagicMock(),
@@ -482,7 +530,9 @@ async def test_process_prompt_template_no_op_when_no_prompt_spec(proxy_logging,
@pytest.mark.asyncio
-async def test_process_prompt_template_applies_when_spec_resolves(proxy_logging, monkeypatch):
+async def test_process_prompt_template_applies_when_spec_resolves(
+ proxy_logging, monkeypatch
+):
from litellm.proxy.prompts import prompt_registry
custom_logger = MagicMock()
@@ -495,7 +545,9 @@ async def test_process_prompt_template_applies_when_spec_resolves(proxy_logging,
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
- prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
+ prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
+ "get_prompt_by_id",
+ lambda *a, **kw: prompt_spec,
)
logging_obj = MagicMock()
@@ -533,7 +585,9 @@ async def test_process_prompt_template_applies_when_spec_resolves(proxy_logging,
@pytest.mark.asyncio
-async def test_process_prompt_template_async_get_prompt_error_raises(proxy_logging, monkeypatch):
+async def test_process_prompt_template_async_get_prompt_error_raises(
+ proxy_logging, monkeypatch
+):
from litellm.proxy.prompts import prompt_registry
custom_logger = MagicMock()
@@ -545,10 +599,14 @@ async def test_process_prompt_template_async_get_prompt_error_raises(proxy_loggi
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
- prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
+ prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
+ "get_prompt_by_id",
+ lambda *a, **kw: prompt_spec,
)
logging_obj = MagicMock()
- logging_obj.async_get_chat_completion_prompt = AsyncMock(side_effect=RuntimeError("bad prompt"))
+ logging_obj.async_get_chat_completion_prompt = AsyncMock(
+ side_effect=RuntimeError("bad prompt")
+ )
with pytest.raises(RuntimeError):
await proxy_logging._process_prompt_template(
data={"messages": [], "model": "m", "prompt_id": "x"},
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_internal_usage_cache.py b/tests/test_litellm/proxy/utils/proxy_logging/test_internal_usage_cache.py
index ff0afa45e36..48cbb6ad329 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_internal_usage_cache.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_internal_usage_cache.py
@@ -42,12 +42,21 @@ def test_internal_usage_cache_init_error_requires_dual_cache():
@pytest.mark.asyncio
async def test_async_get_cache_forwards_args():
inner = MagicMock()
- inner.async_get_cache = AsyncMock(return_value={"hit": True, "value": 42, "source": "redis"})
+ inner.async_get_cache = AsyncMock(
+ return_value={"hit": True, "value": 42, "source": "redis"}
+ )
cache = InternalUsageCache(dual_cache=inner)
- result = await cache.async_get_cache(key="k", litellm_parent_otel_span="span", local_only=True, extra="x")
+ result = await cache.async_get_cache(
+ key="k", litellm_parent_otel_span="span", local_only=True, extra="x"
+ )
forwarded = _kwargs_snapshot(inner.async_get_cache.call_args)
- assert forwarded == {"key": "k", "local_only": True, "parent_otel_span": "span", "extra": "x"}
+ assert forwarded == {
+ "key": "k",
+ "local_only": True,
+ "parent_otel_span": "span",
+ "extra": "x",
+ }
assert result == {"hit": True, "value": 42, "source": "redis"}
@@ -66,7 +75,9 @@ async def test_async_set_cache_forwards_args():
inner.async_set_cache = AsyncMock()
cache = InternalUsageCache(dual_cache=inner)
- await cache.async_set_cache(key="k", value="v", litellm_parent_otel_span="span", local_only=False, ttl=60)
+ await cache.async_set_cache(
+ key="k", value="v", litellm_parent_otel_span="span", local_only=False, ttl=60
+ )
forwarded = _kwargs_snapshot(inner.async_set_cache.call_args)
assert forwarded == {
"key": "k",
@@ -93,7 +104,9 @@ async def test_async_batch_set_cache_forwards_pipeline():
cache = InternalUsageCache(dual_cache=inner)
pairs = [("a", 1), ("b", 2)]
- await cache.async_batch_set_cache(cache_list=pairs, litellm_parent_otel_span=None, local_only=True, ttl=10)
+ await cache.async_batch_set_cache(
+ cache_list=pairs, litellm_parent_otel_span=None, local_only=True, ttl=10
+ )
forwarded = _kwargs_snapshot(inner.async_set_cache_pipeline.call_args)
assert forwarded == {
"cache_list": pairs,
@@ -117,9 +130,15 @@ async def test_async_batch_get_cache_forwards_args():
inner = MagicMock()
inner.async_batch_get_cache = AsyncMock(return_value=[1, 2, 3])
cache = InternalUsageCache(dual_cache=inner)
- result = await cache.async_batch_get_cache(keys=["a", "b", "c"], parent_otel_span="span", local_only=False)
+ result = await cache.async_batch_get_cache(
+ keys=["a", "b", "c"], parent_otel_span="span", local_only=False
+ )
forwarded = _kwargs_snapshot(inner.async_batch_get_cache.call_args)
- assert forwarded == {"keys": ["a", "b", "c"], "parent_otel_span": "span", "local_only": False}
+ assert forwarded == {
+ "keys": ["a", "b", "c"],
+ "parent_otel_span": "span",
+ "local_only": False,
+ }
assert result == [1, 2, 3]
@@ -137,9 +156,16 @@ async def test_async_increment_cache_forwards_args():
inner = MagicMock()
inner.async_increment_cache = AsyncMock(return_value=5.0)
cache = InternalUsageCache(dual_cache=inner)
- result = await cache.async_increment_cache(key="counter", value=1.5, litellm_parent_otel_span="span")
+ result = await cache.async_increment_cache(
+ key="counter", value=1.5, litellm_parent_otel_span="span"
+ )
forwarded = _kwargs_snapshot(inner.async_increment_cache.call_args)
- assert forwarded == {"key": "counter", "value": 1.5, "local_only": False, "parent_otel_span": "span"}
+ assert forwarded == {
+ "key": "counter",
+ "value": 1.5,
+ "local_only": False,
+ "parent_otel_span": "span",
+ }
assert result == 5.0
@@ -149,7 +175,9 @@ async def test_async_increment_cache_propagates_error_raises():
inner.async_increment_cache = AsyncMock(side_effect=OverflowError())
cache = InternalUsageCache(dual_cache=inner)
with pytest.raises(OverflowError):
- await cache.async_increment_cache(key="x", value=1.0, litellm_parent_otel_span=None)
+ await cache.async_increment_cache(
+ key="x", value=1.0, litellm_parent_otel_span=None
+ )
def test_set_cache_forwards_args():
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py
index e33da672599..04b01a04b21 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py
@@ -21,7 +21,6 @@ from litellm.proxy.utils import (
ProxyLogging,
)
-
# ---------------------------------------------------------------------------
# __init__
# ---------------------------------------------------------------------------
@@ -218,9 +217,14 @@ def test_get_proxy_hook_returns_registered_instance(proxy_logging):
"max_parallel_request_limiter": s_parallel,
}
snapshot = {
- "cache_control_check": proxy_logging.get_proxy_hook("cache_control_check") is s_cache,
- "max_budget_limiter": proxy_logging.get_proxy_hook("max_budget_limiter") is s_budget,
- "max_parallel_request_limiter": proxy_logging.get_proxy_hook("max_parallel_request_limiter") is s_parallel,
+ "cache_control_check": proxy_logging.get_proxy_hook("cache_control_check")
+ is s_cache,
+ "max_budget_limiter": proxy_logging.get_proxy_hook("max_budget_limiter")
+ is s_budget,
+ "max_parallel_request_limiter": proxy_logging.get_proxy_hook(
+ "max_parallel_request_limiter"
+ )
+ is s_parallel,
"unknown_returns_none": proxy_logging.get_proxy_hook("unknown") is None,
}
assert snapshot == {
@@ -249,7 +253,9 @@ def test_get_proxy_hook_non_string_key_raises(proxy_logging):
# ---------------------------------------------------------------------------
-def test_init_litellm_callbacks_replaces_string_with_instance(proxy_logging, monkeypatch):
+def test_init_litellm_callbacks_replaces_string_with_instance(
+ proxy_logging, monkeypatch
+):
from litellm.proxy import utils as utils_mod
sentinel_instance = MagicMock(spec=litellm.integrations.custom_logger.CustomLogger)
@@ -278,7 +284,9 @@ def test_init_litellm_callbacks_replaces_string_with_instance(proxy_logging, mon
}
-def test_init_litellm_callbacks_string_resolution_failure_keeps_string(proxy_logging, monkeypatch):
+def test_init_litellm_callbacks_string_resolution_failure_keeps_string(
+ proxy_logging, monkeypatch
+):
from litellm.proxy import utils as utils_mod
monkeypatch.setattr(litellm, "callbacks", ["unknown-logger"])
@@ -293,7 +301,9 @@ def test_init_litellm_callbacks_string_resolution_failure_keeps_string(proxy_log
assert litellm.callbacks[0] == "unknown-logger"
-def test_init_litellm_callbacks_propagates_resolver_error_raises(proxy_logging, monkeypatch):
+def test_init_litellm_callbacks_propagates_resolver_error_raises(
+ proxy_logging, monkeypatch
+):
from litellm.proxy import utils as utils_mod
monkeypatch.setattr(litellm, "callbacks", ["raises-on-init"])
@@ -322,7 +332,9 @@ async def test_update_request_status_when_alerting_set_writes_cache(proxy_loggin
captured.update(kwargs)
proxy_logging.internal_usage_cache.async_set_cache = fake_set_cache # type: ignore[assignment]
- await proxy_logging.update_request_status(litellm_call_id="call-1", status="success")
+ await proxy_logging.update_request_status(
+ litellm_call_id="call-1", status="success"
+ )
snapshot = {
"key": captured["key"],
"value": captured["value"],
@@ -341,14 +353,18 @@ async def test_update_request_status_when_alerting_set_writes_cache(proxy_loggin
async def test_update_request_status_no_alerting_skips_cache(proxy_logging):
proxy_logging.alerting = None
proxy_logging.internal_usage_cache.async_set_cache = AsyncMock()
- await proxy_logging.update_request_status(litellm_call_id="call-1", status="success")
+ await proxy_logging.update_request_status(
+ litellm_call_id="call-1", status="success"
+ )
proxy_logging.internal_usage_cache.async_set_cache.assert_not_called()
@pytest.mark.asyncio
async def test_update_request_status_cache_error_raises(proxy_logging):
proxy_logging.alerting = ["slack"]
- proxy_logging.internal_usage_cache.async_set_cache = AsyncMock(side_effect=ConnectionError("redis"))
+ proxy_logging.internal_usage_cache.async_set_cache = AsyncMock(
+ side_effect=ConnectionError("redis")
+ )
with pytest.raises(ConnectionError):
await proxy_logging.update_request_status(litellm_call_id="x", status="fail")
@@ -358,7 +374,9 @@ async def test_update_request_status_cache_error_raises(proxy_logging):
# ---------------------------------------------------------------------------
-def test_convert_user_api_key_auth_to_dict_pydantic_uses_model_dump(proxy_logging, make_user_api_key_auth):
+def test_convert_user_api_key_auth_to_dict_pydantic_uses_model_dump(
+ proxy_logging, make_user_api_key_auth
+):
auth = make_user_api_key_auth(user_id="u-1", team_id="t-1")
result = proxy_logging._convert_user_api_key_auth_to_dict(auth)
snapshot = {
@@ -385,7 +403,9 @@ def test_convert_user_api_key_auth_to_dict_none_returns_empty_dict(proxy_logging
assert proxy_logging._convert_user_api_key_auth_to_dict(None) == {}
-def test_convert_user_api_key_auth_to_dict_unconvertible_object_returns_empty(proxy_logging):
+def test_convert_user_api_key_auth_to_dict_unconvertible_object_returns_empty(
+ proxy_logging,
+):
class NoDict:
__slots__ = ()
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py b/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py
index 9defb309863..4fef0bed967 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py
@@ -21,13 +21,14 @@ from litellm.types.mcp import (
MCPPreCallResponseObject,
)
-
# ---------------------------------------------------------------------------
# _convert_mcp_to_llm_format
# ---------------------------------------------------------------------------
-def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mcp_request_obj):
+def test_convert_mcp_to_llm_format_returns_synthetic_data(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(tool_name="search", arguments={"q": "hello"})
out = proxy_logging._convert_mcp_to_llm_format(
request_obj=req,
@@ -86,7 +87,9 @@ def test_convert_mcp_to_llm_format_missing_request_obj_raises(proxy_logging):
# ---------------------------------------------------------------------------
-def test_convert_llm_result_to_mcp_response_exception_blocks(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_response_exception_blocks(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj()
result = proxy_logging._convert_llm_result_to_mcp_response(
llm_result=ValueError("boom"),
@@ -98,44 +101,72 @@ def test_convert_llm_result_to_mcp_response_exception_blocks(proxy_logging, make
"error_message": result.error_message,
"modified_arguments": result.modified_arguments,
}
- assert snapshot == {"should_proceed": False, "error_message": "boom", "modified_arguments": None}
+ assert snapshot == {
+ "should_proceed": False,
+ "error_message": "boom",
+ "modified_arguments": None,
+ }
-def test_convert_llm_result_to_mcp_response_blocked_content(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_response_blocked_content(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(tool_name="t", arguments={"a": 1})
llm_result = {"messages": [{"content": "this is blocked"}]}
- result = proxy_logging._convert_llm_result_to_mcp_response(llm_result=llm_result, request_obj=req)
+ result = proxy_logging._convert_llm_result_to_mcp_response(
+ llm_result=llm_result, request_obj=req
+ )
assert isinstance(result, MCPPreCallResponseObject)
assert result.should_proceed is False
assert "blocked" in (result.error_message or "").lower()
-def test_convert_llm_result_to_mcp_response_modified_content_redacted(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_response_modified_content_redacted(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(tool_name="search", arguments={"q": "ssn 123"})
- llm_result = {"messages": [{"content": "Tool: search\nArguments: {\"q\": \"[REDACTED]\"}"}]}
- result = proxy_logging._convert_llm_result_to_mcp_response(llm_result=llm_result, request_obj=req)
+ llm_result = {
+ "messages": [{"content": 'Tool: search\nArguments: {"q": "[REDACTED]"}'}]
+ }
+ result = proxy_logging._convert_llm_result_to_mcp_response(
+ llm_result=llm_result, request_obj=req
+ )
assert isinstance(result, MCPPreCallResponseObject)
snapshot = {
"should_proceed": result.should_proceed,
"modified_q": (result.modified_arguments or {}).get("q"),
"error": result.error_message,
}
- assert snapshot == {"should_proceed": True, "modified_q": "[REDACTED]", "error": None}
+ assert snapshot == {
+ "should_proceed": True,
+ "modified_q": "[REDACTED]",
+ "error": None,
+ }
-def test_convert_llm_result_to_mcp_response_string_blocks(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_response_string_blocks(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj()
- result = proxy_logging._convert_llm_result_to_mcp_response(llm_result="bad input", request_obj=req)
+ result = proxy_logging._convert_llm_result_to_mcp_response(
+ llm_result="bad input", request_obj=req
+ )
assert isinstance(result, MCPPreCallResponseObject)
snapshot = {
"should_proceed": result.should_proceed,
"error_message": result.error_message,
"modified_arguments": result.modified_arguments,
}
- assert snapshot == {"should_proceed": False, "error_message": "bad input", "modified_arguments": None}
+ assert snapshot == {
+ "should_proceed": False,
+ "error_message": "bad input",
+ "modified_arguments": None,
+ }
-def test_convert_llm_result_to_mcp_response_unmodified_returns_none(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_response_unmodified_returns_none(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(tool_name="x", arguments={"a": 1})
same_content = "Tool: x\nArguments: {'a': 1}"
result = proxy_logging._convert_llm_result_to_mcp_response(
@@ -147,7 +178,9 @@ def test_convert_llm_result_to_mcp_response_unmodified_returns_none(proxy_loggin
def test_convert_llm_result_to_mcp_response_no_request_obj_raises(proxy_logging):
with pytest.raises(AttributeError):
- proxy_logging._convert_llm_result_to_mcp_response(llm_result={"messages": [{"content": "x"}]}, request_obj=None)
+ proxy_logging._convert_llm_result_to_mcp_response(
+ llm_result={"messages": [{"content": "x"}]}, request_obj=None
+ )
# ---------------------------------------------------------------------------
@@ -155,16 +188,20 @@ def test_convert_llm_result_to_mcp_response_no_request_obj_raises(proxy_logging)
# ---------------------------------------------------------------------------
-def test_extract_modified_arguments_from_content_parses_json(proxy_logging, make_mcp_request_obj):
+def test_extract_modified_arguments_from_content_parses_json(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj()
out = proxy_logging._extract_modified_arguments_from_content(
- masked_content="Tool: x\nArguments: {\"a\": 1, \"b\": 2, \"c\": 3}",
+ masked_content='Tool: x\nArguments: {"a": 1, "b": 2, "c": 3}',
request_obj=req,
)
assert out == {"a": 1, "b": 2, "c": 3}
-def test_extract_modified_arguments_from_content_no_arguments_line_returns_none(proxy_logging, make_mcp_request_obj):
+def test_extract_modified_arguments_from_content_no_arguments_line_returns_none(
+ proxy_logging, make_mcp_request_obj
+):
out = proxy_logging._extract_modified_arguments_from_content(
masked_content="random content with no arguments",
request_obj=make_mcp_request_obj(),
@@ -172,7 +209,9 @@ def test_extract_modified_arguments_from_content_no_arguments_line_returns_none(
assert out is None
-def test_extract_modified_arguments_from_content_empty_string_returns_none(proxy_logging, make_mcp_request_obj):
+def test_extract_modified_arguments_from_content_empty_string_returns_none(
+ proxy_logging, make_mcp_request_obj
+):
out = proxy_logging._extract_modified_arguments_from_content(
masked_content="",
request_obj=make_mcp_request_obj(),
@@ -180,7 +219,9 @@ def test_extract_modified_arguments_from_content_empty_string_returns_none(proxy
assert out is None
-def test_extract_modified_arguments_from_content_invalid_json_falls_back(proxy_logging, make_mcp_request_obj):
+def test_extract_modified_arguments_from_content_invalid_json_falls_back(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(arguments={"name": "alice"})
out = proxy_logging._extract_modified_arguments_from_content(
masked_content="Tool: x\nArguments: {name: REDACTED}",
@@ -190,9 +231,13 @@ def test_extract_modified_arguments_from_content_invalid_json_falls_back(proxy_l
assert "name" in out
-def test_extract_modified_arguments_from_content_error_swallowed_returns_none(proxy_logging):
+def test_extract_modified_arguments_from_content_error_swallowed_returns_none(
+ proxy_logging,
+):
"""Internal try/except swallows any unexpected error and returns None."""
- out = proxy_logging._extract_modified_arguments_from_content(masked_content=None, request_obj=None)
+ out = proxy_logging._extract_modified_arguments_from_content(
+ masked_content=None, request_obj=None
+ )
assert out is None
@@ -207,13 +252,23 @@ def test_parse_arguments_manually_applies_overrides(proxy_logging):
args_text='"name": "[REDACTED]", "ssn": "[REDACTED]"',
original_args=original,
)
- snapshot = {"name": out["name"], "ssn": out["ssn"], "original_unchanged": original["name"]}
- assert snapshot == {"name": "[REDACTED]", "ssn": "[REDACTED]", "original_unchanged": "alice"}
+ snapshot = {
+ "name": out["name"],
+ "ssn": out["ssn"],
+ "original_unchanged": original["name"],
+ }
+ assert snapshot == {
+ "name": "[REDACTED]",
+ "ssn": "[REDACTED]",
+ "original_unchanged": "alice",
+ }
def test_parse_arguments_manually_returns_original_if_no_match(proxy_logging):
original = {"foo": "bar"}
- out = proxy_logging._parse_arguments_manually(args_text="nothing here", original_args=original)
+ out = proxy_logging._parse_arguments_manually(
+ args_text="nothing here", original_args=original
+ )
assert out == {"foo": "bar"}
@@ -227,7 +282,9 @@ def test_parse_arguments_manually_error_swallowed_returns_none(proxy_logging):
# ---------------------------------------------------------------------------
-def test_convert_llm_result_to_mcp_during_response_exception(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_during_response_exception(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj()
result = proxy_logging._convert_llm_result_to_mcp_during_response(
llm_result=ValueError("during boom"), request_obj=req
@@ -245,7 +302,9 @@ def test_convert_llm_result_to_mcp_during_response_exception(proxy_logging, make
}
-def test_convert_llm_result_to_mcp_during_response_blocked_content(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_during_response_blocked_content(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(tool_name="t", arguments={"a": 1})
result = proxy_logging._convert_llm_result_to_mcp_during_response(
llm_result={"messages": [{"content": "blocked content"}]},
@@ -256,10 +315,14 @@ def test_convert_llm_result_to_mcp_during_response_blocked_content(proxy_logging
assert "blocked" in (result.error_message or "").lower()
-def test_convert_llm_result_to_mcp_during_response_modified_stops(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_during_response_modified_stops(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(tool_name="t", arguments={"a": 1})
result = proxy_logging._convert_llm_result_to_mcp_during_response(
- llm_result={"messages": [{"content": "Tool: t\nArguments: {\"a\": \"[REDACTED]\"}"}]},
+ llm_result={
+ "messages": [{"content": 'Tool: t\nArguments: {"a": "[REDACTED]"}'}]
+ },
request_obj=req,
)
assert isinstance(result, MCPDuringCallResponseObject)
@@ -267,17 +330,24 @@ def test_convert_llm_result_to_mcp_during_response_modified_stops(proxy_logging,
assert "modified" in (result.error_message or "").lower()
-def test_convert_llm_result_to_mcp_during_response_string_blocks(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_during_response_string_blocks(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj()
result = proxy_logging._convert_llm_result_to_mcp_during_response(
llm_result="kill switch", request_obj=req
)
assert isinstance(result, MCPDuringCallResponseObject)
- snapshot = {"should_continue": result.should_continue, "error_message": result.error_message}
+ snapshot = {
+ "should_continue": result.should_continue,
+ "error_message": result.error_message,
+ }
assert snapshot == {"should_continue": False, "error_message": "kill switch"}
-def test_convert_llm_result_to_mcp_during_response_unmodified_returns_none(proxy_logging, make_mcp_request_obj):
+def test_convert_llm_result_to_mcp_during_response_unmodified_returns_none(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(tool_name="t", arguments={"a": 1})
same = "Tool: t\nArguments: {'a': 1}"
assert (
@@ -301,14 +371,18 @@ def test_convert_llm_result_to_mcp_during_response_no_request_obj_raises(proxy_l
# ---------------------------------------------------------------------------
-def test_parse_pre_mcp_call_hook_response_with_modified_args(proxy_logging, make_mcp_request_obj):
+def test_parse_pre_mcp_call_hook_response_with_modified_args(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(arguments={"a": 1})
resp = MCPPreCallResponseObject(
should_proceed=True,
modified_arguments={"a": "x", "b": "y"},
error_message=None,
)
- out = proxy_logging._parse_pre_mcp_call_hook_response(response=resp, original_request=req)
+ out = proxy_logging._parse_pre_mcp_call_hook_response(
+ response=resp, original_request=req
+ )
snapshot = {
"should_proceed": out["should_proceed"],
"modified_arguments": out["modified_arguments"],
@@ -323,16 +397,22 @@ def test_parse_pre_mcp_call_hook_response_with_modified_args(proxy_logging, make
}
-def test_parse_pre_mcp_call_hook_response_no_modifications_uses_original(proxy_logging, make_mcp_request_obj):
+def test_parse_pre_mcp_call_hook_response_no_modifications_uses_original(
+ proxy_logging, make_mcp_request_obj
+):
req = make_mcp_request_obj(arguments={"original": True})
resp = MCPPreCallResponseObject(
should_proceed=True, modified_arguments=None, error_message=None
)
- out = proxy_logging._parse_pre_mcp_call_hook_response(response=resp, original_request=req)
+ out = proxy_logging._parse_pre_mcp_call_hook_response(
+ response=resp, original_request=req
+ )
assert out["modified_arguments"] == {"original": True}
-def test_parse_pre_mcp_call_hook_response_invalid_response_raises(proxy_logging, make_mcp_request_obj):
+def test_parse_pre_mcp_call_hook_response_invalid_response_raises(
+ proxy_logging, make_mcp_request_obj
+):
with pytest.raises(AttributeError):
proxy_logging._parse_pre_mcp_call_hook_response(
response=None, original_request=make_mcp_request_obj()
@@ -344,7 +424,9 @@ def test_parse_pre_mcp_call_hook_response_invalid_response_raises(proxy_logging,
# ---------------------------------------------------------------------------
-def test_create_mcp_request_object_from_kwargs_full(proxy_logging, make_user_api_key_auth):
+def test_create_mcp_request_object_from_kwargs_full(
+ proxy_logging, make_user_api_key_auth
+):
auth = make_user_api_key_auth(user_id="u-1")
obj = proxy_logging._create_mcp_request_object_from_kwargs(
kwargs={
@@ -361,7 +443,12 @@ def test_create_mcp_request_object_from_kwargs_full(proxy_logging, make_user_api
"server_name": obj.server_name,
"auth_user_id": obj.user_api_key_auth.get("user_id"),
}
- assert snapshot == {"tool_name": "calc", "arguments": {"x": 1}, "server_name": "math", "auth_user_id": "u-1"}
+ assert snapshot == {
+ "tool_name": "calc",
+ "arguments": {"x": 1},
+ "server_name": "math",
+ "auth_user_id": "u-1",
+ }
def test_create_mcp_request_object_from_kwargs_empty(proxy_logging):
@@ -413,9 +500,13 @@ def test_convert_mcp_hook_response_to_kwargs_merges_headers(proxy_logging):
assert out["extra_headers"] == {"keep": "yes", "overwrite": "new", "added": "1"}
-def test_convert_mcp_hook_response_to_kwargs_no_response_data_returns_original(proxy_logging):
+def test_convert_mcp_hook_response_to_kwargs_no_response_data_returns_original(
+ proxy_logging,
+):
original = {"a": 1}
- out = proxy_logging._convert_mcp_hook_response_to_kwargs(response_data=None, original_kwargs=original)
+ out = proxy_logging._convert_mcp_hook_response_to_kwargs(
+ response_data=None, original_kwargs=original
+ )
assert out is original
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_module_helpers.py b/tests/test_litellm/proxy/utils/proxy_logging/test_module_helpers.py
index c491f16f2e4..888db70440e 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_module_helpers.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_module_helpers.py
@@ -25,7 +25,6 @@ from litellm.proxy.utils import (
print_verbose,
)
-
# ---------------------------------------------------------------------------
# print_verbose
# ---------------------------------------------------------------------------
@@ -40,7 +39,11 @@ def test_print_verbose_when_set_verbose_true_prints_redacted(monkeypatch, capsys
"out_has_payload": "hello world" in captured.out,
"no_stderr": captured.err == "",
}
- assert snapshot == {"out_has_prefix": True, "out_has_payload": True, "no_stderr": True}
+ assert snapshot == {
+ "out_has_prefix": True,
+ "out_has_payload": True,
+ "no_stderr": True,
+ }
def test_print_verbose_when_set_verbose_false_no_stdout(monkeypatch, capsys):
@@ -193,13 +196,21 @@ def test_enrich_http_exception_adds_guardrail_name_and_mode():
def test_enrich_http_exception_does_not_overwrite_existing_keys():
- detail = {"error": "blocked", "guardrail_name": "explicit", "guardrail_mode": "during_call"}
+ detail = {
+ "error": "blocked",
+ "guardrail_name": "explicit",
+ "guardrail_mode": "during_call",
+ }
exc = HTTPException(status_code=400, detail=detail)
cb = MagicMock()
cb.guardrail_name = "should-not-overwrite"
cb.event_hook = "should-not-overwrite"
_enrich_http_exception_with_guardrail_context(exc, cb)
- assert detail == {"error": "blocked", "guardrail_name": "explicit", "guardrail_mode": "during_call"}
+ assert detail == {
+ "error": "blocked",
+ "guardrail_name": "explicit",
+ "guardrail_mode": "during_call",
+ }
def test_enrich_http_exception_no_op_for_non_http_exception():
@@ -245,7 +256,11 @@ def test_on_backoff_invokes_print_verbose(monkeypatch):
captured = []
monkeypatch.setattr(utils_mod, "print_verbose", lambda s: captured.append(s))
on_backoff({"tries": 3})
- snapshot = {"len": len(captured), "first_has_attempt": "attempt" in captured[0], "first_has_3": "3" in captured[0]}
+ snapshot = {
+ "len": len(captured),
+ "first_has_attempt": "attempt" in captured[0],
+ "first_has_3": "3" in captured[0],
+ }
assert snapshot == {"len": 1, "first_has_attempt": True, "first_has_3": True}
@@ -300,7 +315,9 @@ async def test_lookup_deprecated_key_returns_active_token_id_and_caches(monkeypa
deprecated_row.revoke_at = future
db = MagicMock()
- db.litellm_deprecatedverificationtoken.find_first = AsyncMock(return_value=deprecated_row)
+ db.litellm_deprecatedverificationtoken.find_first = AsyncMock(
+ return_value=deprecated_row
+ )
result = await _lookup_deprecated_key(db=db, hashed_token="hash-abc")
cached_value = fresh.get("hash-abc")
@@ -320,7 +337,9 @@ async def test_lookup_deprecated_key_returns_active_token_id_and_caches(monkeypa
async def test_lookup_deprecated_key_returns_none_when_not_found(monkeypatch):
from litellm.caching.dual_cache import LimitedSizeOrderedDict
- monkeypatch.setattr(utils_mod, "_deprecated_key_cache", LimitedSizeOrderedDict(max_size=10))
+ monkeypatch.setattr(
+ utils_mod, "_deprecated_key_cache", LimitedSizeOrderedDict(max_size=10)
+ )
db = MagicMock()
db.litellm_deprecatedverificationtoken.find_first = AsyncMock(return_value=None)
assert await _lookup_deprecated_key(db=db, hashed_token="missing") is None
@@ -330,9 +349,13 @@ async def test_lookup_deprecated_key_returns_none_when_not_found(monkeypatch):
async def test_lookup_deprecated_key_db_error_returns_none(monkeypatch):
from litellm.caching.dual_cache import LimitedSizeOrderedDict
- monkeypatch.setattr(utils_mod, "_deprecated_key_cache", LimitedSizeOrderedDict(max_size=10))
+ monkeypatch.setattr(
+ utils_mod, "_deprecated_key_cache", LimitedSizeOrderedDict(max_size=10)
+ )
db = MagicMock()
- db.litellm_deprecatedverificationtoken.find_first = AsyncMock(side_effect=RuntimeError("db down"))
+ db.litellm_deprecatedverificationtoken.find_first = AsyncMock(
+ side_effect=RuntimeError("db down")
+ )
result = await _lookup_deprecated_key(db=db, hashed_token="x")
assert result is None
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_failure_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_failure_hook.py
index a2a57931d26..86704b65355 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_failure_hook.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_failure_hook.py
@@ -84,7 +84,8 @@ async def test_post_call_failure_hook_no_callbacks_returns_none(
"out_is_none": out is None,
"litellm_logging_obj_popped": "litellm_logging_obj" not in request_data,
"call_id_preserved": request_data["litellm_call_id"] == "abc",
- "first_api_call_start_time_present": "first_api_call_start_time" in request_data,
+ "first_api_call_start_time_present": "first_api_call_start_time"
+ in request_data,
}
assert snapshot == {
"out_is_none": True,
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py
index 6a339b37a80..20916bfd89d 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py
@@ -32,7 +32,9 @@ def _make_guardrail(name="g", should_run=True, override=None):
@pytest.mark.asyncio
-async def test_post_call_success_hook_returns_response_when_no_callbacks(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled):
+async def test_post_call_success_hook_returns_response_when_no_callbacks(
+ proxy_logging, make_user_api_key_auth, mock_callbacks_disabled
+):
response = {"original": True, "model": "m", "choices": []}
out = await proxy_logging.post_call_success_hook(
data={}, response=response, user_api_key_dict=make_user_api_key_auth()
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py
index 05005dae797..9e809e80661 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py
@@ -37,7 +37,9 @@ async def test_process_pre_call_hook_response_dict_returns_response(proxy_loggin
@pytest.mark.asyncio
-async def test_process_pre_call_hook_response_string_completion_raises_rejected(proxy_logging):
+async def test_process_pre_call_hook_response_string_completion_raises_rejected(
+ proxy_logging,
+):
with pytest.raises(RejectedRequestError):
await proxy_logging.process_pre_call_hook_response(
response="rejected",
@@ -47,7 +49,9 @@ async def test_process_pre_call_hook_response_string_completion_raises_rejected(
@pytest.mark.asyncio
-async def test_process_pre_call_hook_response_string_other_call_type_raises_http(proxy_logging):
+async def test_process_pre_call_hook_response_string_other_call_type_raises_http(
+ proxy_logging,
+):
with pytest.raises(HTTPException) as info:
await proxy_logging.process_pre_call_hook_response(
response="bad",
@@ -80,8 +84,14 @@ async def test_process_pre_call_hook_response_other_type_returns_data(proxy_logg
@pytest.mark.asyncio
-async def test_pre_call_hook_returns_data_when_no_callbacks(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled):
- data = {"messages": [{"role": "user", "content": "hi"}], "model": "m", "temperature": 0.7}
+async def test_pre_call_hook_returns_data_when_no_callbacks(
+ proxy_logging, make_user_api_key_auth, mock_callbacks_disabled
+):
+ data = {
+ "messages": [{"role": "user", "content": "hi"}],
+ "model": "m",
+ "temperature": 0.7,
+ }
proxy_logging.slack_alerting_instance = MagicMock(alerting=None)
out = await proxy_logging.pre_call_hook(
user_api_key_dict=make_user_api_key_auth(),
@@ -92,7 +102,9 @@ async def test_pre_call_hook_returns_data_when_no_callbacks(proxy_logging, make_
@pytest.mark.asyncio
-async def test_pre_call_hook_returns_none_for_none_data(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled):
+async def test_pre_call_hook_returns_none_for_none_data(
+ proxy_logging, make_user_api_key_auth, mock_callbacks_disabled
+):
proxy_logging.slack_alerting_instance = MagicMock(alerting=None)
out = await proxy_logging.pre_call_hook(
user_api_key_dict=make_user_api_key_auth(),
@@ -103,7 +115,9 @@ async def test_pre_call_hook_returns_none_for_none_data(proxy_logging, make_user
@pytest.mark.asyncio
-async def test_pre_call_hook_invokes_pre_call_override(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_pre_call_hook_invokes_pre_call_override(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
captured: Dict[str, Any] = {}
class _Cb(CustomLogger):
@@ -133,7 +147,9 @@ async def test_pre_call_hook_invokes_pre_call_override(proxy_logging, make_user_
@pytest.mark.asyncio
-async def test_pre_call_hook_propagates_callback_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_pre_call_hook_propagates_callback_error_raises(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
class _BadCb(CustomLogger):
async def async_pre_call_hook(self, **kwargs): # type: ignore[override]
raise RuntimeError("rejected")
@@ -149,9 +165,15 @@ async def test_pre_call_hook_propagates_callback_error_raises(proxy_logging, mak
@pytest.mark.asyncio
-async def test_pre_call_hook_processes_guardrail_metadata_when_no_overrides(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled):
+async def test_pre_call_hook_processes_guardrail_metadata_when_no_overrides(
+ proxy_logging, make_user_api_key_auth, mock_callbacks_disabled
+):
"""Even when no callback overrides exist, ``_process_guardrail_metadata`` runs."""
- data = {"messages": [{"role": "user"}], "model": "m", "metadata": {"guardrails": ["g1"]}}
+ data = {
+ "messages": [{"role": "user"}],
+ "model": "m",
+ "metadata": {"guardrails": ["g1"]},
+ }
proxy_logging.slack_alerting_instance = MagicMock(alerting=None)
invoked = {}
diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py
index 65d3c3c8079..46e57f07698 100644
--- a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py
+++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py
@@ -140,7 +140,9 @@ async def test_init_response_taking_too_long_task_runs_when_alerting(proxy_loggi
@pytest.mark.asyncio
-async def test_init_response_taking_too_long_task_no_op_when_alerting_off(proxy_logging):
+async def test_init_response_taking_too_long_task_no_op_when_alerting_off(
+ proxy_logging,
+):
proxy_logging.slack_alerting_instance = MagicMock()
proxy_logging.slack_alerting_instance.alerting = None
proxy_logging.slack_alerting_instance.response_taking_too_long = AsyncMock()
@@ -149,7 +151,9 @@ async def test_init_response_taking_too_long_task_no_op_when_alerting_off(proxy_
proxy_logging.slack_alerting_instance.response_taking_too_long.assert_not_called()
-def test_init_response_taking_too_long_task_no_slack_instance_no_error_raises(proxy_logging):
+def test_init_response_taking_too_long_task_no_slack_instance_no_error_raises(
+ proxy_logging,
+):
proxy_logging.slack_alerting_instance = None
proxy_logging._init_response_taking_too_long_task(data=None)
@@ -160,13 +164,17 @@ def test_init_response_taking_too_long_task_no_slack_instance_no_error_raises(pr
@pytest.mark.asyncio
-async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(proxy_logging):
+async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(
+ proxy_logging,
+):
async def gen():
for ch in ("a", "b", "c"):
yield ch
cb = MagicMock(guardrail_name="g", event_hook="pre_call")
- wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=gen())
+ wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(
+ callback=cb, gen=gen()
+ )
out = [ch async for ch in wrapped]
snapshot = {
"chunks": out,
@@ -183,7 +191,9 @@ async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(pro
@pytest.mark.asyncio
-async def test_wrap_streaming_iterator_with_enrichment_enriches_http_exception_raises(proxy_logging):
+async def test_wrap_streaming_iterator_with_enrichment_enriches_http_exception_raises(
+ proxy_logging,
+):
detail = {"error": "blocked"}
async def boom_gen():
@@ -192,7 +202,9 @@ async def test_wrap_streaming_iterator_with_enrichment_enriches_http_exception_r
raise HTTPException(status_code=400, detail=detail)
cb = MagicMock(guardrail_name="presidio", event_hook="post_call")
- wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=boom_gen())
+ wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(
+ callback=cb, gen=boom_gen()
+ )
with pytest.raises(HTTPException):
async for _ in wrapped:
pass
@@ -206,7 +218,9 @@ async def test_wrap_streaming_iterator_with_enrichment_enriches_http_exception_r
@pytest.mark.asyncio
-async def test_async_post_call_streaming_hook_fast_path_returns_response(proxy_logging, mock_callbacks_disabled, make_user_api_key_auth):
+async def test_async_post_call_streaming_hook_fast_path_returns_response(
+ proxy_logging, mock_callbacks_disabled, make_user_api_key_auth
+):
resp = "chunk-1"
out = await proxy_logging.async_post_call_streaming_hook(
data={}, response=resp, user_api_key_dict=make_user_api_key_auth()
@@ -226,7 +240,9 @@ async def test_async_post_call_streaming_hook_fast_path_returns_response(proxy_l
@pytest.mark.asyncio
-async def test_async_post_call_streaming_hook_invokes_per_chunk_callback(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_async_post_call_streaming_hook_invokes_per_chunk_callback(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
class _Per(CustomLogger):
async def async_post_call_streaming_hook(self, **kwargs): # type: ignore[override]
return "modified-" + str(kwargs.get("response", ""))
@@ -238,7 +254,13 @@ async def test_async_post_call_streaming_hook_invokes_per_chunk_callback(proxy_l
fake_resp = ModelResponse(
id="rid",
- choices=[{"index": 0, "delta": {"role": "assistant", "content": "hi"}, "finish_reason": None}],
+ choices=[
+ {
+ "index": 0,
+ "delta": {"role": "assistant", "content": "hi"},
+ "finish_reason": None,
+ }
+ ],
created=0,
model="gpt-4o-mini",
object="chat.completion.chunk",
@@ -253,7 +275,9 @@ async def test_async_post_call_streaming_hook_invokes_per_chunk_callback(proxy_l
@pytest.mark.asyncio
-async def test_async_post_call_streaming_hook_callback_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_async_post_call_streaming_hook_callback_error_raises(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
class _Per(CustomLogger):
async def async_post_call_streaming_hook(self, **kwargs): # type: ignore[override]
raise RuntimeError("hook-fail")
@@ -264,7 +288,13 @@ async def test_async_post_call_streaming_hook_callback_error_raises(proxy_loggin
fake_resp = ModelResponse(
id="rid",
- choices=[{"index": 0, "delta": {"role": "assistant", "content": "hi"}, "finish_reason": None}],
+ choices=[
+ {
+ "index": 0,
+ "delta": {"role": "assistant", "content": "hi"},
+ "finish_reason": None,
+ }
+ ],
created=0,
model="gpt-4o-mini",
object="chat.completion.chunk",
@@ -283,7 +313,9 @@ async def test_async_post_call_streaming_hook_callback_error_raises(proxy_loggin
@pytest.mark.asyncio
-async def test_async_post_call_streaming_iterator_hook_no_overrides_passes_through(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled):
+async def test_async_post_call_streaming_iterator_hook_no_overrides_passes_through(
+ proxy_logging, make_user_api_key_auth, mock_callbacks_disabled
+):
async def gen():
for ch in ("a", "b"):
yield ch
@@ -308,7 +340,9 @@ async def test_async_post_call_streaming_iterator_hook_no_overrides_passes_throu
@pytest.mark.asyncio
-async def test_async_post_call_streaming_iterator_hook_with_override_chains_callback(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_async_post_call_streaming_iterator_hook_with_override_chains_callback(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
class _IterOverride(CustomLogger):
async def async_post_call_streaming_iterator_hook(self, **kwargs): # type: ignore[override]
async for ch in kwargs["response"]:
@@ -331,7 +365,9 @@ async def test_async_post_call_streaming_iterator_hook_with_override_chains_call
@pytest.mark.asyncio
-async def test_async_post_call_streaming_iterator_hook_upstream_error_raises(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled):
+async def test_async_post_call_streaming_iterator_hook_upstream_error_raises(
+ proxy_logging, make_user_api_key_auth, mock_callbacks_disabled
+):
async def gen():
if False:
yield # pragma: no cover
@@ -362,14 +398,20 @@ async def test_fire_deferred_stream_logging_fires_callback():
logging_obj._on_deferred_stream_complete = deferred
logging_obj._deferred_stream_complete_args = ("payload",)
- ProxyLogging._fire_deferred_stream_logging(request_data={"litellm_logging_obj": logging_obj})
+ ProxyLogging._fire_deferred_stream_logging(
+ request_data={"litellm_logging_obj": logging_obj}
+ )
await asyncio.sleep(0)
snapshot = {
"arg": captured["arg"],
"callback_cleared": logging_obj._on_deferred_stream_complete is None,
"args_cleared": logging_obj._deferred_stream_complete_args is None,
}
- assert snapshot == {"arg": "payload", "callback_cleared": True, "args_cleared": True}
+ assert snapshot == {
+ "arg": "payload",
+ "callback_cleared": True,
+ "args_cleared": True,
+ }
def test_fire_deferred_stream_logging_no_logging_obj_no_error():
@@ -391,13 +433,17 @@ async def test_post_call_response_headers_hook_returns_empty_when_no_callbacks(
proxy_logging, mock_callbacks_disabled, make_user_api_key_auth
):
out = await proxy_logging.post_call_response_headers_hook(
- data={}, user_api_key_dict=make_user_api_key_auth(), response=MagicMock(_hidden_params={})
+ data={},
+ user_api_key_dict=make_user_api_key_auth(),
+ response=MagicMock(_hidden_params={}),
)
assert out == {}
@pytest.mark.asyncio
-async def test_post_call_response_headers_hook_merges_callback_headers(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_post_call_response_headers_hook_merges_callback_headers(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
class _Cb(CustomLogger):
async def async_post_call_response_headers_hook(self, **kwargs): # type: ignore[override]
return {"X-One": "1", "X-Two": "2", "X-Common": "first"}
@@ -416,7 +462,9 @@ async def test_post_call_response_headers_hook_merges_callback_headers(proxy_log
@pytest.mark.asyncio
-async def test_post_call_response_headers_hook_swallows_callback_error(proxy_logging, make_user_api_key_auth, monkeypatch):
+async def test_post_call_response_headers_hook_swallows_callback_error(
+ proxy_logging, make_user_api_key_auth, monkeypatch
+):
"""Errors inside the hook are caught — function returns merged so-far."""
class _Cb(CustomLogger):
diff --git a/tests/test_litellm/responses/test_responses_api_bridge_flag.py b/tests/test_litellm/responses/test_responses_api_bridge_flag.py
index 463af6562f1..ca3a2a8a64e 100644
--- a/tests/test_litellm/responses/test_responses_api_bridge_flag.py
+++ b/tests/test_litellm/responses/test_responses_api_bridge_flag.py
@@ -151,9 +151,7 @@ class TestUseResponsesApiBridgeFlag:
output=[
{"type": "message", "content": [{"type": "text", "text": "Answer"}]}
],
- usage=ResponseAPIUsage(
- input_tokens=10, output_tokens=5, total_tokens=15
- ),
+ usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15),
)
mock_call_aresponses.return_value = mock_response
@@ -202,9 +200,7 @@ class TestUseResponsesApiBridgeFlag:
"arguments": '{"queries": ["test query"]}',
}
],
- usage=ResponseAPIUsage(
- input_tokens=10, output_tokens=5, total_tokens=15
- ),
+ usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15),
)
second_response = ResponsesAPIResponse(
id="resp_second",
@@ -216,9 +212,7 @@ class TestUseResponsesApiBridgeFlag:
"content": [{"type": "text", "text": "Final answer"}],
}
],
- usage=ResponseAPIUsage(
- input_tokens=20, output_tokens=10, total_tokens=30
- ),
+ usage=ResponseAPIUsage(input_tokens=20, output_tokens=10, total_tokens=30),
)
mock_bridge_handler.side_effect = [first_response, second_response]
@@ -267,9 +261,7 @@ class TestUseResponsesApiBridgeFlag:
"content": [{"type": "text", "text": "Native response"}],
}
],
- usage=ResponseAPIUsage(
- input_tokens=10, output_tokens=5, total_tokens=15
- ),
+ usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15),
)
result = await litellm.aresponses(
diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_router_dispatch.py b/tests/test_litellm/router_strategy/adaptive_router/test_router_dispatch.py
index 604155e1221..2703d7b28c6 100644
--- a/tests/test_litellm/router_strategy/adaptive_router/test_router_dispatch.py
+++ b/tests/test_litellm/router_strategy/adaptive_router/test_router_dispatch.py
@@ -454,9 +454,7 @@ def test_finalize_prunes_stale_adaptive_router_hooks_from_callbacks():
Router(model_list=model_list) # simulate hot-reload
adaptive_hooks = [
- cb
- for cb in litellm.callbacks
- if isinstance(cb, AdaptiveRouterPostCallHook)
+ cb for cb in litellm.callbacks if isinstance(cb, AdaptiveRouterPostCallHook)
]
assert len(adaptive_hooks) == 1, (
f"expected exactly one AdaptiveRouterPostCallHook after hot-reload, "
diff --git a/tests/test_litellm/test_acompletion_session_reuse_e2e.py b/tests/test_litellm/test_acompletion_session_reuse_e2e.py
index 79b947bb146..a6f2be09cef 100644
--- a/tests/test_litellm/test_acompletion_session_reuse_e2e.py
+++ b/tests/test_litellm/test_acompletion_session_reuse_e2e.py
@@ -22,7 +22,6 @@ sys.path.insert(0, os.path.abspath("../../.."))
import litellm
-
# ============================================================================
# HELPER FUNCTION
# ============================================================================
diff --git a/tests/test_litellm/test_anthropic_skills_transformation.py b/tests/test_litellm/test_anthropic_skills_transformation.py
index 1b917f08ca9..1b41b05f6d7 100644
--- a/tests/test_litellm/test_anthropic_skills_transformation.py
+++ b/tests/test_litellm/test_anthropic_skills_transformation.py
@@ -22,7 +22,6 @@ from litellm.types.llms.anthropic_skills import (
)
from litellm.types.router import GenericLiteLLMParams
-
FAKE_API_KEY = "sk-ant-test-key-1234"
FAKE_API_BASE = "https://api.anthropic.com"
diff --git a/tests/test_litellm/test_budget_ratchet_check.py b/tests/test_litellm/test_budget_ratchet_check.py
index 9f19944fdba..b05244e4855 100644
--- a/tests/test_litellm/test_budget_ratchet_check.py
+++ b/tests/test_litellm/test_budget_ratchet_check.py
@@ -10,7 +10,9 @@ import subprocess
import sys
from pathlib import Path
-_MODULE_PATH = Path(__file__).resolve().parents[2] / "scripts" / "budget_ratchet_check.py"
+_MODULE_PATH = (
+ Path(__file__).resolve().parents[2] / "scripts" / "budget_ratchet_check.py"
+)
_spec = importlib.util.spec_from_file_location("budget_ratchet_check", _MODULE_PATH)
ratchet = importlib.util.module_from_spec(_spec)
_spec.loader.exec_module(ratchet)
diff --git a/tests/test_litellm/test_check_type_discipline.py b/tests/test_litellm/test_check_type_discipline.py
index 436904b017c..09500530948 100644
--- a/tests/test_litellm/test_check_type_discipline.py
+++ b/tests/test_litellm/test_check_type_discipline.py
@@ -43,7 +43,9 @@ def test_scan_comments_tokenizes_every_comment():
def test_scan_comments_does_not_crash_on_malformed_source():
# A dedent mismatch makes tokenize raise IndentationError (a SyntaxError subclass);
# scan_comments must swallow it, not propagate and crash the whole run.
- comments, violations = checker.scan_comments(Path("x.py"), "if True:\n a = 1\n b = 2\n")
+ comments, violations = checker.scan_comments(
+ Path("x.py"), "if True:\n a = 1\n b = 2\n"
+ )
assert violations == ()
assert comments.cast_ok_lines == frozenset()
@@ -59,7 +61,9 @@ def test_noqa_without_codes_is_flagged(tmp_path):
def test_noqa_with_codes_and_reason_is_clean(tmp_path):
- assert "LIT003" not in _codes(tmp_path, "x = 1 # noqa: TID251 # legacy import, removed in #123\n")
+ assert "LIT003" not in _codes(
+ tmp_path, "x = 1 # noqa: TID251 # legacy import, removed in #123\n"
+ )
def test_ignore_without_reason_is_flagged(tmp_path):
@@ -67,13 +71,18 @@ def test_ignore_without_reason_is_flagged(tmp_path):
def test_ignore_with_codes_and_reason_is_clean(tmp_path):
- assert "LIT004" not in _codes(tmp_path, "x = 1 # pyright: ignore[reportArgumentType] # upstream stub is wrong\n")
+ assert "LIT004" not in _codes(
+ tmp_path,
+ "x = 1 # pyright: ignore[reportArgumentType] # upstream stub is wrong\n",
+ )
def test_ok_suppression_without_reason_is_flagged(tmp_path):
codes = _codes(tmp_path, "y = [] # mutable-ok\n")
assert "LIT005" in codes # reasonless suppression
- assert "LIT002" in codes # and it does not suppress, so the construction still trips
+ assert (
+ "LIT002" in codes
+ ) # and it does not suppress, so the construction still trips
# --------------------------------------------------------------------------- #
@@ -91,8 +100,15 @@ def test_typing_alias_and_forward_ref_annotations_are_flagged(tmp_path):
def test_readonly_annotations_are_clean(tmp_path):
- for ann in ("Mapping[str, int]", "Sequence[int]", "tuple[int, ...]", "frozenset[int]"):
- assert "LIT001" not in _codes(tmp_path, f"from typing import Mapping, Sequence\nx: {ann}\n")
+ for ann in (
+ "Mapping[str, int]",
+ "Sequence[int]",
+ "tuple[int, ...]",
+ "frozenset[int]",
+ ):
+ assert "LIT001" not in _codes(
+ tmp_path, f"from typing import Mapping, Sequence\nx: {ann}\n"
+ )
def test_mutable_construction_is_flagged(tmp_path):
@@ -103,7 +119,8 @@ def test_mutable_construction_is_flagged(tmp_path):
def test_construction_inside_annotation_is_exempt(tmp_path):
# `Callable[[int], str]` carries a list display that is type syntax, not construction.
assert "LIT002" not in _codes(
- tmp_path, "from typing import Callable\ndef f(cb: Callable[[int], str]) -> None:\n return None\n"
+ tmp_path,
+ "from typing import Callable\ndef f(cb: Callable[[int], str]) -> None:\n return None\n",
)
@@ -123,11 +140,16 @@ def test_dict_list_set_method_calls_are_not_construction(tmp_path):
def test_qualified_collections_constructors_still_count(tmp_path):
# collections concretes are rarely method names, so a qualified call still flags.
assert "LIT002" in _codes(tmp_path, "import collections\nq = collections.deque()\n")
- assert "LIT002" in _codes(tmp_path, "import collections\nm = collections.defaultdict(list)\n")
+ assert "LIT002" in _codes(
+ tmp_path, "import collections\nm = collections.defaultdict(list)\n"
+ )
def test_mutable_ok_with_reason_suppresses_both_rules(tmp_path):
- codes = _codes(tmp_path, "x: dict[str, int] = {} # mutable-ok: in-place buffer mutated hot path\n")
+ codes = _codes(
+ tmp_path,
+ "x: dict[str, int] = {} # mutable-ok: in-place buffer mutated hot path\n",
+ )
assert "LIT001" not in codes
assert "LIT002" not in codes
@@ -138,12 +160,15 @@ def test_mutable_ok_with_reason_suppresses_both_rules(tmp_path):
def test_cast_call_is_flagged(tmp_path):
- assert "LIT006" in _codes(tmp_path, "from typing import cast\ny = cast(int, object())\n")
+ assert "LIT006" in _codes(
+ tmp_path, "from typing import cast\ny = cast(int, object())\n"
+ )
def test_cast_ok_with_reason_suppresses(tmp_path):
assert "LIT006" not in _codes(
- tmp_path, "from typing import cast\ny = cast(int, object()) # cast-ok: validated by schema above\n"
+ tmp_path,
+ "from typing import cast\ny = cast(int, object()) # cast-ok: validated by schema above\n",
)
@@ -183,9 +208,12 @@ def test_kwargs_parameter_is_flagged(tmp_path):
def test_typed_args_is_clean_but_kwargs_ok_suppresses(tmp_path):
- assert "LIT008" not in _codes(tmp_path, "def f(*args: int) -> None:\n return None\n")
assert "LIT008" not in _codes(
- tmp_path, "def f(**kwargs: int) -> None: # kwargs-ok: passthrough to a third-party sink\n return None\n"
+ tmp_path, "def f(*args: int) -> None:\n return None\n"
+ )
+ assert "LIT008" not in _codes(
+ tmp_path,
+ "def f(**kwargs: int) -> None: # kwargs-ok: passthrough to a third-party sink\n return None\n",
)
diff --git a/tests/test_litellm/test_dashscope_image_generation.py b/tests/test_litellm/test_dashscope_image_generation.py
index af95e2ca6b4..aa3c955f838 100644
--- a/tests/test_litellm/test_dashscope_image_generation.py
+++ b/tests/test_litellm/test_dashscope_image_generation.py
@@ -18,7 +18,6 @@ from litellm.llms.dashscope.image_generation.transformation import (
from litellm.types.utils import ImageObject, ImageResponse
from litellm.utils import get_llm_provider
-
# ---------------------------------------------------------------------------
# 1. Provider detection
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/test_deepseek_model_metadata.py b/tests/test_litellm/test_deepseek_model_metadata.py
index 4900af5d97d..2b86b61b019 100644
--- a/tests/test_litellm/test_deepseek_model_metadata.py
+++ b/tests/test_litellm/test_deepseek_model_metadata.py
@@ -23,7 +23,6 @@ from litellm.utils import (
supports_response_schema,
)
-
# ---------------------------------------------------------------------------
# Data-level tests – verify the JSON files are in sync
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/test_git_hooks.py b/tests/test_litellm/test_git_hooks.py
index c6980d1f44a..8a0805d6e94 100644
--- a/tests/test_litellm/test_git_hooks.py
+++ b/tests/test_litellm/test_git_hooks.py
@@ -64,7 +64,9 @@ def _run_pre_push(stdin: str) -> subprocess.CompletedProcess:
)
-def _ref_line(branch: str, local_oid: str = _NONZERO_OID, remote_oid: str = _ZERO_OID) -> str:
+def _ref_line(
+ branch: str, local_oid: str = _NONZERO_OID, remote_oid: str = _ZERO_OID
+) -> str:
ref = f"refs/heads/{branch}"
return f"{ref} {local_oid} {ref} {remote_oid}\n"
@@ -97,13 +99,13 @@ def test_commit_msg_accepts_conventional_subjects(tmp_path, subject):
@pytest.mark.parametrize(
"subject",
[
- "add stuff", # no type
- "feat add router strategy", # missing colon
- "feat:add router strategy", # missing space after colon
- "feat():", # empty description
- "ux: thing", # unknown type
- "Feat(router): capital type", # types are lowercase
- "feat(router):", # empty description
+ "add stuff", # no type
+ "feat add router strategy", # missing colon
+ "feat:add router strategy", # missing space after colon
+ "feat():", # empty description
+ "ux: thing", # unknown type
+ "Feat(router): capital type", # types are lowercase
+ "feat(router):", # empty description
# Description must start with a lowercase letter — kept in sync with
# the CI workflow's subjectPattern so the local hook never accepts a
# subject that CI will later reject.
@@ -168,7 +170,9 @@ def test_commit_msg_rejects_empty_message(tmp_path):
def test_commit_msg_skips_comment_only_lines(tmp_path):
# An all-comments file has no subject — should be rejected.
msg_file = tmp_path / "COMMIT_EDITMSG"
- msg_file.write_text("# please enter a commit message\n# above this line\n", encoding="utf-8")
+ msg_file.write_text(
+ "# please enter a commit message\n# above this line\n", encoding="utf-8"
+ )
result = subprocess.run(
["bash", str(_COMMIT_MSG_HOOK), str(msg_file)],
capture_output=True,
@@ -227,10 +231,10 @@ def test_pre_push_accepts_conventional_branches(branch):
[
"random-branch-name",
"litellm_fix/optimize-streaming", # legacy pattern is now rejected
- "ui/navbar-notifications", # not in the allow list
- "feature/", # empty description
- "Feature/foo", # type is case-sensitive
- "feat/foo", # angular commit type, not branch type
+ "ui/navbar-notifications", # not in the allow list
+ "feature/", # empty description
+ "Feature/foo", # type is case-sensitive
+ "feat/foo", # angular commit type, not branch type
],
)
def test_pre_push_rejects_non_conventional_branches(branch):
diff --git a/tests/test_litellm/test_github_triage_workflows.py b/tests/test_litellm/test_github_triage_workflows.py
index ec6e9fc2381..9c54ae7f82c 100644
--- a/tests/test_litellm/test_github_triage_workflows.py
+++ b/tests/test_litellm/test_github_triage_workflows.py
@@ -88,7 +88,9 @@ def _all_run_blocks(workflow: dict) -> list[str]:
@pytest.mark.parametrize("workflow_file,env_var", sorted(DESTRUCTIVE_GATE_ENV.items()))
-def test_should_use_failsafe_equals_true_comparison(workflow_file: str, env_var: str) -> None:
+def test_should_use_failsafe_equals_true_comparison(
+ workflow_file: str, env_var: str
+) -> None:
"""The destructive `--close` gate must use `= "true"` (fail-safe), not
`!= "false"` (which would treat "True", "yes", "1", or any typo as
enabling closure).
@@ -101,9 +103,9 @@ def test_should_use_failsafe_equals_true_comparison(workflow_file: str, env_var:
"""
workflow = _load_workflow(workflow_file)
text = "\n".join(_all_run_blocks(workflow))
- assert env_var in text, (
- f"{workflow_file} no longer references {env_var}; was the gating env var renamed without updating this test?"
- )
+ assert (
+ env_var in text
+ ), f"{workflow_file} no longer references {env_var}; was the gating env var renamed without updating this test?"
accepted_patterns = (
f'"${{{env_var}}}" = "true"',
f'"${{{env_var}:-false}}" = "true"',
@@ -186,14 +188,18 @@ def test_triage_requirements_are_fully_hash_pinned() -> None:
install time. A loosened pin or a missing hash here would silently widen
the supply-chain surface for all the installer workflows.
"""
- assert REQUIREMENTS_FILE.exists(), (
- f"the hash-pinned requirements file the triage workflows install from is missing at {REQUIREMENTS_FILE}"
- )
+ assert (
+ REQUIREMENTS_FILE.exists()
+ ), f"the hash-pinned requirements file the triage workflows install from is missing at {REQUIREMENTS_FILE}"
joined = REQUIREMENTS_FILE.read_text().replace("\\\n", " ")
- entries = [line.strip() for line in joined.splitlines() if line.strip() and not line.strip().startswith("#")]
- assert any(e.split()[0].startswith("openai==") for e in entries), (
- "openai must be pinned to an exact version in the triage requirements"
- )
+ entries = [
+ line.strip()
+ for line in joined.splitlines()
+ if line.strip() and not line.strip().startswith("#")
+ ]
+ assert any(
+ e.split()[0].startswith("openai==") for e in entries
+ ), "openai must be pinned to an exact version in the triage requirements"
for entry in entries:
spec = entry.split()[0]
assert "==" in spec, (
@@ -209,7 +215,10 @@ def test_triage_requirements_are_fully_hash_pinned() -> None:
def _heads_up_run_step() -> dict:
workflow = _load_workflow("triage_rollout_heads_up.yml")
for step in workflow["jobs"]["heads-up"]["steps"]:
- if isinstance(step.get("run"), str) and "triage_rollout_heads_up.py" in step["run"]:
+ if (
+ isinstance(step.get("run"), str)
+ and "triage_rollout_heads_up.py" in step["run"]
+ ):
return step
raise AssertionError("no run step invokes triage_rollout_heads_up.py")
@@ -226,15 +235,15 @@ def test_rollout_heads_up_push_trigger_never_posts() -> None:
post real comments on every push that touches the script.
"""
run = _heads_up_run_step()["run"]
- assert '"${GITHUB_EVENT_NAME:-}" = "workflow_dispatch"' in run, (
- "the real (--close) run must be a manual workflow_dispatch, not the automatic push trigger"
- )
- assert '"${DRY_RUN_INPUT:-true}" = "false"' in run, (
- "the real run must require the dry_run input to be the exact string 'false' (fail-safe); any other value stays dry-run"
- )
- assert run.count("ARGS+=(--close)") == 1, (
- "--close must appear once, inside the manual real-run branch; a second occurrence means the push path posts real comments on merge"
- )
+ assert (
+ '"${GITHUB_EVENT_NAME:-}" = "workflow_dispatch"' in run
+ ), "the real (--close) run must be a manual workflow_dispatch, not the automatic push trigger"
+ assert (
+ '"${DRY_RUN_INPUT:-true}" = "false"' in run
+ ), "the real run must require the dry_run input to be the exact string 'false' (fail-safe); any other value stays dry-run"
+ assert (
+ run.count("ARGS+=(--close)") == 1
+ ), "--close must appear once, inside the manual real-run branch; a second occurrence means the push path posts real comments on merge"
def test_rollout_heads_up_key_is_dispatch_gated() -> None:
@@ -244,9 +253,9 @@ def test_rollout_heads_up_key_is_dispatch_gated() -> None:
key to the automatic push run, which must stay a no-op dry-run preview.
"""
key_expr = (_heads_up_run_step().get("env") or {}).get("OPENAI_API_KEY", "")
- assert "github.event_name == 'workflow_dispatch'" in key_expr, (
- f"OPENAI_API_KEY must be gated on workflow_dispatch so the automatic push trigger gets no key; found: {key_expr!r}"
- )
+ assert (
+ "github.event_name == 'workflow_dispatch'" in key_expr
+ ), f"OPENAI_API_KEY must be gated on workflow_dispatch so the automatic push trigger gets no key; found: {key_expr!r}"
def _reconsider_steps() -> list[dict]:
@@ -266,7 +275,9 @@ def _reaction_steps(steps: list[dict], content: str) -> list[tuple[int, dict]]:
return [
(i, s)
for i, s in enumerate(steps)
- if isinstance(s.get("run"), str) and f"content={content}" in s["run"] and "/reactions" in s["run"]
+ if isinstance(s.get("run"), str)
+ and f"content={content}" in s["run"]
+ and "/reactions" in s["run"]
]
@@ -289,13 +300,15 @@ class TestReconsiderReactions:
assert len(eyes) == 1, "expected exactly one 👀 (eyes) reaction step"
idx, step = eyes[0]
assert idx < run_idx, "👀 must be posted BEFORE the slow triage run, not after"
- assert "github.event.comment.id" in (step.get("env") or {}).get("COMMENT_ID", ""), (
- "👀 must react to the comment that triggered the workflow"
- )
- assert "${COMMENT_ID}" in step["run"], "👀 must react to the triggering comment, not a hardcoded id"
- assert "vars.AGENT_SHIN_ENABLED == 'true'" in step["if"], (
- "👀 must be gated on AGENT_SHIN_ENABLED so dry-run stays inert"
- )
+ assert "github.event.comment.id" in (step.get("env") or {}).get(
+ "COMMENT_ID", ""
+ ), "👀 must react to the comment that triggered the workflow"
+ assert (
+ "${COMMENT_ID}" in step["run"]
+ ), "👀 must react to the triggering comment, not a hardcoded id"
+ assert (
+ "vars.AGENT_SHIN_ENABLED == 'true'" in step["if"]
+ ), "👀 must be gated on AGENT_SHIN_ENABLED so dry-run stays inert"
def test_thumbs_up_reaction_is_posted_after_a_successful_run(self) -> None:
steps = _reconsider_steps()
@@ -304,7 +317,9 @@ class TestReconsiderReactions:
assert len(thumbs) == 1, "expected exactly one 👍 (+1) reaction step"
idx, step = thumbs[0]
assert idx > run_idx, "👍 must come AFTER the triage run"
- assert "success()" in step["if"], "👍 must only fire when the reconsider run succeeded"
- assert "vars.AGENT_SHIN_ENABLED == 'true'" in step["if"], (
- "👍 must be gated on AGENT_SHIN_ENABLED so dry-run stays inert"
- )
+ assert (
+ "success()" in step["if"]
+ ), "👍 must only fire when the reconsider run succeeded"
+ assert (
+ "vars.AGENT_SHIN_ENABLED == 'true'" in step["if"]
+ ), "👍 must be gated on AGENT_SHIN_ENABLED so dry-run stays inert"
diff --git a/tests/test_litellm/test_model_cost_aliases.py b/tests/test_litellm/test_model_cost_aliases.py
index f9f92a85cb2..ee11ebb9702 100644
--- a/tests/test_litellm/test_model_cost_aliases.py
+++ b/tests/test_litellm/test_model_cost_aliases.py
@@ -10,7 +10,6 @@ from unittest.mock import patch
from litellm import verbose_logger
from litellm.litellm_core_utils.get_model_cost_map import _expand_model_aliases
-
# ---------------------------------------------------------------------------
# Core expansion behaviour
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/test_nested_drop_params.py b/tests/test_litellm/test_nested_drop_params.py
index bb1305ffde4..fbcd5302308 100644
--- a/tests/test_litellm/test_nested_drop_params.py
+++ b/tests/test_litellm/test_nested_drop_params.py
@@ -7,7 +7,6 @@ This tests the new JSONPath-like syntax for removing nested fields.
import os
import sys
-
# Add parent directory to path
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))
diff --git a/tests/test_litellm/test_openai_embedding_encoding_format_default.py b/tests/test_litellm/test_openai_embedding_encoding_format_default.py
index 94e4e3c81e5..1cf71139d87 100644
--- a/tests/test_litellm/test_openai_embedding_encoding_format_default.py
+++ b/tests/test_litellm/test_openai_embedding_encoding_format_default.py
@@ -51,9 +51,7 @@ def test_openai_embedding_encoding_format_default(
@pytest.mark.parametrize("env_none", ["none", "NONE", " none "])
-def test_openai_embedding_encoding_format_env_none_omits_param(
- monkeypatch, env_none
-):
+def test_openai_embedding_encoding_format_env_none_omits_param(monkeypatch, env_none):
"""LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT=none omits encoding_format (provider default)."""
monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", env_none)
diff --git a/tests/test_litellm/test_router_google_genai.py b/tests/test_litellm/test_router_google_genai.py
index 81dd7bbdc40..b8179e568a7 100644
--- a/tests/test_litellm/test_router_google_genai.py
+++ b/tests/test_litellm/test_router_google_genai.py
@@ -2,6 +2,7 @@
"""
Test to verify the new Google GenAI router methods
"""
+
import asyncio
import os
import sys
diff --git a/tests/test_litellm/test_router_weighted_failover.py b/tests/test_litellm/test_router_weighted_failover.py
index 8faf6bcd9cf..3e429302123 100644
--- a/tests/test_litellm/test_router_weighted_failover.py
+++ b/tests/test_litellm/test_router_weighted_failover.py
@@ -16,7 +16,6 @@ import pytest
from litellm import Router
from litellm.utils import _get_excluded_filtered_deployments
-
# ---------------------------------------------------------------------------
# Unit tests for _get_excluded_filtered_deployments
# ---------------------------------------------------------------------------
diff --git a/tests/test_litellm/test_ssl_verify_unit.py b/tests/test_litellm/test_ssl_verify_unit.py
index 7cc15703a3b..bd3655b2b80 100644
--- a/tests/test_litellm/test_ssl_verify_unit.py
+++ b/tests/test_litellm/test_ssl_verify_unit.py
@@ -19,7 +19,9 @@ import litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks as _
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail
-from litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks import CatoNetworksGuardrail
+from litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks import (
+ CatoNetworksGuardrail,
+)
class TestBaseAWSLLMSSLVerify:
@@ -156,7 +158,7 @@ class TestCatoNetworksGuardrailSSLVerify:
# Use patch.object on the actual module reference for reliable patching
# across different import orders / CI environments
with patch.object(
- _cato_networks_module, "get_async_httpx_client", return_value=mock_handler
+ _cato_networks_module, "get_async_httpx_client", return_value=mock_handler
) as mock_get_client:
# Initialize with ssl_verify
cert_path = "/path/to/cato_cert.pem"
@@ -179,10 +181,12 @@ class TestCatoNetworksGuardrailSSLVerify:
# Use patch.object on the actual module reference for reliable patching
with patch.object(
- _cato_networks_module, "get_async_httpx_client", return_value=mock_handler
+ _cato_networks_module, "get_async_httpx_client", return_value=mock_handler
) as mock_get_client:
# Initialize without ssl_verify
- CatoNetworksGuardrail(api_key="test_key", api_base="https://test.catonetworks.api")
+ CatoNetworksGuardrail(
+ api_key="test_key", api_base="https://test.catonetworks.api"
+ )
# Should still work, just without custom SSL
assert mock_get_client.called
diff --git a/tests/test_litellm/test_streaming_connection_cleanup.py b/tests/test_litellm/test_streaming_connection_cleanup.py
index 5a81a3ffb17..53bbe8abd3c 100644
--- a/tests/test_litellm/test_streaming_connection_cleanup.py
+++ b/tests/test_litellm/test_streaming_connection_cleanup.py
@@ -19,7 +19,6 @@ from litellm.llms.custom_httpx.aiohttp_transport import (
LiteLLMAiohttpTransport,
)
-
# ── aiohttp transport layer tests ──────────────────────────────
diff --git a/tests/test_litellm/test_thinking_enabled.py b/tests/test_litellm/test_thinking_enabled.py
index 8ba406c395a..94fc4d3bb1f 100644
--- a/tests/test_litellm/test_thinking_enabled.py
+++ b/tests/test_litellm/test_thinking_enabled.py
@@ -14,6 +14,7 @@ class TestIsThinkingEnabled:
@pytest.fixture
def transformer(self):
"""Create a BaseConfig instance for testing."""
+
# BaseConfig is abstract, so we create a minimal concrete subclass
class ConcreteConfig(BaseConfig):
def __init__(self):
@@ -39,6 +40,7 @@ class TestIsThinkingEnabled:
def get_error_class(self, *args, **kwargs):
from litellm.llms.base_llm.chat.transformation import BaseLLMException
+
return BaseLLMException(500, "test error")
return ConcreteConfig()
@@ -69,6 +71,6 @@ class TestIsThinkingEnabled:
def test_is_thinking_enabled(self, transformer, non_default_params, expected):
"""Test is_thinking_enabled with various parameter combinations."""
result = transformer.is_thinking_enabled(non_default_params)
- assert result == expected, (
- f"Expected {expected} for params {non_default_params}, got {result}"
- )
+ assert (
+ result == expected
+ ), f"Expected {expected} for params {non_default_params}, got {result}"
diff --git a/tests/test_litellm/test_type_discipline_gate.py b/tests/test_litellm/test_type_discipline_gate.py
index d7d827685a6..f78395eabc0 100644
--- a/tests/test_litellm/test_type_discipline_gate.py
+++ b/tests/test_litellm/test_type_discipline_gate.py
@@ -8,7 +8,9 @@ drift-safe breach check). Both are pinned here.
import importlib.util
from pathlib import Path
-_MODULE_PATH = Path(__file__).resolve().parents[2] / "scripts" / "type_discipline_gate.py"
+_MODULE_PATH = (
+ Path(__file__).resolve().parents[2] / "scripts" / "type_discipline_gate.py"
+)
_spec = importlib.util.spec_from_file_location("type_discipline_gate", _MODULE_PATH)
gate = importlib.util.module_from_spec(_spec)
_spec.loader.exec_module(gate)
@@ -21,19 +23,28 @@ def _budget(baseline, slack):
def test_over_ceiling_flags_only_counts_above_baseline_plus_slack():
budget = _budget(10, 2) # cap 12
assert gate.over_ceiling({"LIT006": 12}, budget) == frozenset() # at cap
- assert gate.over_ceiling({"LIT006": 13}, budget) == frozenset({"LIT006"}) # over cap
+ assert gate.over_ceiling({"LIT006": 13}, budget) == frozenset(
+ {"LIT006"}
+ ) # over cap
assert gate.over_ceiling({}, budget) == frozenset() # missing rule counts as zero
def test_over_ceiling_is_independent_across_rules():
- budget = {"LIT001": {"baseline": 5, "slack": 0}, "LIT006": {"baseline": 10, "slack": 0}}
- assert gate.over_ceiling({"LIT001": 6, "LIT006": 10}, budget) == frozenset({"LIT001"})
+ budget = {
+ "LIT001": {"baseline": 5, "slack": 0},
+ "LIT006": {"baseline": 10, "slack": 0},
+ }
+ assert gate.over_ceiling({"LIT001": 6, "LIT006": 10}, budget) == frozenset(
+ {"LIT001"}
+ )
def test_evaluate_blames_only_a_rule_over_cap_and_over_base():
budget = _budget(10, 0) # cap 10
# over cap and grown vs base -> breach
- assert [b.rule for b in gate.evaluate({"LIT006": 12}, {"LIT006": 9}, budget)] == ["LIT006"]
+ assert [b.rule for b in gate.evaluate({"LIT006": 12}, {"LIT006": 9}, budget)] == [
+ "LIT006"
+ ]
# over cap but flat vs base (pre-existing drift) -> not blamed
assert gate.evaluate({"LIT006": 12}, {"LIT006": 12}, budget) == []
# within cap -> not blamed regardless of base
diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py
index 44e0b55ee3b..611457934f3 100644
--- a/tests/test_litellm/test_utils.py
+++ b/tests/test_litellm/test_utils.py
@@ -1178,9 +1178,7 @@ def test_check_provider_match_none_value_matches_any_provider():
is True
)
# When custom_llm_provider is also None nothing constrains the match.
- assert (
- litellm.utils._check_provider_match({"litellm_provider": None}, None) is True
- )
+ assert litellm.utils._check_provider_match({"litellm_provider": None}, None) is True
def test_get_provider_rerank_config():
@@ -1545,8 +1543,7 @@ class TestProxyFunctionCalling:
assert result is True, "Resolvable model names work with fallback logic"
# Documentation notes:
- print(
- """
+ print("""
PROXY MODEL RESOLUTION BEHAVIOR:
✅ WORKS (with current fallback logic):
@@ -1561,8 +1558,7 @@ class TestProxyFunctionCalling:
💡 SOLUTION: Use LiteLLM proxy server with proper model_list configuration
that maps custom names to underlying models.
- """
- )
+ """)
@pytest.mark.parametrize(
"proxy_model_with_hints,expected_result",
@@ -1924,8 +1920,7 @@ class TestProxyFunctionCalling:
This test provides documentation on how the proxy server configuration
would typically map custom model names to underlying models.
"""
- print(
- """
+ print("""
REAL-WORLD PROXY SERVER CONFIGURATION EXAMPLE:
===============================================
@@ -1978,8 +1973,7 @@ class TestProxyFunctionCalling:
- Consistent request/response format
- Enhanced streaming support for function calls
- """
- )
+ """)
# Verify that direct underlying models work as expected
bedrock_models = [
@@ -2193,8 +2187,7 @@ class TestProxyFunctionCalling:
This test provides documentation on how the proxy server configuration
would typically map custom model names to underlying models.
"""
- print(
- """
+ print("""
REAL-WORLD PROXY SERVER CONFIGURATION EXAMPLE:
===============================================
@@ -2247,8 +2240,7 @@ class TestProxyFunctionCalling:
- Consistent request/response format
- Enhanced streaming support for function calls
- """
- )
+ """)
# Verify that direct underlying models work as expected
bedrock_models = [
@@ -2462,8 +2454,7 @@ class TestProxyFunctionCalling:
This test provides documentation on how the proxy server configuration
would typically map custom model names to underlying models.
"""
- print(
- """
+ print("""
REAL-WORLD PROXY SERVER CONFIGURATION EXAMPLE:
===============================================
@@ -2516,8 +2507,7 @@ class TestProxyFunctionCalling:
- Consistent request/response format
- Enhanced streaming support for function calls
- """
- )
+ """)
# Verify that direct underlying models work as expected
bedrock_models = [
@@ -4178,7 +4168,9 @@ def test_deepseek_v4_models_in_cost_map():
("deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
- assert info is not None, f"{key} missing from model_prices_and_context_window.json"
+ assert (
+ info is not None
+ ), f"{key} missing from model_prices_and_context_window.json"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
@@ -4194,7 +4186,9 @@ def test_deepseek_v4_models_in_cost_map():
("deepseek/deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
- assert info is not None, f"{key} missing from model_prices_and_context_window.json"
+ assert (
+ info is not None
+ ), f"{key} missing from model_prices_and_context_window.json"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
@@ -4212,7 +4206,11 @@ def test_deepseek_v4_models_in_backup_cost_map():
import json
from pathlib import Path
- json_path = Path(__file__).parents[2] / "litellm" / "model_prices_and_context_window_backup.json"
+ json_path = (
+ Path(__file__).parents[2]
+ / "litellm"
+ / "model_prices_and_context_window_backup.json"
+ )
with open(json_path) as f:
model_cost = json.load(f)
@@ -4556,4 +4554,3 @@ def test_aws_bedrock_project_id_excluded_from_bedrock_optional_params():
assert "aws_bedrock_project_id" not in result
assert result["aws_region_name"] == "us-east-1"
-
diff --git a/tests/test_litellm/types/test_completion.py b/tests/test_litellm/types/test_completion.py
index f24b00df3fc..b753a8e2abb 100644
--- a/tests/test_litellm/types/test_completion.py
+++ b/tests/test_litellm/types/test_completion.py
@@ -1,7 +1,7 @@
"""
Tests for litellm.types.completion module
-This test suite validates the CompletionRequest model and its compatibility with
+This test suite validates the CompletionRequest model and its compatibility with
OpenAI ChatCompletion API message formats.
Usage:
diff --git a/tests/test_litellm/types/test_prometheus_label_value_sanitize.py b/tests/test_litellm/types/test_prometheus_label_value_sanitize.py
index 9ff7eb460e0..d8b90197c63 100644
--- a/tests/test_litellm/types/test_prometheus_label_value_sanitize.py
+++ b/tests/test_litellm/types/test_prometheus_label_value_sanitize.py
@@ -22,7 +22,7 @@ from litellm.types.integrations.prometheus import (
# Escapes per Prometheus text format
('he said "hi"', 'he said \\"hi\\"'),
(r"path\to\file", r"path\\to\\file"),
- (r'quote\"slash\\', r'quote\\\"slash\\\\'),
+ (r"quote\"slash\\", r"quote\\\"slash\\\\"),
# Non-string inputs get coerced to str first
(123, "123"),
(True, "True"),
@@ -31,4 +31,3 @@ from litellm.types.integrations.prometheus import (
)
def test_sanitize_prometheus_label_value_expected_outputs(value, expected):
assert _sanitize_prometheus_label_value(value) == expected
-
diff --git a/tests/test_litellm/types/test_uk_pii_entities.py b/tests/test_litellm/types/test_uk_pii_entities.py
index 378970adf9b..ad35ab092df 100644
--- a/tests/test_litellm/types/test_uk_pii_entities.py
+++ b/tests/test_litellm/types/test_uk_pii_entities.py
@@ -2,7 +2,11 @@
Test UK PII entity types in guardrails module
"""
-from litellm.types.guardrails import PiiEntityType, PiiEntityCategory, PII_ENTITY_CATEGORIES_MAP
+from litellm.types.guardrails import (
+ PiiEntityType,
+ PiiEntityCategory,
+ PII_ENTITY_CATEGORIES_MAP,
+)
class TestUKPiiEntities:
diff --git a/tests/test_passthrough_endpoints.py b/tests/test_passthrough_endpoints.py
index 47ac7511aa1..bf9417912c5 100644
--- a/tests/test_passthrough_endpoints.py
+++ b/tests/test_passthrough_endpoints.py
@@ -10,7 +10,6 @@ import json
import os
import dotenv
-
dotenv.load_dotenv()
diff --git a/tests/test_ratelimit.py b/tests/test_ratelimit.py
index 0469ded3f42..5f3253ae4dc 100644
--- a/tests/test_ratelimit.py
+++ b/tests/test_ratelimit.py
@@ -80,9 +80,7 @@ async def async_call(router: Router, list_of_messages) -> Any:
def sync_call(router: Router, list_of_messages) -> Any:
- return [
- router.completion(model="gpt-5-mini", messages=m) for m in list_of_messages
- ]
+ return [router.completion(model="gpt-5-mini", messages=m) for m in list_of_messages]
class ExpectNoException(Exception):
diff --git a/tests/vector_store_tests/rag/base_rag_tests.py b/tests/vector_store_tests/rag/base_rag_tests.py
index 2c5a2540a7e..c357965da53 100644
--- a/tests/vector_store_tests/rag/base_rag_tests.py
+++ b/tests/vector_store_tests/rag/base_rag_tests.py
@@ -124,9 +124,7 @@ class BaseRAGTest(ABC):
Test document {unique_id} for RAG ingestion and query.
LiteLLM provides a unified interface for 100+ LLMs.
This content should be retrievable via semantic search.
- """.encode(
- "utf-8"
- )
+ """.encode("utf-8")
file_data = (filename, text_content, "text/plain")
ingest_options = self.get_base_ingest_options()
diff --git a/tests/vector_store_tests/rag/test_rag_vertex_ai.py b/tests/vector_store_tests/rag/test_rag_vertex_ai.py
index c99840bb0fe..872383899d8 100644
--- a/tests/vector_store_tests/rag/test_rag_vertex_ai.py
+++ b/tests/vector_store_tests/rag/test_rag_vertex_ai.py
@@ -151,9 +151,7 @@ class TestRAGVertexAI(BaseRAGTest):
Test document {unique_id} for Vertex AI RAG corpus creation.
This tests the automatic corpus creation feature.
The corpus should be created and the file should be uploaded successfully.
- """.encode(
- "utf-8"
- )
+ """.encode("utf-8")
file_data = (filename, text_content, "text/plain")
# Get base options WITHOUT corpus_id to trigger creation
@@ -215,9 +213,7 @@ class TestRAGVertexAI(BaseRAGTest):
text_content = f"""
Test document {unique_id} for existing Vertex AI RAG corpus.
This tests file upload to a pre-existing corpus.
- """.encode(
- "utf-8"
- )
+ """.encode("utf-8")
file_data = (filename, text_content, "text/plain")
ingest_options = self.get_base_ingest_options()
diff --git a/tests/vector_store_tests/test_milvus_vector_store.py b/tests/vector_store_tests/test_milvus_vector_store.py
index 6627f6006d1..a0faf9d4ecd 100644
--- a/tests/vector_store_tests/test_milvus_vector_store.py
+++ b/tests/vector_store_tests/test_milvus_vector_store.py
@@ -12,7 +12,6 @@ import litellm
from litellm.vector_stores import asearch as vector_store_asearch
from litellm.vector_stores import search as vector_store_search
-
# Mock response from actual Milvus API
MOCK_MILVUS_SEARCH_RESPONSE = {
"code": 0,
diff --git a/tests/vector_store_tests/test_vertex_ai_search_api_vector_store.py b/tests/vector_store_tests/test_vertex_ai_search_api_vector_store.py
index 7cace338616..94c92c7fb94 100644
--- a/tests/vector_store_tests/test_vertex_ai_search_api_vector_store.py
+++ b/tests/vector_store_tests/test_vertex_ai_search_api_vector_store.py
@@ -7,7 +7,6 @@ import pytest
from unittest.mock import AsyncMock, MagicMock, patch
import litellm
-
# Mock response from actual Vertex AI Search API
MOCK_VERTEX_SEARCH_RESPONSE = {
"results": [
diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py b/ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py
index 8e92065c696..9ca335430f6 100644
--- a/ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py
+++ b/ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py
@@ -12,7 +12,6 @@ from fastapi import FastAPI, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse
-
app = FastAPI(title="Mock LLM Server")
app.add_middleware(
CORSMiddleware,
diff --git a/ui/litellm-dashboard/scripts/generate_compliance_prompts.py b/ui/litellm-dashboard/scripts/generate_compliance_prompts.py
index d0c29ba321a..a4c7e94613d 100644
--- a/ui/litellm-dashboard/scripts/generate_compliance_prompts.py
+++ b/ui/litellm-dashboard/scripts/generate_compliance_prompts.py
@@ -107,7 +107,7 @@ def main() -> None:
csv_basename = os.path.basename(args.csv)
lines.append(f"// Auto-generated from {csv_basename} — do not edit manually.")
lines.append(
- f"// Regenerate: python scripts/generate_compliance_prompts.py --csv ... --output ..."
+ "// Regenerate: python scripts/generate_compliance_prompts.py --csv ... --output ..."
)
lines.append("")
lines.append(