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(