mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mistral): strip metadata field from messages to prevent extra_forbidden error
This commit is contained in:
parent
15aa40b36e
commit
1c1638ee8c
310 changed files with 3415 additions and 2081 deletions
2
.github/scripts/close_low_quality_prs.py
vendored
2
.github/scripts/close_low_quality_prs.py
vendored
|
|
@ -53,8 +53,6 @@ from agent_shin_shared import ( # noqa: E402 -- sys.path adjusted above
|
|||
ALLOWLIST_LOGINS,
|
||||
GRACE_COMMENT_MARKER,
|
||||
GRACE_PERIOD_SECONDS,
|
||||
GREPTILE_BOT_LOGINS,
|
||||
SCORE_PATTERN,
|
||||
extract_greptile_score,
|
||||
gh,
|
||||
list_open_items,
|
||||
|
|
|
|||
1
.github/scripts/triage_rollout_heads_up.py
vendored
1
.github/scripts/triage_rollout_heads_up.py
vendored
|
|
@ -53,7 +53,6 @@ from agent_shin_shared import ( # noqa: E402
|
|||
)
|
||||
from triage_with_llm import ( # noqa: E402
|
||||
DEFAULT_MODEL,
|
||||
call_llm_judge,
|
||||
fetch_issue,
|
||||
fetch_pr,
|
||||
gh,
|
||||
|
|
|
|||
2
.github/scripts/triage_with_llm.py
vendored
2
.github/scripts/triage_with_llm.py
vendored
|
|
@ -53,8 +53,6 @@ from agent_shin_shared import ( # noqa: E402 -- sys.path adjusted above
|
|||
ALLOWLIST_LOGINS,
|
||||
GRACE_COMMENT_MARKER,
|
||||
GRACE_PERIOD_SECONDS,
|
||||
GREPTILE_BOT_LOGINS,
|
||||
SCORE_PATTERN,
|
||||
extract_greptile_score,
|
||||
gh,
|
||||
parse_iso8601,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
515
.github/workflows/run_llm_translation_tests.py
vendored
515
.github/workflows/run_llm_translation_tests.py
vendored
|
|
@ -12,68 +12,76 @@ import subprocess
|
|||
import xml.etree.ElementTree as ET
|
||||
from collections import defaultdict
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
import json
|
||||
from typing import Dict, List, Tuple, Optional
|
||||
|
||||
|
||||
# ANSI color codes for terminal output
|
||||
class Colors:
|
||||
GREEN = '\033[92m'
|
||||
RED = '\033[91m'
|
||||
YELLOW = '\033[93m'
|
||||
BLUE = '\033[94m'
|
||||
PURPLE = '\033[95m'
|
||||
CYAN = '\033[96m'
|
||||
RESET = '\033[0m'
|
||||
BOLD = '\033[1m'
|
||||
GREEN = "\033[92m"
|
||||
RED = "\033[91m"
|
||||
YELLOW = "\033[93m"
|
||||
BLUE = "\033[94m"
|
||||
PURPLE = "\033[95m"
|
||||
CYAN = "\033[96m"
|
||||
RESET = "\033[0m"
|
||||
BOLD = "\033[1m"
|
||||
|
||||
|
||||
def print_colored(message: str, color: str = Colors.RESET):
|
||||
"""Print colored message to terminal"""
|
||||
print(f"{color}{message}{Colors.RESET}")
|
||||
|
||||
|
||||
def get_provider_from_test_file(test_file: str) -> str:
|
||||
"""Map test file names to provider names"""
|
||||
provider_mapping = {
|
||||
'test_anthropic': 'Anthropic',
|
||||
'test_azure': 'Azure',
|
||||
'test_bedrock': 'AWS Bedrock',
|
||||
'test_openai': 'OpenAI',
|
||||
'test_vertex': 'Google Vertex AI',
|
||||
'test_gemini': 'Google Vertex AI',
|
||||
'test_cohere': 'Cohere',
|
||||
'test_databricks': 'Databricks',
|
||||
'test_groq': 'Groq',
|
||||
'test_together': 'Together AI',
|
||||
'test_mistral': 'Mistral',
|
||||
'test_deepseek': 'DeepSeek',
|
||||
'test_replicate': 'Replicate',
|
||||
'test_huggingface': 'HuggingFace',
|
||||
'test_fireworks': 'Fireworks AI',
|
||||
'test_perplexity': 'Perplexity',
|
||||
'test_cloudflare': 'Cloudflare',
|
||||
'test_voyage': 'Voyage AI',
|
||||
'test_xai': 'xAI',
|
||||
'test_nvidia': 'NVIDIA',
|
||||
'test_watsonx': 'IBM watsonx',
|
||||
'test_azure_ai': 'Azure AI',
|
||||
'test_snowflake': 'Snowflake',
|
||||
'test_infinity': 'Infinity',
|
||||
'test_jina': 'Jina AI',
|
||||
'test_deepgram': 'Deepgram',
|
||||
'test_clarifai': 'Clarifai',
|
||||
'test_triton': 'Triton',
|
||||
"test_anthropic": "Anthropic",
|
||||
"test_azure": "Azure",
|
||||
"test_bedrock": "AWS Bedrock",
|
||||
"test_openai": "OpenAI",
|
||||
"test_vertex": "Google Vertex AI",
|
||||
"test_gemini": "Google Vertex AI",
|
||||
"test_cohere": "Cohere",
|
||||
"test_databricks": "Databricks",
|
||||
"test_groq": "Groq",
|
||||
"test_together": "Together AI",
|
||||
"test_mistral": "Mistral",
|
||||
"test_deepseek": "DeepSeek",
|
||||
"test_replicate": "Replicate",
|
||||
"test_huggingface": "HuggingFace",
|
||||
"test_fireworks": "Fireworks AI",
|
||||
"test_perplexity": "Perplexity",
|
||||
"test_cloudflare": "Cloudflare",
|
||||
"test_voyage": "Voyage AI",
|
||||
"test_xai": "xAI",
|
||||
"test_nvidia": "NVIDIA",
|
||||
"test_watsonx": "IBM watsonx",
|
||||
"test_azure_ai": "Azure AI",
|
||||
"test_snowflake": "Snowflake",
|
||||
"test_infinity": "Infinity",
|
||||
"test_jina": "Jina AI",
|
||||
"test_deepgram": "Deepgram",
|
||||
"test_clarifai": "Clarifai",
|
||||
"test_triton": "Triton",
|
||||
}
|
||||
|
||||
|
||||
for key, provider in provider_mapping.items():
|
||||
if key in test_file:
|
||||
return provider
|
||||
|
||||
|
||||
# For cross-provider test files
|
||||
if any(name in test_file for name in ['test_optional_params', 'test_prompt_factory',
|
||||
'test_router', 'test_text_completion']):
|
||||
return f'Cross-Provider Tests ({test_file})'
|
||||
|
||||
return 'Other Tests'
|
||||
if any(
|
||||
name in test_file
|
||||
for name in [
|
||||
"test_optional_params",
|
||||
"test_prompt_factory",
|
||||
"test_router",
|
||||
"test_text_completion",
|
||||
]
|
||||
):
|
||||
return f"Cross-Provider Tests ({test_file})"
|
||||
|
||||
return "Other Tests"
|
||||
|
||||
|
||||
def format_duration(seconds: float) -> str:
|
||||
"""Format duration in human-readable format"""
|
||||
|
|
@ -89,290 +97,355 @@ def format_duration(seconds: float) -> str:
|
|||
return f"{hours}h {minutes}m"
|
||||
|
||||
|
||||
def generate_markdown_report(junit_xml_path: str, output_path: str, tag: str = None, commit: str = None):
|
||||
def generate_markdown_report(
|
||||
junit_xml_path: str, output_path: str, tag: str = None, commit: str = None
|
||||
):
|
||||
"""Generate a beautiful markdown report from JUnit XML"""
|
||||
try:
|
||||
tree = ET.parse(junit_xml_path)
|
||||
root = tree.getroot()
|
||||
|
||||
|
||||
# Handle both testsuite and testsuites root
|
||||
if root.tag == 'testsuites':
|
||||
suites = root.findall('testsuite')
|
||||
if root.tag == "testsuites":
|
||||
suites = root.findall("testsuite")
|
||||
else:
|
||||
suites = [root]
|
||||
|
||||
|
||||
# Overall statistics
|
||||
total_tests = 0
|
||||
total_failures = 0
|
||||
total_errors = 0
|
||||
total_skipped = 0
|
||||
total_time = 0.0
|
||||
|
||||
|
||||
# Provider breakdown
|
||||
provider_stats = defaultdict(lambda: {'passed': 0, 'failed': 0, 'skipped': 0, 'errors': 0, 'time': 0.0})
|
||||
provider_stats = defaultdict(
|
||||
lambda: {"passed": 0, "failed": 0, "skipped": 0, "errors": 0, "time": 0.0}
|
||||
)
|
||||
provider_tests = defaultdict(list)
|
||||
|
||||
|
||||
for suite in suites:
|
||||
total_tests += int(suite.get('tests', 0))
|
||||
total_failures += int(suite.get('failures', 0))
|
||||
total_errors += int(suite.get('errors', 0))
|
||||
total_skipped += int(suite.get('skipped', 0))
|
||||
total_time += float(suite.get('time', 0))
|
||||
|
||||
for testcase in suite.findall('testcase'):
|
||||
classname = testcase.get('classname', '')
|
||||
test_name = testcase.get('name', '')
|
||||
test_time = float(testcase.get('time', 0))
|
||||
|
||||
total_tests += int(suite.get("tests", 0))
|
||||
total_failures += int(suite.get("failures", 0))
|
||||
total_errors += int(suite.get("errors", 0))
|
||||
total_skipped += int(suite.get("skipped", 0))
|
||||
total_time += float(suite.get("time", 0))
|
||||
|
||||
for testcase in suite.findall("testcase"):
|
||||
classname = testcase.get("classname", "")
|
||||
test_name = testcase.get("name", "")
|
||||
test_time = float(testcase.get("time", 0))
|
||||
|
||||
# Extract test file name from classname
|
||||
if '.' in classname:
|
||||
parts = classname.split('.')
|
||||
test_file = parts[-2] if len(parts) > 1 else 'unknown'
|
||||
if "." in classname:
|
||||
parts = classname.split(".")
|
||||
test_file = parts[-2] if len(parts) > 1 else "unknown"
|
||||
else:
|
||||
test_file = 'unknown'
|
||||
|
||||
test_file = "unknown"
|
||||
|
||||
provider = get_provider_from_test_file(test_file)
|
||||
provider_stats[provider]['time'] += test_time
|
||||
|
||||
provider_stats[provider]["time"] += test_time
|
||||
|
||||
# Check test status
|
||||
if testcase.find('failure') is not None:
|
||||
provider_stats[provider]['failed'] += 1
|
||||
failure = testcase.find('failure')
|
||||
failure_msg = failure.get('message', '') if failure is not None else ''
|
||||
provider_tests[provider].append({
|
||||
'name': test_name,
|
||||
'status': 'FAILED',
|
||||
'time': test_time,
|
||||
'message': failure_msg
|
||||
})
|
||||
elif testcase.find('error') is not None:
|
||||
provider_stats[provider]['errors'] += 1
|
||||
error = testcase.find('error')
|
||||
error_msg = error.get('message', '') if error is not None else ''
|
||||
provider_tests[provider].append({
|
||||
'name': test_name,
|
||||
'status': 'ERROR',
|
||||
'time': test_time,
|
||||
'message': error_msg
|
||||
})
|
||||
elif testcase.find('skipped') is not None:
|
||||
provider_stats[provider]['skipped'] += 1
|
||||
skip = testcase.find('skipped')
|
||||
skip_msg = skip.get('message', '') if skip is not None else ''
|
||||
provider_tests[provider].append({
|
||||
'name': test_name,
|
||||
'status': 'SKIPPED',
|
||||
'time': test_time,
|
||||
'message': skip_msg
|
||||
})
|
||||
if testcase.find("failure") is not None:
|
||||
provider_stats[provider]["failed"] += 1
|
||||
failure = testcase.find("failure")
|
||||
failure_msg = (
|
||||
failure.get("message", "") if failure is not None else ""
|
||||
)
|
||||
provider_tests[provider].append(
|
||||
{
|
||||
"name": test_name,
|
||||
"status": "FAILED",
|
||||
"time": test_time,
|
||||
"message": failure_msg,
|
||||
}
|
||||
)
|
||||
elif testcase.find("error") is not None:
|
||||
provider_stats[provider]["errors"] += 1
|
||||
error = testcase.find("error")
|
||||
error_msg = error.get("message", "") if error is not None else ""
|
||||
provider_tests[provider].append(
|
||||
{
|
||||
"name": test_name,
|
||||
"status": "ERROR",
|
||||
"time": test_time,
|
||||
"message": error_msg,
|
||||
}
|
||||
)
|
||||
elif testcase.find("skipped") is not None:
|
||||
provider_stats[provider]["skipped"] += 1
|
||||
skip = testcase.find("skipped")
|
||||
skip_msg = skip.get("message", "") if skip is not None else ""
|
||||
provider_tests[provider].append(
|
||||
{
|
||||
"name": test_name,
|
||||
"status": "SKIPPED",
|
||||
"time": test_time,
|
||||
"message": skip_msg,
|
||||
}
|
||||
)
|
||||
else:
|
||||
provider_stats[provider]['passed'] += 1
|
||||
provider_tests[provider].append({
|
||||
'name': test_name,
|
||||
'status': 'PASSED',
|
||||
'time': test_time,
|
||||
'message': ''
|
||||
})
|
||||
|
||||
provider_stats[provider]["passed"] += 1
|
||||
provider_tests[provider].append(
|
||||
{
|
||||
"name": test_name,
|
||||
"status": "PASSED",
|
||||
"time": test_time,
|
||||
"message": "",
|
||||
}
|
||||
)
|
||||
|
||||
passed = total_tests - total_failures - total_errors - total_skipped
|
||||
|
||||
|
||||
# Generate the markdown report
|
||||
with open(output_path, 'w') as f:
|
||||
with open(output_path, "w") as f:
|
||||
# Header
|
||||
f.write("# LLM Translation Test Results\n\n")
|
||||
|
||||
|
||||
# Metadata table
|
||||
f.write("## Test Run Information\n\n")
|
||||
f.write("| Field | Value |\n")
|
||||
f.write("|-------|-------|\n")
|
||||
f.write(f"| **Tag** | `{tag or 'N/A'}` |\n")
|
||||
f.write(f"| **Date** | {datetime.utcnow().strftime('%Y-%m-%d %H:%M:%S UTC')} |\n")
|
||||
f.write(
|
||||
f"| **Date** | {datetime.utcnow().strftime('%Y-%m-%d %H:%M:%S UTC')} |\n"
|
||||
)
|
||||
f.write(f"| **Commit** | `{commit or 'N/A'}` |\n")
|
||||
f.write(f"| **Duration** | {format_duration(total_time)} |\n")
|
||||
f.write("\n")
|
||||
|
||||
|
||||
# Overall statistics with visual elements
|
||||
f.write("## Overall Statistics\n\n")
|
||||
|
||||
|
||||
# Summary box
|
||||
f.write("```\n")
|
||||
f.write(f"Total Tests: {total_tests}\n")
|
||||
f.write(f"├── Passed: {passed:>4} ({(passed/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
|
||||
f.write(f"├── Failed: {total_failures:>4} ({(total_failures/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
|
||||
f.write(f"├── Errors: {total_errors:>4} ({(total_errors/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
|
||||
f.write(f"└── Skipped: {total_skipped:>4} ({(total_skipped/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n")
|
||||
f.write(
|
||||
f"├── Passed: {passed:>4} ({(passed/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
|
||||
)
|
||||
f.write(
|
||||
f"├── Failed: {total_failures:>4} ({(total_failures/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
|
||||
)
|
||||
f.write(
|
||||
f"├── Errors: {total_errors:>4} ({(total_errors/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
|
||||
)
|
||||
f.write(
|
||||
f"└── Skipped: {total_skipped:>4} ({(total_skipped/total_tests)*100 if total_tests > 0 else 0:.1f}%)\n"
|
||||
)
|
||||
f.write("```\n\n")
|
||||
|
||||
|
||||
|
||||
# Provider summary table
|
||||
f.write("## Results by Provider\n\n")
|
||||
f.write("| Provider | Total | Pass | Fail | Error | Skip | Pass Rate | Duration |\n")
|
||||
f.write("|----------|-------|------|------|-------|------|-----------|----------|")
|
||||
|
||||
f.write(
|
||||
"| Provider | Total | Pass | Fail | Error | Skip | Pass Rate | Duration |\n"
|
||||
)
|
||||
f.write(
|
||||
"|----------|-------|------|------|-------|------|-----------|----------|"
|
||||
)
|
||||
|
||||
# Sort providers: specific providers first, then cross-provider tests
|
||||
sorted_providers = []
|
||||
cross_provider = []
|
||||
for p in sorted(provider_stats.keys()):
|
||||
if 'Cross-Provider' in p or p == 'Other Tests':
|
||||
if "Cross-Provider" in p or p == "Other Tests":
|
||||
cross_provider.append(p)
|
||||
else:
|
||||
sorted_providers.append(p)
|
||||
|
||||
|
||||
all_providers = sorted_providers + cross_provider
|
||||
|
||||
|
||||
for provider in all_providers:
|
||||
stats = provider_stats[provider]
|
||||
total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped']
|
||||
pass_rate = (stats['passed'] / total * 100) if total > 0 else 0
|
||||
|
||||
f.write(f"\n| {provider} | {total} | {stats['passed']} | {stats['failed']} | ")
|
||||
total = (
|
||||
stats["passed"]
|
||||
+ stats["failed"]
|
||||
+ stats["errors"]
|
||||
+ stats["skipped"]
|
||||
)
|
||||
pass_rate = (stats["passed"] / total * 100) if total > 0 else 0
|
||||
|
||||
f.write(
|
||||
f"\n| {provider} | {total} | {stats['passed']} | {stats['failed']} | "
|
||||
)
|
||||
f.write(f"{stats['errors']} | {stats['skipped']} | {pass_rate:.1f}% | ")
|
||||
f.write(f"{format_duration(stats['time'])} |")
|
||||
|
||||
|
||||
# Detailed test results by provider
|
||||
f.write("\n\n## Detailed Test Results\n\n")
|
||||
|
||||
|
||||
for provider in sorted_providers:
|
||||
if provider_tests[provider]:
|
||||
stats = provider_stats[provider]
|
||||
total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped']
|
||||
|
||||
total = (
|
||||
stats["passed"]
|
||||
+ stats["failed"]
|
||||
+ stats["errors"]
|
||||
+ stats["skipped"]
|
||||
)
|
||||
|
||||
f.write(f"### {provider}\n\n")
|
||||
f.write(f"**Summary:** {stats['passed']}/{total} passed ")
|
||||
f.write(f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%) ")
|
||||
f.write(
|
||||
f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%) "
|
||||
)
|
||||
f.write(f"in {format_duration(stats['time'])}\n\n")
|
||||
|
||||
|
||||
# Group tests by status
|
||||
tests_by_status = defaultdict(list)
|
||||
for test in provider_tests[provider]:
|
||||
tests_by_status[test['status']].append(test)
|
||||
|
||||
tests_by_status[test["status"]].append(test)
|
||||
|
||||
# Show failed tests first (if any)
|
||||
if tests_by_status['FAILED']:
|
||||
if tests_by_status["FAILED"]:
|
||||
f.write("<details>\n<summary>Failed Tests</summary>\n\n")
|
||||
for test in tests_by_status['FAILED']:
|
||||
for test in tests_by_status["FAILED"]:
|
||||
f.write(f"- `{test['name']}` ({test['time']:.2f}s)\n")
|
||||
if test['message']:
|
||||
if test["message"]:
|
||||
# Truncate long error messages
|
||||
msg = test['message'][:200] + '...' if len(test['message']) > 200 else test['message']
|
||||
msg = (
|
||||
test["message"][:200] + "..."
|
||||
if len(test["message"]) > 200
|
||||
else test["message"]
|
||||
)
|
||||
f.write(f" > {msg}\n")
|
||||
f.write("\n</details>\n\n")
|
||||
|
||||
|
||||
# Show errors (if any)
|
||||
if tests_by_status['ERROR']:
|
||||
if tests_by_status["ERROR"]:
|
||||
f.write("<details>\n<summary>Error Tests</summary>\n\n")
|
||||
for test in tests_by_status['ERROR']:
|
||||
for test in tests_by_status["ERROR"]:
|
||||
f.write(f"- `{test['name']}` ({test['time']:.2f}s)\n")
|
||||
f.write("\n</details>\n\n")
|
||||
|
||||
|
||||
# Show passed tests in collapsible section
|
||||
if tests_by_status['PASSED']:
|
||||
if tests_by_status["PASSED"]:
|
||||
f.write("<details>\n<summary>Passed Tests</summary>\n\n")
|
||||
for test in tests_by_status['PASSED']:
|
||||
for test in tests_by_status["PASSED"]:
|
||||
f.write(f"- `{test['name']}` ({test['time']:.2f}s)\n")
|
||||
f.write("\n</details>\n\n")
|
||||
|
||||
|
||||
# Show skipped tests (if any)
|
||||
if tests_by_status['SKIPPED']:
|
||||
if tests_by_status["SKIPPED"]:
|
||||
f.write("<details>\n<summary>Skipped Tests</summary>\n\n")
|
||||
for test in tests_by_status['SKIPPED']:
|
||||
for test in tests_by_status["SKIPPED"]:
|
||||
f.write(f"- `{test['name']}`\n")
|
||||
f.write("\n</details>\n\n")
|
||||
|
||||
|
||||
# Cross-provider tests in a separate section
|
||||
if cross_provider:
|
||||
f.write("### Cross-Provider Tests\n\n")
|
||||
for provider in cross_provider:
|
||||
if provider_tests[provider]:
|
||||
stats = provider_stats[provider]
|
||||
total = stats['passed'] + stats['failed'] + stats['errors'] + stats['skipped']
|
||||
|
||||
total = (
|
||||
stats["passed"]
|
||||
+ stats["failed"]
|
||||
+ stats["errors"]
|
||||
+ stats["skipped"]
|
||||
)
|
||||
|
||||
f.write(f"#### {provider}\n\n")
|
||||
f.write(f"**Summary:** {stats['passed']}/{total} passed ")
|
||||
f.write(f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%)\n\n")
|
||||
|
||||
f.write(
|
||||
f"({(stats['passed']/total)*100 if total > 0 else 0:.1f}%)\n\n"
|
||||
)
|
||||
|
||||
# For cross-provider tests, just show counts
|
||||
f.write(f"- Passed: {stats['passed']}\n")
|
||||
if stats['failed'] > 0:
|
||||
if stats["failed"] > 0:
|
||||
f.write(f"- Failed: {stats['failed']}\n")
|
||||
if stats['errors'] > 0:
|
||||
if stats["errors"] > 0:
|
||||
f.write(f"- Errors: {stats['errors']}\n")
|
||||
if stats['skipped'] > 0:
|
||||
if stats["skipped"] > 0:
|
||||
f.write(f"- Skipped: {stats['skipped']}\n")
|
||||
f.write("\n")
|
||||
|
||||
|
||||
|
||||
print_colored(f"Report generated: {output_path}", Colors.GREEN)
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print_colored(f"Error generating report: {e}", Colors.RED)
|
||||
raise
|
||||
|
||||
def run_tests(test_path: str = "tests/llm_translation/",
|
||||
junit_xml: str = "test-results/junit.xml",
|
||||
report_path: str = "test-results/llm_translation_report.md",
|
||||
tag: str = None,
|
||||
commit: str = None) -> int:
|
||||
|
||||
def run_tests(
|
||||
test_path: str = "tests/llm_translation/",
|
||||
junit_xml: str = "test-results/junit.xml",
|
||||
report_path: str = "test-results/llm_translation_report.md",
|
||||
tag: str = None,
|
||||
commit: str = None,
|
||||
) -> int:
|
||||
"""Run the LLM translation tests and generate report"""
|
||||
|
||||
|
||||
# Create test results directory
|
||||
os.makedirs(os.path.dirname(junit_xml), exist_ok=True)
|
||||
|
||||
|
||||
print_colored("Starting LLM Translation Tests", Colors.BOLD + Colors.BLUE)
|
||||
print_colored(f"Test directory: {test_path}", Colors.CYAN)
|
||||
print_colored(f"Output: {junit_xml}", Colors.CYAN)
|
||||
print()
|
||||
|
||||
|
||||
# Run pytest
|
||||
cmd = [
|
||||
"uv", "run", "--no-sync", "pytest", test_path,
|
||||
"uv",
|
||||
"run",
|
||||
"--no-sync",
|
||||
"pytest",
|
||||
test_path,
|
||||
f"--junitxml={junit_xml}",
|
||||
"-v",
|
||||
"--tb=short",
|
||||
"--maxfail=500",
|
||||
"-n", "auto"
|
||||
"-n",
|
||||
"auto",
|
||||
]
|
||||
|
||||
|
||||
# Add timeout if pytest-timeout is installed
|
||||
try:
|
||||
subprocess.run(["uv", "run", "--no-sync", "python", "-c", "import pytest_timeout"],
|
||||
capture_output=True, check=True)
|
||||
subprocess.run(
|
||||
["uv", "run", "--no-sync", "python", "-c", "import pytest_timeout"],
|
||||
capture_output=True,
|
||||
check=True,
|
||||
)
|
||||
cmd.extend(["--timeout=300"])
|
||||
except:
|
||||
print_colored("Warning: pytest-timeout not installed, skipping timeout option", Colors.YELLOW)
|
||||
|
||||
print_colored(
|
||||
"Warning: pytest-timeout not installed, skipping timeout option",
|
||||
Colors.YELLOW,
|
||||
)
|
||||
|
||||
print_colored("Running pytest with command:", Colors.YELLOW)
|
||||
print(f" {' '.join(cmd)}")
|
||||
print()
|
||||
|
||||
|
||||
# Run the tests
|
||||
result = subprocess.run(cmd, capture_output=False)
|
||||
|
||||
|
||||
# Generate the report regardless of test outcome
|
||||
if os.path.exists(junit_xml):
|
||||
print()
|
||||
print_colored("Generating test report...", Colors.BLUE)
|
||||
generate_markdown_report(junit_xml, report_path, tag, commit)
|
||||
|
||||
|
||||
# Print summary to console
|
||||
print()
|
||||
print_colored("Test Summary:", Colors.BOLD + Colors.PURPLE)
|
||||
|
||||
|
||||
# Parse XML for quick summary
|
||||
tree = ET.parse(junit_xml)
|
||||
root = tree.getroot()
|
||||
|
||||
if root.tag == 'testsuites':
|
||||
suites = root.findall('testsuite')
|
||||
|
||||
if root.tag == "testsuites":
|
||||
suites = root.findall("testsuite")
|
||||
else:
|
||||
suites = [root]
|
||||
|
||||
total = sum(int(s.get('tests', 0)) for s in suites)
|
||||
failures = sum(int(s.get('failures', 0)) for s in suites)
|
||||
errors = sum(int(s.get('errors', 0)) for s in suites)
|
||||
skipped = sum(int(s.get('skipped', 0)) for s in suites)
|
||||
|
||||
total = sum(int(s.get("tests", 0)) for s in suites)
|
||||
failures = sum(int(s.get("failures", 0)) for s in suites)
|
||||
errors = sum(int(s.get("errors", 0)) for s in suites)
|
||||
skipped = sum(int(s.get("skipped", 0)) for s in suites)
|
||||
passed = total - failures - errors - skipped
|
||||
|
||||
|
||||
print(f" Total: {total}")
|
||||
print_colored(f" Passed: {passed}", Colors.GREEN)
|
||||
if failures > 0:
|
||||
|
|
@ -381,59 +454,75 @@ def run_tests(test_path: str = "tests/llm_translation/",
|
|||
print_colored(f" Errors: {errors}", Colors.RED)
|
||||
if skipped > 0:
|
||||
print_colored(f" Skipped: {skipped}", Colors.YELLOW)
|
||||
|
||||
|
||||
if total > 0:
|
||||
pass_rate = (passed / total) * 100
|
||||
color = Colors.GREEN if pass_rate >= 80 else Colors.YELLOW if pass_rate >= 60 else Colors.RED
|
||||
color = (
|
||||
Colors.GREEN
|
||||
if pass_rate >= 80
|
||||
else Colors.YELLOW if pass_rate >= 60 else Colors.RED
|
||||
)
|
||||
print_colored(f" Pass Rate: {pass_rate:.1f}%", color)
|
||||
else:
|
||||
print_colored("No test results found!", Colors.RED)
|
||||
|
||||
|
||||
print()
|
||||
print_colored("Test run complete!", Colors.BOLD + Colors.GREEN)
|
||||
|
||||
|
||||
return result.returncode
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser(description="Run LLM Translation Tests")
|
||||
parser.add_argument("--test-path", default="tests/llm_translation/",
|
||||
help="Path to test directory")
|
||||
parser.add_argument("--junit-xml", default="test-results/junit.xml",
|
||||
help="Path for JUnit XML output")
|
||||
parser.add_argument("--report", default="test-results/llm_translation_report.md",
|
||||
help="Path for markdown report")
|
||||
parser.add_argument(
|
||||
"--test-path", default="tests/llm_translation/", help="Path to test directory"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--junit-xml",
|
||||
default="test-results/junit.xml",
|
||||
help="Path for JUnit XML output",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--report",
|
||||
default="test-results/llm_translation_report.md",
|
||||
help="Path for markdown report",
|
||||
)
|
||||
parser.add_argument("--tag", help="Git tag or version")
|
||||
parser.add_argument("--commit", help="Git commit SHA")
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
# Get git info if not provided
|
||||
if not args.commit:
|
||||
try:
|
||||
result = subprocess.run(["git", "rev-parse", "HEAD"],
|
||||
capture_output=True, text=True)
|
||||
result = subprocess.run(
|
||||
["git", "rev-parse", "HEAD"], capture_output=True, text=True
|
||||
)
|
||||
if result.returncode == 0:
|
||||
args.commit = result.stdout.strip()
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
if not args.tag:
|
||||
try:
|
||||
result = subprocess.run(["git", "describe", "--tags", "--abbrev=0"],
|
||||
capture_output=True, text=True)
|
||||
result = subprocess.run(
|
||||
["git", "describe", "--tags", "--abbrev=0"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
if result.returncode == 0:
|
||||
args.tag = result.stdout.strip()
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
exit_code = run_tests(
|
||||
test_path=args.test_path,
|
||||
junit_xml=args.junit_xml,
|
||||
report_path=args.report,
|
||||
tag=args.tag,
|
||||
commit=args.commit
|
||||
commit=args.commit,
|
||||
)
|
||||
|
||||
|
||||
sys.exit(exit_code)
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ from pathlib import Path
|
|||
|
||||
import testing.postgresql
|
||||
|
||||
|
||||
DESTRUCTIVE_PATTERN = re.compile(r"\bDROP\s+(COLUMN|TABLE|INDEX)\b", re.IGNORECASE)
|
||||
DEFAULT_BASE_BRANCH = "litellm_internal_staging"
|
||||
|
||||
|
|
@ -45,7 +44,7 @@ def _print_freshness_failure(
|
|||
print("", file=out)
|
||||
print("Options:", file=out)
|
||||
print(
|
||||
f" - Fix the above and re-run, OR pass --base-branch <name> if your", file=out
|
||||
" - Fix the above and re-run, OR pass --base-branch <name> if your", file=out
|
||||
)
|
||||
print(
|
||||
f" base branch is not '{base_branch}', OR pass --skip-freshness-check",
|
||||
|
|
|
|||
8
cookbook/LiteLLM_CometAPI.ipynb
vendored
8
cookbook/LiteLLM_CometAPI.ipynb
vendored
|
|
@ -257,7 +257,6 @@
|
|||
],
|
||||
"source": [
|
||||
"from litellm import acompletion\n",
|
||||
"import asyncio\n",
|
||||
"\n",
|
||||
"async def test_get_response():\n",
|
||||
" user_message = \"Hello, how are you?\"\n",
|
||||
|
|
@ -313,8 +312,8 @@
|
|||
}
|
||||
],
|
||||
"source": [
|
||||
"from litellm import acompletion\n",
|
||||
"import asyncio, os, traceback\n",
|
||||
"import os\n",
|
||||
"import traceback\n",
|
||||
"\n",
|
||||
"async def completion_call():\n",
|
||||
" try:\n",
|
||||
|
|
@ -384,7 +383,6 @@
|
|||
}
|
||||
],
|
||||
"source": [
|
||||
"import litellm\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"async def main():\n",
|
||||
|
|
@ -421,9 +419,7 @@
|
|||
}
|
||||
],
|
||||
"source": [
|
||||
"import asyncio\n",
|
||||
"\n",
|
||||
"import litellm\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"async def main():\n",
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ from tabulate import tabulate
|
|||
from termcolor import colored
|
||||
import os
|
||||
|
||||
|
||||
# Define the list of models to benchmark
|
||||
# select any LLM listed here: https://docs.litellm.ai/docs/providers
|
||||
models = ["gpt-3.5-turbo", "claude-2"]
|
||||
|
|
|
|||
1
cookbook/google_adk_litellm_tutorial.ipynb
vendored
1
cookbook/google_adk_litellm_tutorial.ipynb
vendored
|
|
@ -74,7 +74,6 @@
|
|||
"source": [
|
||||
"# Setup environment and API keys\n",
|
||||
"import os\n",
|
||||
"import asyncio\n",
|
||||
"from google.adk.agents import Agent\n",
|
||||
"from google.adk.models.lite_llm import LiteLlm # For multi-model support\n",
|
||||
"from google.adk.sessions import InMemorySessionService\n",
|
||||
|
|
|
|||
|
|
@ -15,14 +15,13 @@ Usage:
|
|||
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import httpx
|
||||
from fastapi import FastAPI, HTTPException, Header, Query
|
||||
from fastapi.responses import JSONResponse
|
||||
import uvicorn
|
||||
|
||||
|
||||
app = FastAPI(
|
||||
title="Braintrust Prompt Wrapper",
|
||||
description="Wrapper server for Braintrust prompts to work with LiteLLM",
|
||||
|
|
@ -264,7 +263,7 @@ def main():
|
|||
print(f"🚀 Starting Braintrust Prompt Wrapper Server on {host}:{port}")
|
||||
print(f"📚 API Documentation available at http://{host}:{port}/docs")
|
||||
print(
|
||||
f"🔑 Make sure to set BRAINTRUST_API_KEY environment variable or pass token in Authorization header"
|
||||
"🔑 Make sure to set BRAINTRUST_API_KEY environment variable or pass token in Authorization header"
|
||||
)
|
||||
|
||||
uvicorn.run(app, host=host, port=port)
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ This example shows how to use the CLI authentication token
|
|||
in your Python scripts after running `litellm-proxy login`.
|
||||
"""
|
||||
|
||||
from textwrap import indent
|
||||
import litellm
|
||||
|
||||
LITELLM_BASE_URL = "http://localhost:4000/"
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
Simple xAI Voice Agent using LiveKit SDK with LiteLLM Gateway
|
||||
|
||||
This example shows how to use LiveKit's xAI realtime plugin through LiteLLM proxy.
|
||||
LiteLLM acts as a unified interface, allowing you to switch between xAI, OpenAI,
|
||||
LiteLLM acts as a unified interface, allowing you to switch between xAI, OpenAI,
|
||||
and Azure realtime APIs without changing your agent code.
|
||||
"""
|
||||
|
||||
|
|
@ -28,7 +28,7 @@ async def run_voice_agent():
|
|||
url = f"ws://{PROXY_URL.replace('http://', '').replace('https://', '')}/v1/realtime?model={MODEL}"
|
||||
headers = {"Authorization": f"Bearer {API_KEY}"}
|
||||
|
||||
print(f"🎙️ Connecting to voice agent...")
|
||||
print("🎙️ Connecting to voice agent...")
|
||||
print(f" Model: {MODEL}")
|
||||
print(f" Proxy: {PROXY_URL}")
|
||||
print()
|
||||
|
|
@ -114,7 +114,7 @@ def main():
|
|||
except Exception as e:
|
||||
print(f"\n❌ Error: {e}")
|
||||
print("\nMake sure LiteLLM proxy is running:")
|
||||
print(f" litellm --config config.yaml --port 4000")
|
||||
print(" litellm --config config.yaml --port 4000")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
|
|
@ -1,14 +1,14 @@
|
|||
"""
|
||||
LiteLLM Migration Script!
|
||||
|
||||
Takes a config.yaml and calls /model/new
|
||||
Takes a config.yaml and calls /model/new
|
||||
|
||||
Inputs:
|
||||
- File path to config.yaml
|
||||
- Proxy base url to your hosted proxy
|
||||
|
||||
Step 1: Reads your config.yaml
|
||||
Step 2: reads `model_list` and loops through all models
|
||||
Step 2: reads `model_list` and loops through all models
|
||||
Step 3: calls `<proxy-base-url>/model/new` for each model
|
||||
"""
|
||||
|
||||
|
|
|
|||
|
|
@ -515,11 +515,10 @@ if __name__ == "__main__":
|
|||
print("=" * 80)
|
||||
print(f"Server starting on: http://{host}:{port}")
|
||||
print(f"Bearer Token: {bearer_token}")
|
||||
print(f"Endpoint: POST /guardrail/{{id}}/version/{{version}}/apply")
|
||||
print("Endpoint: POST /guardrail/{id}/version/{version}/apply")
|
||||
print("=" * 80)
|
||||
print("\nExample curl command:")
|
||||
print(
|
||||
f"""
|
||||
print(f"""
|
||||
curl -X POST "http://{host}:{port}/guardrail/test-guardrail/version/1/apply" \\
|
||||
-H "Authorization: Bearer {bearer_token}" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
|
|
@ -533,8 +532,7 @@ curl -X POST "http://{host}:{port}/guardrail/test-guardrail/version/1/apply" \\
|
|||
}}
|
||||
]
|
||||
}}'
|
||||
"""
|
||||
)
|
||||
""")
|
||||
print("=" * 80)
|
||||
|
||||
uvicorn.run(app, host=host, port=port)
|
||||
|
|
|
|||
|
|
@ -14,12 +14,9 @@ Test the endpoint:
|
|||
curl "http://localhost:8080/beta/litellm_prompt_management?prompt_id=hello-world-prompt"
|
||||
"""
|
||||
|
||||
import os
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Header, Query, status
|
||||
from fastapi.responses import JSONResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# ============================================================================
|
||||
|
|
@ -364,7 +361,7 @@ if __name__ == "__main__":
|
|||
print("=" * 70)
|
||||
print("Mock Prompt Management API Server")
|
||||
print("=" * 70)
|
||||
print(f"\nStarting server on http://localhost:8080")
|
||||
print("\nStarting server on http://localhost:8080")
|
||||
print(f"\nAvailable prompts: {len(PROMPTS_DB)}")
|
||||
for prompt_id in PROMPTS_DB.keys():
|
||||
print(f" - {prompt_id}")
|
||||
|
|
|
|||
|
|
@ -157,7 +157,7 @@ class VeoVideoGenerator:
|
|||
Returns:
|
||||
True if download successful, False otherwise
|
||||
"""
|
||||
print(f"⬇️ Downloading video...")
|
||||
print("⬇️ Downloading video...")
|
||||
print(f"Original URI: {video_uri}")
|
||||
|
||||
# Convert Google URI to LiteLLM proxy URI
|
||||
|
|
@ -198,7 +198,7 @@ class VeoVideoGenerator:
|
|||
if os.path.exists(output_filename):
|
||||
file_size = os.path.getsize(output_filename)
|
||||
if file_size > 0:
|
||||
print(f"✅ Video downloaded successfully!")
|
||||
print("✅ Video downloaded successfully!")
|
||||
print(f"📁 Saved as: {output_filename}")
|
||||
print(f"📏 File size: {file_size / (1024*1024):.2f} MB")
|
||||
return True
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import csv
|
|||
import json
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import Optional, List, Dict, Any
|
||||
from typing import Any
|
||||
|
||||
import os
|
||||
|
||||
|
|
@ -172,7 +172,7 @@ async def migrate_verification_tokens():
|
|||
)
|
||||
continue
|
||||
|
||||
print(f"\nMigration completed!")
|
||||
print("\nMigration completed!")
|
||||
print(f"Successfully processed: {processed_count} records")
|
||||
print(f"Errors encountered: {error_count} records")
|
||||
|
||||
|
|
|
|||
|
|
@ -15,11 +15,10 @@ sys.path.insert(
|
|||
) # Adds the parent directory to the system path
|
||||
import json
|
||||
import sys
|
||||
from typing import Any, List, Literal, Optional
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.logging_utils import (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ sys.path.insert(
|
|||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
import sys
|
||||
from typing import Literal, Optional
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@
|
|||
# Thank you users! We ❤️ you! - Krrish & Ishaan
|
||||
## This provides an LLM Guard Integration for content moderation on the proxy
|
||||
|
||||
from typing import Literal, Optional
|
||||
from typing import Optional
|
||||
|
||||
import aiohttp
|
||||
from fastapi import HTTPException
|
||||
|
|
|
|||
|
|
@ -349,8 +349,10 @@ class BaseEmailLogger(CustomLogger):
|
|||
)
|
||||
|
||||
# Calculate percentage and alert threshold
|
||||
percentage = threshold_pct if threshold_pct is not None else int(
|
||||
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE * 100
|
||||
percentage = (
|
||||
threshold_pct
|
||||
if threshold_pct is not None
|
||||
else int(EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE * 100)
|
||||
)
|
||||
threshold_fraction = percentage / 100.0
|
||||
alert_threshold_str = (
|
||||
|
|
@ -609,9 +611,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
continue
|
||||
|
||||
_id = user_info.token or user_info.user_id or "default_id"
|
||||
_cache_key = (
|
||||
f"email_budget_alerts:max_budget_alert:{threshold_pct}:{_id}"
|
||||
)
|
||||
_cache_key = f"email_budget_alerts:max_budget_alert:{threshold_pct}:{_id}"
|
||||
|
||||
result = await _cache.async_get_cache(key=_cache_key)
|
||||
if result is not None:
|
||||
|
|
@ -630,7 +630,9 @@ class BaseEmailLogger(CustomLogger):
|
|||
continue
|
||||
recipient_emails = list(set(emails))
|
||||
|
||||
event_message = f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
|
||||
event_message = (
|
||||
f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
|
||||
)
|
||||
webhook_event = WebhookEvent(
|
||||
event="max_budget_alert",
|
||||
event_message=event_message,
|
||||
|
|
|
|||
|
|
@ -15,7 +15,6 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
|
||||
from .base_email import BaseEmailLogger
|
||||
|
||||
|
||||
SENDGRID_API_ENDPOINT = "https://api.sendgrid.com/v3/mail/send"
|
||||
|
||||
|
||||
|
|
@ -79,4 +78,4 @@ class SendGridEmailLogger(BaseEmailLogger):
|
|||
verbose_logger.debug(
|
||||
f"SendGrid response status={response.status_code}, body={response.text}"
|
||||
)
|
||||
return
|
||||
return
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ GET - /audit/{id} - Get audit log by id
|
|||
GET - /audit - Get all audit logs
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
#### AUDIT LOGGING ####
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
|
@ -153,11 +153,11 @@ async def get_audit_logs(
|
|||
|
||||
# Return paginated response
|
||||
return PaginatedAuditLogResponse(
|
||||
audit_logs=[
|
||||
AuditLogResponse(**audit_log.model_dump()) for audit_log in audit_logs
|
||||
]
|
||||
if audit_logs
|
||||
else [],
|
||||
audit_logs=(
|
||||
[AuditLogResponse(**audit_log.model_dump()) for audit_log in audit_logs]
|
||||
if audit_logs
|
||||
else []
|
||||
),
|
||||
total=total_count,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
|
|
|
|||
|
|
@ -7,4 +7,4 @@ including custom SSO handlers and advanced authentication features.
|
|||
|
||||
from .custom_sso_handler import EnterpriseCustomSSOHandler
|
||||
|
||||
__all__ = ["EnterpriseCustomSSOHandler"]
|
||||
__all__ = ["EnterpriseCustomSSOHandler"]
|
||||
|
|
|
|||
|
|
@ -53,7 +53,9 @@ class CheckBatchCost:
|
|||
"user_api_key_alias": getattr(user_row, "user_alias", None),
|
||||
}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}")
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}"
|
||||
)
|
||||
return {}
|
||||
|
||||
async def _cleanup_stale_managed_objects(self) -> None:
|
||||
|
|
@ -62,11 +64,22 @@ class CheckBatchCost:
|
|||
in non-terminal states as 'stale_expired'. These will never complete and
|
||||
should not be polled.
|
||||
"""
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(
|
||||
days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS
|
||||
)
|
||||
result = await self.prisma_client.db.litellm_managedobjecttable.update_many(
|
||||
where={
|
||||
"file_purpose": "batch",
|
||||
"status": {"not_in": ["completed", "complete", "failed", "expired", "cancelled", "stale_expired"]},
|
||||
"status": {
|
||||
"not_in": [
|
||||
"completed",
|
||||
"complete",
|
||||
"failed",
|
||||
"expired",
|
||||
"cancelled",
|
||||
"stale_expired",
|
||||
]
|
||||
},
|
||||
"created_at": {"lt": cutoff},
|
||||
},
|
||||
data={"status": "stale_expired"},
|
||||
|
|
@ -120,9 +133,12 @@ class CheckBatchCost:
|
|||
|
||||
try:
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
prom_logger = PrometheusLogger.get_instance()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"CheckBatchCost: could not get Prometheus logger: {e}")
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: could not get Prometheus logger: {e}"
|
||||
)
|
||||
prom_logger = None
|
||||
|
||||
processed_models: List[Tuple[Optional[str], Optional[str]]] = []
|
||||
|
|
@ -161,7 +177,11 @@ class CheckBatchCost:
|
|||
order={"created_at": "asc"},
|
||||
)
|
||||
except Exception as query_err:
|
||||
if "batch_processed" not in str(query_err).lower() and "unknown column" not in str(query_err).lower() and "does not exist" not in str(query_err).lower():
|
||||
if (
|
||||
"batch_processed" not in str(query_err).lower()
|
||||
and "unknown column" not in str(query_err).lower()
|
||||
and "does not exist" not in str(query_err).lower()
|
||||
):
|
||||
raise
|
||||
# Permanent schema gap — cache the result so future cycles skip straight to fallback
|
||||
self._has_batch_processed_column = False
|
||||
|
|
@ -216,14 +236,13 @@ class CheckBatchCost:
|
|||
f"Skipping job {unified_object_id} because of error querying model ID: {model_id} for cost and usage of batch ID: {batch_id}: {e}"
|
||||
)
|
||||
if prom_logger:
|
||||
prom_logger.record_check_batch_cost_error("provider_retrieval_error")
|
||||
prom_logger.record_check_batch_cost_error(
|
||||
"provider_retrieval_error"
|
||||
)
|
||||
continue
|
||||
|
||||
## RETRIEVE THE BATCH JOB OUTPUT FILE
|
||||
if (
|
||||
response.status == "completed"
|
||||
and response.output_file_id is not None
|
||||
):
|
||||
if response.status == "completed" and response.output_file_id is not None:
|
||||
verbose_proxy_logger.info(
|
||||
f"Batch ID: {batch_id} is complete, tracking cost and usage"
|
||||
)
|
||||
|
|
@ -250,20 +269,25 @@ class CheckBatchCost:
|
|||
decoded = _is_base64_encoded_unified_file_id(raw_output_file_id)
|
||||
if decoded:
|
||||
try:
|
||||
raw_output_file_id = decoded.split("llm_output_file_id,")[1].split(";")[0]
|
||||
raw_output_file_id = decoded.split("llm_output_file_id,")[
|
||||
1
|
||||
].split(";")[0]
|
||||
except (IndexError, AttributeError):
|
||||
pass
|
||||
|
||||
credentials = self.llm_router.get_deployment_credentials_with_provider(model_id) or {}
|
||||
credentials = (
|
||||
self.llm_router.get_deployment_credentials_with_provider(model_id)
|
||||
or {}
|
||||
)
|
||||
_file_content = await afile_content(
|
||||
file_id=raw_output_file_id,
|
||||
**credentials,
|
||||
)
|
||||
|
||||
# Access content - handle both direct attribute and method call
|
||||
if hasattr(_file_content, 'content'):
|
||||
if hasattr(_file_content, "content"):
|
||||
content_bytes = _file_content.content # type: ignore[union-attr]
|
||||
elif hasattr(_file_content, 'read'):
|
||||
elif hasattr(_file_content, "read"):
|
||||
content_bytes = await _file_content.read() # type: ignore[misc]
|
||||
else:
|
||||
content_bytes = _file_content # type: ignore[assignment]
|
||||
|
|
@ -290,7 +314,9 @@ class CheckBatchCost:
|
|||
f"Skipping job {unified_object_id} because it is not a valid deployment info"
|
||||
)
|
||||
if prom_logger:
|
||||
prom_logger.record_check_batch_cost_error("deployment_not_found")
|
||||
prom_logger.record_check_batch_cost_error(
|
||||
"deployment_not_found"
|
||||
)
|
||||
continue
|
||||
custom_llm_provider = deployment_info.litellm_params.custom_llm_provider
|
||||
litellm_model_name = deployment_info.litellm_params.model
|
||||
|
|
@ -302,21 +328,32 @@ class CheckBatchCost:
|
|||
|
||||
# CheckBatchCost bypasses async_post_call_success_hook, so convert raw
|
||||
# output/error file IDs to managed base64 IDs before the DB write here.
|
||||
managed_files_hook = self.proxy_logging_obj.get_proxy_hook("managed_files")
|
||||
managed_files_hook = self.proxy_logging_obj.get_proxy_hook(
|
||||
"managed_files"
|
||||
)
|
||||
if managed_files_hook is not None:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
_minimal_auth = UserAPIKeyAuth(
|
||||
user_id=job.created_by or "default-user-id",
|
||||
team_id=getattr(job, "team_id", None),
|
||||
)
|
||||
for _file_attr in ["output_file_id", "error_file_id"]:
|
||||
_raw_file_id = getattr(response, _file_attr, None)
|
||||
if _raw_file_id and not _is_base64_encoded_unified_file_id(_raw_file_id):
|
||||
if _raw_file_id and not _is_base64_encoded_unified_file_id(
|
||||
_raw_file_id
|
||||
):
|
||||
try:
|
||||
_unified_file_id = managed_files_hook.get_unified_output_file_id(
|
||||
output_file_id=_raw_file_id,
|
||||
model_id=model_id,
|
||||
model_name=str(model_name) if model_name else deployment_info.model_name or None,
|
||||
_unified_file_id = (
|
||||
managed_files_hook.get_unified_output_file_id(
|
||||
output_file_id=_raw_file_id,
|
||||
model_id=model_id,
|
||||
model_name=(
|
||||
str(model_name)
|
||||
if model_name
|
||||
else deployment_info.model_name or None
|
||||
),
|
||||
)
|
||||
)
|
||||
await managed_files_hook.store_unified_file_id(
|
||||
file_id=_unified_file_id,
|
||||
|
|
@ -338,7 +375,11 @@ class CheckBatchCost:
|
|||
|
||||
# Pass deployment model_info so custom batch pricing
|
||||
# (input_cost_per_token_batches etc.) is used for cost calc
|
||||
deployment_model_info = deployment_info.model_info.model_dump() if deployment_info.model_info else {}
|
||||
deployment_model_info = (
|
||||
deployment_info.model_info.model_dump()
|
||||
if deployment_info.model_info
|
||||
else {}
|
||||
)
|
||||
batch_cost, batch_usage, batch_models = (
|
||||
await calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=file_content_as_dict,
|
||||
|
|
@ -385,7 +426,9 @@ class CheckBatchCost:
|
|||
|
||||
# Record batch duration (completed_at - created_at)
|
||||
if prom_logger and response.completed_at and response.created_at:
|
||||
duration_seconds = float(response.completed_at - response.created_at)
|
||||
duration_seconds = float(
|
||||
response.completed_at - response.created_at
|
||||
)
|
||||
if duration_seconds >= 0:
|
||||
prom_logger.record_managed_batch_duration(
|
||||
duration_seconds=duration_seconds,
|
||||
|
|
@ -394,7 +437,9 @@ class CheckBatchCost:
|
|||
)
|
||||
|
||||
# Track this job for the final metrics summary
|
||||
processed_models.append((model_name, str(llm_provider) if llm_provider else None))
|
||||
processed_models.append(
|
||||
(model_name, str(llm_provider) if llm_provider else None)
|
||||
)
|
||||
|
||||
# mark the job as complete
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
|
|||
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm import Router, verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -41,7 +40,7 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
):
|
||||
"""
|
||||
Managed vector stores with target_model_names support.
|
||||
|
||||
|
||||
This class provides functionality to:
|
||||
- Create vector stores across multiple models
|
||||
- Retrieve vector stores by unified ID
|
||||
|
|
@ -77,14 +76,14 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> str:
|
||||
"""
|
||||
Generate the format string for the unified vector store ID.
|
||||
|
||||
|
||||
Format:
|
||||
litellm_proxy:vector_store;unified_id,<uuid>;target_model_names,<models>;resource_id,<vs_id>;model_id,<model_id>
|
||||
"""
|
||||
# VectorStoreCreateResponse is a TypedDict, so resource_object is a dictionary
|
||||
# Extract provider resource ID from the response
|
||||
provider_resource_id = resource_object.get("id", "")
|
||||
|
||||
|
||||
# Model ID is stored in hidden params if the response object supports it
|
||||
# For TypedDict responses, we need to check if _hidden_params was added
|
||||
hidden_params: Dict[str, Any] = {}
|
||||
|
|
@ -109,20 +108,18 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> VectorStoreCreateResponse:
|
||||
"""
|
||||
Create a vector store for a specific model.
|
||||
|
||||
|
||||
Args:
|
||||
llm_router: LiteLLM router instance
|
||||
model: Model name to create vector store for
|
||||
request_data: Request data for vector store creation
|
||||
litellm_parent_otel_span: OpenTelemetry span for tracing
|
||||
|
||||
|
||||
Returns:
|
||||
VectorStoreCreateResponse from the provider
|
||||
"""
|
||||
# Use the router to create the vector store
|
||||
response = await llm_router.avector_store_create(
|
||||
model=model, **request_data
|
||||
)
|
||||
response = await llm_router.avector_store_create(model=model, **request_data)
|
||||
return response
|
||||
|
||||
# ============================================================================
|
||||
|
|
@ -139,14 +136,14 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> VectorStoreCreateResponse:
|
||||
"""
|
||||
Create a vector store across multiple models.
|
||||
|
||||
|
||||
Args:
|
||||
create_request: Vector store creation request parameters
|
||||
llm_router: LiteLLM router instance
|
||||
target_model_names_list: List of target model names
|
||||
litellm_parent_otel_span: OpenTelemetry span for tracing
|
||||
user_api_key_dict: User API key authentication details
|
||||
|
||||
|
||||
Returns:
|
||||
VectorStoreCreateResponse with unified ID
|
||||
"""
|
||||
|
|
@ -196,7 +193,7 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
# VectorStoreCreateResponse is a TypedDict, so we need to create a new dict with the unified ID
|
||||
response = responses[0].copy()
|
||||
response["id"] = unified_id
|
||||
|
||||
|
||||
verbose_logger.info(
|
||||
f"Successfully created managed vector store with unified ID: {unified_id}"
|
||||
)
|
||||
|
|
@ -212,13 +209,13 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> Dict[str, Any]:
|
||||
"""
|
||||
List vector stores created by a user.
|
||||
|
||||
|
||||
Args:
|
||||
user_api_key_dict: User API key authentication details
|
||||
limit: Maximum number of vector stores to return
|
||||
after: Cursor for pagination
|
||||
order: Sort order ('asc' or 'desc')
|
||||
|
||||
|
||||
Returns:
|
||||
Dictionary with list of vector stores and pagination info
|
||||
"""
|
||||
|
|
@ -238,23 +235,23 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> bool:
|
||||
"""
|
||||
Check if user has access to a vector store.
|
||||
|
||||
|
||||
Args:
|
||||
vector_store_id: The unified vector store ID
|
||||
user_api_key_dict: User API key authentication details
|
||||
|
||||
|
||||
Returns:
|
||||
True if user has access, False otherwise
|
||||
"""
|
||||
is_unified_id = is_base64_encoded_unified_id(vector_store_id)
|
||||
|
||||
|
||||
if is_unified_id:
|
||||
# Check access for managed vector store
|
||||
return await self.can_user_access_unified_resource_id(
|
||||
vector_store_id,
|
||||
user_api_key_dict,
|
||||
)
|
||||
|
||||
|
||||
# Not a managed vector store, allow access
|
||||
return True
|
||||
|
||||
|
|
@ -263,24 +260,22 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> bool:
|
||||
"""
|
||||
Check if user has access to a managed vector store in request data.
|
||||
|
||||
|
||||
Args:
|
||||
data: Request data containing vector_store_id
|
||||
user_api_key_dict: User API key authentication details
|
||||
|
||||
|
||||
Returns:
|
||||
True if this is a managed vector store and user has access
|
||||
|
||||
|
||||
Raises:
|
||||
HTTPException: If user doesn't have access
|
||||
"""
|
||||
vector_store_id = cast(Optional[str], data.get("vector_store_id"))
|
||||
is_unified_id = (
|
||||
is_base64_encoded_unified_id(vector_store_id)
|
||||
if vector_store_id
|
||||
else False
|
||||
is_base64_encoded_unified_id(vector_store_id) if vector_store_id else False
|
||||
)
|
||||
|
||||
|
||||
if is_unified_id and vector_store_id:
|
||||
if await self.can_user_access_unified_resource_id(
|
||||
vector_store_id, user_api_key_dict
|
||||
|
|
@ -291,7 +286,7 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
status_code=403,
|
||||
detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}",
|
||||
)
|
||||
|
||||
|
||||
return False
|
||||
|
||||
# ============================================================================
|
||||
|
|
@ -307,18 +302,18 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> Union[Exception, str, Dict, None]:
|
||||
"""
|
||||
Pre-call hook to handle vector store operations.
|
||||
|
||||
|
||||
This hook intercepts vector store requests and:
|
||||
- Validates access for managed vector stores
|
||||
- Transforms unified IDs to provider-specific IDs
|
||||
- Adds model routing information
|
||||
|
||||
|
||||
Args:
|
||||
user_api_key_dict: User API key authentication details
|
||||
cache: Cache instance
|
||||
data: Request data
|
||||
call_type: Type of call being made
|
||||
|
||||
|
||||
Returns:
|
||||
Modified request data or None
|
||||
"""
|
||||
|
|
@ -330,40 +325,40 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
# Handle vector store search operations
|
||||
if call_type == "avector_store_search":
|
||||
vector_store_id = data.get("vector_store_id")
|
||||
|
||||
|
||||
if vector_store_id:
|
||||
# Check if it's a managed vector store ID
|
||||
decoded_id = is_base64_encoded_unified_id(vector_store_id)
|
||||
|
||||
|
||||
if decoded_id:
|
||||
verbose_logger.debug(
|
||||
f"Processing managed vector store search: {vector_store_id}"
|
||||
)
|
||||
|
||||
|
||||
# Check access
|
||||
has_access = await self.can_user_access_unified_resource_id(
|
||||
vector_store_id, user_api_key_dict
|
||||
)
|
||||
|
||||
|
||||
if not has_access:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}",
|
||||
)
|
||||
|
||||
|
||||
# Parse the unified ID to extract components
|
||||
parsed_id = parse_unified_id(vector_store_id)
|
||||
|
||||
|
||||
if parsed_id:
|
||||
# Extract the model ID and provider resource ID
|
||||
model_id = parsed_id.get("model_id")
|
||||
provider_resource_id = parsed_id.get("provider_resource_id")
|
||||
target_model_names = parsed_id.get("target_model_names", [])
|
||||
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Decoded vector store - model_id: {model_id}, provider_resource_id: {provider_resource_id}, target_model_names: {target_model_names}"
|
||||
)
|
||||
|
||||
|
||||
# Determine which model to use for routing
|
||||
# Priority: model_id (deployment ID) > first target_model_name
|
||||
routing_model = None
|
||||
|
|
@ -371,28 +366,28 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
routing_model = model_id
|
||||
elif target_model_names and len(target_model_names) > 0:
|
||||
routing_model = target_model_names[0]
|
||||
|
||||
|
||||
# Set the model for routing
|
||||
if routing_model:
|
||||
data["model"] = routing_model
|
||||
verbose_logger.info(
|
||||
f"Routing vector store search to model: {routing_model}"
|
||||
)
|
||||
|
||||
|
||||
# Replace the unified ID with the provider-specific ID
|
||||
if provider_resource_id:
|
||||
data["vector_store_id"] = provider_resource_id
|
||||
verbose_logger.debug(
|
||||
f"Replaced unified ID with provider resource ID: {provider_resource_id}"
|
||||
)
|
||||
|
||||
|
||||
# Handle vector store retrieve/delete operations
|
||||
elif call_type in ("avector_store_retrieve", "avector_store_delete"):
|
||||
await self.check_managed_vector_store_access(data, user_api_key_dict)
|
||||
|
||||
|
||||
# If it's a managed vector store, we'll handle it in the endpoint
|
||||
# No need to transform here as the endpoint will route to the hook
|
||||
|
||||
|
||||
return data
|
||||
|
||||
# ============================================================================
|
||||
|
|
@ -407,15 +402,15 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> Any:
|
||||
"""
|
||||
Post-call hook to transform responses.
|
||||
|
||||
|
||||
This hook can be used to transform responses if needed.
|
||||
For now, it just passes through the response.
|
||||
|
||||
|
||||
Args:
|
||||
data: Request data
|
||||
user_api_key_dict: User API key authentication details
|
||||
response: Response from the provider
|
||||
|
||||
|
||||
Returns:
|
||||
Potentially modified response
|
||||
"""
|
||||
|
|
@ -436,21 +431,21 @@ class _PROXY_LiteLLMManagedVectorStores(
|
|||
) -> List[Dict]:
|
||||
"""
|
||||
Filter deployments based on vector store availability.
|
||||
|
||||
|
||||
This is used by the router to select only deployments that have
|
||||
the vector store available.
|
||||
|
||||
|
||||
Note: This method signature is a compromise between CustomLogger and BaseManagedResource
|
||||
parent classes which have incompatible signatures. The type: ignore[override] is necessary
|
||||
due to this multiple inheritance conflict.
|
||||
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
healthy_deployments: List of healthy deployments
|
||||
messages: Messages (unused for vector stores, required by CustomLogger interface)
|
||||
request_kwargs: Request kwargs containing vector_store_id and mappings
|
||||
parent_otel_span: OpenTelemetry span for tracing
|
||||
|
||||
|
||||
Returns:
|
||||
Filtered list of deployments
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
Enterprise internal user management endpoints
|
||||
"""
|
||||
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
|
|||
|
|
@ -8,7 +8,6 @@ All /vector_store management endpoints
|
|||
/vector_store/list
|
||||
"""
|
||||
|
||||
import copy
|
||||
import json
|
||||
from typing import List, Optional
|
||||
|
||||
|
|
@ -147,12 +146,12 @@ async def list_vector_stores(
|
|||
vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
|
||||
|
||||
# Also clean up in-memory registry to remove any deleted vector stores
|
||||
if litellm.vector_store_registry is not None:
|
||||
db_vector_store_ids = {
|
||||
vs.get("vector_store_id")
|
||||
for vs in vector_stores_from_db
|
||||
vs.get("vector_store_id")
|
||||
for vs in vector_stores_from_db
|
||||
if vs.get("vector_store_id")
|
||||
}
|
||||
# Remove any in-memory vector stores that no longer exist in database
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -106,7 +106,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
# Health & ops
|
||||
"/health",
|
||||
"/metrics",
|
||||
"/watsonx"
|
||||
"/watsonx",
|
||||
)
|
||||
|
||||
GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(
|
||||
|
|
|
|||
|
|
@ -266,6 +266,9 @@ class MistralConfig(OpenAIGPTConfig):
|
|||
for m in messages:
|
||||
m = MistralConfig._handle_name_in_message(m)
|
||||
m = MistralConfig._handle_tool_call_message(m)
|
||||
|
||||
m.pop("metadata", None)
|
||||
|
||||
if MistralConfig._is_empty_assistant_message(m):
|
||||
continue
|
||||
m = strip_none_values_from_message(m) # prevents 'extra_forbidden' error
|
||||
|
|
|
|||
|
|
@ -35,7 +35,7 @@ import httpx
|
|||
class EvalCase:
|
||||
category: str
|
||||
prompt: str
|
||||
ideal: str # criteria the judge checks the response against
|
||||
ideal: str # criteria the judge checks the response against
|
||||
|
||||
|
||||
EVAL_CASES: List[EvalCase] = [
|
||||
|
|
@ -177,14 +177,19 @@ async def evaluate(
|
|||
async with httpx.AsyncClient() as client:
|
||||
for i, case in enumerate(EVAL_CASES, 1):
|
||||
print(f"\n[{i}/{len(EVAL_CASES)}] category={case.category}")
|
||||
print(f" prompt : {case.prompt[:80]}{'…' if len(case.prompt) > 80 else ''}")
|
||||
print(
|
||||
f" prompt : {case.prompt[:80]}{'…' if len(case.prompt) > 80 else ''}"
|
||||
)
|
||||
|
||||
session_id = f"eval-{uuid.uuid4()}"
|
||||
|
||||
# Round 1: single-turn real request — get the actual LLM response to judge.
|
||||
try:
|
||||
response, chosen = await _chat(
|
||||
client, proxy_url, api_key, router,
|
||||
client,
|
||||
proxy_url,
|
||||
api_key,
|
||||
router,
|
||||
[{"role": "user", "content": case.prompt}],
|
||||
session_id=session_id,
|
||||
)
|
||||
|
|
@ -194,16 +199,25 @@ async def evaluate(
|
|||
continue
|
||||
|
||||
print(f" model : {chosen or router}")
|
||||
print(f" response : {response[:120].replace(chr(10), ' ')}{'…' if len(response) > 120 else ''}")
|
||||
print(
|
||||
f" response : {response[:120].replace(chr(10), ' ')}{'…' if len(response) > 120 else ''}"
|
||||
)
|
||||
|
||||
# Judge the real response.
|
||||
judge_msgs = [
|
||||
{"role": "system", "content": JUDGE_SYSTEM},
|
||||
{"role": "user", "content": _judge_user(case.prompt, case.ideal, response)},
|
||||
{
|
||||
"role": "user",
|
||||
"content": _judge_user(case.prompt, case.ideal, response),
|
||||
},
|
||||
]
|
||||
try:
|
||||
verdict, _ = await _chat(
|
||||
client, proxy_url, api_key, judge_model, judge_msgs,
|
||||
client,
|
||||
proxy_url,
|
||||
api_key,
|
||||
judge_model,
|
||||
judge_msgs,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
print(f" ERROR calling judge: {exc}", file=sys.stderr)
|
||||
|
|
@ -227,15 +241,19 @@ async def evaluate(
|
|||
# On PASS → satisfaction follow-up (+alpha). On FAIL → neutral (no signal).
|
||||
follow_up = SATISFY_FOLLOWUP if is_pass else NEUTRAL_FOLLOWUP
|
||||
bandit_msgs = [
|
||||
{"role": "user", "content": case.prompt},
|
||||
{"role": "user", "content": case.prompt},
|
||||
{"role": "assistant", "content": response},
|
||||
{"role": "user", "content": "ok continue"},
|
||||
{"role": "user", "content": "ok continue"},
|
||||
{"role": "assistant", "content": FAB_ASSISTANT},
|
||||
{"role": "user", "content": follow_up},
|
||||
{"role": "user", "content": follow_up},
|
||||
]
|
||||
try:
|
||||
await _chat(
|
||||
client, proxy_url, api_key, router, bandit_msgs,
|
||||
client,
|
||||
proxy_url,
|
||||
api_key,
|
||||
router,
|
||||
bandit_msgs,
|
||||
session_id=session_id,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
|
|
@ -257,11 +275,17 @@ async def evaluate(
|
|||
# Entry point
|
||||
# ---------------------------------------------------------------------------
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description="Evaluate the adaptive router with LLM-as-judge.")
|
||||
ap.add_argument("--proxy-url", default="http://localhost:4000")
|
||||
ap.add_argument("--api-key", required=True, help="proxy API key")
|
||||
ap.add_argument("--router", default="smart-cheap-router", help="adaptive router model name")
|
||||
ap.add_argument("--judge-model", default="smart", help="model name for the judge (via proxy)")
|
||||
ap = argparse.ArgumentParser(
|
||||
description="Evaluate the adaptive router with LLM-as-judge."
|
||||
)
|
||||
ap.add_argument("--proxy-url", default="http://localhost:4000")
|
||||
ap.add_argument("--api-key", required=True, help="proxy API key")
|
||||
ap.add_argument(
|
||||
"--router", default="smart-cheap-router", help="adaptive router model name"
|
||||
)
|
||||
ap.add_argument(
|
||||
"--judge-model", default="smart", help="model name for the judge (via proxy)"
|
||||
)
|
||||
args = ap.parse_args()
|
||||
|
||||
asyncio.run(evaluate(args.proxy_url, args.api_key, args.router, args.judge_model))
|
||||
|
|
|
|||
|
|
@ -72,8 +72,8 @@ PROMPTS: Dict[str, List[str]] = {
|
|||
# so that signals attribute to the right (type, model) bandit cell.
|
||||
SATISFY: Dict[str, str] = {
|
||||
"code_generation": "thanks, that works! now write me a python function that does the inverse",
|
||||
"factual_lookup": "perfect, thanks! who is the current prime minister?",
|
||||
"writing": "great, thanks! now write a follow-up email confirming attendance",
|
||||
"factual_lookup": "perfect, thanks! who is the current prime minister?",
|
||||
"writing": "great, thanks! now write a follow-up email confirming attendance",
|
||||
}
|
||||
|
||||
# Neutral follow-up — does not match any signal regex, does not move the bandit.
|
||||
|
|
@ -83,8 +83,8 @@ NEUTRAL_FOLLOWUP = "ok, noted"
|
|||
# Defaults: smart dominates code/writing; both are fine for factual_lookup.
|
||||
ORACLE: Dict[str, Dict[str, float]] = {
|
||||
"code_generation": {"smart": 0.92, "fast": 0.35},
|
||||
"factual_lookup": {"smart": 0.90, "fast": 0.85},
|
||||
"writing": {"smart": 0.85, "fast": 0.55},
|
||||
"factual_lookup": {"smart": 0.90, "fast": 0.85},
|
||||
"writing": {"smart": 0.85, "fast": 0.55},
|
||||
}
|
||||
|
||||
# Fabricated assistant turn — content doesn't matter for the hook, only the role.
|
||||
|
|
@ -94,11 +94,11 @@ FAB_ASSISTANT = "Got it. Working on that now."
|
|||
def _build_messages(prompt: str, last_user: str) -> List[Dict[str, str]]:
|
||||
"""5-message conversation that passes the SIGNAL_GATE_MIN_MESSAGES=4 gate."""
|
||||
return [
|
||||
{"role": "user", "content": prompt},
|
||||
{"role": "user", "content": prompt},
|
||||
{"role": "assistant", "content": FAB_ASSISTANT},
|
||||
{"role": "user", "content": "ok continue"},
|
||||
{"role": "user", "content": "ok continue"},
|
||||
{"role": "assistant", "content": FAB_ASSISTANT},
|
||||
{"role": "user", "content": last_user},
|
||||
{"role": "user", "content": last_user},
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -155,7 +155,11 @@ async def _drive_one_session(
|
|||
#
|
||||
# Round 1: neutral follow-up → no signal fires, but we learn the pick.
|
||||
ok, chosen = await _send(
|
||||
client, proxy_url, api_key, router, session_id,
|
||||
client,
|
||||
proxy_url,
|
||||
api_key,
|
||||
router,
|
||||
session_id,
|
||||
_build_messages(prompt, NEUTRAL_FOLLOWUP),
|
||||
mock_response=FAB_ASSISTANT,
|
||||
)
|
||||
|
|
@ -171,10 +175,15 @@ async def _drive_one_session(
|
|||
# follow-up matches satisfaction → +alpha for (request_type, chosen).
|
||||
history = _build_messages(prompt, NEUTRAL_FOLLOWUP) + [
|
||||
{"role": "assistant", "content": FAB_ASSISTANT},
|
||||
{"role": "user", "content": follow_up},
|
||||
{"role": "user", "content": follow_up},
|
||||
]
|
||||
await _send(
|
||||
client, proxy_url, api_key, router, session_id, history,
|
||||
client,
|
||||
proxy_url,
|
||||
api_key,
|
||||
router,
|
||||
session_id,
|
||||
history,
|
||||
mock_response=FAB_ASSISTANT,
|
||||
)
|
||||
return chosen
|
||||
|
|
@ -183,13 +192,22 @@ async def _drive_one_session(
|
|||
async def main() -> None:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--proxy-url", default="http://localhost:4000")
|
||||
ap.add_argument("--api-key", required=True, help="proxy key with /v1/chat/completions perms")
|
||||
ap.add_argument("--router", default="smart-cheap-router")
|
||||
ap.add_argument("--rounds", type=int, default=100)
|
||||
ap.add_argument("--rate", type=float, default=0.5,
|
||||
help="seconds between sessions; lower = faster")
|
||||
ap.add_argument("--types", default="code_generation,factual_lookup,writing",
|
||||
help="comma-separated subset of request types to drive")
|
||||
ap.add_argument(
|
||||
"--api-key", required=True, help="proxy key with /v1/chat/completions perms"
|
||||
)
|
||||
ap.add_argument("--router", default="smart-cheap-router")
|
||||
ap.add_argument("--rounds", type=int, default=100)
|
||||
ap.add_argument(
|
||||
"--rate",
|
||||
type=float,
|
||||
default=0.5,
|
||||
help="seconds between sessions; lower = faster",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--types",
|
||||
default="code_generation,factual_lookup,writing",
|
||||
help="comma-separated subset of request types to drive",
|
||||
)
|
||||
args = ap.parse_args()
|
||||
|
||||
types = [t.strip() for t in args.types.split(",") if t.strip() in PROMPTS]
|
||||
|
|
@ -207,7 +225,12 @@ async def main() -> None:
|
|||
rt = random.choice(types)
|
||||
prompt = random.choice(PROMPTS[rt])
|
||||
chosen = await _drive_one_session(
|
||||
client, args.proxy_url, args.api_key, args.router, rt, prompt,
|
||||
client,
|
||||
args.proxy_url,
|
||||
args.api_key,
|
||||
args.router,
|
||||
rt,
|
||||
prompt,
|
||||
)
|
||||
if chosen:
|
||||
counts[(rt, chosen)] = counts.get((rt, chosen), 0) + 1
|
||||
|
|
|
|||
|
|
@ -24,7 +24,6 @@ import shlex
|
|||
import signal
|
||||
import statistics
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -34,7 +33,6 @@ from typing import Any, Optional
|
|||
import aiohttp
|
||||
from aiohttp import web
|
||||
|
||||
|
||||
DEFAULT_MODEL = "perf-test-model"
|
||||
DEFAULT_API_KEY = "sk-1234"
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,6 @@ import statistics
|
|||
|
||||
import aiohttp
|
||||
|
||||
|
||||
REQUEST_BODY = {
|
||||
"model": "db-openai-endpoint",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
|
|
@ -150,7 +149,7 @@ def print_aggregate(results):
|
|||
cov_throughput = (
|
||||
statistics.stdev(run_throughputs) / statistics.mean(run_throughputs) * 100
|
||||
)
|
||||
print(f"\n Run-to-run variance:")
|
||||
print("\n Run-to-run variance:")
|
||||
print(f" Latency CoV: {cov_latency:.1f}%")
|
||||
print(f" Throughput CoV: {cov_throughput:.1f}%")
|
||||
|
||||
|
|
|
|||
|
|
@ -132,7 +132,7 @@ def run_scenario(
|
|||
else:
|
||||
runner = lambda: bench_with_chunk(
|
||||
wrapper, spec["chunk_factory"], iterations
|
||||
) # noqa: E731
|
||||
)
|
||||
|
||||
for _ in range(warmup):
|
||||
runner()
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ USAGE EXAMPLES:
|
|||
export PROVIDER_URL='https://api.openai.com/v1/chat/completions'
|
||||
export LITELLM_PROXY_API_KEY='sk-1234'
|
||||
export PROVIDER_API_KEY='sk-openai-key'
|
||||
|
||||
|
||||
# Run from scripts directory
|
||||
cd scripts
|
||||
python benchmark_proxy_vs_provider.py
|
||||
|
|
@ -266,7 +266,7 @@ async def benchmark_endpoint(
|
|||
print(f"\nStarting benchmark for {url}")
|
||||
|
||||
if warmup:
|
||||
print(f" Warming up with 5 requests...")
|
||||
print(" Warming up with 5 requests...")
|
||||
await warmup_endpoint(
|
||||
url, headers, payload, num_warmup=5, timeout_seconds=timeout_seconds
|
||||
)
|
||||
|
|
@ -359,7 +359,7 @@ def print_results(name: str, results: BenchmarkResults):
|
|||
|
||||
if "latency_stats" in stats:
|
||||
latency = stats["latency_stats"]
|
||||
print(f"\nLatency Statistics (seconds):")
|
||||
print("\nLatency Statistics (seconds):")
|
||||
print(f" Mean: {latency['mean']:.4f}s")
|
||||
print(f" Median (p50): {latency['median']:.4f}s")
|
||||
print(f" Min: {latency['min']:.4f}s")
|
||||
|
|
@ -369,12 +369,12 @@ def print_results(name: str, results: BenchmarkResults):
|
|||
print(f" p99: {latency['p99']:.4f}s")
|
||||
|
||||
if stats["status_codes"]:
|
||||
print(f"\nStatus Codes:")
|
||||
print("\nStatus Codes:")
|
||||
for code, count in sorted(stats["status_codes"].items()):
|
||||
print(f" {code}: {count}")
|
||||
|
||||
if results.errors:
|
||||
print(f"\nErrors (showing first 5 unique):")
|
||||
print("\nErrors (showing first 5 unique):")
|
||||
unique_errors = list(set(results.errors))[:5]
|
||||
for error in unique_errors:
|
||||
count = results.errors.count(error)
|
||||
|
|
@ -439,7 +439,7 @@ def print_run_variance(name: str, results_list: List[BenchmarkResults]):
|
|||
throughputs.append(stats["requests_per_second"])
|
||||
|
||||
if mean_latencies:
|
||||
print(f"\nMean Latency Variance:")
|
||||
print("\nMean Latency Variance:")
|
||||
print(f" Runs: {len(mean_latencies)}")
|
||||
print(f" Mean: {mean(mean_latencies):.4f}s")
|
||||
print(f" Min: {min(mean_latencies):.4f}s")
|
||||
|
|
@ -456,7 +456,7 @@ def print_run_variance(name: str, results_list: List[BenchmarkResults]):
|
|||
)
|
||||
|
||||
if throughputs:
|
||||
print(f"\nThroughput Variance:")
|
||||
print("\nThroughput Variance:")
|
||||
print(f" Mean: {mean(throughputs):.2f} req/s")
|
||||
print(f" Min: {min(throughputs):.2f} req/s")
|
||||
print(f" Max: {max(throughputs):.2f} req/s")
|
||||
|
|
@ -475,18 +475,18 @@ def compare_results(
|
|||
provider_stats = provider_results.calculate_stats()
|
||||
|
||||
print(f"\n{'='*60}")
|
||||
print(f"Comparison: LiteLLM Proxy vs Direct Provider")
|
||||
print("Comparison: LiteLLM Proxy vs Direct Provider")
|
||||
print(f"{'='*60}")
|
||||
|
||||
# Success Rate Comparison
|
||||
print(f"\nSuccess Rate:")
|
||||
print("\nSuccess Rate:")
|
||||
print(f" Proxy: {proxy_stats['success_rate']:.2f}%")
|
||||
print(f" Provider: {provider_stats['success_rate']:.2f}%")
|
||||
diff = proxy_stats["success_rate"] - provider_stats["success_rate"]
|
||||
print(f" Difference: {diff:+.2f}%")
|
||||
|
||||
# Throughput Comparison
|
||||
print(f"\nThroughput (requests/second):")
|
||||
print("\nThroughput (requests/second):")
|
||||
print(f" Proxy: {proxy_stats['requests_per_second']:.2f}")
|
||||
print(f" Provider: {provider_stats['requests_per_second']:.2f}")
|
||||
diff = proxy_stats["requests_per_second"] - provider_stats["requests_per_second"]
|
||||
|
|
@ -494,7 +494,7 @@ def compare_results(
|
|||
|
||||
# Latency Comparison
|
||||
if "latency_stats" in proxy_stats and "latency_stats" in provider_stats:
|
||||
print(f"\nLatency Comparison (seconds):")
|
||||
print("\nLatency Comparison (seconds):")
|
||||
proxy_latency = proxy_stats["latency_stats"]
|
||||
provider_latency = provider_stats["latency_stats"]
|
||||
|
||||
|
|
@ -509,7 +509,7 @@ def compare_results(
|
|||
)
|
||||
|
||||
# Total Time Comparison
|
||||
print(f"\nTotal Time:")
|
||||
print("\nTotal Time:")
|
||||
print(f" Proxy: {proxy_stats['total_time']:.2f}s")
|
||||
print(f" Provider: {provider_stats['total_time']:.2f}s")
|
||||
diff = proxy_stats["total_time"] - provider_stats["total_time"]
|
||||
|
|
@ -662,7 +662,7 @@ Examples:
|
|||
print("=" * 60)
|
||||
print("LiteLLM Proxy vs Provider Benchmark")
|
||||
print("=" * 60)
|
||||
print(f"Configuration (from environment variables):")
|
||||
print("Configuration (from environment variables):")
|
||||
print(f" Proxy URL: {LITELLM_PROXY_URL}")
|
||||
print(f" Provider URL: {PROVIDER_URL}")
|
||||
print(
|
||||
|
|
@ -685,15 +685,15 @@ Examples:
|
|||
)
|
||||
|
||||
if not args.max_concurrent:
|
||||
print(f"\nTip: Use --max-concurrent 100 for more realistic load testing")
|
||||
print(f" (prevents overwhelming the server with all requests at once)")
|
||||
print("\nTip: Use --max-concurrent 100 for more realistic load testing")
|
||||
print(" (prevents overwhelming the server with all requests at once)")
|
||||
|
||||
if args.parallel:
|
||||
print(f"\nWARNING: Running benchmarks in parallel may affect results due to:")
|
||||
print(f" - Shared network bandwidth")
|
||||
print(f" - Provider endpoint receiving double load (via proxy + direct)")
|
||||
print(f" - Potential rate limiting issues")
|
||||
print(f" - Resource contention")
|
||||
print("\nWARNING: Running benchmarks in parallel may affect results due to:")
|
||||
print(" - Shared network bandwidth")
|
||||
print(" - Provider endpoint receiving double load (via proxy + direct)")
|
||||
print(" - Potential rate limiting issues")
|
||||
print(" - Resource contention")
|
||||
|
||||
# Run benchmarks multiple times if requested
|
||||
all_proxy_results = []
|
||||
|
|
@ -703,7 +703,7 @@ Examples:
|
|||
|
||||
if args.runs > 1:
|
||||
print(f"\nRunning {args.runs} benchmark runs for statistical accuracy...")
|
||||
print(f" Results will be averaged across all runs.\n")
|
||||
print(" Results will be averaged across all runs.\n")
|
||||
|
||||
overall_start_time = time.perf_counter()
|
||||
|
||||
|
|
@ -718,7 +718,7 @@ Examples:
|
|||
print(f"{'='*60}")
|
||||
|
||||
if args.parallel:
|
||||
print(f"\nRunning both benchmarks in parallel...")
|
||||
print("\nRunning both benchmarks in parallel...")
|
||||
proxy_results, provider_results = await asyncio.gather(
|
||||
benchmark_endpoint(
|
||||
LITELLM_PROXY_URL,
|
||||
|
|
@ -740,9 +740,9 @@ Examples:
|
|||
),
|
||||
)
|
||||
else:
|
||||
print(f"\nRunning benchmarks sequentially (proxy first, then provider)...")
|
||||
print("\nRunning benchmarks sequentially (proxy first, then provider)...")
|
||||
if run_num == 1:
|
||||
print(f" This ensures accurate results without interference.\n")
|
||||
print(" This ensures accurate results without interference.\n")
|
||||
|
||||
proxy_results = await benchmark_endpoint(
|
||||
LITELLM_PROXY_URL,
|
||||
|
|
@ -755,7 +755,7 @@ Examples:
|
|||
)
|
||||
|
||||
if run_num < args.runs or args.runs == 1:
|
||||
print(f"\nWaiting 3 seconds before starting provider benchmark...")
|
||||
print("\nWaiting 3 seconds before starting provider benchmark...")
|
||||
await asyncio.sleep(3) # Longer pause to ensure clean separation
|
||||
|
||||
provider_results = await benchmark_endpoint(
|
||||
|
|
@ -773,7 +773,7 @@ Examples:
|
|||
|
||||
# Brief pause between runs
|
||||
if run_num < args.runs:
|
||||
print(f"\nWaiting 5 seconds before next run...")
|
||||
print("\nWaiting 5 seconds before next run...")
|
||||
await asyncio.sleep(5)
|
||||
|
||||
overall_benchmark_time = time.perf_counter() - overall_start_time
|
||||
|
|
@ -790,7 +790,7 @@ Examples:
|
|||
raise RuntimeError("Benchmark results not initialized")
|
||||
final_proxy_results = proxy_results
|
||||
final_provider_results = provider_results
|
||||
print(f"\nResults:")
|
||||
print("\nResults:")
|
||||
|
||||
# Print individual results
|
||||
print_results("LiteLLM Proxy", final_proxy_results)
|
||||
|
|
|
|||
|
|
@ -223,6 +223,7 @@ async def drive_async(
|
|||
# Repeat × take-min runner
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class Result:
|
||||
label: str
|
||||
|
|
|
|||
|
|
@ -66,7 +66,12 @@ def _load_head(rel: str) -> dict | None:
|
|||
|
||||
|
||||
def _ref_is_commit(ref: str) -> bool:
|
||||
return _run(["git", "rev-parse", "--verify", "--quiet", f"{ref}^{{commit}}"]).returncode == 0
|
||||
return (
|
||||
_run(
|
||||
["git", "rev-parse", "--verify", "--quiet", f"{ref}^{{commit}}"]
|
||||
).returncode
|
||||
== 0
|
||||
)
|
||||
|
||||
|
||||
def _load_base(rel: str, ref: str) -> dict | None:
|
||||
|
|
@ -102,9 +107,11 @@ def regressions_for(rel: str, base: dict | None, head: dict | None) -> list[Regr
|
|||
Regression(
|
||||
rel,
|
||||
rule,
|
||||
f"rule dropped (ceiling {base_cap} -> removed)"
|
||||
if rule not in head_caps
|
||||
else f"ceiling raised {base_cap} -> {head_caps[rule]}",
|
||||
(
|
||||
f"rule dropped (ceiling {base_cap} -> removed)"
|
||||
if rule not in head_caps
|
||||
else f"ceiling raised {base_cap} -> {head_caps[rule]}"
|
||||
),
|
||||
)
|
||||
for rule, base_cap in sorted(base_caps.items())
|
||||
if rule not in head_caps or head_caps[rule] > base_cap
|
||||
|
|
@ -141,7 +148,9 @@ def main() -> int:
|
|||
regressions.extend(regressions_for(rel, base, head))
|
||||
|
||||
if regressions:
|
||||
print(f"FAIL: budget ceiling(s) loosened vs base {args.base} (merge-base {ref[:12]}):")
|
||||
print(
|
||||
f"FAIL: budget ceiling(s) loosened vs base {args.base} (merge-base {ref[:12]}):"
|
||||
)
|
||||
for reg in regressions:
|
||||
print(f" {reg.budget} {reg.rule}: {reg.detail}")
|
||||
print(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Type-discipline checker: the rules ruff can't enforce.
|
||||
|
||||
|
||||
Rules
|
||||
-----
|
||||
LIT001 Mutable collection in a type annotation, anywhere it appears: function
|
||||
|
|
@ -41,16 +41,16 @@ LIT008 `**kwargs` parameter. The keyword contract is erased and everything it c
|
|||
|
||||
LIT000 Setup failure: a target file could not be read, or contains a syntax error.
|
||||
Reported as a violation rather than crashing the run.
|
||||
|
||||
|
||||
Usage
|
||||
-----
|
||||
python check_type_discipline.py litellm/ tests/
|
||||
|
||||
Exit code 1 if any violation is found. Stdlib only.
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
import ast
|
||||
import io
|
||||
import re
|
||||
|
|
@ -60,27 +60,48 @@ from dataclasses import dataclass
|
|||
from pathlib import Path
|
||||
from collections.abc import Iterable, Iterator, Sequence
|
||||
from typing import NamedTuple
|
||||
|
||||
|
||||
# Mutable collection types, banned in *every* annotation. Name-based, so `dict`,
|
||||
# `typing.Dict`, `collections.deque`, and `collections.abc.MutableMapping` all match
|
||||
# however they were imported. The read-only interfaces (Mapping, Sequence, the
|
||||
# immutable AbstractSet / `abc.Set`, Collection) and the immutable concretes (tuple,
|
||||
# frozenset) are the escape hatch and are deliberately absent -- as is the bare name
|
||||
# `Set`, which collides with the read-only `collections.abc.Set`.
|
||||
MUTABLE_COLLECTIONS = frozenset((
|
||||
"dict", "list", "set",
|
||||
"Dict", "List", "DefaultDict", "OrderedDict", "Counter", "Deque", "ChainMap",
|
||||
"deque", "defaultdict",
|
||||
"MutableMapping", "MutableSequence", "MutableSet",
|
||||
))
|
||||
MUTABLE_COLLECTIONS = frozenset(
|
||||
(
|
||||
"dict",
|
||||
"list",
|
||||
"set",
|
||||
"Dict",
|
||||
"List",
|
||||
"DefaultDict",
|
||||
"OrderedDict",
|
||||
"Counter",
|
||||
"Deque",
|
||||
"ChainMap",
|
||||
"deque",
|
||||
"defaultdict",
|
||||
"MutableMapping",
|
||||
"MutableSequence",
|
||||
"MutableSet",
|
||||
)
|
||||
)
|
||||
|
||||
# Callables whose result is a fresh *mutable* collection (LIT002). `tuple` and
|
||||
# `frozenset` are deliberately absent -- they are the wrappers you reach for, and
|
||||
# a generator expression fed to them is the blessed one-shot build.
|
||||
MUTABLE_CONSTRUCTORS = frozenset((
|
||||
"dict", "list", "set",
|
||||
"deque", "defaultdict", "OrderedDict", "Counter", "ChainMap",
|
||||
))
|
||||
MUTABLE_CONSTRUCTORS = frozenset(
|
||||
(
|
||||
"dict",
|
||||
"list",
|
||||
"set",
|
||||
"deque",
|
||||
"defaultdict",
|
||||
"OrderedDict",
|
||||
"Counter",
|
||||
"ChainMap",
|
||||
)
|
||||
)
|
||||
# A *qualified* call (`x.deque()`) counts as construction only for names that are rarely
|
||||
# method names; `dict`/`list`/`set` are dropped here because `.dict()` / `.set()` / `.list()`
|
||||
# are common methods (e.g. pydantic's `model.dict()`), not collection construction. A
|
||||
|
|
@ -88,7 +109,7 @@ MUTABLE_CONSTRUCTORS = frozenset((
|
|||
QUALIFIED_CONSTRUCTORS = MUTABLE_CONSTRUCTORS - frozenset(("dict", "list", "set"))
|
||||
UNSAFE_GUARDS = frozenset(("TypeGuard", "TypeIs"))
|
||||
MIN_REASON_LEN = 3
|
||||
|
||||
|
||||
NOQA_RE = re.compile(
|
||||
r"#\s*noqa"
|
||||
r"(?P<colon>:\s*(?P<codes>[A-Z]+[0-9]+(?:\s*,\s*[A-Z]+[0-9]+)*))?"
|
||||
|
|
@ -110,18 +131,18 @@ OK_SUPPRESSIONS: tuple[tuple[str, re.Pattern[str]], ...] = (
|
|||
("guard-ok", GUARD_OK_RE),
|
||||
("kwargs-ok", KWARGS_OK_RE),
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
class Violation(NamedTuple):
|
||||
path: Path
|
||||
line: int
|
||||
code: str
|
||||
message: str
|
||||
|
||||
|
||||
def render(self) -> str:
|
||||
return f"{self.path}:{self.line}: {self.code} {self.message}"
|
||||
|
||||
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Comments:
|
||||
"""The lines carrying each valid `*-ok` suppression."""
|
||||
|
|
@ -130,17 +151,17 @@ class Comments:
|
|||
cast_ok_lines: frozenset[int]
|
||||
guard_ok_lines: frozenset[int]
|
||||
kwargs_ok_lines: frozenset[int]
|
||||
|
||||
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Comment scanning (LIT003 / LIT004 / LIT005)
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
|
||||
|
||||
def _reason_of(rest: str) -> str:
|
||||
return rest.strip().lstrip("#-").strip()
|
||||
|
||||
|
||||
|
||||
def _valid_ok(regex: re.Pattern[str], text: str) -> bool:
|
||||
"""True iff `text` carries this suppression with a reason of usable length."""
|
||||
m = regex.search(text)
|
||||
|
|
@ -152,30 +173,55 @@ def _comment_violations(path: Path, line_no: int, text: str) -> Iterator[Violati
|
|||
for token, regex in OK_SUPPRESSIONS:
|
||||
m = regex.search(text)
|
||||
if m and len((m.group("reason") or "").strip()) < MIN_REASON_LEN:
|
||||
yield Violation(path, line_no, "LIT005", f"{token} requires a reason: `# {token}: <reason>`")
|
||||
|
||||
yield Violation(
|
||||
path,
|
||||
line_no,
|
||||
"LIT005",
|
||||
f"{token} requires a reason: `# {token}: <reason>`",
|
||||
)
|
||||
|
||||
m = NOQA_RE.search(text)
|
||||
if m:
|
||||
if not m.group("codes"):
|
||||
yield Violation(path, line_no, "LIT003", "noqa requires rule codes: `# noqa: XXX123 # <reason>`")
|
||||
yield Violation(
|
||||
path,
|
||||
line_no,
|
||||
"LIT003",
|
||||
"noqa requires rule codes: `# noqa: XXX123 # <reason>`",
|
||||
)
|
||||
elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN:
|
||||
yield Violation(path, line_no, "LIT003", "noqa requires a reason: `# noqa: XXX123 # <reason>`")
|
||||
|
||||
yield Violation(
|
||||
path,
|
||||
line_no,
|
||||
"LIT003",
|
||||
"noqa requires a reason: `# noqa: XXX123 # <reason>`",
|
||||
)
|
||||
|
||||
m = IGNORE_RE.search(text)
|
||||
if m:
|
||||
codes = m.group("codes")
|
||||
if not codes or codes == "[]":
|
||||
yield Violation(path, line_no, "LIT004",
|
||||
"ignore requires codes: `# pyright: ignore[ruleName] # <reason>`")
|
||||
yield Violation(
|
||||
path,
|
||||
line_no,
|
||||
"LIT004",
|
||||
"ignore requires codes: `# pyright: ignore[ruleName] # <reason>`",
|
||||
)
|
||||
elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN:
|
||||
yield Violation(path, line_no, "LIT004",
|
||||
"ignore requires a reason: `# pyright: ignore[ruleName] # <reason>`")
|
||||
|
||||
|
||||
yield Violation(
|
||||
path,
|
||||
line_no,
|
||||
"LIT004",
|
||||
"ignore requires a reason: `# pyright: ignore[ruleName] # <reason>`",
|
||||
)
|
||||
|
||||
|
||||
def scan_comments(path: Path, source: str) -> tuple[Comments, tuple[Violation, ...]]:
|
||||
try:
|
||||
tokens = tokenize.generate_tokens(io.StringIO(source).readline)
|
||||
comment_toks = tuple((t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT)
|
||||
comment_toks = tuple(
|
||||
(t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT
|
||||
)
|
||||
except (tokenize.TokenError, SyntaxError):
|
||||
# tokenize raises TokenError (EOF mid-construct) or a SyntaxError subclass
|
||||
# (IndentationError / TabError) on malformed source; defer to ast.parse below,
|
||||
|
|
@ -192,13 +238,17 @@ def scan_comments(path: Path, source: str) -> tuple[Comments, tuple[Violation, .
|
|||
guard_ok_lines=_lines_with(GUARD_OK_RE),
|
||||
kwargs_ok_lines=_lines_with(KWARGS_OK_RE),
|
||||
),
|
||||
tuple(v for line, text in comment_toks for v in _comment_violations(path, line, text)),
|
||||
tuple(
|
||||
v
|
||||
for line, text in comment_toks
|
||||
for v in _comment_violations(path, line, text)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
|
||||
|
||||
def mutable_names_in(annotation: ast.expr) -> Iterator[str]:
|
||||
"""Yield mutable-collection names anywhere inside an annotation expression.
|
||||
|
||||
|
|
@ -219,11 +269,13 @@ def mutable_names_in(annotation: ast.expr) -> Iterator[str]:
|
|||
except SyntaxError:
|
||||
continue
|
||||
yield from mutable_names_in(inner)
|
||||
|
||||
|
||||
|
||||
|
||||
def _mutable_ann(path: Path, line: int, name: str, where: str) -> Violation:
|
||||
return Violation(
|
||||
path, line, "LIT001",
|
||||
path,
|
||||
line,
|
||||
"LIT001",
|
||||
f"mutable `{name}` in {where}: a mutable collection can be grown or rewritten "
|
||||
f"by whoever holds it. Annotate a read-only view -- Mapping[...], Sequence[...], "
|
||||
f"AbstractSet[...], tuple[X, ...], frozenset[X], or a frozen dataclass / "
|
||||
|
|
@ -233,13 +285,19 @@ def _mutable_ann(path: Path, line: int, name: str, where: str) -> Violation:
|
|||
|
||||
|
||||
def _annotation_violations(
|
||||
path: Path, annotation: ast.expr | None, line: int, where: str, ok_lines: frozenset[int]
|
||||
path: Path,
|
||||
annotation: ast.expr | None,
|
||||
line: int,
|
||||
where: str,
|
||||
ok_lines: frozenset[int],
|
||||
) -> Iterator[Violation]:
|
||||
if annotation is None or line in ok_lines:
|
||||
return
|
||||
yield from (_mutable_ann(path, line, name, where) for name in mutable_names_in(annotation))
|
||||
|
||||
|
||||
yield from (
|
||||
_mutable_ann(path, line, name, where) for name in mutable_names_in(annotation)
|
||||
)
|
||||
|
||||
|
||||
def _function_violations(
|
||||
path: Path, node: ast.FunctionDef | ast.AsyncFunctionDef, comments: Comments
|
||||
) -> Iterator[Violation]:
|
||||
|
|
@ -247,14 +305,22 @@ def _function_violations(
|
|||
args = node.args
|
||||
for arg in (*args.posonlyargs, *args.args, *args.kwonlyargs):
|
||||
yield from _annotation_violations(
|
||||
path, arg.annotation, arg.lineno, f"parameter `{arg.arg}` of `{node.name}`", mutable_ok
|
||||
path,
|
||||
arg.annotation,
|
||||
arg.lineno,
|
||||
f"parameter `{arg.arg}` of `{node.name}`",
|
||||
mutable_ok,
|
||||
)
|
||||
|
||||
# *args is allowed when typed (it's just a tuple); ruff ANN002 forces the
|
||||
# annotation, so here we only add the LIT001 mutable-collection check on the element type.
|
||||
if args.vararg is not None:
|
||||
yield from _annotation_violations(
|
||||
path, args.vararg.annotation, args.vararg.lineno, f"`*args` of `{node.name}`", mutable_ok
|
||||
path,
|
||||
args.vararg.annotation,
|
||||
args.vararg.lineno,
|
||||
f"`*args` of `{node.name}`",
|
||||
mutable_ok,
|
||||
)
|
||||
|
||||
# **kwargs is banned outright (LIT008): it erases the keyword contract and forces
|
||||
|
|
@ -262,7 +328,9 @@ def _function_violations(
|
|||
# cannot ban the syntax, so this rule does.
|
||||
if args.kwarg is not None and args.kwarg.lineno not in comments.kwargs_ok_lines:
|
||||
yield Violation(
|
||||
path, args.kwarg.lineno, "LIT008",
|
||||
path,
|
||||
args.kwarg.lineno,
|
||||
"LIT008",
|
||||
f"`**{args.kwarg.arg}` is banned: it erases the keyword contract and forces "
|
||||
f"Any-typing; declare explicit keyword parameters, or accept one frozen payload "
|
||||
f"(frozen dataclass / NamedTuple / ReadOnly TypedDict) "
|
||||
|
|
@ -271,11 +339,17 @@ def _function_violations(
|
|||
|
||||
if node.returns is not None:
|
||||
yield from _annotation_violations(
|
||||
path, node.returns, node.returns.lineno, f"return type of `{node.name}`", mutable_ok
|
||||
path,
|
||||
node.returns,
|
||||
node.returns.lineno,
|
||||
f"return type of `{node.name}`",
|
||||
mutable_ok,
|
||||
)
|
||||
|
||||
|
||||
def iter_annotation_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]:
|
||||
|
||||
|
||||
def iter_annotation_violations(
|
||||
path: Path, tree: ast.AST, comments: Comments
|
||||
) -> Iterator[Violation]:
|
||||
# Every annotation is in scope: signatures (params / *args / return) plus every
|
||||
# `x: T` -- class attribute, local, or module global. The latter three are all
|
||||
# ast.AnnAssign, so one walk covers them; only the signature annotations (which
|
||||
|
|
@ -286,8 +360,11 @@ def iter_annotation_violations(path: Path, tree: ast.AST, comments: Comments) ->
|
|||
elif isinstance(node, ast.AnnAssign):
|
||||
target = node.target.id if isinstance(node.target, ast.Name) else "<target>"
|
||||
yield from _annotation_violations(
|
||||
path, node.annotation, node.lineno,
|
||||
f"the type of `{target}`", comments.mutable_ok_lines,
|
||||
path,
|
||||
node.annotation,
|
||||
node.lineno,
|
||||
f"the type of `{target}`",
|
||||
comments.mutable_ok_lines,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -308,39 +385,54 @@ def _is_cast_call(node: ast.Call) -> bool:
|
|||
)
|
||||
|
||||
|
||||
def iter_cast_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]:
|
||||
def iter_cast_violations(
|
||||
path: Path, tree: ast.AST, comments: Comments
|
||||
) -> Iterator[Violation]:
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Call) and _is_cast_call(node) and node.lineno not in comments.cast_ok_lines:
|
||||
if (
|
||||
isinstance(node, ast.Call)
|
||||
and _is_cast_call(node)
|
||||
and node.lineno not in comments.cast_ok_lines
|
||||
):
|
||||
yield Violation(
|
||||
path, node.lineno, "LIT006",
|
||||
path,
|
||||
node.lineno,
|
||||
"LIT006",
|
||||
"cast() is an unchecked assertion (the type checker takes it on faith); "
|
||||
"validate into a frozen dataclass/NamedTuple/ReadOnly TypedDict at the "
|
||||
"boundary instead (suppress: `# cast-ok: <reason>`)",
|
||||
)
|
||||
|
||||
|
||||
def iter_guard_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]:
|
||||
def iter_guard_violations(
|
||||
path: Path, tree: ast.AST, comments: Comments
|
||||
) -> Iterator[Violation]:
|
||||
# TypeGuard/TypeIs are legal only as a function's return annotation (`-> TypeGuard[int]`),
|
||||
# so the walk is confined to `node.returns`; a runtime name that merely happens to read
|
||||
# `TypeGuard` is not a narrowing predicate. ruff bans the import; this flags the use.
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) or node.returns is None:
|
||||
if (
|
||||
not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
|
||||
or node.returns is None
|
||||
):
|
||||
continue
|
||||
for sub in ast.walk(node.returns):
|
||||
name = (
|
||||
sub.id if isinstance(sub, ast.Name)
|
||||
else sub.attr if isinstance(sub, ast.Attribute)
|
||||
else None
|
||||
sub.id
|
||||
if isinstance(sub, ast.Name)
|
||||
else sub.attr if isinstance(sub, ast.Attribute) else None
|
||||
)
|
||||
if name in UNSAFE_GUARDS and sub.lineno not in comments.guard_ok_lines:
|
||||
yield Violation(
|
||||
path, sub.lineno, "LIT007",
|
||||
path,
|
||||
sub.lineno,
|
||||
"LIT007",
|
||||
f"`{name}` narrowing predicate: the checker never verifies the body, so a "
|
||||
f"wrong guard silently corrupts types; parse into a concrete type instead "
|
||||
f"(suppress: `# guard-ok: <reason>`)",
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Mutable-collection construction (LIT002)
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
|
@ -395,7 +487,9 @@ def _construction_kind(node: ast.expr) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def iter_construction_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]:
|
||||
def iter_construction_violations(
|
||||
path: Path, tree: ast.AST, comments: Comments
|
||||
) -> Iterator[Violation]:
|
||||
in_annotation = _annotation_node_ids(tree)
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.expr) or id(node) in in_annotation:
|
||||
|
|
@ -404,32 +498,37 @@ def iter_construction_violations(path: Path, tree: ast.AST, comments: Comments)
|
|||
if kind is None or node.lineno in comments.mutable_ok_lines:
|
||||
continue
|
||||
yield Violation(
|
||||
path, node.lineno, "LIT002",
|
||||
path,
|
||||
node.lineno,
|
||||
"LIT002",
|
||||
f"mutable {kind}: this builds a collection that can be grown or rewritten. "
|
||||
f"Build it in one shot and freeze it -- a tuple/frozenset wrapping a generator "
|
||||
f"(`tuple(f(x) for x in xs)`), a tuple literal, or a frozen dataclass / NamedTuple "
|
||||
f"/ ReadOnly TypedDict (suppress: `# mutable-ok: <reason>`)",
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Driver
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
|
||||
|
||||
def check_file(path: Path) -> tuple[Violation, ...]:
|
||||
try:
|
||||
source = path.read_text(encoding="utf-8")
|
||||
except (OSError, UnicodeDecodeError) as exc:
|
||||
return (Violation(path, 0, "LIT000", f"could not read file: {exc}"),)
|
||||
|
||||
|
||||
comments, violations = scan_comments(path, source)
|
||||
|
||||
|
||||
try:
|
||||
tree = ast.parse(source, filename=str(path))
|
||||
except SyntaxError as exc:
|
||||
return (*violations, Violation(path, exc.lineno or 0, "LIT000", f"syntax error: {exc.msg}"))
|
||||
|
||||
return (
|
||||
*violations,
|
||||
Violation(path, exc.lineno or 0, "LIT000", f"syntax error: {exc.msg}"),
|
||||
)
|
||||
|
||||
return (
|
||||
*violations,
|
||||
*iter_annotation_violations(path, tree, comments),
|
||||
|
|
@ -437,8 +536,8 @@ def check_file(path: Path) -> tuple[Violation, ...]:
|
|||
*iter_guard_violations(path, tree, comments),
|
||||
*iter_construction_violations(path, tree, comments),
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
def collect_paths(raw: Iterable[str]) -> Iterator[Path]:
|
||||
for item in raw:
|
||||
p = Path(item)
|
||||
|
|
@ -446,24 +545,23 @@ def collect_paths(raw: Iterable[str]) -> Iterator[Path]:
|
|||
yield from sorted(p.rglob("*.py"))
|
||||
elif p.suffix == ".py":
|
||||
yield p
|
||||
|
||||
|
||||
|
||||
|
||||
def main(argv: Sequence[str]) -> int:
|
||||
paths = tuple(a for a in argv if not a.startswith("-"))
|
||||
if not paths:
|
||||
print("usage: check_type_discipline.py <files-or-dirs>...", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
|
||||
violations = sorted(v for path in collect_paths(paths) for v in check_file(path))
|
||||
for v in violations:
|
||||
print(v.render())
|
||||
|
||||
|
||||
if violations:
|
||||
print(f"\n{len(violations)} violation(s).", file=sys.stderr)
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main(sys.argv[1:]))
|
||||
|
||||
|
|
@ -42,8 +42,7 @@ from litellm.types.utils import CallTypes
|
|||
PROBLEMS = [
|
||||
{
|
||||
"id": "has_close_elements",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List
|
||||
|
||||
def has_close_elements(numbers: List[float], threshold: float) -> bool:
|
||||
|
|
@ -54,10 +53,8 @@ PROBLEMS = [
|
|||
>>> has_close_elements([1.0, 2.8, 3.0, 4.0, 5.0, 2.0], 0.3)
|
||||
True
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert has_close_elements([1.0, 2.0, 3.9, 4.0, 5.0, 2.2], 0.3) == True
|
||||
assert has_close_elements([1.0, 2.0, 3.9, 4.0, 5.0, 2.2], 0.05) == False
|
||||
assert has_close_elements([1.0, 2.0, 5.9, 4.0, 5.0], 0.95) == True
|
||||
|
|
@ -65,13 +62,11 @@ PROBLEMS = [
|
|||
assert has_close_elements([1.0, 2.0, 3.0, 4.0, 5.0], 2.0) == True
|
||||
assert has_close_elements([], 0.5) == False
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "separate_paren_groups",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List
|
||||
|
||||
def separate_paren_groups(paren_string: str) -> List[str]:
|
||||
|
|
@ -82,22 +77,18 @@ PROBLEMS = [
|
|||
>>> separate_paren_groups('( ) (( )) (( )( ))')
|
||||
['()', '(())', '(()())']
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert separate_paren_groups('(()()) ((())) () ((())()())') == ['(()())', '((()))', '()', '((())()())']
|
||||
assert separate_paren_groups('() (()) ((())) (((())))') == ['()', '(())', '((()))', '(((())))']
|
||||
assert separate_paren_groups('(()(()))') == ['(()(()))']
|
||||
assert separate_paren_groups('( ) (( )) (( )( ))') == ['()', '(())', '(()())']
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "truncate_number",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
def truncate_number(number: float) -> float:
|
||||
\"\"\"Given a positive floating point number, it can be decomposed into
|
||||
an integer part (largest integer smaller than given number) and decimals
|
||||
|
|
@ -106,21 +97,17 @@ PROBLEMS = [
|
|||
>>> truncate_number(3.5)
|
||||
0.5
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert truncate_number(3.5) == 0.5
|
||||
assert abs(truncate_number(1.33) - 0.33) < 1e-6
|
||||
assert abs(truncate_number(123.456) - 0.456) < 1e-6
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "below_zero",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List
|
||||
|
||||
def below_zero(operations: List[int]) -> bool:
|
||||
|
|
@ -132,10 +119,8 @@ PROBLEMS = [
|
|||
>>> below_zero([1, 2, -4, 5])
|
||||
True
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert below_zero([]) == False
|
||||
assert below_zero([1, 2, -3, 1, 2, -3]) == False
|
||||
assert below_zero([1, 2, -4, 5, 6]) == True
|
||||
|
|
@ -143,13 +128,11 @@ PROBLEMS = [
|
|||
assert below_zero([1, -1, 2, -2, 5, -5, 4, -5]) == True
|
||||
assert below_zero([1, -2]) == True
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "mean_absolute_deviation",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List
|
||||
|
||||
def mean_absolute_deviation(numbers: List[float]) -> float:
|
||||
|
|
@ -161,21 +144,17 @@ PROBLEMS = [
|
|||
>>> mean_absolute_deviation([1.0, 2.0, 3.0, 4.0])
|
||||
1.0
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert abs(mean_absolute_deviation([1.0, 2.0, 3.0, 4.0]) - 1.0) < 1e-6
|
||||
assert abs(mean_absolute_deviation([1.0, 2.0, 3.0, 4.0, 5.0]) - 1.2) < 1e-6
|
||||
assert abs(mean_absolute_deviation([1.0, 1.0, 1.0, 1.0]) - 0.0) < 1e-6
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "intersperse",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List
|
||||
|
||||
def intersperse(numbers: List[int], delimiter: int) -> List[int]:
|
||||
|
|
@ -185,21 +164,17 @@ PROBLEMS = [
|
|||
>>> intersperse([1, 2, 3], 4)
|
||||
[1, 4, 2, 4, 3]
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert intersperse([], 7) == []
|
||||
assert intersperse([5, 6, 3, 2], 8) == [5, 8, 6, 8, 3, 8, 2]
|
||||
assert intersperse([2, 2, 2], 2) == [2, 2, 2, 2, 2]
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "parse_nested_parens",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List
|
||||
|
||||
def parse_nested_parens(paren_string: str) -> List[int]:
|
||||
|
|
@ -209,21 +184,17 @@ PROBLEMS = [
|
|||
>>> parse_nested_parens('(()()) ((())) () ((())())')
|
||||
[2, 3, 1, 3]
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert parse_nested_parens('(()()) ((())) () ((())())') == [2, 3, 1, 3]
|
||||
assert parse_nested_parens('() (()) ((())) (((())))') == [1, 2, 3, 4]
|
||||
assert parse_nested_parens('(()(())((())))') == [4]
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "filter_by_substring",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List
|
||||
|
||||
def filter_by_substring(strings: List[str], substring: str) -> List[str]:
|
||||
|
|
@ -233,22 +204,18 @@ PROBLEMS = [
|
|||
>>> filter_by_substring(['abc', 'bacd', 'cde', 'array'], 'a')
|
||||
['abc', 'bacd', 'array']
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert filter_by_substring([], 'john') == []
|
||||
assert filter_by_substring(['xxx', 'asd', 'xxy', 'john doe', 'xxxuj', 'xxx'], 'xxx') == ['xxx', 'xxxuj', 'xxx']
|
||||
assert filter_by_substring(['xxx', 'asd', 'aaber', 'john doe', 'xxxuj', 'xxx'], 'xx') == ['xxx', 'xxxuj', 'xxx']
|
||||
assert filter_by_substring(['grunt', 'hierarchial', 'abc', 'hierarchial'], 'hi') == ['hierarchial', 'hierarchial']
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "sum_product",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List, Tuple
|
||||
|
||||
def sum_product(numbers: List[int]) -> Tuple[int, int]:
|
||||
|
|
@ -259,23 +226,19 @@ PROBLEMS = [
|
|||
>>> sum_product([1, 2, 3, 4])
|
||||
(10, 24)
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert sum_product([]) == (0, 1)
|
||||
assert sum_product([1, 1, 1]) == (3, 1)
|
||||
assert sum_product([100, 0]) == (100, 0)
|
||||
assert sum_product([3, 5, 7]) == (15, 105)
|
||||
assert sum_product([10]) == (10, 10)
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "max_element",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List
|
||||
|
||||
def max_element(l: List[int]) -> int:
|
||||
|
|
@ -285,21 +248,17 @@ PROBLEMS = [
|
|||
>>> max_element([5, 3, -5, 2, -3, 3, 9, 0, 123, 1, -10])
|
||||
123
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert max_element([1, 2, 3]) == 3
|
||||
assert max_element([5, 3, -5, 2, -3, 3, 9, 0, 124, 1, -10]) == 124
|
||||
assert max_element([-1, -2, -3]) == -1
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "fizz_buzz",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
def fizz_buzz(n: int) -> int:
|
||||
\"\"\"Return the number of times the digit 7 appears in integers less than n which are divisible by 11 or 13.
|
||||
>>> fizz_buzz(50)
|
||||
|
|
@ -309,10 +268,8 @@ PROBLEMS = [
|
|||
>>> fizz_buzz(79)
|
||||
3
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert fizz_buzz(50) == 0
|
||||
assert fizz_buzz(78) == 2
|
||||
assert fizz_buzz(79) == 3
|
||||
|
|
@ -320,13 +277,11 @@ PROBLEMS = [
|
|||
assert fizz_buzz(200) == 6
|
||||
assert fizz_buzz(4000) == 192
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
{
|
||||
"id": "sort_by_binary_len",
|
||||
"prompt": textwrap.dedent(
|
||||
"""\
|
||||
"prompt": textwrap.dedent("""\
|
||||
from typing import List
|
||||
|
||||
def sort_array(arr: List[int]) -> List[int]:
|
||||
|
|
@ -339,10 +294,8 @@ PROBLEMS = [
|
|||
>>> sort_array([1, 0, 2, 3, 4])
|
||||
[0, 1, 2, 4, 3]
|
||||
\"\"\"
|
||||
"""
|
||||
),
|
||||
"tests": textwrap.dedent(
|
||||
"""\
|
||||
"""),
|
||||
"tests": textwrap.dedent("""\
|
||||
assert sort_array([1, 5, 2, 3, 4]) == [1, 2, 4, 3, 5]
|
||||
assert sort_array([-2, -3, -4, -5, -6]) == [-6, -5, -4, -3, -2]
|
||||
assert sort_array([1, 0, 2, 3, 4]) == [0, 1, 2, 4, 3]
|
||||
|
|
@ -350,8 +303,7 @@ PROBLEMS = [
|
|||
assert sort_array([2, 5, 77, 4, 5, 3, 5, 7, 2, 3, 4]) == [2, 2, 4, 4, 3, 3, 5, 5, 5, 7, 77]
|
||||
assert sort_array([3, 6, 44, 12, 32, 5]) == [32, 3, 5, 6, 12, 44]
|
||||
print("PASSED")
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
},
|
||||
]
|
||||
|
||||
|
|
@ -360,8 +312,7 @@ PROBLEMS = [
|
|||
# compressor to identify and drop them.
|
||||
DISTRACTOR_SNIPPETS = [
|
||||
# distractor 0 — database connection pool
|
||||
textwrap.dedent(
|
||||
"""\
|
||||
textwrap.dedent("""\
|
||||
# db_pool.py
|
||||
import threading
|
||||
from contextlib import contextmanager
|
||||
|
|
@ -408,11 +359,9 @@ DISTRACTOR_SNIPPETS = [
|
|||
for conn in self._pool:
|
||||
conn.close()
|
||||
self._pool.clear()
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
# distractor 1 — HTTP retry logic
|
||||
textwrap.dedent(
|
||||
"""\
|
||||
textwrap.dedent("""\
|
||||
# http_retry.py
|
||||
import time
|
||||
import random
|
||||
|
|
@ -456,11 +405,9 @@ DISTRACTOR_SNIPPETS = [
|
|||
resp = requests.get(url, params=params, timeout=30)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
# distractor 2 — LRU cache implementation
|
||||
textwrap.dedent(
|
||||
"""\
|
||||
textwrap.dedent("""\
|
||||
# lru_cache.py
|
||||
from collections import OrderedDict
|
||||
from threading import RLock
|
||||
|
|
@ -509,11 +456,9 @@ DISTRACTOR_SNIPPETS = [
|
|||
|
||||
def __contains__(self, key):
|
||||
return key in self._cache
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
# distractor 3 — CSV report generator
|
||||
textwrap.dedent(
|
||||
"""\
|
||||
textwrap.dedent("""\
|
||||
# report_gen.py
|
||||
import csv
|
||||
import io
|
||||
|
|
@ -566,11 +511,9 @@ DISTRACTOR_SNIPPETS = [
|
|||
except (ValueError, KeyError):
|
||||
return False
|
||||
return self.filter_rows(in_range)
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
# distractor 4 — async task queue
|
||||
textwrap.dedent(
|
||||
"""\
|
||||
textwrap.dedent("""\
|
||||
# task_queue.py
|
||||
import asyncio
|
||||
import logging
|
||||
|
|
@ -640,11 +583,9 @@ DISTRACTOR_SNIPPETS = [
|
|||
async def shutdown(self):
|
||||
for w in self._workers:
|
||||
w.cancel()
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
# distractor 5 — config parser with env var interpolation
|
||||
textwrap.dedent(
|
||||
"""\
|
||||
textwrap.dedent("""\
|
||||
# config_parser.py
|
||||
import os
|
||||
import re
|
||||
|
|
@ -707,8 +648,7 @@ DISTRACTOR_SNIPPETS = [
|
|||
if val is None:
|
||||
raise ConfigError(f"Required config key missing: {key}")
|
||||
return val
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -1036,7 +976,7 @@ def run_benchmark(
|
|||
print(f" Avg total tokens: {base_agg['avg_total_tokens']}")
|
||||
print(f" Avg latency: {base_agg['avg_latency_ms']}ms")
|
||||
|
||||
print(f"\n Compressed (litellm.compress → then call model):")
|
||||
print("\n Compressed (litellm.compress → then call model):")
|
||||
print(
|
||||
f" Pass rate: {comp_agg['pass_rate']}% ({comp_agg['passed']}/{comp_agg['total']})"
|
||||
)
|
||||
|
|
@ -1054,7 +994,7 @@ def run_benchmark(
|
|||
latency_diff = base_agg["avg_latency_ms"] - comp_agg["avg_latency_ms"]
|
||||
pass_diff = comp_agg["pass_rate"] - base_agg["pass_rate"]
|
||||
|
||||
print(f"\n Delta (compressed vs baseline):")
|
||||
print("\n Delta (compressed vs baseline):")
|
||||
print(f" Token savings: {token_savings} tokens ({token_pct}%)")
|
||||
print(f" Latency delta: {latency_diff:+.1f}ms")
|
||||
print(f" Pass rate delta: {pass_diff:+.1f}%")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
#!/usr/bin/env python3
|
||||
"""
|
||||
Bench LiteLLM_HealthCheckTable + PrismaClient
|
||||
Bench LiteLLM_HealthCheckTable + PrismaClient
|
||||
- set DATABASE_URL to your Postgres
|
||||
- Run ```prisma generate``` to install prisma client before running test )
|
||||
- This test writes to the default "public" database. Make sure to run cleanup after testing
|
||||
|
|
|
|||
|
|
@ -360,7 +360,7 @@ class LiteLLMHealthCheckClient:
|
|||
|
||||
# Print detailed results for each model (matching Go output format)
|
||||
print(f"\n{'='*60}", file=sys.stderr)
|
||||
print(f"Starting health check queries\n", file=sys.stderr)
|
||||
print("Starting health check queries\n", file=sys.stderr)
|
||||
|
||||
for model_id, result in results.items():
|
||||
if result.get("healthy"):
|
||||
|
|
@ -383,7 +383,7 @@ class LiteLLMHealthCheckClient:
|
|||
print(f"---- {model_id} ----\n❌ ERROR: {error}\n\n", file=sys.stderr)
|
||||
|
||||
print(f"{'='*60}", file=sys.stderr)
|
||||
print(f"Health Check Summary", file=sys.stderr)
|
||||
print("Health Check Summary", file=sys.stderr)
|
||||
print(f"{'='*60}", file=sys.stderr)
|
||||
print(f"Total models: {len(results)}", file=sys.stderr)
|
||||
print(f"Healthy: {healthy_count}", file=sys.stderr)
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ Notes
|
|||
- `converse-stream` is still JSON-only placeholder (different inner event shapes).
|
||||
- Use real (or any non-empty) AWS creds in the environment of the **proxy**; signing still runs.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
|
|
|
|||
|
|
@ -38,7 +38,7 @@ class SlowHandler(BaseHTTPRequestHandler):
|
|||
return None
|
||||
return self.rfile.read(length)
|
||||
|
||||
def do_POST(self) -> None: # noqa: N802
|
||||
def do_POST(self) -> None:
|
||||
if self.path != "/cygnal/monitor":
|
||||
self.send_error(404, "Not Found")
|
||||
return
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ reader to write tests that kill the survivors.
|
|||
Run after `mutmut run` and `mutmut export-cicd-stats`. Expects mutmut to be
|
||||
invokable as `uv run --no-sync --with mutmut==<version> mutmut <subcommand>`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
|
|
@ -373,9 +374,7 @@ def render(config: dict, survivors: list[str], stats: dict | None) -> str:
|
|||
|
||||
out.append("## Task")
|
||||
out.append("")
|
||||
out.append(
|
||||
dedent(
|
||||
"""\
|
||||
out.append(dedent("""\
|
||||
For each surviving mutant listed above, write a new test in the
|
||||
existing test file (matching its conventions, fixtures, and naming
|
||||
style) that:
|
||||
|
|
@ -388,9 +387,7 @@ def render(config: dict, survivors: list[str], stats: dict | None) -> str:
|
|||
which mutant numbers in the test name or docstring.
|
||||
|
||||
Do not modify the source file. Only add tests.
|
||||
"""
|
||||
).strip()
|
||||
)
|
||||
""").strip())
|
||||
out.append("")
|
||||
|
||||
return "\n".join(out)
|
||||
|
|
|
|||
|
|
@ -98,7 +98,9 @@ def base_counts(ref: str) -> dict:
|
|||
# the body (or the `worktree add` itself) failed. rmtree is already best-effort.
|
||||
subprocess.run(
|
||||
["git", "worktree", "remove", "--force", str(worktree)],
|
||||
cwd=REPO_ROOT, capture_output=True, text=True,
|
||||
cwd=REPO_ROOT,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
shutil.rmtree(parent, ignore_errors=True)
|
||||
|
||||
|
|
@ -110,7 +112,8 @@ def over_ceiling(head: dict, budget: dict) -> frozenset:
|
|||
comparison cannot change the verdict and the base worktree scan can be skipped.
|
||||
"""
|
||||
return frozenset(
|
||||
rule for rule, spec in budget.items()
|
||||
rule
|
||||
for rule, spec in budget.items()
|
||||
if head.get(rule, 0) > spec["baseline"] + spec["slack"]
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -25,7 +25,12 @@ from typing import AsyncIterator, Final
|
|||
import uvicorn
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse, PlainTextResponse, Response, StreamingResponse
|
||||
from starlette.responses import (
|
||||
JSONResponse,
|
||||
PlainTextResponse,
|
||||
Response,
|
||||
StreamingResponse,
|
||||
)
|
||||
from starlette.routing import Route
|
||||
|
||||
_CANNED_CONTENT: Final = "Hello! This is a mock response from the fake OpenAI endpoint."
|
||||
|
|
@ -159,7 +164,14 @@ async def _text_completion_stream(model: str, with_usage: bool) -> AsyncIterator
|
|||
"object": "text_completion",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [{"text": text, "index": 0, "logprobs": None, "finish_reason": finish_reason}],
|
||||
"choices": [
|
||||
{
|
||||
"text": text,
|
||||
"index": 0,
|
||||
"logprobs": None,
|
||||
"finish_reason": finish_reason,
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
yield f"data: {json.dumps(chunk(_CANNED_CONTENT, None))}\n\n"
|
||||
|
|
@ -190,7 +202,10 @@ async def embeddings(request: Request) -> Response:
|
|||
return JSONResponse(
|
||||
{
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": i, "embedding": [0.0] * 1536} for i in range(max(count, 1))],
|
||||
"data": [
|
||||
{"object": "embedding", "index": i, "embedding": [0.0] * 1536}
|
||||
for i in range(max(count, 1))
|
||||
],
|
||||
"model": _requested_model(body),
|
||||
"usage": {"prompt_tokens": 5, "total_tokens": 5},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -19,7 +19,6 @@ from litellm.batches.batch_utils import (
|
|||
from litellm.cost_calculator import batch_cost_calculator
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
|
||||
# --- helpers ---
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -21,7 +21,6 @@ from unittest.mock import patch, MagicMock
|
|||
import httpx
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
|
||||
_BEDROCK_TEST_AWS_ENV = {
|
||||
"AWS_ACCESS_KEY_ID": "test-access-key",
|
||||
"AWS_SECRET_ACCESS_KEY": "test-secret-key",
|
||||
|
|
@ -99,7 +98,9 @@ class _CaptureAsyncHTTPHandler(AsyncHTTPHandler):
|
|||
raw = json if json is not None else (data if data is not None else content)
|
||||
payload = raw if isinstance(raw, dict) else json_module.loads(raw)
|
||||
job_name = payload["jobName"]
|
||||
job_arn = f"arn:aws:bedrock:us-west-2:941277531214:model-invocation-job/{job_name}"
|
||||
job_arn = (
|
||||
f"arn:aws:bedrock:us-west-2:941277531214:model-invocation-job/{job_name}"
|
||||
)
|
||||
self.batch_jobs[job_arn] = {
|
||||
"jobArn": job_arn,
|
||||
"jobName": job_name,
|
||||
|
|
|
|||
|
|
@ -20,7 +20,6 @@ sys.path.insert(0, os.path.abspath("../.."))
|
|||
|
||||
import litellm
|
||||
|
||||
|
||||
SERVER_URL = "https://exampleopenaiendpoint-production-0ee2.up.railway.app/v1"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ import litellm
|
|||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.token_counter import token_counter
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import ast
|
|||
import os
|
||||
from typing import List, Dict, Any
|
||||
|
||||
|
||||
ALLOWED_FILE = os.path.normpath("litellm/_uuid.py")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -23,8 +23,6 @@ def event_loop():
|
|||
loop.close()
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.fixture(scope="function", autouse=True)
|
||||
def setup_and_teardown():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -801,7 +801,9 @@ async def test_list_projects_returns_timestamps():
|
|||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm_enterprise.proxy.management_endpoints.project_endpoints import list_projects
|
||||
from litellm_enterprise.proxy.management_endpoints.project_endpoints import (
|
||||
list_projects,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ProjectTable
|
||||
|
||||
now = datetime(2024, 1, 15, 12, 0, 0, tzinfo=timezone.utc)
|
||||
|
|
|
|||
|
|
@ -14,7 +14,6 @@ from litellm.proxy.guardrails.guardrail_registry import (
|
|||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.akto.akto import AktoGuardrail
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Registry tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ import io
|
|||
import os
|
||||
import sys
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import asyncio
|
||||
|
|
|
|||
|
|
@ -21,7 +21,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
|
|||
ContentFilterCategoryConfig,
|
||||
)
|
||||
|
||||
|
||||
# Test cases: (sentence, expected_result, reason)
|
||||
TEST_CASES = [
|
||||
# ALWAYS BLOCK - Explicit prohibited practices (1-10)
|
||||
|
|
|
|||
|
|
@ -23,7 +23,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
|
|||
ContentFilterCategoryConfig,
|
||||
)
|
||||
|
||||
|
||||
# ── helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
POLICY_DIR = os.path.abspath(
|
||||
|
|
|
|||
|
|
@ -28,7 +28,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
|
|||
ContentFilterCategoryConfig,
|
||||
)
|
||||
|
||||
|
||||
# ── helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
POLICY_DIR = os.path.abspath(
|
||||
|
|
|
|||
|
|
@ -7,7 +7,6 @@ import sys
|
|||
import traceback
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ import os
|
|||
import sys
|
||||
import traceback
|
||||
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
|
|
|||
|
|
@ -278,6 +278,7 @@ def test_embedding_via_proxy(proxy_url: str) -> None:
|
|||
assert len(embedding) >= 64
|
||||
assert all(isinstance(x, (int, float)) for x in embedding)
|
||||
|
||||
|
||||
def test_model_list_advertises_oci_models(proxy_url: str) -> None:
|
||||
"""The /v1/models registry advertises every OCI alias from the config."""
|
||||
r = httpx.get(
|
||||
|
|
@ -288,7 +289,9 @@ def test_model_list_advertises_oci_models(proxy_url: str) -> None:
|
|||
assert r.status_code == 200, r.text
|
||||
advertised = {row["id"] for row in r.json()["data"]}
|
||||
for expected in CHAT_MODELS + ["oci-embed"]:
|
||||
assert expected in advertised, f"{expected} missing from /v1/models: {advertised}"
|
||||
assert (
|
||||
expected in advertised
|
||||
), f"{expected} missing from /v1/models: {advertised}"
|
||||
|
||||
|
||||
def test_chat_completion_no_drop_params(proxy_url_no_drop_params: str) -> None:
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ import pytest
|
|||
from litellm.types.utils import Embedding
|
||||
from litellm.main import bedrock_embedding, embedding, EmbeddingResponse, Usage
|
||||
|
||||
|
||||
_mock_model_id = (
|
||||
"arn:aws:bedrock:us-east-1:123412341234:application-inference-profile/abc123123"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ from litellm.llms.bedrock.common_utils import (
|
|||
)
|
||||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
|
||||
|
||||
NOVA_ARN = "arn:aws:bedrock:us-east-1:123456789012:custom-model-deployment/a1b2c3d4e5f6"
|
||||
NOVA_MODEL = f"bedrock/nova/{NOVA_ARN}"
|
||||
NOVA2_MODEL = f"bedrock/nova-2/{NOVA_ARN}"
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ import pytest
|
|||
|
||||
from litellm import get_model_info
|
||||
|
||||
|
||||
MODEL_NAME = "nvidia.nemotron-super-3-120b"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
BedrockConverseMessagesProcessor,
|
||||
)
|
||||
|
||||
|
||||
MODEL = "anthropic.claude-v2"
|
||||
PROVIDER = "bedrock_converse"
|
||||
|
||||
|
|
@ -546,9 +545,9 @@ def test_bedrock_converse_sorts_text_before_tooluse_sync():
|
|||
tool_indices = [i for i, b in enumerate(content) if "toolUse" in b]
|
||||
|
||||
# All text blocks must come before all toolUse blocks
|
||||
assert max(text_indices) < min(tool_indices), (
|
||||
f"text blocks at {text_indices} should all precede toolUse blocks at {tool_indices}"
|
||||
)
|
||||
assert max(text_indices) < min(
|
||||
tool_indices
|
||||
), f"text blocks at {text_indices} should all precede toolUse blocks at {tool_indices}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -567,9 +566,9 @@ async def test_bedrock_converse_sorts_text_before_tooluse_async():
|
|||
text_indices = [i for i, b in enumerate(content) if "text" in b]
|
||||
tool_indices = [i for i, b in enumerate(content) if "toolUse" in b]
|
||||
|
||||
assert max(text_indices) < min(tool_indices), (
|
||||
f"text blocks at {text_indices} should all precede toolUse blocks at {tool_indices}"
|
||||
)
|
||||
assert max(text_indices) < min(
|
||||
tool_indices
|
||||
), f"text blocks at {text_indices} should all precede toolUse blocks at {tool_indices}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -577,7 +576,9 @@ async def test_bedrock_converse_content_ordering_sync_async_parity():
|
|||
"""Sync and async paths should produce identical content block ordering."""
|
||||
messages = _make_tooluse_before_text_messages()
|
||||
sync_result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
|
||||
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
|
||||
messages, MODEL, PROVIDER
|
||||
async_result = (
|
||||
await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
|
||||
messages, MODEL, PROVIDER
|
||||
)
|
||||
)
|
||||
assert sync_result == async_result
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ def _attrify(d: dict):
|
|||
None)` (et al), which returns None for plain dicts — that would silently
|
||||
skip the row.
|
||||
"""
|
||||
|
||||
class _AttrDict(dict):
|
||||
def __getattr__(self, k):
|
||||
try:
|
||||
|
|
@ -120,7 +121,9 @@ async def test_reset_budget_keys_partial_failure():
|
|||
key1, key2, key3, key4, key5, key6 = (
|
||||
_attrify(k) for k in [key1, key2, key3, key4, key5, key6]
|
||||
)
|
||||
prisma_client.get_data = AsyncMock(return_value=[key1, key2, key3, key4, key5, key6])
|
||||
prisma_client.get_data = AsyncMock(
|
||||
return_value=[key1, key2, key3, key4, key5, key6]
|
||||
)
|
||||
|
||||
async def fake_reset_key(key, current_time):
|
||||
if key["id"] == "key1":
|
||||
|
|
@ -207,7 +210,9 @@ async def test_reset_budget_users_partial_failure():
|
|||
user1, user2, user3, user4, user5, user6 = (
|
||||
_attrify(u) for u in [user1, user2, user3, user4, user5, user6]
|
||||
)
|
||||
prisma_client.get_data = AsyncMock(return_value=[user1, user2, user3, user4, user5, user6])
|
||||
prisma_client.get_data = AsyncMock(
|
||||
return_value=[user1, user2, user3, user4, user5, user6]
|
||||
)
|
||||
|
||||
async def fake_reset_user(user, current_time):
|
||||
if user["id"] == "user1":
|
||||
|
|
|
|||
|
|
@ -2490,9 +2490,7 @@ def test_get_base_model_from_metadata():
|
|||
"litellm_params": {"base_model": "azure/gpt-5-mini"}
|
||||
}
|
||||
result = _get_base_model_from_metadata(model_call_details_with_direct_base_model)
|
||||
assert (
|
||||
result == "azure/gpt-5-mini"
|
||||
), f"Expected 'azure/gpt-5-mini', got {result}"
|
||||
assert result == "azure/gpt-5-mini", f"Expected 'azure/gpt-5-mini', got {result}"
|
||||
|
||||
# Test 4: metadata takes precedence over litellm_metadata
|
||||
model_call_details_with_both = {
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ from litellm.responses.litellm_completion_transformation.transformation import (
|
|||
)
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
Test to reproduce and verify fix for Anthropic tool_result issue with empty call_id.
|
||||
|
||||
This test reproduces the exact error:
|
||||
"messages.0.content.0: unexpected `tool_use_id` found in `tool_result` blocks: tool_use_id.
|
||||
"messages.0.content.0: unexpected `tool_use_id` found in `tool_result` blocks: tool_use_id.
|
||||
Each `tool_result` block must have a corresponding `tool_use` block in the previous message."
|
||||
|
||||
The issue occurs when:
|
||||
|
|
|
|||
|
|
@ -2,12 +2,12 @@
|
|||
Unit tests for BaseResponsesAPIStreamingIterator
|
||||
|
||||
Tests core functionality including:
|
||||
1. Processing chunks and handling ResponseCompletedEvent
|
||||
1. Processing chunks and handling ResponseCompletedEvent
|
||||
2. Ensuring _update_responses_api_response_id_with_model_id is called for final chunk
|
||||
3. Verifying ID update is NOT called for non-final chunks (delta events)
|
||||
4. Edge case handling for invalid JSON, empty chunks, and [DONE] markers
|
||||
|
||||
These tests ensure the streaming iterator correctly processes response chunks
|
||||
These tests ensure the streaming iterator correctly processes response chunks
|
||||
and applies model ID updates only to completed responses, as required for proper
|
||||
response tracking and logging.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -387,9 +387,7 @@ def test_process_chunk_completed_response_updates_id_and_usage_cost(monkeypatch)
|
|||
# Chunk must include a top-level "response" key so BaseResponsesAPIStreamingIterator
|
||||
# runs _update_responses_api_response_id_with_model_id (see streaming_iterator.py).
|
||||
event = iterator._process_chunk(
|
||||
json.dumps(
|
||||
{"type": "response.completed", "response": {"id": "resp_live"}}
|
||||
)
|
||||
json.dumps({"type": "response.completed", "response": {"id": "resp_live"}})
|
||||
)
|
||||
finally:
|
||||
litellm.include_cost_in_streaming_usage = original_include_cost
|
||||
|
|
|
|||
|
|
@ -22,10 +22,8 @@ sys.path.insert(0, os.path.abspath("../.."))
|
|||
import litellm
|
||||
from litellm import completion
|
||||
|
||||
|
||||
# Large document for caching tests (needs 1024+ tokens for Claude models)
|
||||
LARGE_DOCUMENT_FOR_CACHING = (
|
||||
"""
|
||||
LARGE_DOCUMENT_FOR_CACHING = """
|
||||
This is a comprehensive legal agreement between Party A and Party B.
|
||||
|
||||
ARTICLE 1: DEFINITIONS
|
||||
|
|
@ -77,9 +75,7 @@ ARTICLE 9: GENERAL PROVISIONS
|
|||
9.5 Waiver of any provision shall not constitute ongoing waiver.
|
||||
|
||||
IN WITNESS WHEREOF, the parties have executed this Agreement.
|
||||
"""
|
||||
* 8
|
||||
) # Repeat to ensure we have enough tokens (need 1024+ for Claude models)
|
||||
""" * 8 # Repeat to ensure we have enough tokens (need 1024+ for Claude models)
|
||||
|
||||
|
||||
class TestBedrockAnthropicPromptCachingRegression:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
"""
|
||||
Tests for Crusoe provider integration
|
||||
"""
|
||||
|
||||
import os
|
||||
from unittest import mock
|
||||
|
||||
|
|
@ -96,9 +97,9 @@ def test_crusoe_models_configuration():
|
|||
for model in crusoe_models:
|
||||
model_info = get_model_info(model)
|
||||
assert model_info is not None, f"Model info not found for {model}"
|
||||
assert model_info.get("litellm_provider") == "crusoe", (
|
||||
f"{model} should have crusoe as provider"
|
||||
)
|
||||
assert (
|
||||
model_info.get("litellm_provider") == "crusoe"
|
||||
), f"{model} should have crusoe as provider"
|
||||
assert model_info.get("mode") == "chat", f"{model} should be in chat mode"
|
||||
finally:
|
||||
litellm.model_cost = original_model_cost
|
||||
|
|
|
|||
|
|
@ -191,7 +191,11 @@ def test_deepseek_fill_reasoning_content_multiturn():
|
|||
# Case 1: assistant message already has reasoning_content — should be left as-is
|
||||
messages_with_rc = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "Hi", "reasoning_content": "I thought about it"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Hi",
|
||||
"reasoning_content": "I thought about it",
|
||||
},
|
||||
{"role": "user", "content": "Follow up"},
|
||||
]
|
||||
result = config._fill_reasoning_content(messages_with_rc)
|
||||
|
|
@ -259,9 +263,9 @@ def test_deepseek_fill_reasoning_content_guard_in_transform_request():
|
|||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert result["messages"][1].get("reasoning_content") == " ", (
|
||||
"reasoning_content should be injected when thinking is enabled"
|
||||
)
|
||||
assert (
|
||||
result["messages"][1].get("reasoning_content") == " "
|
||||
), "reasoning_content should be injected when thinking is enabled"
|
||||
|
||||
# Case 2: reasoning model + thinking NOT in optional_params -> no injection
|
||||
result = config.transform_request(
|
||||
|
|
@ -271,9 +275,9 @@ def test_deepseek_fill_reasoning_content_guard_in_transform_request():
|
|||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert "reasoning_content" not in result["messages"][1], (
|
||||
"reasoning_content should not be injected when thinking is not enabled"
|
||||
)
|
||||
assert (
|
||||
"reasoning_content" not in result["messages"][1]
|
||||
), "reasoning_content should not be injected when thinking is not enabled"
|
||||
|
||||
# Case 3: non-reasoning model + thinking enabled -> no injection
|
||||
result = config.transform_request(
|
||||
|
|
@ -283,6 +287,6 @@ def test_deepseek_fill_reasoning_content_guard_in_transform_request():
|
|||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert "reasoning_content" not in result["messages"][1], (
|
||||
"reasoning_content should not be injected for non-reasoning models"
|
||||
)
|
||||
assert (
|
||||
"reasoning_content" not in result["messages"][1]
|
||||
), "reasoning_content should not be injected for non-reasoning models"
|
||||
|
|
|
|||
|
|
@ -16,7 +16,6 @@ import litellm
|
|||
from litellm import completion
|
||||
import json
|
||||
|
||||
|
||||
GEMINI_3_IMAGE_SIZE_MAPPINGS = [
|
||||
("512x512", "1:1", "512"),
|
||||
("1024x1024", "1:1", "1K"),
|
||||
|
|
@ -557,9 +556,10 @@ def test_gemini_image_generation_openai_size_auto_uses_google_defaults(size: str
|
|||
map_openai_size_to_gemini_image_config,
|
||||
)
|
||||
|
||||
assert map_openai_size_to_gemini_image_config(
|
||||
size, "gemini-3-pro-image-preview"
|
||||
) is None
|
||||
assert (
|
||||
map_openai_size_to_gemini_image_config(size, "gemini-3-pro-image-preview")
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_gemini_imagen_models_use_predict_endpoint():
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
"""
|
||||
Test for Gemini image generation usage metadata extraction.
|
||||
|
||||
This test verifies the fix for issue #18323 where image_generation()
|
||||
This test verifies the fix for issue #18323 where image_generation()
|
||||
was returning usage=0 while completion() returned proper token usage.
|
||||
"""
|
||||
|
||||
|
|
|
|||
|
|
@ -1683,7 +1683,9 @@ class TestMissingChoicesGuard:
|
|||
|
||||
assert "no 'choices'" in exc_info.value.message
|
||||
|
||||
def test_convert_to_model_response_object_stream_true_no_choices_raises_api_error(self):
|
||||
def test_convert_to_model_response_object_stream_true_no_choices_raises_api_error(
|
||||
self,
|
||||
):
|
||||
"""Missing choices via stream=True path raises APIError when generator is consumed."""
|
||||
from litellm.exceptions import APIError
|
||||
|
||||
|
|
@ -2471,6 +2473,13 @@ class TestConvertToModelResponseObjectCompletion:
|
|||
def test_model_response_none_raises(self):
|
||||
with pytest.raises(Exception):
|
||||
convert_to_model_response_object(
|
||||
response_object={"choices": [{"message": {"content": "hi", "role": "assistant"}, "finish_reason": "stop"}]},
|
||||
response_object={
|
||||
"choices": [
|
||||
{
|
||||
"message": {"content": "hi", "role": "assistant"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
]
|
||||
},
|
||||
model_response_object=None,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import io
|
|||
import os
|
||||
import sys
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import asyncio
|
||||
|
|
|
|||
|
|
@ -38,7 +38,6 @@ from litellm.llms.vertex_ai.gemini.transformation import (
|
|||
)
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
|
||||
|
||||
litellm.num_retries = 3
|
||||
litellm.cache = None
|
||||
user_message = "Write a short poem about the sky"
|
||||
|
|
@ -1104,9 +1103,7 @@ def vertex_httpx_mock_post_valid_response(*args, **kwargs):
|
|||
{
|
||||
"content": {
|
||||
"role": "model",
|
||||
"parts": [
|
||||
{
|
||||
"text": """{
|
||||
"parts": [{"text": """{
|
||||
"recipes": [
|
||||
{"recipe_name": "Chocolate Chip Cookies"},
|
||||
{"recipe_name": "Oatmeal Raisin Cookies"},
|
||||
|
|
@ -1114,9 +1111,7 @@ def vertex_httpx_mock_post_valid_response(*args, **kwargs):
|
|||
{"recipe_name": "Sugar Cookies"},
|
||||
{"recipe_name": "Snickerdoodles"}
|
||||
]
|
||||
}"""
|
||||
}
|
||||
],
|
||||
}"""}],
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
"safetyRatings": [
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Test for preset_cache_key multiple values bug fix.
|
|||
This test verifies that get_cache_key doesn't raise TypeError when kwargs
|
||||
already contains preset_cache_key.
|
||||
|
||||
Issue: When get_cache_key(**kwargs) is called with kwargs containing
|
||||
Issue: When get_cache_key(**kwargs) is called with kwargs containing
|
||||
preset_cache_key, the call to _set_preset_cache_key_in_kwargs() would fail with:
|
||||
TypeError: got multiple values for keyword argument 'preset_cache_key'
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -450,12 +450,9 @@ def c():
|
|||
litellm.enable_caching_on_provider_specific_optional_params = False
|
||||
|
||||
|
||||
embedding_large_text = (
|
||||
"""
|
||||
embedding_large_text = """
|
||||
small text
|
||||
"""
|
||||
* 5
|
||||
)
|
||||
""" * 5
|
||||
|
||||
|
||||
# # test_caching_with_models()
|
||||
|
|
|
|||
|
|
@ -1666,15 +1666,13 @@ def custom_callback(
|
|||
|
||||
#################################################
|
||||
|
||||
print(
|
||||
f"""
|
||||
print(f"""
|
||||
Model: {model},
|
||||
Messages: {messages},
|
||||
User: {user},
|
||||
Seed: {kwargs["seed"]},
|
||||
temperature: {kwargs["temperature"]},
|
||||
"""
|
||||
)
|
||||
""")
|
||||
|
||||
assert kwargs["user"] == "ishaans app"
|
||||
assert kwargs["model"] == "gpt-3.5-turbo-1106"
|
||||
|
|
|
|||
|
|
@ -93,8 +93,7 @@ class testCustomCallbackProxy(CustomLogger):
|
|||
|
||||
print("\n\n in custom callback vars my custom logger, ", vars(my_custom_logger))
|
||||
|
||||
print(
|
||||
f"""
|
||||
print(f"""
|
||||
Model: {model},
|
||||
Messages: {messages},
|
||||
User: {user},
|
||||
|
|
@ -102,8 +101,7 @@ class testCustomCallbackProxy(CustomLogger):
|
|||
Cost: {cost},
|
||||
Response: {response}
|
||||
Proxy Metadata: {metadata}
|
||||
"""
|
||||
)
|
||||
""")
|
||||
return
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
|
|
|
|||
|
|
@ -267,7 +267,6 @@ def test_aaparallel_function_call_with_anthropic_thinking(model):
|
|||
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
|
||||
|
||||
|
||||
_PARALLEL_TOOL_HISTORY_MESSAGES = [
|
||||
{
|
||||
"role": "user",
|
||||
|
|
|
|||
|
|
@ -21,7 +21,6 @@ verbose_logger.setLevel(logging.DEBUG)
|
|||
litellm.set_verbose = True
|
||||
import time
|
||||
|
||||
|
||||
# test_langsmith_logging()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -16,7 +16,6 @@ import litellm
|
|||
from litellm import Router
|
||||
from litellm.caching import RedisCache, RedisClusterCache
|
||||
|
||||
|
||||
## Scenarios
|
||||
## 1. 2 models - openai + azure - 1 model group "gpt-3.5-turbo",
|
||||
## 2. 2 models - openai, azure - 2 diff model groups, 1 caching group
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import io
|
|||
import os
|
||||
import sys
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import asyncio
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import io
|
|||
import os
|
||||
import sys
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import asyncio
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ import io
|
|||
import os
|
||||
import sys
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import asyncio
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import io
|
|||
import os
|
||||
import sys
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import asyncio
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import io
|
|||
import os
|
||||
import sys
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import asyncio
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import io
|
|||
import os
|
||||
import sys
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import asyncio
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue