mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
Merge branch 'BerriAI:litellm_internal_staging' into litellm_internal_staging
This commit is contained in:
commit
ef1737e456
43 changed files with 3173 additions and 202 deletions
1
.github/workflows/test-unit-proxy-db.yml
vendored
1
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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
3
.gitignore
vendored
|
|
@ -100,4 +100,5 @@ STABILIZATION_TODO.md
|
|||
**/playwright-report
|
||||
**/*.storageState.json
|
||||
**/coverage
|
||||
test-config
|
||||
test-config
|
||||
.vscode
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ""),
|
||||
)
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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=(
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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] = (
|
||||
|
|
|
|||
121
litellm/proxy/middleware/request_size_limit_middleware.py
Normal file
121
litellm/proxy/middleware/request_size_limit_middleware.py
Normal 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})
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
276
scripts/mock_bedrock_passthrough_target.py
Normal file
276
scripts/mock_bedrock_passthrough_target.py
Normal 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()
|
||||
29
tests/litellm/test_sambanova_model_metadata.py
Normal file
29
tests/litellm/test_sambanova_model_metadata.py
Normal 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"
|
||||
135
tests/proxy_unit_tests/test_request_size_limit_middleware.py
Normal file
135
tests/proxy_unit_tests/test_request_size_limit_middleware.py
Normal 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
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue