From c631aa69fbcf5cd6dbcaf6b6059d131c26eb4ec1 Mon Sep 17 00:00:00 2001 From: Harsh2KumarPrasad Date: Sun, 26 Apr 2026 18:19:05 +0530 Subject: [PATCH] fix(chatgpt5.4): made changes as requested in review for lint, P2 and patch coverage --- ...to_update_price_and_context_window_file.py | 90 ++- .../workflows/run_llm_translation_tests.py | 512 +++++++------ .../enterprise_callbacks/callback_controls.py | 141 ++-- .../send_emails/base_email.py | 14 +- .../send_emails/sendgrid_email.py | 2 +- .../send_emails/smtp_email.py | 1 + .../litellm_core_utils/litellm_logging.py | 1 + .../proxy/audit_logging_endpoints.py | 10 +- .../litellm_enterprise/proxy/auth/__init__.py | 2 +- .../proxy/auth/custom_sso_handler.py | 37 +- .../proxy/common_utils/check_batch_cost.py | 70 +- .../common_utils/check_responses_cost.py | 33 +- .../proxy/hooks/managed_files.py | 262 ++++--- .../proxy/hooks/managed_vector_stores.py | 92 ++- .../internal_user_endpoints.py | 1 - .../proxy/vector_stores/endpoints.py | 6 +- .../types/enterprise_callbacks/send_emails.py | 12 +- .../llms/chatgpt/responses/transformation.py | 4 - scripts/adaptive_router_demo/eval.py | 54 +- scripts/adaptive_router_demo/traffic.py | 59 +- tests/enterprise/conftest.py | 2 - .../test_prometheus_logging_callbacks.py | 45 +- .../integrations/test_prometheus.py | 86 ++- .../proxy/auth/test_user_api_key_auth.py | 11 +- .../test_apply_guardrail_endpoint.py | 4 +- .../proxy/hooks/test_managed_files.py | 674 ++++++++++-------- .../test_internal_user_endpoints.py | 63 +- .../test_project_endpoints_prisma.py | 4 +- .../vertex_ai/test_gemini_batch_embeddings.py | 16 +- .../test_bedrock_converse_dedup_factory.py | 18 +- .../test_ocr_azure_document_intelligence.py | 5 +- .../send_emails/test_base_email.py | 179 +++-- .../send_emails/test_resend_email.py | 11 +- .../send_emails/test_sendgrid_email.py | 6 +- .../test_callback_controls.py | 332 ++++++--- .../test_batch_retrieve_input_file_id.py | 22 +- ..._retrieve_returns_unified_input_file_id.py | 4 +- .../proxy/test_enterprise_routes.py | 48 +- .../proxy/test_file_deletion_blocking.py | 188 ++--- .../proxy/test_managed_files_access_check.py | 11 +- .../proxy/test_managed_files_hook.py | 18 +- .../integrations/test_custom_guardrail.py | 9 +- .../litellm_core_utils/test_core_helpers.py | 24 +- .../test_streaming_handler.py | 14 +- .../test_anthropic_chat_transformation.py | 8 +- ...al_pass_through_adapters_transformation.py | 4 +- .../test_reasoning_auto_summary_messages.py | 22 +- .../test_anthropic_claude3_transformation.py | 16 +- .../test_chatgpt_responses_transformation.py | 32 + .../custom_httpx/test_llm_http_handler.py | 4 +- .../test_github_copilot_authenticator.py | 64 +- .../test_moonshot_chat_transformation.py | 4 +- .../test_openai_responses_transformation.py | 7 +- ...test_vertex_and_google_ai_studio_gemini.py | 29 +- .../mcp_server/test_mcp_stale_session.py | 116 +-- .../mcp_server/test_semantic_tool_filter.py | 28 +- .../proxy/auth/test_user_api_key_auth.py | 17 +- .../proxy/guardrails/test_llm_as_a_judge.py | 60 +- .../test_internal_user_endpoints.py | 12 +- .../test_llm_pass_through_endpoints.py | 6 +- .../public_endpoints/test_public_endpoints.py | 18 +- .../proxy/test_litellm_pre_call_utils.py | 43 +- .../test_litellm/proxy/test_team_org_move.py | 31 +- .../test_responses_api_bridge_flag.py | 16 +- .../adaptive_router/test_router_dispatch.py | 4 +- tests/test_litellm/test_utils.py | 6 +- .../test_prometheus_label_value_sanitize.py | 3 +- 67 files changed, 2280 insertions(+), 1467 deletions(-) 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..22be769a739 100644 --- a/.github/workflows/run_llm_translation_tests.py +++ b/.github/workflows/run_llm_translation_tests.py @@ -16,64 +16,75 @@ from pathlib import Path import json from typing import Dict, List, Tuple, Optional + # ANSI color codes for terminal output class Colors: - GREEN = '\033[92m' - RED = '\033[91m' - YELLOW = '\033[93m' - BLUE = '\033[94m' - PURPLE = '\033[95m' - CYAN = '\033[96m' - RESET = '\033[0m' - BOLD = '\033[1m' + GREEN = "\033[92m" + RED = "\033[91m" + YELLOW = "\033[93m" + BLUE = "\033[94m" + PURPLE = "\033[95m" + CYAN = "\033[96m" + RESET = "\033[0m" + BOLD = "\033[1m" + def print_colored(message: str, color: str = Colors.RESET): """Print colored message to terminal""" print(f"{color}{message}{Colors.RESET}") + def get_provider_from_test_file(test_file: str) -> str: """Map test file names to provider names""" provider_mapping = { - 'test_anthropic': 'Anthropic', - 'test_azure': 'Azure', - 'test_bedrock': 'AWS Bedrock', - 'test_openai': 'OpenAI', - 'test_vertex': 'Google Vertex AI', - 'test_gemini': 'Google Vertex AI', - 'test_cohere': 'Cohere', - 'test_databricks': 'Databricks', - 'test_groq': 'Groq', - 'test_together': 'Together AI', - 'test_mistral': 'Mistral', - 'test_deepseek': 'DeepSeek', - 'test_replicate': 'Replicate', - 'test_huggingface': 'HuggingFace', - 'test_fireworks': 'Fireworks AI', - 'test_perplexity': 'Perplexity', - 'test_cloudflare': 'Cloudflare', - 'test_voyage': 'Voyage AI', - 'test_xai': 'xAI', - 'test_nvidia': 'NVIDIA', - 'test_watsonx': 'IBM watsonx', - 'test_azure_ai': 'Azure AI', - 'test_snowflake': 'Snowflake', - 'test_infinity': 'Infinity', - 'test_jina': 'Jina AI', - 'test_deepgram': 'Deepgram', - 'test_clarifai': 'Clarifai', - 'test_triton': 'Triton', + "test_anthropic": "Anthropic", + "test_azure": "Azure", + "test_bedrock": "AWS Bedrock", + "test_openai": "OpenAI", + "test_vertex": "Google Vertex AI", + "test_gemini": "Google Vertex AI", + "test_cohere": "Cohere", + "test_databricks": "Databricks", + "test_groq": "Groq", + "test_together": "Together AI", + "test_mistral": "Mistral", + "test_deepseek": "DeepSeek", + "test_replicate": "Replicate", + "test_huggingface": "HuggingFace", + "test_fireworks": "Fireworks AI", + "test_perplexity": "Perplexity", + "test_cloudflare": "Cloudflare", + "test_voyage": "Voyage AI", + "test_xai": "xAI", + "test_nvidia": "NVIDIA", + "test_watsonx": "IBM watsonx", + "test_azure_ai": "Azure AI", + "test_snowflake": "Snowflake", + "test_infinity": "Infinity", + "test_jina": "Jina AI", + "test_deepgram": "Deepgram", + "test_clarifai": "Clarifai", + "test_triton": "Triton", } - + for key, provider in provider_mapping.items(): if key in test_file: return provider - + # For cross-provider test files - if any(name in test_file for name in ['test_optional_params', 'test_prompt_factory', - 'test_router', 'test_text_completion']): - return f'Cross-Provider Tests ({test_file})' - - return 'Other Tests' + if any( + name in test_file + for name in [ + "test_optional_params", + "test_prompt_factory", + "test_router", + "test_text_completion", + ] + ): + return f"Cross-Provider Tests ({test_file})" + + return "Other Tests" + def format_duration(seconds: float) -> str: """Format duration in human-readable format""" @@ -89,290 +100,355 @@ def format_duration(seconds: float) -> str: return f"{hours}h {minutes}m" -def generate_markdown_report(junit_xml_path: str, output_path: str, tag: str = None, commit: str = None): +def generate_markdown_report( + junit_xml_path: str, output_path: str, tag: str = None, commit: str = None +): """Generate a beautiful markdown report from JUnit XML""" try: tree = ET.parse(junit_xml_path) root = tree.getroot() - + # Handle both testsuite and testsuites root - if root.tag == 'testsuites': - suites = root.findall('testsuite') + if root.tag == "testsuites": + suites = root.findall("testsuite") else: suites = [root] - + # Overall statistics total_tests = 0 total_failures = 0 total_errors = 0 total_skipped = 0 total_time = 0.0 - + # Provider breakdown - provider_stats = defaultdict(lambda: {'passed': 0, 'failed': 0, 'skipped': 0, 'errors': 0, 'time': 0.0}) + provider_stats = defaultdict( + lambda: {"passed": 0, "failed": 0, "skipped": 0, "errors": 0, "time": 0.0} + ) provider_tests = defaultdict(list) - + for suite in suites: - total_tests += int(suite.get('tests', 0)) - total_failures += int(suite.get('failures', 0)) - total_errors += int(suite.get('errors', 0)) - total_skipped += int(suite.get('skipped', 0)) - total_time += float(suite.get('time', 0)) - - for testcase in suite.findall('testcase'): - classname = testcase.get('classname', '') - test_name = testcase.get('name', '') - test_time = float(testcase.get('time', 0)) - + total_tests += int(suite.get("tests", 0)) + total_failures += int(suite.get("failures", 0)) + total_errors += int(suite.get("errors", 0)) + total_skipped += int(suite.get("skipped", 0)) + total_time += float(suite.get("time", 0)) + + for testcase in suite.findall("testcase"): + classname = testcase.get("classname", "") + test_name = testcase.get("name", "") + test_time = float(testcase.get("time", 0)) + # Extract test file name from classname - if '.' in classname: - parts = classname.split('.') - test_file = parts[-2] if len(parts) > 1 else 'unknown' + if "." in classname: + parts = classname.split(".") + test_file = parts[-2] if len(parts) > 1 else "unknown" else: - test_file = 'unknown' - + test_file = "unknown" + provider = get_provider_from_test_file(test_file) - provider_stats[provider]['time'] += test_time - + provider_stats[provider]["time"] += test_time + # Check test status - if testcase.find('failure') is not None: - provider_stats[provider]['failed'] += 1 - failure = testcase.find('failure') - failure_msg = failure.get('message', '') if failure is not None else '' - provider_tests[provider].append({ - 'name': test_name, - 'status': 'FAILED', - 'time': test_time, - 'message': failure_msg - }) - elif testcase.find('error') is not None: - provider_stats[provider]['errors'] += 1 - error = testcase.find('error') - error_msg = error.get('message', '') if error is not None else '' - provider_tests[provider].append({ - 'name': test_name, - 'status': 'ERROR', - 'time': test_time, - 'message': error_msg - }) - elif testcase.find('skipped') is not None: - provider_stats[provider]['skipped'] += 1 - skip = testcase.find('skipped') - skip_msg = skip.get('message', '') if skip is not None else '' - provider_tests[provider].append({ - 'name': test_name, - 'status': 'SKIPPED', - 'time': test_time, - 'message': skip_msg - }) + if testcase.find("failure") is not None: + provider_stats[provider]["failed"] += 1 + failure = testcase.find("failure") + failure_msg = ( + failure.get("message", "") if failure is not None else "" + ) + provider_tests[provider].append( + { + "name": test_name, + "status": "FAILED", + "time": test_time, + "message": failure_msg, + } + ) + elif testcase.find("error") is not None: + provider_stats[provider]["errors"] += 1 + error = testcase.find("error") + error_msg = error.get("message", "") if error is not None else "" + provider_tests[provider].append( + { + "name": test_name, + "status": "ERROR", + "time": test_time, + "message": error_msg, + } + ) + elif testcase.find("skipped") is not None: + provider_stats[provider]["skipped"] += 1 + skip = testcase.find("skipped") + skip_msg = skip.get("message", "") if skip is not None else "" + provider_tests[provider].append( + { + "name": test_name, + "status": "SKIPPED", + "time": test_time, + "message": skip_msg, + } + ) else: - provider_stats[provider]['passed'] += 1 - provider_tests[provider].append({ - 'name': test_name, - 'status': 'PASSED', - 'time': test_time, - 'message': '' - }) - + provider_stats[provider]["passed"] += 1 + provider_tests[provider].append( + { + "name": test_name, + "status": "PASSED", + "time": test_time, + "message": "", + } + ) + passed = total_tests - total_failures - total_errors - total_skipped - + # Generate the markdown report - with open(output_path, 'w') as f: + with open(output_path, "w") as f: # Header f.write("# LLM Translation Test Results\n\n") - + # Metadata table f.write("## Test Run Information\n\n") f.write("| Field | Value |\n") f.write("|-------|-------|\n") f.write(f"| **Tag** | `{tag or 'N/A'}` |\n") - f.write(f"| **Date** | {datetime.utcnow().strftime('%Y-%m-%d %H:%M:%S UTC')} |\n") + f.write( + f"| **Date** | {datetime.utcnow().strftime('%Y-%m-%d %H:%M:%S UTC')} |\n" + ) f.write(f"| **Commit** | `{commit or 'N/A'}` |\n") f.write(f"| **Duration** | {format_duration(total_time)} |\n") f.write("\n") - + # Overall statistics with visual elements f.write("## Overall Statistics\n\n") - + # Summary box f.write("```\n") f.write(f"Total Tests: {total_tests}\n") - f.write(f"├── Passed: {passed:>4} ({(passed/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n") - f.write(f"├── Failed: {total_failures:>4} ({(total_failures/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n") - f.write(f"├── Errors: {total_errors:>4} ({(total_errors/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n") - f.write(f"└── Skipped: {total_skipped:>4} ({(total_skipped/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n") + f.write( + f"├── Passed: {passed:>4} ({(passed/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n" + ) + f.write( + f"├── Failed: {total_failures:>4} ({(total_failures/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n" + ) + f.write( + f"├── Errors: {total_errors:>4} ({(total_errors/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n" + ) + f.write( + f"└── Skipped: {total_skipped:>4} ({(total_skipped/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n" + ) f.write("```\n\n") - - + # Provider summary table f.write("## Results by Provider\n\n") - f.write("| Provider | Total | Pass | Fail | Error | Skip | Pass Rate | Duration |\n") - f.write("|----------|-------|------|------|-------|------|-----------|----------|") - + f.write( + "| Provider | Total | Pass | Fail | Error | Skip | Pass Rate | Duration |\n" + ) + f.write( + "|----------|-------|------|------|-------|------|-----------|----------|" + ) + # Sort providers: specific providers first, then cross-provider tests sorted_providers = [] cross_provider = [] for p in sorted(provider_stats.keys()): - if 'Cross-Provider' in p or p == 'Other Tests': + if "Cross-Provider" in p or p == "Other Tests": cross_provider.append(p) else: sorted_providers.append(p) - + all_providers = sorted_providers + cross_provider - + for provider in all_providers: stats = provider_stats[provider] - total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped'] - pass_rate = (stats['passed'] / total * 100) if total > 0 else 0 - - f.write(f"\n| {provider} | {total} | {stats['passed']} | {stats['failed']} | ") + total = ( + stats["passed"] + + stats["failed"] + + stats["errors"] + + stats["skipped"] + ) + pass_rate = (stats["passed"] / total * 100) if total > 0 else 0 + + f.write( + f"\n| {provider} | {total} | {stats['passed']} | {stats['failed']} | " + ) f.write(f"{stats['errors']} | {stats['skipped']} | {pass_rate:.1f}% | ") f.write(f"{format_duration(stats['time'])} |") - + # Detailed test results by provider f.write("\n\n## Detailed Test Results\n\n") - + for provider in sorted_providers: if provider_tests[provider]: stats = provider_stats[provider] - total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped'] - + total = ( + stats["passed"] + + stats["failed"] + + stats["errors"] + + stats["skipped"] + ) + f.write(f"### {provider}\n\n") f.write(f"**Summary:** {stats['passed']}/{total} passed ") - f.write(f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%) ") + f.write( + f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%) " + ) f.write(f"in {format_duration(stats['time'])}\n\n") - + # Group tests by status tests_by_status = defaultdict(list) for test in provider_tests[provider]: - tests_by_status[test['status']].append(test) - + tests_by_status[test["status"]].append(test) + # Show failed tests first (if any) - if tests_by_status['FAILED']: + if tests_by_status["FAILED"]: f.write("
\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 +457,75 @@ def run_tests(test_path: str = "tests/llm_translation/", print_colored(f" Errors: {errors}", Colors.RED) if skipped > 0: print_colored(f" Skipped: {skipped}", Colors.YELLOW) - + if total > 0: pass_rate = (passed / total) * 100 - color = Colors.GREEN if pass_rate >= 80 else Colors.YELLOW if pass_rate >= 60 else Colors.RED + color = ( + Colors.GREEN + if pass_rate >= 80 + else Colors.YELLOW if pass_rate >= 60 else Colors.RED + ) print_colored(f" Pass Rate: {pass_rate:.1f}%", color) else: print_colored("No test results found!", Colors.RED) - + print() print_colored("Test run complete!", Colors.BOLD + Colors.GREEN) - + return result.returncode + if __name__ == "__main__": import argparse - + parser = argparse.ArgumentParser(description="Run LLM Translation Tests") - parser.add_argument("--test-path", default="tests/llm_translation/", - help="Path to test directory") - parser.add_argument("--junit-xml", default="test-results/junit.xml", - help="Path for JUnit XML output") - parser.add_argument("--report", default="test-results/llm_translation_report.md", - help="Path for markdown report") + parser.add_argument( + "--test-path", default="tests/llm_translation/", help="Path to test directory" + ) + parser.add_argument( + "--junit-xml", + default="test-results/junit.xml", + help="Path for JUnit XML output", + ) + parser.add_argument( + "--report", + default="test-results/llm_translation_report.md", + help="Path for markdown report", + ) parser.add_argument("--tag", help="Git tag or version") parser.add_argument("--commit", help="Git commit SHA") - + args = parser.parse_args() - + # Get git info if not provided if not args.commit: try: - result = subprocess.run(["git", "rev-parse", "HEAD"], - capture_output=True, text=True) + result = subprocess.run( + ["git", "rev-parse", "HEAD"], capture_output=True, text=True + ) if result.returncode == 0: args.commit = result.stdout.strip() except: pass - + if not args.tag: try: - result = subprocess.run(["git", "describe", "--tags", "--abbrev=0"], - capture_output=True, text=True) + result = subprocess.run( + ["git", "describe", "--tags", "--abbrev=0"], + capture_output=True, + text=True, + ) if result.returncode == 0: args.tag = result.stdout.strip() except: pass - + exit_code = run_tests( test_path=args.test_path, junit_xml=args.junit_xml, report_path=args.report, tag=args.tag, - commit=args.commit + commit=args.commit, ) - + sys.exit(exit_code) 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/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..2dc158a3cfb 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py @@ -79,4 +79,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..4f2eaa3c468 100644 --- a/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py @@ -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/auth/custom_sso_handler.py b/enterprise/litellm_enterprise/proxy/auth/custom_sso_handler.py index a3682320387..1c74ca3c49c 100644 --- a/enterprise/litellm_enterprise/proxy/auth/custom_sso_handler.py +++ b/enterprise/litellm_enterprise/proxy/auth/custom_sso_handler.py @@ -26,12 +26,12 @@ from litellm.proxy.management_endpoints.types import CustomOpenID class EnterpriseCustomSSOHandler: """ Enterprise Custom SSO Handler for LiteLLM Proxy - + This class provides methods for handling custom SSO authentication flows where users can implement their own authentication logic by processing request headers and returning user information in OpenID format. """ - + @staticmethod async def handle_custom_ui_sso_sign_in( request: Request, @@ -40,16 +40,16 @@ class EnterpriseCustomSSOHandler: Allow a user to execute their custom code to parse incoming request headers and return a OpenID object Use this when you have an OAuth proxy in front of LiteLLM (where the OAuth proxy has already authenticated the user) - + Args: request: The FastAPI request object containing headers and other request data - + Returns: RedirectResponse: Redirect response that sends the user to the LiteLLM UI with authentication token - + Raises: ValueError: If custom_ui_sso_sign_in_handler is not configured - + Example: This method is typically called when a user has already been authenticated by an external OAuth proxy and the proxy has added custom headers containing user information. @@ -63,24 +63,31 @@ class EnterpriseCustomSSOHandler: premium_user, user_custom_ui_sso_sign_in_handler, ) + if premium_user is not True: raise ValueError(CommonProxyErrors.not_premium_user.value) - + if user_custom_ui_sso_sign_in_handler is None: - raise ValueError("custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings.") - - custom_sso_login_handler = cast(CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler) - openid_response: OpenID = await custom_sso_login_handler.handle_custom_ui_sso_sign_in( - request=request, + raise ValueError( + "custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings." + ) + + custom_sso_login_handler = cast( + CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler ) - + openid_response: OpenID = ( + await custom_sso_login_handler.handle_custom_ui_sso_sign_in( + request=request, + ) + ) + # Import here to avoid circular imports from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + return await SSOAuthenticationHandler.get_redirect_response_from_openid( result=openid_response, request=request, received_response=None, generic_client_id=None, ui_access_mode=None, - ) \ No newline at end of file + ) 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 356f6ecd4b5..3aef2bd77cb 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,7 +328,11 @@ class CheckBatchCost: # Pass deployment model_info so custom batch pricing # (input_cost_per_token_batches etc.) is used for cost calc - deployment_model_info = deployment_info.model_info.model_dump() if deployment_info.model_info else {} + deployment_model_info = ( + deployment_info.model_info.model_dump() + if deployment_info.model_info + else {} + ) batch_cost, batch_usage, batch_models = ( await calculate_batch_cost_and_usage( file_content_dictionary=file_content_as_dict, @@ -349,7 +379,9 @@ class CheckBatchCost: # Record batch duration (completed_at - created_at) if prom_logger and response.completed_at and response.created_at: - duration_seconds = float(response.completed_at - response.created_at) + duration_seconds = float( + response.completed_at - response.created_at + ) if duration_seconds >= 0: prom_logger.record_managed_batch_duration( duration_seconds=duration_seconds, @@ -358,7 +390,9 @@ class CheckBatchCost: ) # Track this job for the final metrics summary - processed_models.append((model_name, str(llm_provider) if llm_provider else None)) + processed_models.append( + (model_name, str(llm_provider) if llm_provider else None) + ) # mark the job as complete try: 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_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 60c564072a0..2e7112c4c76 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -125,7 +125,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): db_data["storage_backend"] = hidden_params["storage_backend"] if "storage_url" in hidden_params: db_data["storage_url"] = hidden_params["storage_url"] - + verbose_logger.debug( f"Storage metadata: storage_backend={db_data.get('storage_backend')}, " f"storage_url={db_data.get('storage_url')}" @@ -285,28 +285,28 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): raise Exception( "Filtering by 'target_model_names' is not supported when using managed batches." ) - + where_clause: Dict[str, Any] = {"file_purpose": "batch"} - + # Filter by user who created the batch if user_api_key_dict.user_id: where_clause["created_by"] = user_api_key_dict.user_id - + if after: where_clause["id"] = {"gt": after} - + # Fetch more than needed to allow for post-fetch filtering fetch_limit = limit or 20 if target_model_names: # Fetch extra to account for filtering fetch_limit = max(fetch_limit * 3, 100) - + batches = await self.prisma_client.db.litellm_managedobjecttable.find_many( where=where_clause, take=fetch_limit, order={"created_at": "desc"}, ) - + batch_objects: List[LiteLLMBatch] = [] for batch in batches: try: @@ -314,7 +314,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if len(batch_objects) >= (limit or 20): break - batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object + batch_data = ( + json.loads(batch.file_object) + if isinstance(batch.file_object, str) + else batch.file_object + ) batch_obj = LiteLLMBatch(**batch_data) batch_obj.id = batch.unified_object_id batch_objects.append(batch_obj) @@ -324,7 +328,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): f"Failed to parse batch object {batch.unified_object_id}: {e}" ) continue - + return { "object": "list", "data": batch_objects, @@ -377,11 +381,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): """ Check if the user has access to a list of file IDs. Only checks managed (unified) file IDs. - + Args: file_ids: List of file IDs to check access for user_api_key_dict: User API key authentication details - + Raises: HTTPException: If user doesn't have access to any of the files """ @@ -419,10 +423,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ### HANDLE TRANSFORMATIONS ### # Check both completion and acompletion call types is_completion_call = ( - call_type == CallTypes.completion.value + call_type == CallTypes.completion.value or call_type == CallTypes.acompletion.value ) - + if is_completion_call: messages = data.get("messages") model = data.get("model", "") @@ -431,22 +435,27 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if file_ids: # Check user has access to all managed files await self.check_file_ids_access(file_ids, user_api_key_dict) - + # Check if any files are stored in storage backends and need base64 conversion # This is needed for Vertex AI/Gemini which requires base64 content - is_vertex_ai = model and ("vertex_ai" in model or "gemini" in model.lower()) + is_vertex_ai = model and ( + "vertex_ai" in model or "gemini" in model.lower() + ) if is_vertex_ai: await self._convert_storage_files_to_base64( messages=messages, file_ids=file_ids, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, ) - + model_file_id_mapping = await self.get_model_file_id_mapping( file_ids, user_api_key_dict.parent_otel_span ) data["model_file_id_mapping"] = model_file_id_mapping - elif call_type == CallTypes.aresponses.value or call_type == CallTypes.responses.value: + elif ( + call_type == CallTypes.aresponses.value + or call_type == CallTypes.responses.value + ): # Handle managed files in responses API input and tools file_ids = [] @@ -611,7 +620,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if model_id is None: model_id = cast( Optional[str], - kwargs.get("litellm_metadata", {}).get("model_info", {}).get("id", None), + kwargs.get("litellm_metadata", {}) + .get("model_info", {}) + .get("id", None), ) mapped_file_id: Optional[str] = None if input_file_id and model_file_id_mapping and model_id: @@ -648,7 +659,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) -> List[str]: """ Gets file ids from responses API input. - + The input can be: - A string (no files) - A list of input items, where each item can have: @@ -656,32 +667,35 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): - content: a list that can contain items with type: "input_file" and file_id """ file_ids: List[str] = [] - + if isinstance(input, str): return file_ids - + if not isinstance(input, list): return file_ids - + for item in input: if not isinstance(item, dict): continue - + # Check for direct input_file type if item.get("type") == "input_file": file_id = item.get("file_id") if file_id: file_ids.append(file_id) - + # Check for input_file in content array content = item.get("content") if isinstance(content, list): for content_item in content: - if isinstance(content_item, dict) and content_item.get("type") == "input_file": + if ( + isinstance(content_item, dict) + and content_item.get("type") == "input_file" + ): file_id = content_item.get("file_id") if file_id: file_ids.append(file_id) - + return file_ids def get_file_ids_from_responses_tools( @@ -689,7 +703,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) -> List[str]: """ Gets file ids from responses API tools parameter. - + The tools can contain code_interpreter with container.file_ids: [ { @@ -699,14 +713,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ] """ file_ids: List[str] = [] - + if not isinstance(tools, list): return file_ids - + for tool in tools: if not isinstance(tool, dict): continue - + # Check for code_interpreter with container file_ids if tool.get("type") == "code_interpreter": container = tool.get("container") @@ -716,7 +730,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): for file_id in container_file_ids: if isinstance(file_id, str): file_ids.append(file_id) - + return file_ids def get_vector_store_ids_from_file_search_tools( @@ -916,10 +930,17 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # Emit Prometheus metrics for managed file creation prom_logger = self._get_prometheus_logger() if prom_logger: - first_model = target_model_names_list[0] if target_model_names_list else None + first_model = ( + target_model_names_list[0] if target_model_names_list else None + ) first_provider = "" if responses: - first_provider = getattr(responses[0], "_hidden_params", {}).get("custom_llm_provider") or "" + first_provider = ( + getattr(responses[0], "_hidden_params", {}).get( + "custom_llm_provider" + ) + or "" + ) prom_logger.record_managed_file_created( model=first_model or "", api_provider=first_provider, @@ -1073,16 +1094,24 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): model_name=resolved_model_name, ) setattr(response, file_attr, unified_file_id) - + # Use llm_router credentials when available. Without credentials, # Azure and other auth-required providers return 500/401. file_object = None try: # Import module and use getattr for better testability with mocks import litellm.proxy.proxy_server as proxy_server_module - _llm_router = getattr(proxy_server_module, 'llm_router', None) + + _llm_router = getattr( + proxy_server_module, "llm_router", None + ) if _llm_router is not None and model_id: - _creds = _llm_router.get_deployment_credentials_with_provider(model_id) or {} + _creds = ( + _llm_router.get_deployment_credentials_with_provider( + model_id + ) + or {} + ) file_object = await litellm.afile_retrieve( file_id=original_file_id, **_creds, @@ -1099,7 +1128,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): verbose_logger.warning( f"Failed to retrieve file object for {file_attr}={original_file_id}: {str(e)}. Storing with None and will fetch on-demand." ) - + await self.store_unified_file_id( file_id=unified_file_id, file_object=file_object, @@ -1128,6 +1157,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): from litellm.litellm_core_utils.get_llm_provider_logic import ( get_llm_provider, ) + _, batch_provider, _, _ = get_llm_provider(model=model_name) except Exception: if "/" in model_name: @@ -1199,7 +1229,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # Case 1 : This is not a managed file if not stored_file_object: raise Exception(f"LiteLLM Managed File object with id={file_id} not found") - + # Case 2: Managed file and the file object exists in the database # The stored file_object has the raw provider ID. Replace with the unified ID # so callers see a consistent ID (matching Case 3 which does response.id = file_id). @@ -1217,13 +1247,21 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) try: - model_id, model_file_id = next(iter(stored_file_object.model_mappings.items())) - credentials = llm_router.get_deployment_credentials_with_provider(model_id) or {} - response = await litellm.afile_retrieve(file_id=model_file_id, **credentials) + model_id, model_file_id = next( + iter(stored_file_object.model_mappings.items()) + ) + credentials = ( + llm_router.get_deployment_credentials_with_provider(model_id) or {} + ) + response = await litellm.afile_retrieve( + file_id=model_file_id, **credentials + ) response.id = file_id # Replace with unified ID return response except Exception as e: - raise Exception(f"Failed to retrieve file {file_id} from provider: {str(e)}") from e + raise Exception( + f"Failed to retrieve file {file_id} from provider: {str(e)}" + ) from e async def afile_list( self, @@ -1245,19 +1283,19 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): import litellm.proxy.proxy_server as proxy_server_module # Check if the scheduler has the batch cost checking job registered - scheduler = getattr(proxy_server_module, 'scheduler', None) + scheduler = getattr(proxy_server_module, "scheduler", None) if scheduler is None: return False - + # Check if the check_batch_cost_job exists in the scheduler try: - job = scheduler.get_job('check_batch_cost_job') + job = scheduler.get_job("check_batch_cost_job") if job is not None: return True except Exception: # Job not found or scheduler doesn't support get_job pass - + return False except Exception as e: verbose_logger.warning( @@ -1265,28 +1303,26 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) return False - async def _get_batches_referencing_file( - self, file_id: str - ) -> List[Dict[str, Any]]: + async def _get_batches_referencing_file(self, file_id: str) -> List[Dict[str, Any]]: """ Find batches that reference this file and still need cost tracking. Find batches that are in non-terminal state and have not yet been processed by CheckBatchCost. Args: file_id: The unified file ID to check - + Returns: List of batch objects referencing this file in non-terminal state (max 10 for error message display) """ # Prepare list of file IDs to check (both unified and provider IDs) file_ids_to_check = [file_id] - + # Get model-specific file IDs for this unified file ID if it's a managed file try: model_file_id_mapping = await self.get_model_file_id_mapping( [file_id], litellm_parent_otel_span=None ) - + if model_file_id_mapping and file_id in model_file_id_mapping: # Add all provider file IDs for this unified file provider_file_ids = list(model_file_id_mapping[file_id].values()) @@ -1296,59 +1332,67 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): f"Could not get model file ID mapping for {file_id}: {e}. " f"Will only check unified file ID." ) - MAX_MATCHES_TO_RETURN = 10 - + MAX_MATCHES_TO_RETURN = 10 + batches = await self.prisma_client.db.litellm_managedobjecttable.find_many( where={ "file_purpose": "batch", "batch_processed": False, - "status": {"not_in": ["failed", "expired", "cancelled"]} + "status": {"not_in": ["failed", "expired", "cancelled"]}, }, take=MAX_MATCHES_TO_RETURN, order={"created_at": "desc"}, ) - + referencing_batches = [] for batch in batches: try: # Parse the batch file_object to check for file references - batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object - + batch_data = ( + json.loads(batch.file_object) + if isinstance(batch.file_object, str) + else batch.file_object + ) + # Extract file IDs from batch # Batches typically reference the unified file ID in input_file_id # Output and error files are generated by the provider input_file_id = batch_data.get("input_file_id") output_file_id = batch_data.get("output_file_id") error_file_id = batch_data.get("error_file_id") - - referenced_file_ids = [fid for fid in [input_file_id, output_file_id, error_file_id] if fid] - + + referenced_file_ids = [ + fid for fid in [input_file_id, output_file_id, error_file_id] if fid + ] + # Check if any referenced file ID matches the file we're trying to delete if any(ref_id in file_ids_to_check for ref_id in referenced_file_ids): - referencing_batches.append({ - "batch_id": batch.unified_object_id, - "status": batch.status, - "created_at": batch.created_at, - }) + referencing_batches.append( + { + "batch_id": batch.unified_object_id, + "status": batch.status, + "created_at": batch.created_at, + } + ) except Exception as e: verbose_logger.warning( f"Error parsing batch object {batch.unified_object_id}: {e}" ) continue - + return referencing_batches async def _check_file_deletion_allowed(self, file_id: str) -> None: """ Check if file deletion should be blocked due to batch references. - + Blocks deletion if: 1. File is referenced by any batch in non-terminal state, AND 2. Batch polling is configured (user wants cost tracking) - + Args: file_id: The unified file ID to check - + Raises: HTTPException: If file deletion should be blocked """ @@ -1356,39 +1400,45 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if not self._is_batch_polling_enabled(): # Batch polling not configured, allow deletion return - + # Check if file is referenced by any non-terminal batches referencing_batches = await self._get_batches_referencing_file(file_id) - + if referencing_batches: # File is referenced by non-terminal batches and polling is enabled - MAX_BATCHES_IN_ERROR = 5 # Limit batches shown in error message for readability - + MAX_BATCHES_IN_ERROR = ( + 5 # Limit batches shown in error message for readability + ) + # Show up to MAX_BATCHES_IN_ERROR in the error message batches_to_show = referencing_batches[:MAX_BATCHES_IN_ERROR] - batch_statuses = [f"{b['batch_id']}: {b['status']}" for b in batches_to_show] - + batch_statuses = [ + f"{b['batch_id']}: {b['status']}" for b in batches_to_show + ] + # Determine the count message count_message = f"{len(referencing_batches)}" - if len(referencing_batches) >= 10: # MAX_MATCHES_TO_RETURN from _get_batches_referencing_file + if ( + len(referencing_batches) >= 10 + ): # MAX_MATCHES_TO_RETURN from _get_batches_referencing_file count_message = "10+" - + error_message = ( f"Cannot delete file {file_id}. " f"The file is referenced by {count_message} batch(es) in non-terminal state" ) - + # Add specific batch details if not too many if len(referencing_batches) <= MAX_BATCHES_IN_ERROR: error_message += f": {', '.join(batch_statuses)}. " else: error_message += f" (showing {MAX_BATCHES_IN_ERROR} most recent): {', '.join(batch_statuses)}. " - + error_message += ( f"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. " f"Alternatively, wait for all batches to complete and for cost to be computed (batch_processed=true)." ) - + # Record blocked deletion metric prom_logger = self._get_prometheus_logger() if prom_logger: @@ -1419,7 +1469,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): specific_model_file_id_mapping = model_file_id_mapping.get(file_id) if specific_model_file_id_mapping: # Remove conflicting keys from data to avoid duplicate keyword arguments - filtered_data = {k: v for k, v in data.items() if k not in ("model", "file_id")} + filtered_data = { + k: v for k, v in data.items() if k not in ("model", "file_id") + } for model_id, model_file_id in specific_model_file_id_mapping.items(): delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **filtered_data) # type: ignore @@ -1480,7 +1532,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) -> None: """ Convert files stored in storage backends to base64 format for Vertex AI/Gemini. - + This method checks if any managed files are stored in storage backends, downloads them, and converts them to base64 format in the messages. """ @@ -1488,29 +1540,29 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): for file_id in file_ids: # Check if this is a base64 encoded unified file ID decoded_unified_file_id = _is_base64_encoded_unified_file_id(file_id) - + if not decoded_unified_file_id: continue - + # Check database for storage backend info # IMPORTANT: The database stores the base64 encoded unified_file_id (not the decoded version) # So we query with the original file_id (which is base64 encoded) db_file = await self.prisma_client.db.litellm_managedfiletable.find_first( where={"unified_file_id": file_id} ) - + if not db_file or not db_file.storage_backend or not db_file.storage_url: continue - + # File is stored in a storage backend, download and convert to base64 try: from litellm.llms.base_llm.files.storage_backend_factory import ( get_storage_backend, ) - + storage_backend_name = db_file.storage_backend storage_url = db_file.storage_url - + # Get storage backend (uses same env vars as callback) try: storage_backend = get_storage_backend(storage_backend_name) @@ -1519,18 +1571,22 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): f"Storage backend '{storage_backend_name}' error for file {file_id}: {str(e)}" ) continue - + file_content = await storage_backend.download_file(storage_url) - + # Determine content type from file object - content_type = self._get_content_type_from_file_object(db_file.file_object) - + content_type = self._get_content_type_from_file_object( + db_file.file_object + ) + # Convert to base64 base64_data = base64.b64encode(file_content).decode("utf-8") base64_data_uri = f"data:{content_type};base64,{base64_data}" - + # Update messages to use base64 instead of file_id - self._update_messages_with_base64_data(messages, file_id, base64_data_uri, content_type) + self._update_messages_with_base64_data( + messages, file_id, base64_data_uri, content_type + ) except Exception as e: verbose_logger.exception( f"Error converting file {file_id} from storage backend to base64: {str(e)}" @@ -1541,21 +1597,21 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): def _get_content_type_from_file_object(self, file_object: Optional[Any]) -> str: """ Determine content type from file object. - + Uses the MIME type utility for consistent detection and normalization. - + Args: file_object: The file object from the database (can be dict, JSON string, or None) - + Returns: str: MIME type (defaults to "application/octet-stream" if cannot be determined) """ # Use utility function for detection content_type = get_content_type_from_file_object(file_object) - + # Normalize for Gemini/Vertex AI (requires image/jpeg, not image/jpg) content_type = normalize_mime_type_for_provider(content_type, provider="gemini") - + return content_type def _update_messages_with_base64_data( @@ -1567,7 +1623,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) -> None: """ Update messages to replace file_id with base64 data URI. - + Args: messages: List of messages to update file_id: The file ID to replace @@ -1582,7 +1638,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if element.get("type") == "file": file_element = cast(ChatCompletionFileObject, element) file_element_file = file_element.get("file", {}) - + if file_element_file.get("file_id") == file_id: # Replace file_id with base64 data file_element_file["file_data"] = base64_data_uri @@ -1590,7 +1646,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): file_element_file["format"] = content_type # Remove file_id to ensure only file_data is used file_element_file.pop("file_id", None) - + verbose_logger.debug( f"Converted file {file_id} from storage backend to base64 with format {content_type}" ) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py index 254d816039c..70634537c55 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py @@ -41,7 +41,7 @@ class _PROXY_LiteLLMManagedVectorStores( ): """ Managed vector stores with target_model_names support. - + This class provides functionality to: - Create vector stores across multiple models - Retrieve vector stores by unified ID @@ -77,14 +77,14 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> str: """ Generate the format string for the unified vector store ID. - + Format: litellm_proxy:vector_store;unified_id,;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 +109,18 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> VectorStoreCreateResponse: """ Create a vector store for a specific model. - + Args: llm_router: LiteLLM router instance model: Model name to create vector store for request_data: Request data for vector store creation litellm_parent_otel_span: OpenTelemetry span for tracing - + Returns: VectorStoreCreateResponse from the provider """ # Use the router to create the vector store - response = await llm_router.avector_store_create( - model=model, **request_data - ) + response = await llm_router.avector_store_create(model=model, **request_data) return response # ============================================================================ @@ -139,14 +137,14 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> VectorStoreCreateResponse: """ Create a vector store across multiple models. - + Args: create_request: Vector store creation request parameters llm_router: LiteLLM router instance target_model_names_list: List of target model names litellm_parent_otel_span: OpenTelemetry span for tracing user_api_key_dict: User API key authentication details - + Returns: VectorStoreCreateResponse with unified ID """ @@ -196,7 +194,7 @@ class _PROXY_LiteLLMManagedVectorStores( # VectorStoreCreateResponse is a TypedDict, so we need to create a new dict with the unified ID response = responses[0].copy() response["id"] = unified_id - + verbose_logger.info( f"Successfully created managed vector store with unified ID: {unified_id}" ) @@ -212,13 +210,13 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> Dict[str, Any]: """ List vector stores created by a user. - + Args: user_api_key_dict: User API key authentication details limit: Maximum number of vector stores to return after: Cursor for pagination order: Sort order ('asc' or 'desc') - + Returns: Dictionary with list of vector stores and pagination info """ @@ -238,23 +236,23 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> bool: """ Check if user has access to a vector store. - + Args: vector_store_id: The unified vector store ID user_api_key_dict: User API key authentication details - + Returns: True if user has access, False otherwise """ is_unified_id = is_base64_encoded_unified_id(vector_store_id) - + if is_unified_id: # Check access for managed vector store return await self.can_user_access_unified_resource_id( vector_store_id, user_api_key_dict, ) - + # Not a managed vector store, allow access return True @@ -263,24 +261,22 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> bool: """ Check if user has access to a managed vector store in request data. - + Args: data: Request data containing vector_store_id user_api_key_dict: User API key authentication details - + Returns: True if this is a managed vector store and user has access - + Raises: HTTPException: If user doesn't have access """ vector_store_id = cast(Optional[str], data.get("vector_store_id")) is_unified_id = ( - is_base64_encoded_unified_id(vector_store_id) - if vector_store_id - else False + is_base64_encoded_unified_id(vector_store_id) if vector_store_id else False ) - + if is_unified_id and vector_store_id: if await self.can_user_access_unified_resource_id( vector_store_id, user_api_key_dict @@ -291,7 +287,7 @@ class _PROXY_LiteLLMManagedVectorStores( status_code=403, detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}", ) - + return False # ============================================================================ @@ -307,18 +303,18 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> Union[Exception, str, Dict, None]: """ Pre-call hook to handle vector store operations. - + This hook intercepts vector store requests and: - Validates access for managed vector stores - Transforms unified IDs to provider-specific IDs - Adds model routing information - + Args: user_api_key_dict: User API key authentication details cache: Cache instance data: Request data call_type: Type of call being made - + Returns: Modified request data or None """ @@ -330,40 +326,40 @@ class _PROXY_LiteLLMManagedVectorStores( # Handle vector store search operations if call_type == "avector_store_search": vector_store_id = data.get("vector_store_id") - + if vector_store_id: # Check if it's a managed vector store ID decoded_id = is_base64_encoded_unified_id(vector_store_id) - + if decoded_id: verbose_logger.debug( f"Processing managed vector store search: {vector_store_id}" ) - + # Check access has_access = await self.can_user_access_unified_resource_id( vector_store_id, user_api_key_dict ) - + if not has_access: raise HTTPException( status_code=403, detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}", ) - + # Parse the unified ID to extract components parsed_id = parse_unified_id(vector_store_id) - + if parsed_id: # Extract the model ID and provider resource ID model_id = parsed_id.get("model_id") provider_resource_id = parsed_id.get("provider_resource_id") target_model_names = parsed_id.get("target_model_names", []) - + verbose_logger.debug( f"Decoded vector store - model_id: {model_id}, provider_resource_id: {provider_resource_id}, target_model_names: {target_model_names}" ) - + # Determine which model to use for routing # Priority: model_id (deployment ID) > first target_model_name routing_model = None @@ -371,28 +367,28 @@ class _PROXY_LiteLLMManagedVectorStores( routing_model = model_id elif target_model_names and len(target_model_names) > 0: routing_model = target_model_names[0] - + # Set the model for routing if routing_model: data["model"] = routing_model verbose_logger.info( f"Routing vector store search to model: {routing_model}" ) - + # Replace the unified ID with the provider-specific ID if provider_resource_id: data["vector_store_id"] = provider_resource_id verbose_logger.debug( f"Replaced unified ID with provider resource ID: {provider_resource_id}" ) - + # Handle vector store retrieve/delete operations elif call_type in ("avector_store_retrieve", "avector_store_delete"): await self.check_managed_vector_store_access(data, user_api_key_dict) - + # If it's a managed vector store, we'll handle it in the endpoint # No need to transform here as the endpoint will route to the hook - + return data # ============================================================================ @@ -407,15 +403,15 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> Any: """ Post-call hook to transform responses. - + This hook can be used to transform responses if needed. For now, it just passes through the response. - + Args: data: Request data user_api_key_dict: User API key authentication details response: Response from the provider - + Returns: Potentially modified response """ @@ -436,21 +432,21 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> List[Dict]: """ Filter deployments based on vector store availability. - + This is used by the router to select only deployments that have the vector store available. - + Note: This method signature is a compromise between CustomLogger and BaseManagedResource parent classes which have incompatible signatures. The type: ignore[override] is necessary due to this multiple inheritance conflict. - + Args: model: Model name healthy_deployments: List of healthy deployments messages: Messages (unused for vector stores, required by CustomLogger interface) request_kwargs: Request kwargs containing vector_store_id and mappings parent_otel_span: OpenTelemetry span for tracing - + Returns: Filtered list of deployments """ 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..cf8c38719d7 100644 --- a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py +++ b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py @@ -147,12 +147,12 @@ async def list_vector_stores( vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db( prisma_client=prisma_client ) - + # Also clean up in-memory registry to remove any deleted vector stores if litellm.vector_store_registry is not None: db_vector_store_ids = { - vs.get("vector_store_id") - for vs in vector_stores_from_db + vs.get("vector_store_id") + for vs in vector_stores_from_db if vs.get("vector_store_id") } # Remove any in-memory vector stores that no longer exist in database 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/litellm/llms/chatgpt/responses/transformation.py b/litellm/llms/chatgpt/responses/transformation.py index f85f50b7570..dfc8fa7a1eb 100644 --- a/litellm/llms/chatgpt/responses/transformation.py +++ b/litellm/llms/chatgpt/responses/transformation.py @@ -203,10 +203,6 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): ], } ] - if "created_at" in completed_response_payload: - completed_response_payload["created_at"] = _safe_convert_created_field( - completed_response_payload["created_at"] - ) try: completed_response = ResponsesAPIResponse(**completed_response_payload) except Exception: 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/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/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index d3f7d882da8..494dcd61952 100644 --- a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -289,11 +289,10 @@ async def test_increment_remaining_budget_metrics(prometheus_logger): future_reset_time_team = datetime.now() + timedelta(hours=10) future_reset_time_key = datetime.now() + timedelta(hours=12) # Mock the get_team_object and get_key_object functions to return objects with budget reset times - with patch( - "litellm.proxy.auth.auth_checks.get_team_object" - ) as mock_get_team, patch( - "litellm.proxy.auth.auth_checks.get_key_object" - ) as mock_get_key: + with ( + patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team, + patch("litellm.proxy.auth.auth_checks.get_key_object") as mock_get_key, + ): mock_get_team.return_value = MagicMock(budget_reset_at=future_reset_time_team) mock_get_key.return_value = MagicMock(budget_reset_at=future_reset_time_key) @@ -1518,9 +1517,12 @@ async def test_initialize_remaining_budget_metrics(prometheus_logger): """ litellm.prometheus_initialize_budget_metrics = True # Mock the prisma client and get_paginated_teams function - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams" - ) as mock_get_teams: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams" + ) as mock_get_teams, + ): # Create mock team data with proper datetime objects for budget_reset_at future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now mock_teams = [ @@ -1613,11 +1615,15 @@ async def test_initialize_remaining_budget_metrics_exception_handling( """ litellm.prometheus_initialize_budget_metrics = True # Mock the prisma client and get_paginated_teams function to raise an exception - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams" - ) as mock_get_teams, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" - ) as mock_list_keys: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams" + ) as mock_get_teams, + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" + ) as mock_list_keys, + ): # Make get_paginated_teams raise an exception mock_get_teams.side_effect = Exception("Database error") mock_list_keys.side_effect = Exception("Key listing error") @@ -1636,9 +1642,7 @@ async def test_initialize_remaining_budget_metrics_exception_handling( # Mock litellm_organizationtable to raise an exception for org budget metrics mock_orgtable = MagicMock() - mock_orgtable.find_many = MagicMock( - side_effect=Exception("Org database error") - ) + mock_orgtable.find_many = MagicMock(side_effect=Exception("Org database error")) mock_orgtable.count = MagicMock(side_effect=Exception("Org count error")) mock_db = MagicMock() @@ -1699,9 +1703,12 @@ async def test_initialize_api_key_budget_metrics(prometheus_logger): """ litellm.prometheus_initialize_budget_metrics = True # Mock the prisma client and _list_key_helper function - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" - ) as mock_list_keys: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" + ) as mock_list_keys, + ): # Create mock key data with proper datetime objects for budget_reset_at future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now key1 = UserAPIKeyAuth( diff --git a/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py b/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py index 212c5d4a322..6bdc808c4f3 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/guardrails/test_apply_guardrail_endpoint.py b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py index 0d27df50d15..b7fd665c410 100644 --- a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py +++ b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py @@ -105,7 +105,9 @@ async def test_apply_guardrail_endpoint_with_presidio_guardrail(): mock_guardrail = Mock(spec=CustomGuardrail) # Simulate masking PII entities - returns GenericGuardrailAPIInputs (dict with texts key) mock_guardrail.apply_guardrail = AsyncMock( - return_value={"texts": ["My name is [PERSON] and my email is [EMAIL_ADDRESS]"]} + return_value={ + "texts": ["My name is [PERSON] and my email is [EMAIL_ADDRESS]"] + } ) # Configure the registry to return our mock guardrail diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 9f4ca4ed108..3f9c2486d9c 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -95,9 +95,9 @@ async def test_async_pre_call_deployment_hook_resolves_model_id_from_litellm_met kwargs=kwargs, call_type=CallTypes.acreate_batch ) - assert result["input_file_id"] == provider_file_id, ( - f"Expected provider file ID '{provider_file_id}', got '{result['input_file_id']}'" - ) + assert ( + result["input_file_id"] == provider_file_id + ), f"Expected provider file ID '{provider_file_id}', got '{result['input_file_id']}'" @pytest.mark.asyncio @@ -134,9 +134,9 @@ async def test_async_pre_call_deployment_hook_prefers_top_level_model_info(): kwargs=kwargs, call_type=CallTypes.acreate_batch ) - assert result["input_file_id"] == top_level_provider_file, ( - "Should prefer top-level model_info over litellm_metadata" - ) + assert ( + result["input_file_id"] == top_level_provider_file + ), "Should prefer top-level model_info over litellm_metadata" @pytest.mark.asyncio @@ -162,9 +162,9 @@ async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_uncha kwargs=kwargs, call_type=CallTypes.acreate_batch ) - assert result["input_file_id"] == managed_file_id, ( - "File ID should remain unchanged when model_info is not available" - ) + assert ( + result["input_file_id"] == managed_file_id + ), "File ID should remain unchanged when model_info is not available" # def test_list_managed_files(): @@ -341,7 +341,9 @@ async def test_async_pre_call_hook_for_unified_finetuning_job(): @pytest.mark.asyncio -@pytest.mark.parametrize("call_type", ["afile_content", "afile_delete", "afile_retrieve"]) +@pytest.mark.parametrize( + "call_type", ["afile_content", "afile_delete", "afile_retrieve"] +) async def test_can_user_call_unified_file_id(call_type): """ Test that on file retrieve, delete, and content we check if the user has access to the file @@ -601,7 +603,7 @@ async def test_error_file_id_for_failed_batch(): "litellm_model_name": "gpt-4o", "unified_batch_id": "litellm_proxy;model_id:test-model-id;llm_batch_id:batch_abc123", } - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=AsyncMock() ) @@ -620,12 +622,11 @@ async def test_error_file_id_for_failed_batch(): # Mock the afile_retrieve to simulate retrieving error file metadata with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_retrieve: mock_retrieve.return_value = error_file_object - + user_api_key_dict = UserAPIKeyAuth( - user_id="test-user-123", - parent_otel_span=MagicMock() + user_id="test-user-123", parent_otel_span=MagicMock() ) - + response = await proxy_managed_files.async_post_call_success_hook( data={}, user_api_key_dict=user_api_key_dict, @@ -636,7 +637,9 @@ async def test_error_file_id_for_failed_batch(): assert cast(LiteLLMBatch, response).error_file_id is not None assert not cast(LiteLLMBatch, response).error_file_id.startswith("error-") # Verify it's a base64 encoded managed file ID - assert _is_base64_encoded_unified_file_id(cast(LiteLLMBatch, response).error_file_id) + assert _is_base64_encoded_unified_file_id( + cast(LiteLLMBatch, response).error_file_id + ) @pytest.mark.asyncio @@ -650,7 +653,7 @@ async def test_async_post_call_success_hook_twice_assert_no_unique_violation(): # Use AsyncMock instead of real database connection prisma_client = AsyncMock() - + batch = LiteLLMBatch( id="bGl0ZWxsbV9wcm94eTttb2RlbF9pZDoxMjM0NTY3OTtsbG1fYmF0Y2hfaWQ6YmF0Y2hfNjg1YzVlNWQ2Mzk4ODE5MGI4NWJkYjIxNDdiYTEzMWQ", completion_window="24h", @@ -678,8 +681,10 @@ async def test_async_post_call_success_hook_twice_assert_no_unique_violation(): # first retrieve batch tasks = [] first_create_task = asyncio.create_task - with patch('asyncio.create_task') as mock_create_task: - mock_create_task.side_effect = lambda coro: tasks.append(first_create_task(coro)) or tasks[-1] + with patch("asyncio.create_task") as mock_create_task: + mock_create_task.side_effect = ( + lambda coro: tasks.append(first_create_task(coro)) or tasks[-1] + ) response = await proxy_managed_files.async_post_call_success_hook( data={}, @@ -700,8 +705,10 @@ async def test_async_post_call_success_hook_twice_assert_no_unique_violation(): # second retrieve batch tasks = [] second_create_task = asyncio.create_task - with patch('asyncio.create_task') as mock_create_task: - mock_create_task.side_effect = lambda coro: tasks.append(second_create_task(coro)) or tasks[-1] + with patch("asyncio.create_task") as mock_create_task: + mock_create_task.side_effect = ( + lambda coro: tasks.append(second_create_task(coro)) or tasks[-1] + ) await proxy_managed_files.async_post_call_success_hook( data={}, @@ -728,7 +735,7 @@ def test_update_responses_input_with_unified_file_id(): # Create a base64-encoded unified file ID # This decodes to: litellm_proxy:application/pdf;unified_id,6c0b5890-8914-48e0-b8f4-0ae5ed3c14a5;target_model_names,gpt-4o;llm_output_file_id,file-ECBPW7ML9g7XHdwGgUPZaM;llm_output_file_model_id,e26453f9e76e7993680d0068d98c1f4cc205bbad0967a33c664893568ca743c2 unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" - + # Test input with unified file ID in content array input_data = [ { @@ -745,15 +752,18 @@ def test_update_responses_input_with_unified_file_id(): ], } ] - + # Update the input updated_input = update_responses_input_with_model_file_ids(input=input_data) - + # Verify the file_id was updated to the provider-specific file ID assert updated_input[0]["content"][0]["type"] == "input_file" assert updated_input[0]["content"][0]["file_id"] == "file-ECBPW7ML9g7XHdwGgUPZaM" assert updated_input[0]["content"][1]["type"] == "input_text" - assert updated_input[0]["content"][1]["text"] == "What is the first dragon in the book?" + assert ( + updated_input[0]["content"][1]["text"] + == "What is the first dragon in the book?" + ) def test_update_responses_input_with_regular_file_id(): @@ -767,7 +777,7 @@ def test_update_responses_input_with_regular_file_id(): # Regular OpenAI file ID (not a unified file ID) regular_file_id = "file-abc123xyz" - + input_data = [ { "role": "user", @@ -783,10 +793,10 @@ def test_update_responses_input_with_regular_file_id(): ], } ] - + # Update the input updated_input = update_responses_input_with_model_file_ids(input=input_data) - + # Verify the file_id was kept unchanged (regular OpenAI file ID) assert updated_input[0]["content"][0]["type"] == "input_file" assert updated_input[0]["content"][0]["file_id"] == regular_file_id @@ -800,11 +810,11 @@ def test_update_responses_input_with_string_input(): from litellm.litellm_core_utils.prompt_templates.common_utils import ( update_responses_input_with_model_file_ids, ) - + input_data = "What is AI?" - + updated_input = update_responses_input_with_model_file_ids(input=input_data) - + assert updated_input == input_data assert isinstance(updated_input, str) @@ -822,7 +832,7 @@ def test_update_responses_input_with_multiple_file_ids(): unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" # Regular OpenAI file ID regular_file_id = "file-regular123" - + input_data = [ { "role": "user", @@ -842,9 +852,9 @@ def test_update_responses_input_with_multiple_file_ids(): ], } ] - + updated_input = update_responses_input_with_model_file_ids(input=input_data) - + # Verify unified file ID was updated assert updated_input[0]["content"][0]["file_id"] == "file-ECBPW7ML9g7XHdwGgUPZaM" # Verify regular file ID was kept unchanged @@ -864,7 +874,7 @@ def test_update_responses_input_with_model_file_id_mapping(): # Managed file ID (unified) managed_file_id = "litellm_proxy_file_123" - + # Model file ID mapping model_file_id_mapping = { managed_file_id: { @@ -872,7 +882,7 @@ def test_update_responses_input_with_model_file_id_mapping(): "model_id_2": "azure_file_xyz", } } - + input_data = [ { "role": "user", @@ -888,24 +898,24 @@ def test_update_responses_input_with_model_file_id_mapping(): ], } ] - + # Update input with model_id_1 mapping updated_input = update_responses_input_with_model_file_ids( input=input_data, model_id="model_id_1", model_file_id_mapping=model_file_id_mapping, ) - + # Verify the file_id was mapped to the correct provider-specific file ID assert updated_input[0]["content"][0]["file_id"] == "openai_file_abc" - + # Test with different model_id updated_input_2 = update_responses_input_with_model_file_ids( input=input_data, model_id="model_id_2", model_file_id_mapping=model_file_id_mapping, ) - + assert updated_input_2[0]["content"][0]["file_id"] == "azure_file_xyz" @@ -913,7 +923,7 @@ def test_update_responses_tools_with_model_file_id_mapping(): """ Test that update_responses_tools_with_model_file_ids correctly maps file IDs in code_interpreter tools with container.file_ids. - + This is a regression test for the issue where managed file IDs in tools.container.file_ids were not being replaced with provider-specific file IDs, causing "string too long" errors from OpenAI. @@ -925,7 +935,7 @@ def test_update_responses_tools_with_model_file_id_mapping(): # Managed file IDs managed_file_id_1 = "litellm_proxy_file_123" managed_file_id_2 = "litellm_proxy_file_456" - + # Model file ID mapping model_file_id_mapping = { managed_file_id_1: { @@ -935,7 +945,7 @@ def test_update_responses_tools_with_model_file_id_mapping(): "model_id_1": "openai_file_def", }, } - + tools = [ { "type": "code_interpreter", @@ -945,17 +955,20 @@ def test_update_responses_tools_with_model_file_id_mapping(): }, } ] - + # Update tools with model mapping updated_tools = update_responses_tools_with_model_file_ids( tools=tools, model_id="model_id_1", model_file_id_mapping=model_file_id_mapping, ) - + # Verify the file IDs were mapped to provider-specific file IDs assert updated_tools[0]["type"] == "code_interpreter" - assert updated_tools[0]["container"]["file_ids"] == ["openai_file_abc", "openai_file_def"] + assert updated_tools[0]["container"]["file_ids"] == [ + "openai_file_abc", + "openai_file_def", + ] def test_update_responses_tools_without_mapping(): @@ -968,7 +981,7 @@ def test_update_responses_tools_without_mapping(): ) regular_file_id = "file-abc123" - + tools = [ { "type": "code_interpreter", @@ -978,14 +991,14 @@ def test_update_responses_tools_without_mapping(): }, } ] - + # Update tools without mapping updated_tools = update_responses_tools_with_model_file_ids( tools=tools, model_id=None, model_file_id_mapping=None, ) - + # Verify the file ID was kept unchanged assert updated_tools[0]["container"]["file_ids"] == [regular_file_id] @@ -1001,13 +1014,13 @@ def test_update_responses_tools_with_mixed_file_ids(): managed_file_id = "litellm_proxy_file_123" regular_file_id = "file-abc123" - + model_file_id_mapping = { managed_file_id: { "model_id_1": "openai_file_abc", }, } - + tools = [ { "type": "code_interpreter", @@ -1017,16 +1030,19 @@ def test_update_responses_tools_with_mixed_file_ids(): }, } ] - + # Update tools updated_tools = update_responses_tools_with_model_file_ids( tools=tools, model_id="model_id_1", model_file_id_mapping=model_file_id_mapping, ) - + # Verify managed file ID was mapped and regular file ID was kept - assert updated_tools[0]["container"]["file_ids"] == ["openai_file_abc", regular_file_id] + assert updated_tools[0]["container"]["file_ids"] == [ + "openai_file_abc", + regular_file_id, + ] def test_get_file_ids_from_responses_tools(): @@ -1037,7 +1053,7 @@ def test_get_file_ids_from_responses_tools(): proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) - + tools = [ { "type": "code_interpreter", @@ -1047,9 +1063,9 @@ def test_get_file_ids_from_responses_tools(): }, } ] - + file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools) - + assert file_ids == ["file-123", "file-456"] @@ -1060,7 +1076,7 @@ def test_get_file_ids_from_responses_tools_multiple_tools(): proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) - + tools = [ { "type": "code_interpreter", @@ -1080,9 +1096,9 @@ def test_get_file_ids_from_responses_tools_multiple_tools(): }, }, ] - + file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools) - + # Should extract file IDs only from code_interpreter tools assert file_ids == ["file-123", "file-456", "file-789"] @@ -1094,15 +1110,15 @@ def test_get_file_ids_from_responses_tools_empty(): proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) - + # Test with None file_ids = proxy_managed_files.get_file_ids_from_responses_tools(None) assert file_ids == [] - + # Test with empty list file_ids = proxy_managed_files.get_file_ids_from_responses_tools([]) assert file_ids == [] - + # Test with tools without file_ids tools = [{"type": "file_search"}] file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools) @@ -1119,30 +1135,30 @@ async def test_check_file_ids_access_with_unified_file_ids(): # Create a unified file ID unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" regular_file_id = "file-abc123" - + # Mock the access check to return True prisma_client = AsyncMock() internal_usage_cache = MagicMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock can_user_call_unified_file_id to return True proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True) - + user_api_key_dict = UserAPIKeyAuth( user_id="test_user_123", parent_otel_span=MagicMock(), ) - + # Should not raise an exception for accessible files await proxy_managed_files.check_file_ids_access( [unified_file_id, regular_file_id], user_api_key_dict, ) - + # Verify can_user_call_unified_file_id was called for the unified file ID proxy_managed_files.can_user_call_unified_file_id.assert_called_once_with( unified_file_id, user_api_key_dict @@ -1155,32 +1171,32 @@ async def test_check_file_ids_access_denied(): Test that check_file_ids_access raises HTTPException when user doesn't have access. """ from litellm.proxy._types import UserAPIKeyAuth - + unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" - + prisma_client = AsyncMock() internal_usage_cache = MagicMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock can_user_call_unified_file_id to return False (access denied) proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=False) - + user_api_key_dict = UserAPIKeyAuth( user_id="test_user_123", parent_otel_span=MagicMock(), ) - + # Should raise HTTPException with 403 status code with pytest.raises(HTTPException) as exc_info: await proxy_managed_files.check_file_ids_access( [unified_file_id], user_api_key_dict, ) - + assert exc_info.value.status_code == 403 assert "does not have access to the file" in exc_info.value.detail @@ -1191,32 +1207,32 @@ async def test_check_file_ids_access_with_regular_files_only(): Test that check_file_ids_access doesn't check access for regular (non-unified) file IDs. """ from litellm.proxy._types import UserAPIKeyAuth - + regular_file_id_1 = "file-abc123" regular_file_id_2 = "file-xyz789" - + prisma_client = AsyncMock() internal_usage_cache = MagicMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock can_user_call_unified_file_id (should not be called for regular files) proxy_managed_files.can_user_call_unified_file_id = AsyncMock() - + user_api_key_dict = UserAPIKeyAuth( user_id="test_user_123", parent_otel_span=MagicMock(), ) - + # Should not raise exception and should not call can_user_call_unified_file_id await proxy_managed_files.check_file_ids_access( [regular_file_id_1, regular_file_id_2], user_api_key_dict, ) - + # Verify can_user_call_unified_file_id was NOT called proxy_managed_files.can_user_call_unified_file_id.assert_not_called() @@ -1227,31 +1243,31 @@ async def test_completion_with_file_access_check(): Test that completion call type checks file access before processing. """ from litellm.proxy._types import UserAPIKeyAuth - + unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" - + prisma_client = AsyncMock() prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) - + internal_usage_cache = MagicMock() internal_usage_cache.async_get_cache = AsyncMock(return_value=None) - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock the get_model_file_id_mapping to return empty dict proxy_managed_files.get_model_file_id_mapping = AsyncMock(return_value={}) - + # Mock access check to allow access proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True) - + user_api_key_dict = UserAPIKeyAuth( user_id="test_user_123", parent_otel_span=MagicMock(), ) - + data = { "messages": [ { @@ -1267,7 +1283,7 @@ async def test_completion_with_file_access_check(): ], "model": "gpt-4", } - + # Should not raise exception result = await proxy_managed_files.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -1275,7 +1291,7 @@ async def test_completion_with_file_access_check(): data=data, call_type="acompletion", ) - + # Verify access check was called proxy_managed_files.can_user_call_unified_file_id.assert_called_once() @@ -1286,32 +1302,32 @@ async def test_responses_with_file_access_check(): Test that responses API checks file access for files in both input and tools. """ from litellm.proxy._types import UserAPIKeyAuth - + unified_file_id_1 = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" unified_file_id_2 = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsNzc3Nzc3Nzc7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1YWVo7bGxtX291dHB1dF9maWxlX21vZGVsX2lkLG1vZGVsXzEyMw" - + prisma_client = AsyncMock() prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) - + internal_usage_cache = MagicMock() internal_usage_cache.async_get_cache = AsyncMock(return_value=None) - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock the get_model_file_id_mapping to return empty dict proxy_managed_files.get_model_file_id_mapping = AsyncMock(return_value={}) - + # Mock access check to allow access proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True) - + user_api_key_dict = UserAPIKeyAuth( user_id="test_user_123", parent_otel_span=MagicMock(), ) - + data = { "input": [ { @@ -1333,7 +1349,7 @@ async def test_responses_with_file_access_check(): ], "model": "gpt-4", } - + # Should not raise exception result = await proxy_managed_files.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -1341,7 +1357,7 @@ async def test_responses_with_file_access_check(): data=data, call_type="aresponses", ) - + # Verify access check was called for both file IDs assert proxy_managed_files.can_user_call_unified_file_id.call_count == 2 @@ -1353,17 +1369,19 @@ async def test_store_unified_file_id_with_none_file_object(): (e.g., for batch output files that are stored before file metadata is available). """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - prisma_client.db.litellm_managedfiletable.create = AsyncMock(return_value=MagicMock()) + prisma_client.db.litellm_managedfiletable.create = AsyncMock( + return_value=MagicMock() + ) internal_usage_cache = MagicMock() internal_usage_cache.async_set_cache = AsyncMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Store with file_object=None (simulating batch output file storage) await proxy_managed_files.store_unified_file_id( file_id="test-unified-file-id", @@ -1372,7 +1390,7 @@ async def test_store_unified_file_id_with_none_file_object(): model_mappings={"model-123": "file-provider-xyz"}, user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), ) - + # Verify DB create was called with expected data (without file_object) prisma_client.db.litellm_managedfiletable.create.assert_called_once() call_args = prisma_client.db.litellm_managedfiletable.create.call_args @@ -1387,34 +1405,38 @@ async def test_afile_delete_returns_provider_response_when_stored_file_object_no stored file_object is None (e.g., for batch output files). """ from litellm.types.llms.openai import OpenAIFileObject - + unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsdGVzdC1pZDt0YXJnZXRfbW9kZWxfbmFtZXMsZ3B0LTRvO2xsbV9vdXRwdXRfZmlsZV9pZCxmaWxlLXByb3ZpZGVyLXh5ejtsbG1fb3V0cHV0X2ZpbGVfbW9kZWxfaWQsbW9kZWwtMTIz" - + prisma_client = AsyncMock() db_record = MagicMock() db_record.model_mappings = '{"model-123": "file-provider-xyz"}' - prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=db_record) + prisma_client.db.litellm_managedfiletable.find_first = AsyncMock( + return_value=db_record + ) prisma_client.db.litellm_managedfiletable.delete = AsyncMock() - + internal_usage_cache = MagicMock() - internal_usage_cache.async_get_cache = AsyncMock(return_value={ - "unified_file_id": unified_file_id, - "model_mappings": {"model-123": "file-provider-xyz"}, - "flat_model_file_ids": ["file-provider-xyz"], - "file_object": None, - "created_by": "test-user", - "updated_by": "test-user", - }) + internal_usage_cache.async_get_cache = AsyncMock( + return_value={ + "unified_file_id": unified_file_id, + "model_mappings": {"model-123": "file-provider-xyz"}, + "flat_model_file_ids": ["file-provider-xyz"], + "file_object": None, + "created_by": "test-user", + "updated_by": "test-user", + } + ) internal_usage_cache.async_set_cache = AsyncMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock the delete_unified_file_id to return None (simulating file_object=None) proxy_managed_files.delete_unified_file_id = AsyncMock(return_value=None) - + # Mock router response provider_delete_response = OpenAIFileObject( id="file-provider-xyz", @@ -1424,16 +1446,16 @@ async def test_afile_delete_returns_provider_response_when_stored_file_object_no filename="test.jsonl", purpose="batch", ) - + mock_router = MagicMock() mock_router.afile_delete = AsyncMock(return_value=provider_delete_response) - + result = await proxy_managed_files.afile_delete( file_id=unified_file_id, litellm_parent_otel_span=None, llm_router=mock_router, ) - + # Should return the provider response with the unified file ID assert result is not None assert result.id == unified_file_id @@ -1446,21 +1468,21 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none(): file_object is None (e.g., for batch output files). """ from litellm.types.llms.openai import OpenAIFileObject - + prisma_client = AsyncMock() internal_usage_cache = MagicMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock get_unified_file_id to return a stored object with file_object=None stored_file = MagicMock() stored_file.file_object = None stored_file.model_mappings = {"model-123": "file-provider-xyz"} proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file) - + # Mock the router and provider response provider_file_response = OpenAIFileObject( id="file-provider-xyz", @@ -1470,23 +1492,25 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none(): filename="output.jsonl", purpose="batch_output", ) - + mock_router = MagicMock() - mock_router.get_deployment_credentials_with_provider = MagicMock(return_value={ - "api_key": "test-key", - "api_base": "https://api.openai.com", - }) - + mock_router.get_deployment_credentials_with_provider = MagicMock( + return_value={ + "api_key": "test-key", + "api_base": "https://api.openai.com", + } + ) + with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_afile_retrieve: mock_afile_retrieve.return_value = provider_file_response - + unified_file_id = "test-unified-file-id" result = await proxy_managed_files.afile_retrieve( file_id=unified_file_id, litellm_parent_otel_span=None, llm_router=mock_router, ) - + # Should return the provider response with the unified file ID assert result is not None assert result.id == unified_file_id @@ -1501,27 +1525,27 @@ async def test_afile_retrieve_raises_error_when_no_router_and_file_object_none() """ prisma_client = AsyncMock() internal_usage_cache = MagicMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock get_unified_file_id to return a stored object with file_object=None stored_file = MagicMock() stored_file.file_object = None stored_file.model_mappings = {"model-123": "file-provider-xyz"} proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file) - + unified_file_id = "test-unified-file-id" - + with pytest.raises(Exception) as exc_info: await proxy_managed_files.afile_retrieve( file_id=unified_file_id, litellm_parent_otel_span=None, llm_router=None, ) - + assert "llm_router is required" in str(exc_info.value) @@ -1532,15 +1556,15 @@ async def test_afile_retrieve_returns_stored_file_object_when_exists(): (the normal case for user-uploaded files). """ from litellm.types.llms.openai import OpenAIFileObject - + prisma_client = AsyncMock() internal_usage_cache = MagicMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock get_unified_file_id to return a stored object WITH file_object stored_file_object = OpenAIFileObject( id="test-unified-file-id", @@ -1553,13 +1577,13 @@ async def test_afile_retrieve_returns_stored_file_object_when_exists(): stored_file = MagicMock() stored_file.file_object = stored_file_object proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file) - + result = await proxy_managed_files.afile_retrieve( file_id="test-unified-file-id", litellm_parent_otel_span=None, llm_router=None, ) - + # Should return the stored file object directly assert result == stored_file_object @@ -1572,21 +1596,21 @@ async def test_afile_retrieve_raises_error_for_non_managed_file(): """ prisma_client = AsyncMock() internal_usage_cache = MagicMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) - + # Mock get_unified_file_id to return None (file not found) proxy_managed_files.get_unified_file_id = AsyncMock(return_value=None) - + with pytest.raises(Exception) as exc_info: await proxy_managed_files.afile_retrieve( file_id="non-existent-file-id", litellm_parent_otel_span=None, ) - + assert "not found" in str(exc_info.value) @@ -1597,54 +1621,58 @@ async def test_list_batches_from_managed_objects_table(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() - + batch_record_1 = MagicMock() batch_record_1.unified_object_id = "unified-batch-id-1" - batch_record_1.file_object = json.dumps({ - "id": "batch_abc123", - "object": "batch", - "endpoint": "/v1/chat/completions", - "completion_window": "24h", - "status": "completed", - "created_at": 1234567890, - "input_file_id": "file-input-1", - "request_counts": {"total": 1, "completed": 1, "failed": 0}, - }) - + batch_record_1.file_object = json.dumps( + { + "id": "batch_abc123", + "object": "batch", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "status": "completed", + "created_at": 1234567890, + "input_file_id": "file-input-1", + "request_counts": {"total": 1, "completed": 1, "failed": 0}, + } + ) + batch_record_2 = MagicMock() batch_record_2.unified_object_id = "unified-batch-id-2" - batch_record_2.file_object = json.dumps({ - "id": "batch_xyz789", - "object": "batch", - "endpoint": "/v1/chat/completions", - "completion_window": "24h", - "status": "in_progress", - "created_at": 1234567891, - "input_file_id": "file-input-2", - "request_counts": {"total": 5, "completed": 2, "failed": 0}, - }) - + batch_record_2.file_object = json.dumps( + { + "id": "batch_xyz789", + "object": "batch", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "status": "in_progress", + "created_at": 1234567891, + "input_file_id": "file-input-2", + "request_counts": {"total": 5, "completed": 2, "failed": 0}, + } + ) + prisma_client.db.litellm_managedobjecttable.find_many.return_value = [ batch_record_1, batch_record_2, ] - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) - + result = await proxy_managed_files.list_user_batches( user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), limit=10, ) - + assert result["object"] == "list" assert len(result["data"]) == 2 assert result["data"][0].id == "unified-batch-id-1" assert result["data"][1].id == "unified-batch-id-2" assert result["first_id"] == "unified-batch-id-1" assert result["last_id"] == "unified-batch-id-2" - + # Should filter by user_id (created_by) prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with( where={"file_purpose": "batch", "created_by": "test-user"}, @@ -1659,21 +1687,21 @@ async def test_list_batches_from_managed_objects_table_empty_list(): prisma_client = AsyncMock() prisma_client.db.litellm_managedobjecttable.find_many.return_value = [] - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) - + result = await proxy_managed_files.list_user_batches( user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), ) - + assert result["object"] == "list" assert len(result["data"]) == 0 assert result["first_id"] is None assert result["last_id"] is None assert result["has_more"] is False - + # Verify where clause includes created_by filter # Default take is 20 when no limit is provided prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with( @@ -1685,6 +1713,7 @@ async def test_list_batches_from_managed_objects_table_empty_list(): def _create_unified_batch_id(model_id: str, batch_id: str) -> str: import base64 + unified_str = f"litellm_proxy;model_id:{model_id};llm_batch_id:{batch_id}" return base64.urlsafe_b64encode(unified_str.encode()).decode().rstrip("=") @@ -1694,11 +1723,11 @@ async def test_list_batches_from_managed_objects_table_provider_filter_raises_ex from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) - + # Filtering by provider should raise Exception with pytest.raises(Exception) as exc_info: await proxy_managed_files.list_user_batches( @@ -1706,11 +1735,11 @@ async def test_list_batches_from_managed_objects_table_provider_filter_raises_ex limit=10, provider="openai", ) - + assert str(exc_info.value) == ( "Filtering by 'provider' is not supported when using managed batches." ) - + # Verify find_many was NOT called since exception is raised before database query prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called() @@ -1720,7 +1749,7 @@ async def test_list_batches_from_managed_objects_table_target_model_name_filter_ from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -1732,59 +1761,64 @@ async def test_list_batches_from_managed_objects_table_target_model_name_filter_ limit=10, target_model_names="gpt-4o,gpt-3.5", ) - + assert str(exc_info.value) == ( "Filtering by 'target_model_names' is not supported when using managed batches." ) - + # Verify find_many was NOT called since exception is raised before database query prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called() + @pytest.mark.asyncio async def test_list_batches_from_managed_objects_table_filters_by_created_by(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() - + # Create batch for user1 batch_user1 = MagicMock() batch_user1.unified_object_id = "unified-batch-user1" - batch_user1.file_object = json.dumps({ - "id": "batch_user1_abc", - "object": "batch", - "endpoint": "/v1/chat/completions", - "completion_window": "24h", - "status": "completed", - "created_at": 1234567890, - "input_file_id": "file-input-user1", - "request_counts": {"total": 1, "completed": 1, "failed": 0}, - }) - + batch_user1.file_object = json.dumps( + { + "id": "batch_user1_abc", + "object": "batch", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "status": "completed", + "created_at": 1234567890, + "input_file_id": "file-input-user1", + "request_counts": {"total": 1, "completed": 1, "failed": 0}, + } + ) + # Create batch for user2 batch_user2 = MagicMock() batch_user2.unified_object_id = "unified-batch-user2" - batch_user2.file_object = json.dumps({ - "id": "batch_user2_xyz", - "object": "batch", - "endpoint": "/v1/chat/completions", - "completion_window": "24h", - "status": "completed", - "created_at": 1234567891, - "input_file_id": "file-input-user2", - "request_counts": {"total": 2, "completed": 2, "failed": 0}, - }) - + batch_user2.file_object = json.dumps( + { + "id": "batch_user2_xyz", + "object": "batch", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "status": "completed", + "created_at": 1234567891, + "input_file_id": "file-input-user2", + "request_counts": {"total": 2, "completed": 2, "failed": 0}, + } + ) + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) - + # Query with user1's API key - should only return user1's batch prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user1] result_user1 = await proxy_managed_files.list_user_batches( user_api_key_dict=UserAPIKeyAuth(user_id="user1"), limit=10, ) - + assert len(result_user1["data"]) == 1 assert result_user1["data"][0].id == "unified-batch-user1" prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with( @@ -1792,14 +1826,14 @@ async def test_list_batches_from_managed_objects_table_filters_by_created_by(): take=10, order={"created_at": "desc"}, ) - + # Query with user2's API key - should only return user2's batch prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user2] result_user2 = await proxy_managed_files.list_user_batches( user_api_key_dict=UserAPIKeyAuth(user_id="user2"), limit=10, ) - + assert len(result_user2["data"]) == 1 assert result_user2["data"][0].id == "unified-batch-user2" prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with( @@ -1822,7 +1856,7 @@ async def test_return_unified_file_id_includes_expires_at(): filename="test.jsonl", purpose="batch", status="uploaded", - expires_at=1234657890, + expires_at=1234657890, ) file_object._hidden_params = {"model_id": "test-model-id"} @@ -1862,25 +1896,27 @@ async def test_return_unified_file_id_includes_expires_at(): async def test_user_b_cannot_retrieve_user_a_batch(): """ Test that User B cannot retrieve a batch created by User A. - + This verifies batch isolation between users at the database/hook level. """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - + # Mock database to return User A as the creator batch_record = MagicMock() batch_record.created_by = "user_a_id" prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) - + # User B tries to retrieve User A's batch - unified_batch_id = "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz" - + unified_batch_id = ( + "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz" + ) + with pytest.raises(HTTPException) as exc_info: await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( @@ -1890,7 +1926,7 @@ async def test_user_b_cannot_retrieve_user_a_batch(): data={"batch_id": unified_batch_id}, call_type="aretrieve_batch", ) - + # Should raise 403 Permission Denied assert exc_info.value.status_code == 403 @@ -1901,21 +1937,23 @@ async def test_user_b_cannot_cancel_user_a_batch(): Test that User B cannot cancel a batch created by User A. """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - + # Mock database to return User A as the creator batch_record = MagicMock() batch_record.created_by = "user_a_id" prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) - + # User B tries to cancel User A's batch - unified_batch_id = "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz" - + unified_batch_id = ( + "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz" + ) + with pytest.raises(HTTPException) as exc_info: await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( @@ -1925,7 +1963,7 @@ async def test_user_b_cannot_cancel_user_a_batch(): data={"batch_id": unified_batch_id}, call_type="acancel_batch", ) - + # Should raise 403 Permission Denied assert exc_info.value.status_code == 403 @@ -1934,26 +1972,28 @@ async def test_user_b_cannot_cancel_user_a_batch(): async def test_user_a_can_retrieve_own_batch(): """ Test that User A can successfully retrieve their own batch. - + This is a positive test case to ensure permission checks don't block legitimate access. """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - + # Mock database to return User A as the creator batch_record = MagicMock() batch_record.created_by = "user_a_id" prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) - + # User A retrieves their own batch - unified_batch_id = "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz" - + unified_batch_id = ( + "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz" + ) + # Should not raise an exception result = await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( @@ -1963,7 +2003,7 @@ async def test_user_a_can_retrieve_own_batch(): data={"batch_id": unified_batch_id}, call_type="aretrieve_batch", ) - + # Should successfully return the decoded batch_id assert "batch_id" in result assert result["model"] == "my-model" @@ -1975,21 +2015,23 @@ async def test_user_b_cannot_retrieve_user_a_file(): Test that User B cannot retrieve a file created by User A. """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - + # Mock database to return User A as the creator file_record = MagicMock() file_record.created_by = "user_a_id" prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) - + # User B tries to retrieve User A's file - unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM" - + unified_file_id = ( + "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM" + ) + with pytest.raises(HTTPException) as exc_info: await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( @@ -1999,7 +2041,7 @@ async def test_user_b_cannot_retrieve_user_a_file(): data={"file_id": unified_file_id}, call_type="afile_retrieve", ) - + # Should raise 403 Permission Denied assert exc_info.value.status_code == 403 @@ -2010,21 +2052,23 @@ async def test_user_b_cannot_download_user_a_file_content(): Test that User B cannot download file content for User A's file. """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - + # Mock database to return User A as the creator file_record = MagicMock() file_record.created_by = "user_a_id" prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) - + # User B tries to download User A's file content - unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM" - + unified_file_id = ( + "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM" + ) + with pytest.raises(HTTPException) as exc_info: await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( @@ -2034,7 +2078,7 @@ async def test_user_b_cannot_download_user_a_file_content(): data={"file_id": unified_file_id}, call_type="afile_content", ) - + # Should raise 403 Permission Denied assert exc_info.value.status_code == 403 @@ -2045,21 +2089,23 @@ async def test_user_b_cannot_delete_user_a_file(): Test that User B cannot delete a file created by User A. """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - + # Mock database to return User A as the creator file_record = MagicMock() file_record.created_by = "user_a_id" prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) - + # User B tries to delete User A's file - unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM" - + unified_file_id = ( + "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM" + ) + with pytest.raises(HTTPException) as exc_info: await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( @@ -2069,7 +2115,7 @@ async def test_user_b_cannot_delete_user_a_file(): data={"file_id": unified_file_id}, call_type="afile_delete", ) - + # Should raise 403 Permission Denied assert exc_info.value.status_code == 403 @@ -2078,34 +2124,38 @@ async def test_user_b_cannot_delete_user_a_file(): async def test_user_a_can_retrieve_own_file(): """ Test that User A can successfully retrieve their own file. - + Positive test case to ensure permission checks work correctly for the owner. """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - + # Mock database to return User A as the creator file_record = MagicMock() file_record.created_by = "user_a_id" file_record.model_mappings = '{"model-123": "file-abc123"}' - file_record.file_object = json.dumps({ - "id": "file-abc123", - "object": "file", - "bytes": 1234, - "created_at": 1234567890, - "filename": "test.jsonl", - "purpose": "batch", - }) + file_record.file_object = json.dumps( + { + "id": "file-abc123", + "object": "file", + "bytes": 1234, + "created_at": 1234567890, + "filename": "test.jsonl", + "purpose": "batch", + } + ) prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) - + # User A retrieves their own file - unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM" - + unified_file_id = ( + "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM" + ) + # Should not raise an exception result = await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( @@ -2115,7 +2165,7 @@ async def test_user_a_can_retrieve_own_file(): data={"file_id": unified_file_id}, call_type="afile_retrieve", ) - + # Should successfully return the decoded file_id assert "file_id" in result @@ -2124,44 +2174,46 @@ async def test_user_a_can_retrieve_own_file(): async def test_list_batches_only_returns_user_own_batches(): """ Test that list_user_batches only returns batches created by the requesting user. - + This ensures users cannot see other users' batches in list operations. """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - + # Create batches for User A batch_user_a = MagicMock() batch_user_a.unified_object_id = "batch-user-a" - batch_user_a.file_object = json.dumps({ - "id": "batch_a", - "object": "batch", - "endpoint": "/v1/chat/completions", - "completion_window": "24h", - "status": "completed", - "created_at": 1234567890, - "input_file_id": "file-a", - "request_counts": {"total": 1, "completed": 1, "failed": 0}, - }) - + batch_user_a.file_object = json.dumps( + { + "id": "batch_a", + "object": "batch", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "status": "completed", + "created_at": 1234567890, + "input_file_id": "file-a", + "request_counts": {"total": 1, "completed": 1, "failed": 0}, + } + ) + # Mock database to only return User A's batches prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user_a] - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) - + # User A requests their batches result = await proxy_managed_files.list_user_batches( user_api_key_dict=UserAPIKeyAuth(user_id="user_a_id"), limit=10, ) - + # Should only return User A's batches assert len(result["data"]) == 1 assert result["data"][0].id == "batch-user-a" - + # Verify the database query filtered by user_id prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with( where={"file_purpose": "batch", "created_by": "user_a_id"}, @@ -2174,51 +2226,49 @@ async def test_list_batches_only_returns_user_own_batches(): async def test_same_user_different_keys_can_access_batch(): """ Test that different API keys for the same user can access the same batch. - + This verifies that permission checks are based on user_id, not API key, allowing users to have multiple keys that can all access their resources. """ from litellm.proxy._types import UserAPIKeyAuth - + prisma_client = AsyncMock() - + # Mock database to return the user_id as creator batch_record = MagicMock() batch_record.created_by = "user_a_id" prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record - + proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) - - unified_batch_id = "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz" - + + unified_batch_id = ( + "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz" + ) + # First API key for User A retrieves the batch result1 = await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( - user_id="user_a_id", - api_key="key-1", - parent_otel_span=MagicMock() + user_id="user_a_id", api_key="key-1", parent_otel_span=MagicMock() ), cache=MagicMock(), data={"batch_id": unified_batch_id}, call_type="aretrieve_batch", ) - + assert "batch_id" in result1 - + # Second API key for the same User A retrieves the batch result2 = await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( - user_id="user_a_id", - api_key="key-2", - parent_otel_span=MagicMock() + user_id="user_a_id", api_key="key-2", parent_otel_span=MagicMock() ), cache=MagicMock(), data={"batch_id": unified_batch_id}, call_type="aretrieve_batch", ) - + assert "batch_id" in result2 # Both keys should get the same result assert result1["batch_id"] == result2["batch_id"] diff --git a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py index 69c3b4cb59a..699ab8d1c4c 100644 --- a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py @@ -31,12 +31,16 @@ class TestAvailableEnterpriseUsers: self, client, mock_user_api_key_auth ): """Test when max_users is set and user count is within limit""" - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.premium_user", - True, - ), patch( - "litellm.proxy.proxy_server.premium_user_data", - {"max_users": 10}, + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.proxy_server.premium_user", + True, + ), + patch( + "litellm.proxy.proxy_server.premium_user_data", + {"max_users": 10}, + ), ): # Mock database count mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=5) @@ -66,12 +70,16 @@ class TestAvailableEnterpriseUsers: self, client, mock_user_api_key_auth ): """Test when max_users is not set (premium_user_data is None or doesn't contain max_users)""" - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.premium_user", - True, - ), patch( - "litellm.proxy.proxy_server.premium_user_data", - None, + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.proxy_server.premium_user", + True, + ), + patch( + "litellm.proxy.proxy_server.premium_user_data", + None, + ), ): # Mock database count mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=3) @@ -99,12 +107,16 @@ class TestAvailableEnterpriseUsers: self, client, mock_user_api_key_auth ): """Test the current bug where total_users_remaining can be negative""" - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.premium_user", - True, - ), patch( - "litellm.proxy.proxy_server.premium_user_data", - {"key": "value"}, + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.proxy_server.premium_user", + True, + ), + patch( + "litellm.proxy.proxy_server.premium_user_data", + {"key": "value"}, + ), ): # Mock database count higher than max_users to trigger the bug mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=8) @@ -140,12 +152,15 @@ class TestAvailableEnterpriseUsers: """Test when prisma_client is None (no database connection)""" from litellm.proxy._types import CommonProxyErrors - with patch( - "litellm.proxy.proxy_server.prisma_client", - None, - ), patch( - "litellm.proxy.proxy_server.premium_user", - True, + with ( + patch( + "litellm.proxy.proxy_server.prisma_client", + None, + ), + patch( + "litellm.proxy.proxy_server.premium_user", + True, + ), ): # Override the dependency client.app.dependency_overrides[mock_user_api_key_auth] = lambda: { 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 52cb94ff346..8e0a3655b57 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/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py b/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py index 54ea41a6450..fde468618ea 100644 --- a/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py +++ b/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py @@ -581,17 +581,21 @@ def test_vertex_ai_text_only_embedding_uses_embed_content(): def test_filter_embed_params_drops_unsupported(): """Unsupported params like max_tokens should be filtered out.""" - result = _filter_embed_params({"dimensions": 768, "max_tokens": 256, "temperature": 0.5}) + result = _filter_embed_params( + {"dimensions": 768, "max_tokens": 256, "temperature": 0.5} + ) assert result == {"outputDimensionality": 768} def test_filter_embed_params_keeps_supported(): """All supported Gemini embedding params should pass through.""" - result = _filter_embed_params({ - "dimensions": 768, - "task_type": "RETRIEVAL_DOCUMENT", - "title": "My doc", - }) + result = _filter_embed_params( + { + "dimensions": 768, + "task_type": "RETRIEVAL_DOCUMENT", + "title": "My doc", + } + ) assert result == { "outputDimensionality": 768, "taskType": "RETRIEVAL_DOCUMENT", 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..a24a29d0ab6 100644 --- a/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py +++ b/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py @@ -546,9 +546,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 +567,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 +577,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/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/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py index af5e2341406..23a5d9ec990 100644 --- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -30,6 +30,7 @@ def no_invitation_wait(monkeypatch): monkeypatch.setattr(BaseEmailLogger, "_wait_for_invitation_creation", _noop) + @pytest.fixture def base_email_logger(): return BaseEmailLogger() @@ -283,7 +284,10 @@ async def test_send_key_created_email_without_key( mock_send_email.assert_called_once() call_args = mock_send_email.call_args[1] assert "sk-secret-key-456" not in call_args["html_body"] - assert "[Key hidden for security - retrieve from dashboard]" in call_args["html_body"] + assert ( + "[Key hidden for security - retrieve from dashboard]" + in call_args["html_body"] + ) @pytest.mark.asyncio @@ -317,7 +321,10 @@ async def test_send_key_rotated_email_without_key( mock_send_email.assert_called_once() call_args = mock_send_email.call_args[1] assert "sk-secret-rotated-789" not in call_args["html_body"] - assert "[Key hidden for security - retrieve from dashboard]" in call_args["html_body"] + assert ( + "[Key hidden for security - retrieve from dashboard]" + in call_args["html_body"] + ) @pytest.mark.asyncio @@ -371,52 +378,52 @@ async def test_get_invitation_link_creates_new_when_none_exist(base_email_logger """Test that _get_invitation_link creates a new invitation when none exist""" # Mock prisma client with no existing invitation rows mock_prisma = mock.MagicMock() - + # Mock find_many to return empty list (no existing invitations) async def mock_find_many_empty(*args, **kwargs): return [] - + mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_empty - + # Mock the create_invitation_for_user function mock_created_invitation = mock.MagicMock() mock_created_invitation.id = "new-invitation-id" - + with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with mock.patch( "litellm.proxy.management_helpers.user_invitation.create_invitation_for_user", - return_value=mock_created_invitation + return_value=mock_created_invitation, ) as mock_create_invitation: # Execute result = await base_email_logger._get_invitation_link( user_id="test-user", base_url="http://test.com" ) - + # Verify that create_invitation_for_user was called mock_create_invitation.assert_called_once() call_args = mock_create_invitation.call_args[1] assert call_args["data"].user_id == "test-user" assert call_args["user_api_key_dict"].user_id == "test-user" - + # Verify the returned link uses the new invitation ID assert result == "http://test.com/ui?invitation_id=new-invitation-id" -@pytest.mark.asyncio +@pytest.mark.asyncio async def test_get_invitation_link_uses_existing_when_available(base_email_logger): """Test that _get_invitation_link uses existing invitation when available""" # Mock prisma client with existing invitation row mock_invitation_row = mock.MagicMock() mock_invitation_row.id = "existing-invitation-id" - + mock_prisma = mock.MagicMock() - + # Mock find_many to return existing invitation async def mock_find_many_existing(*args, **kwargs): return [mock_invitation_row] - + mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_existing - + with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with mock.patch( "litellm.proxy.management_helpers.user_invitation.create_invitation_for_user" @@ -425,10 +432,10 @@ async def test_get_invitation_link_uses_existing_when_available(base_email_logge result = await base_email_logger._get_invitation_link( user_id="test-user", base_url="http://test.com" ) - + # Verify that create_invitation_for_user was NOT called mock_create_invitation.assert_not_called() - + # Verify the returned link uses the existing invitation ID assert result == "http://test.com/ui?invitation_id=existing-invitation-id" @@ -438,33 +445,33 @@ async def test_get_invitation_link_creates_new_when_list_is_none(base_email_logg """Test that _get_invitation_link creates a new invitation when invitation_rows is None""" # Mock prisma client to return None mock_prisma = mock.MagicMock() - + # Mock find_many to return None async def mock_find_many_none(*args, **kwargs): return None - + mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_none - + # Mock the create_invitation_for_user function mock_created_invitation = mock.MagicMock() mock_created_invitation.id = "new-invitation-from-none" - + with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with mock.patch( "litellm.proxy.management_helpers.user_invitation.create_invitation_for_user", - return_value=mock_created_invitation + return_value=mock_created_invitation, ) as mock_create_invitation: # Execute result = await base_email_logger._get_invitation_link( user_id="test-user", base_url="http://test.com" ) - + # Verify that create_invitation_for_user was called mock_create_invitation.assert_called_once() call_args = mock_create_invitation.call_args[1] assert call_args["data"].user_id == "test-user" assert call_args["user_api_key_dict"].user_id == "test-user" - + # Verify the returned link uses the new invitation ID assert result == "http://test.com/ui?invitation_id=new-invitation-from-none" @@ -495,13 +502,15 @@ async def test_get_email_params_user_invitation( user_email="test@example.com", ) - assert result.logo_url == "https://litellm-listing.s3.amazonaws.com/litellm_logo.png" + assert ( + result.logo_url + == "https://litellm-listing.s3.amazonaws.com/litellm_logo.png" + ) assert result.support_contact == "support@berri.ai" assert result.base_url == "http://test.com/ui?invitation_id=test-id" assert result.recipient_email == "test@example.com" - @pytest.fixture def mock_env_vars(monkeypatch): """Set up test environment variables""" @@ -513,69 +522,74 @@ def mock_env_vars(monkeypatch): monkeypatch.setenv("PROXY_BASE_URL", "http://test.com") monkeypatch.setenv("PROXY_API_URL", "https://test.com") + @pytest.mark.asyncio async def test_get_email_params_custom_templates_premium_user(mock_env_vars): """Test that _get_email_params returns correct values with custom templates for premium users""" # Mock premium_user as True with patch("litellm.proxy.proxy_server.premium_user", True): email_logger = BaseEmailLogger() - + # Test invitation email params invitation_params = await email_logger._get_email_params( email_event=EmailEvent.new_user_invitation, user_id="testid", user_email="test@example.com", - event_message="New User Invitation" + event_message="New User Invitation", ) - + assert invitation_params.subject == "Welcome to Test Company!" assert invitation_params.signature == "Best regards,\nTest Company Team" assert invitation_params.logo_url == "https://test-company.com/logo.png" assert invitation_params.support_contact == "support@test-company.com" assert invitation_params.base_url == "http://test.com" - + # Test key created email params key_params = await email_logger._get_email_params( email_event=EmailEvent.virtual_key_created, user_id="testid", user_email="test@example.com", - event_message="API Key Created" + event_message="API Key Created", ) - + assert key_params.subject == "Your Test Company API Key" assert key_params.signature == "Best regards,\nTest Company Team" + @pytest.mark.asyncio async def test_get_email_params_non_premium_user(mock_env_vars): """Test that non-premium users get default templates even when custom ones are provided""" # Mock premium_user as False with patch("litellm.proxy.proxy_server.premium_user", False): email_logger = BaseEmailLogger() - + # Test invitation email params email_params = await email_logger._get_email_params( email_event=EmailEvent.new_user_invitation, user_email="test@example.com", - event_message="New User Invitation" + event_message="New User Invitation", ) - + # Should use default values even though custom values are set in env assert email_params.subject == "LiteLLM: New User Invitation" assert email_params.signature == EMAIL_FOOTER - assert email_params.logo_url == "https://litellm-listing.s3.amazonaws.com/litellm_logo.png" + assert ( + email_params.logo_url + == "https://litellm-listing.s3.amazonaws.com/litellm_logo.png" + ) assert email_params.support_contact == "support@berri.ai" - # Test key created email params key_params = await email_logger._get_email_params( email_event=EmailEvent.virtual_key_created, user_email="test@example.com", - event_message="API Key Created" + event_message="API Key Created", ) - + assert key_params.subject == "LiteLLM: API Key Created" assert key_params.signature == EMAIL_FOOTER + @pytest.mark.asyncio async def test_get_email_params_default_templates(monkeypatch): """Test that _get_email_params uses default templates when custom ones aren't provided""" @@ -583,28 +597,28 @@ async def test_get_email_params_default_templates(monkeypatch): monkeypatch.delenv("EMAIL_SUBJECT_INVITATION", raising=False) monkeypatch.delenv("EMAIL_SUBJECT_KEY_CREATED", raising=False) monkeypatch.delenv("EMAIL_SIGNATURE", raising=False) - + # Mock premium_user as True (shouldn't matter since no custom values are set) with patch("litellm.proxy.proxy_server.premium_user", True): email_logger = BaseEmailLogger() - + # Test invitation email params with default template invitation_params = await email_logger._get_email_params( email_event=EmailEvent.new_user_invitation, user_email="test@example.com", - event_message="New User Invitation" + event_message="New User Invitation", ) - + assert invitation_params.subject == "LiteLLM: New User Invitation" assert invitation_params.signature == EMAIL_FOOTER - + # Test key created email params with default template key_params = await email_logger._get_email_params( email_event=EmailEvent.virtual_key_created, user_email="test@example.com", - event_message="API Key Created" + event_message="API Key Created", ) - + assert key_params.subject == "LiteLLM: API Key Created" assert key_params.signature == EMAIL_FOOTER @@ -639,7 +653,10 @@ async def test_send_soft_budget_alert_email( call_args = mock_send_email.call_args[1] assert call_args["from_email"] == BaseEmailLogger.DEFAULT_LITELLM_EMAIL assert call_args["to_email"] == ["test@example.com"] - assert call_args["subject"] == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0" + assert ( + call_args["subject"] + == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0" + ) assert "$100.0" in call_args["html_body"] # soft_budget assert "$105.0" in call_args["html_body"] # spend assert "$200.0" in call_args["html_body"] # max_budget @@ -673,13 +690,13 @@ async def test_send_soft_budget_alert_email_no_max_budget( call_args = mock_send_email.call_args[1] assert "$100.0" in call_args["html_body"] # soft_budget assert "$105.0" in call_args["html_body"] # spend - assert "Maximum Budget" not in call_args["html_body"] # max_budget should not be shown + assert ( + "Maximum Budget" not in call_args["html_body"] + ) # max_budget should not be shown @pytest.mark.asyncio -async def test_budget_alerts_soft_budget_crossed( - base_email_logger, mock_send_email -): +async def test_budget_alerts_soft_budget_crossed(base_email_logger, mock_send_email): """Test that budget_alerts sends email when soft budget is crossed""" user_info = CallInfo( user_id="test_user", @@ -708,11 +725,14 @@ async def test_budget_alerts_soft_budget_crossed( mock_send_email.assert_called_once() call_args = mock_send_email.call_args[1] assert call_args["to_email"] == ["test@example.com"] - + # Verify cache was set to prevent duplicate alerts mock_cache.async_set_cache.assert_called_once() cache_call_args = mock_cache.async_set_cache.call_args[1] - assert cache_call_args["key"] == "email_budget_alerts:soft_budget_crossed:test_user" + assert ( + cache_call_args["key"] + == "email_budget_alerts:soft_budget_crossed:test_user" + ) assert cache_call_args["value"] == "SENT" assert cache_call_args["ttl"] == EMAIL_BUDGET_ALERT_TTL @@ -766,9 +786,7 @@ async def test_budget_alerts_soft_budget_duplicate_prevention( @pytest.mark.asyncio -async def test_budget_alerts_no_budgets( - base_email_logger, mock_send_email -): +async def test_budget_alerts_no_budgets(base_email_logger, mock_send_email): """Test that budget_alerts returns early when no budgets are set""" user_info = CallInfo( user_id="test_user", @@ -817,7 +835,10 @@ async def test_budget_alerts_uses_token_for_cache_key( # Verify cache key uses token instead of user_id mock_cache.async_set_cache.assert_called_once() cache_call_args = mock_cache.async_set_cache.call_args[1] - assert cache_call_args["key"] == "email_budget_alerts:soft_budget_crossed:hashed_token_123" + assert ( + cache_call_args["key"] + == "email_budget_alerts:soft_budget_crossed:hashed_token_123" + ) @pytest.mark.asyncio @@ -838,7 +859,9 @@ async def test_get_email_params_soft_budget_crossed( ) # Should use default subject template for soft_budget_crossed - assert result.subject == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0" + assert ( + result.subject == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0" + ) assert result.recipient_email == "test@example.com" assert result.base_url == "http://test.com" @@ -867,16 +890,20 @@ async def test_budget_alerts_max_budget_alert_crossed( "PROXY_BASE_URL": "http://test.com", }, ): - await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info) + await base_email_logger.budget_alerts( + type="max_budget_alert", user_info=user_info + ) mock_send_email.assert_called_once() call_args = mock_send_email.call_args[1] assert call_args["to_email"] == ["test@example.com"] assert "Max Budget Alert" in call_args["subject"] - + mock_cache.async_set_cache.assert_called_once() cache_call_args = mock_cache.async_set_cache.call_args[1] - assert cache_call_args["key"] == "email_budget_alerts:max_budget_alert:test_user" + assert ( + cache_call_args["key"] == "email_budget_alerts:max_budget_alert:test_user" + ) assert cache_call_args["value"] == "SENT" assert cache_call_args["ttl"] == EMAIL_BUDGET_ALERT_TTL @@ -906,15 +933,15 @@ async def test_multi_threshold_sends_crossed_thresholds( base_email_logger.internal_usage_cache = mock_cache with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): - await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info) + await base_email_logger.budget_alerts( + type="max_budget_alert", user_info=user_info + ) # spend=80 crosses 50% ($50) and 75% ($75), but not 100% ($100) assert mock_send_email.call_count == 2 # Check cache keys include threshold percentage - cache_keys = [ - c[1]["key"] for c in mock_cache.async_set_cache.call_args_list - ] + cache_keys = [c[1]["key"] for c in mock_cache.async_set_cache.call_args_list] assert "email_budget_alerts:max_budget_alert:50:hashed_key_1" in cache_keys assert "email_budget_alerts:max_budget_alert:75:hashed_key_1" in cache_keys @@ -949,7 +976,9 @@ async def test_multi_threshold_dedup_cache_prevents_resend( base_email_logger.internal_usage_cache = mock_cache with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): - await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info) + await base_email_logger.budget_alerts( + type="max_budget_alert", user_info=user_info + ) # Only 75% should fire assert mock_send_email.call_count == 1 @@ -980,7 +1009,9 @@ async def test_multi_threshold_owner_email_auto_included( base_email_logger.internal_usage_cache = mock_cache with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): - await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info) + await base_email_logger.budget_alerts( + type="max_budget_alert", user_info=user_info + ) mock_send_email.assert_called_once() to_emails = mock_send_email.call_args[1]["to_email"] @@ -1002,7 +1033,7 @@ async def test_multi_threshold_malformed_keys_skipped( event_group=Litellm_EntityType.KEY, max_budget_alert_emails={ "fifty": ["finance@co.com"], # invalid - "50": ["finance@co.com"], # valid, crossed + "50": ["finance@co.com"], # valid, crossed }, ) @@ -1012,7 +1043,9 @@ async def test_multi_threshold_malformed_keys_skipped( base_email_logger.internal_usage_cache = mock_cache with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): - await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info) + await base_email_logger.budget_alerts( + type="max_budget_alert", user_info=user_info + ) # Only the valid "50" threshold should fire assert mock_send_email.call_count == 1 @@ -1041,7 +1074,9 @@ async def test_multi_threshold_empty_emails_only_owner( base_email_logger.internal_usage_cache = mock_cache with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): - await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info) + await base_email_logger.budget_alerts( + type="max_budget_alert", user_info=user_info + ) mock_send_email.assert_called_once() to_emails = mock_send_email.call_args[1]["to_email"] @@ -1067,11 +1102,13 @@ async def test_no_map_preserves_old_single_threshold( base_email_logger.internal_usage_cache = mock_cache with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): - await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info) + await base_email_logger.budget_alerts( + type="max_budget_alert", user_info=user_info + ) mock_send_email.assert_called_once() call_args = mock_send_email.call_args[1] assert call_args["to_email"] == ["test@example.com"] # Old path cache key has no threshold percentage cache_key = mock_cache.async_set_cache.call_args[1]["key"] - assert cache_key == "email_budget_alerts:max_budget_alert:test_user" \ No newline at end of file + assert cache_key == "email_budget_alerts:max_budget_alert:test_user" diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py index b07216921eb..dfb51c8bbc8 100644 --- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py +++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py @@ -87,7 +87,7 @@ async def test_send_email_success(mock_env_vars): async def test_send_email_missing_api_key(): # Remove the API key from environment before initializing logger original_key = os.environ.pop("RESEND_API_KEY", None) - + try: # Initialize the logger after removing the API key logger = ResendEmailLogger() @@ -104,16 +104,19 @@ async def test_send_email_missing_api_key(): mock_response.raise_for_status.return_value = None mock_response.status_code = 200 mock_response.json.return_value = {"id": "test_email_id"} - + mock_async_client = mock.AsyncMock() mock_async_client.post.return_value = mock_response - + # Directly inject the mock client to bypass any caching logger.async_httpx_client = mock_async_client # Send email await logger.send_email( - from_email=from_email, to_email=to_email, subject=subject, html_body=html_body + from_email=from_email, + to_email=to_email, + subject=subject, + html_body=html_body, ) # Verify the HTTP client was called with None as the API key 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/enterprise_callbacks/test_callback_controls.py b/tests/test_litellm/enterprise/enterprise_callbacks/test_callback_controls.py index b160ca5130c..9743e7bedc9 100644 --- a/tests/test_litellm/enterprise/enterprise_callbacks/test_callback_controls.py +++ b/tests/test_litellm/enterprise/enterprise_callbacks/test_callback_controls.py @@ -18,168 +18,282 @@ from litellm.types.utils import StandardCallbackDynamicParams class TestEnterpriseCallbackControls: - + @pytest.fixture def mock_premium_user(self): """Fixture to mock premium user check as True""" - with patch.object(EnterpriseCallbackControls, '_should_allow_dynamic_callback_disabling', return_value=True): + with patch.object( + EnterpriseCallbackControls, + "_should_allow_dynamic_callback_disabling", + return_value=True, + ): yield - - @pytest.fixture + + @pytest.fixture def mock_non_premium_user(self): """Fixture to mock premium user check as False""" - with patch.object(EnterpriseCallbackControls, '_should_allow_dynamic_callback_disabling', return_value=False): + with patch.object( + EnterpriseCallbackControls, + "_should_allow_dynamic_callback_disabling", + return_value=False, + ): yield @pytest.fixture def mock_request_headers(self): """Fixture to mock get_proxy_server_request_headers""" - with patch('enterprise.litellm_enterprise.enterprise_callbacks.callback_controls.get_proxy_server_request_headers') as mock_headers: + with patch( + "enterprise.litellm_enterprise.enterprise_callbacks.callback_controls.get_proxy_server_request_headers" + ) as mock_headers: yield mock_headers - def test_callback_disabled_langfuse_string(self, mock_premium_user, mock_request_headers): + def test_callback_disabled_langfuse_string( + self, mock_premium_user, mock_request_headers + ): """Test that 'langfuse' string callback is disabled when specified in headers""" mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) + + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) assert result is True - def test_callback_disabled_langfuse_customlogger(self, mock_premium_user, mock_request_headers): + def test_callback_disabled_langfuse_customlogger( + self, mock_premium_user, mock_request_headers + ): """Test that LangfusePromptManagement CustomLogger instance is disabled when 'langfuse' specified in headers""" mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - + langfuse_logger = LangfusePromptManagement() - result = EnterpriseCallbackControls.is_callback_disabled_dynamically(langfuse_logger, litellm_params, standard_callback_dynamic_params) + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + langfuse_logger, litellm_params, standard_callback_dynamic_params + ) assert result is True - def test_callback_disabled_s3_v2_string(self, mock_premium_user, mock_request_headers): + def test_callback_disabled_s3_v2_string( + self, mock_premium_user, mock_request_headers + ): """Test that 's3_v2' string callback is disabled when specified in headers""" mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "s3_v2"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("s3_v2", litellm_params, standard_callback_dynamic_params) + + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "s3_v2", litellm_params, standard_callback_dynamic_params + ) assert result is True - def test_callback_disabled_s3_v2_customlogger(self, mock_premium_user, mock_request_headers): + def test_callback_disabled_s3_v2_customlogger( + self, mock_premium_user, mock_request_headers + ): """Test that S3Logger CustomLogger instance is disabled when 's3_v2' specified in headers""" mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "s3_v2"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - + # Mock S3Logger to avoid async initialization issues - with patch('litellm.integrations.s3_v2.S3Logger.__init__', return_value=None): + with patch("litellm.integrations.s3_v2.S3Logger.__init__", return_value=None): s3_logger = S3Logger() - result = EnterpriseCallbackControls.is_callback_disabled_dynamically(s3_logger, litellm_params, standard_callback_dynamic_params) + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + s3_logger, litellm_params, standard_callback_dynamic_params + ) assert result is True - def test_callback_disabled_datadog_string(self, mock_premium_user, mock_request_headers): + def test_callback_disabled_datadog_string( + self, mock_premium_user, mock_request_headers + ): """Test that 'datadog' string callback is disabled when specified in headers""" mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "datadog"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) + + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "datadog", litellm_params, standard_callback_dynamic_params + ) assert result is True - def test_callback_disabled_datadog_customlogger(self, mock_premium_user, mock_request_headers): + def test_callback_disabled_datadog_customlogger( + self, mock_premium_user, mock_request_headers + ): """Test that DataDogLogger CustomLogger instance is disabled when 'datadog' specified in headers""" mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "datadog"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - + # Mock DataDogLogger to avoid async initialization issues - with patch('litellm.integrations.datadog.datadog.DataDogLogger.__init__', return_value=None): + with patch( + "litellm.integrations.datadog.datadog.DataDogLogger.__init__", + return_value=None, + ): datadog_logger = DataDogLogger() - result = EnterpriseCallbackControls.is_callback_disabled_dynamically(datadog_logger, litellm_params, standard_callback_dynamic_params) + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + datadog_logger, litellm_params, standard_callback_dynamic_params + ) assert result is True def test_multiple_callbacks_disabled(self, mock_premium_user, mock_request_headers): """Test that multiple callbacks can be disabled with comma-separated list""" - mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse,datadog,s3_v2"} + mock_request_headers.return_value = { + X_LITELLM_DISABLE_CALLBACKS: "langfuse,datadog,s3_v2" + } litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - # Test each callback is disabled - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) is True - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) is True - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("s3_v2", litellm_params, standard_callback_dynamic_params) is True - - # Test non-disabled callback is not disabled - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("prometheus", litellm_params, standard_callback_dynamic_params) is False - def test_callback_not_disabled_when_not_in_list(self, mock_premium_user, mock_request_headers): + # Test each callback is disabled + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) + is True + ) + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "datadog", litellm_params, standard_callback_dynamic_params + ) + is True + ) + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "s3_v2", litellm_params, standard_callback_dynamic_params + ) + is True + ) + + # Test non-disabled callback is not disabled + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "prometheus", litellm_params, standard_callback_dynamic_params + ) + is False + ) + + def test_callback_not_disabled_when_not_in_list( + self, mock_premium_user, mock_request_headers + ): """Test that callbacks not in the disabled list are not disabled""" mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) + + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "datadog", litellm_params, standard_callback_dynamic_params + ) assert result is False - def test_callback_not_disabled_when_no_header(self, mock_premium_user, mock_request_headers): + def test_callback_not_disabled_when_no_header( + self, mock_premium_user, mock_request_headers + ): """Test that callbacks are not disabled when the header is not present""" mock_request_headers.return_value = {} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) + + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) assert result is False - def test_callback_not_disabled_when_header_none(self, mock_premium_user, mock_request_headers): + def test_callback_not_disabled_when_header_none( + self, mock_premium_user, mock_request_headers + ): """Test that callbacks are not disabled when the header value is None""" mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: None} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) + + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) assert result is False - def test_non_premium_user_cannot_disable_callbacks(self, mock_non_premium_user, mock_request_headers): + def test_non_premium_user_cannot_disable_callbacks( + self, mock_non_premium_user, mock_request_headers + ): """Test that non-premium users cannot disable callbacks even with the header""" mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) + + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) assert result is False - def test_case_insensitive_callback_matching(self, mock_premium_user, mock_request_headers): + def test_case_insensitive_callback_matching( + self, mock_premium_user, mock_request_headers + ): """Test that callback matching is case insensitive""" - mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "LANGFUSE,DataDog"} + mock_request_headers.return_value = { + X_LITELLM_DISABLE_CALLBACKS: "LANGFUSE,DataDog" + } litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - + # Test lowercase callbacks are disabled - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) is True - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) is True + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) + is True + ) + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "datadog", litellm_params, standard_callback_dynamic_params + ) + is True + ) - def test_whitespace_handling_in_disabled_callbacks(self, mock_premium_user, mock_request_headers): + def test_whitespace_handling_in_disabled_callbacks( + self, mock_premium_user, mock_request_headers + ): """Test that whitespace around callback names is handled correctly""" - mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: " langfuse , datadog , s3_v2 "} + mock_request_headers.return_value = { + X_LITELLM_DISABLE_CALLBACKS: " langfuse , datadog , s3_v2 " + } litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) is True - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) is True - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("s3_v2", litellm_params, standard_callback_dynamic_params) is True - def test_custom_logger_not_in_registry(self, mock_premium_user, mock_request_headers): + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) + is True + ) + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "datadog", litellm_params, standard_callback_dynamic_params + ) + is True + ) + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "s3_v2", litellm_params, standard_callback_dynamic_params + ) + is True + ) + + def test_custom_logger_not_in_registry( + self, mock_premium_user, mock_request_headers + ): """Test that CustomLogger not in registry is not disabled""" - mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "unknown_logger"} + mock_request_headers.return_value = { + X_LITELLM_DISABLE_CALLBACKS: "unknown_logger" + } litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - + # Create a mock CustomLogger that's not in the registry class UnknownLogger(CustomLogger): pass - + unknown_logger = UnknownLogger() - result = EnterpriseCallbackControls.is_callback_disabled_dynamically(unknown_logger, litellm_params, standard_callback_dynamic_params) + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + unknown_logger, litellm_params, standard_callback_dynamic_params + ) assert result is False def test_exception_handling(self, mock_premium_user, mock_request_headers): @@ -188,32 +302,64 @@ class TestEnterpriseCallbackControls: mock_request_headers.side_effect = Exception("Test exception") litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) + + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) assert result is False - def test_callback_disabled_via_request_body_langfuse(self, mock_premium_user, mock_request_headers): + def test_callback_disabled_via_request_body_langfuse( + self, mock_premium_user, mock_request_headers + ): """Test that callbacks can be disabled via request body litellm_disabled_callbacks""" mock_request_headers.return_value = {} # No headers litellm_params = {"proxy_server_request": {"url": "test"}} - standard_callback_dynamic_params = StandardCallbackDynamicParams(litellm_disabled_callbacks=["langfuse"]) - - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) + standard_callback_dynamic_params = StandardCallbackDynamicParams( + litellm_disabled_callbacks=["langfuse"] + ) + + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) assert result is True - def test_callback_disabled_via_request_body_multiple(self, mock_premium_user, mock_request_headers): + def test_callback_disabled_via_request_body_multiple( + self, mock_premium_user, mock_request_headers + ): """Test that multiple callbacks can be disabled via request body""" mock_request_headers.return_value = {} # No headers litellm_params = {"proxy_server_request": {"url": "test"}} - standard_callback_dynamic_params = StandardCallbackDynamicParams(litellm_disabled_callbacks=["langfuse", "datadog", "s3_v2"]) - + standard_callback_dynamic_params = StandardCallbackDynamicParams( + litellm_disabled_callbacks=["langfuse", "datadog", "s3_v2"] + ) + # Test each callback is disabled - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) is True - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("datadog", litellm_params, standard_callback_dynamic_params) is True - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("s3_v2", litellm_params, standard_callback_dynamic_params) is True - + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) + is True + ) + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "datadog", litellm_params, standard_callback_dynamic_params + ) + is True + ) + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "s3_v2", litellm_params, standard_callback_dynamic_params + ) + is True + ) + # Test non-disabled callback is not disabled - assert EnterpriseCallbackControls.is_callback_disabled_dynamically("prometheus", litellm_params, standard_callback_dynamic_params) is False + assert ( + EnterpriseCallbackControls.is_callback_disabled_dynamically( + "prometheus", litellm_params, standard_callback_dynamic_params + ) + is False + ) def test_admin_can_disable_dynamic_callback_disabling(self, mock_request_headers): """ @@ -223,11 +369,13 @@ class TestEnterpriseCallbackControls: mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - + # Mock litellm.allow_dynamic_callback_disabling set to False - with patch('litellm.allow_dynamic_callback_disabling', False): - with patch('litellm.proxy.proxy_server.premium_user', True): - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) + with patch("litellm.allow_dynamic_callback_disabling", False): + with patch("litellm.proxy.proxy_server.premium_user", True): + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) assert result is False def test_admin_can_enable_dynamic_callback_disabling(self, mock_request_headers): @@ -238,14 +386,18 @@ class TestEnterpriseCallbackControls: mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - + # Mock litellm.allow_dynamic_callback_disabling set to True - with patch('litellm.allow_dynamic_callback_disabling', True): - with patch('litellm.proxy.proxy_server.premium_user', True): - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) + with patch("litellm.allow_dynamic_callback_disabling", True): + with patch("litellm.proxy.proxy_server.premium_user", True): + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) assert result is True - def test_default_admin_setting_allows_dynamic_callback_disabling(self, mock_request_headers): + def test_default_admin_setting_allows_dynamic_callback_disabling( + self, mock_request_headers + ): """ Test that when allow_dynamic_callback_disabling is not set, it defaults to True and allows dynamic callback disabling for premium users @@ -253,8 +405,10 @@ class TestEnterpriseCallbackControls: mock_request_headers.return_value = {X_LITELLM_DISABLE_CALLBACKS: "langfuse"} litellm_params = {"proxy_server_request": {"url": "test"}} standard_callback_dynamic_params = StandardCallbackDynamicParams() - + # litellm.allow_dynamic_callback_disabling should default to True - with patch('litellm.proxy.proxy_server.premium_user', True): - result = EnterpriseCallbackControls.is_callback_disabled_dynamically("langfuse", litellm_params, standard_callback_dynamic_params) + with patch("litellm.proxy.proxy_server.premium_user", True): + result = EnterpriseCallbackControls.is_callback_disabled_dynamically( + "langfuse", litellm_params, standard_callback_dynamic_params + ) assert result is True 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..74a18a9b652 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 @@ -18,11 +18,17 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( DECODED_UNIFIED_INPUT_FILE_ID = "litellm_proxy:application/octet-stream;unified_id,test-uuid;target_model_names,azure-gpt-4" -B64_UNIFIED_INPUT_FILE_ID = base64.urlsafe_b64encode(DECODED_UNIFIED_INPUT_FILE_ID.encode()).decode().rstrip("=") +B64_UNIFIED_INPUT_FILE_ID = ( + base64.urlsafe_b64encode(DECODED_UNIFIED_INPUT_FILE_ID.encode()) + .decode() + .rstrip("=") +) RAW_INPUT_FILE_ID = "file-raw-provider-abc123" DECODED_UNIFIED_BATCH_ID = "litellm_proxy;model_id:model-xyz;llm_batch_id:batch-123" -B64_UNIFIED_BATCH_ID = base64.urlsafe_b64encode(DECODED_UNIFIED_BATCH_ID.encode()).decode().rstrip("=") +B64_UNIFIED_BATCH_ID = ( + base64.urlsafe_b64encode(DECODED_UNIFIED_BATCH_ID.encode()).decode().rstrip("=") +) @pytest.mark.asyncio @@ -55,10 +61,16 @@ async def test_should_resolve_raw_input_file_id_to_unified(): mock_managed_file.unified_file_id = B64_UNIFIED_INPUT_FILE_ID mock_prisma = MagicMock() - mock_prisma.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=mock_db_object) - mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=mock_managed_file) + mock_prisma.db.litellm_managedobjecttable.find_first = AsyncMock( + return_value=mock_db_object + ) + mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock( + return_value=mock_managed_file + ) - from litellm.proxy.openai_files_endpoints.common_utils import get_batch_from_database + from litellm.proxy.openai_files_endpoints.common_utils import ( + get_batch_from_database, + ) _, response = await get_batch_from_database( batch_id=B64_UNIFIED_BATCH_ID, 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/enterprise/proxy/test_managed_files_access_check.py b/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py index 8cb642b7a44..18b90a56a6e 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py @@ -144,6 +144,7 @@ async def test_check_batch_cost_should_call_afile_content_directly_with_credenti # Mock the batch response (completed, with output file) from litellm.types.utils import LiteLLMBatch + batch_response = LiteLLMBatch( id="batch-123", completion_window="24h", @@ -201,9 +202,11 @@ async def test_check_batch_cost_should_call_afile_content_directly_with_credenti # Verify the DB update writes batch_processed, status, and file_object mock_prisma.db.litellm_managedobjecttable.update.assert_called_once() - update_call_kwargs = mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs + update_call_kwargs = ( + mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs + ) assert update_call_kwargs["data"]["batch_processed"] is True assert update_call_kwargs["data"]["status"] == "complete" - assert "file_object" in update_call_kwargs["data"], ( - "file_object must be written to DB so list_batches reads updated status" - ) + assert ( + "file_object" in update_call_kwargs["data"] + ), "file_object must be written to DB so list_batches reads updated status" diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index 9526304aff0..d4855d5b130 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -110,10 +110,9 @@ async def test_should_pass_credentials_to_afile_retrieve(): mock_afile_retrieve = AsyncMock(return_value=_make_file_object("file-output-abc")) - with patch( - "litellm.afile_retrieve", mock_afile_retrieve - ), patch( - "litellm.proxy.proxy_server.llm_router", mock_router + with ( + patch("litellm.afile_retrieve", mock_afile_retrieve), + patch("litellm.proxy.proxy_server.llm_router", mock_router), ): await managed_files.async_post_call_success_hook( data={}, @@ -128,7 +127,9 @@ async def test_should_pass_credentials_to_afile_retrieve(): f"afile_retrieve must receive api_key from router credentials. " f"Got kwargs: {call_kwargs.kwargs}" ) - assert call_kwargs.kwargs.get("api_base") == "https://my-azure.openai.azure.com/", ( + assert ( + call_kwargs.kwargs.get("api_base") == "https://my-azure.openai.azure.com/" + ), ( f"afile_retrieve must receive api_base from router credentials. " f"Got kwargs: {call_kwargs.kwargs}" ) @@ -150,10 +151,9 @@ async def test_should_fallback_when_no_router(): mock_afile_retrieve = AsyncMock(return_value=_make_file_object("file-output-abc")) - with patch( - "litellm.afile_retrieve", mock_afile_retrieve - ), patch( - "litellm.proxy.proxy_server.llm_router", None + with ( + patch("litellm.afile_retrieve", mock_afile_retrieve), + patch("litellm.proxy.proxy_server.llm_router", None), ): await managed_files.async_post_call_success_hook( data={}, diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index d09c4ac2c38..ff6c87132a6 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -1086,9 +1086,12 @@ class TestCustomGuardrailSpendLogMatchRedaction: ][0]["match"] == "[REDACTED]" ) - assert raw["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][0][ - "match" - ] == "GG" + assert ( + raw["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][0][ + "match" + ] + == "GG" + ) def test_add_standard_logging_redacts_regex_field(self): cg = CustomGuardrail(guardrail_name="test-rail") 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_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index d6281703a0a..7913f587b28 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -2003,23 +2003,19 @@ async def test_azure_streaming_role_preserved_with_include_usage(sync_mode: bool chunks.append(chunk) # The prompt_filter chunk should be forwarded with choices=[] - assert len(chunks[0].choices) == 0, ( - f"Expected prompt_filter chunk with choices=[], got {len(chunks[0].choices)} choices" - ) + assert ( + len(chunks[0].choices) == 0 + ), f"Expected prompt_filter chunk with choices=[], got {len(chunks[0].choices)} choices" # At least one chunk must have role='assistant' in its delta has_role = any( - len(c.choices) > 0 - and getattr(c.choices[0].delta, "role", None) == "assistant" + len(c.choices) > 0 and getattr(c.choices[0].delta, "role", None) == "assistant" for c in chunks ) assert has_role, ( "No chunk contained role='assistant' in delta (issue #24221). " "Chunk deltas: " - + str([ - c.choices[0].delta if c.choices else "no choices" - for c in chunks - ]) + + str([c.choices[0].delta if c.choices else "no choices" for c in chunks]) ) diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index e1fe4befb57..a409ebbe922 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -1653,9 +1653,7 @@ def test_max_effort_rejected_for_opus_45(): messages = [{"role": "user", "content": "Test"}] - with pytest.raises( - ValueError, match="effort='max' is not supported by this model" - ): + with pytest.raises(ValueError, match="effort='max' is not supported by this model"): optional_params = {"output_config": {"effort": "max"}} config.transform_request( model="claude-opus-4-5-20251101", @@ -2217,9 +2215,7 @@ def test_max_effort_rejected_for_sonnet_46(): config = AnthropicConfig() messages = [{"role": "user", "content": "Test"}] - with pytest.raises( - ValueError, match="effort='max' is not supported by this model" - ): + with pytest.raises(ValueError, match="effort='max' is not supported by this model"): config.transform_request( model="claude-sonnet-4-6-20260219", messages=messages, diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index 42efde90926..3fc0dea4748 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -2327,7 +2327,8 @@ class TestAnthropicStreamWrapperToolArgs: def _find_tool_deltas(self, events): return [ - e for e in events + e + for e in events if isinstance(e, dict) and e.get("type") == "content_block_delta" and isinstance(e.get("delta"), dict) @@ -2375,7 +2376,6 @@ class TestAnthropicStreamWrapperToolArgs: assert parsed == {"city": "Tokyo"} - def test_translate_anthropic_tool_choice_none(): """ Regression test for issue #24443. 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/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index 7a2a6f56d6f..108983c6766 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -802,7 +802,10 @@ async def test_promote_message_start_cache_when_message_stop_omits_cache_fields( "delta": {"stop_reason": "end_turn", "stop_sequence": None}, "usage": {"input_tokens": 10, "output_tokens": 181}, } - yield {"type": "message_stop", "usage": {"input_tokens": 10, "output_tokens": 181}} + yield { + "type": "message_stop", + "usage": {"input_tokens": 10, "output_tokens": 181}, + } merged: list[dict] = [] async for chunk in cfg._promote_message_stop_usage(_stream()): @@ -856,7 +859,11 @@ async def test_unified_bedrock_messages_cache_on_start_only_never_negative_cost( }, }, } - yield {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}} + yield { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + } yield { "type": "content_block_delta", "index": 0, @@ -868,7 +875,10 @@ async def test_unified_bedrock_messages_cache_on_start_only_never_negative_cost( "delta": {"stop_reason": "end_turn", "stop_sequence": None}, "usage": {"output_tokens": 181, "input_tokens": 10}, } - yield {"type": "message_stop", "usage": {"input_tokens": 10, "output_tokens": 181}} + yield { + "type": "message_stop", + "usage": {"input_tokens": 10, "output_tokens": 181}, + } logging_obj = LiteLLMLoggingObj( model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", diff --git a/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py b/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py index 70707c95de9..8f2ca197c22 100644 --- a/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py +++ b/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py @@ -233,3 +233,35 @@ class TestChatGPTResponsesAPITransformation: ) assert parsed.output_text == "Hello!" + + def test_chatgpt_non_stream_sse_response_ignores_non_string_delta(self): + config = ChatGPTResponsesAPIConfig() + response_payload = { + "id": "resp_test", + "object": "response", + "created_at": 1700000000, + "status": "completed", + "model": "gpt-5.4", + "output": [], + } + sse_body = "\n".join( + [ + f"data: {json.dumps({'type': 'response.output_text.delta', 'delta': 123})}", + f"data: {json.dumps({'type': 'response.output_text.delta', 'delta': 'Hello'})}", + f"data: {json.dumps({'type': 'response.completed', 'response': response_payload})}", + "data: [DONE]", + "", + ] + ) + raw_response = httpx.Response( + 200, headers={"content-type": "text/event-stream"}, text=sse_body + ) + logging_obj = MagicMock() + + parsed = config.transform_response_api_response( + model="chatgpt/gpt-5.4", + raw_response=raw_response, + logging_obj=logging_obj, + ) + + assert parsed.output_text == "Hello" diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 6924eb8d3d9..6f20d132a3a 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -99,7 +99,9 @@ def test_fingerprint_agentic_tools_is_deterministic(): tools_a = {"tool_calls": [{"id": "1", "input": {"q": "abc"}, "name": "web_search"}]} tools_b = {"tool_calls": [{"name": "web_search", "input": {"q": "abc"}, "id": "1"}]} - assert handler._fingerprint_agentic_tools(tools_a) == handler._fingerprint_agentic_tools(tools_b) + assert handler._fingerprint_agentic_tools( + tools_a + ) == handler._fingerprint_agentic_tools(tools_b) @pytest.mark.asyncio 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/moonshot/test_moonshot_chat_transformation.py b/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py index b4744a7ed18..e398699bfdd 100644 --- a/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py +++ b/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py @@ -666,7 +666,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.""" diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py index dae87842832..35005c57369 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py @@ -547,7 +547,12 @@ class TestOpenAIResponsesAPIConfig: """Base helper strips ``namespace`` from custom_tool_call for every provider path.""" inp = [ {"type": "function_call", "call_id": "a", "name": "f", "namespace": "keep"}, - {"type": "custom_tool_call", "call_id": "b", "name": "c", "namespace": "drop"}, + { + "type": "custom_tool_call", + "call_id": "b", + "name": "c", + "namespace": "drop", + }, ] out = BaseResponsesAPIConfig.strip_custom_tool_call_namespace_from_responses_input( inp diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 0118b39ddd1..c056a347376 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -3499,7 +3499,12 @@ def test_video_metadata_supported_for_all_gemini_models(): } ] - for model in ["gemini-1.5-pro", "gemini-2.5-flash", "gemini-2.5-pro", "gemini-3-pro-preview"]: + for model in [ + "gemini-1.5-pro", + "gemini-2.5-flash", + "gemini-2.5-pro", + "gemini-3-pro-preview", + ]: contents = _gemini_convert_messages_with_history(messages=messages, model=model) file_part = None @@ -3509,19 +3514,25 @@ def test_video_metadata_supported_for_all_gemini_models(): break assert file_part is not None, f"{model}: file part should exist" - assert "video_metadata" in file_part, f"{model}: video_metadata should be present" + assert ( + "video_metadata" in file_part + ), f"{model}: video_metadata should be present" assert file_part["video_metadata"]["fps"] == 5, f"{model}: fps should be 5" # Per-part media_resolution is Gemini 3+ only; 2.x uses generation_config global for model in ["gemini-3-pro-preview"]: contents = _gemini_convert_messages_with_history(messages=messages, model=model) file_part = next(p for p in contents[0]["parts"] if "file_data" in p) - assert "media_resolution" in file_part, f"{model}: media_resolution should be present" + assert ( + "media_resolution" in file_part + ), f"{model}: media_resolution should be present" for model in ["gemini-1.5-pro", "gemini-2.5-flash", "gemini-2.5-pro"]: contents = _gemini_convert_messages_with_history(messages=messages, model=model) file_part = next(p for p in contents[0]["parts"] if "file_data" in p) - assert "media_resolution" not in file_part, f"{model}: per-part media_resolution should not be set" + assert ( + "media_resolution" not in file_part + ), f"{model}: per-part media_resolution should not be set" def test_chunk_parser_handles_prompt_feedback_block(): @@ -4154,8 +4165,9 @@ def test_vertex_ai_usage_metadata_with_document_tokens_in_prompt(): # DOCUMENT tokens should be included in text_tokens: 8 (TEXT) + 774 (DOCUMENT) = 782 assert result.prompt_tokens_details is not None - assert result.prompt_tokens_details.text_tokens == 782, \ - "DOCUMENT modality tokens should be added to text_tokens (8 TEXT + 774 DOCUMENT = 782)" + assert ( + result.prompt_tokens_details.text_tokens == 782 + ), "DOCUMENT modality tokens should be added to text_tokens (8 TEXT + 774 DOCUMENT = 782)" # Verify completion token details assert result.completion_tokens_details is not None @@ -4190,8 +4202,9 @@ def test_vertex_ai_usage_metadata_with_document_tokens_cached(): # DOCUMENT cached tokens map to cached_text_tokens, so: # text_tokens = (8 TEXT + 774 DOCUMENT) - 400 cached = 382 - assert result.prompt_tokens_details.text_tokens == 382, \ - "text_tokens should be (8 + 774) - 400 cached = 382" + assert ( + result.prompt_tokens_details.text_tokens == 382 + ), "text_tokens should be (8 + 774) - 400 cached = 382" assert result.prompt_tokens_details.cached_tokens == 400 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 a0bfbff4222..4cc764f3c6d 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 @@ -497,31 +497,39 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401(): oauth_server.auth_type = MCPAuth.oauth2 oauth_server.needs_user_oauth_token = True - with patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - new_callable=AsyncMock, - return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", - ), patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", - new_callable=AsyncMock, - return_value=False, - ), patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", - new_callable=AsyncMock, - return_value=None, - ) as mock_get_stored_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", - return_value=oauth_server, - ), patch.object( - session_manager, - "handle_request", - new_callable=AsyncMock, - ) as mock_handle_request: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + new_callable=AsyncMock, + return_value=None, + ) as mock_get_stored_token, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + return_value=oauth_server, + ), + patch.object( + session_manager, + "handle_request", + new_callable=AsyncMock, + ) as mock_handle_request, + ): with pytest.raises(HTTPException) as exc_info: await handle_streamable_http_mcp(scope, receive, send) @@ -562,31 +570,39 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401(): oauth_server.auth_type = MCPAuth.oauth2 oauth_server.needs_user_oauth_token = True - with patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - new_callable=AsyncMock, - return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", - ), patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", - new_callable=AsyncMock, - return_value=False, - ), patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", - new_callable=AsyncMock, - return_value={"Authorization": "Bearer cached-token"}, - ) as mock_get_stored_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", - return_value=oauth_server, - ), patch.object( - session_manager, - "handle_request", - new_callable=AsyncMock, - ) as mock_handle_request: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + new_callable=AsyncMock, + return_value={"Authorization": "Bearer cached-token"}, + ) as mock_get_stored_token, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + return_value=oauth_server, + ), + patch.object( + session_manager, + "handle_request", + new_callable=AsyncMock, + ) as mock_handle_request, + ): await handle_streamable_http_mcp(scope, receive, send) assert mock_get_stored_token.await_count == 1 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/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 9c43ebcbe79..2ed8bab3285 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -1752,7 +1752,11 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): from starlette.datastructures import URL from starlette.requests import Request - from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy._types import ( + LiteLLM_TeamTableCachedObj, + LitellmUserRoles, + UserAPIKeyAuth, + ) from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder api_key = "sk-test-team-metadata-refresh" @@ -1833,16 +1837,17 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): request_data={}, ) - assert result.team_metadata == {"guardrails": ["test-guardrail-333"]}, ( - f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}" - ) + assert result.team_metadata == { + "guardrails": ["test-guardrail-333"] + }, f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}" finally: for k, v in _originals.items(): setattr(_proxy_server_mod, k, v) - + + # --------------------------------------------------------------------------- - + # _run_centralized_common_checks — centralized authz gate # --------------------------------------------------------------------------- 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..6a46bcb2adc 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 @@ -39,8 +39,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 +102,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 +118,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 +187,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 +208,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 +226,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/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index a1ba7ecd677..2ed3b87c13c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -1949,11 +1949,13 @@ async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker): assert exc.value.status_code == 403 # Critical: no delete_many calls should have executed. - assert not hasattr( - mock_prisma_client.db.litellm_verificationtoken.delete_many, "mock_calls" - ) or len( - mock_prisma_client.db.litellm_verificationtoken.delete_many.mock_calls - ) == 0 + assert ( + not hasattr( + mock_prisma_client.db.litellm_verificationtoken.delete_many, "mock_calls" + ) + or len(mock_prisma_client.db.litellm_verificationtoken.delete_many.mock_calls) + == 0 + ) @pytest.mark.asyncio 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 c9be00afed0..0e7e0ba3211 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 @@ -1092,9 +1092,9 @@ class TestVertexAIPassThroughHandler: assert result is not None assert result["result"] is not None - assert result["kwargs"].get("custom_llm_provider") == "gemini", ( - "Google AI Studio embedContent URLs must set custom_llm_provider=gemini, not vertex_ai" - ) + assert ( + result["kwargs"].get("custom_llm_provider") == "gemini" + ), "Google AI Studio embedContent URLs must set custom_llm_provider=gemini, not vertex_ai" assert result["kwargs"].get("model") == "gemini-embedding-2-preview" mock_completion_cost.assert_called_once() diff --git a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py index a65462d3f9b..959201eda28 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -166,9 +166,9 @@ def test_anthropic_provider_fields_support_byok(): "Anthropic api_key must be optional so admins can configure BYOK models " "without entering a key. See BYOK tutorial." ) - assert fields_by_key["api_key"].get("tooltip"), ( - "Anthropic api_key must have a tooltip explaining the BYOK use case." - ) + assert fields_by_key["api_key"].get( + "tooltip" + ), "Anthropic api_key must have a tooltip explaining the BYOK use case." assert "api_base" in fields_by_key, ( "Anthropic provider form must expose api_base so cloud customers " "can override the upstream URL without env var access." @@ -176,16 +176,16 @@ def test_anthropic_provider_fields_support_byok(): api_base_field = fields_by_key["api_base"] assert api_base_field["required"] is False assert api_base_field["field_type"] == "text" - assert api_base_field.get("tooltip"), ( - "api_base should have a tooltip explaining it is optional." - ) + assert api_base_field.get( + "tooltip" + ), "api_base should have a tooltip explaining it is optional." # UI forms render fields in credential_fields order; api_base should come first # so an admin sees the URL override before the key field. field_order = [f["key"] for f in anthropic["credential_fields"]] - assert field_order.index("api_base") < field_order.index("api_key"), ( - "api_base must appear before api_key in credential_fields (matches AI21 and ANTHROPIC_TEXT convention)." - ) + assert field_order.index("api_base") < field_order.index( + "api_key" + ), "api_base must appear before api_key in credential_fields (matches AI21 and ANTHROPIC_TEXT convention)." def test_public_model_hub_with_healthy_model(): 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 1e5c688e366..794051877f3 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -3326,7 +3326,9 @@ async def test_team_guardrail_merges_with_global_policy(): policy_registry = get_policy_registry() policy_registry._policies = { "global-policy": Policy( - guardrails=PolicyGuardrails(add=["policy-guardrail-1", "policy-guardrail-2"]), + guardrails=PolicyGuardrails( + add=["policy-guardrail-1", "policy-guardrail-2"] + ), ), } policy_registry._initialized = True @@ -3347,14 +3349,18 @@ async def test_team_guardrail_merges_with_global_policy(): guardrails = data["metadata"].get("guardrails", []) - assert "team-direct-guardrail" in guardrails, \ - f"Team guardrail missing from merged list: {guardrails}" - assert "policy-guardrail-1" in guardrails, \ - f"policy-guardrail-1 missing: {guardrails}" - assert "policy-guardrail-2" in guardrails, \ - f"policy-guardrail-2 missing: {guardrails}" - assert len(guardrails) == len(set(guardrails)), \ - f"Duplicates in guardrails list: {guardrails}" + assert ( + "team-direct-guardrail" in guardrails + ), f"Team guardrail missing from merged list: {guardrails}" + assert ( + "policy-guardrail-1" in guardrails + ), f"policy-guardrail-1 missing: {guardrails}" + assert ( + "policy-guardrail-2" in guardrails + ), f"policy-guardrail-2 missing: {guardrails}" + assert len(guardrails) == len( + set(guardrails) + ), f"Duplicates in guardrails list: {guardrails}" # Verify get_guardrail_from_metadata returns the merged list even # when litellm_metadata is present (the bug: it returned [] before fix) @@ -3365,9 +3371,9 @@ async def test_team_guardrail_merges_with_global_policy(): dummy = _DummyGuardrail(guardrail_name="team-direct-guardrail") returned = dummy.get_guardrail_from_metadata(data) - assert "team-direct-guardrail" in returned, ( - f"get_guardrail_from_metadata shadowed by litellm_metadata; got: {returned}" - ) + assert ( + "team-direct-guardrail" in returned + ), f"get_guardrail_from_metadata shadowed by litellm_metadata; got: {returned}" finally: policy_registry._policies = {} @@ -3396,9 +3402,10 @@ async def test_get_guardrail_from_metadata_prefers_metadata_over_litellm_metadat } result = dummy.get_guardrail_from_metadata(data) - assert result == ["my-guardrail", "other-guardrail"], ( - f"Expected guardrails from metadata, got: {result}" - ) + assert result == [ + "my-guardrail", + "other-guardrail", + ], f"Expected guardrails from metadata, got: {result}" def test_get_guardrail_from_metadata_reads_litellm_metadata_when_no_metadata(): @@ -3419,6 +3426,6 @@ def test_get_guardrail_from_metadata_reads_litellm_metadata_when_no_metadata(): } result = dummy.get_guardrail_from_metadata(data) - assert result == ["my-guardrail"], ( - f"Expected guardrails from litellm_metadata fallback, got: {result}" - ) + assert result == [ + "my-guardrail" + ], f"Expected guardrails from litellm_metadata fallback, got: {result}" 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/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_utils.py b/tests/test_litellm/test_utils.py index 93c61e003d9..5c6150d6c9b 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -2825,9 +2825,9 @@ def test_gemini_embedding_2_ga_in_cost_map(): assert info.get("input_cost_per_audio_per_second") == 0.00016 assert info.get("input_cost_per_video_per_second") == 0.00079 if provider in ("vertex_ai-embedding-models", "vertex_ai"): - assert info.get("uses_embed_content") is True, ( - f"{key} must have uses_embed_content=true for correct Vertex AI routing" - ) + assert ( + info.get("uses_embed_content") is True + ), f"{key} must have uses_embed_content=true for correct Vertex AI routing" def test_gemini_lyria_3_preview_models_in_cost_map(): 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 -