mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
style: run black formatter on files from staging merge
This commit is contained in:
parent
94fbcac6f4
commit
af24a3f720
43 changed files with 1920 additions and 1216 deletions
|
|
@ -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()
|
||||
|
|
|
|||
512
.github/workflows/run_llm_translation_tests.py
vendored
512
.github/workflows/run_llm_translation_tests.py
vendored
|
|
@ -16,64 +16,75 @@ from pathlib import Path
|
|||
import json
|
||||
from typing import Dict, List, Tuple, Optional
|
||||
|
||||
|
||||
# ANSI color codes for terminal output
|
||||
class Colors:
|
||||
GREEN = '\033[92m'
|
||||
RED = '\033[91m'
|
||||
YELLOW = '\033[93m'
|
||||
BLUE = '\033[94m'
|
||||
PURPLE = '\033[95m'
|
||||
CYAN = '\033[96m'
|
||||
RESET = '\033[0m'
|
||||
BOLD = '\033[1m'
|
||||
GREEN = "\033[92m"
|
||||
RED = "\033[91m"
|
||||
YELLOW = "\033[93m"
|
||||
BLUE = "\033[94m"
|
||||
PURPLE = "\033[95m"
|
||||
CYAN = "\033[96m"
|
||||
RESET = "\033[0m"
|
||||
BOLD = "\033[1m"
|
||||
|
||||
|
||||
def print_colored(message: str, color: str = Colors.RESET):
|
||||
"""Print colored message to terminal"""
|
||||
print(f"{color}{message}{Colors.RESET}")
|
||||
|
||||
|
||||
def get_provider_from_test_file(test_file: str) -> str:
|
||||
"""Map test file names to provider names"""
|
||||
provider_mapping = {
|
||||
'test_anthropic': 'Anthropic',
|
||||
'test_azure': 'Azure',
|
||||
'test_bedrock': 'AWS Bedrock',
|
||||
'test_openai': 'OpenAI',
|
||||
'test_vertex': 'Google Vertex AI',
|
||||
'test_gemini': 'Google Vertex AI',
|
||||
'test_cohere': 'Cohere',
|
||||
'test_databricks': 'Databricks',
|
||||
'test_groq': 'Groq',
|
||||
'test_together': 'Together AI',
|
||||
'test_mistral': 'Mistral',
|
||||
'test_deepseek': 'DeepSeek',
|
||||
'test_replicate': 'Replicate',
|
||||
'test_huggingface': 'HuggingFace',
|
||||
'test_fireworks': 'Fireworks AI',
|
||||
'test_perplexity': 'Perplexity',
|
||||
'test_cloudflare': 'Cloudflare',
|
||||
'test_voyage': 'Voyage AI',
|
||||
'test_xai': 'xAI',
|
||||
'test_nvidia': 'NVIDIA',
|
||||
'test_watsonx': 'IBM watsonx',
|
||||
'test_azure_ai': 'Azure AI',
|
||||
'test_snowflake': 'Snowflake',
|
||||
'test_infinity': 'Infinity',
|
||||
'test_jina': 'Jina AI',
|
||||
'test_deepgram': 'Deepgram',
|
||||
'test_clarifai': 'Clarifai',
|
||||
'test_triton': 'Triton',
|
||||
"test_anthropic": "Anthropic",
|
||||
"test_azure": "Azure",
|
||||
"test_bedrock": "AWS Bedrock",
|
||||
"test_openai": "OpenAI",
|
||||
"test_vertex": "Google Vertex AI",
|
||||
"test_gemini": "Google Vertex AI",
|
||||
"test_cohere": "Cohere",
|
||||
"test_databricks": "Databricks",
|
||||
"test_groq": "Groq",
|
||||
"test_together": "Together AI",
|
||||
"test_mistral": "Mistral",
|
||||
"test_deepseek": "DeepSeek",
|
||||
"test_replicate": "Replicate",
|
||||
"test_huggingface": "HuggingFace",
|
||||
"test_fireworks": "Fireworks AI",
|
||||
"test_perplexity": "Perplexity",
|
||||
"test_cloudflare": "Cloudflare",
|
||||
"test_voyage": "Voyage AI",
|
||||
"test_xai": "xAI",
|
||||
"test_nvidia": "NVIDIA",
|
||||
"test_watsonx": "IBM watsonx",
|
||||
"test_azure_ai": "Azure AI",
|
||||
"test_snowflake": "Snowflake",
|
||||
"test_infinity": "Infinity",
|
||||
"test_jina": "Jina AI",
|
||||
"test_deepgram": "Deepgram",
|
||||
"test_clarifai": "Clarifai",
|
||||
"test_triton": "Triton",
|
||||
}
|
||||
|
||||
|
||||
for key, provider in provider_mapping.items():
|
||||
if key in test_file:
|
||||
return provider
|
||||
|
||||
|
||||
# For cross-provider test files
|
||||
if any(name in test_file for name in ['test_optional_params', 'test_prompt_factory',
|
||||
'test_router', 'test_text_completion']):
|
||||
return f'Cross-Provider Tests ({test_file})'
|
||||
|
||||
return 'Other Tests'
|
||||
if any(
|
||||
name in test_file
|
||||
for name in [
|
||||
"test_optional_params",
|
||||
"test_prompt_factory",
|
||||
"test_router",
|
||||
"test_text_completion",
|
||||
]
|
||||
):
|
||||
return f"Cross-Provider Tests ({test_file})"
|
||||
|
||||
return "Other Tests"
|
||||
|
||||
|
||||
def format_duration(seconds: float) -> str:
|
||||
"""Format duration in human-readable format"""
|
||||
|
|
@ -89,290 +100,355 @@ def format_duration(seconds: float) -> str:
|
|||
return f"{hours}h {minutes}m"
|
||||
|
||||
|
||||
def generate_markdown_report(junit_xml_path: str, output_path: str, tag: str = None, commit: str = None):
|
||||
def generate_markdown_report(
|
||||
junit_xml_path: str, output_path: str, tag: str = None, commit: str = None
|
||||
):
|
||||
"""Generate a beautiful markdown report from JUnit XML"""
|
||||
try:
|
||||
tree = ET.parse(junit_xml_path)
|
||||
root = tree.getroot()
|
||||
|
||||
|
||||
# Handle both testsuite and testsuites root
|
||||
if root.tag == 'testsuites':
|
||||
suites = root.findall('testsuite')
|
||||
if root.tag == "testsuites":
|
||||
suites = root.findall("testsuite")
|
||||
else:
|
||||
suites = [root]
|
||||
|
||||
|
||||
# Overall statistics
|
||||
total_tests = 0
|
||||
total_failures = 0
|
||||
total_errors = 0
|
||||
total_skipped = 0
|
||||
total_time = 0.0
|
||||
|
||||
|
||||
# Provider breakdown
|
||||
provider_stats = defaultdict(lambda: {'passed': 0, 'failed': 0, 'skipped': 0, 'errors': 0, 'time': 0.0})
|
||||
provider_stats = defaultdict(
|
||||
lambda: {"passed": 0, "failed": 0, "skipped": 0, "errors": 0, "time": 0.0}
|
||||
)
|
||||
provider_tests = defaultdict(list)
|
||||
|
||||
|
||||
for suite in suites:
|
||||
total_tests += int(suite.get('tests', 0))
|
||||
total_failures += int(suite.get('failures', 0))
|
||||
total_errors += int(suite.get('errors', 0))
|
||||
total_skipped += int(suite.get('skipped', 0))
|
||||
total_time += float(suite.get('time', 0))
|
||||
|
||||
for testcase in suite.findall('testcase'):
|
||||
classname = testcase.get('classname', '')
|
||||
test_name = testcase.get('name', '')
|
||||
test_time = float(testcase.get('time', 0))
|
||||
|
||||
total_tests += int(suite.get("tests", 0))
|
||||
total_failures += int(suite.get("failures", 0))
|
||||
total_errors += int(suite.get("errors", 0))
|
||||
total_skipped += int(suite.get("skipped", 0))
|
||||
total_time += float(suite.get("time", 0))
|
||||
|
||||
for testcase in suite.findall("testcase"):
|
||||
classname = testcase.get("classname", "")
|
||||
test_name = testcase.get("name", "")
|
||||
test_time = float(testcase.get("time", 0))
|
||||
|
||||
# Extract test file name from classname
|
||||
if '.' in classname:
|
||||
parts = classname.split('.')
|
||||
test_file = parts[-2] if len(parts) > 1 else 'unknown'
|
||||
if "." in classname:
|
||||
parts = classname.split(".")
|
||||
test_file = parts[-2] if len(parts) > 1 else "unknown"
|
||||
else:
|
||||
test_file = 'unknown'
|
||||
|
||||
test_file = "unknown"
|
||||
|
||||
provider = get_provider_from_test_file(test_file)
|
||||
provider_stats[provider]['time'] += test_time
|
||||
|
||||
provider_stats[provider]["time"] += test_time
|
||||
|
||||
# Check test status
|
||||
if testcase.find('failure') is not None:
|
||||
provider_stats[provider]['failed'] += 1
|
||||
failure = testcase.find('failure')
|
||||
failure_msg = failure.get('message', '') if failure is not None else ''
|
||||
provider_tests[provider].append({
|
||||
'name': test_name,
|
||||
'status': 'FAILED',
|
||||
'time': test_time,
|
||||
'message': failure_msg
|
||||
})
|
||||
elif testcase.find('error') is not None:
|
||||
provider_stats[provider]['errors'] += 1
|
||||
error = testcase.find('error')
|
||||
error_msg = error.get('message', '') if error is not None else ''
|
||||
provider_tests[provider].append({
|
||||
'name': test_name,
|
||||
'status': 'ERROR',
|
||||
'time': test_time,
|
||||
'message': error_msg
|
||||
})
|
||||
elif testcase.find('skipped') is not None:
|
||||
provider_stats[provider]['skipped'] += 1
|
||||
skip = testcase.find('skipped')
|
||||
skip_msg = skip.get('message', '') if skip is not None else ''
|
||||
provider_tests[provider].append({
|
||||
'name': test_name,
|
||||
'status': 'SKIPPED',
|
||||
'time': test_time,
|
||||
'message': skip_msg
|
||||
})
|
||||
if testcase.find("failure") is not None:
|
||||
provider_stats[provider]["failed"] += 1
|
||||
failure = testcase.find("failure")
|
||||
failure_msg = (
|
||||
failure.get("message", "") if failure is not None else ""
|
||||
)
|
||||
provider_tests[provider].append(
|
||||
{
|
||||
"name": test_name,
|
||||
"status": "FAILED",
|
||||
"time": test_time,
|
||||
"message": failure_msg,
|
||||
}
|
||||
)
|
||||
elif testcase.find("error") is not None:
|
||||
provider_stats[provider]["errors"] += 1
|
||||
error = testcase.find("error")
|
||||
error_msg = error.get("message", "") if error is not None else ""
|
||||
provider_tests[provider].append(
|
||||
{
|
||||
"name": test_name,
|
||||
"status": "ERROR",
|
||||
"time": test_time,
|
||||
"message": error_msg,
|
||||
}
|
||||
)
|
||||
elif testcase.find("skipped") is not None:
|
||||
provider_stats[provider]["skipped"] += 1
|
||||
skip = testcase.find("skipped")
|
||||
skip_msg = skip.get("message", "") if skip is not None else ""
|
||||
provider_tests[provider].append(
|
||||
{
|
||||
"name": test_name,
|
||||
"status": "SKIPPED",
|
||||
"time": test_time,
|
||||
"message": skip_msg,
|
||||
}
|
||||
)
|
||||
else:
|
||||
provider_stats[provider]['passed'] += 1
|
||||
provider_tests[provider].append({
|
||||
'name': test_name,
|
||||
'status': 'PASSED',
|
||||
'time': test_time,
|
||||
'message': ''
|
||||
})
|
||||
|
||||
provider_stats[provider]["passed"] += 1
|
||||
provider_tests[provider].append(
|
||||
{
|
||||
"name": test_name,
|
||||
"status": "PASSED",
|
||||
"time": test_time,
|
||||
"message": "",
|
||||
}
|
||||
)
|
||||
|
||||
passed = total_tests - total_failures - total_errors - total_skipped
|
||||
|
||||
|
||||
# Generate the markdown report
|
||||
with open(output_path, 'w') as f:
|
||||
with open(output_path, "w") as f:
|
||||
# Header
|
||||
f.write("# LLM Translation Test Results\n\n")
|
||||
|
||||
|
||||
# Metadata table
|
||||
f.write("## Test Run Information\n\n")
|
||||
f.write("| Field | Value |\n")
|
||||
f.write("|-------|-------|\n")
|
||||
f.write(f"| **Tag** | `{tag or 'N/A'}` |\n")
|
||||
f.write(f"| **Date** | {datetime.utcnow().strftime('%Y-%m-%d %H:%M:%S UTC')} |\n")
|
||||
f.write(
|
||||
f"| **Date** | {datetime.utcnow().strftime('%Y-%m-%d %H:%M:%S UTC')} |\n"
|
||||
)
|
||||
f.write(f"| **Commit** | `{commit or 'N/A'}` |\n")
|
||||
f.write(f"| **Duration** | {format_duration(total_time)} |\n")
|
||||
f.write("\n")
|
||||
|
||||
|
||||
# Overall statistics with visual elements
|
||||
f.write("## Overall Statistics\n\n")
|
||||
|
||||
|
||||
# Summary box
|
||||
f.write("```\n")
|
||||
f.write(f"Total Tests: {total_tests}\n")
|
||||
f.write(f"├── Passed: {passed:>4} ({(passed/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
|
||||
f.write(f"├── Failed: {total_failures:>4} ({(total_failures/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
|
||||
f.write(f"├── Errors: {total_errors:>4} ({(total_errors/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
|
||||
f.write(f"└── Skipped: {total_skipped:>4} ({(total_skipped/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
|
||||
f.write(
|
||||
f"├── Passed: {passed:>4} ({(passed/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
|
||||
)
|
||||
f.write(
|
||||
f"├── Failed: {total_failures:>4} ({(total_failures/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
|
||||
)
|
||||
f.write(
|
||||
f"├── Errors: {total_errors:>4} ({(total_errors/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
|
||||
)
|
||||
f.write(
|
||||
f"└── Skipped: {total_skipped:>4} ({(total_skipped/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
|
||||
)
|
||||
f.write("```\n\n")
|
||||
|
||||
|
||||
|
||||
# Provider summary table
|
||||
f.write("## Results by Provider\n\n")
|
||||
f.write("| Provider | Total | Pass | Fail | Error | Skip | Pass Rate | Duration |\n")
|
||||
f.write("|----------|-------|------|------|-------|------|-----------|----------|")
|
||||
|
||||
f.write(
|
||||
"| Provider | Total | Pass | Fail | Error | Skip | Pass Rate | Duration |\n"
|
||||
)
|
||||
f.write(
|
||||
"|----------|-------|------|------|-------|------|-----------|----------|"
|
||||
)
|
||||
|
||||
# Sort providers: specific providers first, then cross-provider tests
|
||||
sorted_providers = []
|
||||
cross_provider = []
|
||||
for p in sorted(provider_stats.keys()):
|
||||
if 'Cross-Provider' in p or p == 'Other Tests':
|
||||
if "Cross-Provider" in p or p == "Other Tests":
|
||||
cross_provider.append(p)
|
||||
else:
|
||||
sorted_providers.append(p)
|
||||
|
||||
|
||||
all_providers = sorted_providers + cross_provider
|
||||
|
||||
|
||||
for provider in all_providers:
|
||||
stats = provider_stats[provider]
|
||||
total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped']
|
||||
pass_rate = (stats['passed'] / total * 100) if total > 0 else 0
|
||||
|
||||
f.write(f"\n| {provider} | {total} | {stats['passed']} | {stats['failed']} | ")
|
||||
total = (
|
||||
stats["passed"]
|
||||
+ stats["failed"]
|
||||
+ stats["errors"]
|
||||
+ stats["skipped"]
|
||||
)
|
||||
pass_rate = (stats["passed"] / total * 100) if total > 0 else 0
|
||||
|
||||
f.write(
|
||||
f"\n| {provider} | {total} | {stats['passed']} | {stats['failed']} | "
|
||||
)
|
||||
f.write(f"{stats['errors']} | {stats['skipped']} | {pass_rate:.1f}% | ")
|
||||
f.write(f"{format_duration(stats['time'])} |")
|
||||
|
||||
|
||||
# Detailed test results by provider
|
||||
f.write("\n\n## Detailed Test Results\n\n")
|
||||
|
||||
|
||||
for provider in sorted_providers:
|
||||
if provider_tests[provider]:
|
||||
stats = provider_stats[provider]
|
||||
total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped']
|
||||
|
||||
total = (
|
||||
stats["passed"]
|
||||
+ stats["failed"]
|
||||
+ stats["errors"]
|
||||
+ stats["skipped"]
|
||||
)
|
||||
|
||||
f.write(f"### {provider}\n\n")
|
||||
f.write(f"**Summary:** {stats['passed']}/{total} passed ")
|
||||
f.write(f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%) ")
|
||||
f.write(
|
||||
f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%) "
|
||||
)
|
||||
f.write(f"in {format_duration(stats['time'])}\n\n")
|
||||
|
||||
|
||||
# Group tests by status
|
||||
tests_by_status = defaultdict(list)
|
||||
for test in provider_tests[provider]:
|
||||
tests_by_status[test['status']].append(test)
|
||||
|
||||
tests_by_status[test["status"]].append(test)
|
||||
|
||||
# Show failed tests first (if any)
|
||||
if tests_by_status['FAILED']:
|
||||
if tests_by_status["FAILED"]:
|
||||
f.write("<details>\n<summary>Failed Tests</summary>\n\n")
|
||||
for test in tests_by_status['FAILED']:
|
||||
for test in tests_by_status["FAILED"]:
|
||||
f.write(f"- `{test['name']}` ({test['time']:.2f}s)\n")
|
||||
if test['message']:
|
||||
if test["message"]:
|
||||
# Truncate long error messages
|
||||
msg = test['message'][:200] + '...' if len(test['message']) > 200 else test['message']
|
||||
msg = (
|
||||
test["message"][:200] + "..."
|
||||
if len(test["message"]) > 200
|
||||
else test["message"]
|
||||
)
|
||||
f.write(f" > {msg}\n")
|
||||
f.write("\n</details>\n\n")
|
||||
|
||||
|
||||
# Show errors (if any)
|
||||
if tests_by_status['ERROR']:
|
||||
if tests_by_status["ERROR"]:
|
||||
f.write("<details>\n<summary>Error Tests</summary>\n\n")
|
||||
for test in tests_by_status['ERROR']:
|
||||
for test in tests_by_status["ERROR"]:
|
||||
f.write(f"- `{test['name']}` ({test['time']:.2f}s)\n")
|
||||
f.write("\n</details>\n\n")
|
||||
|
||||
|
||||
# Show passed tests in collapsible section
|
||||
if tests_by_status['PASSED']:
|
||||
if tests_by_status["PASSED"]:
|
||||
f.write("<details>\n<summary>Passed Tests</summary>\n\n")
|
||||
for test in tests_by_status['PASSED']:
|
||||
for test in tests_by_status["PASSED"]:
|
||||
f.write(f"- `{test['name']}` ({test['time']:.2f}s)\n")
|
||||
f.write("\n</details>\n\n")
|
||||
|
||||
|
||||
# Show skipped tests (if any)
|
||||
if tests_by_status['SKIPPED']:
|
||||
if tests_by_status["SKIPPED"]:
|
||||
f.write("<details>\n<summary>Skipped Tests</summary>\n\n")
|
||||
for test in tests_by_status['SKIPPED']:
|
||||
for test in tests_by_status["SKIPPED"]:
|
||||
f.write(f"- `{test['name']}`\n")
|
||||
f.write("\n</details>\n\n")
|
||||
|
||||
|
||||
# Cross-provider tests in a separate section
|
||||
if cross_provider:
|
||||
f.write("### Cross-Provider Tests\n\n")
|
||||
for provider in cross_provider:
|
||||
if provider_tests[provider]:
|
||||
stats = provider_stats[provider]
|
||||
total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped']
|
||||
|
||||
total = (
|
||||
stats["passed"]
|
||||
+ stats["failed"]
|
||||
+ stats["errors"]
|
||||
+ stats["skipped"]
|
||||
)
|
||||
|
||||
f.write(f"#### {provider}\n\n")
|
||||
f.write(f"**Summary:** {stats['passed']}/{total} passed ")
|
||||
f.write(f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%)\n\n")
|
||||
|
||||
f.write(
|
||||
f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%)\n\n"
|
||||
)
|
||||
|
||||
# For cross-provider tests, just show counts
|
||||
f.write(f"- Passed: {stats['passed']}\n")
|
||||
if stats['failed'] > 0:
|
||||
if stats["failed"] > 0:
|
||||
f.write(f"- Failed: {stats['failed']}\n")
|
||||
if stats['errors'] > 0:
|
||||
if stats["errors"] > 0:
|
||||
f.write(f"- Errors: {stats['errors']}\n")
|
||||
if stats['skipped'] > 0:
|
||||
if stats["skipped"] > 0:
|
||||
f.write(f"- Skipped: {stats['skipped']}\n")
|
||||
f.write("\n")
|
||||
|
||||
|
||||
|
||||
print_colored(f"Report generated: {output_path}", Colors.GREEN)
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print_colored(f"Error generating report: {e}", Colors.RED)
|
||||
raise
|
||||
|
||||
def run_tests(test_path: str = "tests/llm_translation/",
|
||||
junit_xml: str = "test-results/junit.xml",
|
||||
report_path: str = "test-results/llm_translation_report.md",
|
||||
tag: str = None,
|
||||
commit: str = None) -> int:
|
||||
|
||||
def run_tests(
|
||||
test_path: str = "tests/llm_translation/",
|
||||
junit_xml: str = "test-results/junit.xml",
|
||||
report_path: str = "test-results/llm_translation_report.md",
|
||||
tag: str = None,
|
||||
commit: str = None,
|
||||
) -> int:
|
||||
"""Run the LLM translation tests and generate report"""
|
||||
|
||||
|
||||
# Create test results directory
|
||||
os.makedirs(os.path.dirname(junit_xml), exist_ok=True)
|
||||
|
||||
|
||||
print_colored("Starting LLM Translation Tests", Colors.BOLD + Colors.BLUE)
|
||||
print_colored(f"Test directory: {test_path}", Colors.CYAN)
|
||||
print_colored(f"Output: {junit_xml}", Colors.CYAN)
|
||||
print()
|
||||
|
||||
|
||||
# Run pytest
|
||||
cmd = [
|
||||
"uv", "run", "--no-sync", "pytest", test_path,
|
||||
"uv",
|
||||
"run",
|
||||
"--no-sync",
|
||||
"pytest",
|
||||
test_path,
|
||||
f"--junitxml={junit_xml}",
|
||||
"-v",
|
||||
"--tb=short",
|
||||
"--maxfail=500",
|
||||
"-n", "auto"
|
||||
"-n",
|
||||
"auto",
|
||||
]
|
||||
|
||||
|
||||
# Add timeout if pytest-timeout is installed
|
||||
try:
|
||||
subprocess.run(["uv", "run", "--no-sync", "python", "-c", "import pytest_timeout"],
|
||||
capture_output=True, check=True)
|
||||
subprocess.run(
|
||||
["uv", "run", "--no-sync", "python", "-c", "import pytest_timeout"],
|
||||
capture_output=True,
|
||||
check=True,
|
||||
)
|
||||
cmd.extend(["--timeout=300"])
|
||||
except:
|
||||
print_colored("Warning: pytest-timeout not installed, skipping timeout option", Colors.YELLOW)
|
||||
|
||||
print_colored(
|
||||
"Warning: pytest-timeout not installed, skipping timeout option",
|
||||
Colors.YELLOW,
|
||||
)
|
||||
|
||||
print_colored("Running pytest with command:", Colors.YELLOW)
|
||||
print(f" {' '.join(cmd)}")
|
||||
print()
|
||||
|
||||
|
||||
# Run the tests
|
||||
result = subprocess.run(cmd, capture_output=False)
|
||||
|
||||
|
||||
# Generate the report regardless of test outcome
|
||||
if os.path.exists(junit_xml):
|
||||
print()
|
||||
print_colored("Generating test report...", Colors.BLUE)
|
||||
generate_markdown_report(junit_xml, report_path, tag, commit)
|
||||
|
||||
|
||||
# Print summary to console
|
||||
print()
|
||||
print_colored("Test Summary:", Colors.BOLD + Colors.PURPLE)
|
||||
|
||||
|
||||
# Parse XML for quick summary
|
||||
tree = ET.parse(junit_xml)
|
||||
root = tree.getroot()
|
||||
|
||||
if root.tag == 'testsuites':
|
||||
suites = root.findall('testsuite')
|
||||
|
||||
if root.tag == "testsuites":
|
||||
suites = root.findall("testsuite")
|
||||
else:
|
||||
suites = [root]
|
||||
|
||||
total = sum(int(s.get('tests', 0)) for s in suites)
|
||||
failures = sum(int(s.get('failures', 0)) for s in suites)
|
||||
errors = sum(int(s.get('errors', 0)) for s in suites)
|
||||
skipped = sum(int(s.get('skipped', 0)) for s in suites)
|
||||
|
||||
total = sum(int(s.get("tests", 0)) for s in suites)
|
||||
failures = sum(int(s.get("failures", 0)) for s in suites)
|
||||
errors = sum(int(s.get("errors", 0)) for s in suites)
|
||||
skipped = sum(int(s.get("skipped", 0)) for s in suites)
|
||||
passed = total - failures - errors - skipped
|
||||
|
||||
|
||||
print(f" Total: {total}")
|
||||
print_colored(f" Passed: {passed}", Colors.GREEN)
|
||||
if failures > 0:
|
||||
|
|
@ -381,59 +457,75 @@ def run_tests(test_path: str = "tests/llm_translation/",
|
|||
print_colored(f" Errors: {errors}", Colors.RED)
|
||||
if skipped > 0:
|
||||
print_colored(f" Skipped: {skipped}", Colors.YELLOW)
|
||||
|
||||
|
||||
if total > 0:
|
||||
pass_rate = (passed / total) * 100
|
||||
color = Colors.GREEN if pass_rate >= 80 else Colors.YELLOW if pass_rate >= 60 else Colors.RED
|
||||
color = (
|
||||
Colors.GREEN
|
||||
if pass_rate >= 80
|
||||
else Colors.YELLOW if pass_rate >= 60 else Colors.RED
|
||||
)
|
||||
print_colored(f" Pass Rate: {pass_rate:.1f}%", color)
|
||||
else:
|
||||
print_colored("No test results found!", Colors.RED)
|
||||
|
||||
|
||||
print()
|
||||
print_colored("Test run complete!", Colors.BOLD + Colors.GREEN)
|
||||
|
||||
|
||||
return result.returncode
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser(description="Run LLM Translation Tests")
|
||||
parser.add_argument("--test-path", default="tests/llm_translation/",
|
||||
help="Path to test directory")
|
||||
parser.add_argument("--junit-xml", default="test-results/junit.xml",
|
||||
help="Path for JUnit XML output")
|
||||
parser.add_argument("--report", default="test-results/llm_translation_report.md",
|
||||
help="Path for markdown report")
|
||||
parser.add_argument(
|
||||
"--test-path", default="tests/llm_translation/", help="Path to test directory"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--junit-xml",
|
||||
default="test-results/junit.xml",
|
||||
help="Path for JUnit XML output",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--report",
|
||||
default="test-results/llm_translation_report.md",
|
||||
help="Path for markdown report",
|
||||
)
|
||||
parser.add_argument("--tag", help="Git tag or version")
|
||||
parser.add_argument("--commit", help="Git commit SHA")
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
# Get git info if not provided
|
||||
if not args.commit:
|
||||
try:
|
||||
result = subprocess.run(["git", "rev-parse", "HEAD"],
|
||||
capture_output=True, text=True)
|
||||
result = subprocess.run(
|
||||
["git", "rev-parse", "HEAD"], capture_output=True, text=True
|
||||
)
|
||||
if result.returncode == 0:
|
||||
args.commit = result.stdout.strip()
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
if not args.tag:
|
||||
try:
|
||||
result = subprocess.run(["git", "describe", "--tags", "--abbrev=0"],
|
||||
capture_output=True, text=True)
|
||||
result = subprocess.run(
|
||||
["git", "describe", "--tags", "--abbrev=0"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
if result.returncode == 0:
|
||||
args.tag = result.stdout.strip()
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
exit_code = run_tests(
|
||||
test_path=args.test_path,
|
||||
junit_xml=args.junit_xml,
|
||||
report_path=args.report,
|
||||
tag=args.tag,
|
||||
commit=args.commit
|
||||
commit=args.commit,
|
||||
)
|
||||
|
||||
|
||||
sys.exit(exit_code)
|
||||
|
|
|
|||
12
cache_demo_config.yaml
Normal file
12
cache_demo_config.yaml
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
model_list:
|
||||
- model_name: bedrock-claude-haiku
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0
|
||||
aws_region_name: us-east-1
|
||||
|
||||
litellm_settings:
|
||||
success_callback: []
|
||||
failure_callback: []
|
||||
|
||||
general_settings:
|
||||
store_model_in_db: false
|
||||
75
cache_demo_request.py
Normal file
75
cache_demo_request.py
Normal file
|
|
@ -0,0 +1,75 @@
|
|||
"""
|
||||
Demo script: back-to-back Bedrock streaming requests with prompt caching.
|
||||
Request 1: populates the cache (cache_creation_input_tokens)
|
||||
Request 2: reads from cache (cache_read_input_tokens)
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
|
||||
import httpx
|
||||
|
||||
PROXY_URL = "http://localhost:4001"
|
||||
API_KEY = "sk-1234"
|
||||
|
||||
# ~5000 token system prompt (above claude-haiku-4-5's 2048-token min for caching on Bedrock)
|
||||
LARGE_SYSTEM_PROMPT = (
|
||||
"AWS Bedrock provides managed ML infrastructure for enterprise workloads. "
|
||||
"Anthropic Claude models support prompt caching for cost optimization. "
|
||||
) * 200
|
||||
|
||||
|
||||
def make_streaming_request(req_num: int, label: str) -> None:
|
||||
print(f"\n{'='*60}")
|
||||
print(f"Request {req_num}: {label}")
|
||||
print(f"{'='*60}")
|
||||
|
||||
payload = {
|
||||
"model": "bedrock-claude-haiku",
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": LARGE_SYSTEM_PROMPT,
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": f"Say only: 'Request {req_num} done'"},
|
||||
],
|
||||
"stream": True,
|
||||
"max_tokens": 20,
|
||||
}
|
||||
|
||||
full_response = ""
|
||||
with httpx.Client(timeout=60) as client:
|
||||
with client.stream(
|
||||
"POST",
|
||||
f"{PROXY_URL}/v1/chat/completions",
|
||||
json=payload,
|
||||
headers={"Authorization": f"Bearer {API_KEY}"},
|
||||
) as r:
|
||||
r.raise_for_status()
|
||||
for line in r.iter_lines():
|
||||
if line.startswith("data: ") and line != "data: [DONE]":
|
||||
chunk = json.loads(line[6:])
|
||||
delta = chunk.get("choices", [{}])[0].get("delta", {})
|
||||
if content := delta.get("content"):
|
||||
full_response += content
|
||||
|
||||
print(f"Response: {full_response!r}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("Sending request 1 (cache write)...")
|
||||
make_streaming_request(1, "cache WRITE (populates cache)")
|
||||
|
||||
print("\nWaiting 2s between requests...")
|
||||
time.sleep(2)
|
||||
|
||||
print("Sending request 2 (cache read)...")
|
||||
make_streaming_request(2, "cache READ (hits cache)")
|
||||
|
||||
print("\n\nDone. Check SpendLogs in the DB or the UI at http://localhost:4001")
|
||||
|
|
@ -14,53 +14,74 @@ from litellm.types.utils import StandardCallbackDynamicParams
|
|||
class EnterpriseCallbackControls:
|
||||
@staticmethod
|
||||
def is_callback_disabled_dynamically(
|
||||
callback: litellm.CALLBACK_TYPES,
|
||||
litellm_params: dict,
|
||||
standard_callback_dynamic_params: StandardCallbackDynamicParams
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a callback is disabled via the x-litellm-disable-callbacks header or via `litellm_disabled_callbacks` in standard_callback_dynamic_params.
|
||||
|
||||
Args:
|
||||
callback: The callback to check (can be string, CustomLogger instance, or callable)
|
||||
litellm_params: Parameters containing proxy server request info
|
||||
|
||||
Returns:
|
||||
bool: True if the callback should be disabled, False otherwise
|
||||
"""
|
||||
from litellm.litellm_core_utils.custom_logger_registry import (
|
||||
CustomLoggerRegistry,
|
||||
)
|
||||
callback: litellm.CALLBACK_TYPES,
|
||||
litellm_params: dict,
|
||||
standard_callback_dynamic_params: StandardCallbackDynamicParams,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a callback is disabled via the x-litellm-disable-callbacks header or via `litellm_disabled_callbacks` in standard_callback_dynamic_params.
|
||||
|
||||
Args:
|
||||
callback: The callback to check (can be string, CustomLogger instance, or callable)
|
||||
litellm_params: Parameters containing proxy server request info
|
||||
|
||||
Returns:
|
||||
bool: True if the callback should be disabled, False otherwise
|
||||
"""
|
||||
from litellm.litellm_core_utils.custom_logger_registry import (
|
||||
CustomLoggerRegistry,
|
||||
)
|
||||
|
||||
try:
|
||||
disabled_callbacks = EnterpriseCallbackControls.get_disabled_callbacks(
|
||||
litellm_params, standard_callback_dynamic_params
|
||||
)
|
||||
verbose_logger.debug(
|
||||
f"Dynamically disabled callbacks from {X_LITELLM_DISABLE_CALLBACKS}: {disabled_callbacks}"
|
||||
)
|
||||
verbose_logger.debug(
|
||||
f"Checking if {callback} is disabled via headers. Disable callbacks from headers: {disabled_callbacks}"
|
||||
)
|
||||
if disabled_callbacks is not None:
|
||||
#########################################################
|
||||
# premium user check
|
||||
#########################################################
|
||||
if (
|
||||
not EnterpriseCallbackControls._should_allow_dynamic_callback_disabling()
|
||||
):
|
||||
return False
|
||||
#########################################################
|
||||
if isinstance(callback, str):
|
||||
if callback.lower() in disabled_callbacks:
|
||||
verbose_logger.debug(
|
||||
f"Not logging to {callback} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}"
|
||||
)
|
||||
return True
|
||||
elif isinstance(callback, CustomLogger):
|
||||
# get the string name of the callback
|
||||
callback_str = (
|
||||
CustomLoggerRegistry.get_callback_str_from_class_type(
|
||||
callback.__class__
|
||||
)
|
||||
)
|
||||
if (
|
||||
callback_str is not None
|
||||
and callback_str.lower() in disabled_callbacks
|
||||
):
|
||||
verbose_logger.debug(
|
||||
f"Not logging to {callback_str} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}"
|
||||
)
|
||||
return True
|
||||
return False
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Error checking disabled callbacks header: {str(e)}")
|
||||
return False
|
||||
|
||||
try:
|
||||
disabled_callbacks = EnterpriseCallbackControls.get_disabled_callbacks(litellm_params, standard_callback_dynamic_params)
|
||||
verbose_logger.debug(f"Dynamically disabled callbacks from {X_LITELLM_DISABLE_CALLBACKS}: {disabled_callbacks}")
|
||||
verbose_logger.debug(f"Checking if {callback} is disabled via headers. Disable callbacks from headers: {disabled_callbacks}")
|
||||
if disabled_callbacks is not None:
|
||||
#########################################################
|
||||
# premium user check
|
||||
#########################################################
|
||||
if not EnterpriseCallbackControls._should_allow_dynamic_callback_disabling():
|
||||
return False
|
||||
#########################################################
|
||||
if isinstance(callback, str):
|
||||
if callback.lower() in disabled_callbacks:
|
||||
verbose_logger.debug(f"Not logging to {callback} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}")
|
||||
return True
|
||||
elif isinstance(callback, CustomLogger):
|
||||
# get the string name of the callback
|
||||
callback_str = CustomLoggerRegistry.get_callback_str_from_class_type(callback.__class__)
|
||||
if callback_str is not None and callback_str.lower() in disabled_callbacks:
|
||||
verbose_logger.debug(f"Not logging to {callback_str} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}")
|
||||
return True
|
||||
return False
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
f"Error checking disabled callbacks header: {str(e)}"
|
||||
)
|
||||
return False
|
||||
@staticmethod
|
||||
def get_disabled_callbacks(litellm_params: dict, standard_callback_dynamic_params: StandardCallbackDynamicParams) -> Optional[List[str]]:
|
||||
def get_disabled_callbacks(
|
||||
litellm_params: dict,
|
||||
standard_callback_dynamic_params: StandardCallbackDynamicParams,
|
||||
) -> Optional[List[str]]:
|
||||
"""
|
||||
Get the disabled callbacks from the standard callback dynamic params.
|
||||
"""
|
||||
|
|
@ -71,18 +92,24 @@ class EnterpriseCallbackControls:
|
|||
request_headers = get_proxy_server_request_headers(litellm_params)
|
||||
disabled_callbacks = request_headers.get(X_LITELLM_DISABLE_CALLBACKS, None)
|
||||
if disabled_callbacks is not None:
|
||||
disabled_callbacks = set([cb.strip().lower() for cb in disabled_callbacks.split(",")])
|
||||
disabled_callbacks = set(
|
||||
[cb.strip().lower() for cb in disabled_callbacks.split(",")]
|
||||
)
|
||||
return list(disabled_callbacks)
|
||||
|
||||
|
||||
#########################################################
|
||||
# check if disabled via request body
|
||||
#########################################################
|
||||
if standard_callback_dynamic_params.get("litellm_disabled_callbacks", None) is not None:
|
||||
return standard_callback_dynamic_params.get("litellm_disabled_callbacks", None)
|
||||
|
||||
if (
|
||||
standard_callback_dynamic_params.get("litellm_disabled_callbacks", None)
|
||||
is not None
|
||||
):
|
||||
return standard_callback_dynamic_params.get(
|
||||
"litellm_disabled_callbacks", None
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _should_allow_dynamic_callback_disabling():
|
||||
import litellm
|
||||
|
|
@ -90,10 +117,14 @@ class EnterpriseCallbackControls:
|
|||
|
||||
# Check if admin has disabled this feature
|
||||
if litellm.allow_dynamic_callback_disabling is not True:
|
||||
verbose_logger.debug("Dynamic callback disabling is disabled by admin via litellm.allow_dynamic_callback_disabling")
|
||||
verbose_logger.debug(
|
||||
"Dynamic callback disabling is disabled by admin via litellm.allow_dynamic_callback_disabling"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
if premium_user:
|
||||
return True
|
||||
verbose_logger.warning(f"Disabling callbacks using request headers is an enterprise feature. {CommonProxyErrors.not_premium_user.value}")
|
||||
return False
|
||||
verbose_logger.warning(
|
||||
f"Disabling callbacks using request headers is an enterprise feature. {CommonProxyErrors.not_premium_user.value}"
|
||||
)
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -79,4 +79,4 @@ class SendGridEmailLogger(BaseEmailLogger):
|
|||
verbose_logger.debug(
|
||||
f"SendGrid response status={response.status_code}, body={response.text}"
|
||||
)
|
||||
return
|
||||
return
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
"""
|
||||
This is the litellm SMTP email integration
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import List
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
"""
|
||||
Enterprise specific logging utils
|
||||
"""
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingMetadata
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -7,4 +7,4 @@ including custom SSO handlers and advanced authentication features.
|
|||
|
||||
from .custom_sso_handler import EnterpriseCustomSSOHandler
|
||||
|
||||
__all__ = ["EnterpriseCustomSSOHandler"]
|
||||
__all__ = ["EnterpriseCustomSSOHandler"]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -41,7 +41,7 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
):
|
||||
"""
|
||||
Managed vector stores with target_model_names support.
|
||||
|
||||
|
||||
This class provides functionality to:
|
||||
- Create vector stores across multiple models
|
||||
- Retrieve vector stores by unified ID
|
||||
|
|
@ -77,14 +77,14 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> str:
|
||||
"""
|
||||
Generate the format string for the unified vector store ID.
|
||||
|
||||
|
||||
Format:
|
||||
litellm_proxy:vector_store;unified_id,<uuid>;target_model_names,<models>;resource_id,<vs_id>;model_id,<model_id>
|
||||
"""
|
||||
# VectorStoreCreateResponse is a TypedDict, so resource_object is a dictionary
|
||||
# Extract provider resource ID from the response
|
||||
provider_resource_id = resource_object.get("id", "")
|
||||
|
||||
|
||||
# Model ID is stored in hidden params if the response object supports it
|
||||
# For TypedDict responses, we need to check if _hidden_params was added
|
||||
hidden_params: Dict[str, Any] = {}
|
||||
|
|
@ -109,20 +109,18 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> VectorStoreCreateResponse:
|
||||
"""
|
||||
Create a vector store for a specific model.
|
||||
|
||||
|
||||
Args:
|
||||
llm_router: LiteLLM router instance
|
||||
model: Model name to create vector store for
|
||||
request_data: Request data for vector store creation
|
||||
litellm_parent_otel_span: OpenTelemetry span for tracing
|
||||
|
||||
|
||||
Returns:
|
||||
VectorStoreCreateResponse from the provider
|
||||
"""
|
||||
# Use the router to create the vector store
|
||||
response = await llm_router.avector_store_create(
|
||||
model=model, **request_data
|
||||
)
|
||||
response = await llm_router.avector_store_create(model=model, **request_data)
|
||||
return response
|
||||
|
||||
# ============================================================================
|
||||
|
|
@ -139,14 +137,14 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> VectorStoreCreateResponse:
|
||||
"""
|
||||
Create a vector store across multiple models.
|
||||
|
||||
|
||||
Args:
|
||||
create_request: Vector store creation request parameters
|
||||
llm_router: LiteLLM router instance
|
||||
target_model_names_list: List of target model names
|
||||
litellm_parent_otel_span: OpenTelemetry span for tracing
|
||||
user_api_key_dict: User API key authentication details
|
||||
|
||||
|
||||
Returns:
|
||||
VectorStoreCreateResponse with unified ID
|
||||
"""
|
||||
|
|
@ -196,7 +194,7 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
# VectorStoreCreateResponse is a TypedDict, so we need to create a new dict with the unified ID
|
||||
response = responses[0].copy()
|
||||
response["id"] = unified_id
|
||||
|
||||
|
||||
verbose_logger.info(
|
||||
f"Successfully created managed vector store with unified ID: {unified_id}"
|
||||
)
|
||||
|
|
@ -212,13 +210,13 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> Dict[str, Any]:
|
||||
"""
|
||||
List vector stores created by a user.
|
||||
|
||||
|
||||
Args:
|
||||
user_api_key_dict: User API key authentication details
|
||||
limit: Maximum number of vector stores to return
|
||||
after: Cursor for pagination
|
||||
order: Sort order ('asc' or 'desc')
|
||||
|
||||
|
||||
Returns:
|
||||
Dictionary with list of vector stores and pagination info
|
||||
"""
|
||||
|
|
@ -238,23 +236,23 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> bool:
|
||||
"""
|
||||
Check if user has access to a vector store.
|
||||
|
||||
|
||||
Args:
|
||||
vector_store_id: The unified vector store ID
|
||||
user_api_key_dict: User API key authentication details
|
||||
|
||||
|
||||
Returns:
|
||||
True if user has access, False otherwise
|
||||
"""
|
||||
is_unified_id = is_base64_encoded_unified_id(vector_store_id)
|
||||
|
||||
|
||||
if is_unified_id:
|
||||
# Check access for managed vector store
|
||||
return await self.can_user_access_unified_resource_id(
|
||||
vector_store_id,
|
||||
user_api_key_dict,
|
||||
)
|
||||
|
||||
|
||||
# Not a managed vector store, allow access
|
||||
return True
|
||||
|
||||
|
|
@ -263,24 +261,22 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> bool:
|
||||
"""
|
||||
Check if user has access to a managed vector store in request data.
|
||||
|
||||
|
||||
Args:
|
||||
data: Request data containing vector_store_id
|
||||
user_api_key_dict: User API key authentication details
|
||||
|
||||
|
||||
Returns:
|
||||
True if this is a managed vector store and user has access
|
||||
|
||||
|
||||
Raises:
|
||||
HTTPException: If user doesn't have access
|
||||
"""
|
||||
vector_store_id = cast(Optional[str], data.get("vector_store_id"))
|
||||
is_unified_id = (
|
||||
is_base64_encoded_unified_id(vector_store_id)
|
||||
if vector_store_id
|
||||
else False
|
||||
is_base64_encoded_unified_id(vector_store_id) if vector_store_id else False
|
||||
)
|
||||
|
||||
|
||||
if is_unified_id and vector_store_id:
|
||||
if await self.can_user_access_unified_resource_id(
|
||||
vector_store_id, user_api_key_dict
|
||||
|
|
@ -291,7 +287,7 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
status_code=403,
|
||||
detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}",
|
||||
)
|
||||
|
||||
|
||||
return False
|
||||
|
||||
# ============================================================================
|
||||
|
|
@ -307,18 +303,18 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> Union[Exception, str, Dict, None]:
|
||||
"""
|
||||
Pre-call hook to handle vector store operations.
|
||||
|
||||
|
||||
This hook intercepts vector store requests and:
|
||||
- Validates access for managed vector stores
|
||||
- Transforms unified IDs to provider-specific IDs
|
||||
- Adds model routing information
|
||||
|
||||
|
||||
Args:
|
||||
user_api_key_dict: User API key authentication details
|
||||
cache: Cache instance
|
||||
data: Request data
|
||||
call_type: Type of call being made
|
||||
|
||||
|
||||
Returns:
|
||||
Modified request data or None
|
||||
"""
|
||||
|
|
@ -330,40 +326,40 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
# Handle vector store search operations
|
||||
if call_type == "avector_store_search":
|
||||
vector_store_id = data.get("vector_store_id")
|
||||
|
||||
|
||||
if vector_store_id:
|
||||
# Check if it's a managed vector store ID
|
||||
decoded_id = is_base64_encoded_unified_id(vector_store_id)
|
||||
|
||||
|
||||
if decoded_id:
|
||||
verbose_logger.debug(
|
||||
f"Processing managed vector store search: {vector_store_id}"
|
||||
)
|
||||
|
||||
|
||||
# Check access
|
||||
has_access = await self.can_user_access_unified_resource_id(
|
||||
vector_store_id, user_api_key_dict
|
||||
)
|
||||
|
||||
|
||||
if not has_access:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}",
|
||||
)
|
||||
|
||||
|
||||
# Parse the unified ID to extract components
|
||||
parsed_id = parse_unified_id(vector_store_id)
|
||||
|
||||
|
||||
if parsed_id:
|
||||
# Extract the model ID and provider resource ID
|
||||
model_id = parsed_id.get("model_id")
|
||||
provider_resource_id = parsed_id.get("provider_resource_id")
|
||||
target_model_names = parsed_id.get("target_model_names", [])
|
||||
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Decoded vector store - model_id: {model_id}, provider_resource_id: {provider_resource_id}, target_model_names: {target_model_names}"
|
||||
)
|
||||
|
||||
|
||||
# Determine which model to use for routing
|
||||
# Priority: model_id (deployment ID) > first target_model_name
|
||||
routing_model = None
|
||||
|
|
@ -371,28 +367,28 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
routing_model = model_id
|
||||
elif target_model_names and len(target_model_names) > 0:
|
||||
routing_model = target_model_names[0]
|
||||
|
||||
|
||||
# Set the model for routing
|
||||
if routing_model:
|
||||
data["model"] = routing_model
|
||||
verbose_logger.info(
|
||||
f"Routing vector store search to model: {routing_model}"
|
||||
)
|
||||
|
||||
|
||||
# Replace the unified ID with the provider-specific ID
|
||||
if provider_resource_id:
|
||||
data["vector_store_id"] = provider_resource_id
|
||||
verbose_logger.debug(
|
||||
f"Replaced unified ID with provider resource ID: {provider_resource_id}"
|
||||
)
|
||||
|
||||
|
||||
# Handle vector store retrieve/delete operations
|
||||
elif call_type in ("avector_store_retrieve", "avector_store_delete"):
|
||||
await self.check_managed_vector_store_access(data, user_api_key_dict)
|
||||
|
||||
|
||||
# If it's a managed vector store, we'll handle it in the endpoint
|
||||
# No need to transform here as the endpoint will route to the hook
|
||||
|
||||
|
||||
return data
|
||||
|
||||
# ============================================================================
|
||||
|
|
@ -407,15 +403,15 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> Any:
|
||||
"""
|
||||
Post-call hook to transform responses.
|
||||
|
||||
|
||||
This hook can be used to transform responses if needed.
|
||||
For now, it just passes through the response.
|
||||
|
||||
|
||||
Args:
|
||||
data: Request data
|
||||
user_api_key_dict: User API key authentication details
|
||||
response: Response from the provider
|
||||
|
||||
|
||||
Returns:
|
||||
Potentially modified response
|
||||
"""
|
||||
|
|
@ -436,21 +432,21 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> List[Dict]:
|
||||
"""
|
||||
Filter deployments based on vector store availability.
|
||||
|
||||
|
||||
This is used by the router to select only deployments that have
|
||||
the vector store available.
|
||||
|
||||
|
||||
Note: This method signature is a compromise between CustomLogger and BaseManagedResource
|
||||
parent classes which have incompatible signatures. The type: ignore[override] is necessary
|
||||
due to this multiple inheritance conflict.
|
||||
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
healthy_deployments: List of healthy deployments
|
||||
messages: Messages (unused for vector stores, required by CustomLogger interface)
|
||||
request_kwargs: Request kwargs containing vector_store_id and mappings
|
||||
parent_otel_span: OpenTelemetry span for tracing
|
||||
|
||||
|
||||
Returns:
|
||||
Filtered list of deployments
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
Enterprise internal user management endpoints
|
||||
"""
|
||||
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -39,15 +39,23 @@ class EmailEvent(str, enum.Enum):
|
|||
soft_budget_crossed = "Soft Budget Crossed"
|
||||
max_budget_alert = "Max Budget Alert"
|
||||
|
||||
|
||||
class EmailEventSettings(BaseModel):
|
||||
event: EmailEvent
|
||||
enabled: bool
|
||||
|
||||
|
||||
class EmailEventSettingsUpdateRequest(BaseModel):
|
||||
settings: List[EmailEventSettings]
|
||||
|
||||
|
||||
class EmailEventSettingsResponse(BaseModel):
|
||||
settings: List[EmailEventSettings]
|
||||
|
||||
|
||||
class DefaultEmailSettings(BaseModel):
|
||||
"""Default settings for email events"""
|
||||
|
||||
settings: Dict[EmailEvent, bool] = Field(
|
||||
default_factory=lambda: {
|
||||
EmailEvent.virtual_key_created: True, # On by default
|
||||
|
|
@ -57,10 +65,12 @@ class DefaultEmailSettings(BaseModel):
|
|||
EmailEvent.max_budget_alert: True, # On by default
|
||||
}
|
||||
)
|
||||
|
||||
def to_dict(self) -> Dict[str, bool]:
|
||||
"""Convert to dictionary with string keys for storage"""
|
||||
return {event.value: enabled for event, enabled in self.settings.items()}
|
||||
|
||||
@classmethod
|
||||
def get_defaults(cls) -> Dict[str, bool]:
|
||||
"""Get the default settings as a dictionary with string keys"""
|
||||
return cls().to_dict()
|
||||
return cls().to_dict()
|
||||
|
|
|
|||
|
|
@ -51,8 +51,7 @@ class PrometheusLabelFactoryContext:
|
|||
self.enum_values = enum_values
|
||||
enum_dict = enum_values.model_dump()
|
||||
self._sanitized_enum: Dict[str, Optional[str]] = {
|
||||
k: _sanitize_prometheus_label_value(v)
|
||||
for k, v in enum_dict.items()
|
||||
k: _sanitize_prometheus_label_value(v) for k, v in enum_dict.items()
|
||||
}
|
||||
self._custom_by_sanitized_key: Dict[str, Optional[str]] = {}
|
||||
if enum_values.custom_metadata_labels is not None:
|
||||
|
|
|
|||
|
|
@ -294,9 +294,7 @@ class Authenticator:
|
|||
access_token_url = os.getenv(
|
||||
"GITHUB_COPILOT_ACCESS_TOKEN_URL", DEFAULT_GITHUB_ACCESS_TOKEN_URL
|
||||
)
|
||||
client_id = os.getenv(
|
||||
"GITHUB_COPILOT_CLIENT_ID", DEFAULT_GITHUB_CLIENT_ID
|
||||
)
|
||||
client_id = os.getenv("GITHUB_COPILOT_CLIENT_ID", DEFAULT_GITHUB_CLIENT_ID)
|
||||
|
||||
for attempt in range(max_attempts):
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -79,7 +79,9 @@ class BasePassthroughUtils:
|
|||
for header_name, header_value in request_headers.items():
|
||||
if header_name.lower().startswith(PASS_THROUGH_HEADER_PREFIX):
|
||||
# Strip the 'x-pass-' prefix and normalize to lowercase
|
||||
actual_header_name = header_name[len(PASS_THROUGH_HEADER_PREFIX) :].lower()
|
||||
actual_header_name = header_name[
|
||||
len(PASS_THROUGH_HEADER_PREFIX) :
|
||||
].lower()
|
||||
if actual_header_name in _PASS_THROUGH_PROTECTED_HEADERS or any(
|
||||
actual_header_name.startswith(p)
|
||||
for p in _PASS_THROUGH_PROTECTED_HEADER_PREFIXES
|
||||
|
|
|
|||
|
|
@ -784,7 +784,7 @@ class UserAPIKeyLabelValues:
|
|||
org_id: Optional[str] = None
|
||||
org_alias: Optional[str] = None
|
||||
|
||||
#Added for test compatibility.
|
||||
# Added for test compatibility.
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
"""
|
||||
Match former Pydantic behavior: unknown keys are ignored; ``api_key_hash`` maps to
|
||||
|
|
|
|||
|
|
@ -23,8 +23,6 @@ def event_loop():
|
|||
loop.close()
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.fixture(scope="function", autouse=True)
|
||||
def setup_and_teardown():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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: {
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ def no_invitation_wait(monkeypatch):
|
|||
|
||||
monkeypatch.setattr(BaseEmailLogger, "_wait_for_invitation_creation", _noop)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def base_email_logger():
|
||||
return BaseEmailLogger()
|
||||
|
|
@ -283,7 +284,10 @@ async def test_send_key_created_email_without_key(
|
|||
mock_send_email.assert_called_once()
|
||||
call_args = mock_send_email.call_args[1]
|
||||
assert "sk-secret-key-456" not in call_args["html_body"]
|
||||
assert "[Key hidden for security - retrieve from dashboard]" in call_args["html_body"]
|
||||
assert (
|
||||
"[Key hidden for security - retrieve from dashboard]"
|
||||
in call_args["html_body"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -317,7 +321,10 @@ async def test_send_key_rotated_email_without_key(
|
|||
mock_send_email.assert_called_once()
|
||||
call_args = mock_send_email.call_args[1]
|
||||
assert "sk-secret-rotated-789" not in call_args["html_body"]
|
||||
assert "[Key hidden for security - retrieve from dashboard]" in call_args["html_body"]
|
||||
assert (
|
||||
"[Key hidden for security - retrieve from dashboard]"
|
||||
in call_args["html_body"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -371,52 +378,52 @@ async def test_get_invitation_link_creates_new_when_none_exist(base_email_logger
|
|||
"""Test that _get_invitation_link creates a new invitation when none exist"""
|
||||
# Mock prisma client with no existing invitation rows
|
||||
mock_prisma = mock.MagicMock()
|
||||
|
||||
|
||||
# Mock find_many to return empty list (no existing invitations)
|
||||
async def mock_find_many_empty(*args, **kwargs):
|
||||
return []
|
||||
|
||||
|
||||
mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_empty
|
||||
|
||||
|
||||
# Mock the create_invitation_for_user function
|
||||
mock_created_invitation = mock.MagicMock()
|
||||
mock_created_invitation.id = "new-invitation-id"
|
||||
|
||||
|
||||
with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
with mock.patch(
|
||||
"litellm.proxy.management_helpers.user_invitation.create_invitation_for_user",
|
||||
return_value=mock_created_invitation
|
||||
return_value=mock_created_invitation,
|
||||
) as mock_create_invitation:
|
||||
# Execute
|
||||
result = await base_email_logger._get_invitation_link(
|
||||
user_id="test-user", base_url="http://test.com"
|
||||
)
|
||||
|
||||
|
||||
# Verify that create_invitation_for_user was called
|
||||
mock_create_invitation.assert_called_once()
|
||||
call_args = mock_create_invitation.call_args[1]
|
||||
assert call_args["data"].user_id == "test-user"
|
||||
assert call_args["user_api_key_dict"].user_id == "test-user"
|
||||
|
||||
|
||||
# Verify the returned link uses the new invitation ID
|
||||
assert result == "http://test.com/ui?invitation_id=new-invitation-id"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_invitation_link_uses_existing_when_available(base_email_logger):
|
||||
"""Test that _get_invitation_link uses existing invitation when available"""
|
||||
# Mock prisma client with existing invitation row
|
||||
mock_invitation_row = mock.MagicMock()
|
||||
mock_invitation_row.id = "existing-invitation-id"
|
||||
|
||||
|
||||
mock_prisma = mock.MagicMock()
|
||||
|
||||
|
||||
# Mock find_many to return existing invitation
|
||||
async def mock_find_many_existing(*args, **kwargs):
|
||||
return [mock_invitation_row]
|
||||
|
||||
|
||||
mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_existing
|
||||
|
||||
|
||||
with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
with mock.patch(
|
||||
"litellm.proxy.management_helpers.user_invitation.create_invitation_for_user"
|
||||
|
|
@ -425,10 +432,10 @@ async def test_get_invitation_link_uses_existing_when_available(base_email_logge
|
|||
result = await base_email_logger._get_invitation_link(
|
||||
user_id="test-user", base_url="http://test.com"
|
||||
)
|
||||
|
||||
|
||||
# Verify that create_invitation_for_user was NOT called
|
||||
mock_create_invitation.assert_not_called()
|
||||
|
||||
|
||||
# Verify the returned link uses the existing invitation ID
|
||||
assert result == "http://test.com/ui?invitation_id=existing-invitation-id"
|
||||
|
||||
|
|
@ -438,33 +445,33 @@ async def test_get_invitation_link_creates_new_when_list_is_none(base_email_logg
|
|||
"""Test that _get_invitation_link creates a new invitation when invitation_rows is None"""
|
||||
# Mock prisma client to return None
|
||||
mock_prisma = mock.MagicMock()
|
||||
|
||||
|
||||
# Mock find_many to return None
|
||||
async def mock_find_many_none(*args, **kwargs):
|
||||
return None
|
||||
|
||||
|
||||
mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_none
|
||||
|
||||
|
||||
# Mock the create_invitation_for_user function
|
||||
mock_created_invitation = mock.MagicMock()
|
||||
mock_created_invitation.id = "new-invitation-from-none"
|
||||
|
||||
|
||||
with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
with mock.patch(
|
||||
"litellm.proxy.management_helpers.user_invitation.create_invitation_for_user",
|
||||
return_value=mock_created_invitation
|
||||
return_value=mock_created_invitation,
|
||||
) as mock_create_invitation:
|
||||
# Execute
|
||||
result = await base_email_logger._get_invitation_link(
|
||||
user_id="test-user", base_url="http://test.com"
|
||||
)
|
||||
|
||||
|
||||
# Verify that create_invitation_for_user was called
|
||||
mock_create_invitation.assert_called_once()
|
||||
call_args = mock_create_invitation.call_args[1]
|
||||
assert call_args["data"].user_id == "test-user"
|
||||
assert call_args["user_api_key_dict"].user_id == "test-user"
|
||||
|
||||
|
||||
# Verify the returned link uses the new invitation ID
|
||||
assert result == "http://test.com/ui?invitation_id=new-invitation-from-none"
|
||||
|
||||
|
|
@ -495,13 +502,15 @@ async def test_get_email_params_user_invitation(
|
|||
user_email="test@example.com",
|
||||
)
|
||||
|
||||
assert result.logo_url == "https://litellm-listing.s3.amazonaws.com/litellm_logo.png"
|
||||
assert (
|
||||
result.logo_url
|
||||
== "https://litellm-listing.s3.amazonaws.com/litellm_logo.png"
|
||||
)
|
||||
assert result.support_contact == "support@berri.ai"
|
||||
assert result.base_url == "http://test.com/ui?invitation_id=test-id"
|
||||
assert result.recipient_email == "test@example.com"
|
||||
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_env_vars(monkeypatch):
|
||||
"""Set up test environment variables"""
|
||||
|
|
@ -513,69 +522,74 @@ def mock_env_vars(monkeypatch):
|
|||
monkeypatch.setenv("PROXY_BASE_URL", "http://test.com")
|
||||
monkeypatch.setenv("PROXY_API_URL", "https://test.com")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_email_params_custom_templates_premium_user(mock_env_vars):
|
||||
"""Test that _get_email_params returns correct values with custom templates for premium users"""
|
||||
# Mock premium_user as True
|
||||
with patch("litellm.proxy.proxy_server.premium_user", True):
|
||||
email_logger = BaseEmailLogger()
|
||||
|
||||
|
||||
# Test invitation email params
|
||||
invitation_params = await email_logger._get_email_params(
|
||||
email_event=EmailEvent.new_user_invitation,
|
||||
user_id="testid",
|
||||
user_email="test@example.com",
|
||||
event_message="New User Invitation"
|
||||
event_message="New User Invitation",
|
||||
)
|
||||
|
||||
|
||||
assert invitation_params.subject == "Welcome to Test Company!"
|
||||
assert invitation_params.signature == "Best regards,\nTest Company Team"
|
||||
assert invitation_params.logo_url == "https://test-company.com/logo.png"
|
||||
assert invitation_params.support_contact == "support@test-company.com"
|
||||
assert invitation_params.base_url == "http://test.com"
|
||||
|
||||
|
||||
# Test key created email params
|
||||
key_params = await email_logger._get_email_params(
|
||||
email_event=EmailEvent.virtual_key_created,
|
||||
user_id="testid",
|
||||
user_email="test@example.com",
|
||||
event_message="API Key Created"
|
||||
event_message="API Key Created",
|
||||
)
|
||||
|
||||
|
||||
assert key_params.subject == "Your Test Company API Key"
|
||||
assert key_params.signature == "Best regards,\nTest Company Team"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_email_params_non_premium_user(mock_env_vars):
|
||||
"""Test that non-premium users get default templates even when custom ones are provided"""
|
||||
# Mock premium_user as False
|
||||
with patch("litellm.proxy.proxy_server.premium_user", False):
|
||||
email_logger = BaseEmailLogger()
|
||||
|
||||
|
||||
# Test invitation email params
|
||||
email_params = await email_logger._get_email_params(
|
||||
email_event=EmailEvent.new_user_invitation,
|
||||
user_email="test@example.com",
|
||||
event_message="New User Invitation"
|
||||
event_message="New User Invitation",
|
||||
)
|
||||
|
||||
|
||||
# Should use default values even though custom values are set in env
|
||||
assert email_params.subject == "LiteLLM: New User Invitation"
|
||||
assert email_params.signature == EMAIL_FOOTER
|
||||
assert email_params.logo_url == "https://litellm-listing.s3.amazonaws.com/litellm_logo.png"
|
||||
assert (
|
||||
email_params.logo_url
|
||||
== "https://litellm-listing.s3.amazonaws.com/litellm_logo.png"
|
||||
)
|
||||
assert email_params.support_contact == "support@berri.ai"
|
||||
|
||||
|
||||
# Test key created email params
|
||||
key_params = await email_logger._get_email_params(
|
||||
email_event=EmailEvent.virtual_key_created,
|
||||
user_email="test@example.com",
|
||||
event_message="API Key Created"
|
||||
event_message="API Key Created",
|
||||
)
|
||||
|
||||
|
||||
assert key_params.subject == "LiteLLM: API Key Created"
|
||||
assert key_params.signature == EMAIL_FOOTER
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_email_params_default_templates(monkeypatch):
|
||||
"""Test that _get_email_params uses default templates when custom ones aren't provided"""
|
||||
|
|
@ -583,28 +597,28 @@ async def test_get_email_params_default_templates(monkeypatch):
|
|||
monkeypatch.delenv("EMAIL_SUBJECT_INVITATION", raising=False)
|
||||
monkeypatch.delenv("EMAIL_SUBJECT_KEY_CREATED", raising=False)
|
||||
monkeypatch.delenv("EMAIL_SIGNATURE", raising=False)
|
||||
|
||||
|
||||
# Mock premium_user as True (shouldn't matter since no custom values are set)
|
||||
with patch("litellm.proxy.proxy_server.premium_user", True):
|
||||
email_logger = BaseEmailLogger()
|
||||
|
||||
|
||||
# Test invitation email params with default template
|
||||
invitation_params = await email_logger._get_email_params(
|
||||
email_event=EmailEvent.new_user_invitation,
|
||||
user_email="test@example.com",
|
||||
event_message="New User Invitation"
|
||||
event_message="New User Invitation",
|
||||
)
|
||||
|
||||
|
||||
assert invitation_params.subject == "LiteLLM: New User Invitation"
|
||||
assert invitation_params.signature == EMAIL_FOOTER
|
||||
|
||||
|
||||
# Test key created email params with default template
|
||||
key_params = await email_logger._get_email_params(
|
||||
email_event=EmailEvent.virtual_key_created,
|
||||
user_email="test@example.com",
|
||||
event_message="API Key Created"
|
||||
event_message="API Key Created",
|
||||
)
|
||||
|
||||
|
||||
assert key_params.subject == "LiteLLM: API Key Created"
|
||||
assert key_params.signature == EMAIL_FOOTER
|
||||
|
||||
|
|
@ -639,7 +653,10 @@ async def test_send_soft_budget_alert_email(
|
|||
call_args = mock_send_email.call_args[1]
|
||||
assert call_args["from_email"] == BaseEmailLogger.DEFAULT_LITELLM_EMAIL
|
||||
assert call_args["to_email"] == ["test@example.com"]
|
||||
assert call_args["subject"] == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
|
||||
assert (
|
||||
call_args["subject"]
|
||||
== "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
|
||||
)
|
||||
assert "$100.0" in call_args["html_body"] # soft_budget
|
||||
assert "$105.0" in call_args["html_body"] # spend
|
||||
assert "$200.0" in call_args["html_body"] # max_budget
|
||||
|
|
@ -673,13 +690,13 @@ async def test_send_soft_budget_alert_email_no_max_budget(
|
|||
call_args = mock_send_email.call_args[1]
|
||||
assert "$100.0" in call_args["html_body"] # soft_budget
|
||||
assert "$105.0" in call_args["html_body"] # spend
|
||||
assert "Maximum Budget" not in call_args["html_body"] # max_budget should not be shown
|
||||
assert (
|
||||
"Maximum Budget" not in call_args["html_body"]
|
||||
) # max_budget should not be shown
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_budget_alerts_soft_budget_crossed(
|
||||
base_email_logger, mock_send_email
|
||||
):
|
||||
async def test_budget_alerts_soft_budget_crossed(base_email_logger, mock_send_email):
|
||||
"""Test that budget_alerts sends email when soft budget is crossed"""
|
||||
user_info = CallInfo(
|
||||
user_id="test_user",
|
||||
|
|
@ -708,11 +725,14 @@ async def test_budget_alerts_soft_budget_crossed(
|
|||
mock_send_email.assert_called_once()
|
||||
call_args = mock_send_email.call_args[1]
|
||||
assert call_args["to_email"] == ["test@example.com"]
|
||||
|
||||
|
||||
# Verify cache was set to prevent duplicate alerts
|
||||
mock_cache.async_set_cache.assert_called_once()
|
||||
cache_call_args = mock_cache.async_set_cache.call_args[1]
|
||||
assert cache_call_args["key"] == "email_budget_alerts:soft_budget_crossed:test_user"
|
||||
assert (
|
||||
cache_call_args["key"]
|
||||
== "email_budget_alerts:soft_budget_crossed:test_user"
|
||||
)
|
||||
assert cache_call_args["value"] == "SENT"
|
||||
assert cache_call_args["ttl"] == EMAIL_BUDGET_ALERT_TTL
|
||||
|
||||
|
|
@ -766,9 +786,7 @@ async def test_budget_alerts_soft_budget_duplicate_prevention(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_budget_alerts_no_budgets(
|
||||
base_email_logger, mock_send_email
|
||||
):
|
||||
async def test_budget_alerts_no_budgets(base_email_logger, mock_send_email):
|
||||
"""Test that budget_alerts returns early when no budgets are set"""
|
||||
user_info = CallInfo(
|
||||
user_id="test_user",
|
||||
|
|
@ -817,7 +835,10 @@ async def test_budget_alerts_uses_token_for_cache_key(
|
|||
# Verify cache key uses token instead of user_id
|
||||
mock_cache.async_set_cache.assert_called_once()
|
||||
cache_call_args = mock_cache.async_set_cache.call_args[1]
|
||||
assert cache_call_args["key"] == "email_budget_alerts:soft_budget_crossed:hashed_token_123"
|
||||
assert (
|
||||
cache_call_args["key"]
|
||||
== "email_budget_alerts:soft_budget_crossed:hashed_token_123"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -838,7 +859,9 @@ async def test_get_email_params_soft_budget_crossed(
|
|||
)
|
||||
|
||||
# Should use default subject template for soft_budget_crossed
|
||||
assert result.subject == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
|
||||
assert (
|
||||
result.subject == "LiteLLM: Soft Budget Crossed - Total Soft Budget: $100.0"
|
||||
)
|
||||
assert result.recipient_email == "test@example.com"
|
||||
assert result.base_url == "http://test.com"
|
||||
|
||||
|
|
@ -867,15 +890,19 @@ async def test_budget_alerts_max_budget_alert_crossed(
|
|||
"PROXY_BASE_URL": "http://test.com",
|
||||
},
|
||||
):
|
||||
await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info)
|
||||
await base_email_logger.budget_alerts(
|
||||
type="max_budget_alert", user_info=user_info
|
||||
)
|
||||
|
||||
mock_send_email.assert_called_once()
|
||||
call_args = mock_send_email.call_args[1]
|
||||
assert call_args["to_email"] == ["test@example.com"]
|
||||
assert "Max Budget Alert" in call_args["subject"]
|
||||
|
||||
|
||||
mock_cache.async_set_cache.assert_called_once()
|
||||
cache_call_args = mock_cache.async_set_cache.call_args[1]
|
||||
assert cache_call_args["key"] == "email_budget_alerts:max_budget_alert:test_user"
|
||||
assert (
|
||||
cache_call_args["key"] == "email_budget_alerts:max_budget_alert:test_user"
|
||||
)
|
||||
assert cache_call_args["value"] == "SENT"
|
||||
assert cache_call_args["ttl"] == EMAIL_BUDGET_ALERT_TTL
|
||||
assert cache_call_args["ttl"] == EMAIL_BUDGET_ALERT_TTL
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue