Merge branch 'BerriAI:litellm_internal_staging' into litellm_internal_staging

This commit is contained in:
jamiexiami 2026-05-07 11:29:12 +08:00 • committed by GitHub
commit ef1737e456
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
43 changed files with 3173 additions and 202 deletions

View file

@ -141,6 +141,7 @@ jobs:
tests/proxy_unit_tests/test_server_root_path.py
tests/proxy_unit_tests/test_proxy_pass_user_config.py
tests/proxy_unit_tests/test_proxy_token_counter.py
tests/proxy_unit_tests/test_request_size_limit_middleware.py
workers: 4
dist: loadscope
timeout: 15

3
.gitignore vendored
View file

@ -100,4 +100,5 @@ STABILIZATION_TODO.md
**/playwright-report
**/*.storageState.json
**/coverage
test-config
test-config
.vscode

View file

@ -1,3 +1,8 @@
codecov:
require_ci_to_pass: false # post coverage status even if CI has unrelated failures
notify:
wait_for_ci: false # post as soon as expected uploads arrive, don't wait on CI
component_management:
individual_components:
- component_id: "Router"
@ -28,7 +33,7 @@ coverage:
project:
default:
target: auto
threshold: 1% # at maximum allow project coverage to drop by 1%
threshold: 0% # do not allow project coverage to drop
patch:
default:
target: auto

View file

@ -414,6 +414,9 @@ custom_prometheus_metadata_labels: List[str] = []
custom_prometheus_tags: List[str] = []
prometheus_metrics_config: Optional[List] = None
prometheus_emit_stream_label: bool = False
prometheus_end_user_metrics_max_series_per_metric: Optional[int] = 10000
prometheus_end_user_metrics_ttl_seconds: Optional[float] = 3600.0
prometheus_end_user_metrics_cleanup_interval_seconds: Optional[float] = 60.0
disable_add_prefix_to_prompt: bool = (
False # used by anthropic, to disable adding prefix to prompt
)

View file

@ -1457,6 +1457,12 @@ KEY_ROTATION_JOB_NAME = "litellm_key_rotation_job"
EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME = "litellm_expired_ui_session_key_cleanup_job"
SPEND_LOG_RUN_LOOPS = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500))
SPEND_LOG_CLEANUP_BATCH_SIZE = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000))
SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(
os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3)
)
SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS = float(
os.getenv("SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.5)
)
SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100))
SPEND_LOG_QUEUE_POLL_INTERVAL = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0))
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = int(

View file

@ -14,16 +14,18 @@ For batching specific details see CustomBatchLogger class
import asyncio
import os
import time
import traceback
from typing import List, Optional
from typing import List, Optional, Union
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.utils import StandardLoggingPayload
from litellm.types.utils import StandardAuditLogPayload, StandardLoggingPayload
class AzureSentinelLogger(CustomBatchLogger):
@ -39,6 +41,7 @@ class AzureSentinelLogger(CustomBatchLogger):
tenant_id: Optional[str] = None,
client_id: Optional[str] = None,
client_secret: Optional[str] = None,
audit_stream_name: Optional[str] = None,
**kwargs,
):
"""
@ -57,57 +60,77 @@ class AzureSentinelLogger(CustomBatchLogger):
If not provided, will use AZURE_SENTINEL_CLIENT_ID or AZURE_CLIENT_ID env var.
client_secret (str, optional): Azure Client Secret for OAuth2 authentication.
If not provided, will use AZURE_SENTINEL_CLIENT_SECRET or AZURE_CLIENT_SECRET env var.
audit_stream_name (str, optional): Stream name from DCR for audit logs.
If not provided, audit logs use the standard stream name.
"""
self.async_httpx_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
self.dcr_immutable_id = dcr_immutable_id or os.getenv(
resolved_dcr_immutable_id = dcr_immutable_id or os.getenv(
"AZURE_SENTINEL_DCR_IMMUTABLE_ID"
)
self.stream_name = stream_name or os.getenv(
"AZURE_SENTINEL_STREAM_NAME", "Custom-LiteLLM"
resolved_stream_name = (
stream_name or os.getenv("AZURE_SENTINEL_STREAM_NAME") or "Custom-LiteLLM"
)
self.endpoint = endpoint or os.getenv("AZURE_SENTINEL_ENDPOINT")
self.tenant_id = (
resolved_audit_stream_name = audit_stream_name or resolved_stream_name
resolved_endpoint = endpoint or os.getenv("AZURE_SENTINEL_ENDPOINT")
resolved_tenant_id = (
tenant_id
or os.getenv("AZURE_SENTINEL_TENANT_ID")
or os.getenv("AZURE_TENANT_ID")
)
self.client_id = (
resolved_client_id = (
client_id
or os.getenv("AZURE_SENTINEL_CLIENT_ID")
or os.getenv("AZURE_CLIENT_ID")
)
self.client_secret = (
resolved_client_secret = (
client_secret
or os.getenv("AZURE_SENTINEL_CLIENT_SECRET")
or os.getenv("AZURE_CLIENT_SECRET")
)
if not self.dcr_immutable_id:
if not resolved_dcr_immutable_id:
raise ValueError(
"AZURE_SENTINEL_DCR_IMMUTABLE_ID is required. Set it as an environment variable or pass dcr_immutable_id parameter."
)
if not self.endpoint:
if not resolved_endpoint:
raise ValueError(
"AZURE_SENTINEL_ENDPOINT is required. Set it as an environment variable or pass endpoint parameter."
)
if not self.tenant_id:
if not resolved_tenant_id:
raise ValueError(
"AZURE_SENTINEL_TENANT_ID or AZURE_TENANT_ID is required. Set it as an environment variable or pass tenant_id parameter."
)
if not self.client_id:
if not resolved_client_id:
raise ValueError(
"AZURE_SENTINEL_CLIENT_ID or AZURE_CLIENT_ID is required. Set it as an environment variable or pass client_id parameter."
)
if not self.client_secret:
if not resolved_client_secret:
raise ValueError(
"AZURE_SENTINEL_CLIENT_SECRET or AZURE_CLIENT_SECRET is required. Set it as an environment variable or pass client_secret parameter."
)
self.dcr_immutable_id = resolved_dcr_immutable_id
self.stream_name = resolved_stream_name
self.audit_stream_name = resolved_audit_stream_name
self.endpoint = resolved_endpoint
self.tenant_id = resolved_tenant_id
self.client_id = resolved_client_id
self.client_secret = resolved_client_secret
# Build API endpoint: {Endpoint}/dataCollectionRules/{DCR Immutable ID}/streams/{Stream Name}?api-version=2023-01-01
self.api_endpoint = f"{self.endpoint.rstrip('/')}/dataCollectionRules/{self.dcr_immutable_id}/streams/{self.stream_name}?api-version=2023-01-01"
self.api_endpoint = self._build_api_endpoint(
endpoint=resolved_endpoint,
dcr_immutable_id=resolved_dcr_immutable_id,
stream_name=resolved_stream_name,
)
self.audit_api_endpoint = self._build_api_endpoint(
endpoint=resolved_endpoint,
dcr_immutable_id=resolved_dcr_immutable_id,
stream_name=resolved_audit_stream_name,
)
# OAuth2 scope for Azure Monitor
self.oauth_scope = "https://monitor.azure.com/.default"
@ -118,6 +141,13 @@ class AzureSentinelLogger(CustomBatchLogger):
super().__init__(**kwargs, flush_lock=self.flush_lock)
asyncio.create_task(self.periodic_flush())
self.log_queue: List[StandardLoggingPayload] = []
self.audit_log_queue: List[StandardAuditLogPayload] = []
@staticmethod
def _build_api_endpoint(
endpoint: str, dcr_immutable_id: str, stream_name: str
) -> str:
return f"{endpoint.rstrip('/')}/dataCollectionRules/{dcr_immutable_id}/streams/{stream_name}?api-version=2023-01-01"
async def _get_oauth_token(self) -> str:
"""
@ -126,9 +156,6 @@ class AzureSentinelLogger(CustomBatchLogger):
Returns:
Bearer token string
"""
# Check if we have a valid cached token
import time
if (
self.oauth_token
and self.oauth_token_expires_at
@ -170,9 +197,6 @@ class AzureSentinelLogger(CustomBatchLogger):
if not self.oauth_token:
raise Exception("OAuth2 token response did not contain access_token")
# Cache token expiry time
import time
self.oauth_token_expires_at = time.time() + expires_in
return self.oauth_token
@ -246,6 +270,34 @@ class AzureSentinelLogger(CustomBatchLogger):
)
pass
async def async_log_audit_log_event(
self, audit_log: StandardAuditLogPayload
) -> None:
"""
Async log LiteLLM audit log events to Azure Sentinel.
Audit logs are queued separately from standard LLM logs so mixed callback
usage never sends schema-mismatched records in the same ingestion batch.
"""
try:
verbose_logger.debug(
"Azure Sentinel: Logging audit event id=%s action=%s table=%s",
audit_log.get("id"),
audit_log.get("action"),
audit_log.get("table_name"),
)
self.audit_log_queue.append(audit_log)
if len(self.audit_log_queue) >= self.batch_size:
await self.async_send_audit_batch()
except Exception as e:
verbose_logger.exception(
f"Azure Sentinel Audit Log Layer Error - {str(e)}\n{traceback.format_exc()}"
)
pass
async def async_send_batch(self):
"""
Sends the batch of logs to Azure Monitor Logs Ingestion API
@ -253,22 +305,42 @@ class AzureSentinelLogger(CustomBatchLogger):
Raises:
Raises a NON Blocking verbose_logger.exception if an error occurs
"""
await self._async_send_batch_to_api(
log_queue=self.log_queue,
api_endpoint=self.api_endpoint,
log_type="logs",
)
async def async_send_audit_batch(self):
"""
Sends the batch of audit logs to Azure Monitor Logs Ingestion API
"""
await self._async_send_batch_to_api(
log_queue=self.audit_log_queue,
api_endpoint=self.audit_api_endpoint,
log_type="audit logs",
)
async def _async_send_batch_to_api(
self,
log_queue: List[Union[StandardLoggingPayload, StandardAuditLogPayload]],
api_endpoint: str,
log_type: str,
) -> None:
try:
if not self.log_queue:
if not log_queue:
return
verbose_logger.debug(
"Azure Sentinel - about to flush %s events", len(self.log_queue)
"Azure Sentinel - about to flush %s %s", len(log_queue), log_type
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
# Get OAuth2 token
bearer_token = await self._get_oauth_token()
# Convert log queue to JSON array format expected by Logs Ingestion API
# Each log entry should be a JSON object in the array
body = safe_dumps(self.log_queue)
body = safe_dumps(log_queue)
# Set headers for Logs Ingestion API
headers = {
@ -278,7 +350,7 @@ class AzureSentinelLogger(CustomBatchLogger):
# Send the request
response = await self.async_httpx_client.post(
url=self.api_endpoint, data=body.encode("utf-8"), headers=headers
url=api_endpoint, data=body.encode("utf-8"), headers=headers
)
if response.status_code not in [200, 204]:
@ -301,4 +373,15 @@ class AzureSentinelLogger(CustomBatchLogger):
f"Azure Sentinel Error sending batch API - {str(e)}\n{traceback.format_exc()}"
)
finally:
self.log_queue.clear()
log_queue.clear()
async def flush_queue(self):
if self.flush_lock is None:
return
async with self.flush_lock:
if self.log_queue:
await self.async_send_batch()
if self.audit_log_queue:
await self.async_send_audit_batch()
self.last_flush_time = time.time()

View file

@ -25,6 +25,9 @@ from typing import (
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import (
BoundedPrometheusSeriesTracker,
)
from litellm.integrations.prometheus_helpers import (
PrometheusLabelFactoryContext,
_get_cached_end_user_id_for_cost_tracking,
@ -81,6 +84,7 @@ class PrometheusLogger(CustomLogger):
if _custom_buckets is not None
else LATENCY_BUCKETS
)
self._bounded_prometheus_series_tracker = BoundedPrometheusSeriesTracker()
# Create metric factory functions
self._counter_factory = self._create_metric_factory(Counter)
@ -984,6 +988,40 @@ class PrometheusLogger(CustomLogger):
return filtered_labels
def _track_end_user_metric_series(
self,
metric: Any,
metric_name: DEFINED_PROMETHEUS_METRICS,
labels: Dict[str, Optional[str]],
) -> None:
"""
Cap the cardinality of metrics that include the ``end_user`` label.
Called *after* ``metric.labels(...).inc()/observe()`` so the emission is
recorded in prometheus-client's child map before any eviction runs.
Series that get evicted before the next scrape lose updates accrued
since the last scrape — this is inherent to any cardinality cap.
"""
labelnames = self.get_labels_for_metric(metric_name)
if UserAPIKeyLabelNames.END_USER.value not in labelnames:
return
if labels.get(UserAPIKeyLabelNames.END_USER.value) is None:
return
max_series = litellm.prometheus_end_user_metrics_max_series_per_metric
ttl_seconds = litellm.prometheus_end_user_metrics_ttl_seconds
if max_series is None and ttl_seconds is None:
return
self._bounded_prometheus_series_tracker.track_series(
metric=metric,
metric_name=metric_name,
label_values=tuple(labels.get(label) for label in labelnames),
max_series=max_series,
ttl_seconds=ttl_seconds,
cleanup_interval_seconds=litellm.prometheus_end_user_metrics_cleanup_interval_seconds,
)
def _inc_labeled_counter(
self,
counter: Any,
@ -998,6 +1036,7 @@ class PrometheusLogger(CustomLogger):
label_context=label_context,
)
counter.labels(**_labels).inc(amount)
self._track_end_user_metric_series(counter, metric_name, _labels)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
# Define prometheus client
@ -1404,12 +1443,12 @@ class PrometheusLogger(CustomLogger):
)
remaining_tokens_variable_name = f"litellm-key-remaining-tokens-{model_group}"
remaining_requests = (
metadata.get(remaining_requests_variable_name, sys.maxsize) or sys.maxsize
)
remaining_tokens = (
metadata.get(remaining_tokens_variable_name, sys.maxsize) or sys.maxsize
)
remaining_requests = metadata.get(remaining_requests_variable_name)
if remaining_requests is None:
remaining_requests = sys.maxsize
remaining_tokens = metadata.get(remaining_tokens_variable_name)
if remaining_tokens is None:
remaining_tokens = sys.maxsize
enum_values = UserAPIKeyLabelValues(
hashed_api_key=user_api_key,
@ -1479,6 +1518,11 @@ class PrometheusLogger(CustomLogger):
self.litellm_llm_api_time_to_first_token_metric.labels(
**_ttft_labels
).observe(time_to_first_token_seconds)
self._track_end_user_metric_series(
self.litellm_llm_api_time_to_first_token_metric,
"litellm_llm_api_time_to_first_token_metric",
_ttft_labels,
)
else:
verbose_logger.debug(
"Time to first token metric not emitted, stream option in model_parameters is not True"
@ -1499,6 +1543,11 @@ class PrometheusLogger(CustomLogger):
self.litellm_llm_api_latency_metric.labels(**_labels).observe(
api_call_total_time_seconds
)
self._track_end_user_metric_series(
self.litellm_llm_api_latency_metric,
"litellm_llm_api_latency_metric",
_labels,
)
# total request latency
total_time_seconds = self._safe_duration_seconds(
@ -1516,6 +1565,11 @@ class PrometheusLogger(CustomLogger):
self.litellm_request_total_latency_metric.labels(**_labels).observe(
total_time_seconds
)
self._track_end_user_metric_series(
self.litellm_request_total_latency_metric,
"litellm_request_total_latency_metric",
_labels,
)
# request queue time (time from arrival to processing start)
_litellm_params = kwargs.get("litellm_params", {}) or {}
@ -1533,6 +1587,11 @@ class PrometheusLogger(CustomLogger):
self.litellm_request_queue_time_metric.labels(**_labels).observe(
queue_time_seconds
)
self._track_end_user_metric_series(
self.litellm_request_queue_time_metric,
"litellm_request_queue_time_seconds",
_labels,
)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
verbose_logger.debug(

View file

@ -0,0 +1,107 @@
from __future__ import annotations
import time
from collections import OrderedDict
from threading import RLock
from typing import Any, Dict, Optional
class BoundedPrometheusSeriesTracker:
"""
Tracks Prometheus child series and removes stale/excess labelsets.
The tracker is label-agnostic: callers decide which series should be tracked
and pass the full label tuple used by the Prometheus metric.
"""
def __init__(self) -> None:
self._series: Dict[str, OrderedDict[tuple[Optional[str], ...], float]] = {}
self._last_ttl_cleanup: Dict[str, float] = {}
self.lock = RLock()
def track_series(
self,
metric: Any,
metric_name: str,
label_values: tuple[Optional[str], ...],
max_series: Optional[int],
ttl_seconds: Optional[float],
cleanup_interval_seconds: Optional[float],
) -> None:
if max_series is None and ttl_seconds is None:
return
now = time.monotonic()
with self.lock:
series = self._series.setdefault(metric_name, OrderedDict())
series[label_values] = now
series.move_to_end(label_values)
if ttl_seconds is not None and self._should_run_ttl_cleanup(
metric_name=metric_name,
now=now,
cleanup_interval_seconds=cleanup_interval_seconds,
):
expired_label_values = [
tracked_label_values
for tracked_label_values, last_seen in series.items()
if now - last_seen > ttl_seconds
]
for tracked_label_values in expired_label_values:
self._remove_metric_series(metric, series, tracked_label_values)
# max_series <= 0 is treated as "unlimited" so a misconfigured zero
# value cannot silently drop every emission for this metric.
if max_series is not None and max_series > 0:
while len(series) > max_series:
tracked_label_values = next(iter(series))
if not self._remove_metric_child(metric, tracked_label_values):
break
del series[tracked_label_values]
def _should_run_ttl_cleanup(
self,
metric_name: str,
now: float,
cleanup_interval_seconds: Optional[float],
) -> bool:
if cleanup_interval_seconds is None or cleanup_interval_seconds <= 0:
self._last_ttl_cleanup[metric_name] = now
return True
last_cleanup = self._last_ttl_cleanup.get(metric_name)
if last_cleanup is None or now - last_cleanup >= cleanup_interval_seconds:
self._last_ttl_cleanup[metric_name] = now
return True
return False
def _remove_metric_series(
self,
metric: Any,
series: OrderedDict[tuple[Optional[str], ...], float],
label_values: tuple[Optional[str], ...],
) -> None:
if self._remove_metric_child(metric, label_values):
series.pop(label_values, None)
@staticmethod
def _remove_metric_child(
metric: Any, label_values: tuple[Optional[str], ...]
) -> bool:
"""
Remove the Prometheus child for ``label_values`` and report whether the
tracker should commit the matching state change.
Returns ``True`` when the child is no longer present in Prometheus
(either it was just removed or it was already gone), and ``False`` when
``metric.remove()`` raised an unexpected error and the child likely
still exists.
"""
try:
metric.remove(*label_values)
return True
except KeyError:
return True
except (AttributeError, ValueError):
return False

View file

@ -14,6 +14,7 @@ from litellm import _custom_logger_compatible_callbacks_literal
from litellm.integrations.agentops import AgentOps
from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook
from litellm.integrations.argilla import ArgillaLogger
from litellm.integrations.azure_sentinel.azure_sentinel import AzureSentinelLogger
from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLogger
from litellm.integrations.bitbucket import BitBucketPromptManager
from litellm.integrations.braintrust_logging import BraintrustLogger
@ -73,6 +74,7 @@ class CustomLoggerRegistry:
"opik": OpikLogger,
"argilla": ArgillaLogger,
"opentelemetry": OpenTelemetry,
"azure_sentinel": AzureSentinelLogger,
"azure_storage": AzureBlobStorageLogger,
"humanloop": HumanloopLogger,
# OTEL compatible loggers

View file

@ -579,6 +579,10 @@ class ModelResponseIterator:
# Accumulate compaction blocks for multi-turn reconstruction
self.compaction_blocks: List[Dict[str, Any]] = []
# Accumulate streamed thinking text so final usage can split reasoning
# tokens from regular output tokens.
self.reasoning_content_chunks: List[str] = []
# Track server tool use inputs and results for code_interpreter_results
self._server_tool_inputs: Dict[str, Any] = {}
self.tool_results: List[Dict[str, Any]] = []
@ -609,9 +613,14 @@ class ModelResponseIterator:
return False
def _handle_usage(self, anthropic_usage_chunk: Union[dict, UsageDelta]) -> Usage:
reasoning_content = (
"".join(self.reasoning_content_chunks)
if self.reasoning_content_chunks
else None
)
return AnthropicConfig().calculate_usage(
usage_object=cast(dict, anthropic_usage_chunk),
reasoning_content=None,
reasoning_content=reasoning_content,
speed=self.speed,
)
@ -658,10 +667,13 @@ class ModelResponseIterator:
"thinking" in content_block["delta"]
or "signature" in content_block["delta"]
):
thinking_content = content_block["delta"].get("thinking")
if isinstance(thinking_content, str) and thinking_content:
self.reasoning_content_chunks.append(thinking_content)
thinking_blocks = [
ChatCompletionThinkingBlock(
type="thinking",
thinking=content_block["delta"].get("thinking") or "",
thinking=thinking_content or "",
signature=str(content_block["delta"].get("signature") or ""),
)
]

View file

@ -2156,8 +2156,16 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
speed: Optional[str] = None,
) -> Usage:
# NOTE: Sometimes the usage object has None set explicitly for token counts, meaning .get() & key access returns None, and we need to account for this
prompt_tokens = usage_object.get("input_tokens", 0) or 0
completion_tokens = usage_object.get("output_tokens", 0) or 0
raw_prompt_tokens = usage_object.get("input_tokens", 0) or 0
prompt_tokens: int = (
int(raw_prompt_tokens) if isinstance(raw_prompt_tokens, (int, float)) else 0
)
raw_completion_tokens = usage_object.get("output_tokens", 0) or 0
completion_tokens: int = (
int(raw_completion_tokens)
if isinstance(raw_completion_tokens, (int, float))
else 0
)
_usage = usage_object
cache_creation_input_tokens: int = 0
cache_read_input_tokens: int = 0
@ -2226,11 +2234,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
text_tokens=raw_input_tokens,
)
# Always populate completion_token_details, not just when there's reasoning_content
reasoning_tokens = (
estimated_reasoning_tokens = (
token_counter(text=reasoning_content, count_response_tokens=True)
if reasoning_content
else 0
)
reasoning_tokens = min(estimated_reasoning_tokens, completion_tokens)
completion_token_details = CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens if reasoning_tokens > 0 else 0,
text_tokens=(

View file

@ -299,29 +299,9 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
)
def _get_response_stream_shape(self):
"""Get the response stream shape for parsing, reusing existing logic."""
try:
# Try to reuse the cached shape from the existing decoder
from litellm.llms.bedrock.chat.invoke_handler import (
get_response_stream_shape,
)
from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE
return get_response_stream_shape()
except ImportError:
# Fallback: create our own shape
try:
from botocore.loaders import Loader
from botocore.model import ServiceModel
loader = Loader()
bedrock_service_dict = loader.load_service_model(
"bedrock-runtime", "service-2"
)
bedrock_service_model = ServiceModel(bedrock_service_dict)
return bedrock_service_model.shape_for("ResponseStream")
except Exception as e:
verbose_logger.warning(f"Could not load response stream shape: {e}")
return None
return BEDROCK_RESPONSE_STREAM_SHAPE
def _extract_response_content(self, events: InvokeAgentEventList) -> str:
"""Extract the final response content from parsed events."""

View file

@ -67,9 +67,13 @@ from litellm.types.utils import (
from litellm.utils import CustomStreamWrapper, get_secret
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError, ModelResponseIterator, get_bedrock_tool_name
from ..common_utils import (
BEDROCK_RESPONSE_STREAM_SHAPE,
BedrockError,
ModelResponseIterator,
get_bedrock_tool_name,
)
_response_stream_shape_cache = None
bedrock_tool_name_mappings: InMemoryCache = InMemoryCache(
max_size_in_memory=50, default_ttl=600
)
@ -1391,20 +1395,6 @@ class BedrockLLM(BaseAWSLLM):
return None
def get_response_stream_shape():
global _response_stream_shape_cache
if _response_stream_shape_cache is None:
from botocore.loaders import Loader
from botocore.model import ServiceModel
loader = Loader()
bedrock_service_dict = loader.load_service_model("bedrock-runtime", "service-2")
bedrock_service_model = ServiceModel(bedrock_service_dict)
_response_stream_shape_cache = bedrock_service_model.shape_for("ResponseStream")
return _response_stream_shape_cache
class AWSEventStreamDecoder:
def __init__(self, model: str, json_mode: Optional[bool] = False) -> None:
from botocore.parsers import EventStreamJSONParser
@ -1838,8 +1828,18 @@ class AWSEventStreamDecoder:
yield self._chunk_parser(chunk_data=_data)
def _parse_message_from_event(self, event) -> Optional[str]:
if BEDROCK_RESPONSE_STREAM_SHAPE is None:
raise BedrockError(
status_code=500,
message=(
"Bedrock event-stream shape could not be loaded from botocore. "
"Ensure botocore is correctly installed."
),
)
response_dict = event.to_response_dict()
parsed_response = self.parser.parse(response_dict, get_response_stream_shape())
parsed_response = self.parser.parse(
response_dict, BEDROCK_RESPONSE_STREAM_SHAPE
)
if response_dict["status_code"] != 200:
decoded_body = response_dict["body"].decode()

View file

@ -14,6 +14,7 @@ if TYPE_CHECKING:
import httpx
import litellm
from litellm import verbose_logger
from litellm.llms.base_llm.anthropic_messages.transformation import (
BaseAnthropicMessagesConfig,
)
@ -917,38 +918,57 @@ def get_bedrock_chat_config(model: str):
return litellm.AmazonInvokeConfig()
def _load_bedrock_response_stream_shape():
"""
Load the ResponseStream shape from botocore's bundled bedrock-runtime schema.
Called once at module import time; the result is stored in
``BEDROCK_RESPONSE_STREAM_SHAPE`` and reused for the process lifetime.
Returns ``None`` if botocore is unavailable or the service model cannot be
loaded, so the module still imports cleanly.
"""
try:
from botocore.loaders import Loader
from botocore.model import ServiceModel
loader = Loader()
service_dict = loader.load_service_model("bedrock-runtime", "service-2")
return ServiceModel(service_dict).shape_for("ResponseStream")
except Exception as e:
verbose_logger.warning(
"litellm: could not pre-load bedrock-runtime response stream shape "
"— Bedrock event-stream decoding will be unavailable. Error: %s",
e,
)
return None
# Eagerly resolved once per process — avoids per-instance or per-request disk I/O.
BEDROCK_RESPONSE_STREAM_SHAPE = _load_bedrock_response_stream_shape()
class BedrockEventStreamDecoderBase:
"""
Base class for event stream decoding for Bedrock
"""
_response_stream_shape_cache = None
def __init__(self):
from botocore.parsers import EventStreamJSONParser
self.parser = EventStreamJSONParser()
def get_response_stream_shape(self):
if self._response_stream_shape_cache is None:
from botocore.loaders import Loader
from botocore.model import ServiceModel
loader = Loader()
bedrock_service_dict = loader.load_service_model(
"bedrock-runtime", "service-2"
)
bedrock_service_model = ServiceModel(bedrock_service_dict)
self._response_stream_shape_cache = bedrock_service_model.shape_for(
"ResponseStream"
)
return self._response_stream_shape_cache
def _parse_message_from_event(self, event) -> Optional[str]:
if BEDROCK_RESPONSE_STREAM_SHAPE is None:
raise BedrockError(
status_code=500,
message=(
"Bedrock event-stream shape could not be loaded from botocore. "
"Ensure botocore is correctly installed."
),
)
response_dict = event.to_response_dict()
parsed_response = self.parser.parse(
response_dict, self.get_response_stream_shape()
response_dict, BEDROCK_RESPONSE_STREAM_SHAPE
)
if response_dict["status_code"] != 200:

View file

@ -1,4 +1,5 @@
import asyncio
import concurrent.futures
import inspect
import os
import socket
@ -133,6 +134,11 @@ _DEFAULT_TIMEOUT = httpx.Timeout(
timeout=COMPLETION_HTTP_FALLBACK_SECONDS,
connect=HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS,
)
_STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS = 5.0
_STREAMING_ERROR_BODY_READ_EXECUTOR = concurrent.futures.ThreadPoolExecutor(
max_workers=50,
thread_name_prefix="litellm-streaming-error-body-read",
)
def _prepare_request_data_and_content(
@ -386,17 +392,30 @@ def _safe_get_response_text(response: httpx.Response) -> str:
return ""
async def _safe_aread_response(response: httpx.Response) -> bytes:
async def _safe_aread_response(
response: httpx.Response, timeout: Optional[float] = None
) -> bytes:
"""Safely read async response body, falling back to empty bytes on errors."""
try:
if timeout is not None:
return await asyncio.wait_for(response.aread(), timeout=timeout)
return await response.aread()
except Exception:
return b""
def _safe_read_response(response: httpx.Response) -> bytes:
def _safe_read_response(
response: httpx.Response, timeout: Optional[float] = None
) -> bytes:
"""Safely read sync response body, falling back to empty bytes on errors."""
try:
if timeout is not None:
future = _STREAMING_ERROR_BODY_READ_EXECUTOR.submit(response.read)
try:
return future.result(timeout=timeout)
except Exception:
response.close()
return b""
return response.read()
except Exception:
return b""
@ -405,8 +424,19 @@ def _safe_read_response(response: httpx.Response) -> bytes:
def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> None:
"""Raise a MaskedHTTPStatusError for sync HTTP handlers."""
if stream:
_body = mask_sensitive_info(_safe_read_response(e.response))
raise MaskedHTTPStatusError(e, message=_body, text=_body) from None
try:
_body = mask_sensitive_info(
_safe_read_response(
e.response,
timeout=_STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS,
)
)
raise MaskedHTTPStatusError(e, message=_body, text=_body) from None
finally:
try:
e.response.close()
except Exception:
pass
_text = mask_sensitive_info(_safe_get_response_text(e.response))
raise MaskedHTTPStatusError(e, message=_text, text=_text) from None
@ -414,8 +444,19 @@ def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> None:
async def _raise_masked_async_error(e: httpx.HTTPStatusError, stream: bool) -> None:
"""Raise a MaskedHTTPStatusError for async HTTP handlers."""
if stream:
_body = mask_sensitive_info(await _safe_aread_response(e.response))
raise MaskedHTTPStatusError(e, message=_body, text=_body) from None
try:
_body = mask_sensitive_info(
await _safe_aread_response(
e.response,
timeout=_STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS,
)
)
raise MaskedHTTPStatusError(e, message=_body, text=_body) from None
finally:
try:
await e.response.aclose()
except Exception:
pass
_text = mask_sensitive_info(_safe_get_response_text(e.response))
raise MaskedHTTPStatusError(e, message=_text, text=_text) from None

View file

@ -9,7 +9,27 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.utils import GenericStreamingChunk as GChunk
from litellm.types.utils import StreamingChatCompletionChunk
_response_stream_shape_cache = None
def _load_sagemaker_response_stream_shape():
try:
from botocore.loaders import Loader
from botocore.model import ServiceModel
loader = Loader()
service_dict = loader.load_service_model("sagemaker-runtime", "service-2")
return ServiceModel(service_dict).shape_for(
"InvokeEndpointWithResponseStreamOutput"
)
except Exception as e:
verbose_logger.warning(
"litellm: could not pre-load sagemaker-runtime response stream shape "
"— SageMaker event-stream decoding will be unavailable. Error: %s",
e,
)
return None
SAGEMAKER_RESPONSE_STREAM_SHAPE = _load_sagemaker_response_stream_shape()
class SagemakerError(BaseLLMException):
@ -187,8 +207,18 @@ class AWSEventStreamDecoder:
verbose_logger.error(f"Final error parsing accumulated JSON: {e}")
def _parse_message_from_event(self, event) -> Optional[str]:
if SAGEMAKER_RESPONSE_STREAM_SHAPE is None:
raise SagemakerError(
status_code=500,
message=(
"SageMaker event-stream shape could not be loaded from botocore. "
"Ensure botocore is correctly installed."
),
)
response_dict = event.to_response_dict()
parsed_response = self.parser.parse(response_dict, get_response_stream_shape())
parsed_response = self.parser.parse(
response_dict, SAGEMAKER_RESPONSE_STREAM_SHAPE
)
if response_dict["status_code"] != 200:
raise ValueError(f"Bad response code, expected 200: {response_dict}")
@ -204,20 +234,3 @@ class AWSEventStreamDecoder:
return None
return chunk.decode() # type: ignore[no-any-return]
def get_response_stream_shape():
global _response_stream_shape_cache
if _response_stream_shape_cache is None:
from botocore.loaders import Loader
from botocore.model import ServiceModel
loader = Loader()
sagemaker_service_dict = loader.load_service_model(
"sagemaker-runtime", "service-2"
)
sagemaker_service_model = ServiceModel(sagemaker_service_dict)
_response_stream_shape_cache = sagemaker_service_model.shape_for(
"InvokeEndpointWithResponseStreamOutput"
)
return _response_stream_shape_cache

View file

@ -28874,6 +28874,19 @@
"mode": "chat",
"output_cost_per_token": 0.0
},
"sambanova/MiniMax-M2.7": {
"input_cost_per_token": 3e-07,
"litellm_provider": "sambanova",
"max_input_tokens": 204800,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"source": "https://cloud.sambanova.ai/plans/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"sambanova/DeepSeek-R1": {
"input_cost_per_token": 5e-06,
"litellm_provider": "sambanova",

View file

@ -768,7 +768,9 @@ class MCPServerManager:
)
return new_server
async def _maybe_register_openapi_tools(self, server: MCPServer):
async def _maybe_register_openapi_tools(
self, server: MCPServer, *, initialize_mapping: bool = True
):
"""Register OpenAPI tools if the server has a spec_path configured."""
if server.spec_path:
verbose_logger.info(
@ -779,7 +781,8 @@ class MCPServerManager:
server=server,
base_url=server.url or "",
)
self.initialize_tool_name_to_mcp_server_name_mapping()
if initialize_mapping:
self.initialize_tool_name_to_mcp_server_name_mapping()
async def add_server(self, mcp_server: LiteLLM_MCPServerTable):
try:
@ -1978,7 +1981,11 @@ class MCPServerManager:
_SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024
def _assign_unique_short_prefix(self, server: MCPServer) -> None:
def _assign_unique_short_prefix(
self,
server: MCPServer,
registry: Optional[Dict[str, MCPServer]] = None,
) -> None:
"""Resolve and cache a collision-free short tool prefix on ``server``.
Called at registration time for every MCP server entering the
@ -2002,7 +2009,8 @@ class MCPServerManager:
return
used: Dict[str, str] = {}
for other in self.get_registry().values():
registry_for_collision_check = registry or self.get_registry()
for other in registry_for_collision_check.values():
if other.server_id == server.server_id:
continue
if other.short_prefix:
@ -2916,46 +2924,72 @@ class MCPServerManager:
# against the *full* set so dedup is deterministic regardless of
# iteration order.
for server in db_mcp_servers:
existing_server = previous_registry.get(server.server_id)
try:
existing_server = previous_registry.get(server.server_id)
if (
existing_server is not None
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
new_registry[server.server_id] = existing_server
continue
if (
existing_server is not None
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
new_registry[server.server_id] = existing_server
continue
_warn_on_server_name_fields(
server_id=server.server_id,
alias=getattr(server, "alias", None),
server_name=getattr(server, "server_name", None),
)
verbose_logger.debug(
f"Building server from DB: {server.server_id} ({server.server_name})"
)
new_server = await self.build_mcp_server_from_table(server)
# Carry the cached short_prefix from the previous registry entry
# (if any) so the prefix is stable across reloads.
if existing_server is not None and existing_server.short_prefix:
new_server.short_prefix = existing_server.short_prefix
new_registry[server.server_id] = new_server
_warn_on_server_name_fields(
server_id=server.server_id,
alias=getattr(server, "alias", None),
server_name=getattr(server, "server_name", None),
)
verbose_logger.debug(
f"Building server from DB: {server.server_id} ({server.server_name})"
)
new_server = await self.build_mcp_server_from_table(server)
# Carry the cached short_prefix from the previous registry entry
# (if any) so the prefix is stable across reloads.
if existing_server is not None and existing_server.short_prefix:
new_server.short_prefix = existing_server.short_prefix
new_registry[server.server_id] = new_server
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
server.server_id,
getattr(server, "alias", None),
e,
)
# Swap in the new registry first so _assign_unique_short_prefix
# sees the complete set when checking for collisions.
self.registry = new_registry
for new_server in new_registry.values():
self._assign_unique_short_prefix(new_server)
# Register OpenAPI tools *after* the final short prefix is assigned
# so the tools are stored in the global registry under the same
# prefix that lookups will use.
await self._maybe_register_openapi_tools(new_server)
# Assign short prefixes against the full candidate set without
# publishing the staged registry to concurrent callers.
registered_registry: Dict[str, MCPServer] = {}
registered_openapi_tools = False
for server_id, new_server in new_registry.items():
try:
self._assign_unique_short_prefix(new_server, registry=new_registry)
# Register OpenAPI tools *after* the final short prefix is assigned
# so the tools are stored in the global registry under the same
# prefix that lookups will use.
await self._maybe_register_openapi_tools(
new_server, initialize_mapping=False
)
registered_registry[server_id] = new_server
if new_server.spec_path:
registered_openapi_tools = True
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
new_server.server_id,
getattr(new_server, "alias", None),
e,
)
self.registry = registered_registry
if registered_openapi_tools:
self.initialize_tool_name_to_mcp_server_name_mapping()
verbose_logger.debug(
"MCP registry refreshed (%s servers in registry)", len(new_registry)
"MCP registry refreshed (%s servers in registry)", len(registered_registry)
)
def get_mcp_servers_from_ids(self, server_ids: List[str]) -> List[MCPServer]:

View file

@ -3348,6 +3348,19 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
],
)
azure_sentinel: CallbackOnUI = CallbackOnUI(
litellm_callback_name="azure_sentinel",
ui_callback_name="Azure Sentinel",
litellm_callback_params=[
"AZURE_SENTINEL_DCR_IMMUTABLE_ID",
"AZURE_SENTINEL_ENDPOINT",
"AZURE_SENTINEL_TENANT_ID",
"AZURE_SENTINEL_CLIENT_ID",
"AZURE_SENTINEL_CLIENT_SECRET",
"AZURE_SENTINEL_STREAM_NAME",
],
)
openmeter: CallbackOnUI = CallbackOnUI(
litellm_callback_name="openmeter",
ui_callback_name="OpenMeter",

View file

@ -3516,7 +3516,6 @@ async def _check_team_member_budget(
if (
team_object is not None
and team_object.team_id is not None
and user_object is not None
and valid_token is not None
and valid_token.user_id is not None
):
@ -3619,6 +3618,7 @@ async def _check_team_member_model_access(
llm_router=llm_router,
models=member_allowed_models,
object_type="team",
team_id=team_object.team_id,
)
except ProxyException:
raise ProxyException(

View file

@ -5,8 +5,10 @@ from typing import Optional
from litellm._logging import verbose_proxy_logger
from litellm.caching import RedisCache
from litellm.constants import (
SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS,
SPEND_LOG_CLEANUP_BATCH_SIZE,
SPEND_LOG_CLEANUP_JOB_NAME,
SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES,
SPEND_LOG_RUN_LOOPS,
)
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
@ -74,6 +76,7 @@ class SpendLogCleanup:
"""
total_deleted = 0
run_count = 0
consecutive_failures = 0
while True:
if run_count > SPEND_LOG_RUN_LOOPS:
verbose_proxy_logger.info(
@ -82,18 +85,50 @@ class SpendLogCleanup:
break
# Step 1: Find logs and delete them in one go without fetching to application
# Delete in batches, limited by self.batch_size
deleted_result = await prisma_client.db.execute_raw(
"""
DELETE FROM "LiteLLM_SpendLogs"
WHERE "request_id" IN (
SELECT "request_id" FROM "LiteLLM_SpendLogs"
WHERE "startTime" < $1::timestamptz
LIMIT $2
try:
deleted_result = await prisma_client.db.execute_raw(
"""
DELETE FROM "LiteLLM_SpendLogs"
WHERE "request_id" IN (
SELECT "request_id" FROM "LiteLLM_SpendLogs"
WHERE "startTime" < $1::timestamptz
LIMIT $2
)
""",
cutoff_date,
self.batch_size,
)
""",
cutoff_date,
self.batch_size,
)
except Exception as batch_exc:
# A single batch failure (e.g. Prisma/DB timeout) must not abort
# the whole run — subsequent batches may still succeed.
consecutive_failures += 1
verbose_proxy_logger.exception(
"Spend log cleanup batch failed "
"(run_count=%d, consecutive_failures=%d, batch_size=%d, "
"cutoff=%s, total_deleted_so_far=%d): %s: %s",
run_count,
consecutive_failures,
self.batch_size,
cutoff_date.isoformat(),
total_deleted,
type(batch_exc).__name__,
batch_exc,
)
if (
consecutive_failures
>= SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES
):
verbose_proxy_logger.error(
"Aborting spend log cleanup after %d consecutive batch "
"failures; total deleted before abort: %d",
consecutive_failures,
total_deleted,
)
break
await asyncio.sleep(SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS)
continue
consecutive_failures = 0
deleted_count = 0
if isinstance(deleted_result, int):
@ -168,7 +203,13 @@ class SpendLogCleanup:
verbose_proxy_logger.info(f"Deleted {total_deleted} logs")
except Exception as e:
verbose_proxy_logger.error(f"Error during cleanup: {str(e)}")
# .exception() captures the traceback; str(e) alone on a Prisma/DB
# timeout is often empty and gives operators no signal to diagnose.
verbose_proxy_logger.exception(
"Error during spend log cleanup: %s: %s",
type(e).__name__,
e,
)
return # Return after error handling
finally:
# Only release the lock if it was actually acquired

View file

@ -4,6 +4,7 @@
This is an enterprise feature and requires a premium license.
"""
import re
from typing import Any, Dict, List, Optional, Set, Tuple
from fastapi import (
@ -843,6 +844,18 @@ async def get_service_provider_config(request: Request):
return SCIMServiceProviderConfig(meta=meta)
def _parse_scim_eq_filter(scim_filter: str) -> Optional[Tuple[str, str]]:
"""Parse the SCIM equality filters Okta uses before user lifecycle changes."""
match = re.match(
r"""\s*([\w.]+)\s+eq\s+(['"]?)(.*?)\2\s*$""",
scim_filter,
flags=re.IGNORECASE,
)
if not match:
return None
return match.group(1).lower(), match.group(3)
# User Endpoints
@scim_router.get(
"/Users",
@ -867,15 +880,21 @@ async def get_users(
try:
prisma_client = await _get_prisma_client_or_raise_exception()
# Parse filter if provided (basic support)
where_conditions = {}
where_conditions: Dict[str, Any] = {}
if filter:
# Very basic filter support - only handling userName eq and emails.value eq
if "userName eq" in filter:
user_id = filter.split("userName eq ")[1].strip("\"'")
where_conditions["user_id"] = user_id
elif "emails.value eq" in filter:
email = filter.split("emails.value eq ")[1].strip("\"'")
where_conditions["user_email"] = email
# Okta locates users by userName before deprovisioning. LiteLLM
# exposes SCIM userName from user_email, while older SCIM-created
# users may still have user_id == userName, so support both.
parsed_filter = _parse_scim_eq_filter(filter)
if parsed_filter:
filter_attribute, filter_value = parsed_filter
if filter_attribute == "username":
where_conditions["OR"] = [
{"user_email": filter_value},
{"user_id": filter_value},
]
elif filter_attribute == "emails.value":
where_conditions["user_email"] = filter_value
# Get users from database
users: List[LiteLLM_UserTable] = (

View file

@ -0,0 +1,121 @@
import json
from typing import Callable, Optional, Union
from starlette.types import ASGIApp, Message, Receive, Scope, Send
MaxRequestSizeGetter = Callable[[], Optional[Union[int, float]]]
RequestSizeLimitEnabledGetter = Callable[[], bool]
class RequestEntityTooLarge(Exception):
pass
class RequestSizeLimitMiddleware:
"""
Reject oversized requests before downstream auth/routes parse the body.
Content-Length can be rejected without reading any body bytes. Requests
without Content-Length are counted as the ASGI stream is consumed, limiting
memory exposure to the configured threshold plus the current chunk.
"""
def __init__(
self,
app: ASGIApp,
get_max_request_size_mb: MaxRequestSizeGetter,
is_request_size_limit_enabled: RequestSizeLimitEnabledGetter,
) -> None:
self.app = app
self.get_max_request_size_mb = get_max_request_size_mb
self.is_request_size_limit_enabled = is_request_size_limit_enabled
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] != "http":
await self.app(scope, receive, send)
return
max_request_size_mb = self.get_max_request_size_mb()
max_request_size_bytes = _mb_to_bytes(max_request_size_mb)
if max_request_size_bytes is None or not self.is_request_size_limit_enabled():
await self.app(scope, receive, send)
return
content_length = _get_content_length(scope=scope)
if content_length is not None and content_length > max_request_size_bytes:
await _send_request_too_large(
send=send, max_request_size_mb=max_request_size_mb
)
return
received_body_bytes = 0
response_started = False
async def limited_receive() -> Message:
nonlocal received_body_bytes
message = await receive()
if message["type"] != "http.request":
return message
received_body_bytes += len(message.get("body", b""))
if received_body_bytes > max_request_size_bytes:
raise RequestEntityTooLarge
return message
async def tracking_send(message: Message) -> None:
nonlocal response_started
if message["type"] == "http.response.start":
response_started = True
await send(message)
try:
await self.app(scope, limited_receive, tracking_send)
except RequestEntityTooLarge:
if response_started:
raise
await _send_request_too_large(
send=send, max_request_size_mb=max_request_size_mb
)
def _mb_to_bytes(max_request_size_mb: Optional[Union[int, float]]) -> Optional[int]:
if max_request_size_mb is None:
return None
if max_request_size_mb <= 0:
return None
return int(max_request_size_mb * 1024 * 1024)
def _get_content_length(scope: Scope) -> Optional[int]:
headers = dict(scope.get("headers") or [])
raw_content_length = headers.get(b"content-length")
if raw_content_length is None:
return None
try:
return int(raw_content_length)
except ValueError:
return None
async def _send_request_too_large(
send: Send,
max_request_size_mb: Optional[Union[int, float]],
) -> None:
body = json.dumps(
{"error": f"Request size is too large. Max size is {max_request_size_mb} MB"},
separators=(",", ":"),
).encode("utf-8")
await send(
{
"type": "http.response.start",
"status": 413,
"headers": [
(b"content-type", b"application/json"),
(b"content-length", str(len(body)).encode("latin-1")),
],
}
)
await send({"type": "http.response.body", "body": body, "more_body": False})

View file

@ -169,6 +169,66 @@ class ProxyInitializationHelpers:
)
return uvicorn_args
@staticmethod
def _get_reload_options(config_path: Optional[str]) -> dict:
"""Build uvicorn reload kwargs so --reload also reacts to YAML edits."""
options: dict = {"reload": True}
if not config_path:
return options
config_abs = os.path.abspath(config_path)
config_dir = os.path.dirname(config_abs)
cwd = os.path.abspath(os.getcwd())
reload_dirs = [cwd]
if config_dir and config_dir != cwd:
reload_dirs.append(config_dir)
options["reload_dirs"] = reload_dirs
# Must be a basename, not an absolute path: uvicorn's
# resolve_reload_patterns() calls pathlib.Path.glob(), which raises
# NotImplementedError on absolute patterns (uvicorn discussion #2156).
options["reload_includes"] = ["*.py", os.path.basename(config_abs)]
return options
@staticmethod
def _patch_statreload_for_config(config_path: str) -> bool:
"""Make uvicorn's StatReload reloader notice YAML config changes.
Uvicorn uses WatchFilesReload when the optional `watchfiles` package
is installed, otherwise StatReload. StatReload hard-codes `*.py` in
`iter_py_files()` and silently ignores `reload_includes`, so the
kwargs from `_get_reload_options` alone don't trigger reloads on YAML
edits. We monkey-patch `iter_py_files` to also yield the config path.
Idempotent across calls and a no-op for the WatchFilesReload path.
"""
try:
from uvicorn.supervisors.statreload import StatReload
except ImportError: # pragma: no cover - uvicorn is a hard dep
return False
if not config_path:
return False
from pathlib import Path
config_abs = Path(config_path).resolve()
patched_paths = getattr(StatReload, "_litellm_patched_config_paths", None)
if patched_paths is None:
original_iter = StatReload.iter_py_files
patched_paths = set()
def _iter_with_config(self): # type: ignore[no-untyped-def]
yield from original_iter(self)
for path in StatReload._litellm_patched_config_paths:
if path.exists():
yield path
StatReload.iter_py_files = _iter_with_config # type: ignore[assignment]
StatReload._litellm_patched_config_paths = patched_paths # type: ignore[attr-defined]
patched_paths.add(config_abs)
return True
@staticmethod
def _init_hypercorn_server(
app: FastAPI,
@ -619,7 +679,7 @@ class ProxyInitializationHelpers:
"--reload",
is_flag=True,
default=False,
help="Enable uvicorn hot reload (dev only). Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.",
help="Enable uvicorn hot reload (dev only). Also reloads when the --config YAML file changes. Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.",
)
def run_server( # noqa: PLR0915
host,
@ -1028,7 +1088,11 @@ def run_server( # noqa: PLR0915
uvicorn_args["loop"] = loop_type
if reload:
uvicorn_args["reload"] = True
uvicorn_args.update(
ProxyInitializationHelpers._get_reload_options(config)
)
if config:
ProxyInitializationHelpers._patch_statreload_for_config(config)
uvicorn.run(
**uvicorn_args,

View file

@ -403,6 +403,9 @@ from litellm.proxy.middleware.in_flight_requests_middleware import (
InFlightRequestsMiddleware,
)
from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMiddleware
from litellm.proxy.middleware.request_size_limit_middleware import (
RequestSizeLimitMiddleware,
)
from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router
from litellm.proxy.openai_files_endpoints.files_endpoints import (
router as openai_files_router,
@ -14881,6 +14884,11 @@ app.include_router(ui_discovery_endpoints_router)
app.include_router(google_router)
attach_lazy_features(app)
app.add_middleware(
RequestSizeLimitMiddleware,
get_max_request_size_mb=lambda: general_settings.get("max_request_size_mb"),
is_request_size_limit_enabled=lambda: premium_user is True,
)
async def _stream_mcp_asgi_response(

View file

@ -28956,6 +28956,19 @@
"mode": "chat",
"output_cost_per_token": 0.0
},
"sambanova/MiniMax-M2.7": {
"input_cost_per_token": 3e-07,
"litellm_provider": "sambanova",
"max_input_tokens": 204800,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"source": "https://cloud.sambanova.ai/plans/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"sambanova/DeepSeek-R1": {
"input_cost_per_token": 5e-06,
"litellm_provider": "sambanova",

View file

@ -0,0 +1,276 @@
#!/usr/bin/env python3
"""
Minimal HTTP target for testing LiteLLM **Bedrock pass-through** (`/bedrock/...` on the proxy).
What it does
- Serves a tiny Converse-shaped JSON (and optional invoke-shaped) response so the proxy can
complete a round trip without calling AWS.
- Does **not** verify SigV4 (Bedrock does); any Authorization header is accepted.
How to run
uv run python scripts/mock_bedrock_passthrough_target.py --host 127.0.0.1 --port 9999
Wire LiteLLM to this host (use **one** of these patterns):
1) model_list (recommended) — set the Bedrock runtime base to the mock:
model_list:
- model_name: mock-bedrock-claude
litellm_params:
model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0
custom_llm_provider: bedrock
aws_region_name: us-west-2
api_base: "http://127.0.0.1:9999"
2) Environment (see litellm BaseAWSLLM.get_runtime_endpoint)::
export AWS_BEDROCK_RUNTIME_ENDPOINT="http://127.0.0.1:9999"
Then call the proxy, e.g. (model_name must match config)::
curl -sS -X POST "http://127.0.0.1:4000/bedrock/model/mock-bedrock-claude/converse" \
-H "Authorization: Bearer $LITELLM_KEY" -H "Content-Type: application/json" \
-d '{"messages":[{"role":"user","content":[{"text":"hi"}]}]}'
The proxy will forward to: {api_base}/model/<resolved model id>/converse (SigV4-signed).
This mock implements POST .../converse and returns a minimal valid Converse response.
Notes
- `invoke-with-response-stream` returns a real **binary** AWS event stream
(`application/vnd.amazon.eventstream`) with Anthropic-style JSON payloads inside each
`PayloadPart`, matching Bedrock's InvokeModelWithResponseStream wire format. See
https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_InvokeModelWithResponseStream.html
and https://docs.aws.amazon.com/awstreams/latest/devguide/message-formats.html
- `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
import base64
import json
from binascii import crc32
from struct import pack
from typing import Any, Dict, Iterator, List
from fastapi import FastAPI, Request
from fastapi.responses import JSONResponse
from starlette.responses import StreamingResponse
app = FastAPI(title="Mock Bedrock runtime (pass-through test target)")
# Minimal structure compatible with Converse: https://docs.aws.amazon.com/bedrock/latest/APIReference/API_Converse.html
def _converse_response_body() -> Dict[str, Any]:
return {
"output": {
"message": {
"role": "assistant",
"content": [
{"text": "mock: ok from mock_bedrock_passthrough_target.py"}
],
}
},
"stopReason": "end_turn",
"usage": {
"inputTokens": 1,
"outputTokens": 2,
"totalTokens": 3,
},
}
# Minimal invoke (Anthropic messages on bedrock) style — adjust if you test /invoke
def _invoke_response_body() -> Dict[str, Any]:
return {
"id": "msg_mock",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "mock invoke response"}],
"model": "mock",
"stop_reason": "end_turn",
"usage": {"input_tokens": 1, "output_tokens": 2},
}
def _encode_event_stream_message(headers: Dict[str, str], payload: bytes) -> bytes:
"""Single AWS binary event-stream frame (same layout botocore's ``EventStreamBuffer`` parses)."""
header_blob = b""
for name, value in headers.items():
nb = name.encode("utf-8")
vb = value.encode("utf-8")
header_blob += bytes([len(nb)]) + nb + bytes([7]) + pack("!H", len(vb)) + vb
headers_length = len(header_blob)
payload_length = len(payload)
total_length = 12 + headers_length + payload_length + 4
prelude_wo_crc = pack("!II", total_length, headers_length)
prelude_crc_val = crc32(prelude_wo_crc) & 0xFFFFFFFF
prelude = prelude_wo_crc + pack("!I", prelude_crc_val)
wo_msg_crc = prelude + header_blob + payload
msg_crc_val = crc32(wo_msg_crc[8:], prelude_crc_val) & 0xFFFFFFFF
return wo_msg_crc + pack("!I", msg_crc_val)
def _bedrock_payload_part(inner_event: Dict[str, Any]) -> bytes:
"""Outer JSON expected by bedrock-runtime ``ResponseStream`` / ``PayloadPart``."""
inner_bytes = json.dumps(inner_event, separators=(",", ":")).encode("utf-8")
outer = {
"chunk": {
"bytes": base64.b64encode(inner_bytes).decode("ascii"),
}
}
return json.dumps(outer, separators=(",", ":")).encode("utf-8")
def _anthropic_invoke_stream_events(
model_id: str, assistant_text: str
) -> List[Dict[str, Any]]:
"""
Minimal Anthropic Messages stream events as returned inside Bedrock stream chunks.
Mirrors the sequence Amazon emits for Claude on ``invoke-with-response-stream``.
"""
msg_id = "msg_mock_bedrock_stream"
input_tokens = 3
output_tokens = max(1, len(assistant_text) // 4)
events: List[Dict[str, Any]] = [
{
"type": "message_start",
"message": {
"model": model_id,
"id": msg_id,
"type": "message",
"role": "assistant",
"content": [],
"stop_reason": None,
"stop_sequence": None,
"usage": {
"input_tokens": input_tokens,
"output_tokens": 1,
"cache_creation_input_tokens": 0,
"cache_read_input_tokens": 0,
"cache_creation": {
"ephemeral_5m_input_tokens": 0,
"ephemeral_1h_input_tokens": 0,
},
},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""},
},
]
# Split text into small deltas so downstream streaming behavior is visible.
step = 24
for i in range(0, len(assistant_text), step):
events.append(
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "text_delta",
"text": assistant_text[i : i + step],
},
}
)
events.append({"type": "content_block_stop", "index": 0})
events.append(
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"cache_creation_input_tokens": 0,
"cache_read_input_tokens": 0,
},
}
)
events.append(
{
"type": "message_stop",
"amazon-bedrock-invocationMetrics": {
"inputTokenCount": input_tokens,
"outputTokenCount": output_tokens,
"invocationLatency": 42,
"firstByteLatency": 10,
},
}
)
return events
def _iter_invoke_with_response_stream(model_id: str) -> Iterator[bytes]:
text = (
"mock streaming: ok from scripts/mock_bedrock_passthrough_target.py "
"(invoke-with-response-stream)."
)
headers = {
":event-type": "chunk",
":content-type": "application/json",
":message-type": "event",
}
for ev in _anthropic_invoke_stream_events(model_id, text):
yield _encode_event_stream_message(headers, _bedrock_payload_part(ev))
@app.get("/health")
def health() -> Dict[str, str]:
return {"status": "ok"}
@app.post("/model/{model_path:path}/converse")
async def converse(model_path: str, request: Request) -> JSONResponse:
# Optional: log body for debugging
_ = await request.body()
return JSONResponse(content=_converse_response_body())
@app.post("/model/{model_path:path}/converse-stream")
async def converse_stream(model_path: str, request: Request) -> JSONResponse:
"""
Not a real AWS event stream — returns JSON for quick smoke tests only.
"""
_ = await request.body()
return JSONResponse(
content={
"note": "This mock does not implement application/vnd.amazon.eventstream; use /converse for basic tests."
}
)
@app.post("/model/{model_path:path}/invoke")
async def invoke(model_path: str, request: Request) -> JSONResponse:
_ = await request.body()
return JSONResponse(content=_invoke_response_body())
@app.post("/model/{model_path:path}/invoke-with-response-stream")
async def invoke_with_response_stream(
model_path: str, request: Request
) -> StreamingResponse:
"""
Binary ``application/vnd.amazon.eventstream`` body compatible with boto3/botocore
``InvokeModelWithResponseStream`` / LiteLLM's Bedrock invoke streaming path.
"""
_ = await request.body()
return StreamingResponse(
_iter_invoke_with_response_stream(model_id=model_path),
media_type="application/vnd.amazon.eventstream",
)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=9999)
args = parser.parse_args()
import uvicorn
uvicorn.run(app, host=args.host, port=args.port, log_level="info")
if __name__ == "__main__":
main()

View file

@ -0,0 +1,29 @@
import json
from pathlib import Path
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
def test_sambanova_minimax_m27_model_info():
model = "sambanova/MiniMax-M2.7"
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
info = model_cost.get(model)
assert (
info is not None
), f"{model} not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "sambanova"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] > 0
assert info["output_cost_per_token"] > 0
assert info["max_input_tokens"] == 204800
assert info["max_output_tokens"] == 131072
assert info["supports_function_calling"] is True
assert info["supports_reasoning"] is True
assert info["supports_tool_choice"] is True
routed_model, provider, _, _ = get_llm_provider(model=model)
assert routed_model == "MiniMax-M2.7"
assert provider == "sambanova"

View file

@ -0,0 +1,135 @@
import pytest
from starlette.responses import JSONResponse
from starlette.testclient import TestClient
from starlette.types import Message
from litellm.proxy.middleware.request_size_limit_middleware import (
RequestSizeLimitMiddleware,
)
def test_request_size_limit_middleware_rejects_content_length_before_body_read():
downstream_called = False
async def app(scope, receive, send):
nonlocal downstream_called
downstream_called = True
response = JSONResponse({"ok": True})
await response(scope, receive, send)
client = TestClient(
RequestSizeLimitMiddleware(
app,
get_max_request_size_mb=lambda: 1,
is_request_size_limit_enabled=lambda: True,
)
)
response = client.post(
"/chat/completions",
content=b"x" * (1024 * 1024 + 1),
headers={"content-type": "application/json"},
)
assert response.status_code == 413
assert response.json() == {"error": "Request size is too large. Max size is 1 MB"}
assert response.headers["content-length"] == str(len(response.content))
assert downstream_called is False
def test_request_size_limit_middleware_zero_limit_disables_guard():
downstream_called = False
async def app(scope, receive, send):
nonlocal downstream_called
downstream_called = True
response = JSONResponse({"ok": True})
await response(scope, receive, send)
client = TestClient(
RequestSizeLimitMiddleware(
app,
get_max_request_size_mb=lambda: 0,
is_request_size_limit_enabled=lambda: True,
)
)
response = client.post(
"/chat/completions",
content=b"x",
headers={"content-type": "application/json"},
)
assert response.status_code == 200
assert response.json() == {"ok": True}
assert downstream_called is True
@pytest.mark.asyncio
async def test_request_size_limit_middleware_rejects_streamed_body_without_content_length():
received_body_bytes = 0
async def app(scope, receive, send):
nonlocal received_body_bytes
while True:
message = await receive()
if message["type"] == "http.disconnect":
break
received_body_bytes += len(message.get("body", b""))
if not message.get("more_body", False):
break
response = JSONResponse({"ok": True})
await response(scope, receive, send)
middleware = RequestSizeLimitMiddleware(
app,
get_max_request_size_mb=lambda: 1,
is_request_size_limit_enabled=lambda: True,
)
sent_messages: list[Message] = []
receive_messages: list[Message] = [
{
"type": "http.request",
"body": b"x" * (1024 * 1024),
"more_body": True,
},
{
"type": "http.request",
"body": b"y",
"more_body": False,
},
]
async def receive():
return receive_messages.pop(0)
async def send(message):
sent_messages.append(message)
await middleware(
{
"type": "http",
"method": "POST",
"path": "/chat/completions",
"headers": [(b"content-type", b"application/json")],
},
receive,
send,
)
expected_body = b'{"error":"Request size is too large. Max size is 1 MB"}'
assert sent_messages[0] == {
"type": "http.response.start",
"status": 413,
"headers": [
(b"content-type", b"application/json"),
(b"content-length", str(len(expected_body)).encode("latin-1")),
],
}
assert sent_messages[1] == {
"type": "http.response.body",
"body": expected_body,
"more_body": False,
}
assert received_body_bytes == 1024 * 1024

View file

@ -2,13 +2,18 @@
Test Azure Sentinel logging integration
"""
import datetime
from unittest.mock import AsyncMock, patch
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.integrations.azure_sentinel.azure_sentinel import AzureSentinelLogger
from litellm.types.utils import StandardLoggingPayload
from litellm.types.utils import StandardAuditLogPayload, StandardLoggingPayload
def _close_periodic_flush_task(coro):
coro.close()
return None
@pytest.mark.asyncio
@ -20,7 +25,7 @@ async def test_azure_sentinel_oauth_and_send_batch():
test_client_id = "test-client-id"
test_client_secret = "test-client-secret"
with patch("asyncio.create_task"):
with patch("asyncio.create_task", side_effect=_close_periodic_flush_task):
logger = AzureSentinelLogger(
dcr_immutable_id=test_dcr_id,
endpoint=test_endpoint,
@ -42,9 +47,6 @@ async def test_azure_sentinel_oauth_and_send_batch():
# Add to queue
logger.log_queue.append(standard_payload)
# Mock OAuth token response
from unittest.mock import MagicMock
mock_token_response = MagicMock()
mock_token_response.status_code = 200
mock_token_response.json = MagicMock(
@ -91,3 +93,173 @@ async def test_azure_sentinel_oauth_and_send_batch():
# Verify queue is cleared
assert len(logger.log_queue) == 0
@pytest.mark.asyncio
async def test_azure_sentinel_queues_audit_log_event():
"""Test that Azure Sentinel supports direct audit log callbacks"""
with patch("asyncio.create_task", side_effect=_close_periodic_flush_task):
logger = AzureSentinelLogger(
dcr_immutable_id="dcr-test123456789",
endpoint="https://test-dce.eastus-1.ingest.monitor.azure.com",
tenant_id="test-tenant-id",
client_id="test-client-id",
client_secret="test-client-secret",
)
logger.batch_size = 2
logger.async_send_audit_batch = AsyncMock()
audit_log = StandardAuditLogPayload(
id="audit-123",
updated_at="2026-05-06T04:39:00+00:00",
changed_by="user-1",
changed_by_api_key="sk-test",
action="created",
table_name="LiteLLM_TeamTable",
object_id="team-1",
before_value=None,
updated_values='{"team_alias": "sentinel-demo"}',
)
await logger.async_log_audit_log_event(audit_log)
assert logger.audit_log_queue == [audit_log]
logger.async_send_audit_batch.assert_not_called()
await logger.async_log_audit_log_event(audit_log)
assert logger.audit_log_queue == [audit_log, audit_log]
logger.async_send_audit_batch.assert_awaited_once()
@pytest.mark.asyncio
async def test_azure_sentinel_sends_audit_log_payload_to_ingestion_api():
"""Test that queued audit logs are sent to Azure Monitor Logs Ingestion"""
with patch("asyncio.create_task", side_effect=_close_periodic_flush_task):
logger = AzureSentinelLogger(
dcr_immutable_id="dcr-test123456789",
endpoint="https://test-dce.eastus-1.ingest.monitor.azure.com",
tenant_id="test-tenant-id",
client_id="test-client-id",
client_secret="test-client-secret",
)
audit_log = StandardAuditLogPayload(
id="audit-123",
updated_at="2026-05-06T04:39:00+00:00",
changed_by="user-1",
changed_by_api_key="sk-test",
action="created",
table_name="LiteLLM_TeamTable",
object_id="team-1",
before_value=None,
updated_values='{"team_alias": "sentinel-demo"}',
)
await logger.async_log_audit_log_event(audit_log)
mock_token_response = MagicMock()
mock_token_response.status_code = 200
mock_token_response.json = MagicMock(
return_value={
"access_token": "test-bearer-token",
"expires_in": 3600,
}
)
mock_token_response.text = "Success"
mock_api_response = MagicMock()
mock_api_response.status_code = 204
mock_api_response.text = "Success"
async def mock_post(*args, **kwargs):
if "oauth2/v2.0/token" in kwargs.get("url", ""):
return mock_token_response
return mock_api_response
logger.async_httpx_client.post = AsyncMock(side_effect=mock_post)
await logger.flush_queue()
api_call_args = logger.async_httpx_client.post.call_args_list[-1]
body = json.loads(api_call_args.kwargs["data"].decode("utf-8"))
assert body == [audit_log]
assert "dcr-test123456789" in api_call_args.kwargs["url"]
assert "Custom-LiteLLM" in api_call_args.kwargs["url"]
assert len(logger.audit_log_queue) == 0
@pytest.mark.asyncio
async def test_azure_sentinel_flushes_standard_and_audit_logs_separately():
"""Test mixed callback roles do not send schema-mismatched batches."""
with patch("asyncio.create_task", side_effect=_close_periodic_flush_task):
logger = AzureSentinelLogger(
dcr_immutable_id="dcr-test123456789",
stream_name="Custom-LiteLLM-Standard",
audit_stream_name="Custom-LiteLLM-Audit",
endpoint="https://test-dce.eastus-1.ingest.monitor.azure.com",
tenant_id="test-tenant-id",
client_id="test-client-id",
client_secret="test-client-secret",
)
standard_payload = StandardLoggingPayload(
id="standard-123",
call_type="completion",
model="gpt-3.5-turbo",
status="success",
messages=[{"role": "user", "content": "Hello"}],
response={"choices": [{"message": {"content": "Hi"}}]},
)
audit_log = StandardAuditLogPayload(
id="audit-123",
updated_at="2026-05-06T04:39:00+00:00",
changed_by="user-1",
changed_by_api_key="sk-test",
action="created",
table_name="LiteLLM_TeamTable",
object_id="team-1",
before_value=None,
updated_values='{"team_alias": "sentinel-demo"}',
)
logger.log_queue.append(standard_payload)
await logger.async_log_audit_log_event(audit_log)
mock_token_response = MagicMock()
mock_token_response.status_code = 200
mock_token_response.json = MagicMock(
return_value={
"access_token": "test-bearer-token",
"expires_in": 3600,
}
)
mock_token_response.text = "Success"
mock_api_response = MagicMock()
mock_api_response.status_code = 204
mock_api_response.text = "Success"
async def mock_post(*args, **kwargs):
if "oauth2/v2.0/token" in kwargs.get("url", ""):
return mock_token_response
return mock_api_response
logger.async_httpx_client.post = AsyncMock(side_effect=mock_post)
await logger.flush_queue()
ingestion_calls = [
call
for call in logger.async_httpx_client.post.call_args_list
if "dataCollectionRules" in call.kwargs["url"]
]
assert len(ingestion_calls) == 2
standard_call, audit_call = ingestion_calls
assert "Custom-LiteLLM-Standard" in standard_call.kwargs["url"]
assert json.loads(standard_call.kwargs["data"].decode("utf-8")) == [
standard_payload
]
assert "Custom-LiteLLM-Audit" in audit_call.kwargs["url"]
assert json.loads(audit_call.kwargs["data"].decode("utf-8")) == [audit_log]

View file

@ -1,4 +1,5 @@
import logging
import sys
import pytest
from prometheus_client import REGISTRY
@ -123,3 +124,36 @@ def test_virtual_key_rate_limit_metrics_accept_custom_metadata_labels(
and sample.value == 3
for sample in samples
)
def test_virtual_key_rate_limit_metrics_preserve_zero_remaining_values(
monkeypatch: pytest.MonkeyPatch,
):
prometheus_logger = _create_prometheus_logger_with_custom_labels(monkeypatch)
metadata = {
"model_group": "gpt-4o-mini",
"litellm-key-remaining-requests-gpt-4o-mini": 0,
"litellm-key-remaining-tokens-gpt-4o-mini": 0,
}
kwargs = {
"litellm_params": {
"metadata": metadata,
},
"standard_logging_object": _standard_logging_payload_with_requester_metadata(),
}
prometheus_logger._set_virtual_key_rate_limit_metrics(
user_api_key="test-hash",
user_api_key_alias="test-alias",
kwargs=kwargs,
metadata=metadata,
model_id="model-123",
)
request_samples = _metric_samples("litellm_remaining_api_key_requests_for_model")
token_samples = _metric_samples("litellm_remaining_api_key_tokens_for_model")
assert any(sample.value == 0 for sample in request_samples)
assert any(sample.value == 0 for sample in token_samples)
assert not any(sample.value == sys.maxsize for sample in request_samples)
assert not any(sample.value == sys.maxsize for sample in token_samples)

View file

@ -0,0 +1,181 @@
from time import monotonic
import pytest
from prometheus_client import REGISTRY
import litellm
from litellm.integrations.prometheus import PrometheusLogger
from litellm.integrations.prometheus_helpers import bounded_prometheus_series_tracker
from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import (
BoundedPrometheusSeriesTracker,
)
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
@pytest.fixture(autouse=True)
def cleanup_prometheus_registry():
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
try:
REGISTRY.unregister(collector)
except Exception:
pass
old_enable_end_user = litellm.enable_end_user_cost_tracking_prometheus_only
old_metrics_config = litellm.prometheus_metrics_config
old_max_series = litellm.prometheus_end_user_metrics_max_series_per_metric
old_ttl_seconds = litellm.prometheus_end_user_metrics_ttl_seconds
old_cleanup_interval_seconds = (
litellm.prometheus_end_user_metrics_cleanup_interval_seconds
)
yield
litellm.enable_end_user_cost_tracking_prometheus_only = old_enable_end_user
litellm.prometheus_metrics_config = old_metrics_config
litellm.prometheus_end_user_metrics_max_series_per_metric = old_max_series
litellm.prometheus_end_user_metrics_ttl_seconds = old_ttl_seconds
litellm.prometheus_end_user_metrics_cleanup_interval_seconds = (
old_cleanup_interval_seconds
)
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
try:
REGISTRY.unregister(collector)
except Exception:
pass
def test_prometheus_end_user_series_are_capped_per_metric():
litellm.enable_end_user_cost_tracking_prometheus_only = True
litellm.prometheus_metrics_config = [
{
"group": "end-user-spend",
"metrics": ["litellm_spend_metric"],
"include_labels": ["end_user"],
}
]
litellm.prometheus_end_user_metrics_max_series_per_metric = 3
litellm.prometheus_end_user_metrics_ttl_seconds = None
logger = PrometheusLogger()
for index in range(6):
PrometheusLogger._inc_labeled_counter(
logger,
logger.litellm_spend_metric,
"litellm_spend_metric",
UserAPIKeyLabelValues(end_user=f"end-user-{index}"),
amount=0.01,
)
assert len(logger.litellm_spend_metric._metrics) == 3
assert set(logger.litellm_spend_metric._metrics) == {
("end-user-3",),
("end-user-4",),
("end-user-5",),
}
def test_bounded_prometheus_series_tracker_is_label_agnostic():
class FakeMetric:
def __init__(self):
self.removed_label_values = []
def remove(self, *label_values):
self.removed_label_values.append(label_values)
metric = FakeMetric()
tracker = BoundedPrometheusSeriesTracker()
for index in range(4):
tracker.track_series(
metric=metric,
metric_name="generic_metric",
label_values=(f"route-{index}", "200"),
max_series=2,
ttl_seconds=None,
cleanup_interval_seconds=60.0,
)
assert metric.removed_label_values == [
("route-0", "200"),
("route-1", "200"),
]
def test_bounded_prometheus_series_tracker_treats_zero_max_as_unlimited():
# A misconfigured ``max_series=0`` must not silently evict every emission.
class FakeMetric:
def __init__(self):
self.removed_label_values = []
def remove(self, *label_values):
self.removed_label_values.append(label_values)
metric = FakeMetric()
tracker = BoundedPrometheusSeriesTracker()
for index in range(3):
tracker.track_series(
metric=metric,
metric_name="generic_metric",
label_values=(f"end-user-{index}",),
max_series=0,
ttl_seconds=None,
cleanup_interval_seconds=60.0,
)
assert metric.removed_label_values == []
def test_prometheus_end_user_series_expire_by_ttl(monkeypatch):
litellm.enable_end_user_cost_tracking_prometheus_only = True
litellm.prometheus_metrics_config = [
{
"group": "end-user-spend",
"metrics": ["litellm_spend_metric"],
"include_labels": ["end_user"],
}
]
litellm.prometheus_end_user_metrics_max_series_per_metric = None
litellm.prometheus_end_user_metrics_ttl_seconds = 10.0
litellm.prometheus_end_user_metrics_cleanup_interval_seconds = 0.0
logger = PrometheusLogger()
current_time = [monotonic()]
monkeypatch.setattr(
bounded_prometheus_series_tracker.time,
"monotonic",
lambda: current_time[0],
)
PrometheusLogger._inc_labeled_counter(
logger,
logger.litellm_spend_metric,
"litellm_spend_metric",
UserAPIKeyLabelValues(end_user="stale-end-user"),
amount=0.01,
)
current_time[0] += 11.0
PrometheusLogger._inc_labeled_counter(
logger,
logger.litellm_spend_metric,
"litellm_spend_metric",
UserAPIKeyLabelValues(end_user="fresh-end-user"),
amount=0.01,
)
assert set(logger.litellm_spend_metric._metrics) == {("fresh-end-user",)}
def test_prometheus_end_user_not_tracked_by_default():
litellm.enable_end_user_cost_tracking_prometheus_only = None
labels = PrometheusLogger().get_labels_for_metric("litellm_spend_metric")
assert "end_user" in labels
label_values = UserAPIKeyLabelValues(end_user="not-exported")
from litellm.integrations.prometheus import prometheus_label_factory
prometheus_labels = prometheus_label_factory(labels, label_values)
assert prometheus_labels["end_user"] is None

View file

@ -1,7 +1,11 @@
import json
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from unittest.mock import AsyncMock, MagicMock
import pytest
import litellm
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.llms.anthropic.chat.handler import ModelResponseIterator, make_call
from litellm.types.llms.openai import (
@ -343,6 +347,289 @@ def test_text_only_streaming_has_index_zero():
), f"Expected index=0, got {parsed.choices[0].index}"
def test_streaming_thinking_deltas_count_reasoning_tokens_in_usage():
"""Anthropic streaming usage should account for emitted thinking deltas."""
chunks = [
{
"type": "message_start",
"message": {
"id": "msg_123",
"type": "message",
"role": "assistant",
"content": [],
"usage": {"input_tokens": 10, "output_tokens": 1},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "thinking", "thinking": ""},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "thinking_delta",
"thinking": "First I need to count the favorable outcomes. ",
},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "thinking_delta",
"thinking": "Then I compare that count with all possible outcomes.",
},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "signature_delta", "signature": "sig_123"},
},
{"type": "content_block_stop", "index": 0},
{
"type": "content_block_start",
"index": 1,
"content_block": {"type": "text", "text": ""},
},
{
"type": "content_block_delta",
"index": 1,
"delta": {"type": "text_delta", "text": "The probability is 3/8."},
},
{"type": "content_block_stop", "index": 1},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 50},
},
]
iterator = ModelResponseIterator(None, sync_stream=True)
final_usage = None
reasoning_deltas = []
for chunk in chunks:
parsed = iterator.chunk_parser(chunk)
reasoning_content = getattr(parsed.choices[0].delta, "reasoning_content", None)
if reasoning_content:
reasoning_deltas.append(reasoning_content)
if parsed.usage is not None:
final_usage = parsed.usage
assert reasoning_deltas == [
"First I need to count the favorable outcomes. ",
"Then I compare that count with all possible outcomes.",
]
assert final_usage is not None
completion_tokens_details = final_usage.completion_tokens_details
assert completion_tokens_details is not None
assert completion_tokens_details.reasoning_tokens > 0
assert completion_tokens_details.text_tokens == (
final_usage.completion_tokens - completion_tokens_details.reasoning_tokens
)
def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinking():
"""The completion API should preserve Anthropic thinking usage in streaming mode."""
thinking_parts = [
"First I need to count the favorable outcomes. ",
"Then I compare that count with all possible outcomes.",
]
thinking_text = "".join(thinking_parts)
answer_text = "The probability is 3/8."
requests_seen = []
class MockAnthropicHandler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def log_message(self, format, *args): # type: ignore[no-untyped-def]
return
def do_POST(self): # type: ignore[no-untyped-def]
content_length = int(self.headers.get("content-length", "0"))
payload = json.loads(self.rfile.read(content_length).decode("utf-8"))
requests_seen.append(
{
"path": self.path,
"model": payload.get("model"),
"stream": payload.get("stream", False),
"thinking": payload.get("thinking"),
}
)
if payload.get("stream"):
events = [
{
"type": "message_start",
"message": {
"id": "msg_mock",
"type": "message",
"role": "assistant",
"model": payload.get("model"),
"content": [],
"usage": {"input_tokens": 10, "output_tokens": 1},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "thinking", "thinking": ""},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "thinking_delta",
"thinking": thinking_parts[0],
},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "thinking_delta",
"thinking": thinking_parts[1],
},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "signature_delta",
"signature": "sig_mock",
},
},
{"type": "content_block_stop", "index": 0},
{
"type": "content_block_start",
"index": 1,
"content_block": {"type": "text", "text": ""},
},
{
"type": "content_block_delta",
"index": 1,
"delta": {"type": "text_delta", "text": answer_text},
},
{"type": "content_block_stop", "index": 1},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 50},
},
{"type": "message_stop"},
]
self._write_response(
content_type="text/event-stream",
body="".join(
f"data: {json.dumps(event)}\n\n" for event in events
).encode("utf-8"),
)
return
self._write_response(
content_type="application/json",
body=json.dumps(
{
"id": "msg_mock",
"type": "message",
"role": "assistant",
"model": payload.get("model"),
"content": [
{
"type": "thinking",
"thinking": thinking_text,
"signature": "sig_mock",
},
{"type": "text", "text": answer_text},
],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 50},
}
).encode("utf-8"),
)
def _write_response(self, content_type: str, body: bytes) -> None:
self.send_response(200)
self.send_header("content-type", content_type)
self.send_header("content-length", str(len(body)))
self.end_headers()
self.wfile.write(body)
server = ThreadingHTTPServer(("127.0.0.1", 0), MockAnthropicHandler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
request_kwargs = {
"model": "anthropic/claude-sonnet-4-6",
"api_base": f"http://127.0.0.1:{server.server_port}",
"api_key": "test",
"messages": [
{
"role": "user",
"content": "Solve a probability problem and show thinking.",
}
],
"thinking": {"type": "adaptive"},
"max_tokens": 128,
}
non_stream_response = litellm.completion(**request_kwargs, stream=False)
non_stream_details = non_stream_response.usage.completion_tokens_details
assert non_stream_details is not None
assert non_stream_details.reasoning_tokens > 0
reasoning_chunks = []
content_chunks = []
stream_usage = None
for chunk in litellm.completion(
**request_kwargs,
stream=True,
stream_options={"include_usage": True},
):
chunk_dict = chunk.model_dump(exclude_none=True)
choices = chunk_dict.get("choices") or []
if choices:
delta = choices[0].get("delta") or {}
if delta.get("reasoning_content"):
reasoning_chunks.append(delta["reasoning_content"])
if delta.get("content"):
content_chunks.append(delta["content"])
if chunk_dict.get("usage"):
stream_usage = chunk_dict["usage"]
assert reasoning_chunks == thinking_parts
assert content_chunks == [answer_text]
assert stream_usage is not None
stream_completion_details = stream_usage["completion_tokens_details"]
assert (
stream_completion_details["reasoning_tokens"]
== non_stream_details.reasoning_tokens
)
assert stream_completion_details["text_tokens"] == (
stream_usage["completion_tokens"]
- stream_completion_details["reasoning_tokens"]
)
assert requests_seen == [
{
"path": "/v1/messages",
"model": "claude-sonnet-4-6",
"stream": False,
"thinking": {"type": "adaptive"},
},
{
"path": "/v1/messages",
"model": "claude-sonnet-4-6",
"stream": True,
"thinking": {"type": "adaptive"},
},
]
finally:
server.shutdown()
def test_text_and_tool_streaming_has_index_zero():
"""Test that mixed text and tool streaming responses have choice index=0"""
chunks = [

View file

@ -97,6 +97,34 @@ def test_calculate_usage():
assert usage._cache_read_input_tokens == 0
def test_calculate_usage_clamps_text_tokens_when_reasoning_estimate_exceeds_output():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={"input_tokens": 10, "output_tokens": 1},
reasoning_content="This reasoning text intentionally tokenizes above one output token.",
)
assert usage.completion_tokens == 1
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens == usage.completion_tokens
assert usage.completion_tokens_details.text_tokens == 0
def test_calculate_usage_handles_mocked_output_tokens_with_reasoning_content():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={"input_tokens": 10, "output_tokens": MagicMock()},
reasoning_content="mocked response reasoning",
)
assert usage.completion_tokens == 0
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens == 0
assert usage.completion_tokens_details.text_tokens == 0
@pytest.mark.parametrize(
"usage_object,expected_usage",
[

View file

@ -13,6 +13,110 @@ sys.path.insert(
from litellm.llms.bedrock.common_utils import BedrockModelInfo
# --------------------------------------------------------------------------- #
# BEDROCK_RESPONSE_STREAM_SHAPE eager-load tests #
# --------------------------------------------------------------------------- #
def test_bedrock_response_stream_shape_loaded_at_import():
"""
BEDROCK_RESPONSE_STREAM_SHAPE is resolved at module import time.
In a standard environment with botocore installed it must be non-None.
"""
from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE
assert BEDROCK_RESPONSE_STREAM_SHAPE is not None
def test_bedrock_response_stream_shape_load_failure_returns_none():
"""
If botocore's Loader raises (e.g. missing data files), _load_bedrock_response_stream_shape
should return None rather than propagating the exception, so the module
still imports cleanly.
"""
from unittest.mock import patch
import litellm.llms.bedrock.common_utils as mod
with patch(
"botocore.loaders.Loader.load_service_model",
side_effect=Exception("no data"),
):
shape = mod._load_bedrock_response_stream_shape()
assert shape is None
def test_bedrock_response_stream_shape_is_structure_shape():
"""
The loaded shape should be the botocore StructureShape for ResponseStream,
not a plain dict or any other type.
"""
from botocore.model import StructureShape
from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE
assert BEDROCK_RESPONSE_STREAM_SHAPE is not None, (
"BEDROCK_RESPONSE_STREAM_SHAPE is None — botocore may not be installed"
)
shape: StructureShape = BEDROCK_RESPONSE_STREAM_SHAPE # remove Optional
assert isinstance(shape, StructureShape)
assert shape.name == "ResponseStream"
def test_bedrock_response_stream_shape_same_object_across_imports():
"""
Both bedrock modules that use the shape must reference the identical object —
confirming the constant is not re-loaded per import.
"""
from litellm.llms.bedrock.chat.invoke_handler import (
BEDROCK_RESPONSE_STREAM_SHAPE as invoke_shape,
)
from litellm.llms.bedrock.common_utils import (
BEDROCK_RESPONSE_STREAM_SHAPE as common_shape,
)
assert common_shape is invoke_shape
def test_bedrock_event_stream_decoder_base_uses_module_shape():
"""
BedrockEventStreamDecoderBase instances no longer carry their own
per-instance cache — _parse_message_from_event uses the module constant
directly, so there is no instance-level _response_stream_shape_cache attr.
"""
from litellm.llms.bedrock.common_utils import BedrockEventStreamDecoderBase
decoder_a = BedrockEventStreamDecoderBase()
decoder_b = BedrockEventStreamDecoderBase()
assert "_response_stream_shape_cache" not in decoder_a.__dict__
assert "_response_stream_shape_cache" not in decoder_b.__dict__
def test_bedrock_parse_message_from_event_raises_on_none_shape():
"""
When BEDROCK_RESPONSE_STREAM_SHAPE is None (botocore unavailable),
_parse_message_from_event must raise BedrockError before touching the
botocore parser — not an opaque AttributeError from inside botocore.
"""
from unittest.mock import MagicMock, patch
import litellm.llms.bedrock.common_utils as mod
from litellm.llms.bedrock.common_utils import BedrockError, BedrockEventStreamDecoderBase
decoder = BedrockEventStreamDecoderBase()
mock_event = MagicMock()
with patch.object(mod, "BEDROCK_RESPONSE_STREAM_SHAPE", None):
with pytest.raises(BedrockError) as exc_info:
decoder._parse_message_from_event(mock_event)
assert exc_info.value.status_code == 500
assert "botocore" in str(exc_info.value.message).lower()
# The botocore parser must never have been called
mock_event.to_response_dict.assert_not_called()
def test_deepseek_cris():
"""
Test that DeepSeek models with cross-region inference prefix use converse route

View file

@ -1,8 +1,10 @@
import asyncio
import io
import os
import pathlib
import ssl
import sys
import threading
from unittest.mock import MagicMock, patch
import certifi
@ -18,11 +20,111 @@ from litellm.llms.custom_httpx.aiohttp_transport import LiteLLMAiohttpTransport
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
MaskedHTTPStatusError,
_get_httpx_client,
get_ssl_configuration,
)
@pytest.mark.asyncio
async def test_async_post_streaming_status_error_should_not_wait_forever_for_body(
monkeypatch,
):
"""
Vertex Anthropic streamRawPredict can return a pre-stream 4xx where the
streamed error body never terminates. The handler must still surface the
status promptly instead of blocking the downstream client.
"""
class HangingErrorStream(httpx.AsyncByteStream):
async def __aiter__(self):
await asyncio.Event().wait()
if False:
yield b""
async def mock_handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(
400,
request=request,
headers={"content-type": "application/json"},
stream=HangingErrorStream(),
)
monkeypatch.setattr(
"litellm.llms.custom_httpx.http_handler._STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS",
0.01,
)
litellm_handler = AsyncHTTPHandler()
await litellm_handler.client.aclose()
litellm_handler.client = httpx.AsyncClient(
transport=httpx.MockTransport(mock_handler)
)
try:
with pytest.raises(MaskedHTTPStatusError) as exc_info:
await asyncio.wait_for(
litellm_handler.post(
"https://vertex.example/streamRawPredict",
stream=True,
),
timeout=0.2,
)
assert exc_info.value.status_code == 400
assert exc_info.value.response.status_code == 400
finally:
await litellm_handler.close()
def test_sync_post_streaming_status_error_should_not_wait_forever_for_body(
monkeypatch,
):
"""
Keep the sync streaming error path aligned with the async path so a
non-terminating streamed error body cannot block a worker thread forever.
"""
class HangingSyncErrorStream(httpx.SyncByteStream):
def __init__(self):
self.closed_event = threading.Event()
def __iter__(self):
self.closed_event.wait()
if False:
yield b""
def close(self):
self.closed_event.set()
def mock_handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(
400,
request=request,
headers={"content-type": "application/json"},
stream=HangingSyncErrorStream(),
)
monkeypatch.setattr(
"litellm.llms.custom_httpx.http_handler._STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS",
0.01,
)
litellm_handler = HTTPHandler()
litellm_handler.client.close()
litellm_handler.client = httpx.Client(transport=httpx.MockTransport(mock_handler))
try:
with pytest.raises(MaskedHTTPStatusError) as exc_info:
litellm_handler.post(
"https://vertex.example/streamRawPredict",
stream=True,
)
assert exc_info.value.status_code == 400
assert exc_info.value.response.status_code == 400
finally:
litellm_handler.close()
@pytest.mark.asyncio
async def test_ssl_security_level(monkeypatch):
# Ensure aiohttp transport is enabled for this test

View file

@ -11,6 +11,102 @@ from litellm.llms.sagemaker.common_utils import AWSEventStreamDecoder
from litellm.llms.sagemaker.completion.transformation import SagemakerConfig
# --------------------------------------------------------------------------- #
# SAGEMAKER_RESPONSE_STREAM_SHAPE eager-load tests #
# --------------------------------------------------------------------------- #
def test_sagemaker_response_stream_shape_loaded_at_import():
"""
SAGEMAKER_RESPONSE_STREAM_SHAPE is resolved at module import time.
In a standard environment with botocore installed it must be non-None.
"""
from litellm.llms.sagemaker.common_utils import SAGEMAKER_RESPONSE_STREAM_SHAPE
assert SAGEMAKER_RESPONSE_STREAM_SHAPE is not None
def test_sagemaker_response_stream_shape_load_failure_returns_none():
"""
If botocore's Loader raises (e.g. missing data files), _load_sagemaker_response_stream_shape
should return None rather than propagating the exception, so the module
still imports cleanly.
"""
from unittest.mock import patch
import litellm.llms.sagemaker.common_utils as mod
with patch(
"botocore.loaders.Loader.load_service_model",
side_effect=Exception("no data"),
):
shape = mod._load_sagemaker_response_stream_shape()
assert shape is None
def test_sagemaker_response_stream_shape_is_structure_shape():
"""
The loaded shape should be the botocore StructureShape for
InvokeEndpointWithResponseStreamOutput, not a plain dict or any other type.
"""
from botocore.model import StructureShape
from litellm.llms.sagemaker.common_utils import SAGEMAKER_RESPONSE_STREAM_SHAPE
assert SAGEMAKER_RESPONSE_STREAM_SHAPE is not None, (
"SAGEMAKER_RESPONSE_STREAM_SHAPE is None — botocore may not be installed"
)
shape: StructureShape = SAGEMAKER_RESPONSE_STREAM_SHAPE # remove Optional
assert isinstance(shape, StructureShape)
assert shape.name == "InvokeEndpointWithResponseStreamOutput"
def test_sagemaker_response_stream_shape_not_reloaded_on_new_decoder():
"""
Creating multiple AWSEventStreamDecoder instances must not trigger
additional botocore Loader calls — the shape is resolved once at import
time and reused.
"""
from litellm.llms.sagemaker.common_utils import SAGEMAKER_RESPONSE_STREAM_SHAPE
decoder_a = AWSEventStreamDecoder(model="test-model-a")
decoder_b = AWSEventStreamDecoder(model="test-model-b")
# Both decoders should use the same pre-loaded shape object (identity check)
assert "_response_stream_shape_cache" not in decoder_a.__dict__
assert "_response_stream_shape_cache" not in decoder_b.__dict__
# The module constant is still the same object
from litellm.llms.sagemaker.common_utils import (
SAGEMAKER_RESPONSE_STREAM_SHAPE as shape_after,
)
assert SAGEMAKER_RESPONSE_STREAM_SHAPE is shape_after
def test_sagemaker_parse_message_from_event_raises_on_none_shape():
"""
When SAGEMAKER_RESPONSE_STREAM_SHAPE is None (botocore unavailable),
_parse_message_from_event must raise ValueError before touching the
botocore parser — not an opaque AttributeError from inside botocore.
"""
from unittest.mock import MagicMock, patch
import litellm.llms.sagemaker.common_utils as mod
from litellm.llms.sagemaker.common_utils import SagemakerError
decoder = AWSEventStreamDecoder(model="test-model")
mock_event = MagicMock()
with patch.object(mod, "SAGEMAKER_RESPONSE_STREAM_SHAPE", None):
with pytest.raises(SagemakerError) as exc_info:
decoder._parse_message_from_event(mock_event)
assert exc_info.value.status_code == 500
assert "botocore" in str(exc_info.value.message).lower()
# The botocore parser must never have been called
mock_event.to_response_dict.assert_not_called()
@pytest.mark.asyncio
async def test_aiter_bytes_unicode_decode_error():
"""

View file

@ -2282,6 +2282,144 @@ class TestMCPServerManagerReload:
mock_build.assert_awaited_once_with(db_row)
assert manager.registry["server-1"] is rebuilt_server
@pytest.mark.asyncio
async def test_skips_server_when_build_from_database_fails(self, caplog):
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
except ImportError:
pytest.skip("MCP server not available")
manager = MCPServerManager()
timestamp = datetime.utcnow()
healthy_row = _make_db_mcp_server("healthy-server", timestamp)
bad_row = _make_db_mcp_server("bad-server", timestamp)
another_healthy_row = _make_db_mcp_server("another-healthy-server", timestamp)
healthy_server = MCPServer(
server_id="healthy-server",
name="healthy",
transport=MCPTransport.http,
updated_at=timestamp,
)
another_healthy_server = MCPServer(
server_id="another-healthy-server",
name="another-healthy",
transport=MCPTransport.http,
updated_at=timestamp,
)
async def build_server(db_row):
if db_row.server_id == "bad-server":
raise RuntimeError("transient build failure")
if db_row.server_id == "healthy-server":
return healthy_server
return another_healthy_server
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[healthy_row, bad_row, another_healthy_row]
)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma,
),
patch.object(
manager,
"build_mcp_server_from_table",
AsyncMock(side_effect=build_server),
),
patch.object(manager, "_maybe_register_openapi_tools", AsyncMock()),
caplog.at_level("ERROR", logger="LiteLLM"),
):
await manager.reload_servers_from_database()
assert set(manager.registry) == {"healthy-server", "another-healthy-server"}
assert manager.registry["healthy-server"] is healthy_server
assert manager.registry["another-healthy-server"] is another_healthy_server
assert "Skipping MCP server bad-server" in caplog.text
@pytest.mark.asyncio
async def test_skips_server_when_openapi_registration_fails(self, caplog):
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
except ImportError:
pytest.skip("MCP server not available")
manager = MCPServerManager()
timestamp = datetime.utcnow()
healthy_row = _make_db_mcp_server("healthy-server", timestamp)
bad_openapi_row = _make_db_mcp_server("bad-openapi-server", timestamp)
existing_server = MCPServer(
server_id="existing-server",
name="existing",
transport=MCPTransport.http,
updated_at=timestamp,
)
manager.registry = {existing_server.server_id: existing_server}
healthy_server = MCPServer(
server_id="healthy-server",
name="healthy",
transport=MCPTransport.http,
updated_at=timestamp,
)
bad_openapi_server = MCPServer(
server_id="bad-openapi-server",
name="bad-openapi",
transport=MCPTransport.http,
spec_path="https://example.invalid/openapi.json",
updated_at=timestamp,
)
async def build_server(db_row):
if db_row.server_id == "healthy-server":
return healthy_server
return bad_openapi_server
observed_registries = []
async def register_openapi_tools(server, **kwargs):
observed_registries.append(set(manager.registry))
assert kwargs == {"initialize_mapping": False}
if server.server_id == "bad-openapi-server":
raise RuntimeError("blocked address")
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[healthy_row, bad_openapi_row]
)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma,
),
patch.object(
manager,
"build_mcp_server_from_table",
AsyncMock(side_effect=build_server),
),
patch.object(
manager,
"_maybe_register_openapi_tools",
AsyncMock(side_effect=register_openapi_tools),
),
caplog.at_level("ERROR", logger="LiteLLM"),
):
await manager.reload_servers_from_database()
assert set(manager.registry) == {"healthy-server"}
assert manager.registry["healthy-server"] is healthy_server
assert observed_registries == [
{"existing-server"},
{"existing-server"},
]
assert "Skipping MCP server bad-openapi-server" in caplog.text
@pytest.mark.asyncio
async def test_call_mcp_tool_logs_failure_via_post_call_failure_hook():
@ -2946,7 +3084,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow():
"""
P1 Regression: list_tools path must apply _resolve_oauth2_flow to legacy DB
rows where oauth2_flow is NULL but M2M credentials are present.
Without this fix, has_client_credentials returns False and the caller's
Authorization header is forwarded upstream instead of being blocked.
"""
@ -3044,7 +3182,7 @@ async def test_call_tool_empty_extra_headers_returns_none():
"""
P2 Regression: When all configured extra_headers are filtered out (e.g.
Authorization for M2M), the resulting extra_headers should be None, not {}.
Downstream code that checks `if extra_headers is None` will behave
differently if an empty dict is passed instead.
"""
@ -3071,7 +3209,10 @@ async def test_call_tool_empty_extra_headers_returns_none():
extra_headers=["Authorization"], # Will be filtered out for M2M
)
raw_headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
raw_headers = {
"Authorization": "Bearer sk-1234",
"Content-Type": "application/json",
}
captured_extra_headers = None
@ -3108,8 +3249,8 @@ async def test_call_tool_empty_extra_headers_returns_none():
pass # We only care about the captured headers
# With P2 fix: extra_headers should be None (not {}) when all headers filtered
assert captured_extra_headers is None, (
"P2 API consistency issue: expected None for empty extra_headers, got: "
+ str(captured_extra_headers)
assert (
captured_extra_headers is None
), "P2 API consistency issue: expected None for empty extra_headers, got: " + str(
captured_extra_headers
)

View file

@ -906,6 +906,285 @@ def test_can_object_call_model_no_access_to_alias_or_underlying():
assert "my-fake-gpt" in str(exc_info.value.message)
# -- Team-member access-group resolution with team-scoped DB models -----------
def _make_team_scoped_router(team_id: str = "team-a"):
"""
Build a Router whose model_list looks like what the proxy creates for
team-scoped BYOK DB models: the internal model_name is
``<public_name>_<team_id>_<uuid>`` and the public name lives in
``model_info.team_public_model_name``. Two models belong to the
access group ``fast-models``; one (``mock-power``) does not.
"""
from litellm import Router
model_list = [
{
"model_name": f"mock-fast-1_{team_id}_aaa",
"litellm_params": {
"model": "openai/mock-fast-1",
"api_key": "fake",
},
"model_info": {
"id": f"demo-mock-fast-1-{team_id}",
"team_id": team_id,
"team_public_model_name": "mock-fast-1",
"access_groups": ["fast-models"],
},
},
{
"model_name": f"mock-fast-2_{team_id}_bbb",
"litellm_params": {
"model": "openai/mock-fast-2",
"api_key": "fake",
},
"model_info": {
"id": f"demo-mock-fast-2-{team_id}",
"team_id": team_id,
"team_public_model_name": "mock-fast-2",
"access_groups": ["fast-models"],
},
},
{
"model_name": f"mock-power_{team_id}_ccc",
"litellm_params": {
"model": "openai/mock-power",
"api_key": "fake",
},
"model_info": {
"id": f"demo-mock-power-{team_id}",
"team_id": team_id,
"team_public_model_name": "mock-power",
},
},
]
return Router(model_list=model_list)
def test_can_object_call_model_access_group_with_team_id():
"""
When team_id is passed, _can_object_call_model should resolve
model_info.access_groups for team-scoped DB models and allow
access via group name.
"""
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = _make_team_scoped_router()
result = _can_object_call_model(
model="mock-fast-1",
llm_router=router,
models=["fast-models", "mock-power"],
object_type="team",
team_id="team-a",
)
assert result is True
def test_can_object_call_model_access_group_without_team_id_fails():
"""
Without team_id the router cannot find team-scoped DB models, so
access group resolution fails and the call is denied.
This is the pre-fix behavior.
"""
from litellm.proxy._types import ProxyException
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = _make_team_scoped_router()
with pytest.raises(ProxyException):
_can_object_call_model(
model="mock-fast-1",
llm_router=router,
models=["fast-models", "mock-power"],
object_type="team",
# team_id intentionally omitted
)
def test_can_object_call_model_literal_name_with_team_id():
"""
Literal model name matching should still work when team_id is
passed — no regression from adding team_id.
"""
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = _make_team_scoped_router()
result = _can_object_call_model(
model="mock-power",
llm_router=router,
models=["fast-models", "mock-power"],
object_type="team",
team_id="team-a",
)
assert result is True
def test_can_object_call_model_denied_model_with_team_id():
"""
A model not in the allowed list (by name or access group) should
still be denied even when team_id is passed.
"""
from litellm.proxy._types import ProxyException
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = _make_team_scoped_router()
with pytest.raises(ProxyException):
_can_object_call_model(
model="mock-vision",
llm_router=router,
models=["fast-models", "mock-power"],
object_type="team",
team_id="team-a",
)
def test_can_object_call_model_second_group_member_with_team_id():
"""
Both models in the access group should be reachable, not just
the first one.
"""
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = _make_team_scoped_router()
result = _can_object_call_model(
model="mock-fast-2",
llm_router=router,
models=["fast-models"],
object_type="team",
team_id="team-a",
)
assert result is True
@pytest.mark.asyncio
async def test_check_team_member_model_access_with_access_group():
"""
End-to-end test of _check_team_member_model_access: a member whose
allowed_models contains an access group name should be allowed to
call models in that group for team-scoped DB models.
"""
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import _check_team_member_model_access
router = _make_team_scoped_router()
team = LiteLLM_TeamTable(team_id="team-a")
token = UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a")
membership = LiteLLM_TeamMembership(
user_id="alice",
team_id="team-a",
litellm_budget_table=LiteLLM_BudgetTable(
allowed_models=["fast-models", "mock-power"],
),
)
with patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
return_value=membership,
):
# Should not raise — mock-fast-1 is in the fast-models group
await _check_team_member_model_access(
model="mock-fast-1",
team_object=team,
valid_token=token,
llm_router=router,
prisma_client=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
)
@pytest.mark.asyncio
async def test_check_team_member_model_access_denied_model():
"""
A member with per-member allowed_models should be denied access to
a model that is neither listed by name nor covered by an access group.
"""
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import _check_team_member_model_access
router = _make_team_scoped_router()
team = LiteLLM_TeamTable(team_id="team-a")
token = UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a")
membership = LiteLLM_TeamMembership(
user_id="alice",
team_id="team-a",
litellm_budget_table=LiteLLM_BudgetTable(
allowed_models=["fast-models", "mock-power"],
),
)
with patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
return_value=membership,
):
with pytest.raises(ProxyException) as exc_info:
await _check_team_member_model_access(
model="mock-vision",
team_object=team,
valid_token=token,
llm_router=router,
prisma_client=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
)
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
@pytest.mark.asyncio
async def test_check_team_member_model_access_no_override_inherits_team():
"""
When a member has no allowed_models (empty budget table), the function
should return without raising — the team-level check applies instead.
"""
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import _check_team_member_model_access
router = _make_team_scoped_router()
team = LiteLLM_TeamTable(team_id="team-a")
token = UserAPIKeyAuth(token="sk-test", user_id="bob", team_id="team-a")
membership = LiteLLM_TeamMembership(
user_id="bob",
team_id="team-a",
litellm_budget_table=LiteLLM_BudgetTable(),
)
with patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
return_value=membership,
):
# Should return without raising — no per-member restriction
await _check_team_member_model_access(
model="mock-vision",
team_object=team,
valid_token=token,
llm_router=router,
prisma_client=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
)
# Tag Budget Enforcement Tests

View file

@ -306,6 +306,79 @@ async def test_team_member_budget_check_no_team_membership():
assert result is True
@pytest.mark.asyncio
async def test_team_member_budget_check_blocks_regenerated_key_after_old_key_exhausts_budget():
"""Deleting an exhausted key and creating a new key must not reset a user's team budget."""
request_body = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "test"}],
}
team_object = LiteLLM_TeamTable(
team_id="test-team-1",
team_alias="Test Team",
spend=0.0,
max_budget=None,
)
# The spend below represents usage accumulated by an earlier key that was
# later deleted. The new key must still be checked against the same
# user/team membership spend instead of receiving a fresh per-key budget.
regenerated_token = UserAPIKeyAuth(
token="new-regenerated-token",
user_id="test-user-1",
team_id="test-team-1",
models=["gpt-3.5-turbo"],
)
team_membership = LiteLLM_TeamMembership(
user_id="test-user-1",
team_id="test-team-1",
spend=0.0000002,
litellm_budget_table=LiteLLM_BudgetTable(
max_budget=0.0000001,
),
)
mock_request = MagicMock(spec=Request)
mock_prisma_client = MagicMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
with (
patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
new_callable=AsyncMock,
return_value=team_membership,
) as mock_get_team_membership,
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await common_checks(
request_body=request_body,
team_object=team_object,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=mock_proxy_logging_obj,
valid_token=regenerated_token,
request=mock_request,
)
mock_get_team_membership.assert_any_await(
user_id="test-user-1",
team_id="test-team-1",
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
)
assert "Budget has been exceeded" in str(exc_info.value)
assert "test-user-1" in str(exc_info.value)
assert "test-team-1" in str(exc_info.value)
@pytest.mark.asyncio
async def test_team_member_budget_check_personal_key_not_team():
"""Test that team member budget check is skipped for personal keys (no team)."""

View file

@ -4,6 +4,7 @@ import pytest
from fastapi import HTTPException
from litellm.proxy._types import (
LiteLLM_UserTable,
LitellmUserRoles,
NewUserRequest,
NewUserResponse,
@ -16,8 +17,8 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import (
_process_group_patch_operations,
create_group,
create_user,
get_users,
get_service_provider_config,
patch_group,
patch_user,
update_group,
update_user,
@ -259,6 +260,124 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp
)
@pytest.mark.asyncio
async def test_get_users_filters_username_by_exposed_scim_username_for_okta(mocker):
"""
Okta deprovisioning first locates a user with `userName eq "<email>"`.
LiteLLM exposes SCIM userName from user_email, so the lookup must match
user_email even when the internal user_id is a UUID.
"""
user = LiteLLM_UserTable(
user_id="internal-user-id",
user_email="okta.user@example.com",
user_alias="Okta User",
teams=[],
metadata={},
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[user])
mock_prisma_client.db.litellm_usertable.count = AsyncMock(return_value=1)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
AsyncMock(
return_value=SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
id="internal-user-id",
userName="okta.user@example.com",
emails=[SCIMUserEmail(value="okta.user@example.com")],
)
),
)
response = await get_users(
startIndex=1,
count=10,
filter='userName eq "okta.user@example.com"',
)
expected_where = {
"OR": [
{"user_email": "okta.user@example.com"},
{"user_id": "okta.user@example.com"},
]
}
mock_prisma_client.db.litellm_usertable.find_many.assert_awaited_once_with(
where=expected_where,
skip=0,
take=10,
order={"created_at": "desc"},
)
mock_prisma_client.db.litellm_usertable.count.assert_awaited_once_with(
where=expected_where
)
assert response.totalResults == 1
assert response.Resources[0].id == "internal-user-id"
@pytest.mark.asyncio
async def test_get_users_filters_email_value_by_user_email(mocker):
"""
SCIM clients can locate users with `emails.value eq "<email>"`; keep that
filter as a direct user_email lookup alongside the userName fallback query.
"""
user = LiteLLM_UserTable(
user_id="internal-user-id",
user_email="scim.user@example.com",
user_alias="SCIM User",
teams=[],
metadata={},
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[user])
mock_prisma_client.db.litellm_usertable.count = AsyncMock(return_value=1)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
AsyncMock(
return_value=SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
id="internal-user-id",
userName="scim.user@example.com",
emails=[SCIMUserEmail(value="scim.user@example.com")],
)
),
)
response = await get_users(
startIndex=1,
count=10,
filter='emails.value eq "scim.user@example.com"',
)
expected_where = {"user_email": "scim.user@example.com"}
mock_prisma_client.db.litellm_usertable.find_many.assert_awaited_once_with(
where=expected_where,
skip=0,
take=10,
order={"created_at": "desc"},
)
mock_prisma_client.db.litellm_usertable.count.assert_awaited_once_with(
where=expected_where
)
assert response.totalResults == 1
assert response.Resources[0].id == "internal-user-id"
@pytest.mark.asyncio
async def test_handle_existing_user_by_email_no_email(mocker):
"""Should return None when new_user_request has no email"""
@ -1337,7 +1456,7 @@ async def test_create_group_with_nonexistent_users_creates_when_flag_true(
)
# Execute the create_group function - should succeed
result = await create_group(group=scim_group)
await create_group(group=scim_group)
# Verify users were created
assert mock_create_user.call_count == 2

View file

@ -133,6 +133,87 @@ class TestProxyInitializationHelpers:
)
assert args["timeout_worker_healthcheck"] == 15
def test_get_reload_options_no_config(self):
opts = ProxyInitializationHelpers._get_reload_options(None)
assert opts == {"reload": True}
def test_get_reload_options_with_config_in_cwd(self, tmp_path, monkeypatch):
config_file = tmp_path / "config.yaml"
config_file.write_text("model_list: []\n")
monkeypatch.chdir(tmp_path)
opts = ProxyInitializationHelpers._get_reload_options("config.yaml")
assert opts["reload"] is True
assert opts["reload_dirs"] == [str(tmp_path)]
assert opts["reload_includes"] == ["*.py", "config.yaml"]
def test_get_reload_options_with_config_outside_cwd(self, tmp_path, monkeypatch):
cwd_dir = tmp_path / "work"
cwd_dir.mkdir()
elsewhere = tmp_path / "configs"
elsewhere.mkdir()
config_file = elsewhere / "proxy.yaml"
config_file.write_text("model_list: []\n")
monkeypatch.chdir(cwd_dir)
opts = ProxyInitializationHelpers._get_reload_options(str(config_file))
assert opts["reload"] is True
assert opts["reload_dirs"] == [str(cwd_dir), str(elsewhere)]
assert opts["reload_includes"] == ["*.py", "proxy.yaml"]
def test_patch_statreload_for_config_yields_yaml(self, tmp_path):
from pathlib import Path
from uvicorn.supervisors.statreload import StatReload
if hasattr(StatReload, "_litellm_patched_config_paths"):
StatReload._litellm_patched_config_paths.clear()
config_file = tmp_path / "config.yaml"
config_file.write_text("model_list: []\n")
py_file = tmp_path / "module.py"
py_file.write_text("x = 1\n")
applied = ProxyInitializationHelpers._patch_statreload_for_config(
str(config_file)
)
assert applied is True
fake_self = types.SimpleNamespace(
config=types.SimpleNamespace(reload_dirs=[tmp_path])
)
yielded_paths = {Path(p).resolve() for p in StatReload.iter_py_files(fake_self)}
assert config_file.resolve() in yielded_paths
assert py_file.resolve() in yielded_paths
def test_patch_statreload_for_config_is_idempotent(self, tmp_path):
from pathlib import Path
from uvicorn.supervisors.statreload import StatReload
if hasattr(StatReload, "_litellm_patched_config_paths"):
StatReload._litellm_patched_config_paths.clear()
config_file = tmp_path / "config.yaml"
config_file.write_text("model_list: []\n")
py_file = tmp_path / "only.py"
py_file.write_text("x = 1\n")
for _ in range(3):
ProxyInitializationHelpers._patch_statreload_for_config(str(config_file))
fake_self = types.SimpleNamespace(
config=types.SimpleNamespace(reload_dirs=[tmp_path])
)
yielded = list(StatReload.iter_py_files(fake_self))
assert len(yielded) == len(set(map(str, yielded)))
yielded_paths = {Path(p).resolve() for p in yielded}
assert config_file.resolve() in yielded_paths
assert py_file.resolve() in yielded_paths
@patch("asyncio.run")
@patch("builtins.print")
def test_init_hypercorn_server(self, mock_print, mock_asyncio_run):

View file

@ -331,6 +331,172 @@ async def test_delete_old_logs_continues_on_valid_int_return():
assert total_deleted == 800
@pytest.mark.asyncio
async def test_delete_old_logs_continues_after_single_batch_failure(monkeypatch):
"""A single batch failure (e.g. DB timeout) must not abort the whole run —
subsequent batches should still execute and their counts accumulate."""
import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module
# Zero out the failure backoff so the test doesn't take ~0.5s of real sleep.
monkeypatch.setattr(
cleanup_module, "SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.0
)
mock_prisma_client = MagicMock()
mock_db = MagicMock()
# batch 1 succeeds, batch 2 raises (one-off DB timeout), batches 3-4 succeed,
# batch 5 returns 0 → loop exits naturally.
mock_db.execute_raw = AsyncMock(
side_effect=[100, TimeoutError("simulated DB timeout"), 200, 50, 0]
)
mock_prisma_client.db = mock_db
cleaner = cleanup_module.SpendLogCleanup(
general_settings={"maximum_spend_logs_retention_period": "7d"}
)
cutoff_date = datetime.now(timezone.utc) - timedelta(days=7)
total_deleted = await cleaner._delete_old_logs(mock_prisma_client, cutoff_date)
# All 5 batches should have been attempted; 100 + 200 + 50 = 350 deleted.
assert mock_db.execute_raw.call_count == 5
assert total_deleted == 350
@pytest.mark.asyncio
async def test_delete_old_logs_aborts_after_consecutive_failures(monkeypatch):
"""If batch failures persist for SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES
in a row (e.g. DB is down), the loop must abort instead of hot-looping."""
import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module
# Lower the threshold so the test is fast and deterministic.
monkeypatch.setattr(cleanup_module, "SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3)
monkeypatch.setattr(
cleanup_module, "SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.0
)
mock_prisma_client = MagicMock()
mock_db = MagicMock()
# Every batch raises — must abort after exactly 3 attempts, not loop forever.
mock_db.execute_raw = AsyncMock(
side_effect=ConnectionError("simulated persistent DB outage")
)
mock_prisma_client.db = mock_db
cleaner = cleanup_module.SpendLogCleanup(
general_settings={"maximum_spend_logs_retention_period": "7d"}
)
cutoff_date = datetime.now(timezone.utc) - timedelta(days=7)
total_deleted = await cleaner._delete_old_logs(mock_prisma_client, cutoff_date)
assert mock_db.execute_raw.call_count == 3
assert total_deleted == 0
@pytest.mark.asyncio
async def test_delete_old_logs_resets_consecutive_failures_on_success(monkeypatch):
"""A success between failures must reset the consecutive-failure counter so
intermittent timeouts don't trip the abort threshold."""
import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module
monkeypatch.setattr(cleanup_module, "SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3)
monkeypatch.setattr(
cleanup_module, "SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.0
)
mock_prisma_client = MagicMock()
mock_db = MagicMock()
# Pattern: fail, fail, success (resets counter), fail, fail, success, done.
# Without reset, three of these would trip abort; with reset, they don't.
mock_db.execute_raw = AsyncMock(
side_effect=[
TimeoutError("t1"),
TimeoutError("t2"),
100,
TimeoutError("t3"),
TimeoutError("t4"),
50,
0,
]
)
mock_prisma_client.db = mock_db
cleaner = cleanup_module.SpendLogCleanup(
general_settings={"maximum_spend_logs_retention_period": "7d"}
)
cutoff_date = datetime.now(timezone.utc) - timedelta(days=7)
total_deleted = await cleaner._delete_old_logs(mock_prisma_client, cutoff_date)
assert mock_db.execute_raw.call_count == 7
assert total_deleted == 150
@pytest.mark.asyncio
async def test_cleanup_uses_logger_exception_for_full_traceback(monkeypatch):
"""The outer error handler must call logger.exception() (not .error(str(e)))
so Prisma/DB timeouts surface a full traceback and exception type."""
import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module
mock_logger = MagicMock()
monkeypatch.setattr(cleanup_module, "verbose_proxy_logger", mock_logger)
mock_prisma_client = MagicMock()
# Force the outer try/except to fire by making _should_delete_spend_logs raise.
cleaner = cleanup_module.SpendLogCleanup(
general_settings={"maximum_spend_logs_retention_period": "7d"}
)
cleaner.pod_lock_manager = None
def boom():
raise RuntimeError("simulated prisma timeout")
cleaner._should_delete_spend_logs = boom # type: ignore[assignment]
await cleaner.cleanup_old_spend_logs(mock_prisma_client)
assert mock_logger.exception.called, "expected logger.exception() to be called"
# The exception type name must appear in the formatted args so operators can
# tell *what* failed, not just "Error during cleanup:".
call_args = mock_logger.exception.call_args
formatted = call_args[0][0] % call_args[0][1:]
assert "RuntimeError" in formatted
assert "simulated prisma timeout" in formatted
@pytest.mark.asyncio
async def test_cleanup_releases_lock_after_persistent_batch_failures(monkeypatch):
"""Even when batch deletion aborts due to consecutive failures, the pod lock
must still be released so the next scheduled run isn't permanently blocked."""
import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module
monkeypatch.setattr(cleanup_module, "SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 2)
monkeypatch.setattr(
cleanup_module, "SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.0
)
mock_prisma_client = MagicMock()
mock_db = MagicMock()
mock_db.execute_raw = AsyncMock(side_effect=TimeoutError("DB down"))
mock_prisma_client.db = mock_db
mock_pod_lock_manager = MagicMock()
mock_pod_lock_manager.redis_cache = MagicMock()
mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True)
mock_pod_lock_manager.release_lock = AsyncMock()
cleaner = cleanup_module.SpendLogCleanup(
general_settings={"maximum_spend_logs_retention_period": "7d"}
)
cleaner.pod_lock_manager = mock_pod_lock_manager
await cleaner.cleanup_old_spend_logs(mock_prisma_client)
# Cleanup didn't crash; the abort-after-failures path returned cleanly.
mock_pod_lock_manager.release_lock.assert_awaited_once()
def test_cleanup_batch_size_env_var(monkeypatch):
"""Ensure batch size is configurable via environment variable"""
import importlib