diff --git a/litellm/__init__.py b/litellm/__init__.py index e6c30e12286..e8422952ab4 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -43,6 +43,7 @@ from typing import ( Type, ) from litellm.types.integrations.datadog import DatadogInitParams +from litellm.types.integrations.newrelic import NewRelicInitParams from litellm._logging import ( set_verbose, _turn_on_debug, @@ -154,10 +155,12 @@ _custom_logger_compatible_callbacks_literal = Literal[ "gitlab", "cloudzero", "focus", + "mavvrik", "vantage", "posthog", "levo", "compression_interception", + "newrelic", ] cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None @@ -412,6 +415,7 @@ s3_callback_params: Optional[Dict] = None s3_audit_callback_params: Optional[Dict] = None datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]] = None datadog_params: Optional[Union[DatadogInitParams, Dict]] = None +newrelic_params: Optional[Union[NewRelicInitParams, Dict]] = None aws_sqs_callback_params: Optional[Dict] = None generic_logger_headers: Optional[Dict] = None default_key_generate_params: Optional[Dict] = None @@ -1373,6 +1377,7 @@ from .search.main import * from .realtime_api.main import ( _arealtime, acreate_realtime_client_secret, + acreate_realtime_transcription_session, arealtime_calls, ) from .responses.main import _aresponses_websocket diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 87c26b776e8..d27cfefda73 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -171,6 +171,8 @@ class ResponsesToCompletionBridgeHandler: model_response = validated_kwargs["model_response"] logging_obj = validated_kwargs["logging_obj"] custom_llm_provider = validated_kwargs["custom_llm_provider"] + if kwargs.get("stream") is True and "stream" not in optional_params: + optional_params = {**optional_params, "stream": True} request_data = self.transformation_handler.transform_request( model=model, @@ -263,6 +265,8 @@ class ResponsesToCompletionBridgeHandler: model_response = validated_kwargs["model_response"] logging_obj = validated_kwargs["logging_obj"] custom_llm_provider = validated_kwargs["custom_llm_provider"] + if kwargs.get("stream") is True and "stream" not in optional_params: + optional_params = {**optional_params, "stream": True} try: request_data = self.transformation_handler.transform_request( diff --git a/litellm/constants.py b/litellm/constants.py index a5e3926aa8b..663afb87fb5 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -421,6 +421,7 @@ REPLICATE_POLLING_DELAY_SECONDS = float( DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS = int( os.getenv("DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS", 4096) ) +DEFAULT_OCI_CHAT_MAX_TOKENS = 4096 TOGETHER_AI_4_B = int(os.getenv("TOGETHER_AI_4_B", 4)) TOGETHER_AI_8_B = int(os.getenv("TOGETHER_AI_8_B", 8)) TOGETHER_AI_21_B = int(os.getenv("TOGETHER_AI_21_B", 21)) @@ -1483,6 +1484,7 @@ DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job" DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME = "db_daily_tag_spend_update_job" PROMETHEUS_EMIT_BUDGET_METRICS_JOB_NAME = "prometheus_emit_budget_metrics" CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME = "cloudzero_export_usage_data" +MAVVRIK_FOCUS_EXPORT_JOB_NAME = "mavvrik_focus_export_usage_data" CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int( os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000) ) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 88029615ba8..e934c6a6f83 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -2488,6 +2488,11 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor): ) +_TRANSCRIPTION_COMPLETED_EVENT_TYPE = ( + "conversation.item.input_audio_transcription.completed" +) + + def handle_realtime_stream_cost_calculation( results: OpenAIRealtimeStreamList, combined_usage_object: Usage, @@ -2533,4 +2538,99 @@ def handle_realtime_stream_cost_calculation( break # exit if we find a valid model total_cost = input_cost_per_token + output_cost_per_token + if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results): + total_cost += handle_realtime_transcription_cost_calculation( + results=results, + custom_llm_provider=custom_llm_provider, + litellm_model_name=litellm_model_name, + ) + return total_cost + + +def handle_realtime_transcription_cost_calculation( + results: OpenAIRealtimeStreamList, + custom_llm_provider: str, + litellm_model_name: str, +) -> float: + """ + Cost for realtime transcription sessions (e.g. gpt-realtime-whisper). + + Transcription sessions emit no `response.done` events; instead each + `conversation.item.input_audio_transcription.completed` event carries a + `usage` object billed by the ASR model. The usage is one of: + - {"type": "duration", "seconds": } → priced via input_cost_per_second + - {"type": "tokens", "input_tokens": ...} → priced via input/audio token cost + """ + completed_events = [ + cast(dict, result) + for result in results + if result.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE + ] + if not completed_events: + return 0.0 + + model_name = ( + _get_transcription_model_name_from_results(results) or litellm_model_name + ) + try: + model_info = litellm.get_model_info( + model=model_name, custom_llm_provider=custom_llm_provider + ) + except Exception: + model_info = None + + total_cost = 0.0 + for event in completed_events: + usage = event.get("usage") or {} + total_cost += _transcription_usage_cost(usage, model_info) + return total_cost + + +def _get_transcription_model_name_from_results( + results: OpenAIRealtimeStreamList, +) -> Optional[str]: + """Resolve the ASR model from a transcription_session.* / session.* event.""" + for result in results: + if result.get("type") in ( + "transcription_session.created", + "transcription_session.updated", + "session.created", + "session.updated", + ): + session = cast(dict, result).get("session", {}) or {} + transcription = ( + (session.get("audio", {}) or {}).get("input", {}) or {} + ).get("transcription", {}) or session.get("input_audio_transcription", {}) + model = (transcription or {}).get("model") or session.get("model") + if model: + return model + return None + + +def _transcription_usage_cost(usage: dict, model_info: Optional[ModelInfo]) -> float: + if model_info is None: + return 0.0 + usage_type = usage.get("type") + if usage_type == "duration": + seconds = usage.get("seconds") or 0.0 + per_second = model_info.get("input_cost_per_second") or 0.0 + return float(seconds) * float(per_second) + if usage_type == "tokens": + input_token_details = usage.get("input_token_details") or {} + audio_tokens = input_token_details.get("audio_tokens") or 0 + text_tokens = input_token_details.get("text_tokens") or 0 + output_tokens = usage.get("output_tokens") or 0 + audio_cost = float(audio_tokens) * float( + model_info.get("input_cost_per_audio_token") + or model_info.get("input_cost_per_token") + or 0.0 + ) + text_cost = float(text_tokens) * float( + model_info.get("input_cost_per_token") or 0.0 + ) + output_cost = float(output_tokens) * float( + model_info.get("output_cost_per_token") or 0.0 + ) + return audio_cost + text_cost + output_cost + return 0.0 diff --git a/litellm/integrations/SlackAlerting/budget_alert_types.py b/litellm/integrations/SlackAlerting/budget_alert_types.py index ea80b258540..2a19ec0b7fa 100644 --- a/litellm/integrations/SlackAlerting/budget_alert_types.py +++ b/litellm/integrations/SlackAlerting/budget_alert_types.py @@ -1,7 +1,7 @@ from abc import ABC, abstractmethod from typing import Literal -from litellm.proxy._types import CallInfo +from litellm.proxy._types import CallInfo, Litellm_EntityType class BaseBudgetAlertType(ABC): @@ -31,6 +31,8 @@ class SoftBudgetAlert(BaseBudgetAlertType): return "Soft Budget Crossed: " def get_id(self, user_info: CallInfo) -> str: + if user_info.event_group == Litellm_EntityType.TEAM: + return user_info.team_id or "default_id" return user_info.token or "default_id" diff --git a/litellm/integrations/callback_configs.json b/litellm/integrations/callback_configs.json index 3a69c9a7936..590c848767a 100644 --- a/litellm/integrations/callback_configs.json +++ b/litellm/integrations/callback_configs.json @@ -290,6 +290,21 @@ }, "description": "Langsmith Logging Integration" }, + { + "id": "newrelic", + "displayName": "New Relic", + "logo": "newrelic.png", + "supports_key_team_logging": false, + "dynamic_params": { + "NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED": { + "type": "text", + "ui_name": "Record AI Content (default: true)", + "description": "Whether to record AI message content. Set to false to disable.", + "required": false + } + }, + "description": "New Relic AI Monitoring Integration" + }, { "id": "openmeter", "displayName": "OpenMeter", diff --git a/litellm/integrations/focus/destinations/__init__.py b/litellm/integrations/focus/destinations/__init__.py index e0cd90c1d61..21945c9b457 100644 --- a/litellm/integrations/focus/destinations/__init__.py +++ b/litellm/integrations/focus/destinations/__init__.py @@ -4,6 +4,7 @@ from .base import FocusDestination, FocusTimeWindow from .factory import FocusDestinationFactory from .gcs_destination import FocusGCSDestination from .s3_destination import FocusS3Destination +from .mavvrik_destination import FocusMavvrikDestination from .vantage_destination import FocusVantageDestination __all__ = [ @@ -12,5 +13,6 @@ __all__ = [ "FocusGCSDestination", "FocusTimeWindow", "FocusS3Destination", + "FocusMavvrikDestination", "FocusVantageDestination", ] diff --git a/litellm/integrations/focus/destinations/factory.py b/litellm/integrations/focus/destinations/factory.py index 7ce21d4040a..cd25a87729f 100644 --- a/litellm/integrations/focus/destinations/factory.py +++ b/litellm/integrations/focus/destinations/factory.py @@ -8,6 +8,7 @@ from typing import Any, Dict, Optional from .base import FocusDestination from .gcs_destination import FocusGCSDestination from .s3_destination import FocusS3Destination +from .mavvrik_destination import FocusMavvrikDestination from .vantage_destination import FocusVantageDestination @@ -32,6 +33,8 @@ class FocusDestinationFactory: return FocusVantageDestination(prefix=prefix, config=normalized_config) if provider_lower == "gcs": return FocusGCSDestination(prefix=prefix, config=normalized_config) + if provider_lower == "mavvrik": + return FocusMavvrikDestination(prefix=prefix, config=normalized_config) raise NotImplementedError( f"Provider '{provider}' not supported for Focus export" ) @@ -87,6 +90,15 @@ class FocusDestinationFactory: "FOCUS_GCS_BUCKET_NAME must be provided for GCS exports" ) return {k: v for k, v in resolved.items() if v is not None} + if provider == "mavvrik": + resolved = { + "api_key": overrides.get("api_key") or os.getenv("MAVVRIK_API_KEY"), + "api_endpoint": overrides.get("api_endpoint") + or os.getenv("MAVVRIK_API_ENDPOINT"), + "connection_id": overrides.get("connection_id") + or os.getenv("MAVVRIK_CONNECTION_ID"), + } + return {k: v for k, v in resolved.items() if v is not None} raise NotImplementedError( f"Provider '{provider}' not supported for Focus export configuration" ) diff --git a/litellm/integrations/focus/destinations/mavvrik_destination.py b/litellm/integrations/focus/destinations/mavvrik_destination.py new file mode 100644 index 00000000000..1e3c98b9a70 --- /dev/null +++ b/litellm/integrations/focus/destinations/mavvrik_destination.py @@ -0,0 +1,345 @@ +"""Mavvrik GCS destination for FOCUS export. + +Flow: + 1. GET /metrics/agent/ai/{connection_id}/upload-url → GCS signed URL + 2. PUT with CSV content +""" + +from __future__ import annotations + +import gzip +from typing import Any, Optional +from urllib.parse import urlparse + +from litellm._logging import verbose_logger +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + get_async_httpx_client, + httpxSpecialProvider, +) + +from .base import FocusDestination, FocusTimeWindow + +_MAVVRIK_ALLOWED_SUFFIXES = (".mavvrik.dev", ".mavvrik.ai", ".mavvrik.app") + +# GCS requires intermediate chunks to be a multiple of 256 KB. +# 8 MB gives a good balance between round-trips and memory pressure. +_GCS_CHUNK_SIZE = 8 * 1024 * 1024 # 8 MB + + +def _validate_api_endpoint(api_endpoint: str) -> None: + if not api_endpoint.startswith("https://"): + raise ValueError("MAVVRIK_API_ENDPOINT must be an HTTPS URL") + hostname = (urlparse(api_endpoint).hostname or "").lower() + if not any(hostname.endswith(suffix) for suffix in _MAVVRIK_ALLOWED_SUFFIXES): + raise ValueError( + "MAVVRIK_API_ENDPOINT host must be a Mavvrik domain " + "(e.g. https://api.mavvrik.dev/)" + ) + + +def _validate_gcs_url(url: str, label: str) -> None: + parsed = urlparse(url) + if parsed.scheme != "https": + raise ValueError( + f"Mavvrik FOCUS destination: {label} must be HTTPS, got scheme '{parsed.scheme}'" + ) + hostname = (parsed.hostname or "").lower() + if not ( + hostname == "storage.googleapis.com" + or hostname.endswith(".storage.googleapis.com") + ): + raise ValueError( + f"Mavvrik FOCUS destination: {label} must be a GCS endpoint " + f"(storage.googleapis.com), got '{hostname}'" + ) + + +class FocusMavvrikDestination(FocusDestination): + """Upload FOCUS CSV exports to Mavvrik via GCS signed URL.""" + + def __init__( + self, + *, + prefix: str, + config: Optional[dict[str, Any]] = None, + ) -> None: + config = config or {} + api_key = config.get("api_key") + api_endpoint = config.get("api_endpoint") + connection_id = config.get("connection_id") + + if not api_key: + raise ValueError( + "MAVVRIK_API_KEY must be provided for Mavvrik FOCUS destination " + "(set MAVVRIK_API_KEY env var or pass in destination_config)" + ) + if not api_endpoint: + raise ValueError( + "MAVVRIK_API_ENDPOINT must be provided for Mavvrik FOCUS destination " + "(set MAVVRIK_API_ENDPOINT env var or pass in destination_config)" + ) + if not connection_id: + raise ValueError( + "MAVVRIK_CONNECTION_ID must be provided for Mavvrik FOCUS destination " + "(set MAVVRIK_CONNECTION_ID env var or pass in destination_config)" + ) + + _validate_api_endpoint(api_endpoint) + + self.api_key = api_key + self.api_endpoint = api_endpoint.rstrip("/") + self.connection_id = connection_id + self.prefix = prefix + self._http: AsyncHTTPHandler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.LoggingCallback + ) + self._registered = False + + @property + def _agent_url(self) -> str: + return f"{self.api_endpoint}/metrics/agent/ai/{self.connection_id}" + + @property + def _upload_url_endpoint(self) -> str: + return f"{self.api_endpoint}/metrics/agent/ai/{self.connection_id}/upload-url" + + @property + def _auth_headers(self) -> dict[str, str]: + return {"Content-Type": "application/json", "x-api-key": self.api_key} + + async def _ensure_registered(self) -> Optional[int]: + """POST agent endpoint to register/initialize the connector (once per instance). + + Returns metricsMarker from the Mavvrik response — the last date index + Mavvrik has successfully processed. Used by the logger to catch up any + dates that were missed due to previous export failures. + + Returns None if the connector was already registered (cached). + """ + if self._registered: + return None + resp = await self._http.client.request( + method="POST", + url=self._agent_url, + headers=self._auth_headers, + json={"name": self.connection_id}, + timeout=30.0, + ) + if resp.status_code == 410: + # Connector has been disconnected in Mavvrik — reset flag so next + # delivery attempt re-registers after it becomes active again. + self._registered = False + raise RuntimeError( + "Mavvrik FOCUS destination: connector is disconnected (410). " + "Re-enable the connection in the Mavvrik dashboard." + ) + if resp.status_code >= 400: + raise RuntimeError( + f"Mavvrik FOCUS destination: register failed " + f"({resp.status_code}): {resp.text[:200]}" + ) + self._registered = True + metrics_marker = resp.json().get("metricsMarker", 0) + verbose_logger.debug( + "Mavvrik FOCUS destination: connector registered (metricsMarker=%s)", + metrics_marker, + ) + return metrics_marker + + async def _get_signed_url(self, date_str: str) -> str: + """GET upload-url endpoint → GCS signed URL for the given date.""" + params = {"name": date_str, "type": "metrics", "datetime": date_str} + resp = await self._http.client.request( + method="GET", + url=self._upload_url_endpoint, + headers=self._auth_headers, + params=params, + timeout=30.0, + ) + if resp.status_code >= 400: + raise RuntimeError( + f"Mavvrik FOCUS destination: failed to get signed URL " + f"({resp.status_code}): {resp.text[:200]}" + ) + signed_url = resp.json().get("url") + if not signed_url: + raise RuntimeError( + f"Mavvrik FOCUS destination: response missing 'url' field: {resp.json()}" + ) + _validate_gcs_url(signed_url, "signed URL") + verbose_logger.debug( + "Mavvrik FOCUS destination: got signed URL for date %s", date_str + ) + return signed_url + + async def _upload_to_gcs(self, signed_url: str, content: bytes) -> None: + """Upload gzip-compressed CSV to GCS via chunked resumable upload. + + The full CSV is gzip-compressed first, then uploaded in _GCS_CHUNK_SIZE + chunks using the GCS resumable upload protocol. GCS assembles the chunks + server-side into a single complete object — the bucket receives one file + regardless of how many chunks were sent. + + Intermediate chunks: Content-Range: bytes X-Y/* → expect 308 + Final chunk: Content-Range: bytes X-Y/T → expect 200/201 + + This handles exports larger than available memory for a single PUT while + keeping the destination code self-contained (no changes to the FOCUS + pipeline upstream). + """ + gzip_bytes = gzip.compress(content) + total = len(gzip_bytes) + + # Step 1: initiate resumable upload session + metadata = b'{"contentEncoding":"gzip","contentDisposition":"attachment"}' + init_resp = await self._http.client.request( + method="POST", + url=signed_url, + headers={ + "Content-Type": "application/gzip", + "x-goog-resumable": "start", + }, + content=metadata, + timeout=30.0, + ) + if init_resp.status_code not in (200, 201): + raise RuntimeError( + f"Mavvrik FOCUS destination: GCS session init failed " + f"({init_resp.status_code}): {init_resp.text[:400]}" + ) + + session_uri = init_resp.headers.get("Location") + if not session_uri: + raise RuntimeError( + "Mavvrik FOCUS destination: GCS session init missing Location header" + ) + _validate_gcs_url(session_uri, "session URI") + + verbose_logger.debug( + "Mavvrik FOCUS destination: GCS session started, uploading %d gzip bytes " + "in %d chunk(s)", + total, + max(1, -(-total // _GCS_CHUNK_SIZE)), # ceiling division + ) + + # Step 2: upload in chunks; cancel session on any failure to avoid + # lingering GCS sessions (they stay open for ~1 week otherwise). + offset = 0 + try: + while offset < total: + chunk = gzip_bytes[offset : offset + _GCS_CHUNK_SIZE] + chunk_end = offset + len(chunk) - 1 + is_final = (offset + len(chunk)) >= total + content_range = ( + f"bytes {offset}-{chunk_end}/{total}" + if is_final + else f"bytes {offset}-{chunk_end}/*" + ) + expected_statuses = {200, 201} if is_final else {308} + + resp = await self._http.client.request( + method="PUT", + url=session_uri, + headers={ + "Content-Type": "application/gzip", + "Content-Range": content_range, + }, + content=chunk, + timeout=120.0, + ) + if resp.status_code not in expected_statuses: + raise RuntimeError( + f"Mavvrik FOCUS destination: GCS chunk upload failed " + f"(chunk offset={offset}, expected={expected_statuses}, " + f"got={resp.status_code}): {resp.text[:400]}" + ) + offset += len(chunk) + verbose_logger.debug( + "Mavvrik FOCUS destination: uploaded chunk offset=%d/%d", + offset, + total, + ) + except Exception: + # Cancel the open GCS session so it doesn't linger for up to 1 week. + try: + await self._http.client.request( + method="DELETE", url=session_uri, timeout=10.0 + ) + verbose_logger.debug( + "Mavvrik FOCUS destination: cancelled GCS session after error" + ) + except Exception: + pass + raise + + async def get_metrics_marker(self) -> Optional[int]: + """Register with Mavvrik and return the current metricsMarker. + + The metricsMarker is a Unix timestamp (seconds) representing the last + date Mavvrik has successfully ingested. Called on every scheduled run + so the logger can detect and catch up any dates missed due to previous + export failures. + + Always calls the Mavvrik register API — unlike deliver() which skips + registration once _registered is True, catch-up requires a fresh + marker value on every run. + """ + resp = await self._http.client.request( + method="POST", + url=self._agent_url, + headers=self._auth_headers, + json={"name": self.connection_id}, + timeout=30.0, + ) + if resp.status_code == 410: + self._registered = False + raise RuntimeError( + "Mavvrik FOCUS destination: connector is disconnected (410). " + "Re-enable the connection in the Mavvrik dashboard." + ) + if resp.status_code >= 400: + raise RuntimeError( + f"Mavvrik FOCUS destination: register failed " + f"({resp.status_code}): {resp.text[:200]}" + ) + self._registered = True + metrics_marker = resp.json().get("metricsMarker", 0) + verbose_logger.debug( + "Mavvrik FOCUS destination: got metricsMarker=%s", metrics_marker + ) + return metrics_marker + + async def deliver( + self, + *, + content: bytes, + time_window: FocusTimeWindow, + filename: str, + ) -> None: + """Upload FOCUS CSV to Mavvrik via GCS signed URL. + + Uses the start date of the time window as the object date key. + """ + if not content: + verbose_logger.debug( + "Mavvrik FOCUS destination: empty content, skipping upload" + ) + return + + date_str = time_window.start_time.strftime("%Y-%m-%d") + + verbose_logger.debug( + "Mavvrik FOCUS destination: uploading %d bytes for date=%s (%s)", + len(content), + date_str, + filename, + ) + + await self._ensure_registered() + signed_url = await self._get_signed_url(date_str) + await self._upload_to_gcs(signed_url, content) + + verbose_logger.debug( + "Mavvrik FOCUS destination: upload complete for date=%s", date_str + ) diff --git a/litellm/integrations/mavvrik_focus/__init__.py b/litellm/integrations/mavvrik_focus/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py new file mode 100644 index 00000000000..47d3e1da7bc --- /dev/null +++ b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py @@ -0,0 +1,272 @@ +"""MavvrikFocusLogger — FOCUS-based Mavvrik export logger. + +Usage in config.yaml: + litellm_settings: + callbacks: ["mavvrik"] + +Required env vars: + MAVVRIK_API_KEY + MAVVRIK_API_ENDPOINT + MAVVRIK_CONNECTION_ID + +Optional env vars: + MAVVRIK_FOCUS_MAX_ROWS — row cap per export window (default: 500000) + +Only daily frequency is supported. The Mavvrik ingestion protocol stores one +file per calendar date (metrics/YYYY-MM-DD). Hourly or interval exports would +overwrite each other within the same day, producing incomplete data. +""" + +from __future__ import annotations + +import os +from datetime import datetime, timedelta, timezone +from typing import TYPE_CHECKING, Any, List, Optional + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.constants import MAVVRIK_FOCUS_EXPORT_JOB_NAME +from litellm.integrations.focus.destinations.base import FocusTimeWindow +from litellm.integrations.focus.focus_logger import FocusLogger + +if TYPE_CHECKING: + from apscheduler.schedulers.asyncio import AsyncIOScheduler +else: + AsyncIOScheduler = Any + + +def _parse_metrics_marker( + marker: Optional[object], +) -> Optional[datetime]: + """Parse metricsMarker from Mavvrik register response into a UTC datetime. + + Handles both formats Mavvrik may return: + - Unix timestamp (int/float): e.g. 1749340800 + - ISO date string: e.g. "2026-06-09" or "2026-06-09T00:00:00Z" + + Returns None for falsy values (0, None, empty string) which indicate + no data has been ingested yet. + """ + if not marker: + return None + try: + if isinstance(marker, (int, float)): + return datetime.fromtimestamp(float(marker), tz=timezone.utc).replace( + hour=0, minute=0, second=0, microsecond=0 + ) + if isinstance(marker, str): + marker = marker.strip() + if not marker: + return None + # Try ISO date first (YYYY-MM-DD), then full ISO datetime + for fmt in ("%Y-%m-%d", "%Y-%m-%dT%H:%M:%SZ", "%Y-%m-%dT%H:%M:%S"): + try: + return datetime.strptime(marker, fmt).replace(tzinfo=timezone.utc) + except ValueError: + continue + except Exception: + pass + verbose_proxy_logger.warning( + "Mavvrik FOCUS: could not parse metricsMarker %r — skipping catch-up", marker + ) + return None + + +class MavvrikFocusLogger(FocusLogger): + """FOCUS-based export logger that routes to the Mavvrik destination.""" + + def __init__(self, **kwargs: Any) -> None: + frequency = os.getenv("MAVVRIK_FOCUS_FREQUENCY", "daily").lower() + if frequency != "daily": + raise ValueError( + f"MAVVRIK_FOCUS_FREQUENCY='{frequency}' is not supported. " + "Only 'daily' is allowed -- the Mavvrik ingestion protocol stores one " + "file per calendar date (metrics/YYYY-MM-DD). Hourly or interval " + "exports would overwrite each other within the same day." + ) + super().__init__( + provider="mavvrik", + export_format="csv", + frequency="daily", + prefix="mavvrik_focus_exports", + destination_config={ + "api_key": os.getenv("MAVVRIK_API_KEY"), + "api_endpoint": os.getenv("MAVVRIK_API_ENDPOINT"), + "connection_id": os.getenv("MAVVRIK_CONNECTION_ID"), + }, + **kwargs, + ) + raw = os.getenv("MAVVRIK_FOCUS_MAX_ROWS") + self._max_rows: Optional[int] = int(raw) if raw else 500_000 + + async def _export_window( + self, + *, + window: FocusTimeWindow, + limit: Optional[int], + ) -> None: + """Export with Mavvrik row cap applied when no explicit limit is passed.""" + effective_limit = limit if limit is not None else self._max_rows + engine = self._ensure_engine() + data = await engine._database.get_usage_data( + limit=effective_limit, + start_time_utc=window.start_time, + end_time_utc=window.end_time, + ) + if effective_limit is not None and len(data) >= effective_limit: + verbose_proxy_logger.warning( + "Mavvrik FOCUS export: row cap reached (%d rows). " + "Some data for window %s→%s may be excluded. " + "Increase MAVVRIK_FOCUS_MAX_ROWS to export all rows.", + effective_limit, + window.start_time.date(), + window.end_time.date(), + ) + if data.is_empty(): + verbose_proxy_logger.debug( + "Mavvrik FOCUS export: no usage data for window %s", window + ) + return + normalized = engine._transformer.transform(data) + if normalized.is_empty(): + return + payload = engine._serializer.serialize(normalized) + if not payload: + return + await engine._destination.deliver( + content=payload, + time_window=window, + filename=engine._build_filename(window), + ) + + # Maximum number of days to catch up in a single run. Prevents runaway + # loops if the connector was disabled for a long time, and avoids querying + # data that has likely been cleaned up from LiteLLM_DailyUserSpend. + _MAX_CATCHUP_DAYS = 7 + + async def _run_scheduled_export(self) -> None: + """Export today's window, catching up any dates Mavvrik has not yet received. + + On each run: + 1. Register with Mavvrik → get metricsMarker (last successfully ingested date) + 2. If metricsMarker is behind yesterday, catch up missed dates (capped at + _MAX_CATCHUP_DAYS to avoid runaway loops on long outages) + 3. Export yesterday (today's daily window) + + This ensures a failed export on day N is automatically retried on day N+1 + without any manual intervention. + """ + engine = self._ensure_engine() + from litellm.integrations.focus.destinations.mavvrik_destination import ( # noqa: PLC0415 + FocusMavvrikDestination, + ) + + destination = engine._destination + if not isinstance(destination, FocusMavvrikDestination): + await super()._run_scheduled_export() + return + + # Register and get the last date Mavvrik has processed. + # metricsMarker may be a Unix timestamp (int/float) or an ISO date string. + marker = await destination.get_metrics_marker() + + now = datetime.now(timezone.utc) + yesterday = now.replace(hour=0, minute=0, second=0, microsecond=0) - timedelta( + days=1 + ) + + last_ingested = _parse_metrics_marker(marker) + + # Catch up missed dates, capped at _MAX_CATCHUP_DAYS + if last_ingested and last_ingested < yesterday: + # Never go further back than _MAX_CATCHUP_DAYS from yesterday + earliest_catchup = yesterday - timedelta(days=self._MAX_CATCHUP_DAYS - 1) + catch_up_date = max(last_ingested + timedelta(days=1), earliest_catchup) + + if last_ingested + timedelta(days=1) < earliest_catchup: + verbose_proxy_logger.warning( + "Mavvrik FOCUS export: metricsMarker is more than %d days behind " + "(%s). Catching up from %s only; earlier data will not be re-exported.", + self._MAX_CATCHUP_DAYS, + last_ingested.date(), + catch_up_date.date(), + ) + + while catch_up_date < yesterday: + verbose_proxy_logger.info( + "Mavvrik FOCUS export: catching up missed date %s", + catch_up_date.date(), + ) + window = FocusTimeWindow( + start_time=catch_up_date, + end_time=catch_up_date + timedelta(days=1), + frequency="daily", + ) + await self._export_window(window=window, limit=None) + catch_up_date += timedelta(days=1) + + # Export yesterday's window (the normal daily run) + window = FocusTimeWindow( + start_time=yesterday, + end_time=yesterday + timedelta(days=1), + frequency="daily", + ) + await self._export_window(window=window, limit=None) + + async def initialize_mavvrik_focus_export_job(self) -> None: + """Scheduler entry point — uses Mavvrik-specific pod-lock key.""" + from litellm.proxy.proxy_server import proxy_logging_obj # noqa: PLC0415 + + pod_lock_manager = None + if proxy_logging_obj is not None: + writer = getattr(proxy_logging_obj, "db_spend_update_writer", None) + if writer is not None: + pod_lock_manager = getattr(writer, "pod_lock_manager", None) + + if pod_lock_manager and pod_lock_manager.redis_cache: + acquired = await pod_lock_manager.acquire_lock( + cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME + ) + if not acquired: + verbose_proxy_logger.debug( + "Mavvrik FOCUS export: unable to acquire pod lock" + ) + return + try: + await self._run_scheduled_export() + finally: + await pod_lock_manager.release_lock( + cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME + ) + else: + await self._run_scheduled_export() + + @staticmethod + async def init_mavvrik_focus_background_job( + scheduler: AsyncIOScheduler, + ) -> None: + """Register the Mavvrik FOCUS export job on the provided scheduler.""" + loggers: List[MavvrikFocusLogger] = [ + cb + for cb in litellm.logging_callback_manager.get_custom_loggers_for_type( + callback_type=MavvrikFocusLogger + ) + if type(cb) is MavvrikFocusLogger + ] + if not loggers: + verbose_proxy_logger.debug( + "No MavvrikFocusLogger registered; skipping scheduler" + ) + return + + logger = loggers[0] + trigger_kwargs = logger._build_scheduler_trigger() + scheduler.add_job( # type: ignore[attr-defined] + logger.initialize_mavvrik_focus_export_job, + id=MAVVRIK_FOCUS_EXPORT_JOB_NAME, + replace_existing=True, + **trigger_kwargs, + ) + verbose_proxy_logger.info( + "mavvrik_focus: background export job scheduled (%s)", trigger_kwargs + ) diff --git a/litellm/integrations/newrelic/__init__.py b/litellm/integrations/newrelic/__init__.py new file mode 100644 index 00000000000..5b0f5b9cb24 --- /dev/null +++ b/litellm/integrations/newrelic/__init__.py @@ -0,0 +1,10 @@ +""" +New Relic AI Monitoring Integration for LiteLLM + +This module provides integration with New Relic's AI Monitoring feature to track +LLM requests, responses, and usage metrics. +""" + +from litellm.integrations.newrelic.newrelic import NewRelicLogger + +__all__ = ["NewRelicLogger"] diff --git a/litellm/integrations/newrelic/newrelic.py b/litellm/integrations/newrelic/newrelic.py new file mode 100644 index 00000000000..753b8520337 --- /dev/null +++ b/litellm/integrations/newrelic/newrelic.py @@ -0,0 +1,926 @@ +""" +New Relic AI Monitoring Integration for LiteLLM + +This module provides integration with New Relic's AI Monitoring feature to track +LLM requests, responses, and usage metrics. + +Environment Variables (consumed by the New Relic agent at process bootstrap - +set via container env, or before invoking `newrelic-admin run-program`): + NEW_RELIC_LICENSE_KEY: Your New Relic license key (required) + NEW_RELIC_APP_NAME: Your application name (required) + +UI- and runtime-toggleable: + NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED: Whether to record message + content (optional, default: true) + +Configuration: + Message logging can be controlled via (both must agree to record): + 1. turn_off_message_logging parameter - pass via callback initialization or config YAML + 2. NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED env var + + Default behavior: Messages ARE recorded unless explicitly disabled by either method + Either method can disable recording - both must enable for recording to occur + +Usage - Python SDK: + import litellm + litellm.callbacks = ["newrelic"] + + # Or with explicit configuration: + from litellm.integrations.newrelic import NewRelicLogger + litellm.callbacks = [NewRelicLogger(turn_off_message_logging=True)] + +Usage - Proxy Server (config.yaml): + litellm_settings: + callbacks: ["newrelic"] + newrelic_params: + turn_off_message_logging: true # Disable message content recording + + # Or disable via environment variable: + # export NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED=false + + # Ensure New Relic agent is initialized (use newrelic-admin or initialize manually) + # newrelic-admin run-program python your_app.py +""" + +import json +import os +import threading +import time +import uuid +from typing import Any, Dict, List, Optional, Tuple, Union + +import litellm +from litellm._logging import verbose_logger +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.redact_messages import should_redact_message_logging +from litellm.types.integrations.newrelic import NewRelicInitParams +from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus +from litellm.types.utils import ModelResponse, Message, StandardLoggingPayload + +try: + import newrelic.agent as _newrelic_agent +except ImportError: + _newrelic_agent = None # type: ignore + + +class NewRelicLogger(CustomLogger): + """ + New Relic logger for LiteLLM to send AI monitoring events. + + This logger creates two types of New Relic custom events: + 1. LlmChatCompletionSummary - One per completion request + 2. LlmChatCompletionMessage - One per message (request and response) + """ + + # Class-level state for supportability metric emission, shared across all instances. + # Protected by _metric_lock to ensure thread-safe access. + _last_metric_emission_time: float = 0.0 + _metric_lock = threading.Lock() + + def __init__(self, **kwargs): + ######################################################### + # Handle newrelic_params set as litellm.newrelic_params + ######################################################### + dict_newrelic_params = self._get_newrelic_params() + + # Use setdefault so constructor kwargs take priority over global params. + # model_dump() always returns all fields (including defaults), so update() + # would silently overwrite explicit constructor args like turn_off_message_logging=True. + for k, v in dict_newrelic_params.items(): + kwargs.setdefault(k, v) + + # CustomLogger.__init__ will set self.turn_off_message_logging from kwargs + super().__init__(**kwargs) + + # Check for required environment variables + self.license_key = os.getenv("NEW_RELIC_LICENSE_KEY") + self.app_name = os.getenv("NEW_RELIC_APP_NAME") + + # Validate configuration + if not self.license_key or not self.app_name: + verbose_logger.warning( + "New Relic integration requires NEW_RELIC_LICENSE_KEY and " + "NEW_RELIC_APP_NAME environment variables. Integration will be disabled." + ) + self.enabled = False + elif _newrelic_agent is None: + verbose_logger.error( + "New Relic Python agent not installed. Review the New Relic integration documentation at https://docs.litellm.ai/docs/observability/newrelic." + ) + self.enabled = False + else: + try: + # timeout=0 forces non-blocking startup: the agent connects in a + # background thread regardless of newrelic.ini / NEW_RELIC_STARTUP_TIMEOUT. + _newrelic_agent.register_application(timeout=0) + + self.enabled = True + verbose_logger.info( + f"New Relic AI Monitoring initialized for app: {self.app_name}, " + f"content recording: {self.record_content}" + ) + except Exception as e: + verbose_logger.error( + f"Failed to initialize New Relic agent: {e}. " + "Integration will be disabled." + ) + self.enabled = False + + def _get_newrelic_params(self) -> Dict: + """ + Get the newrelic_params from litellm.newrelic_params + + These are params specific to initializing the NewRelicLogger e.g. turn_off_message_logging + """ + dict_newrelic_params: Dict = {} + if litellm.newrelic_params is not None: + if isinstance(litellm.newrelic_params, NewRelicInitParams): + dict_newrelic_params = litellm.newrelic_params.model_dump() + elif isinstance(litellm.newrelic_params, Dict): + # only allow params that are of NewRelicInitParams + dict_newrelic_params = NewRelicInitParams( + **litellm.newrelic_params + ).model_dump() + return dict_newrelic_params + + @property + def record_content(self) -> bool: + """Whether to record message content in New Relic. + + Both turn_off_message_logging param AND NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED + env var must agree to record content. If either disables recording, content will not + be recorded. Read at call time so UI config changes take effect without a restart. + Default: True (record content) unless explicitly disabled by either method. + """ + return (not self.turn_off_message_logging) and self._parse_bool_env( + "NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED", True + ) + + def _parse_bool_env(self, var_name: str, default: bool = False) -> bool: + """Parse a boolean environment variable. + + Accepts true/false, 1/0, yes/no, on/off (case-insensitive, + whitespace-tolerant) — matching the convention used in + ``litellm/__init__.py`` and the standard library's + ``configparser.BOOLEAN_STATES``. Unrecognised values log a + warning and fall back to ``default`` rather than silently + flipping user intent. + """ + raw = os.getenv(var_name) + if not raw: + return default + value = raw.strip().lower() + if value in ("1", "true", "yes", "on"): + return True + if value in ("0", "false", "no", "off"): + return False + verbose_logger.warning( + f"{var_name}={raw!r} is not a recognised boolean " + f"(accepts true/false, 1/0, yes/no, on/off). " + f"Falling back to default ({default})." + ) + return default + + def _get_litellm_version(self) -> str: + """ + Get litellm version for supportability metrics. + + Returns: + Version string (e.g., "1.80.0") or "unknown" if unable to determine + """ + try: + from importlib.metadata import version + + return version("litellm") + except Exception as e: + verbose_logger.warning(f"Unable to determine litellm version: {e}") + return "unknown" + + def _emit_supportability_metric(self): + """ + Emit New Relic supportability metric for LiteLLM usage. + + Per spec, this metric should be emitted at least once every 27 hours + to indicate the library is in use. Format: + Supportability/Python/ML/LiteLLM/{version} + + This method updates _last_metric_emission_time and should + be called within a lock when checking periodic emission. + """ + try: + litellm_version = self._get_litellm_version() + metric_name = f"Supportability/Python/ML/LiteLLM/{litellm_version}" + + # Record metric with value of 1 (will be aggregated by New Relic) + app = _newrelic_agent.application() + + # Always update the timestamp so the 27-hour back-off applies + # regardless of whether the app is ready, preventing lock contention + # on every request when the agent is slow to register or never starts. + NewRelicLogger._last_metric_emission_time = time.time() + + if app and app.enabled: + app.record_custom_metric(metric_name, 1) + verbose_logger.info( + f"Emitted New Relic supportability metric: {metric_name}" + ) + else: + verbose_logger.info( + "New Relic application is not enabled; skipping metric recording." + ) + + except Exception as e: + verbose_logger.warning(f"Failed to emit supportability metric: {e}") + + def _check_and_emit_periodic_metric(self): + """ + Check if 27 hours have passed since last metric emission and re-emit if needed. + + Uses a mutex to ensure only one thread emits the metric even if multiple + requests are being processed concurrently. + """ + # Quick check without lock to avoid unnecessary locking + current_time = time.time() + time_since_last_emission = ( + current_time - NewRelicLogger._last_metric_emission_time + ) + + if time_since_last_emission >= 97200: # 27 hours = 97200 seconds + # Acquire lock to ensure only one thread emits + with NewRelicLogger._metric_lock: + # Double-check inside lock in case another thread just emitted + current_time = time.time() + time_since_last_emission = ( + current_time - NewRelicLogger._last_metric_emission_time + ) + + if time_since_last_emission >= 97200: + self._emit_supportability_metric() + + def _get_trace_context( + self, + kwargs: Dict, + standard_logging_object: Optional[StandardLoggingPayload] = None, + ) -> str: + """ + Get the New Relic trace ID for AI monitoring events. + + This integration runs in LiteLLM's async logging worker, outside the + New Relic agent's current transaction. Because we can't call + `newrelic.agent.current_trace_id()` to let the agent populate the + trace_id on AIM custom events, we manually simulate what the agent + would do. An AIM event without a trace_id is malformed per the NR + schema, so this method always returns a valid string. + + Resolution order: + 1. W3C traceparent header (litellm_params.metadata.headers.traceparent) - + what the agent would link to if we were in-transaction. + 2. StandardLoggingPayload.trace_id - LiteLLM's internal trace for + retry/fallback grouping. + 3. Generated UUID - synthetic grouping key when upstream context is + absent or parsing it fails. + + Span IDs are intentionally not emitted: any span ID recoverable from + the inbound traceparent is the caller's parent span, not ours. + + Returns: + trace_id: always a non-empty string. + """ + trace_id: Optional[str] = None + try: + litellm_params = kwargs.get("litellm_params") or {} + metadata = litellm_params.get("metadata") or {} + headers = metadata.get("headers") or {} + # Normalize header key lookup to be case-insensitive per W3C spec + traceparent = next( + (v for k, v in headers.items() if k.lower() == "traceparent"), None + ) + + if traceparent: + # Extract trace_id from traceparent header if available + # traceparent format: "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00" + parts = traceparent.split("-") + if len(parts) == 4: + trace_id = parts[1] + + if not trace_id and standard_logging_object: + slo_trace_id = standard_logging_object.get("trace_id") + if slo_trace_id: + trace_id = slo_trace_id + + except Exception as e: + verbose_logger.warning( + f"Unable to parse New Relic trace context from upstream sources: {e}" + ) + + if not trace_id: + trace_id = uuid.uuid4().hex + verbose_logger.debug( + f"New Relic trace_id not available from distributed tracing headers or " + f"StandardLoggingPayload. Generated trace_id={trace_id} for AI monitoring " + f"event grouping." + ) + + return trace_id + + def _extract_completion_id(self, kwargs: Dict, response_obj: ModelResponse) -> str: + """ + Extract completion ID from kwargs or response_obj, or generate one. + """ + completion_id = None + + if response_obj: + completion_id = response_obj.get("id") + + if not completion_id: + completion_id = kwargs.get("litellm_call_id") + + # If still not found, generate UUID and log warning per spec + if not completion_id: + completion_id = str(uuid.uuid4()) + + return completion_id + + def _get_vendor( + self, + kwargs: Dict, + standard_logging_object: Optional[StandardLoggingPayload] = None, + ) -> str: + """Extract vendor/provider, preferring StandardLoggingPayload.""" + if standard_logging_object: + vendor = standard_logging_object.get("custom_llm_provider") + if vendor: + return vendor + litellm_params = kwargs.get("litellm_params", {}) or {} + return litellm_params.get("custom_llm_provider") or "litellm" + + def _get_model_names( + self, + kwargs: Dict, + response_obj: ModelResponse, + standard_logging_object: Optional[StandardLoggingPayload] = None, + ) -> Tuple[str, str]: + """ + Extract request and response model names, preferring StandardLoggingPayload + for the request model. + + Returns: + Tuple of (request_model, response_model) + """ + request_model = None + if standard_logging_object: + slo_model = standard_logging_object.get("model") + if slo_model: + request_model = str(slo_model) + if not request_model: + request_model = str(kwargs.get("model") or "unknown") + response_model: str = str(response_obj.get("model") or request_model) + return request_model, response_model + + def _extract_usage( + self, + response_obj: ModelResponse, + standard_logging_object: Optional[StandardLoggingPayload] = None, + ) -> Dict[str, int]: + """Extract usage statistics, preferring StandardLoggingPayload.""" + if standard_logging_object: + prompt = standard_logging_object.get("prompt_tokens") + completion = standard_logging_object.get("completion_tokens") + total = standard_logging_object.get("total_tokens") + if any(x is not None for x in [prompt, completion, total]): + return { + "prompt_tokens": prompt or 0, + "completion_tokens": completion or 0, + "total_tokens": total or 0, + } + + usage = response_obj.get("usage", None) + if not usage: + return {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0} + + return { + "prompt_tokens": usage.get("prompt_tokens") or 0, + "completion_tokens": usage.get("completion_tokens") or 0, + "total_tokens": usage.get("total_tokens") or 0, + } + + def _get_finish_reason(self, response_obj: ModelResponse) -> str: + """ + Extract finish reason from first choice in the response. + + Returns "unknown" if choices are not present or finish_reason is not found. + """ + choices = response_obj.get("choices") or [] + if choices and len(choices) > 0: + return choices[0].get("finish_reason") or "unknown" + return "unknown" + + def _to_epoch_ms(self, t: Any) -> float: + """Convert a datetime or float timestamp to epoch milliseconds.""" + if hasattr(t, "timestamp"): + return t.timestamp() * 1000.0 + return float(t) * 1000.0 + + def _get_duration( + self, + kwargs: Dict, + start_time: Any, + end_time: Any, + standard_logging_object: Optional[StandardLoggingPayload] = None, + ) -> Optional[float]: + """ + Extract duration in milliseconds. + + Resolution order: + 1. StandardLoggingPayload.response_time (already computed by LiteLLM) + 2. llm_api_duration_ms from kwargs + 3. Calculated from start_time and end_time + """ + if standard_logging_object: + response_time = standard_logging_object.get("response_time") + if response_time is not None: + return ( + float(response_time) * 1000.0 + ) # SLO stores seconds; convert to ms + + duration_ms = kwargs.get("llm_api_duration_ms") + if duration_ms is not None: + return float(duration_ms) + + if start_time is not None and end_time is not None: + return self._to_epoch_ms(end_time) - self._to_epoch_ms(start_time) + + return None + + def _get_request_params( + self, + kwargs: Dict, + standard_logging_object: Optional[StandardLoggingPayload] = None, + ) -> Dict[str, Any]: + """ + Extract request parameters like temperature and max_tokens, preferring + StandardLoggingPayload.model_parameters. + + Returns dict with available parameters, omitting those not present. + """ + if standard_logging_object: + source_params = standard_logging_object.get("model_parameters") or {} + else: + source_params = kwargs.get("optional_params") or {} + + params = {} + + temperature = source_params.get("temperature") + if temperature is not None: + params["temperature"] = temperature + + max_tokens = source_params.get("max_tokens") + if max_tokens is not None: + params["max_tokens"] = max_tokens + + return params + + def _extract_message_content(self, message: Union[Message, Dict]) -> str: + """ + Extract content from a message, handling various formats. + + Handles tool calls, multimodal content (as JSON), and standard text content. + Returns empty string if content is None or missing. + """ + content = message.get("content") + + # Handle tool calls + if message.get("tool_calls"): + try: + return json.dumps(message["tool_calls"]) + except Exception: + return str(message["tool_calls"]) + + # Handle None or missing content + if content is None: + return "" + + # Handle list content (multimodal) + if isinstance(content, list): + try: + return json.dumps(content) + except Exception: + return str(content) + + # Handle non-string content + if not isinstance(content, str): + return str(content) + + return content + + def _extract_all_messages( + self, + kwargs: Dict, + response_obj: ModelResponse, + response_model: str, + vendor: str, + standard_logging_object: Optional[StandardLoggingPayload] = None, + ) -> List[Dict[str, Any]]: + """ + Extract all messages (request + response) with sequence numbers and timestamps. + + Processes request messages from StandardLoggingPayload.messages (preferred) or + kwargs["messages"] (fallback), and response messages from response_obj["choices"]. + Assigns sequential numbers starting at 0. + Adds timestamps from StandardLoggingPayload (preferred) or kwargs if available + (converted to epoch milliseconds). + """ + messages = [] + sequence = 0 + + # Extract timestamps, preferring StandardLoggingPayload + start_time = None + if standard_logging_object: + start_time = standard_logging_object.get("startTime") + if not start_time: + start_time = kwargs.get("start_time") + + end_time = None + if standard_logging_object: + end_time = standard_logging_object.get("endTime") + if not end_time: + end_time = kwargs.get("end_time") + + # Content is recorded only when the NR-specific switches allow it AND + # LiteLLM's wider redaction decision (turn_off_message_logging, dynamic + # params, headers) does not require redaction. Async streaming hands the + # callback an unredacted async_complete_streaming_response, so without + # this gate generated content would still reach NR even when the user + # has globally disabled message logging. + record_content = self.record_content and not should_redact_message_logging( + kwargs + ) + + # Extract request messages, preferring StandardLoggingPayload. + # SLO messages can be a string (serialized/redacted), so only use it when it's a list. + slo_messages = ( + standard_logging_object.get("messages") if standard_logging_object else None + ) + if isinstance(slo_messages, list): + request_messages = slo_messages + else: + request_messages = kwargs.get("messages") or [] + for msg in request_messages: + message_data = { + "role": msg.get("role") or "user", + "sequence": sequence, + "response.model": response_model, + "vendor": vendor, + } + + # Add timestamp for request message if available (convert to milliseconds) + if start_time is not None: + message_data["timestamp"] = int(self._to_epoch_ms(start_time)) + + if record_content: + message_data["content"] = self._extract_message_content(msg) + + messages.append(message_data) + sequence += 1 + + # Extract response messages from choices + choices = response_obj.get("choices") or [] + if choices and len(choices) > 0: + for choice in choices: + # Prefer "message" (non-streaming); fall back to "delta" (streaming-assembled) + message = choice.get("message", None) or choice.get("delta", None) + if message: + message_data = { + "role": message.get("role") or "assistant", + "sequence": sequence, + "response.model": response_model, + "vendor": vendor, + "is_response": True, + } + + # Add timestamp for response message if available (convert to milliseconds) + if end_time is not None: + message_data["timestamp"] = int(self._to_epoch_ms(end_time)) + + if record_content: + message_data["content"] = self._extract_message_content(message) + + messages.append(message_data) + sequence += 1 + + return messages + + def _record_summary_event( + self, + request_id: str, + trace_id: Optional[str], + request_model: str, + response_model: str, + vendor: str, + finish_reason: str, + num_messages: int, + usage: Dict[str, int], + duration: Optional[float] = None, + request_params: Optional[Dict[str, Any]] = None, + ): + """Record LlmChatCompletionSummary event to New Relic.""" + try: + event_data = { + "id": request_id, + "request_id": request_id, + "request.model": request_model, + "response.model": response_model, + "response.choices.finish_reason": finish_reason, + "response.number_of_messages": num_messages, + "vendor": vendor, + "ingest_source": "litellm", + "response.usage.prompt_tokens": usage["prompt_tokens"], + "response.usage.completion_tokens": usage["completion_tokens"], + "response.usage.total_tokens": usage["total_tokens"], + } + + # Add optional attributes if present + if trace_id: + event_data["trace_id"] = trace_id + + if duration is not None: + event_data["duration"] = duration + + # Add request parameters if present + if request_params: + if "temperature" in request_params: + event_data["request.temperature"] = request_params["temperature"] + if "max_tokens" in request_params: + event_data["request.max_tokens"] = request_params["max_tokens"] + + app = _newrelic_agent.application() + + if app and app.enabled: + app.record_custom_event("LlmChatCompletionSummary", event_data) + else: + verbose_logger.warning( + "New Relic application is not enabled; skipping summary event recording." + ) + + except Exception as e: + verbose_logger.warning(f"Failed to record New Relic summary event: {e}") + self.handle_callback_failure("newrelic") + + def _record_message_events( + self, + request_id: str, + llm_response_id: str, + trace_id: Optional[str], + messages: List[Dict[str, Any]], + ): + """Record LlmChatCompletionMessage events to New Relic. + + Args: + request_id: Agent-generated UUID that links to Summary event's id + llm_response_id: LLM's response ID (e.g., "chatcmpl-...") for message id format + trace_id: Trace ID for distributed tracing (None if not available) + messages: List of message dicts to record + """ + try: + app = _newrelic_agent.application() + + if not (app and app.enabled): + verbose_logger.warning( + "New Relic application is not enabled; skipping message event recording." + ) + return + + for message in messages: + sequence = message["sequence"] + event_data = { + "id": f"{llm_response_id}-{sequence}", + "request_id": request_id, + "completion_id": request_id, + "role": message["role"], + "sequence": sequence, + "response.model": message["response.model"], + "vendor": message["vendor"], + "ingest_source": "litellm", + "token_count": 0, # Per-message token counts are not available from LiteLLM + } + + # Add trace context if available + if trace_id: + event_data["trace_id"] = trace_id + + # Add content only if it was included in the message data + if "content" in message: + event_data["content"] = message["content"] + + # Add is_response only if True (per spec, omit for request messages) + if message.get("is_response"): + event_data["is_response"] = True + + # Forward actual request/response timestamp (ms) so NR uses the + # real LLM call window rather than the async-logger fire time. + # Requires newrelic>=11.2.0 which reads params["timestamp"] as + # the intrinsic event timestamp. + if "timestamp" in message: + event_data["timestamp"] = message["timestamp"] + + app.record_custom_event("LlmChatCompletionMessage", event_data) + + except Exception as e: + verbose_logger.warning(f"Failed to record New Relic message events: {e}") + self.handle_callback_failure("newrelic") + + def _record_error_metric(self): + """Record error metric to New Relic.""" + try: + if not self.enabled: + return + + self._check_and_emit_periodic_metric() + + app = _newrelic_agent.application() + if app and app.enabled: + app.record_custom_metric("LLM/LiteLLM/Error", 1) + except Exception as e: + verbose_logger.warning(f"Failed to record New Relic error metric: {e}") + self.handle_callback_failure("newrelic") + + def _process_success( + self, + kwargs: Dict, + response_obj: ModelResponse, + start_time: Optional[float] = None, + end_time: Optional[float] = None, + ): + """ + Core logic for processing successful LLM calls. + Used by both sync and async success event handlers. + """ + # Early exit if not enabled + if not self.enabled: + return + + # Check and emit periodic supportability metric if 27 hours have passed + self._check_and_emit_periodic_metric() + + # Use StandardLoggingPayload where available for normalized, pre-computed values + standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get( + "standard_logging_object" + ) + + # Get trace context + trace_id = self._get_trace_context(kwargs, standard_logging_object) + + # Generate unique request ID for this request (used as Summary event id) + request_id = str(uuid.uuid4()) + + # Extract data from response + llm_response_id = self._extract_completion_id(kwargs, response_obj) + vendor = self._get_vendor(kwargs, standard_logging_object) + request_model, response_model = self._get_model_names( + kwargs, response_obj, standard_logging_object + ) + usage = self._extract_usage(response_obj, standard_logging_object) + finish_reason = self._get_finish_reason(response_obj) + + # Extract additional summary event fields + duration = self._get_duration( + kwargs, start_time, end_time, standard_logging_object + ) + request_params = self._get_request_params(kwargs, standard_logging_object) + + # Extract all messages + messages = self._extract_all_messages( + kwargs, response_obj, response_model, vendor, standard_logging_object + ) + + # Record summary event + self._record_summary_event( + request_id=request_id, + trace_id=trace_id, + request_model=request_model, + response_model=response_model, + vendor=vendor, + finish_reason=finish_reason, + num_messages=len(messages), + usage=usage, + duration=duration, + request_params=request_params, + ) + + # Record message events + self._record_message_events( + request_id=request_id, + llm_response_id=llm_response_id, + trace_id=trace_id, + messages=messages, + ) + + async def async_health_check(self) -> IntegrationHealthCheckStatus: + """ + Check if the New Relic integration is healthy. + + Verifies that the integration is enabled and the New Relic agent + has an active, connected application, then records a small + `LiteLLMConnectionTest` custom event so the user can confirm the + end-to-end pipeline in the New Relic UI via NRQL: + `SELECT * FROM LiteLLMConnectionTest SINCE 1 hour ago`. + + The `LiteLLMConnectionTest` event type is intentionally outside the + `Llm*` family that AI Monitoring queries, so test events do not + appear in AI Monitoring dashboards. + """ + if not self.enabled: + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message="New Relic integration is disabled. Check that " + "NEW_RELIC_LICENSE_KEY and NEW_RELIC_APP_NAME are set and the " + "newrelic package is installed.", + ) + + try: + app = _newrelic_agent.application() + if not (app and app.enabled): + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message=( + "New Relic Python agent not installed. Review the New Relic integration documentation at https://docs.litellm.ai/docs/observability/newrelic." + ), + ) + + app.record_custom_event( + "LiteLLMConnectionTest", + { + "is_test_event": True, + "app_name": self.app_name, + "source": "litellm-proxy", + "timestamp": time.time(), + }, + ) + return IntegrationHealthCheckStatus(status="healthy", error_message=None) + except Exception as e: + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message=str(e), + ) + + # CustomLogger interface implementation + + def log_pre_api_call(self, model, messages, kwargs): + """Unused per spec.""" + pass + + def log_post_api_call(self, kwargs, response_obj, start_time, end_time): + """Unused per spec.""" + pass + + def log_success_event(self, kwargs, response_obj, start_time, end_time): + """ + Main success path for non-streaming requests. + + Note: New Relic's record_custom_event is synchronous but non-blocking + (in-memory operation), so it's safe to call from sync context. + """ + try: + self._process_success(kwargs, response_obj, start_time, end_time) + except Exception as e: + verbose_logger.warning(f"Error in New Relic log_success_event: {e}") + self.handle_callback_failure("newrelic") + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + """ + Main success path for async/streaming requests. + + Note: New Relic's SDK is thread-safe and record_custom_event is fast, + so we can call it directly without asyncio.to_thread(). + """ + try: + self._process_success(kwargs, response_obj, start_time, end_time) + except Exception as e: + verbose_logger.warning(f"Error in New Relic async_log_success_event: {e}") + self.handle_callback_failure("newrelic") + + def log_failure_event(self, kwargs, response_obj, start_time, end_time): + """ + Log error metric for failed LLM calls (sync). + + Per spec: Do not send AI events on failure, only record error metric. + """ + try: + self._record_error_metric() + + except Exception as e: + verbose_logger.warning(f"Error in New Relic log_failure_event: {e}") + self.handle_callback_failure("newrelic") + + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + """ + Log error metric for failed LLM calls (async). + + Per spec: Do not send AI events on failure, only record error metric. + """ + try: + self._record_error_metric() + + except Exception as e: + verbose_logger.warning(f"Error in New Relic async_log_failure_event: {e}") + self.handle_callback_failure("newrelic") diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index fd402b90d88..a7fae104c92 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -25,6 +25,7 @@ from litellm.integrations.datadog.datadog_metrics import DatadogMetricsLogger from litellm.integrations.deepeval import DeepEvalLogger from litellm.integrations.dotprompt import DotpromptManager from litellm.integrations.focus.focus_logger import FocusLogger +from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import MavvrikFocusLogger from litellm.integrations.vantage.vantage_logger import VantageLogger from litellm.integrations.galileo import GalileoObserve from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger @@ -39,6 +40,7 @@ from litellm.integrations.langsmith import LangsmithLogger from litellm.integrations.litellm_agent import LiteLLMAgentModelResolver from litellm.integrations.literal_ai import LiteralAILogger from litellm.integrations.mlflow import MlflowLogger +from litellm.integrations.newrelic import NewRelicLogger from litellm.integrations.openmeter import OpenMeterLogger from litellm.integrations.opentelemetry import OpenTelemetry from litellm.integrations.opik.opik import OpikLogger @@ -102,8 +104,10 @@ class CustomLoggerRegistry: "gitlab": GitLabPromptManager, "cloudzero": CloudZeroLogger, "focus": FocusLogger, + "mavvrik": MavvrikFocusLogger, "vantage": VantageLogger, "posthog": PostHogLogger, + "newrelic": NewRelicLogger, } try: diff --git a/litellm/litellm_core_utils/duration_parser.py b/litellm/litellm_core_utils/duration_parser.py index 6d2b4226ff4..036d691c686 100644 --- a/litellm/litellm_core_utils/duration_parser.py +++ b/litellm/litellm_core_utils/duration_parser.py @@ -131,6 +131,8 @@ def get_next_standardized_reset_time( # Handle different time units if unit == "d": return _handle_day_reset(current_time, base_midnight, value, tz) + elif unit == "w": + return _handle_day_reset(current_time, base_midnight, value * 7, tz) elif unit == "h": return _handle_hour_reset(current_time, base_midnight, value) elif unit == "m": diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index dbfcf55d75d..b2db334d5ff 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -158,6 +158,7 @@ from ..integrations.litellm_agent import LiteLLMAgentModelResolver from ..integrations.literal_ai import LiteralAILogger from ..integrations.logfire_logger import LogfireLevel, LogfireLogger from ..integrations.lunary import LunaryLogger +from ..integrations.newrelic import NewRelicLogger from ..integrations.openmeter import OpenMeterLogger from ..integrations.opik.opik import OpikLogger from ..integrations.posthog import PostHogLogger @@ -3507,9 +3508,7 @@ class Logging(LiteLLMLoggingBaseClass): else: return None - def _handle_anthropic_messages_response_logging( - self, result: Any - ) -> Union[ModelResponse, ResponsesAPIResponse]: + def _handle_anthropic_messages_response_logging(self, result: Any) -> ModelResponse: """ Handles logging for Anthropic messages responses. @@ -3528,15 +3527,14 @@ class Logging(LiteLLMLoggingBaseClass): return result elif isinstance(result, ModelResponse): return result - elif isinstance( + + if isinstance( result, (ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent), ): - # anthropic_messages() can route to OpenAI Responses API; in that path - # the assembled streaming result is one of these terminal events rather than - # a ModelResponse. Return the inner response so downstream handlers - # (_transform_usage_objects, normalize_logging_result) can process it. - return result.response + result = result.response + if isinstance(result, ResponsesAPIResponse): + return self._translate_responses_api_response_to_model_response(result) httpx_response = self.model_call_details.get("httpx_response", None) if httpx_response and isinstance(httpx_response, httpx.Response): @@ -3570,6 +3568,55 @@ class Logging(LiteLLMLoggingBaseClass): ) return result + def _translate_responses_api_response_to_model_response( + self, result: ResponsesAPIResponse + ) -> ModelResponse: + """ + Convert a Responses API response into a ModelResponse for spend_logs. + + The proxy UI parses spend_log rows expecting chat-completion shape + (response.choices[0].message); a raw ResponsesAPIResponse dump (output[...]) + would render as empty in the Logs tab. Translation also yields full + choices/message detail downstream consumers can rely on. + """ + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + try: + return LiteLLMResponsesTransformationHandler().transform_response( + model=self.model, + raw_response=result, + model_response=litellm.ModelResponse(), + logging_obj=self, + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=litellm.encoding, + ) + except Exception as e: + verbose_logger.debug( + "Responses API -> ModelResponse translation failed for " + "anthropic_messages logging (%s); falling back to minimal " + "usage-only ModelResponse to keep the spend_logs row.", + str(e), + ) + model_response = litellm.ModelResponse() + model_response.model = self.model + usage = getattr(result, "usage", None) + if usage is not None and ResponseAPILoggingUtils._is_response_api_usage( + usage + ): + setattr( + model_response, + "usage", + ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( + usage + ), + ) + return model_response + def _handle_non_streaming_google_genai_generate_content_response_logging( self, result: Any ) -> ModelResponse: @@ -4124,6 +4171,17 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 focus_logger = FocusLogger() _in_memory_loggers.append(focus_logger) return focus_logger # type: ignore + elif logging_integration == "mavvrik": + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + MavvrikFocusLogger, + ) + + for callback in _in_memory_loggers: + if type(callback) is MavvrikFocusLogger: + return callback # type: ignore + mavvrik_focus_logger = MavvrikFocusLogger() + _in_memory_loggers.append(mavvrik_focus_logger) + return mavvrik_focus_logger # type: ignore elif logging_integration == "vantage": from litellm.integrations.vantage.vantage_logger import VantageLogger @@ -4419,6 +4477,13 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 gitlab_logger = GitLabPromptManager(gitlab_config=gitlab_config) _in_memory_loggers.append(gitlab_logger) return gitlab_logger # type: ignore + elif logging_integration == "newrelic": + for callback in _in_memory_loggers: + if isinstance(callback, NewRelicLogger): + return callback # type: ignore + newrelic_logger = NewRelicLogger() + _in_memory_loggers.append(newrelic_logger) + return newrelic_logger # type: ignore return None except Exception as e: verbose_logger.exception( @@ -4720,6 +4785,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915 for callback in _in_memory_loggers: if isinstance(callback, SMTPEmailLogger): return callback + elif logging_integration == "newrelic": + for callback in _in_memory_loggers: + if isinstance(callback, NewRelicLogger): + return callback return None except Exception as e: diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index b09f2bb130e..5059e612f2f 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -4741,6 +4741,12 @@ class BedrockConverseMessagesProcessor: guardContent={"text": {"text": element["text"]}} ) _parts.append(_part) + elif element["type"] in ("grounding_source", "query"): + # Contextual grounding tags are guardrail metadata; the + # model only needs the underlying text, so render them + # as plain text on the generate path. + _part = BedrockContentBlock(text=element["text"]) + _parts.append(_part) elif element["type"] == "image_url": format: Optional[str] = None if isinstance(element["image_url"], dict): @@ -5173,6 +5179,12 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 guardContent={"text": {"text": element["text"]}} ) _parts.append(_part) + elif element["type"] in ("grounding_source", "query"): + # Contextual grounding tags are guardrail metadata; the + # model only needs the underlying text, so render them as + # plain text on the generate path. + _part = BedrockContentBlock(text=element["text"]) + _parts.append(_part) elif element["type"] == "image_url": format: Optional[str] = None if isinstance(element["image_url"], dict): diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 0d5d4a2086c..c8f87d96e2f 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -47,6 +47,7 @@ class RealTimeStreaming: user_api_key_dict: Optional[Any] = None, request_data: Optional[Dict] = None, backend_uses_beta_protocol: Optional[bool] = None, + force_transcription_model: Optional[str] = None, ): self.websocket = websocket self.backend_ws = backend_ws @@ -100,6 +101,11 @@ class RealTimeStreaming: self._flushing_pending_messages_until_setup: bool = False self._pending_messages_until_setup: List[str] = [] self._pending_messages_byte_total: int = 0 + # Whether this is a transcription-only session (session.type == "transcription", + # e.g. gpt-realtime-whisper). Such sessions must not be sent response.create and + # their input_audio_transcription.completed usage drives duration-based cost. + self._force_transcription_model = force_transcription_model + self._is_transcription_session: bool = force_transcription_model is not None # Per-connection caps for pre-setup audio frames (message count + total bytes). _MAX_BUFFERED_MESSAGES: int = 200 @@ -211,6 +217,8 @@ class RealTimeStreaming: self.session_tools = tools # GA: session.type is required; log it for traceability but no action needed verbose_logger.debug(f"Realtime session.type: {session.get('type')}") + if session.get("type") == "transcription": + self._is_transcription_session = True except (json.JSONDecodeError, AttributeError, TypeError): pass @@ -227,6 +235,55 @@ class RealTimeStreaming: except (AttributeError, TypeError): pass + def _detect_transcription_session_from_backend( + self, event_obj: Union[dict, OpenAIRealtimeEvents] + ) -> None: + """Flag transcription-only sessions from backend session events.""" + try: + event_type = event_obj.get("type", "") + if event_type in ( + "transcription_session.created", + "transcription_session.updated", + ): + self._is_transcription_session = True + elif event_type in ("session.created", "session.updated"): + session = cast(dict, event_obj).get("session", {}) or {} + if session.get("type") == "transcription": + self._is_transcription_session = True + except (AttributeError, TypeError): + pass + + def _capture_transcription_usage( + self, event_obj: Union[dict, OpenAIRealtimeEvents] + ) -> None: + """ + Append a usage-only transcription completed event to the logged results so + the cost calculator can bill it by audio duration. The default logged event + types exclude this event, so it is captured here directly for transcription + sessions rather than widening logging for every realtime session. Only the + type and usage are kept — the transcript is already captured separately in + input_messages, so it is not duplicated into the response log here. + """ + try: + usage = event_obj.get("usage") + if usage is None: + return + # If this event type is already captured by store_message (e.g. the user + # logs all realtime events), don't append a second copy. + if self._should_store_message(event_obj): + return + self.messages.append( + cast( + OpenAIRealtimeEvents, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": usage, + }, + ) + ) + except (AttributeError, TypeError): + pass + def _collect_tool_calls_from_response_done( self, event_obj: Union[dict, OpenAIRealtimeEvents] ) -> None: @@ -287,6 +344,7 @@ class RealTimeStreaming: backend, False if the provider transformation produced no output and the message was effectively dropped. """ + message = self._enforce_transcription_session_model(message) if self.provider_config: transformed = self.provider_config.transform_realtime_request( message, self.model, self.session_configuration_request @@ -306,6 +364,80 @@ class RealTimeStreaming: await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined] return True + def _enforce_transcription_session_model(self, message: str) -> str: + """Force client transcription session updates to the authorized model. + + `/v1/realtime?intent=transcription` may intentionally omit `model` from + the upstream URL for Azure compatibility, but the proxy still authorizes + a resolved LiteLLM model before opening the backend websocket. If a + client later sends a transcription `session.update`, any model embedded + in that update must be rewritten to the same authorized model instead of + allowing a post-auth model/deployment switch. + + Normal realtime sessions keep their independent nested transcription + model behavior because `_force_transcription_model` is only set for + transcription-intent websocket routes. + """ + if self._force_transcription_model is None: + return message + + try: + message_obj = json.loads(message) + except (json.JSONDecodeError, TypeError): + return message + + if message_obj.get("type") not in ( + "session.update", + "transcription_session.update", + ): + return message + + session = message_obj.get("session") + if not isinstance(session, dict): + return message + + if session.get("type") == "transcription": + self._is_transcription_session = True + + authorized_model = self._force_transcription_model + changed = False + + transcription = session.get("input_audio_transcription") + if ( + isinstance(transcription, dict) + and transcription.get("model") != authorized_model + ): + session["input_audio_transcription"] = { + **transcription, + "model": authorized_model, + } + changed = True + + audio = session.get("audio") + if isinstance(audio, dict): + audio_input = audio.get("input") + if isinstance(audio_input, dict): + nested_transcription = audio_input.get("transcription") + if ( + isinstance(nested_transcription, dict) + and nested_transcription.get("model") != authorized_model + ): + session["audio"] = { + **audio, + "input": { + **audio_input, + "transcription": { + **nested_transcription, + "model": authorized_model, + }, + }, + } + changed = True + + if not changed: + return message + return json.dumps(message_obj) + def _uses_deferred_backend_setup(self) -> bool: """True when setup is deferred until the client's first session.update.""" if self.provider_config is None: @@ -792,6 +924,8 @@ class RealTimeStreaming: """ event_type = event_obj.get("type") + self._detect_transcription_session_from_backend(event_obj) + # Send session.created to the client FIRST so it stays in sync, then inject # the disable-auto-response session.update; otherwise a backend error could # reach the client before it sees session.created. @@ -809,6 +943,14 @@ class RealTimeStreaming: self._collect_user_input_from_backend_event(event_obj) self.store_message(event_obj) await self.websocket.send_text(raw_response) + + # Transcription-only sessions (e.g. gpt-realtime-whisper) have no + # assistant turn: capture audio-duration usage for cost and never + # trigger response.create. + if self._is_transcription_session: + self._capture_transcription_usage(event_obj) + return True + blocked = await self.run_realtime_guardrails( transcript, item_id=event_obj.get("item_id"), diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index 8c20f4c430e..f049abcf47f 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -469,12 +469,16 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): if should_start_new_block and not self.sent_content_block_finish: # Queue the sequence: content_block_stop -> content_block_start - # For text blocks the trigger chunk is not emitted as a separate - # delta because content_block_start carries the information. - # For tool_use blocks we must also emit the trigger chunk's delta - # when it carries input_json_delta data, because some providers - # (e.g. xAI, Gemini) include tool arguments in the same streaming - # chunk as the function name/id. + # -> (optionally) the trigger chunk's delta. + # + # The synthesized content_block_start always carries an + # empty body, so the chunk that *triggered* the transition + # also carries the new block's first delta. It must be + # re-emitted or the first token of the new block is lost. + # This applies to text_delta and thinking_delta (the first + # non-empty text/thinking token) as well as input_json_delta + # (providers like xAI/Gemini bundle tool arguments with the + # function name/id in a single chunk). # 1. Stop current content block self.chunk_queue.append( @@ -493,14 +497,9 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): } ) - # 3. If the trigger chunk carries tool argument data, queue it - # so the input_json_delta is not silently dropped. - if ( - processed_chunk.get("type") == "content_block_delta" - and isinstance(processed_chunk.get("delta"), dict) - and processed_chunk["delta"].get("type") == "input_json_delta" - and processed_chunk["delta"].get("partial_json") - ): + # 3. If the trigger chunk carries delta content, queue it + # so the first delta of the new block is not silently dropped. + if self._trigger_delta_has_content(processed_chunk): self.chunk_queue.append(processed_chunk) self.sent_content_block_finish = False @@ -711,12 +710,16 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): if not self.queued_usage_chunk: if should_start_new_block and not self.sent_content_block_finish: # Queue the sequence: content_block_stop -> content_block_start - # For text blocks the trigger chunk is not emitted as a separate - # delta because content_block_start carries the information. - # For tool_use blocks we must also emit the trigger chunk's delta - # when it carries input_json_delta data, because some providers - # (e.g. xAI, Gemini) include tool arguments in the same streaming - # chunk as the function name/id. + # -> (optionally) the trigger chunk's delta. + # + # The synthesized content_block_start always carries an + # empty body, so the chunk that *triggered* the transition + # also carries the new block's first delta. It must be + # re-emitted or the first token of the new block is lost. + # This applies to text_delta and thinking_delta (the + # first non-empty text/thinking token) as well as + # input_json_delta (providers like xAI/Gemini bundle tool + # arguments with the function name/id in a single chunk). # 1. Stop current content block self.chunk_queue.append( @@ -733,15 +736,9 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): } ) - # 3. If the trigger chunk carries tool argument data, queue it - # so the input_json_delta is not silently dropped. - if ( - processed_chunk.get("type") == "content_block_delta" - and isinstance(processed_chunk.get("delta"), dict) - and processed_chunk["delta"].get("type") - == "input_json_delta" - and processed_chunk["delta"].get("partial_json") - ): + # 3. If the trigger chunk carries delta content, queue it + # so the first delta of the new block is not silently dropped. + if self._trigger_delta_has_content(processed_chunk): self.chunk_queue.append(processed_chunk) # Reset state for new block @@ -898,6 +895,38 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): def _increment_content_block_index(self): self.current_content_block_index += 1 + @staticmethod + def _trigger_delta_has_content(processed_chunk: Dict[str, Any]) -> bool: + """Return True if a translated trigger chunk carries a non-empty + ``content_block_delta`` payload that must be re-emitted after a + block transition. + + When an upstream chunk both *triggers* a new content block (its type + differs from the active block) and *carries* delta content, that + content belongs to the new block. The synthesized + ``content_block_start`` only ever carries an empty body — see + ``_translate_streaming_openai_chunk_to_anthropic_content_block``, + which returns an empty ``TextBlock``/``ToolUseBlock``/thinking block — + so the trigger chunk's delta must be re-queued or the first token of + the new block (the first non-empty text/thinking delta, or bundled + tool arguments) is silently dropped. + """ + if processed_chunk.get("type") != "content_block_delta": + return False + delta = processed_chunk.get("delta") + if not isinstance(delta, dict): + return False + delta_type = delta.get("type") + if delta_type == "text_delta": + return bool(delta.get("text")) + if delta_type == "input_json_delta": + return bool(delta.get("partial_json")) + if delta_type == "thinking_delta": + return bool(delta.get("thinking")) + if delta_type == "signature_delta": + return bool(delta.get("signature")) + return False + def _should_start_new_content_block(self, chunk: "ModelResponseStream") -> bool: """ Determine if we should start a new content block based on the processed chunk. diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/handler.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/handler.py index f8c827ab057..70855afa81c 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/handler.py @@ -102,9 +102,9 @@ def _build_responses_kwargs( from litellm.types.utils import CallTypes if isinstance(value, LiteLLMLoggingObject): - # Reclassify as acompletion so the success handler doesn't try to - # validate the Responses API event as an AnthropicResponse. - # (Mirrors the pattern used in LiteLLMMessagesToCompletionTransformationHandler.) + # Keep call_type as anthropic_messages so spend_logs are billed + # against /v1/messages; the success handler translates the + # Responses API result back to a ModelResponse for the row. setattr(value, "call_type", CallTypes.anthropic_messages.value) responses_kwargs[key] = value elif key not in excluded and key not in responses_kwargs and value is not None: diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index 1f3357fd788..9c8de6c06a1 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -8,6 +8,7 @@ from typing import Any, Optional, cast from litellm._logging import _redact_string, verbose_proxy_logger from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES +from litellm.types.realtime import RealtimeQueryParams from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from ....litellm_core_utils.realtime_streaming import RealTimeStreaming @@ -35,6 +36,7 @@ class AzureOpenAIRealtime(AzureChatCompletion): model: str, api_version: Optional[str], realtime_protocol: Optional[str] = None, + query_params: Optional[RealtimeQueryParams] = None, ) -> str: """ Construct Azure realtime WebSocket URL. @@ -46,6 +48,7 @@ class AzureOpenAIRealtime(AzureChatCompletion): realtime_protocol: Protocol version to use: - "GA" or "v1": Uses /openai/v1/realtime (GA path) - "beta" or None: Uses /openai/realtime (beta path, default) + query_params: Extra query params to forward (e.g. intent=transcription). Returns: WebSocket URL string @@ -54,6 +57,8 @@ class AzureOpenAIRealtime(AzureChatCompletion): beta/default: "wss://.../openai/realtime?api-version=2024-10-01-preview&deployment=gpt-4o-realtime-preview" GA/v1: "wss://.../openai/v1/realtime?model=gpt-realtime-deployment" """ + from urllib.parse import urlencode + api_base = api_base.replace("https://", "wss://") # Determine path based on realtime_protocol (case-insensitive) @@ -61,13 +66,25 @@ class AzureOpenAIRealtime(AzureChatCompletion): "GA", "V1", ) + intent = (query_params or {}).get("intent") + if _is_ga: path = "/openai/v1/realtime" - return f"{api_base}{path}?model={model}" + query_parts = [] + if intent != "transcription" and ( + query_params is None or "model" in query_params + ): + query_parts.append(urlencode({"model": model})) else: # Default to beta path for backwards compatibility path = "/openai/realtime" - return f"{api_base}{path}?api-version={api_version}&deployment={model}" + query_parts = [urlencode({"api-version": api_version, "deployment": model})] + + if intent: + query_parts.append(urlencode({"intent": intent})) + + qs = "&".join(query_parts) + return f"{api_base}{path}?{qs}" if qs else f"{api_base}{path}" async def async_realtime( self, @@ -81,6 +98,7 @@ class AzureOpenAIRealtime(AzureChatCompletion): client: Optional[Any] = None, timeout: Optional[float] = None, realtime_protocol: Optional[str] = None, + query_params: Optional[RealtimeQueryParams] = None, user_api_key_dict: Optional[Any] = None, litellm_metadata: Optional[dict] = None, ): @@ -96,7 +114,11 @@ class AzureOpenAIRealtime(AzureChatCompletion): raise ValueError("api_version is required for Azure OpenAI calls") url = self._construct_url( - api_base, model, api_version, realtime_protocol=realtime_protocol + api_base, + model, + api_version, + realtime_protocol=realtime_protocol, + query_params=query_params, ) try: @@ -113,9 +135,15 @@ class AzureOpenAIRealtime(AzureChatCompletion): websocket, cast(ClientConnection, backend_ws), logging_obj, + model=model, user_api_key_dict=user_api_key_dict, request_data={"litellm_metadata": litellm_metadata or {}}, backend_uses_beta_protocol=backend_uses_beta_protocol, + force_transcription_model=( + model + if (query_params or {}).get("intent") == "transcription" + else None + ), ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/llms/azure/realtime/http_transformation.py b/litellm/llms/azure/realtime/http_transformation.py index df1e2707af2..d6bdbd24db4 100644 --- a/litellm/llms/azure/realtime/http_transformation.py +++ b/litellm/llms/azure/realtime/http_transformation.py @@ -40,6 +40,13 @@ class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig): version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17" return f"{base}/openai/realtime/calls?api-version={version}" + def get_transcription_session_url( + self, api_base: Optional[str], model: str, api_version: Optional[str] = None + ) -> str: + base = self.get_api_base(api_base).rstrip("/") + version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17" + return f"{base}/openai/realtime/transcription_sessions?api-version={version}" + def get_realtime_calls_headers(self, ephemeral_key: str) -> dict: return { "api-key": ephemeral_key, diff --git a/litellm/llms/base_llm/realtime/http_transformation.py b/litellm/llms/base_llm/realtime/http_transformation.py index 712ec42380f..be1413a3c0b 100644 --- a/litellm/llms/base_llm/realtime/http_transformation.py +++ b/litellm/llms/base_llm/realtime/http_transformation.py @@ -59,6 +59,15 @@ class BaseRealtimeHTTPConfig(ABC): ) -> str: """Return the full URL for POST /realtime/client_secrets.""" + def get_transcription_session_url( + self, api_base: Optional[str], model: str, api_version: Optional[str] = None + ) -> str: + """Return the full URL for POST /realtime/transcription_sessions.""" + base = (api_base or "").rstrip("/") + if base.endswith("/v1"): + base = base[:-3] + return f"{base}/v1/realtime/transcription_sessions" + @abstractmethod def validate_environment( self, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 25424feaeb4..a3281655e9c 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -125,6 +125,7 @@ from litellm.types.vector_stores import ( VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, ) +from litellm.types.realtime import RealtimeQueryParams from litellm.types.videos.main import VideoObject from litellm.utils import ( CustomStreamWrapper, @@ -2305,6 +2306,7 @@ class BaseLLMHTTPHandler: if extra_body: data.update(extra_body) + stream = bool(stream or data.get("stream")) # Preserve the OpenAI-style request context (not sent to the provider) for streaming # hooks/metadata; the streaming iterator now consumes this to run deployment hooks @@ -2467,6 +2469,7 @@ class BaseLLMHTTPHandler: if extra_body: data.update(extra_body) + stream = bool(stream or data.get("stream")) # Preserve the OpenAI-style request context (not sent to the provider) for streaming # hooks/metadata; the streaming iterator now consumes this to run deployment hooks @@ -5315,6 +5318,23 @@ class BaseLLMHTTPHandler: headers=error_headers, ) + @staticmethod + def _append_query_params( + url: str, query_params: Optional[RealtimeQueryParams] + ) -> str: + """Append query_params to url, skipping keys already present in the URL.""" + if not query_params: + return url + from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse + + parsed = urlparse(url) + existing = dict(parse_qsl(parsed.query)) + extras = {k: v for k, v in query_params.items() if k not in existing} + if not extras: + return url + new_query = parsed.query + ("&" if parsed.query else "") + urlencode(extras) + return urlunparse(parsed._replace(query=new_query)) + async def async_realtime( self, model: str, @@ -5328,11 +5348,14 @@ class BaseLLMHTTPHandler: timeout: Optional[float] = None, user_api_key_dict: Optional[Any] = None, litellm_metadata: Optional[Dict[str, Any]] = None, + query_params: Optional[RealtimeQueryParams] = None, ): import websockets from websockets.asyncio.client import ClientConnection - url = provider_config.get_complete_url(api_base, model, api_key) + url = self._append_query_params( + provider_config.get_complete_url(api_base, model, api_key), query_params + ) headers = provider_config.validate_environment( headers=headers, model=model, @@ -5373,6 +5396,11 @@ class BaseLLMHTTPHandler: model, user_api_key_dict=user_api_key_dict, request_data=_request_data, + force_transcription_model=( + model + if (query_params or {}).get("intent") == "transcription" + else None + ), ) if _session_config: realtime_streaming.session_configuration_request = _session_config @@ -5437,6 +5465,69 @@ class BaseLLMHTTPHandler: """ Forward POST /v1/realtime/client_secrets to upstream provider. + Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and + header auth when available; falls back to the legacy OpenAI-style defaults. + """ + return await self._async_realtime_session_post( + endpoint="client_secrets", + api_base=api_base, + api_key=api_key, + request_data=request_data, + logging_obj=logging_obj, + timeout=timeout, + provider_config=provider_config, + model=model, + extra_headers=extra_headers, + client=client, + api_version=api_version, + ) + + async def async_realtime_transcription_session_handler( + self, + api_base: str, + api_key: str, + request_data: Dict[str, Any], + logging_obj: LiteLLMLoggingObj, + timeout: Union[float, httpx.Timeout], + provider_config: Optional[Any] = None, + model: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + api_version: Optional[str] = None, + ) -> httpx.Response: + """Forward POST /v1/realtime/transcription_sessions to upstream provider.""" + return await self._async_realtime_session_post( + endpoint="transcription_sessions", + api_base=api_base, + api_key=api_key, + request_data=request_data, + logging_obj=logging_obj, + timeout=timeout, + provider_config=provider_config, + model=model, + extra_headers=extra_headers, + client=client, + api_version=api_version, + ) + + async def _async_realtime_session_post( + self, + endpoint: Literal["client_secrets", "transcription_sessions"], + api_base: str, + api_key: str, + request_data: Dict[str, Any], + logging_obj: LiteLLMLoggingObj, + timeout: Union[float, httpx.Timeout], + provider_config: Optional[Any] = None, + model: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + api_version: Optional[str] = None, + ) -> httpx.Response: + """ + Shared POST flow for the realtime HTTP session endpoints + (client_secrets and transcription_sessions). + Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and header auth when available; falls back to the legacy OpenAI-style defaults. """ @@ -5448,14 +5539,19 @@ class BaseLLMHTTPHandler: async_httpx_client = client if provider_config is not None: - url = provider_config.get_complete_url( - api_base=api_base, model=model or "", api_version=api_version - ) + if endpoint == "transcription_sessions": + url = provider_config.get_transcription_session_url( + api_base=api_base, model=model or "", api_version=api_version + ) + else: + url = provider_config.get_complete_url( + api_base=api_base, model=model or "", api_version=api_version + ) headers: Dict[str, Any] = provider_config.validate_environment( headers={}, model=model or "", api_key=api_key ) else: - url = f"{api_base.rstrip('/')}/v1/realtime/client_secrets" + url = f"{api_base.rstrip('/')}/v1/realtime/{endpoint}" headers = { "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index 3406538c774..299f346a7eb 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -101,6 +101,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): def __init__(self) -> None: super().__init__() self.authenticator = Authenticator() + self._stream_item_ids_by_output_index: Dict[int, str] = {} @property def custom_llm_provider(self) -> LlmProviders: @@ -129,6 +130,61 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): """ return dict(response_api_optional_params) + def transform_streaming_response( + self, + model: str, + parsed_chunk: dict, + logging_obj: LiteLLMLoggingObj, + ) -> Any: + parsed_chunk = self._normalize_stream_item_id(parsed_chunk) + return super().transform_streaming_response( + model=model, + parsed_chunk=parsed_chunk, + logging_obj=logging_obj, + ) + + def _normalize_stream_item_id(self, parsed_chunk: dict) -> dict: + """Rewrite streamed item ids to one stable id per output_index. + + GitHub Copilot tags each event of a single output item with a different + item id, so clients that key streaming state by item id (e.g. the Vercel + AI SDK) crash with "reasoning part not found" / "text part not + found". Every sub-event carries a top-level ``item_id`` (whatever the + item type), so its presence is the rewrite signal; output_item.added / + .done instead nest the id under ``item``. The anchor is keyed by + output_index and taken from output_item.added, which the protocol always + emits first, so it is written before any sub-event reads it. Copilot + accepts that id paired with the final encrypted_content next turn, so + multi-turn replay is unaffected. + + State is keyed by output_index on this config, which + ProviderConfigManager builds fresh per request, so it is stream-scoped. + """ + output_index = parsed_chunk.get("output_index") + if not isinstance(output_index, int): + return parsed_chunk + + if parsed_chunk.get("type") == "response.output_item.added": + item = parsed_chunk.get("item") + if isinstance(item, dict) and isinstance(item.get("id"), str): + self._stream_item_ids_by_output_index[output_index] = item["id"] + return parsed_chunk + + stable_id = self._stream_item_ids_by_output_index.get(output_index) + if stable_id is None: + return parsed_chunk + + if isinstance(parsed_chunk.get("item_id"), str): + parsed_chunk = dict(parsed_chunk) + parsed_chunk["item_id"] = stable_id + elif parsed_chunk.get("type") == "response.output_item.done": + item = parsed_chunk.get("item") + if isinstance(item, dict): + parsed_chunk = dict(parsed_chunk) + parsed_chunk["item"] = {**item, "id": stable_id} + + return parsed_chunk + def validate_environment( self, headers: dict, diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index f050f9eea36..d1248b6e518 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -25,6 +25,7 @@ from typing import ( import httpx import litellm +from litellm.constants import DEFAULT_OCI_CHAT_MAX_TOKENS from litellm.litellm_core_utils.logging_utils import track_llm_api_timing from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.custom_httpx.http_handler import ( @@ -87,15 +88,20 @@ STREAMING_TIMEOUT = 60 * 5 def _model_uses_max_completion_tokens(model: str) -> bool: """Return True for OCI-hosted models that require ``maxCompletionTokens``. - Reasoning models on OCI (e.g. the OpenAI GPT-5 family) reject ``maxTokens`` - with HTTP 400 and require ``maxCompletionTokens`` per OpenAI's reasoning-API - convention. Driven by ``supports_reasoning`` in - ``model_prices_and_context_window.json`` so new model families are picked - up via a catalog update rather than a code change. + OpenAI commercial models proxied through OCI (``openai.*``) reject + ``maxTokens`` with HTTP 400 on the reasoning families (gpt-5.x, o-series) + and accept ``maxCompletionTokens`` everywhere, so route the whole vendor + prefix to it rather than chasing each new release in + ``model_prices_and_context_window.json``. The ``openai.gpt-oss-*`` open + weights are served by OCI's own stack and keep ``maxTokens``. Any other + vendor falls back to the catalog's ``supports_reasoning`` flag. """ if not model: return False name = model[4:] if model.lower().startswith("oci/") else model + lowered = name.lower() + if lowered.startswith("openai."): + return not lowered.startswith("openai.gpt-oss") return supports_reasoning(model=name, custom_llm_provider="oci") @@ -193,19 +199,49 @@ def _normalize_response_format(selected_params: Dict, vendor: OCIVendors) -> Non rf = selected_params.get("responseFormat") if not isinstance(rf, dict) or "type" not in rf: return - rf_payload = dict(rf) - selected_params["responseFormat"] = rf_payload - response_type = rf_payload["type"] - if "json_schema" in rf_payload: - raw_schema = rf_payload.pop("json_schema") - rf_payload["jsonSchema"] = ( - dict(raw_schema) if isinstance(raw_schema, dict) else raw_schema - ) + + rf_type = str(rf["type"]).lower() + raw_schema = rf.get("json_schema") + json_schema = raw_schema if isinstance(raw_schema, dict) else None + + if rf_type == "text": + selected_params["responseFormat"] = {"type": "TEXT"} + return + if vendor == OCIVendors.COHERE: - rf_payload["type"] = response_type - else: - fmt = response_type.upper() - rf_payload["type"] = "JSON_OBJECT" if fmt == "JSON" else fmt + # OCI Cohere has no JSON_SCHEMA type; a schema rides on JSON_OBJECT. + payload: Dict[str, Any] = {"type": "JSON_OBJECT"} + if json_schema is not None and json_schema.get("schema") is not None: + payload["schema"] = json_schema["schema"] + selected_params["responseFormat"] = payload + return + + if rf_type == "json_schema": + if json_schema is None: + raise OCIError( + status_code=400, + message="response_format type 'json_schema' requires a 'json_schema' object", + ) + # OCI's ResponseJsonSchema accepts only name/description/schema/isStrict. + # OpenAI sends `strict` instead of `isStrict`; forwarding it (or any + # other extra key) makes OCI reject the whole request with HTTP 400. + oci_schema: Dict[str, Any] = {"name": json_schema.get("name") or "response"} + if json_schema.get("description") is not None: + oci_schema["description"] = json_schema["description"] + if json_schema.get("schema") is not None: + oci_schema["schema"] = json_schema["schema"] + if json_schema.get("strict") is not None: + oci_schema["isStrict"] = json_schema["strict"] + selected_params["responseFormat"] = { + "type": "JSON_SCHEMA", + "jsonSchema": oci_schema, + } + return + + fmt = rf_type.upper() + selected_params["responseFormat"] = { + "type": "JSON_OBJECT" if fmt == "JSON" else fmt + } def get_vendor_from_model(model: str) -> OCIVendors: @@ -297,6 +333,11 @@ class OCIChatConfig(BaseConfig): if get_vendor_from_model(model) == OCIVendors.COHERE else self.openai_to_oci_generic_param_map ) + # `n` is intentionally not advertised for Cohere even though n=1 is + # tolerated: Cohere has no numGenerations field, so n>1 cannot be + # honoured and advertising it would be misleading. Callers that gate on + # this list strip n=1 (a no-op, matching what map_openai_params does); + # callers that bypass it have n=1 dropped there. Both paths converge. return [key for key, value in param_map.items() if value] def map_openai_params( @@ -317,6 +358,19 @@ class OCIChatConfig(BaseConfig): for key, value in {**non_default_params, **optional_params}.items(): alias = param_map.get(key) if alias is False: + # max_retries is a litellm-level control param (litellm applies + # retries itself); it is never a generation param OCI accepts, so + # drop it silently. The litellm proxy injects it on every request, + # which otherwise 500s OCI calls unless drop_params is set. + if key == "max_retries": + continue + # n=1 (or None) is the OpenAI default: a single generation, which + # every OCI model produces anyway. Drop it silently so standard + # clients that always send n=1 (e.g. the MLflow gateway) are not + # rejected; only n>1 is genuinely unsupported on Cohere, which + # has no numGenerations field. + if key == "n" and (value is None or value == 1): + continue if drop_params or litellm.drop_params: continue raise OCIError( @@ -451,6 +505,13 @@ class OCIChatConfig(BaseConfig): elif oci_alias in optional_params: selected_params[target] = optional_params[oci_alias] # type: ignore[index] + # OCI's server-side default token cap is tiny (~20 tokens), so an + # omitted max_tokens silently truncates the response mid-string. Most + # callers never send a limit (MLflow judges among them), so inject a + # sane default when one is absent, mirroring litellm's Anthropic config. + if max_tokens_key not in selected_params: + selected_params[max_tokens_key] = DEFAULT_OCI_CHAT_MAX_TOKENS + # OCI expects uppercase reasoning levels (LOW/MEDIUM/HIGH/NONE); OpenAI # clients send lowercase. OpenAI's "disable" maps to OCI's "NONE". if "reasoningEffort" in selected_params: diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index f34dae2df09..6751004f1b1 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -157,8 +157,14 @@ class OpenAIRealtime(OpenAIChatCompletion): websocket, cast(ClientConnection, backend_ws), logging_obj, + model=model, user_api_key_dict=user_api_key_dict, request_data={"litellm_metadata": litellm_metadata or {}}, + force_transcription_model=( + model + if (query_params or {}).get("intent") == "transcription" + else None + ), ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/llms/openai/realtime/http_transformation.py b/litellm/llms/openai/realtime/http_transformation.py index 1663fcd1fcd..7a6af39ba65 100644 --- a/litellm/llms/openai/realtime/http_transformation.py +++ b/litellm/llms/openai/realtime/http_transformation.py @@ -41,6 +41,14 @@ class OpenAIRealtimeHTTPConfig(BaseRealtimeHTTPConfig): base = base[:-3] return f"{base}/v1/realtime/calls" + def get_transcription_session_url( + self, api_base: Optional[str], model: str, api_version: Optional[str] = None + ) -> str: + base = self.get_api_base(api_base).rstrip("/") + if base.endswith("/v1"): + base = base[:-3] + return f"{base}/v1/realtime/transcription_sessions" + def validate_environment( self, headers: dict, diff --git a/litellm/llms/parallel_ai/search/transformation.py b/litellm/llms/parallel_ai/search/transformation.py index 12d570f1733..85602bf1d86 100644 --- a/litellm/llms/parallel_ai/search/transformation.py +++ b/litellm/llms/parallel_ai/search/transformation.py @@ -1,7 +1,7 @@ """ -Calls Parallel AI's /search endpoint to search the web. +Calls Parallel AI's /v1/search endpoint to search the web. -Parallel AI API Reference: https://docs.parallel.ai/api-reference/search-and-extract-api-beta/search +Parallel AI API Reference: https://docs.parallel.ai/api-reference/search/search """ from typing import Dict, List, Optional, TypedDict, Union @@ -18,36 +18,43 @@ from litellm.secret_managers.main import get_secret_str class _ParallelAISourcePolicy(TypedDict, total=False): - """Source policy for Parallel AI search results.""" - - allowed_domains: List[str] # Optional - list of allowed domains - disallowed_domains: List[str] # Optional - list of disallowed domains + include_domains: List[str] + exclude_domains: List[str] + after_date: str -class _ParallelAISearchRequestRequired(TypedDict): - """Required fields for Parallel AI Search API request.""" - - # Note: At least one of objective or search_queries must be provided - pass +class _ParallelAIExcerptSettings(TypedDict, total=False): + max_chars_per_result: int -class ParallelAISearchRequest(_ParallelAISearchRequestRequired, total=False): +class _ParallelAIAdvancedSettings(TypedDict, total=False): + source_policy: _ParallelAISourcePolicy + excerpt_settings: _ParallelAIExcerptSettings + fetch_policy: Dict + location: str + max_results: int + + +class ParallelAISearchRequest(TypedDict, total=False): """ - Parallel AI Search API request format. - Based on: https://docs.parallel.ai/api-reference/search-and-extract-api-beta/search + Parallel AI v1 Search API request format. + Based on: https://docs.parallel.ai/api-reference/search/search """ + search_queries: List[str] # Required - at least one keyword search query objective: str # Optional - natural-language description of search goal - search_queries: List[str] # Optional - list of keyword search queries - processor: str # Optional - search processor ('base', 'pro'), default 'base' - max_results: int # Optional - maximum number of results, default 10 - max_chars_per_result: int # Optional - max characters per result excerpt - source_policy: _ParallelAISourcePolicy # Optional - source policy for allowed/disallowed domains + mode: str # Optional - 'turbo', 'basic', or 'advanced' (default 'advanced') + max_chars_total: int # Optional - upper bound on total excerpt characters + session_id: str # Optional - tracks calls across search/extract requests + client_model: str # Optional - model consuming the results + advanced_settings: _ParallelAIAdvancedSettings + + +LEGACY_PROCESSOR_TO_MODE = {"base": "basic", "pro": "advanced"} class ParallelAISearchConfig(BaseSearchConfig): PARALLEL_AI_API_BASE = "https://api.parallel.ai" - PARALLEL_HEADER_SEARCH_EXTRACT_VALUE = "search-extract-2025-10-10" @staticmethod def ui_friendly_name() -> str: @@ -60,9 +67,6 @@ class ParallelAISearchConfig(BaseSearchConfig): api_base: Optional[str] = None, **kwargs, ) -> Dict: - """ - Validate environment and return headers. - """ api_key = ( api_key or get_secret_str("PARALLEL_AI_API_KEY") @@ -74,7 +78,6 @@ class ParallelAISearchConfig(BaseSearchConfig): ) headers["x-api-key"] = api_key headers["Content-Type"] = "application/json" - headers["parallel-beta"] = self.PARALLEL_HEADER_SEARCH_EXTRACT_VALUE return headers def get_complete_url( @@ -84,32 +87,18 @@ class ParallelAISearchConfig(BaseSearchConfig): data: Optional[Union[Dict, List[Dict]]] = None, **kwargs, ) -> str: - """ - Get complete URL for Search endpoint. - """ api_base = ( api_base or get_secret_str("PARALLEL_AI_API_BASE") or self.PARALLEL_AI_API_BASE ) - # Parallel AI search endpoint is at /v1beta/search - if not api_base.endswith("/v1beta/search"): - if api_base.endswith("/"): - api_base = f"{api_base}v1beta/search" - else: - api_base = f"{api_base}/v1beta/search" + api_base = api_base.rstrip("/") + if not api_base.endswith("/v1/search"): + api_base = f"{api_base.removesuffix('/v1')}/v1/search" return api_base - def _transform_query_to_objective(self, query: Union[str, List[str]]) -> str: - """ - Transform query to objective. - """ - if isinstance(query, list): - return " ".join(query) - return query - def transform_search_request( self, query: Union[str, List[str]], @@ -117,57 +106,78 @@ class ParallelAISearchConfig(BaseSearchConfig): **kwargs, ) -> Dict: """ - Transform Search request to Parallel AI API format. + Transform Search request to Parallel AI v1 API format. Args: query: Search query (string or list of strings) - - If string: maps to `objective` (natural language) + - If string: maps to `search_queries` (single item) and `objective` - If list: maps to `search_queries` (keyword queries) optional_params: Optional parameters for the request - - max_results: Maximum number of search results (default 10) - - search_domain_filter: List of domains to include -> maps to `source_policy.allowed_domains` - - exclude_domains: List of domains to exclude -> maps to `source_policy.disallowed_domains` - - processor: Search processor ('base', 'pro') - - max_chars_per_result: Max characters per result excerpt + - mode: Search mode ('turbo', 'basic', 'advanced'); defaults to 'basic' + - processor: Legacy v1beta param; 'base' maps to mode 'basic', 'pro' to 'advanced' + - max_results: Maximum number of search results -> `advanced_settings.max_results` + - search_domain_filter: Domains to include -> `advanced_settings.source_policy.include_domains` + - exclude_domains: Domains to exclude -> `advanced_settings.source_policy.exclude_domains` + - country: ISO 3166-1 alpha-2 code -> `advanced_settings.location` + - max_chars_per_result: -> `advanced_settings.excerpt_settings.max_chars_per_result` + - Any other params are passed through to the request body as-is Returns: - Dict with typed request data following ParallelAISearchRequest spec + Dict with request data following the v1 search request spec """ + params = dict(optional_params) + request_data: ParallelAISearchRequest = {} - # Map query to objective (string or list both become objective) if isinstance(query, list): - request_data["objective"] = self._transform_query_to_objective(query) + request_data["search_queries"] = query else: + request_data["search_queries"] = [query] request_data["objective"] = query - # Transform Perplexity unified spec parameters to Parallel AI format - if "max_results" in optional_params: - request_data["max_results"] = optional_params["max_results"] + mode = params.pop("mode", None) + processor = params.pop("processor", None) + if mode is None and processor is not None: + mode = LEGACY_PROCESSOR_TO_MODE.get(processor, processor) + # the v1 API defaults to 'advanced' when mode is omitted; default to 'basic' + # instead to keep v1beta's default tier (processor 'base') and litellm's + # $0.004/query cost map entry for `parallel_ai/search` accurate + request_data["mode"] = mode or "basic" + + advanced_settings: _ParallelAIAdvancedSettings = {} + + if "max_results" in params: + advanced_settings["max_results"] = params.pop("max_results") + + if "country" in params: + advanced_settings["location"] = params.pop("country") + + if "max_chars_per_result" in params: + advanced_settings["excerpt_settings"] = { + "max_chars_per_result": params.pop("max_chars_per_result") + } - # Map domain filters to source_policy source_policy: _ParallelAISourcePolicy = {} - if "search_domain_filter" in optional_params: - source_policy["allowed_domains"] = optional_params["search_domain_filter"] + if "search_domain_filter" in params: + source_policy["include_domains"] = params.pop("search_domain_filter") - if "exclude_domains" in optional_params: - source_policy["disallowed_domains"] = optional_params["exclude_domains"] + if "exclude_domains" in params: + source_policy["exclude_domains"] = params.pop("exclude_domains") if source_policy: - request_data["source_policy"] = source_policy + advanced_settings["source_policy"] = source_policy - # Convert to dict before dynamic key assignments - result_data = dict(request_data) + advanced_settings.update(params.pop("advanced_settings", {})) - # pass through all other parameters as-is - for param, value in optional_params.items(): - if ( - param not in self.get_supported_perplexity_optional_params() - and param not in result_data - ): - result_data[param] = value + if advanced_settings: + request_data["advanced_settings"] = advanced_settings + # unified-spec param with no v1 equivalent + params.pop("max_tokens_per_page", None) + + result_data: Dict = dict(request_data) + result_data.update(params) return result_data def transform_search_response( @@ -177,36 +187,27 @@ class ParallelAISearchConfig(BaseSearchConfig): **kwargs, ) -> SearchResponse: """ - Transform Parallel AI API response to LiteLLM unified SearchResponse format. + Transform Parallel AI v1 API response to LiteLLM unified SearchResponse format. - Parallel AI → LiteLLM mappings: - - results[].title → SearchResult.title - - results[].url → SearchResult.url - - results[].excerpts (array) → SearchResult.snippet (joined string) - - No date/last_updated fields in Parallel AI response (set to None) - - Args: - raw_response: Raw httpx response from Parallel AI API - logging_obj: Logging object for tracking - - Returns: - SearchResponse with standardized format + Parallel AI -> LiteLLM mappings: + - results[].title -> SearchResult.title + - results[].url -> SearchResult.url + - results[].excerpts (array) -> SearchResult.snippet (joined string) + - results[].publish_date -> SearchResult.date """ response_json = raw_response.json() - # Transform results to SearchResult objects results = [] for result in response_json.get("results", []): - # Join excerpts array into a single snippet string - excerpts = result.get("excerpts", []) + excerpts = result.get("excerpts") or [] snippet = " ... ".join(excerpts) if excerpts else "" search_result = SearchResult( - title=result.get("title", ""), - url=result.get("url", ""), + title=result.get("title") or "", + url=result.get("url") or "", snippet=snippet, - date=None, # Parallel AI doesn't provide date in response - last_updated=None, # Parallel AI doesn't provide last_updated in response + date=result.get("publish_date"), + last_updated=None, ) results.append(search_result) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 430a789d2a0..3ec7b0814dd 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -3422,14 +3422,18 @@ class ModelResponseIterator: self.has_seen_tool_calls = True break - # Handle final chunk with finishReason but no content. - # _process_candidates skips candidates without "content", - # so the finish_reason from the final chunk is lost. + # _process_candidates skips candidates without a "content" part, so a + # content-less chunk leaves choices empty and the downstream streaming + # handler hits IndexError on choices[0]. This covers the final chunk + # (finishReason, no content) and mid-stream metadata-only chunks + # (grounding/web-search/thought, no content and no finishReason — seen + # with web_search + reasoning) by emitting an empty-delta choice. if not model_response.choices and _candidates: from litellm.types.utils import Delta, StreamingChoices for candidate in _candidates: finish_reason_str = candidate.get("finishReason") + mapped_finish_reason = None if finish_reason_str is not None: if self.has_seen_tool_calls: mapped_finish_reason = "tool_calls" @@ -3437,14 +3441,14 @@ class ModelResponseIterator: mapped_finish_reason = VertexGeminiConfig._check_finish_reason( None, finish_reason_str ) - choice = StreamingChoices( - finish_reason=mapped_finish_reason, - index=candidate.get("index", 0), - delta=Delta(content=None, role=None), - logprobs=None, - enhancements=None, - ) - model_response.choices.append(choice) + choice = StreamingChoices( + finish_reason=mapped_finish_reason, + index=candidate.get("index", 0), + delta=Delta(content=None, role=None), + logprobs=None, + enhancements=None, + ) + model_response.choices.append(choice) # Also handle the case where the final chunk has empty # content (e.g. text:"") WITH finishReason. In this case diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index aab0e4264d0..01a01ea7a76 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -4409,6 +4409,23 @@ "/v1/audio/transcriptions" ] }, + "azure/gpt-realtime-whisper": { + "input_cost_per_second": 0.0002833333333333333, + "litellm_provider": "azure", + "mode": "audio_transcription", + "source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/gpt-realtime-whisper", + "supported_endpoints": [ + "/v1/realtime", + "/v1/realtime/transcription_sessions" + ], + "supported_modalities": [ + "audio" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true + }, "azure/gpt-5.1-2025-11-13": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_priority": 2.5e-07, @@ -7557,6 +7574,45 @@ "supports_function_calling": true, "supports_tool_choice": true }, + "azure_ai/deepseek-v3.1": { + "input_cost_per_token": 1.23e-06, + "litellm_provider": "azure_ai", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.94e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, + "azure_ai/deepseek-v4-pro": { + "input_cost_per_token": 1.74e-06, + "litellm_provider": "azure_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 3.48e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, + "azure_ai/deepseek-v4-flash": { + "input_cost_per_token": 1.9e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 5.1e-07, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "azure_ai/embed-v-4-0": { "input_cost_per_token": 1.2e-07, "litellm_provider": "azure_ai", @@ -40916,6 +40972,23 @@ "supports_system_messages": true, "supports_tool_choice": true }, + "gpt-realtime-whisper": { + "input_cost_per_second": 0.0002833333333333333, + "litellm_provider": "openai", + "mode": "audio_transcription", + "source": "https://platform.openai.com/docs/models/gpt-realtime-whisper", + "supported_endpoints": [ + "/v1/realtime", + "/v1/realtime/transcription_sessions" + ], + "supported_modalities": [ + "audio" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true + }, "sora-2": { "litellm_provider": "openai", "mode": "video_generation", diff --git a/litellm/proxy/_experimental/mcp_server/exceptions.py b/litellm/proxy/_experimental/mcp_server/exceptions.py index fd8fc3d5e58..a00e797a6bd 100644 --- a/litellm/proxy/_experimental/mcp_server/exceptions.py +++ b/litellm/proxy/_experimental/mcp_server/exceptions.py @@ -10,8 +10,9 @@ class MCPUpstreamAuthError(Exception): (typically HTTP 401) and the gateway should surface it transparently to the client instead of swallowing it. - Only relevant for pass-through MCP servers (see - ``MCPServer.is_oauth_passthrough``). The gateway converts this exception + Relevant for MCP servers that delegate OAuth to the upstream server, + including pass-through servers and OAuth2 servers with + ``delegate_auth_to_upstream`` enabled. The gateway converts this exception into an HTTP 401 response on single-server routes, preserving any ``WWW-Authenticate`` challenge emitted by the upstream so standards- compliant MCP clients can trigger the upstream OAuth flow. diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 73935beeb3a..5e419b5c0a3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2777,28 +2777,40 @@ class MCPServerManager: Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details. - For pass-through MCP servers (``MCPServer.is_oauth_passthrough``) an + For OAuth pass-through and upstream-delegated OAuth2 MCP servers, an upstream HTTP 401 is converted into :class:`MCPUpstreamAuthError` instead of being swallowed to an empty tool list. That lets the single-server HTTP routes surface a proper 401 + ``WWW-Authenticate`` challenge so standards-compliant MCP clients trigger the upstream - OAuth flow. Non-pass-through servers keep today's swallow-and-log - behaviour so the multi-server ``/mcp`` aggregator doesn't get - tainted by a single bad server. + OAuth flow. Other servers keep today's swallow-and-log behaviour so + the multi-server ``/mcp`` aggregator doesn't get tainted by a single + bad server. Args: client: MCP client instance server_name: Name of the server for logging - server: Optional MCPServer; when pass-through, auth errors are - re-raised as :class:`MCPUpstreamAuthError`. + server: Optional MCPServer; when upstream auth is delegated, auth + errors are re-raised as :class:`MCPUpstreamAuthError`. Returns: List of tools from the server """ - is_passthrough = bool(server is not None and server.is_oauth_passthrough) + should_surface_upstream_auth = bool( + server is not None + and ( + server.is_oauth_passthrough + or ( + server.auth_type == MCPAuth.oauth2 + and getattr(server, "delegate_auth_to_upstream", False) is True + and not server.has_client_credentials + ) + ) + ) try: with anyio.fail_after(MCP_TOOL_LISTING_TIMEOUT): - tools = await client.list_tools(raise_on_error=is_passthrough) + tools = await client.list_tools( + raise_on_error=should_surface_upstream_auth + ) verbose_logger.debug(f"Tools from {server_name}: {tools}") return tools except TimeoutError: @@ -2815,12 +2827,12 @@ class MCPServerManager: ) return [] except Exception as e: - if is_passthrough: + if should_surface_upstream_auth: auth_info = _extract_upstream_auth_failure(e) if auth_info is not None: status_code, www_authenticate = auth_info verbose_logger.info( - f"Upstream auth failure from pass-through MCP server " + f"Upstream auth failure from MCP server " f"{server_name}: HTTP {status_code}" ) raise MCPUpstreamAuthError( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 731493b1337..6288ef149b0 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2552,6 +2552,7 @@ if MCP_AVAILABLE: standard_logging_mcp_tool_call ) litellm_logging_obj.model = f"MCP: {name}" + litellm_logging_obj.model_call_details["model"] = f"MCP: {name}" # Resolve the MCP server early so BYOK checks and credential injection # apply to ALL dispatch paths (local tool registry AND managed MCP server). if mcp_server is None: @@ -3426,6 +3427,8 @@ if MCP_AVAILABLE: ) if stored_oauth_headers: continue + if getattr(server, "delegate_auth_to_upstream", False) is True: + continue request = StarletteRequest(scope) base_url = get_request_base_url(request) @@ -3960,7 +3963,7 @@ if MCP_AVAILABLE: ): _stateful_session_locks.pop(active_request_session_id, None) except MCPUpstreamAuthError as e: - # Pass-through server returned 401 — surface it to the client so + # Upstream delegated auth returned 401; surface it to the client so # standards-compliant MCP clients trigger the upstream OAuth flow. raise e.to_http_exception( base_url=get_request_base_url(StarletteRequest(scope)), @@ -4076,7 +4079,7 @@ if MCP_AVAILABLE: ): await sse_session_manager.handle_request(scope, receive, send) except MCPUpstreamAuthError as e: - # Pass-through server returned 401 — surface it to the client so + # Upstream delegated auth returned 401; surface it to the client so # standards-compliant MCP clients trigger the upstream OAuth flow. raise e.to_http_exception( base_url=get_request_base_url(StarletteRequest(scope)), diff --git a/litellm/proxy/_experimental/out/assets/logos/newrelic.png b/litellm/proxy/_experimental/out/assets/logos/newrelic.png new file mode 100644 index 00000000000..c841e3e7136 Binary files /dev/null and b/litellm/proxy/_experimental/out/assets/logos/newrelic.png differ diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 35e9e0cd74b..a9b8319745c 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1856,6 +1856,10 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase): key: str # required +class BlockModelRequest(LiteLLMPydanticObjectBase): + model_id: str # required + + class AddTeamCallback(LiteLLMPydanticObjectBase): callback_name: str callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = ( @@ -2231,8 +2235,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): health_check_concurrency: Optional[int] = Field( None, description=( - "limit concurrent health checks per cycle; when unset, " - "health checks run without a concurrency cap" + "limit concurrent health checks per cycle; when unset, health checks run without a concurrency cap" ), ) health_check_skip_disabled_background_models: bool = Field( @@ -3094,6 +3097,14 @@ class AllCallbacks(LiteLLMPydanticObjectBase): ui_callback_name="Galileo", ) + newrelic: CallbackOnUI = CallbackOnUI( + litellm_callback_name="newrelic", + ui_callback_name="New Relic", + litellm_callback_params=[ + "NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED", + ], + ) + class SpendLogsMetadata(TypedDict): """ diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 6eae9d0d475..aa967732a90 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3155,6 +3155,98 @@ async def can_key_call_model( raise +async def can_key_call_resolved_model( + model: str, + llm_model_list: Optional[list], + valid_token: UserAPIKeyAuth, + llm_router: Optional[litellm.Router], +) -> None: + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + skip_key_model_check = valid_token.config or ( + isinstance(valid_token.models, list) + and SpecialModelNames.all_team_models.value in valid_token.models + ) + if not skip_key_model_check: + await can_key_call_model( + model=model, + llm_model_list=llm_model_list, + valid_token=valid_token, + llm_router=llm_router, + ) + + team_object: Optional[LiteLLM_TeamTableCachedObj] = None + team_object_from_lookup = False + if valid_token.team_id is not None: + try: + team_object = await get_team_object( + team_id=valid_token.team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=valid_token.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + team_object_from_lookup = True + except Exception: + team_object = LiteLLM_TeamTableCachedObj( + team_id=valid_token.team_id, + models=valid_token.team_models, + blocked=valid_token.team_blocked, + team_alias=valid_token.team_alias, + metadata=valid_token.team_metadata, + object_permission_id=valid_token.team_object_permission_id, + object_permission=valid_token.team_object_permission, + ) + + if team_object is not None: + try: + await can_team_access_model( + model=model, + team_object=team_object, + llm_router=llm_router, + team_model_aliases=valid_token.team_model_aliases, + ) + except ProxyException as team_denial: + if team_denial.type != ProxyErrorTypes.team_model_access_denied: + raise + if not await _key_access_group_grants_model( + model=model, + valid_token=valid_token, + team_object=team_object, + llm_router=llm_router, + ): + raise + + if valid_token.user_id is not None and team_object_from_lookup: + await _check_team_member_model_access( + model=model, + team_object=team_object, + valid_token=valid_token, + llm_router=llm_router, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + if valid_token.project_id is not None: + project_object = await get_project_object( + project_id=valid_token.project_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if project_object is not None and len(project_object.models) > 0: + can_project_access_model( + model=model, + project_object=project_object, + llm_router=llm_router, + ) + + def can_org_access_model( model: str, org_object: Optional[LiteLLM_OrganizationTable], diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index b9a9f3cebb7..c257073b4fa 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2075,6 +2075,17 @@ class ProxyBaseLLMRequestProcessing: ) ) yield serialize_chunk(chunk) + except (asyncio.CancelledError, GeneratorExit): + # Client disconnected mid-stream. CancelledError / GeneratorExit + # are BaseException and bypass the success/failure logging + # callbacks that release the pre-call max_parallel_requests +1; + # release it here. This is the outermost generator Starlette closes + # on disconnect, so the nested iterator hook (which only sees + # GeneratorExit on GC) cannot own the refund. + proxy_logging_obj._release_max_parallel_requests_on_disconnect( + user_api_key_dict + ) + raise except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format( diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 765c419479e..4a550cb73a4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -49,6 +49,7 @@ from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessag from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( BedrockContentItem, BedrockGuardrailOutput, + BedrockGuardrailQualifier, BedrockGuardrailResponse, BedrockRequest, BedrockTextContent, @@ -74,6 +75,29 @@ from litellm.types.utils import ( GUARDRAIL_NAME = "bedrock" _BEDROCK_DYNAMIC_BODY_DENYLIST = frozenset({"content", "source"}) +# Maps an OpenAI message content-block ``type`` to the Bedrock guardrail qualifier +# it represents, so callers can drive contextual grounding by tagging their content. +# The model response is qualified as ``guard_content`` directly by the OUTPUT builder; +# the existing ``guarded_text`` marker is intentionally left unmapped here so its +# guardrail-hook payload is unchanged by this feature. +_CONTENT_TYPE_TO_QUALIFIER: Dict[str, BedrockGuardrailQualifier] = { + "grounding_source": "grounding_source", + "query": "query", +} + +# Roles whose ``grounding_source`` blocks are trusted as reference material for the +# contextual-grounding check. Only app-authored roles qualify: ``tool``/``function`` +# results and ``user`` content can carry caller- or externally-influenced text, which +# must not be graded against as if it were the application's own source material. +_GROUNDING_SOURCE_TRUSTED_ROLES = frozenset({"system", "developer"}) + + +class QualifiedTextBlock(NamedTuple): + """A piece of message text paired with its Bedrock grounding qualifier (if any).""" + + text: str + qualifier: Optional[BedrockGuardrailQualifier] + class GuardrailMessageFilterResult(NamedTuple): payload_messages: Optional[List[AllMessageValues]] @@ -164,41 +188,71 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if messages is None: return bedrock_request for message in messages: - message_text_content: Optional[List[str]] = self.get_content_for_message( - message=message - ) - if message_text_content is None: + blocks = self.get_content_items_for_message(message=message) + if blocks is None: continue - for text_content in message_text_content: - bedrock_content_item = BedrockContentItem( - text=BedrockTextContent(text=text_content) + for block in blocks: + # INPUT scans send plain text only. Grounding qualifiers are attached + # exclusively when assembling the OUTPUT request, so a caller cannot use + # a grounding_source/query tag to change how input-safety policies treat + # their content (which would be an input-guardrail bypass). + bedrock_request_content.append( + BedrockContentItem(text=BedrockTextContent(text=block.text)) ) - bedrock_request_content.append(bedrock_content_item) bedrock_request["content"] = bedrock_request_content return bedrock_request def _create_bedrock_output_content_request( - self, response: Union[Any, ModelResponse] + self, + response: Union[Any, ModelResponse], + messages: Optional[List[AllMessageValues]] = None, ) -> BedrockRequest: """ Create a bedrock request for the output content - the LLM response. + + Contextual grounding grades the response against the reference source and + the user query from the request. When the request tagged any + ``grounding_source``/``query`` blocks, they are emitted first and the + response is qualified as ``guard_content`` so Bedrock can score grounding. + Without such tags the payload is the legacy single response block. """ bedrock_request: BedrockRequest = BedrockRequest(source="OUTPUT") - bedrock_request_content: List[BedrockContentItem] = [] - if isinstance(response, litellm.ModelResponse): - for choice in response.choices: - if isinstance(choice, litellm.Choices): - if choice.message.content and isinstance( - choice.message.content, str - ): - bedrock_content_item = BedrockContentItem( - text=BedrockTextContent(text=choice.message.content) - ) - bedrock_request_content.append(bedrock_content_item) - bedrock_request["content"] = bedrock_request_content + grounding_blocks = self._collect_grounding_blocks(messages) + bedrock_request_content: List[BedrockContentItem] = [ + self._build_content_item(block) for block in grounding_blocks + ] + has_grounding = len(bedrock_request_content) > 0 + # Append the response (the content to guard) after any grounding blocks; assign + # unconditionally so harvested grounding blocks survive a non-ModelResponse input. + bedrock_request_content.extend( + self._build_response_content_items(response, has_grounding=has_grounding) + ) + bedrock_request["content"] = bedrock_request_content return bedrock_request + def _build_response_content_items( + self, response: Union[Any, ModelResponse], has_grounding: bool + ) -> List[BedrockContentItem]: + """Build content item(s) from the model response. When the request supplied + grounding, the response is qualified ``guard_content`` so Bedrock can score it. + """ + items: List[BedrockContentItem] = [] + if not isinstance(response, litellm.ModelResponse): + return items + for choice in response.choices: + if ( + isinstance(choice, litellm.Choices) + and isinstance(choice.message.content, str) + and choice.message.content + ): + block = QualifiedTextBlock( + text=choice.message.content, + qualifier="guard_content" if has_grounding else None, + ) + items.append(self._build_content_item(block)) + return items + def convert_to_bedrock_format( self, source: Literal["INPUT", "OUTPUT"], @@ -221,10 +275,68 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) elif source == "OUTPUT": bedrock_request = self._create_bedrock_output_content_request( - response=response + response=response, messages=messages ) return bedrock_request + def get_content_items_for_message( + self, message: AllMessageValues + ) -> Optional[List[QualifiedTextBlock]]: + """ + Flatten a message into text blocks, preserving any contextual-grounding + qualifier carried by the content-block ``type`` (grounding_source / query). + Untagged text keeps ``qualifier=None`` so the payload is unchanged for + callers that do not use grounding. + """ + content = message.get("content") + if content is None: + return None + blocks: List[QualifiedTextBlock] = [] + if isinstance(content, str): + blocks.append(QualifiedTextBlock(text=content, qualifier=None)) + elif isinstance(content, list): + for item in content: + if isinstance(item, dict) and "text" in item: + qualifier = _CONTENT_TYPE_TO_QUALIFIER.get(item.get("type", "")) + blocks.append( + QualifiedTextBlock(text=item["text"], qualifier=qualifier) + ) + elif isinstance(item, str): + blocks.append(QualifiedTextBlock(text=item, qualifier=None)) + return blocks + + def _build_content_item(self, block: QualifiedTextBlock) -> BedrockContentItem: + """Build a Bedrock content item, attaching qualifiers only when present.""" + text_content = BedrockTextContent(text=block.text) + if block.qualifier is not None: + text_content["qualifiers"] = [block.qualifier] + return BedrockContentItem(text=text_content) + + def _collect_grounding_blocks( + self, messages: Optional[List[AllMessageValues]] + ) -> List[QualifiedTextBlock]: + """Harvest grounding_source/query blocks from the request for an OUTPUT scan. + + ``grounding_source`` is honored only from app-authored roles (system / + developer). A grounding_source tag on a ``user``, ``tool`` or ``function`` + message is ignored, so neither a forwarded end-user message nor a tool/function + result carrying externally-influenced content can supply fake evidence for the + contextual-grounding check to grade the response against. ``query`` is accepted + from any role (it is the user's question). + """ + grounding: List[QualifiedTextBlock] = [] + for message in messages or []: + role = message.get("role") + for block in self.get_content_items_for_message(message=message) or []: + if block.qualifier == "query": + grounding.append(block) + elif ( + block.qualifier == "grounding_source" + and role in _GROUNDING_SOURCE_TRUSTED_ROLES + ): + grounding.append(block) + return grounding + def _prepare_guardrail_messages_for_role( self, messages: Optional[List[AllMessageValues]], @@ -1169,6 +1281,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): output_content_bedrock = await self.make_bedrock_api_request( source="OUTPUT", response=response, + messages=new_messages, request_data=data, logging_event_type=GuardrailEventHooks.post_call, ) @@ -1281,6 +1394,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): output_guardrail_response = await self.make_bedrock_api_request( source="OUTPUT", response=assembled_model_response, + messages=request_data.get("messages"), request_data=request_data, logging_event_type=GuardrailEventHooks.post_call, ) @@ -1414,28 +1528,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return new_content, masking_index - def get_content_for_message(self, message: AllMessageValues) -> Optional[List[str]]: - """ - Get the content for a message. - - For bedrock guardrails we create a list of all the text content in the message. - - If a message has a list of content items, we flatten the list and return a list of text content. - """ - message_text_content = [] - content = message.get("content") - if content is None: - return None - if isinstance(content, str): - message_text_content.append(content) - elif isinstance(content, list): - for item in content: - if isinstance(item, dict) and "text" in item: - message_text_content.append(item["text"]) - elif isinstance(item, str): - message_text_content.append(item) - return message_text_content - def _apply_masking_to_response( self, response: Union[ModelResponse, Any], diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py new file mode 100644 index 00000000000..b73572e4ed5 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py @@ -0,0 +1,46 @@ +"""Ovalix guardrail hook: registration and initialization for the proxy.""" + +from typing import TYPE_CHECKING + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .ovalix import OvalixGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): + """Create and register an Ovalix guardrail callback from proxy config.""" + import litellm + + tracker_api_base = getattr(litellm_params, "tracker_api_base", None) + tracker_api_key = getattr(litellm_params, "tracker_api_key", None) + application_id = getattr(litellm_params, "application_id", None) + pre_checkpoint_id = getattr(litellm_params, "pre_checkpoint_id", None) + post_checkpoint_id = getattr(litellm_params, "post_checkpoint_id", None) + + _ovalix_callback = OvalixGuardrail( + guardrail_name=guardrail.get("guardrail_name", ""), + tracker_api_base=tracker_api_base, + tracker_api_key=tracker_api_key, + application_id=application_id, + pre_checkpoint_id=pre_checkpoint_id, + post_checkpoint_id=post_checkpoint_id, + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + litellm.logging_callback_manager.add_litellm_callback(_ovalix_callback) + + return _ovalix_callback + + +# Registry of guardrail name -> initializer for proxy config loading. +guardrail_initializer_registry = { + SupportedGuardrailIntegrations.OVALIX.value: initialize_guardrail, +} + +# Registry of guardrail name -> guardrail class (e.g. for apply_guardrail API). +guardrail_class_registry = { + SupportedGuardrailIntegrations.OVALIX.value: OvalixGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py new file mode 100644 index 00000000000..2ebbeb31c0b --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py @@ -0,0 +1,330 @@ +"""Ovalix guardrail integration: pre- and post-call checks via the Tracker service. + +Use Ovalix Guardrails for your LLM calls. Supports pre_call (user input) and +post_call (model output) checkpoints with optional correction/blocking. +""" + +import datetime +import hashlib +import os +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.exceptions import GuardrailRaisedException +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs + +if TYPE_CHECKING: + from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel + + +BLOCKED_BY_OVALIX_FALLBACK_MESSAGE = "This message was blocked by Ovalix" +BLOCKED_ACTION_TYPE = "block" + + +class OvalixGuardrailMissingSecrets(Exception): + """Raised when required Ovalix config (API base, key, application/checkpoint IDs) is missing.""" + + pass + + +class OvalixGuardrailBlockedException(GuardrailRaisedException): + """ + Raised when Ovalix blocks a message. Sets status_code=400 so the proxy + returns 400 and HTTP clients do not retry (they retry on 5xx). + """ + + status_code = 400 + + def __init__( + self, + guardrail_name: Optional[str] = None, + message: str = "", + should_wrap_with_default_message: bool = True, + ): + super().__init__( + guardrail_name=guardrail_name, + message=message, + should_wrap_with_default_message=should_wrap_with_default_message, + ) + + +class OvalixGuardrail(CustomGuardrail): + """ + Ovalix guardrail: pre-prompt (pre_call) and post-prompt (post_call) checks + via the Tracker service, with application and checkpoint resolution from the + Monolith backend. + """ + + def __init__( + self, + tracker_api_base: Optional[str] = None, + tracker_api_key: Optional[str] = None, + application_id: Optional[str] = None, + pre_checkpoint_id: Optional[str] = None, + post_checkpoint_id: Optional[str] = None, + **kwargs: Any, + ): + self._tracker_api_base = tracker_api_base or os.environ.get( + "OVALIX_TRACKER_API_BASE" + ) + self._tracker_api_key = tracker_api_key or os.environ.get( + "OVALIX_TRACKER_API_KEY" + ) + self._application_id = application_id or os.environ.get("OVALIX_APPLICATION_ID") + self._pre_checkpoint_id = pre_checkpoint_id or os.environ.get( + "OVALIX_PRE_CHECKPOINT_ID" + ) + self._post_checkpoint_id = post_checkpoint_id or os.environ.get( + "OVALIX_POST_CHECKPOINT_ID" + ) + + if "supported_event_hooks" not in kwargs: + kwargs["supported_event_hooks"] = [] + + self._validate_config(kwargs["supported_event_hooks"]) + + self._tracker_headers = httpx.Headers( + { + "Authorization": f"Bearer {self._tracker_api_key}", + "Content-Type": "application/json", + }, + encoding="utf-8", + ) + + self._async_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) + + super().__init__(**kwargs) + verbose_proxy_logger.debug( + "Ovalix Guardrail initialized: tracker=%s, application_id=%s, pre_checkpoint_id=%s, post_checkpoint_id=%s", + self._tracker_api_base, + self._application_id, + self._pre_checkpoint_id, + self._post_checkpoint_id, + ) + + def _validate_config( + self, supported_event_hooks: List[GuardrailEventHooks] + ) -> None: + """Ensure required secrets and checkpoint IDs are set; auto-add hooks when IDs are present.""" + errors: List[str] = [] + + if not self._tracker_api_base: + errors.append( + "Tracker API base, set OVALIX_TRACKER_API_BASE or pass tracker_api_base" + ) + if not self._tracker_api_key: + errors.append( + "Tracker API key, set OVALIX_TRACKER_API_KEY or pass tracker_api_key" + ) + if not self._application_id: + errors.append( + "Application ID, set OVALIX_APPLICATION_ID or pass application_id" + ) + if ( + not self._pre_checkpoint_id + and GuardrailEventHooks.pre_call in supported_event_hooks + ): + errors.append( + "Pre-checkpoint ID, set OVALIX_PRE_CHECKPOINT_ID or pass pre_checkpoint_id" + ) + if ( + not self._post_checkpoint_id + and GuardrailEventHooks.post_call in supported_event_hooks + ): + errors.append( + "Post-checkpoint ID, set OVALIX_POST_CHECKPOINT_ID or pass post_checkpoint_id" + ) + if not self._pre_checkpoint_id and not self._post_checkpoint_id: + errors.append( + "Pre-checkpoint ID or Post-checkpoint ID, set OVALIX_PRE_CHECKPOINT_ID or OVALIX_POST_CHECKPOINT_ID or pass pre_checkpoint_id or post_checkpoint_id" + ) + + if errors: + raise OvalixGuardrailMissingSecrets( + "Missing Ovalix guardrail configuration errors: " + ". ".join(errors) + ) + + # auto-add hooks when checkpoint IDs are present + if ( + self._pre_checkpoint_id + and GuardrailEventHooks.pre_call not in supported_event_hooks + ): + supported_event_hooks.append(GuardrailEventHooks.pre_call) + if ( + self._post_checkpoint_id + and GuardrailEventHooks.post_call not in supported_event_hooks + ): + supported_event_hooks.append(GuardrailEventHooks.post_call) + + def _get_actor(self, data: dict) -> str: + """Return a stable actor identifier from request metadata (e.g. user email or id).""" + metadata = data.get("metadata") or data.get("litellm_metadata") or {} + if metadata.get("user_api_key_user_email"): + return metadata["user_api_key_user_email"] + if metadata.get("user_api_key_user_id"): + return metadata["user_api_key_user_id"] + return "unknown" + + def _get_tracker_actor_id(self, data: dict) -> str: + """Normalize the actor string into a short, stable id for Tracker API payloads.""" + # NOTE: this hash is purely for normalization — it collapses an arbitrary actor + # string (email, user id, or "unknown") into a compact, fixed-length, consistent + # key. It is not a privacy/security measure and the actor value is not sensitive, + # so a plain SHA-256 (truncated) is sufficient; no salting/KDF is needed here. + actor_id = self._get_actor(data).encode() + normalized_actor_id = hashlib.sha256(actor_id).hexdigest()[:8] + return normalized_actor_id + + def _get_session_id(self, data: dict) -> str: + """Return a unique identifier for the chat/session (actor + date + application_id).""" + actor_hash = self._get_tracker_actor_id(data) + today = datetime.datetime.now(datetime.timezone.utc).strftime("%Y-%m-%d") + return f"{actor_hash}_{today}_{self._application_id}" + + async def _call_checkpoint( + self, + content: str, + checkpoint_id: str, + actor: str, + session_id: str, + ) -> Dict[str, Any]: + """Call the Ovalix Tracker checkpoint API and return the JSON response.""" + application_id = self._application_id + if not application_id or not checkpoint_id: + raise ValueError("Ovalix: application_id or checkpoint_id not resolved") + + url = f"{self._tracker_api_base}/tracking/custom_application/checkpoint" + headers = dict(self._tracker_headers) + payload = { + "application_id": application_id, + "checkpoint_id": checkpoint_id, + "actor": actor, + "session_id": session_id, + "data_type": "TEXT", + "data": {"content": content}, + } + response = await self._async_handler.post(url, headers=headers, json=payload) + response.raise_for_status() + return response.json() + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + """ + Apply Ovalix guardrail to the given inputs (request or response text). + + Used by the unified guardrail flow and the /apply_guardrail API. + For "request", uses the pre-checkpoint; for "response", uses the post-checkpoint. + + Args: + inputs: Guardrail API inputs (e.g. texts to check). + request_data: Full request payload (messages, metadata, response). + input_type: "request" (pre_call) or "response" (post_call). + logging_obj: Optional logging context. + + Returns: + Updated inputs (e.g. with replaced/corrected texts, or unchanged). + """ + if not self._pre_checkpoint_id and not self._post_checkpoint_id: + return inputs + + tracker_actor_id = self._get_tracker_actor_id(request_data) + session_id = self._get_session_id(request_data) + texts = inputs.get("texts") or [] + if not texts or not isinstance(texts, list): + return inputs + + if input_type == "response": + if not self._post_checkpoint_id: + return inputs + corrected_llm_responses = await self._generate_post_guardrail_llm_texts( + texts, tracker_actor_id, session_id, self._post_checkpoint_id + ) + return {**inputs, "texts": corrected_llm_responses} + + if self._pre_checkpoint_id: + post_guardrail_texts = await self._generate_post_guardrail_llm_texts( + texts, tracker_actor_id, session_id, self._pre_checkpoint_id + ) + return {**inputs, "texts": post_guardrail_texts} + return inputs + + async def _generate_post_guardrail_llm_texts( + self, texts: List[str], actor: str, session_id: str, checkpoint_id: str + ) -> List[str]: + """Generate post-guardrail LLM responses for the given LLM responses.""" + post_guardrail_texts: List[str] = [] + + is_first_response = True + for llm_response in reversed(texts): + try: + resp = await self._call_checkpoint( + llm_response, checkpoint_id, actor, session_id + ) + except Exception as e: + verbose_proxy_logger.exception( + "Ovalix apply_guardrail checkpoint call failed: %s", e + ) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"Ovalix guardrail error: {e!s}", + should_wrap_with_default_message=False, + ) from e + + action_type = (resp.get("action_type") or "").lower() + blocking_message = ( + self._get_trackers_corrected_message(resp) + or BLOCKED_BY_OVALIX_FALLBACK_MESSAGE + ) + if action_type == BLOCKED_ACTION_TYPE and is_first_response: + self._block_current_message(blocking_message) + elif action_type == BLOCKED_ACTION_TYPE: + post_guardrail_texts.insert(0, blocking_message) + else: + corrected_text = ( + self._get_trackers_corrected_message(resp) or llm_response + ) + post_guardrail_texts.insert(0, corrected_text) + is_first_response = False + return post_guardrail_texts + + def _block_current_message(self, blocking_message: str) -> None: + """Raise OvalixGuardrailBlockedException with the given message (no default wrapper).""" + raise OvalixGuardrailBlockedException( + guardrail_name=self.guardrail_name, + message=blocking_message, + should_wrap_with_default_message=False, + ) + + def _get_trackers_corrected_message(self, resp: dict) -> Optional[str]: + """Extract corrected/blocking message content from Tracker checkpoint response.""" + modified = resp.get("modified_data") + if isinstance(modified, dict) and "content" in modified: + return modified["content"] + return None + + @staticmethod + def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( + OvalixGuardrailConfigModel, + ) + + return OvalixGuardrailConfigModel diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 6ef8bbc4006..507e8e4d4da 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -130,6 +130,7 @@ services = Union[ "generic_api", "arize", "galileo", + "newrelic", "sqs", ], str, @@ -208,6 +209,7 @@ async def health_services_endpoint( # noqa: PLR0915 "generic_api", "arize", "galileo", + "newrelic", "sqs", ]: raise HTTPException( @@ -325,6 +327,26 @@ async def health_services_endpoint( # noqa: PLR0915 "status": "success", "message": "Mock LLM request made - check langfuse.", } + elif service == "newrelic": + if not _is_proxy_admin(user_api_key_dict): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": "Only proxy admins can trigger the New Relic test event." + }, + ) + from litellm.integrations.newrelic.newrelic import NewRelicLogger + + newrelic_logger = NewRelicLogger() + response = await newrelic_logger.async_health_check() + return { + "status": response["status"], + "message": ( + response["error_message"] + if response["status"] == "unhealthy" + else "New Relic is healthy — test event sent" + ), + } if service == "webhook": user_info = CallInfo( diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index f45c63d1380..85d034b7a41 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -2954,6 +2954,46 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): f"Error in rate limit failure event: {str(e)}" ) + async def async_release_max_parallel_requests_on_disconnect( + self, user_api_key_dict: UserAPIKeyAuth + ) -> None: + """ + Release the api-key ``max_parallel_requests`` slot that + ``async_pre_call_hook`` reserved, for a request that ended without + either logging callback firing. + + The +1 is normally undone by ``async_log_success_event`` (natural + stream completion) or ``async_log_failure_event`` (LLM error). When a + client cancels a stream mid-flight, the cancellation surfaces as + ``asyncio.CancelledError`` / ``GeneratorExit`` and neither callback + runs, so without this the counter leaks one slot per cancelled stream + until the key wedges at its limit. + """ + if ( + not user_api_key_dict.api_key + or user_api_key_dict.max_parallel_requests is None + ): + return + + await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( + increment_list=[ + RedisPipelineIncrementOperation( + key=self.create_rate_limit_keys( + key="api_key", + value=user_api_key_dict.api_key, + rate_limit_type="max_parallel_requests", + ), + increment_value=-1, + # Refresh the window TTL on the decrement, matching the + # failure path. max_parallel_requests is a concurrency + # gauge, not a rolling-window count, so the key must + # outlive in-flight requests rather than expire mid-stream. + ttl=self.window_size, + ) + ], + litellm_parent_otel_span=None, + ) + async def async_post_call_success_hook( self, data: dict, user_api_key_dict: UserAPIKeyAuth, response ): diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 1ef0ad9fd47..13107b68864 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -924,11 +924,21 @@ async def get_daily_activity( where=where_conditions ) - # Fetch paginated results + # Fetch paginated results. + # ``date`` alone is not a unique sort key -- a busy tenant has many + # rows per date (one per api_key, model, model_group, provider, + # endpoint, ...), so offset pagination over ``date desc`` lands on + # arbitrary boundaries and the same row can be skipped on one page + # and returned on another. A client that pages through and sums the + # per-page metrics (the Usage dashboard) then gets a non-deterministic + # total. Adding ``id`` (the row's UUID primary key, present on both + # LiteLLM_DailyUserSpend and LiteLLM_DailyTeamSpend) as a tiebreaker + # gives every page a stable cursor (#30164). daily_spend_data = await getattr(prisma_client.db, table_name).find_many( where=where_conditions, order=[ {"date": "desc"}, + {"id": "asc"}, ], skip=(page - 1) * page_size, take=page_size, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index eba16c077b0..2f239c8da84 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1860,18 +1860,23 @@ async def prepare_key_update_data( non_default_values["budget_reset_at"] = key_reset_at non_default_values["budget_duration"] = budget_duration - if "budget_limits" in non_default_values and non_default_values["budget_limits"]: - from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time - + if "budget_limits" in non_default_values: raw_windows = non_default_values["budget_limits"] - initialized_windows = [] - for window in raw_windows: - w = window if isinstance(window, dict) else window.model_dump() - w["reset_at"] = get_budget_reset_time( - budget_duration=w["budget_duration"] - ).isoformat() - initialized_windows.append(w) - non_default_values["budget_limits"] = json.dumps(initialized_windows) + if raw_windows: + from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time + + initialized_windows = [] + for window in raw_windows: + w = window if isinstance(window, dict) else window.model_dump() + w["reset_at"] = get_budget_reset_time( + budget_duration=w["budget_duration"] + ).isoformat() + initialized_windows.append(w) + non_default_values["budget_limits"] = json.dumps(initialized_windows) + else: + # [] / None clears the field; prisma-client-py has no DbNull + # sentinel for Json? columns, so store the JSON literal null + non_default_values["budget_limits"] = json.dumps(None) if "object_permission" in non_default_values: non_default_values = await _handle_update_object_permission( @@ -2248,14 +2253,18 @@ async def _validate_update_key_data( # - Anyone else (non-PROXY_ADMIN, not the owner, not a team member # on a team key): must pass _check_key_admin_access (PROXY_ADMIN # / key-owner / team-admin / org-admin of the key). - # - max_budget / spend: always require the admin check, even for the - # key owner or a team member (matches the existing admin-only - # budget semantics). + # - max_budget / spend / budget_limits: always require the admin + # check, even for the key owner or a team member (matches the + # existing admin-only budget semantics). budget_limits uses + # model_fields_set because an explicit null/[] clears the field + # and must gate the same as setting or changing it. _is_budget_change = ( - data.max_budget is not None and data.max_budget != existing_key_row.max_budget - ) or ( - data.spend is not None - and data.spend != getattr(existing_key_row, "spend", None) + (data.max_budget is not None and data.max_budget != existing_key_row.max_budget) + or ( + data.spend is not None + and data.spend != getattr(existing_key_row, "spend", None) + ) + or "budget_limits" in data.model_fields_set ) # Personal-key bypass: the caller both created the key AND still owns it diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 0be476469a6..def6e271635 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -15,13 +15,14 @@ import datetime import json from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast -from fastapi import APIRouter, Depends, HTTPException, Request, status +from fastapi import APIRouter, Depends, HTTPException, Header, Request, status from pydantic import BaseModel, ConfigDict, Field from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.constants import LITELLM_PROXY_ADMIN_NAME from litellm.proxy._types import ( + BlockModelRequest, CommonProxyErrors, LiteLLM_ProxyModelTable, LiteLLM_TeamTable, @@ -331,6 +332,168 @@ async def patch_model( ) +async def _set_model_blocked_status( + data: BlockModelRequest, + user_api_key_dict: UserAPIKeyAuth, + blocked: bool, + action: Literal["blocked", "unblocked"], + litellm_changed_by: Optional[str], +) -> Optional[LiteLLM_ProxyModelTable]: + from litellm.proxy.proxy_server import ( + litellm_proxy_admin_name, + llm_router, + prisma_client, + store_model_in_db, + ) + + try: + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + if store_model_in_db is not True: + raise ProxyException( + message="Model updates only supported for DB-stored models", + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param=None, + ) + + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise ProxyException( + message="Only proxy admins can change a model's blocked flag.", + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param="blocked", + ) + + db_model = await get_db_model( + model_id=data.model_id, + prisma_client=prisma_client, + ) + + if db_model is None: + if ( + llm_router + and llm_router.get_deployment(model_id=data.model_id) is not None + ): + raise ProxyException( + message="Cannot edit config-based model. Store model in DB via /model/new first.", + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param=None, + ) + raise ProxyException( + message=f"Model {data.model_id} not found on proxy.", + type=ProxyErrorTypes.not_found_error, + code=status.HTTP_404_NOT_FOUND, + param=None, + ) + + updated_model = await ModelRepository(prisma_client).table.update( + where={"model_id": data.model_id}, + data={ + "blocked": blocked, + "updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name, + "updated_at": cast(str, get_utc_datetime()), + }, + ) + + await clear_cache() + + asyncio.create_task( + create_object_audit_log( + object_id=data.model_id, + action=action, + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, + before_value=db_model.model_dump_json(exclude_none=True), + after_value=( + updated_model.model_dump_json(exclude_none=True) + if isinstance(updated_model, BaseModel) + else None + ), + litellm_changed_by=litellm_changed_by, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ) + ) + + return updated_model + + except Exception as e: + verbose_proxy_logger.exception(f"Error in model {action}: {str(e)}") + + if isinstance(e, (HTTPException, ProxyException)): + raise e + + raise ProxyException( + message=f"Error updating model blocked status: {str(e)}", + type=ProxyErrorTypes.internal_server_error, + code=status.HTTP_500_INTERNAL_SERVER_ERROR, + param=None, + ) + + +@router.post( + "/model/block", + tags=["model management"], + dependencies=[Depends(user_api_key_auth)], +) +async def block_model( + data: BlockModelRequest, + http_request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + litellm_changed_by: Optional[str] = Header( + None, + description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", + ), +) -> Optional[LiteLLM_ProxyModelTable]: + """ + Block a DB-stored model deployment from serving requests. + + Parameters: + - model_id: str - The model deployment id to block. + """ + return await _set_model_blocked_status( + data=data, + user_api_key_dict=user_api_key_dict, + blocked=True, + action="blocked", + litellm_changed_by=litellm_changed_by, + ) + + +@router.post( + "/model/unblock", + tags=["model management"], + dependencies=[Depends(user_api_key_auth)], +) +async def unblock_model( + data: BlockModelRequest, + http_request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + litellm_changed_by: Optional[str] = Header( + None, + description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", + ), +) -> Optional[LiteLLM_ProxyModelTable]: + """ + Unblock a DB-stored model deployment so it can serve requests again. + + Parameters: + - model_id: str - The model deployment id to unblock. + """ + return await _set_model_blocked_status( + data=data, + user_api_key_dict=user_api_key_dict, + blocked=False, + action="unblocked", + litellm_changed_by=litellm_changed_by, + ) + + ################################# Helper Functions ################################# #################################################################################### #################################################################################### diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index 82d90be3bec..a912a88a993 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -517,6 +517,13 @@ class AnthropicPassthroughLoggingHandler: # Process each individual event for event_str in individual_events: try: + # Skip OpenAI-style [DONE] sentinels some Anthropic-compatible + # providers emit. Match the whole SSE line so a valid chunk whose + # text payload happens to contain "[DONE]" is not dropped. + if any( + line.strip() == "data: [DONE]" for line in event_str.split("\n") + ): + continue transformed_openai_chunk = anthropic_model_response_iterator.convert_str_chunk_to_generic_chunk( chunk=event_str ) @@ -525,6 +532,14 @@ class AnthropicPassthroughLoggingHandler: except (StopIteration, StopAsyncIteration): break + except json.JSONDecodeError: + # Some upstreams emit non-JSON SSE lines; skip them so the + # logging pipeline is not broken by a single bad frame. + verbose_proxy_logger.debug( + "Skipping non-JSON SSE event: %s", + event_str[:200], + ) + continue complete_streaming_response = litellm.stream_chunk_builder( chunks=all_openai_chunks, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index d2b848e3c33..04c9256fdf8 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -174,9 +174,7 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915 data["adapter_id"] = adapter_id - verbose_proxy_logger.debug( - "Request received by LiteLLM:\n{}".format(json.dumps(data, indent=4)), - ) + verbose_proxy_logger.debug("Request received by LiteLLM:\n%s", data) data["model"] = ( general_settings.get("completion_model", None) # server default or user_model # model name passed via cli args @@ -298,7 +296,7 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915 ) ) - verbose_proxy_logger.debug("\nResponse from Litellm:\n{}".format(response)) + verbose_proxy_logger.debug("\nResponse from Litellm:\n%s", response) return response except Exception as e: await proxy_logging_obj.post_call_failure_hook( @@ -811,9 +809,10 @@ async def pass_through_request( # noqa: PLR0915 else: _parsed_body = await _read_request_body(request) verbose_proxy_logger.debug( - "Pass through endpoint sending request to \nURL {}\nheaders: {}\nbody: {}\n".format( - url, headers, _parsed_body - ) + "Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n", + url, + headers, + _parsed_body, ) ### COLLECT GUARDRAILS FOR PASSTHROUGH ENDPOINT ### diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 267e388112d..ea2aa8fb01e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -252,6 +252,7 @@ from litellm.proxy.analytics_endpoints.analytics_endpoints import ( ) from litellm.proxy.auth.auth_checks import ( ExperimentalUIJWTToken, + can_key_call_resolved_model, get_team_object, log_db_metrics, ) @@ -7112,6 +7113,17 @@ async def async_data_generator( # noqa: PLR0915 if not request_data.get("_litellm_skip_openai_stream_done"): done_message = "[DONE]" yield f"data: {done_message}\n\n" + except (asyncio.CancelledError, GeneratorExit): + # Client disconnected mid-stream. CancelledError / GeneratorExit are + # BaseException, so they bypass the success/failure logging callbacks + # that normally release the pre-call max_parallel_requests +1; release + # it here. This is the outermost generator Starlette closes on + # disconnect, so it fires reliably regardless of needs_iterator_wrap + # (a nested iterator hook would only see GeneratorExit on GC). + proxy_logging_obj._release_max_parallel_requests_on_disconnect( + user_api_key_dict + ) + raise except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format( @@ -7839,6 +7851,15 @@ class ProxyStartupEvent: ) await VantageLogger.init_vantage_background_job(scheduler=scheduler) + ######################################################## + # Mavvrik FOCUS Background Job + ######################################################## + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( # noqa: PLC0415 + MavvrikFocusLogger, + ) + + await MavvrikFocusLogger.init_mavvrik_focus_background_job(scheduler=scheduler) + ######################################################## # Prometheus Background Job ######################################################## @@ -8224,6 +8245,7 @@ async def model_list( include_metadata: Optional[bool] = False, fallback_type: Optional[str] = None, scope: Optional[str] = None, + healthy_only: Optional[bool] = False, ): """ Use `/model/info` - to get detailed model information, example - pricing, mode, etc. @@ -8237,6 +8259,15 @@ async def model_list( - scope: Optional scope parameter. Currently only accepts "expand". When scope=expand is passed, proxy admins, team admins, and org admins will receive all proxy models as if they are a proxy admin. + - healthy_only: When true, hide models whose backing deployments are all marked + unhealthy by background health checks. Requires + `background_health_checks: true` in general_settings; without + health state the listing is returned unfiltered (fail open). + Models expanded from wildcard routes (e.g. `openai/*`) are not + filtered, and nothing is hidden when `allowed_fails_policy` is + configured (cooldown remains the sole exclusion mechanism). + Hiding is presentation-only: a hidden model can still be + called directly. """ global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj @@ -8270,6 +8301,19 @@ async def model_list( llm_router.get_fully_blocked_model_names() if llm_router is not None else set() ) + # Opt-in: also hide models whose deployments are all unhealthy per background + # health checks. Empty when health state is unavailable or stale (fail open). + unhealthy_names: Set[str] = set() + if healthy_only and llm_router is not None: + unhealthy_names = await llm_router.async_get_fully_unhealthy_model_names() + if not unhealthy_names: + verbose_proxy_logger.debug( + "healthy_only=true but no unhealthy deployment state is available " + "(requires background_health_checks); returning unfiltered model list" + ) + + hidden_names = blocked_names | unhealthy_names + # If scope=expand and user has admin privileges, return all proxy models if should_expand_scope: # Get all proxy models as if user is a proxy admin @@ -8302,9 +8346,9 @@ async def model_list( only_model_access_groups=only_model_access_groups or False, ) - # Hide paused models from the public listing (admins manage them via /model/info) - if blocked_names: - all_models = [m for m in all_models if m not in blocked_names] + # Hide paused/unhealthy models from the public listing + if hidden_names: + all_models = [m for m in all_models if m not in hidden_names] # Build response data with all proxy models model_data = [] @@ -8339,9 +8383,9 @@ async def model_list( user_api_key_cache=user_api_key_cache, ) - # Hide paused models from the public listing (admins manage them via /model/info) - if blocked_names: - all_models = [m for m in all_models if m not in blocked_names] + # Hide paused/unhealthy models from the public listing + if hidden_names: + all_models = [m for m in all_models if m not in hidden_names] # Build response data model_data = [] @@ -9458,13 +9502,15 @@ async def vertex_ai_live_passthrough_endpoint( @lru_cache(maxsize=_REALTIME_BODY_CACHE_SIZE) def _realtime_query_params_template( - model: str, intent: Optional[str] + model: Optional[str], intent: Optional[str] ) -> Tuple[Tuple[str, str], ...]: """ Build a hashable representation of the realtime query params so we can cache the repetitive model/intent combinations. """ - params: List[Tuple[str, str]] = [("model", model)] + params: List[Tuple[str, str]] = [] + if model is not None: + params.append(("model", model)) if intent is not None: params.append(("intent", intent)) return tuple(params) @@ -9475,8 +9521,10 @@ def _realtime_query_params_template( @app.websocket("/realtime") async def realtime_websocket_endpoint( websocket: WebSocket, - model: str, - intent: str = fastapi.Query( + model: Optional[str] = fastapi.Query( + None, description="The model to use for the websocket connection." + ), + intent: Optional[str] = fastapi.Query( None, description="The intent of the websocket connection." ), guardrails: Optional[str] = fastapi.Query( @@ -9493,6 +9541,25 @@ async def realtime_websocket_endpoint( accept_kwargs: dict = {} if requested_protocols: accept_kwargs["subprotocol"] = requested_protocols[0] + + route_model = model + if route_model is None: + if intent == "transcription": + route_model = "gpt-realtime-whisper" + else: + await websocket.close(code=1008, reason="model query parameter is required") + return + assert route_model is not None + try: + await can_key_call_resolved_model( + model=route_model, + llm_model_list=llm_model_list, + valid_token=user_api_key_dict, + llm_router=llm_router, + ) + except ProxyException as e: + await websocket.close(code=1008, reason=e.message[:120]) + return await websocket.accept(**accept_kwargs) # Only use explicit parameters, not all query params @@ -9501,7 +9568,7 @@ async def realtime_websocket_endpoint( ) data: Dict[str, Any] = { - "model": model, + "model": route_model, "websocket": websocket, "query_params": query_params, # Only explicit params } @@ -9521,7 +9588,7 @@ async def realtime_websocket_endpoint( request._url = websocket.url async def return_body(): - return _realtime_request_body(model) + return _realtime_request_body(route_model) request.body = return_body # type: ignore @@ -9547,7 +9614,7 @@ async def realtime_websocket_endpoint( user_request_timeout=user_request_timeout, user_max_tokens=user_max_tokens, user_api_base=user_api_base, - model=model, + model=route_model, route_type="_arealtime", ) except Exception as e: diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index 14d004d977e..a953dbec6b7 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -10,6 +10,7 @@ from fastapi import status as http_status from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, @@ -19,11 +20,143 @@ from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.types.realtime import ( RealtimeClientSecretRequest, RealtimeClientSecretResponse, + RealtimeTranscriptionSessionRequest, + RealtimeTranscriptionSessionResponse, ) router = APIRouter() _REALTIME_TOKEN_VERSION = "realtime_v1" +_DEFAULT_REALTIME_MODEL = "gpt-4o-realtime-preview" +_DEFAULT_TRANSCRIPTION_MODEL = "gpt-realtime-whisper" +_ALLOWED_SESSION_TYPES = ("realtime", "transcription") + + +def _coerce_realtime_session_type(session_type: Optional[str]) -> str: + if session_type in _ALLOWED_SESSION_TYPES: + return session_type + return "realtime" + + +def _append_model_candidate(candidates: list[str], model: Any) -> None: + if isinstance(model, str) and model and model not in candidates: + candidates.append(model) + + +def _transcription_model_candidates_from_session(session: dict) -> list[str]: + candidates: list[str] = [] + + audio = session.get("audio") + if isinstance(audio, dict): + audio_input = audio.get("input") + if isinstance(audio_input, dict): + nested_transcription = audio_input.get("transcription") + if isinstance(nested_transcription, dict): + _append_model_candidate( + candidates, + nested_transcription.get("model"), + ) + + flat_transcription = session.get("input_audio_transcription") + if isinstance(flat_transcription, dict): + _append_model_candidate(candidates, flat_transcription.get("model")) + + return candidates + + +def _set_transcription_model_on_session( + session: dict, + model: str, + create_if_missing: bool = False, +) -> None: + updated_existing_config = False + + flat_transcription = session.get("input_audio_transcription") + if isinstance(flat_transcription, dict): + session["input_audio_transcription"] = { + **flat_transcription, + "model": model, + } + updated_existing_config = True + + audio = session.get("audio") + if isinstance(audio, dict): + audio_input = audio.get("input") + if isinstance(audio_input, dict): + nested_transcription = audio_input.get("transcription") + if isinstance(nested_transcription, dict): + session["audio"] = { + **audio, + "input": { + **audio_input, + "transcription": { + **nested_transcription, + "model": model, + }, + }, + } + updated_existing_config = True + + if updated_existing_config or not create_if_missing: + return + + audio = audio if isinstance(audio, dict) else {} + audio_input = audio.get("input") + audio_input = audio_input if isinstance(audio_input, dict) else {} + session["audio"] = { + **audio, + "input": { + **audio_input, + "transcription": {"model": model}, + }, + } + + +async def _prepare_client_secret_session( + req: RealtimeClientSecretRequest, + user_api_key_dict: UserAPIKeyAuth, + llm_model_list: Optional[list], + llm_router: Any, +) -> tuple[str, Optional[dict], str]: + session_type = _coerce_realtime_session_type( + req.session.type if req.session else None + ) + session_data: Optional[dict] = ( + req.session.model_dump(exclude_none=True) if req.session else None + ) + if session_data is not None: + session_data["type"] = session_type + + session_model = req.session.model if req.session else None + model: str = session_model or req.model or _DEFAULT_REALTIME_MODEL + if session_type != "transcription": + return model, session_data, session_type + + transcription_model_candidates = _transcription_model_candidates_from_session( + session_data or {} + ) + if not transcription_model_candidates: + _append_model_candidate(transcription_model_candidates, session_model) + _append_model_candidate(transcription_model_candidates, req.model) + if not transcription_model_candidates: + transcription_model_candidates.append(_DEFAULT_TRANSCRIPTION_MODEL) + + model = transcription_model_candidates[0] + for transcription_model in transcription_model_candidates: + await can_key_call_resolved_model( + model=transcription_model, + valid_token=user_api_key_dict, + llm_model_list=llm_model_list, + llm_router=llm_router, + ) + if session_data is not None: + _set_transcription_model_on_session( + session=session_data, + model=model, + create_if_missing=True, + ) + session_data.pop("model", None) + return model, session_data, session_type def _encode_realtime_token_payload( @@ -32,6 +165,7 @@ def _encode_realtime_token_payload( user_id: Optional[str], team_id: Optional[str], expires_at: Optional[int], + session_type: str = "realtime", ) -> str: """ Encode metadata with the upstream ephemeral key so /realtime/calls can @@ -44,6 +178,7 @@ def _encode_realtime_token_payload( "user_id": user_id or "", "team_id": team_id or "", "expires_at": expires_at, + "session_type": session_type, } return json.dumps(payload, separators=(",", ":")) @@ -94,6 +229,7 @@ async def create_realtime_client_secret( add_litellm_data_to_request, general_settings, llm_router, + llm_model_list, proxy_config, proxy_logging_obj, route_request, @@ -106,17 +242,18 @@ async def create_realtime_client_secret( body = await _read_request_body(request=request) req = RealtimeClientSecretRequest(**body) - model: str = ( - (req.session.model if req.session else None) - or req.model - or "gpt-4o-realtime-preview" + model, session_data, session_type = await _prepare_client_secret_session( + req=req, + user_api_key_dict=user_api_key_dict, + llm_model_list=llm_model_list, + llm_router=llm_router, ) data = {"model": model} # If session is provided, use it; otherwise create one from model - if req.session: - data["session"] = req.session.model_dump(exclude_none=True) + if session_data is not None: + data["session"] = session_data elif req.model: # User provided model at root level, convert to session format data["session"] = {"type": "realtime", "model": model} @@ -161,6 +298,8 @@ async def create_realtime_client_secret( "litellm.proxy.realtime_endpoints.webrtc.create_realtime_client_secret(): Exception - %s", str(e), ) + if isinstance(e, ProxyException): + raise e if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "message", str(e)), @@ -199,6 +338,7 @@ async def create_realtime_client_secret( user_id=getattr(user_api_key_dict, "user_id", None), team_id=getattr(user_api_key_dict, "team_id", None), expires_at=expires_at if isinstance(expires_at, int) else None, + session_type=session_type, ) encrypted_token: str = encrypt_value_helper(token_payload) upstream_json["value"] = encrypted_token @@ -279,16 +419,20 @@ async def proxy_realtime_calls( model = ( decoded_payload.get("model_id") or request.query_params.get("model") - or "gpt-4o-realtime-preview" + or _DEFAULT_REALTIME_MODEL ) user_id = decoded_payload.get("user_id") or None team_id = decoded_payload.get("team_id") or None + session_type = _coerce_realtime_session_type( + decoded_payload.get("session_type") + ) else: # Backward compatibility: older tokens contained only encrypted upstream key. openai_ephemeral_key = decrypted_token_value - model = request.query_params.get("model", "gpt-4o-realtime-preview") + model = request.query_params.get("model", _DEFAULT_REALTIME_MODEL) user_id = None team_id = None + session_type = "realtime" # Build a minimal UserAPIKeyAuth with user/team IDs from the token # so spend tracking and budget enforcement work correctly. @@ -299,11 +443,17 @@ async def proxy_realtime_calls( data: dict = {} try: - # Build session config for the multipart form data session_config = { - "type": "realtime", - "model": model, + "type": session_type, } + if session_type == "transcription": + _set_transcription_model_on_session( + session=session_config, + model=model, + create_if_missing=True, + ) + else: + session_config["model"] = model data = { "model": model, @@ -366,3 +516,145 @@ async def proxy_realtime_calls( status_code=upstream_resp.status_code, media_type=upstream_resp.headers.get("content-type", "application/sdp"), ) + + +@router.post( + "/v1/realtime/transcription_sessions", + dependencies=[Depends(user_api_key_auth)], + tags=["realtime"], +) +@router.post( + "/realtime/transcription_sessions", + dependencies=[Depends(user_api_key_auth)], + tags=["realtime"], +) +@router.post( + "/openai/v1/realtime/transcription_sessions", + dependencies=[Depends(user_api_key_auth)], + tags=["realtime"], +) +async def create_realtime_transcription_session( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> RealtimeTranscriptionSessionResponse: + """ + Create an ephemeral Realtime transcription session + (POST /v1/realtime/transcription_sessions) for the WebRTC/WebSocket flow. + + Mirrors the client_secrets route but targets the transcription_sessions + endpoint and encrypts the ephemeral key returned under `client_secret.value`. + """ + from litellm.proxy.proxy_server import ( + add_litellm_data_to_request, + general_settings, + llm_router, + llm_model_list, + proxy_config, + proxy_logging_obj, + route_request, + user_model, + version, + ) + + data: dict = {} + try: + body = await _read_request_body(request=request) + req = RealtimeTranscriptionSessionRequest(**body) + + model: str = req.resolved_model() or "gpt-realtime-whisper" + await can_key_call_resolved_model( + model=model, + valid_token=user_api_key_dict, + llm_model_list=llm_model_list, + llm_router=llm_router, + ) + + transcription_session = {k: v for k, v in body.items() if k != "model"} + data = {"model": model, "transcription_session": transcription_session} + + data = await add_litellm_data_to_request( + data=data, + request=request, + general_settings=general_settings, + user_api_key_dict=user_api_key_dict, + version=version, + proxy_config=proxy_config, + ) + + data = await proxy_logging_obj.pre_call_hook( + user_api_key_dict=user_api_key_dict, + data=data, + call_type="acreate_realtime_transcription_session", + ) + + verbose_proxy_logger.debug( + "Realtime: /v1/realtime/transcription_sessions (model=%s)", model + ) + + llm_call = await route_request( + data=data, + route_type="acreate_realtime_transcription_session", + llm_router=llm_router, + user_model=user_model, + ) + upstream_resp: httpx.Response = await llm_call # type: ignore + + except Exception as e: + await proxy_logging_obj.post_call_failure_hook( + user_api_key_dict=user_api_key_dict, + original_exception=e, + request_data=data, + ) + verbose_proxy_logger.error( + "litellm.proxy.realtime_endpoints.create_realtime_transcription_session(): Exception - %s", + str(e), + ) + if isinstance(e, ProxyException): + raise e + if isinstance(e, HTTPException): + raise ProxyException( + message=getattr(e, "detail", getattr(e, "message", str(e))), + type=getattr(e, "type", "None"), + param=getattr(e, "param", "None"), + code=getattr(e, "status_code", http_status.HTTP_400_BAD_REQUEST), + ) + raise ProxyException( + message=getattr(e, "message", str(e)), + type=getattr(e, "type", "None"), + param=getattr(e, "param", "None"), + code=getattr(e, "status_code", 500), + ) + + if upstream_resp.status_code != 200: + verbose_proxy_logger.error( + "Realtime transcription_sessions upstream error %s: %s", + upstream_resp.status_code, + upstream_resp.text, + ) + return Response( # type: ignore[return-value] + content=upstream_resp.content, + status_code=upstream_resp.status_code, + media_type="application/json", + ) + + upstream_json: dict = upstream_resp.json() + + # Encrypt the ephemeral key (returned under client_secret.value) with routing + # metadata so the follow-up /realtime/calls request can recover the model. + client_secret = upstream_json.get("client_secret") + if isinstance(client_secret, dict) and "value" in client_secret: + raw_value: str = client_secret.get("value", "") + expires_at = client_secret.get("expires_at") + token_payload = _encode_realtime_token_payload( + ephemeral_key=raw_value, + model_id=model, + user_id=getattr(user_api_key_dict, "user_id", None), + team_id=getattr(user_api_key_dict, "team_id", None), + expires_at=expires_at if isinstance(expires_at, int) else None, + session_type="transcription", + ) + client_secret["value"] = encrypt_value_helper(token_payload) + upstream_json["client_secret"] = client_secret + + return RealtimeTranscriptionSessionResponse(**upstream_json) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 8f6f7084a0c..3626a21516d 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -1,6 +1,7 @@ import asyncio from typing import TYPE_CHECKING, Any, Literal, Optional +import httpx from fastapi import HTTPException, status import litellm @@ -46,6 +47,30 @@ def _is_a2a_agent_model(model_name: Any) -> bool: return isinstance(model_name, str) and model_name.startswith("a2a/") +def _raise_if_model_fully_blocked( + llm_router: LitellmRouter, model_name: Any, team_id: Optional[str] +) -> None: + if not isinstance(model_name, str) or not model_name: + return + if not isinstance(llm_router, litellm.Router): + return + deployments = ( + llm_router.get_model_list(model_name=model_name, team_id=team_id) or [] + ) + if llm_router._are_all_deployments_blocked(deployments): + raise litellm.PermissionDeniedError( + message="Model is blocked", + model=model_name, + llm_provider="", + response=httpx.Response( + status_code=403, + request=httpx.Request( + method="POST", url="https://github.com/BerriAI/litellm" + ), + ), + ) + + ROUTE_ENDPOINT_MAPPING = { "acompletion": "/chat/completions", "atext_completion": "/completions", @@ -74,6 +99,7 @@ ROUTE_ENDPOINT_MAPPING = { "avideo_extension": "/videos/extensions", "acreate_realtime_client_secret": "/realtime/client_secrets", "arealtime_calls": "/realtime/calls", + "acreate_realtime_transcription_session": "/realtime/transcription_sessions", "acreate_container": "/containers", "alist_containers": "/containers", "aretrieve_container": "/containers/{container_id}", @@ -261,6 +287,7 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin "_arealtime", # private function for realtime API "acreate_realtime_client_secret", "arealtime_calls", + "acreate_realtime_transcription_session", "_aresponses_websocket", # private function for responses WebSocket mode "aimage_edit", "agenerate_content", @@ -411,6 +438,9 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin else: return getattr(litellm, f"{route_type}")(**data) elif llm_router is not None: + _raise_if_model_fully_blocked( + llm_router=llm_router, model_name=data.get("model"), team_id=team_id + ) # Evals API: always route to litellm directly (not through router) # But extract model credentials if a model is provided if route_type in [ @@ -427,6 +457,7 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin "adelete_run", "acreate_realtime_client_secret", "arealtime_calls", + "acreate_realtime_transcription_session", ]: # If a model is provided, get its credentials from the router model = data.get("model") diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index aa85be6671a..ef06adb27fc 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -2142,6 +2142,20 @@ async def ui_view_spend_logs( # noqa: PLR0915 data = await prisma_client.db.query_raw(sql_query, *sql_params) + # query_raw returns the JSONB `metadata` column as a string (the Prisma + # serialiser bypasses the model-layer JSON hydration we get on the ORM + # path). The UI reads `metadata.status` / `metadata.error_information` + # as object fields, so failure rows looked like successes (#29674). + # Re-hydrate to dict here. + for row in data: + if isinstance(row, dict): + md = row.get("metadata") + if isinstance(md, str): + try: + row["metadata"] = json.loads(md) + except (ValueError, TypeError): + row["metadata"] = {} + # Calculate total pages total_pages = (total_records + page_size - 1) // page_size diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index ebd5b5d90cd..4aa555164b0 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -137,6 +137,9 @@ from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter from litellm.proxy.hooks.parallel_request_limiter import ( _PROXY_MaxParallelRequestsHandler, ) +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, +) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor from litellm.repositories.budget_repository import BudgetRepository @@ -2688,6 +2691,42 @@ class ProxyLogging: logging_obj._deferred_stream_complete_args = None asyncio.create_task(_deferred_cb(*_args)) + def _release_max_parallel_requests_on_disconnect( + self, user_api_key_dict: UserAPIKeyAuth + ) -> None: + """ + Release the api-key max_parallel_requests slot when a streaming + response is cancelled mid-flight (client disconnect). Neither the + success nor failure logging callback fires on the resulting + CancelledError / GeneratorExit, so the pre-call +1 would otherwise + leak. + + Must be called from the outermost streaming generator (the one + Starlette drives and closes on disconnect). A nested iterator-hook + generator only receives GeneratorExit when it is garbage collected, + which is non-deterministic, so the refund cannot live there. + + Scheduled fire-and-forget (no await) because awaiting is not + permitted while unwinding a GeneratorExit. + """ + limiter = self.get_proxy_hook("parallel_request_limiter") + if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + return + try: + asyncio.create_task( + limiter.async_release_max_parallel_requests_on_disconnect( + user_api_key_dict + ) + ) + except RuntimeError: + # No running event loop (e.g. interpreter/loop shutdown); the + # counter's window TTL will reclaim the slot. + verbose_proxy_logger.warning( + "parallel_request_limiter_v3: could not schedule " + "max_parallel_requests release on disconnect; no running " + "event loop. Slot will be reclaimed when its window TTL expires" + ) + def _init_response_taking_too_long_task(self, data: Optional[dict] = None): """ Initialize the response taking too long task if user is using slack alerting diff --git a/litellm/realtime_api/README.md b/litellm/realtime_api/README.md index 6b467c056a6..d810de2f24f 100644 --- a/litellm/realtime_api/README.md +++ b/litellm/realtime_api/README.md @@ -1 +1,9 @@ -Abstraction / Routing logic for OpenAI's `/v1/realtime` endpoint. \ No newline at end of file +Abstraction / Routing logic for OpenAI's `/v1/realtime` endpoints. + +Supported endpoints: +- WebSocket: `/v1/realtime` (with `intent=transcription` for transcription-only sessions) +- HTTP: `/v1/realtime/client_secrets`, `/v1/realtime/transcription_sessions` + +Supported providers: OpenAI, Azure OpenAI, Bedrock, Vertex AI, xAI. + +For user-facing documentation and usage examples, see the litellm-docs repo. \ No newline at end of file diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 95d6f7c3e03..7031ecaa1a0 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -15,6 +15,7 @@ from litellm.types.realtime import ( RealtimeExpiresAfter, RealtimeQueryParams, RealtimeSessionConfig, + RealtimeTranscriptionSessionRequest, ) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import LlmProviders @@ -159,6 +160,78 @@ async def acreate_realtime_client_secret( ) +@wrapper_client +async def acreate_realtime_transcription_session( + model: Optional[str] = None, + transcription_session: Optional[Dict[str, Any]] = None, + timeout: Optional[float] = None, + **kwargs, +): + """ + Create an ephemeral transcription session via POST + /v1/realtime/transcription_sessions. + + ``transcription_session`` is the upstream request body (input_audio_format, + input_audio_transcription, turn_detection, …). ``model`` is a LiteLLM-only + routing hint; the provider model lives in + ``transcription_session.input_audio_transcription.model``. + """ + req = RealtimeTranscriptionSessionRequest( + model=model, + **(transcription_session or {}), + ) + model_name = req.resolved_model() or "gpt-realtime-whisper" + litellm_logging_obj: LiteLLMLogging = kwargs.get("litellm_logging_obj") # type: ignore + litellm_params = GenericLiteLLMParams(**kwargs) + + ( + model_name, + custom_llm_provider, + dynamic_api_key, + dynamic_api_base, + ) = get_llm_provider( + model=model_name, + api_base=litellm_params.api_base, + api_key=litellm_params.api_key, + ) + ( + provider_config, + resolved_api_base, + resolved_api_key, + ) = _get_realtime_http_provider_config( + custom_llm_provider=custom_llm_provider, + dynamic_api_base=dynamic_api_base, + dynamic_api_key=dynamic_api_key, + litellm_params=litellm_params, + ) + litellm_logging_obj.update_from_kwargs( + kwargs=kwargs, + model=model_name, + optional_params={"transcription_session": transcription_session}, + litellm_params={"api_base": resolved_api_base}, + custom_llm_provider=custom_llm_provider, + ) + request_data = req.model_dump(exclude_none=True, exclude={"model"}) + # Ensure the upstream body's input_audio_transcription.model matches the + # authorized routing model. This prevents a caller from supplying an allowed + # top-level model for auth while sneaking a different model into the nested + # transcription config that gets forwarded to the provider. + if isinstance(request_data.get("input_audio_transcription"), dict): + request_data["input_audio_transcription"]["model"] = model_name + return await base_llm_http_handler.async_realtime_transcription_session_handler( + api_base=resolved_api_base, + api_key=resolved_api_key, + request_data=request_data, + logging_obj=litellm_logging_obj, + timeout=timeout or request_timeout, + provider_config=provider_config, + model=model_name, + extra_headers=kwargs.get("extra_headers"), + client=kwargs.get("client"), + api_version=litellm_params.api_version, + ) + + @wrapper_client async def arealtime_calls( openai_ephemeral_key: str, @@ -246,9 +319,13 @@ async def _arealtime( # noqa: PLR0915 api_key=api_key, ) - # Ensure query params use the normalized provider model (no proxy aliases). + # If the client supplied `model` in the URL, ensure it uses the normalized + # provider model (no proxy aliases). If they omitted it, preserve that shape + # for transcription-only sessions like OpenAI's `?intent=transcription`. if query_params is not None: - query_params = {**query_params, "model": model} + query_params = {**query_params} + if "model" in query_params: + query_params["model"] = model litellm_logging_obj.update_from_kwargs( kwargs=kwargs, @@ -278,6 +355,7 @@ async def _arealtime( # noqa: PLR0915 headers=headers, user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata(kwargs), + query_params=query_params, ) elif _custom_llm_provider == "azure": api_base = ( @@ -300,8 +378,13 @@ async def _arealtime( # noqa: PLR0915 kwargs.get("realtime_protocol") or litellm_params.get("realtime_protocol") or os.environ.get("LITELLM_AZURE_REALTIME_PROTOCOL") - or "beta" ) + if ( + realtime_protocol is None + and (query_params or {}).get("intent") == "transcription" + ): + realtime_protocol = "GA" + realtime_protocol = realtime_protocol or "beta" await azure_realtime.async_realtime( model=model, websocket=websocket, @@ -313,6 +396,7 @@ async def _arealtime( # noqa: PLR0915 timeout=timeout, logging_obj=litellm_logging_obj, realtime_protocol=realtime_protocol, + query_params=query_params, user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata(kwargs), ) @@ -450,6 +534,7 @@ async def _arealtime( # noqa: PLR0915 headers=headers, user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata(kwargs), + query_params=query_params, ) else: raise ValueError(f"Unsupported model: {model}") diff --git a/litellm/router.py b/litellm/router.py index d1c8e227bea..9d614dd21e2 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3381,7 +3381,10 @@ class Router: # Request Number X, Model Number Y _tasks.append( _async_completion_no_exceptions_return_idx( - model=model, idx=idx, messages=message, **kwargs # type: ignore + model=model, + idx=idx, + messages=message, # type: ignore[arg-type] + **kwargs, ) ) responses = await asyncio.gather(*_tasks) @@ -3544,7 +3547,7 @@ class Router: self, model: str, messages: List[AllMessageValues], priority: int, stream: Literal[False] = False, **kwargs ) -> ModelResponse: ... - + @overload async def schedule_acompletion( self, model: str, messages: List[AllMessageValues], priority: int, stream: Literal[True], **kwargs @@ -7937,8 +7940,7 @@ class Router: re.compile(pattern) except re.error as exc: raise ValueError( - f"Invalid regex in tag_regex for model '{deployment.model_name}': " - f"{pattern!r} — {exc}" + f"Invalid regex in tag_regex for model '{deployment.model_name}': {pattern!r} — {exc}" ) from exc deployment = self._add_deployment(deployment=deployment) @@ -8196,8 +8198,7 @@ class Router: if deployment.model_name in self.adaptive_routers: raise ValueError( - f"Adaptive-router deployment {deployment.model_name} already exists. " - "Please use a different model name." + f"Adaptive-router deployment {deployment.model_name} already exists. Please use a different model name." ) adaptive_router = AdaptiveRouter( @@ -9407,8 +9408,7 @@ class Router: ): model_group_info.supports_parallel_function_calling = True if ( - model_info.get("supports_vision", None) is not None - and model_info["supports_vision"] is True # type: ignore + model_info.get("supports_vision", None) is not None and model_info["supports_vision"] is True # type: ignore ): model_group_info.supports_vision = True if ( @@ -9428,8 +9428,7 @@ class Router: model_group_info.supports_url_context = True if ( - model_info.get("supports_reasoning", None) is not None - and model_info["supports_reasoning"] is True # type: ignore + model_info.get("supports_reasoning", None) is not None and model_info["supports_reasoning"] is True # type: ignore ): model_group_info.supports_reasoning = True if ( @@ -10052,6 +10051,76 @@ class Router: name for name, fully_blocked in blocked_by_name.items() if fully_blocked } + @staticmethod + def _are_all_deployments_blocked( + deployments: List[DeploymentTypedDict], + ) -> bool: + return len(deployments) > 0 and all( + (deployment.get("model_info") or {}).get("blocked") is True + for deployment in deployments + ) + + def _is_model_fully_blocked(self, model: str) -> bool: + deployments = self.get_model_list(model_name=model) or [] + return self._are_all_deployments_blocked(deployments=deployments) + + async def async_get_fully_unhealthy_model_names(self) -> Set[str]: + """ + Returns the set of model names where every backing deployment is currently + marked unhealthy by background health checks (and the health state is not stale). + + Used by `/v1/models?healthy_only=true` to hide models that cannot serve any + request. A model with at least one healthy (or unknown-health) deployment + remains visible. Returns an empty set when no health state is available, so + callers fail open to the unfiltered listing. + + Notes: + - Mirrors `_async_filter_health_check_unhealthy_deployments`: when + `allowed_fails_policy` is set, cooldown is the sole routing exclusion + mechanism, so nothing is hidden here either. + - Team-specific public model names (`team_public_model_name`) are + aggregated alongside `model_name`, so team aliases of fully-unhealthy + deployments are hidden too (unlike `get_fully_blocked_model_names`, + which matches `model_name` only). + - Wildcard routes (e.g. `openai/*`) are matched by their literal + deployment name only; models expanded from a wildcard route are not + hidden (fail open). + - Intentionally diverges from the routing-time safety net (which + bypasses the health filter when every candidate is unhealthy and + still attempts the request): hiding here is presentation-only — + it answers "should this model be advertised?", not "should a + request for it still be attempted?". A hidden model can still be + called directly. + """ + if self.allowed_fails_policy is not None: + return set() + unhealthy_ids = ( + await self.health_state_cache.async_get_unhealthy_deployment_ids() + ) + if not unhealthy_ids: + return set() + deployments = self.get_model_list() or [] + unhealthy_by_name: Dict[str, bool] = {} + for deployment in deployments: + model_info = deployment.get("model_info") or {} + names = [deployment.get("model_name") or ""] + team_public_model_name = model_info.get("team_public_model_name") + if team_public_model_name: + names.append(team_public_model_name) + is_unhealthy = model_info.get("id") in unhealthy_ids + for name in names: + if not name: + continue + if name in unhealthy_by_name: + unhealthy_by_name[name] = unhealthy_by_name[name] and is_unhealthy + else: + unhealthy_by_name[name] = is_unhealthy + return { + name + for name, fully_unhealthy in unhealthy_by_name.items() + if fully_unhealthy + } + def _get_team_specific_model( self, deployment: DeploymentTypedDict, team_id: Optional[str] = None ) -> Optional[str]: diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 0d81e25592d..5aeb9e366d1 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -23,6 +23,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( ContentFilterCategoryConfig, ) +from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( + OvalixGuardrailConfigModel, +) from litellm.types.proxy.guardrails.guardrail_hooks.promptguard import ( PromptGuardConfigModel, ) @@ -97,6 +100,7 @@ class SupportedGuardrailIntegrations(Enum): GENERIC_GUARDRAIL_API = "generic_guardrail_api" QUALIFIRE = "qualifire" CUSTOM_CODE = "custom_code" + OVALIX = "ovalix" MICROSOFT_PURVIEW = "microsoft_purview" SEMANTIC_GUARD = "semantic_guard" MCP_END_USER_PERMISSION = "mcp_end_user_permission" @@ -852,6 +856,7 @@ class LitellmParams( BaseLitellmParams, EnkryptAIGuardrailConfigs, IBMGuardrailsBaseConfigModel, + OvalixGuardrailConfigModel, QualifireGuardrailConfigModel, BlockCodeExecutionGuardrailConfigModel, HiddenlayerGuardrailConfigModel, diff --git a/litellm/types/integrations/newrelic.py b/litellm/types/integrations/newrelic.py new file mode 100644 index 00000000000..2de9769b181 --- /dev/null +++ b/litellm/types/integrations/newrelic.py @@ -0,0 +1,9 @@ +from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams + + +class NewRelicInitParams(StandardCustomLoggerInitParams): + """ + Params for initializing a New Relic logger on litellm + """ + + pass diff --git a/litellm/types/llms/oci.py b/litellm/types/llms/oci.py index df551d8a8c6..621f40aa31d 100644 --- a/litellm/types/llms/oci.py +++ b/litellm/types/llms/oci.py @@ -291,25 +291,6 @@ class CohereToolResult(BaseModel): outputs: List[Dict[str, Any]] -class CohereResponseFormat(BaseModel): - """Response format for Cohere.""" - - type: str - - -class CohereResponseTextFormat(CohereResponseFormat): - """Text response format for Cohere.""" - - type: Literal["text"] = "text" - - -class CohereResponseJSONSchemaFormat(CohereResponseFormat): - """JSON schema response format for Cohere.""" - - type: Literal["json_schema"] = "json_schema" - jsonSchema: Dict[str, Any] - - class CohereChatRequest(BaseModel): """Cohere chat request model.""" @@ -336,13 +317,10 @@ class CohereChatRequest(BaseModel): # ``OCIChatConfig.openai_to_oci_cohere_param_map`` which marks # ``tool_choice`` as unsupported. The field is intentionally absent here # so it isn't silently dropped or surfaced as a supported feature. - responseFormat: Optional[ - Union[ - CohereResponseTextFormat, - CohereResponseJSONSchemaFormat, - CohereResponseFormat, - ] - ] = None + # OCI Cohere responseFormat is {"type": "TEXT" | "JSON_OBJECT", "schema"?: ...}; + # there is no JSON_SCHEMA type. The shape is built in + # OCIChatConfig._normalize_response_format. + responseFormat: Optional[Dict[str, Any]] = None preambleOverride: Optional[str] = None documents: Optional[List[Dict[str, Any]]] = None searchQueriesOnly: Optional[bool] = None diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 0c854d89bb1..cbb316eec75 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -802,6 +802,8 @@ ValidUserMessageContentTypes = [ "audio_url", "document", "guarded_text", + "grounding_source", + "query", "video_url", "file", ] # used for validating user messages. Prevent users from accidentally sending anthropic messages. @@ -813,6 +815,8 @@ ValidUserMessageContentTypesLiteral = Literal[ "audio_url", "document", "guarded_text", + "grounding_source", + "query", "video_url", "file", ] @@ -824,6 +828,8 @@ ValidUserMessageContentTypes = [ "audio_url", "document", "guarded_text", + "grounding_source", + "query", "video_url", "file", ] # used for validating user messages. Prevent users from accidentally sending anthropic messages. @@ -851,6 +857,8 @@ ValidChatCompletionMessageContentTypesLiteral = Literal[ "audio_url", "document", "guarded_text", + "grounding_source", + "query", "video_url", "file", "thinking", @@ -864,6 +872,8 @@ ValidChatCompletionMessageContentTypes = [ "audio_url", "document", "guarded_text", + "grounding_source", + "query", "video_url", "file", "thinking", diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 5ff39930cb9..74d4616cddd 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -2,9 +2,14 @@ from typing import Any, Dict, List, Literal, Optional, Union from typing_extensions import TypedDict +# Bedrock contextual grounding tags each content block so the guardrail knows +# which text is the reference source, the user question, and the content to grade. +BedrockGuardrailQualifier = Literal["grounding_source", "query", "guard_content"] + class BedrockTextContent(TypedDict, total=False): text: str + qualifiers: List[BedrockGuardrailQualifier] class BedrockContentItem(TypedDict, total=False): diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/ovalix.py b/litellm/types/proxy/guardrails/guardrail_hooks/ovalix.py new file mode 100644 index 00000000000..7417d1a00c9 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/ovalix.py @@ -0,0 +1,37 @@ +"""Pydantic config model for the Ovalix guardrail (Tracker API, application and checkpoint IDs).""" + +from typing import Optional + +from pydantic import Field + +from .base import GuardrailConfigModel + + +class OvalixGuardrailConfigModel(GuardrailConfigModel): + """Configuration parameters for the Ovalix guardrail (pre/post call checkpoints).""" + + tracker_api_base: Optional[str] = Field( + default=None, + description="Base URL for the Ovalix Tracker service.", + ) + tracker_api_key: Optional[str] = Field( + default=None, + description="API key for the Ovalix Tracker service.", + ) + application_id: Optional[str] = Field( + default=None, + description="Application ID for the Ovalix Tracker service.", + ) + pre_checkpoint_id: Optional[str] = Field( + default=None, + description="Pre-checkpoint ID for the Ovalix Tracker service.", + ) + post_checkpoint_id: Optional[str] = Field( + default=None, + description="Post-checkpoint ID for the Ovalix Tracker service.", + ) + + @staticmethod + def ui_friendly_name() -> str: + """Display name for this guardrail in the proxy UI.""" + return "Ovalix Guardrail" diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py index 62e4044061b..0db8232a54d 100644 --- a/litellm/types/realtime.py +++ b/litellm/types/realtime.py @@ -115,3 +115,40 @@ class RealtimeClientSecretResponse(BaseModel): expires_at: Optional[int] = None value: str session: Optional[Dict[str, Any]] = None + + +class RealtimeTranscriptionSessionRequest(BaseModel): + """ + Request body for POST /v1/realtime/transcription_sessions. + + Mirrors OpenAI's RealtimeTranscriptionSessionCreateRequest. The model used + for routing is taken from the LiteLLM-only top-level `model` hint, falling + back to `input_audio_transcription.model`. All other fields pass through + unchanged to the provider. + """ + + model_config = {"extra": "allow"} + + # LiteLLM-only routing hint — stripped before forwarding upstream. + model: Optional[str] = None + input_audio_transcription: Optional[Dict[str, Any]] = None + + def resolved_model(self) -> Optional[str]: + if self.model: + return self.model + if self.input_audio_transcription: + return self.input_audio_transcription.get("model") + return None + + +class RealtimeTranscriptionSessionResponse(BaseModel): + """ + Response from POST /v1/realtime/transcription_sessions. + + `client_secret.value` contains the encrypted token instead of the raw + ephemeral key. Unknown fields pass through unchanged. + """ + + model_config = {"extra": "allow"} + + client_secret: Optional[Dict[str, Any]] = None diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 21eb0c9a173..a11ea965f34 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -533,6 +533,7 @@ CallTypesLiteral = Literal[ "acreate_skill", "acreate_realtime_client_secret", "arealtime_calls", + "acreate_realtime_transcription_session", ] # Mapping of API routes to their corresponding call types @@ -2493,6 +2494,7 @@ class LoggedLiteLLMParams(TypedDict, total=False): litellm_call_id: Optional[str] model_alias_map: Optional[dict] metadata: Optional[dict] + litellm_metadata: Optional[dict] model_info: Optional[dict] proxy_server_request: Optional[dict] acompletion: Optional[bool] diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index f0b2432ddc8..8cdde5ac82a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -4409,6 +4409,23 @@ "/v1/audio/transcriptions" ] }, + "azure/gpt-realtime-whisper": { + "input_cost_per_second": 0.0002833333333333333, + "litellm_provider": "azure", + "mode": "audio_transcription", + "source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/gpt-realtime-whisper", + "supported_endpoints": [ + "/v1/realtime", + "/v1/realtime/transcription_sessions" + ], + "supported_modalities": [ + "audio" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true + }, "azure/gpt-5.1-2025-11-13": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_priority": 2.5e-07, @@ -7557,6 +7574,45 @@ "supports_function_calling": true, "supports_tool_choice": true }, + "azure_ai/deepseek-v3.1": { + "input_cost_per_token": 1.23e-06, + "litellm_provider": "azure_ai", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.94e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, + "azure_ai/deepseek-v4-pro": { + "input_cost_per_token": 1.74e-06, + "litellm_provider": "azure_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 3.48e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, + "azure_ai/deepseek-v4-flash": { + "input_cost_per_token": 1.9e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 5.1e-07, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "azure_ai/embed-v-4-0": { "input_cost_per_token": 1.2e-07, "litellm_provider": "azure_ai", @@ -40956,6 +41012,23 @@ "supports_system_messages": true, "supports_tool_choice": true }, + "gpt-realtime-whisper": { + "input_cost_per_second": 0.0002833333333333333, + "litellm_provider": "openai", + "mode": "audio_transcription", + "source": "https://platform.openai.com/docs/models/gpt-realtime-whisper", + "supported_endpoints": [ + "/v1/realtime", + "/v1/realtime/transcription_sessions" + ], + "supported_modalities": [ + "audio" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true + }, "sora-2": { "litellm_provider": "openai", "mode": "video_generation", diff --git a/scripts/benchmark_streaming_chunk_overhead.py b/scripts/benchmark_streaming_chunk_overhead.py index 948be096bec..11fbea6a6a3 100644 --- a/scripts/benchmark_streaming_chunk_overhead.py +++ b/scripts/benchmark_streaming_chunk_overhead.py @@ -170,48 +170,59 @@ def _make_wrapper( ) -def drive_sync(provider_key: str, chunks_per_stream: int, n_streams: int) -> float: +@dataclass +class TimingSample: + wall_s: float + cpu_s: float + + +def drive_sync( + provider_key: str, chunks_per_stream: int, n_streams: int +) -> TimingSample: provider, factory = PROVIDERS[provider_key] # Pre-build the chunk lists; we only measure wrapper iteration cost. chunk_lists = [factory(chunks_per_stream) for _ in range(n_streams)] gc.collect() gc.disable() try: - start = time.perf_counter() + wall_start = time.perf_counter() + cpu_start = time.process_time() for chunks in chunk_lists: wrapper = _make_wrapper(chunks, provider, async_stream=False) for _ in wrapper: pass - elapsed = time.perf_counter() - start + wall_elapsed = time.perf_counter() - wall_start + cpu_elapsed = time.process_time() - cpu_start finally: gc.enable() - return elapsed + return TimingSample(wall_s=wall_elapsed, cpu_s=cpu_elapsed) async def drive_async( provider_key: str, chunks_per_stream: int, n_streams: int -) -> float: +) -> TimingSample: provider, factory = PROVIDERS[provider_key] chunk_lists = [factory(chunks_per_stream) for _ in range(n_streams)] gc.collect() gc.disable() try: - start = time.perf_counter() + wall_start = time.perf_counter() + cpu_start = time.process_time() for chunks in chunk_lists: wrapper = _make_wrapper(chunks, provider, async_stream=True) async for _ in wrapper: pass - elapsed = time.perf_counter() - start + wall_elapsed = time.perf_counter() - wall_start + cpu_elapsed = time.process_time() - cpu_start finally: gc.enable() - return elapsed + return TimingSample(wall_s=wall_elapsed, cpu_s=cpu_elapsed) # --------------------------------------------------------------------------- # Repeat × take-min runner # --------------------------------------------------------------------------- - @dataclass class Result: label: str @@ -222,7 +233,11 @@ class Result: total_chunks: int elapsed_min_s: float elapsed_median_s: float + cpu_at_min_wall_s: float + cpu_median_s: float per_chunk_us: float + cpu_per_chunk_us: float + cpu_to_wall_ratio: float chunks_per_sec: float streams_per_sec: float @@ -260,11 +275,16 @@ def run_case( else: raise ValueError(f"unknown mode {mode!r}") - elapsed_min = min(samples) - elapsed_median = statistics.median(samples) + best_sample = min(samples, key=lambda s: s.wall_s) + elapsed_min = best_sample.wall_s + elapsed_median = statistics.median(s.wall_s for s in samples) + cpu_at_min_wall = best_sample.cpu_s + cpu_median = statistics.median(s.cpu_s for s in samples) # Each stream emits chunks_per_stream text chunks + 1 finish/usage chunk. total_chunks = n_streams * (chunks_per_stream + 1) per_chunk_us = (elapsed_min * 1_000_000) / total_chunks + cpu_per_chunk_us = (cpu_at_min_wall * 1_000_000) / total_chunks + cpu_to_wall_ratio = cpu_at_min_wall / elapsed_min if elapsed_min > 0 else 0.0 chunks_per_sec = total_chunks / elapsed_min if elapsed_min > 0 else 0.0 streams_per_sec = n_streams / elapsed_min if elapsed_min > 0 else 0.0 @@ -277,7 +297,11 @@ def run_case( total_chunks=total_chunks, elapsed_min_s=elapsed_min, elapsed_median_s=elapsed_median, + cpu_at_min_wall_s=cpu_at_min_wall, + cpu_median_s=cpu_median, per_chunk_us=per_chunk_us, + cpu_per_chunk_us=cpu_per_chunk_us, + cpu_to_wall_ratio=cpu_to_wall_ratio, chunks_per_sec=chunks_per_sec, streams_per_sec=streams_per_sec, ) @@ -289,6 +313,8 @@ def format_result(r: Result) -> str: f"min={r.elapsed_min_s*1000:8.2f} ms " f"median={r.elapsed_median_s*1000:8.2f} ms " f"per-chunk={r.per_chunk_us:7.2f} μs " + f"cpu/chunk={r.cpu_per_chunk_us:7.2f} μs " + f"cpu/wall={r.cpu_to_wall_ratio:5.2f}x " f"chunks/s={r.chunks_per_sec:>10,.0f} " f"streams/s={r.streams_per_sec:>8,.1f}" ) diff --git a/tests/documentation_tests/test_env_keys.py b/tests/documentation_tests/test_env_keys.py index b1324d5dee0..681cd536259 100644 --- a/tests/documentation_tests/test_env_keys.py +++ b/tests/documentation_tests/test_env_keys.py @@ -18,6 +18,12 @@ env_keys = set() # Terminal/environment detection variables that should not be documented # These are internal variables used for terminal detection, not user-configurable settings +# Guard-only env vars: read solely to raise on invalid values; the only valid +# value is the default, so there is nothing meaningful to document. +EXCLUDED_GUARD_ONLY_VARS = { + "MAVVRIK_FOCUS_FREQUENCY", +} + EXCLUDED_TERMINAL_VARS = { "TERM", "TERM_PROGRAM", @@ -64,6 +70,7 @@ for root, dirs, files in os.walk(repo_base): match for match in getenv_matches if match not in EXCLUDED_TERMINAL_VARS + and match not in EXCLUDED_GUARD_ONLY_VARS ) # Extract only the key part, excluding terminal vars # Find all keys using litellm.get_secret() diff --git a/tests/guardrails_tests/test_bedrock_guardrails.py b/tests/guardrails_tests/test_bedrock_guardrails.py index 6e78a8c4284..a23e89e576c 100644 --- a/tests/guardrails_tests/test_bedrock_guardrails.py +++ b/tests/guardrails_tests/test_bedrock_guardrails.py @@ -1160,7 +1160,10 @@ async def test_convert_to_bedrock_format_post_call_streaming_hook(): output_call = bedrock_calls[0] assert output_call["source"] == "OUTPUT" assert output_call["response"] is not None - assert output_call["messages"] is None # OUTPUT calls don't need messages + # OUTPUT forwards the request messages so contextual grounding can pull + # grounding_source/query blocks from them even on streamed responses. A + # plain-text (non-grounding) request still yields the single-block payload. + assert output_call["messages"] == request_data["messages"] # Verify that the response content was masked # The streaming chunks should now contain the masked content diff --git a/tests/integration/test_oci_proxy_integration.py b/tests/integration/test_oci_proxy_integration.py index 8bfcdd90486..f41e4826a04 100644 --- a/tests/integration/test_oci_proxy_integration.py +++ b/tests/integration/test_oci_proxy_integration.py @@ -30,6 +30,7 @@ locally-running proxy. from __future__ import annotations +import json import os import socket import subprocess @@ -41,7 +42,6 @@ from typing import Iterator import httpx import pytest - # --------------------------------------------------------------------------- # Skip gate # --------------------------------------------------------------------------- @@ -79,7 +79,9 @@ def _wait_for_health(base_url: str, proc: subprocess.Popen, deadline: float) -> except httpx.HTTPError: pass time.sleep(0.5) - raise RuntimeError(f"litellm proxy did not become ready within {STARTUP_TIMEOUT_S}s") + raise RuntimeError( + f"litellm proxy did not become ready within {STARTUP_TIMEOUT_S}s" + ) def _oci_env_from_profile() -> dict[str, str]: @@ -106,38 +108,35 @@ def _oci_env_from_profile() -> dict[str, str]: } -@pytest.fixture(scope="module") -def proxy_url() -> Iterator[str]: - oci_env = _oci_env_from_profile() - - port = _free_port() - base_url = f"http://127.0.0.1:{port}" - +def _serve(config_path: str) -> Iterator[str]: + """Boot the litellm proxy with the given config and yield its base URL.""" env = os.environ.copy() - env.update(oci_env) + env.update(_oci_env_from_profile()) # Avoid pulling in DB-backed features for this lightweight smoke run. env.pop("DATABASE_URL", None) env["STORE_MODEL_IN_DB"] = "False" + port = _free_port() + base_url = f"http://127.0.0.1:{port}" + # Prefer the `litellm` console script that lives next to the active # Python so we inherit the test virtualenv. Fall back to PATH. cli = Path(sys.executable).parent / "litellm" if not cli.exists(): cli = "litellm" - cmd = [ - str(cli), - "--config", - str(CONFIG_PATH), - "--port", - str(port), - "--host", - "127.0.0.1", - "--num_workers", - "1", - ] proc = subprocess.Popen( - cmd, + [ + str(cli), + "--config", + config_path, + "--port", + str(port), + "--host", + "127.0.0.1", + "--num_workers", + "1", + ], env=env, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, @@ -155,6 +154,27 @@ def proxy_url() -> Iterator[str]: proc.wait(timeout=5) +@pytest.fixture(scope="module") +def proxy_url() -> Iterator[str]: + yield from _serve(str(CONFIG_PATH)) + + +@pytest.fixture(scope="module") +def proxy_url_no_drop_params(tmp_path_factory) -> Iterator[str]: + """A proxy WITHOUT drop_params, to prove benign params the proxy injects + (e.g. max_retries) don't break OCI calls.""" + cfg = tmp_path_factory.mktemp("oci_nodrop") / "config.yaml" + cfg.write_text( + "model_list:\n" + " - model_name: oci-cohere-command\n" + " litellm_params:\n" + " model: oci/cohere.command-latest\n" + "general_settings:\n" + f" master_key: {MASTER_KEY}\n" + ) + yield from _serve(str(cfg)) + + # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- @@ -206,9 +226,7 @@ def test_chat_completion_via_proxy(proxy_url: str, model: str) -> None: # Reasoning models may return empty content if their budget covers only # the thinking turn — accept either text or a non-empty reasoning field. has_content = bool(msg.get("content")) - has_reasoning = bool(msg.get("reasoning_content")) or bool( - msg.get("reasoning") - ) + has_reasoning = bool(msg.get("reasoning_content")) or bool(msg.get("reasoning")) assert has_content or has_reasoning, f"empty assistant message for {model}: {msg}" usage = body.get("usage") or {} assert usage.get("total_tokens", 0) > 0 @@ -232,7 +250,7 @@ def test_chat_completion_streaming_via_proxy(proxy_url: str, model: str) -> None continue if not line.startswith("data:"): continue - payload = line[len("data:"):].strip() + payload = line[len("data:") :].strip() if payload == "[DONE]": saw_done = True break @@ -260,7 +278,6 @@ def test_embedding_via_proxy(proxy_url: str) -> None: assert len(embedding) >= 64 assert all(isinstance(x, (int, float)) for x in embedding) - def test_model_list_advertises_oci_models(proxy_url: str) -> None: """The /v1/models registry advertises every OCI alias from the config.""" r = httpx.get( @@ -272,3 +289,124 @@ def test_model_list_advertises_oci_models(proxy_url: str) -> None: advertised = {row["id"] for row in r.json()["data"]} for expected in CHAT_MODELS + ["oci-embed"]: assert expected in advertised, f"{expected} missing from /v1/models: {advertised}" + + +def test_chat_completion_no_drop_params(proxy_url_no_drop_params: str) -> None: + """A plain chat completion succeeds through a proxy without drop_params. + + Regression for the HTTP 500 ``param `max_retries` is not supported on OCI``: + the proxy injects max_retries on every request, so without this fix any OCI + call through the proxy failed unless drop_params was set. + """ + r = httpx.post( + f"{proxy_url_no_drop_params}/v1/chat/completions", + headers=_auth_headers(), + json=_chat_payload("oci-cohere-command"), + timeout=REQUEST_TIMEOUT_S, + ) + assert r.status_code == 200, f"no-drop_params -> {r.status_code}: {r.text}" + body = r.json() + assert body["object"] == "chat.completion" + assert body["choices"][0]["message"].get("content") is not None + + +def test_cohere_default_n_via_proxy(proxy_url: str) -> None: + """A Cohere request carrying the default n=1 succeeds through the gateway. + + Regression for the HTTP 500 ``param `n` is not supported on OCI`` that + rejected every client which always sends n=1 (e.g. the MLflow gateway), + since OCI Cohere has no numGenerations field. + """ + payload = {**_chat_payload("oci-cohere-command"), "n": 1} + r = httpx.post( + f"{proxy_url}/v1/chat/completions", + headers=_auth_headers(), + json=payload, + timeout=REQUEST_TIMEOUT_S, + ) + assert r.status_code == 200, f"n=1 -> {r.status_code}: {r.text}" + body = r.json() + assert body["object"] == "chat.completion" + assert body["choices"][0]["message"].get("content") is not None + + +@pytest.mark.parametrize("model", ["oci-cohere-command", "oci-llama"]) +def test_response_format_json_schema_via_proxy(proxy_url: str, model: str) -> None: + """A response_format json_schema succeeds through the gateway for both a + Cohere and a generic OCI model. + Regression for the HTTP 400 ``Please pass in correct format of request`` + that rejected every json_schema request (which MLflow LLM judges always + send): generic models choke on OpenAI's ``strict`` key, and Cohere has no + JSON_SCHEMA type. + """ + r = httpx.post( + f"{proxy_url}/v1/chat/completions", + headers=_auth_headers(), + json={ + "model": model, + "messages": [ + { + "role": "user", + "content": "Rate the answer 4 to 2+2. Give an integer score and a short rationale.", + } + ], + "max_tokens": 200, + "response_format": { + "type": "json_schema", + "json_schema": { + "name": "judgment", + "strict": True, + "schema": { + "type": "object", + "properties": { + "score": {"type": "integer"}, + "rationale": {"type": "string"}, + }, + "required": ["score", "rationale"], + "additionalProperties": False, + }, + }, + }, + }, + timeout=REQUEST_TIMEOUT_S, + ) + assert r.status_code == 200, f"{model} json_schema -> {r.status_code}: {r.text}" + content = r.json()["choices"][0]["message"]["content"] + assert content is not None + assert "score" in json.loads(content) + + +def test_omitted_max_tokens_not_truncated(proxy_url: str) -> None: + """A request that omits max_tokens completes instead of being cut off. + Regression for OCI's tiny server-side maxTokens default (~20 tokens): without + an injected default, a request that doesn't set max_tokens came back with + finish_reason "length" after ~19 tokens, so structured outputs (e.g. MLflow + judge JSON) arrived as unterminated strings. The OCI provider now injects a + sane default when the caller omits one. + """ + r = httpx.post( + f"{proxy_url}/v1/chat/completions", + headers=_auth_headers(), + json={ + "model": "oci-cohere-command", + "messages": [ + { + "role": "user", + "content": "In four or five complete sentences, explain why the sky appears blue.", + } + ], + }, + timeout=REQUEST_TIMEOUT_S, + ) + assert r.status_code == 200, f"omitted max_tokens -> {r.status_code}: {r.text}" + body = r.json() + choice = body["choices"][0] + assert ( + choice["finish_reason"] != "length" + ), f"response truncated by token cap: {choice}" + assert choice["finish_reason"] == "stop" + content = choice["message"].get("content") or "" + assert content.strip(), f"empty content: {choice}" + # The ~20-token server default truncated well before this; a complete + # four-to-five sentence answer comfortably exceeds it. + assert body["usage"]["completion_tokens"] > 50, body["usage"] diff --git a/tests/llm_translation/realtime/test_openai_realtime.py b/tests/llm_translation/realtime/test_openai_realtime.py index fc9f938b4cd..0e50e2792d6 100644 --- a/tests/llm_translation/realtime/test_openai_realtime.py +++ b/tests/llm_translation/realtime/test_openai_realtime.py @@ -393,3 +393,38 @@ async def test_realtime_query_params_use_normalized_model_name(monkeypatch): called_kwargs = mock_async_realtime.call_args.kwargs assert called_kwargs["query_params"]["model"] == "gpt-4o-realtime-preview" assert called_kwargs["query_params"]["intent"] == "chat" + + +@pytest.mark.asyncio +async def test_realtime_query_params_preserve_missing_model(monkeypatch): + """ + OpenAI-compatible transcription clients can connect with only + ?intent=transcription and send the model in session.update. Do not add + model= back into the upstream query params when the client omitted it. + """ + from litellm.realtime_api import main as realtime_main + + mock_async_realtime = AsyncMock() + monkeypatch.setattr( + realtime_main, + "openai_realtime", + MagicMock(async_realtime=mock_async_realtime), + ) + + def fake_get_llm_provider(model, api_base=None, api_key=None): + return ("gpt-realtime-whisper", "openai", None, None) + + monkeypatch.setattr(realtime_main, "get_llm_provider", fake_get_llm_provider) + + query_params: RealtimeQueryParams = {"intent": "transcription"} + + await realtime_main._arealtime( + model="gpt-realtime-whisper", + websocket=MagicMock(), + api_key="sk-test", + query_params=query_params, + litellm_logging_obj=MagicMock(), + ) + + called_kwargs = mock_async_realtime.call_args.kwargs + assert called_kwargs["query_params"] == {"intent": "transcription"} diff --git a/tests/proxy_unit_tests/test_realtime_cache.py b/tests/proxy_unit_tests/test_realtime_cache.py index c4cb4ea8e02..8316ed1d29a 100644 --- a/tests/proxy_unit_tests/test_realtime_cache.py +++ b/tests/proxy_unit_tests/test_realtime_cache.py @@ -44,10 +44,14 @@ def test_realtime_query_params_template_caches_each_pair_separately(): params_with_intent_first = _realtime_query_params_template("gpt-4o", "intent-a") params_with_intent_second = _realtime_query_params_template("gpt-4o", "intent-a") params_without_intent = _realtime_query_params_template("gpt-4o", None) + params_transcription_without_model = _realtime_query_params_template( + None, "transcription" + ) assert params_with_intent_first is params_with_intent_second assert params_with_intent_first == (("model", "gpt-4o"), ("intent", "intent-a")) assert params_without_intent == (("model", "gpt-4o"),) + assert params_transcription_without_model == (("intent", "transcription"),) assert params_with_intent_first is not params_without_intent diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py index 734033ed6be..a5bc01c2b74 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py @@ -125,6 +125,33 @@ def test_completion_skips_rewrapping_preformatted_cached_chat_stream(): assert result is stream +def test_completion_preserves_top_level_stream_flag_in_responses_request(): + stream = MagicMock(spec=CustomStreamWrapper) + stream.custom_llm_provider = "cached_response" + bridge = ResponsesToCompletionBridgeHandler() + kwargs = _bridge_kwargs(stream=False) + kwargs["stream"] = True + kwargs["optional_params"].pop("stream") + + with ( + patch.object( + bridge.transformation_handler, + "transform_request", + return_value={"model": "gpt-5.4", "input": "hi"}, + ) as transform_request, + patch("litellm.responses", return_value=stream), + patch.object( + bridge, + "_apply_post_stream_processing", + side_effect=lambda s, *a, **kw: s, + ), + ): + result = bridge.completion(**kwargs) + + assert result is stream + assert transform_request.call_args.kwargs["optional_params"]["stream"] is True + + @pytest.mark.asyncio async def test_acompletion_skips_rewrapping_preformatted_cached_chat_stream(): stream = MagicMock(spec=CustomStreamWrapper) @@ -148,3 +175,31 @@ async def test_acompletion_skips_rewrapping_preformatted_cached_chat_stream(): post.assert_called_once() assert result is stream + + +@pytest.mark.asyncio +async def test_acompletion_preserves_top_level_stream_flag_in_responses_request(): + stream = MagicMock(spec=CustomStreamWrapper) + stream.custom_llm_provider = "cached_response" + bridge = ResponsesToCompletionBridgeHandler() + kwargs = _bridge_kwargs(stream=False) + kwargs["stream"] = True + kwargs["optional_params"].pop("stream") + + with ( + patch.object( + bridge.transformation_handler, + "transform_request", + return_value={"model": "gpt-5.4", "input": "hi"}, + ) as transform_request, + patch("litellm.aresponses", new=AsyncMock(return_value=stream)), + patch.object( + bridge, + "_apply_post_stream_processing", + side_effect=lambda s, *a, **kw: s, + ), + ): + result = await bridge.acompletion(**kwargs) + + assert result is stream + assert transform_request.call_args.kwargs["optional_params"]["stream"] is True diff --git a/tests/test_litellm/integrations/SlackAlerting/test_budget_alert_types.py b/tests/test_litellm/integrations/SlackAlerting/test_budget_alert_types.py index efb8c1c4b28..52b7cc983a7 100644 --- a/tests/test_litellm/integrations/SlackAlerting/test_budget_alert_types.py +++ b/tests/test_litellm/integrations/SlackAlerting/test_budget_alert_types.py @@ -28,6 +28,31 @@ class TestSoftBudgetAlert: result = alert.get_id(user_info) assert result == "default_id" + def test_get_id_returns_team_id_for_team_event_group(self): + """Team soft budget alerts dedupe by team, not by the calling key's token""" + alert = SoftBudgetAlert() + user_info = CallInfo( + spend=120.0, + token="test_token_123", + team_id="team_456", + event_group=Litellm_EntityType.TEAM, + ) + + result = alert.get_id(user_info) + assert result == "team_456" + + def test_get_id_returns_default_id_for_team_event_group_without_team_id(self): + alert = SoftBudgetAlert() + user_info = CallInfo( + spend=120.0, + token="test_token_123", + team_id=None, + event_group=Litellm_EntityType.TEAM, + ) + + result = alert.get_id(user_info) + assert result == "default_id" + def test_get_id_with_empty_token(self): """Test that get_id returns 'default_id' when token is empty string""" alert = SoftBudgetAlert() diff --git a/tests/test_litellm/integrations/focus/test_mavvrik_destination.py b/tests/test_litellm/integrations/focus/test_mavvrik_destination.py new file mode 100644 index 00000000000..797238ae238 --- /dev/null +++ b/tests/test_litellm/integrations/focus/test_mavvrik_destination.py @@ -0,0 +1,751 @@ +"""Tests for FocusMavvrikDestination.""" + +from __future__ import annotations + +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.integrations.focus.destinations.base import FocusTimeWindow +from litellm.integrations.focus.destinations.mavvrik_destination import ( + FocusMavvrikDestination, + _validate_api_endpoint, +) + +VALID_ENDPOINT = "https://api.mavvrik.ai/tenant123" + + +def _make_window() -> FocusTimeWindow: + return FocusTimeWindow( + start_time=datetime(2026, 1, 1, 0, 0, 0, tzinfo=timezone.utc), + end_time=datetime(2026, 1, 2, 0, 0, 0, tzinfo=timezone.utc), + frequency="daily", + ) + + +def _dest(**overrides) -> FocusMavvrikDestination: + config = { + "api_key": "test-key", + "api_endpoint": VALID_ENDPOINT, + "connection_id": "conn-123", + } + config.update(overrides) + return FocusMavvrikDestination(prefix="mavvrik_focus_exports", config=config) + + +def test_missing_api_key_raises(): + with pytest.raises(ValueError, match="MAVVRIK_API_KEY"): + FocusMavvrikDestination( + prefix="p", + config={"api_endpoint": VALID_ENDPOINT, "connection_id": "c"}, + ) + + +def test_missing_api_endpoint_raises(): + with pytest.raises(ValueError, match="MAVVRIK_API_ENDPOINT"): + FocusMavvrikDestination( + prefix="p", + config={"api_key": "k", "connection_id": "c"}, + ) + + +def test_missing_connection_id_raises(): + with pytest.raises(ValueError, match="MAVVRIK_CONNECTION_ID"): + FocusMavvrikDestination( + prefix="p", + config={"api_key": "k", "api_endpoint": VALID_ENDPOINT}, + ) + + +def test_non_https_endpoint_raises(): + with pytest.raises(ValueError, match="HTTPS"): + _validate_api_endpoint("http://api.mavvrik.ai/tenant") + + +def test_non_mavvrik_domain_raises(): + with pytest.raises(ValueError, match="Mavvrik domain"): + _validate_api_endpoint("https://evil.com/tenant") + + +def test_valid_mavvrik_domains_accepted(): + for domain in ( + "https://api.mavvrik.ai/tenant", + "https://api.mavvrik.dev/tenant", + "https://api.mavvrik.app/tenant", + ): + _validate_api_endpoint(domain) # must not raise + + +def test_initializes_with_not_registered(): + dest = _dest() + assert dest._registered is False + + +@pytest.mark.asyncio +async def test_deliver_skips_empty_content(): + dest = _dest() + await dest.deliver(content=b"", time_window=_make_window(), filename="usage.csv") + # _registered still False — _ensure_registered was never called + assert dest._registered is False + + +@pytest.mark.asyncio +async def test_large_content_uploads_in_multiple_chunks(): + """Content larger than _GCS_CHUNK_SIZE must be uploaded in multiple chunks. + + GCS assembles intermediate chunks (308) + final chunk (200) into one object. + The destination must send Content-Range headers for each chunk correctly. + """ + from litellm.integrations.focus.destinations.mavvrik_destination import ( + FocusMavvrikDestination, + _GCS_CHUNK_SIZE, + ) + + dest = FocusMavvrikDestination( + prefix="p", + config={"api_key": "k", "api_endpoint": VALID_ENDPOINT, "connection_id": "c"}, + ) + + register_resp = MagicMock() + register_resp.status_code = 200 + + signed_url_resp = MagicMock() + signed_url_resp.status_code = 200 + signed_url_resp.json.return_value = { + "url": "https://storage.googleapis.com/upload?sig=x" + } + + init_resp = MagicMock() + init_resp.status_code = 200 + init_resp.headers = {"Location": "https://storage.googleapis.com/session"} + + # First chunk → 308, second (final) chunk → 200 + chunk1_resp = MagicMock() + chunk1_resp.status_code = 308 + + chunk2_resp = MagicMock() + chunk2_resp.status_code = 200 + + mock_http = MagicMock() + mock_http.client = MagicMock() + mock_http.client.request = AsyncMock( + side_effect=[ + register_resp, + signed_url_resp, + init_resp, + chunk1_resp, + chunk2_resp, + ] + ) + dest._http = mock_http + + # Build content that when gzipped exceeds one chunk. + # Use incompressible random-ish bytes to ensure gzip doesn't shrink it below the chunk size. + import os as _os + + raw = b"col1,col2\n" + _os.urandom(_GCS_CHUNK_SIZE + 1024) + + await dest.deliver( + content=raw, + time_window=_make_window(), + filename="usage.csv", + ) + + # register + get_signed_url + init + 2 chunk PUTs = 5 calls + assert mock_http.client.request.call_count == 5 + + # Check Content-Range headers + put_calls = mock_http.client.request.call_args_list[3:] + assert "bytes" in put_calls[0].kwargs["headers"]["Content-Range"] + assert "/*" in put_calls[0].kwargs["headers"]["Content-Range"] # intermediate + assert "/*" not in put_calls[1].kwargs["headers"]["Content-Range"] # final + + +@pytest.mark.asyncio +async def test_deliver_calls_register_get_url_and_upload(): + dest = _dest() + + register_resp = MagicMock() + register_resp.status_code = 200 + + signed_url_resp = MagicMock() + signed_url_resp.status_code = 200 + signed_url_resp.json.return_value = {"url": "https://storage.googleapis.com/signed"} + + init_resp = MagicMock() + init_resp.status_code = 200 + init_resp.headers = {"Location": "https://storage.googleapis.com/session-uri"} + + upload_resp = MagicMock() + upload_resp.status_code = 200 + + mock_http = MagicMock() + mock_http.client = MagicMock() + # All 4 calls go through self._http.client.request: + # 1. register, 2. get_signed_url, 3. GCS session init POST, 4. GCS PUT + mock_http.client.request = AsyncMock( + side_effect=[register_resp, signed_url_resp, init_resp, upload_resp] + ) + dest._http = mock_http + + await dest.deliver( + content=b"header\nrow1\n", + time_window=_make_window(), + filename="usage.csv", + ) + + assert dest._registered is True + assert mock_http.client.request.call_count == 4 + # Verify Content-Range header was set on the PUT + put_call = mock_http.client.request.call_args_list[3] + assert "Content-Range" in put_call.kwargs["headers"] + + +@pytest.mark.asyncio +async def test_register_called_only_once_across_multiple_deliveries(): + dest = _dest() + + register_resp = MagicMock() + register_resp.status_code = 200 + + def _signed_url_resp(): + r = MagicMock() + r.status_code = 200 + r.json.return_value = {"url": "https://storage.googleapis.com/signed"} + return r + + init_resp = MagicMock() + init_resp.status_code = 200 + init_resp.headers = {"Location": "https://storage.googleapis.com/session-uri"} + + upload_resp = MagicMock() + upload_resp.status_code = 200 + + mock_http = MagicMock() + mock_http.client = MagicMock() + # First delivery: register, get_signed_url, GCS init, GCS PUT + # Second delivery: get_signed_url, GCS init, GCS PUT (register skipped) + mock_http.client.request = AsyncMock( + side_effect=[ + register_resp, + _signed_url_resp(), + init_resp, + upload_resp, + _signed_url_resp(), + init_resp, + upload_resp, + ] + ) + dest._http = mock_http + + window = _make_window() + await dest.deliver(content=b"header\nrow1\n", time_window=window, filename="1.csv") + await dest.deliver(content=b"header\nrow2\n", time_window=window, filename="2.csv") + + # 7 total: register(1) + [get_url+init+put](2) × 2 deliveries + assert mock_http.client.request.call_count == 7 + # First call was register + first_call = mock_http.client.request.call_args_list[0] + assert first_call.kwargs["method"] == "POST" + assert "/upload-url" not in first_call.kwargs["url"] + + +@pytest.mark.asyncio +async def test_deliver_raises_on_register_failure(): + dest = _dest() + + fail_resp = MagicMock() + fail_resp.status_code = 403 + fail_resp.text = "Forbidden" + + mock_http = MagicMock() + mock_http.client = MagicMock() + mock_http.client.request = AsyncMock(return_value=fail_resp) + dest._http = mock_http + + with pytest.raises(RuntimeError, match="register failed"): + await dest.deliver( + content=b"data", + time_window=_make_window(), + filename="usage.csv", + ) + + +@pytest.mark.asyncio +async def test_deliver_raises_on_signed_url_api_error(): + """_get_signed_url must raise RuntimeError when the API returns a 4xx.""" + dest = _dest() + + register_resp = MagicMock() + register_resp.status_code = 200 + + fail_resp = MagicMock() + fail_resp.status_code = 500 + fail_resp.text = "Internal Server Error" + + mock_http = MagicMock() + mock_http.client = MagicMock() + mock_http.client.request = AsyncMock(side_effect=[register_resp, fail_resp]) + dest._http = mock_http + + with pytest.raises(RuntimeError, match="failed to get signed URL"): + await dest.deliver( + content=b"data", + time_window=_make_window(), + filename="usage.csv", + ) + + +@pytest.mark.asyncio +async def test_deliver_raises_on_missing_signed_url(): + dest = _dest() + + register_resp = MagicMock() + register_resp.status_code = 200 + + bad_url_resp = MagicMock() + bad_url_resp.status_code = 200 + bad_url_resp.json.return_value = {} # no 'url' field + + mock_http = MagicMock() + mock_http.client = MagicMock() + mock_http.client.request = AsyncMock(side_effect=[register_resp, bad_url_resp]) + dest._http = mock_http + + with pytest.raises(RuntimeError, match="missing 'url' field"): + await dest.deliver( + content=b"data", + time_window=_make_window(), + filename="usage.csv", + ) + + +@pytest.mark.asyncio +async def test_deliver_raises_on_non_gcs_signed_url(): + """Signed URL pointing to a non-GCS host must be rejected before any upload.""" + from litellm.integrations.focus.destinations.mavvrik_destination import ( + _validate_gcs_url, + ) + + with pytest.raises(ValueError, match="GCS endpoint"): + _validate_gcs_url("https://evil.com/upload?token=abc", "signed URL") + + +@pytest.mark.asyncio +async def test_deliver_raises_on_non_gcs_session_uri(): + """Session URI from Location header pointing to a non-GCS host must be rejected.""" + dest = _dest() + + register_resp = MagicMock() + register_resp.status_code = 200 + + signed_url_resp = MagicMock() + signed_url_resp.status_code = 200 + # signed URL is valid GCS + signed_url_resp.json.return_value = { + "url": "https://storage.googleapis.com/upload?sig=abc" + } + + # Location header points to a non-GCS host + init_resp = MagicMock() + init_resp.status_code = 200 + init_resp.headers = {"Location": "https://evil.com/session-uri"} + + mock_http = MagicMock() + mock_http.client = MagicMock() + # register, get_signed_url, GCS session init (returns bad Location) + mock_http.client.request = AsyncMock( + side_effect=[register_resp, signed_url_resp, init_resp] + ) + dest._http = mock_http + + with pytest.raises(ValueError, match="GCS endpoint"): + await dest.deliver( + content=b"data", + time_window=_make_window(), + filename="usage.csv", + ) + + +def test_factory_creates_mavvrik_destination(monkeypatch): + monkeypatch.setenv("MAVVRIK_API_KEY", "k") + monkeypatch.setenv("MAVVRIK_API_ENDPOINT", VALID_ENDPOINT) + monkeypatch.setenv("MAVVRIK_CONNECTION_ID", "c") + + from litellm.integrations.focus.destinations.factory import FocusDestinationFactory + + dest = FocusDestinationFactory.create(provider="mavvrik", prefix="p") + + assert isinstance(dest, FocusMavvrikDestination) + assert dest.api_key == "k" + assert dest.connection_id == "c" + + +def test_only_daily_frequency_is_supported(): + """MavvrikFocusLogger must raise ValueError for non-daily frequencies.""" + import importlib + + for freq in ("hourly", "interval"): + + def _make(f=freq, monkeypatch=None): + import os + + old = os.environ.get("MAVVRIK_FOCUS_FREQUENCY") + os.environ["MAVVRIK_FOCUS_FREQUENCY"] = f + try: + from litellm.integrations.mavvrik_focus import mavvrik_focus_logger + + importlib.reload(mavvrik_focus_logger) + with pytest.raises(ValueError, match="Only 'daily' is allowed"): + mavvrik_focus_logger.MavvrikFocusLogger() + finally: + if old is None: + os.environ.pop("MAVVRIK_FOCUS_FREQUENCY", None) + else: + os.environ["MAVVRIK_FOCUS_FREQUENCY"] = old + + _make() + + +def test_max_rows_defaults_to_500k(): + """MAVVRIK_FOCUS_MAX_ROWS defaults to 500_000 when not set.""" + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + MavvrikFocusLogger, + ) + + logger = MavvrikFocusLogger() + assert logger._max_rows == 500_000 + + +def test_max_rows_reads_from_env(monkeypatch): + """MAVVRIK_FOCUS_MAX_ROWS env var is respected.""" + monkeypatch.setenv("MAVVRIK_FOCUS_MAX_ROWS", "100000") + + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + MavvrikFocusLogger, + ) + + logger = MavvrikFocusLogger() + assert logger._max_rows == 100_000 + + +@pytest.mark.asyncio +async def test_export_window_passes_max_rows_as_limit(monkeypatch): + """_export_window must pass _max_rows as limit to get_usage_data.""" + monkeypatch.setenv("MAVVRIK_FOCUS_MAX_ROWS", "1000") + + import polars as pl + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + MavvrikFocusLogger, + ) + from litellm.integrations.focus.destinations.base import FocusTimeWindow + from datetime import datetime, timezone + + logger = MavvrikFocusLogger() + assert logger._max_rows == 1000 + + # Mock the engine internals so _export_window runs through our new code path + db_mock = MagicMock() + db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) # empty → no upload + + engine_mock = MagicMock() + engine_mock._database = db_mock + logger._engine = engine_mock + + window = FocusTimeWindow( + start_time=datetime(2026, 1, 1, tzinfo=timezone.utc), + end_time=datetime(2026, 1, 2, tzinfo=timezone.utc), + frequency="daily", + ) + await logger._export_window(window=window, limit=None) + + db_mock.get_usage_data.assert_called_once_with( + limit=1000, + start_time_utc=window.start_time, + end_time_utc=window.end_time, + ) + + +@pytest.mark.asyncio +async def test_run_scheduled_export_catches_up_missed_dates(): + """If metricsMarker is 2 days behind, _run_scheduled_export exports missed dates first.""" + import polars as pl + from datetime import datetime, timedelta, timezone + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + MavvrikFocusLogger, + ) + from litellm.integrations.focus.destinations.mavvrik_destination import ( + FocusMavvrikDestination, + ) + + logger = MavvrikFocusLogger() + + # metricsMarker = 3 days ago → 2 missed dates (day-2 and day-1) + today's run + now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0) + yesterday = now - timedelta(days=1) + two_days_ago = now - timedelta(days=2) + three_days_ago = now - timedelta(days=3) + + marker_ts = int(three_days_ago.timestamp()) + + # Mock destination + dest_mock = MagicMock(spec=FocusMavvrikDestination) + dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts) + + # Mock engine + db_mock = MagicMock() + db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) + engine_mock = MagicMock() + engine_mock._database = db_mock + engine_mock._destination = dest_mock + logger._engine = engine_mock + + await logger._run_scheduled_export() + + # Should have queried DB 3 times: day-2, day-1 (yesterday), and the normal yesterday window + # Actually: catch-up covers [three_days_ago+1 .. yesterday) = [two_days_ago, yesterday) + # = two_days_ago only (1 missed), then normal yesterday = 2 total calls + calls = db_mock.get_usage_data.call_args_list + assert len(calls) == 2 + # First call is the catch-up (two_days_ago) + assert calls[0].kwargs["start_time_utc"].date() == two_days_ago.date() + # Second call is yesterday's normal daily run + assert calls[1].kwargs["start_time_utc"].date() == yesterday.date() + + +@pytest.mark.asyncio +async def test_run_scheduled_export_no_catchup_when_marker_is_current(): + """If metricsMarker = yesterday, no catch-up needed — just export yesterday.""" + import polars as pl + from datetime import datetime, timedelta, timezone + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + MavvrikFocusLogger, + ) + from litellm.integrations.focus.destinations.mavvrik_destination import ( + FocusMavvrikDestination, + ) + + logger = MavvrikFocusLogger() + + now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0) + yesterday = now - timedelta(days=1) + marker_ts = int(yesterday.timestamp()) + + dest_mock = MagicMock(spec=FocusMavvrikDestination) + dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts) + + db_mock = MagicMock() + db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) + engine_mock = MagicMock() + engine_mock._database = db_mock + engine_mock._destination = dest_mock + logger._engine = engine_mock + + await logger._run_scheduled_export() + + # Only one call — yesterday's normal run, no catch-up + assert db_mock.get_usage_data.call_count == 1 + assert ( + db_mock.get_usage_data.call_args.kwargs["start_time_utc"].date() + == yesterday.date() + ) + + +@pytest.mark.asyncio +async def test_metrics_marker_always_calls_api(): + """get_metrics_marker must call the register API every time to get a fresh marker. + + This is the key difference from deliver() — catch-up requires the current + metricsMarker on every scheduled run, not just the first one. + """ + dest = _dest() + + register_resp = MagicMock() + register_resp.status_code = 200 + register_resp.json.return_value = { + "id": "litellm-conn-123", + "metricsMarker": 1749340800, + } + + mock_http = MagicMock() + mock_http.client = MagicMock() + mock_http.client.request = AsyncMock(return_value=register_resp) + dest._http = mock_http + + # First call + marker = await dest.get_metrics_marker() + assert marker == 1749340800 + assert dest._registered is True + + # Second call — must call API again to get fresh marker (not return None) + marker2 = await dest.get_metrics_marker() + assert marker2 == 1749340800 + assert mock_http.client.request.call_count == 2 # API called both times + + +def test_parse_metrics_marker_handles_unix_timestamp(): + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + _parse_metrics_marker, + ) + from datetime import datetime, timezone + + # Use a known date and compute its timestamp to avoid hardcoding + known_date = datetime(2026, 6, 9, 0, 0, 0, tzinfo=timezone.utc) + ts = int(known_date.timestamp()) + + result = _parse_metrics_marker(ts) + assert result is not None + assert result.date().isoformat() == "2026-06-09" + assert result.tzinfo == timezone.utc + + +def test_parse_metrics_marker_handles_iso_date_string(): + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + _parse_metrics_marker, + ) + + result = _parse_metrics_marker("2026-06-09") + assert result is not None + assert result.date().isoformat() == "2026-06-09" + + +def test_parse_metrics_marker_handles_iso_datetime_string(): + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + _parse_metrics_marker, + ) + + result = _parse_metrics_marker("2026-06-09T00:00:00Z") + assert result is not None + assert result.date().isoformat() == "2026-06-09" + + +def test_parse_metrics_marker_returns_none_for_zero(): + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + _parse_metrics_marker, + ) + + assert _parse_metrics_marker(0) is None + assert _parse_metrics_marker(None) is None + assert _parse_metrics_marker("") is None + + +def test_parse_metrics_marker_returns_none_for_garbage(): + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + _parse_metrics_marker, + ) + + # Should not raise — logs warning and returns None + assert _parse_metrics_marker("not-a-date") is None + + +@pytest.mark.asyncio +async def test_catchup_capped_at_max_catchup_days(): + """Catch-up must not go further back than _MAX_CATCHUP_DAYS.""" + import polars as pl + from datetime import datetime, timedelta, timezone + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + MavvrikFocusLogger, + ) + from litellm.integrations.focus.destinations.mavvrik_destination import ( + FocusMavvrikDestination, + ) + + logger = MavvrikFocusLogger() + max_days = MavvrikFocusLogger._MAX_CATCHUP_DAYS + + now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0) + yesterday = now - timedelta(days=1) + # Marker is 30 days ago — well beyond the cap + thirty_days_ago = now - timedelta(days=30) + marker_ts = int(thirty_days_ago.timestamp()) + + dest_mock = MagicMock(spec=FocusMavvrikDestination) + dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts) + + db_mock = MagicMock() + db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) + engine_mock = MagicMock() + engine_mock._database = db_mock + engine_mock._destination = dest_mock + logger._engine = engine_mock + + await logger._run_scheduled_export() + + # Should have queried at most _MAX_CATCHUP_DAYS times + # (max_days - 1 catch-up dates + 1 yesterday = max_days total) + assert db_mock.get_usage_data.call_count <= max_days + + # First catch-up date must not be earlier than (yesterday - max_days + 1) + earliest_allowed = yesterday - timedelta(days=max_days - 1) + first_call_start = db_mock.get_usage_data.call_args_list[0].kwargs["start_time_utc"] + assert first_call_start.date() >= earliest_allowed.date() + + +@pytest.mark.asyncio +async def test_register_resets_on_410(): + """_registered flag must be False after a 410 so next run re-registers.""" + dest = _dest() + + resp_410 = MagicMock() + resp_410.status_code = 410 + resp_410.text = "Gone" + + mock_http = MagicMock() + mock_http.client = MagicMock() + mock_http.client.request = AsyncMock(return_value=resp_410) + dest._http = mock_http + dest._registered = False # not yet registered — trigger the call + + with pytest.raises(RuntimeError, match="disconnected"): + await dest._ensure_registered() + + assert dest._registered is False + + +@pytest.mark.asyncio +async def test_gcs_session_cancelled_on_chunk_failure(): + """GCS session must be cancelled (DELETE) when a chunk PUT fails.""" + dest = _dest() + + register_resp = MagicMock() + register_resp.status_code = 200 + + signed_url_resp = MagicMock() + signed_url_resp.status_code = 200 + signed_url_resp.json.return_value = { + "url": "https://storage.googleapis.com/upload?sig=x" + } + + init_resp = MagicMock() + init_resp.status_code = 200 + init_resp.headers = {"Location": "https://storage.googleapis.com/session"} + + # Chunk PUT fails with 500 + fail_resp = MagicMock() + fail_resp.status_code = 500 + fail_resp.text = "Internal Server Error" + + # DELETE (session cancel) + delete_resp = MagicMock() + delete_resp.status_code = 200 + + mock_http = MagicMock() + mock_http.client = MagicMock() + mock_http.client.request = AsyncMock( + side_effect=[register_resp, signed_url_resp, init_resp, fail_resp, delete_resp] + ) + dest._http = mock_http + + with pytest.raises(RuntimeError, match="GCS chunk upload failed"): + await dest.deliver( + content=b"header\nrow1\n", + time_window=_make_window(), + filename="usage.csv", + ) + + # Verify DELETE was called to cancel the session + calls = mock_http.client.request.call_args_list + delete_call = calls[4] + assert delete_call.kwargs["method"] == "DELETE" + assert "storage.googleapis.com/session" in delete_call.kwargs["url"] diff --git a/tests/test_litellm/integrations/newrelic/test_newrelic.py b/tests/test_litellm/integrations/newrelic/test_newrelic.py new file mode 100644 index 00000000000..541b271cb77 --- /dev/null +++ b/tests/test_litellm/integrations/newrelic/test_newrelic.py @@ -0,0 +1,1351 @@ +import os +import sys +from datetime import datetime, timezone +from unittest.mock import MagicMock, patch + +import pytest + +# newrelic is a proxy-runtime dependency (pyproject.toml) and is not installed +# in the CI Python environment. Mock it in sys.modules before importing the +# integration so that deferred `import newrelic.agent` calls inside NewRelicLogger +# methods resolve to these mocks rather than failing with ModuleNotFoundError. +_mock_newrelic = MagicMock() +_mock_newrelic_agent = MagicMock() +# Explicitly link so _mock_newrelic.agent IS _mock_newrelic_agent. Without this, +# the first getattr(_mock_newrelic, 'agent') auto-creates a different child mock, +# causing patch("newrelic.agent.xxx") to patch the wrong object. +_mock_newrelic.agent = _mock_newrelic_agent +sys.modules["newrelic"] = _mock_newrelic +sys.modules["newrelic.agent"] = _mock_newrelic_agent + +import litellm +import litellm.integrations.newrelic.newrelic as nr_module +from litellm.integrations.newrelic.newrelic import NewRelicLogger + +# The module may have been imported before sys.modules was patched (e.g. via +# litellm's own startup imports), leaving _newrelic_agent=None. Point it at +# the mock agent so all tests see a non-None agent. +nr_module._newrelic_agent = _mock_newrelic_agent + + +# --------------------------------------------------------------------------- +# Shared fixtures +# --------------------------------------------------------------------------- + +NR_ENV = { + "NEW_RELIC_LICENSE_KEY": "test-license-key", + "NEW_RELIC_APP_NAME": "test-app", +} + + +def make_logger(**kwargs) -> NewRelicLogger: + """Instantiate NewRelicLogger with NR agent calls mocked out.""" + with patch.dict(os.environ, NR_ENV): + return NewRelicLogger(**kwargs) + + +def make_kwargs( + model="gpt-4", + provider="openai", + messages=None, + optional_params=None, + traceparent=None, +) -> dict: + """Build a minimal kwargs dict representative of a litellm callback invocation.""" + headers = {} + if traceparent: + headers["traceparent"] = traceparent + + return { + "model": model, + "messages": messages or [{"role": "user", "content": "Hello"}], + "optional_params": optional_params or {}, + "litellm_params": { + "custom_llm_provider": provider, + "metadata": {"headers": headers}, + }, + "start_time": 1_000_000.0, + "end_time": 1_000_001.5, + "llm_api_duration_ms": 1500.0, + } + + +def make_response( + model="gpt-4", + response_id="chatcmpl-abc123", + content="Hello there!", + finish_reason="stop", + prompt_tokens=10, + completion_tokens=20, +): + """Build a minimal ModelResponse-like dict.""" + return { + "id": response_id, + "model": model, + "choices": [ + { + "message": {"role": "assistant", "content": content}, + "finish_reason": finish_reason, + } + ], + "usage": { + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "total_tokens": prompt_tokens + completion_tokens, + }, + } + + +def make_slo(**overrides): + """Build a StandardLoggingPayload-like dict with sentinel values distinct from + make_kwargs/make_response defaults, so tests can prove the SLO branch won.""" + base = { + "trace_id": "slo-trace-abc", + "custom_llm_provider": "slo-provider", + "model": "slo-model", + "prompt_tokens": 100, + "completion_tokens": 200, + "total_tokens": 300, + "response_time": 1.5, # seconds; converted to ms by _get_duration + "model_parameters": {"temperature": 0.7, "max_tokens": 500}, + "startTime": 2_000_000.0, + "endTime": 2_000_001.5, + "messages": [{"role": "user", "content": "from-slo"}], + } + base.update(overrides) + return base + + +# --------------------------------------------------------------------------- +# Init / configuration +# --------------------------------------------------------------------------- + + +class TestNewRelicLoggerInit: + def test_disabled_when_license_key_missing(self): + with patch("newrelic.agent.register_application"): + with patch.dict(os.environ, {"NEW_RELIC_APP_NAME": "app"}, clear=True): + logger = NewRelicLogger() + assert logger.enabled is False + + def test_disabled_when_app_name_missing(self): + with patch("newrelic.agent.register_application"): + with patch.dict(os.environ, {"NEW_RELIC_LICENSE_KEY": "key"}, clear=True): + logger = NewRelicLogger() + assert logger.enabled is False + + def test_enabled_with_valid_env_vars(self): + with patch("newrelic.agent.register_application"): + with patch.dict(os.environ, NR_ENV): + logger = NewRelicLogger() + assert logger.enabled is True + + def test_disabled_on_import_error(self): + with patch.object( + _mock_newrelic_agent, "register_application", side_effect=ImportError + ): + with patch.dict(os.environ, NR_ENV): + logger = NewRelicLogger() + assert logger.enabled is False + + def test_disabled_on_agent_startup_error(self): + with patch.object( + _mock_newrelic_agent, + "register_application", + side_effect=RuntimeError("agent startup failed"), + ): + with patch.dict(os.environ, NR_ENV): + logger = NewRelicLogger() + assert logger.enabled is False + + def test_disabled_when_agent_package_missing(self): + with patch.object(nr_module, "_newrelic_agent", None): + with patch.dict(os.environ, NR_ENV): + logger = NewRelicLogger() + assert logger.enabled is False + + def test_record_content_default_true(self): + logger = make_logger() + assert logger.record_content is True + + def test_record_content_disabled_by_param(self): + logger = make_logger(turn_off_message_logging=True) + assert logger.record_content is False + + def test_record_content_disabled_by_env_var(self): + with patch("newrelic.agent.register_application"): + with patch.dict( + os.environ, + {**NR_ENV, "NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED": "false"}, + ): + logger = NewRelicLogger() + assert logger.record_content is False + + def test_record_content_requires_both_enabled(self): + """param says record, but env var says no — result is False.""" + with patch("newrelic.agent.register_application"): + with patch.dict( + os.environ, + {**NR_ENV, "NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED": "false"}, + ): + logger = NewRelicLogger(turn_off_message_logging=False) + assert logger.record_content is False + + def test_constructor_kwargs_take_priority_over_global_params(self): + """Constructor turn_off_message_logging=True must not be overwritten by + litellm.newrelic_params which defaults turn_off_message_logging to False.""" + from litellm.types.integrations.newrelic import NewRelicInitParams + + with patch("newrelic.agent.register_application"): + with patch.dict(os.environ, NR_ENV): + with patch( + "litellm.newrelic_params", + NewRelicInitParams(turn_off_message_logging=False), + ): + logger = NewRelicLogger(turn_off_message_logging=True) + assert logger.record_content is False + + def test_newrelic_params_plain_dict_branch(self): + """litellm.newrelic_params can be a plain dict; it should be validated + through NewRelicInitParams and its values applied to the logger.""" + with patch("newrelic.agent.register_application"): + with patch.dict(os.environ, NR_ENV): + with patch( + "litellm.newrelic_params", + {"turn_off_message_logging": True}, + ): + logger = NewRelicLogger() + assert logger.turn_off_message_logging is True + + +# --------------------------------------------------------------------------- +# _parse_bool_env +# --------------------------------------------------------------------------- + + +class TestParseBoolEnv: + def setup_method(self): + self.logger = make_logger() + + @pytest.mark.parametrize("raw", ["true", "TRUE", "True", "1", "yes", "on", "ON"]) + def test_truthy_values(self, raw): + with patch.dict(os.environ, {"MY_VAR": raw}): + assert self.logger._parse_bool_env("MY_VAR") is True + + @pytest.mark.parametrize("raw", ["false", "FALSE", "0", "no", "off", "Off"]) + def test_falsy_values(self, raw): + with patch.dict(os.environ, {"MY_VAR": raw}): + assert self.logger._parse_bool_env("MY_VAR") is False + + @pytest.mark.parametrize("raw", [" true ", " 1\t", "\nyes"]) + def test_whitespace_tolerance_truthy(self, raw): + with patch.dict(os.environ, {"MY_VAR": raw}): + assert self.logger._parse_bool_env("MY_VAR") is True + + @pytest.mark.parametrize("raw", [" false ", " 0\t", "\nno"]) + def test_whitespace_tolerance_falsy(self, raw): + with patch.dict(os.environ, {"MY_VAR": raw}): + assert self.logger._parse_bool_env("MY_VAR") is False + + def test_missing_uses_default(self): + with patch.dict(os.environ, {}, clear=True): + assert self.logger._parse_bool_env("MY_VAR", default=True) is True + assert self.logger._parse_bool_env("MY_VAR", default=False) is False + + def test_empty_string_uses_default(self): + with patch.dict(os.environ, {"MY_VAR": ""}): + assert self.logger._parse_bool_env("MY_VAR", default=True) is True + assert self.logger._parse_bool_env("MY_VAR", default=False) is False + + @pytest.mark.parametrize("raw", ["maybe", "2", "enabled", "tru"]) + def test_unrecognised_value_falls_back_to_default_with_warning(self, raw): + with ( + patch.dict(os.environ, {"MY_VAR": raw}), + patch.object(nr_module.verbose_logger, "warning") as mock_warn, + ): + assert self.logger._parse_bool_env("MY_VAR", default=True) is True + assert self.logger._parse_bool_env("MY_VAR", default=False) is False + assert mock_warn.call_count == 2 + # Warning should mention the variable name and the raw value + for call in mock_warn.call_args_list: + assert "MY_VAR" in call.args[0] + assert repr(raw) in call.args[0] + + +# --------------------------------------------------------------------------- +# _get_trace_context +# --------------------------------------------------------------------------- + + +class TestGetTraceContext: + def setup_method(self): + self.logger = make_logger() + + def test_extracts_trace_id_from_traceparent(self): + kwargs = make_kwargs( + traceparent="00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00" + ) + trace_id = self.logger._get_trace_context(kwargs) + assert trace_id == "4bf92f3577b34da6a3ce929d0e0e4736" + + def test_generates_uuid_when_no_headers(self): + kwargs = make_kwargs() + trace_id = self.logger._get_trace_context(kwargs) + assert trace_id is not None + assert ( + len(trace_id) == 32 + ) # 32-char lowercase hex, matches W3C traceparent format + + def test_generates_uuid_when_traceparent_malformed(self): + kwargs = make_kwargs(traceparent="not-valid") + trace_id = self.logger._get_trace_context(kwargs) + # Falls back to a 32-char lowercase hex, matching W3C traceparent format + assert trace_id is not None + assert len(trace_id) == 32 + + def test_extracts_trace_id_from_mixed_case_traceparent_header(self): + # Callers passing headers directly may not normalise case; per W3C spec + # header names are case-insensitive, so "Traceparent" must work too. + kwargs = make_kwargs() + kwargs["litellm_params"]["metadata"]["headers"] = { + "Traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00" + } + trace_id = self.logger._get_trace_context(kwargs) + assert trace_id == "4bf92f3577b34da6a3ce929d0e0e4736" + + def test_parse_failure_falls_through_to_synthetic_uuid(self): + """When parsing upstream sources raises, emit a synthetic UUID rather + than dropping the event. NR schema requires every AIM event carry a + trace_id; this method's contract is to always return a valid string. + """ + # Non-dict headers value forces .items() to raise inside the try + kwargs = {"litellm_params": {"metadata": {"headers": "not-a-dict"}}} + trace_id = self.logger._get_trace_context(kwargs) + assert trace_id is not None + assert len(trace_id) == 32 # 32-char lowercase hex fallback + + +# --------------------------------------------------------------------------- +# _extract_message_content edge cases +# --------------------------------------------------------------------------- + + +class TestExtractMessageContent: + def setup_method(self): + self.logger = make_logger() + + def test_plain_text(self): + assert self.logger._extract_message_content({"content": "hello"}) == "hello" + + def test_none_content_returns_empty_string(self): + assert self.logger._extract_message_content({"content": None}) == "" + + def test_missing_content_returns_empty_string(self): + assert self.logger._extract_message_content({}) == "" + + def test_tool_calls_serialized_as_json(self): + msg = { + "content": None, + "tool_calls": [{"id": "call_1", "function": {"name": "get_weather"}}], + } + result = self.logger._extract_message_content(msg) + assert "get_weather" in result + assert "call_1" in result + + def test_multimodal_list_serialized_as_json(self): + msg = { + "content": [ + {"type": "text", "text": "describe this"}, + {"type": "image_url"}, + ] + } + result = self.logger._extract_message_content(msg) + assert "describe this" in result + assert "image_url" in result + + def test_non_string_content_coerced_to_str(self): + """Numeric/bool content passes the None and list guards; final branch coerces to str.""" + assert self.logger._extract_message_content({"content": 123}) == "123" + assert self.logger._extract_message_content({"content": True}) == "True" + + +# --------------------------------------------------------------------------- +# _extract_all_messages — record_content=False path +# --------------------------------------------------------------------------- + + +class TestExtractAllMessagesContentDisabled: + def test_no_content_key_when_recording_disabled(self): + logger = make_logger(turn_off_message_logging=True) + kwargs = make_kwargs(messages=[{"role": "user", "content": "secret"}]) + response = make_response(content="also secret") + + messages = logger._extract_all_messages( + kwargs, response, response_model="gpt-4", vendor="openai" + ) + + for msg in messages: + assert "content" not in msg + + +class TestExtractAllMessagesRespectsLitellmRedaction: + """Regression tests for the async-streaming redaction bypass. + + NR-specific switches alone are insufficient: when + ``litellm.turn_off_message_logging=True`` (or the per-request equivalents), + async streaming callbacks receive an unredacted + ``async_complete_streaming_response``. Without consulting LiteLLM's + redaction decision the integration would still write generated content + into NR events. + """ + + def _assert_no_content(self, logger, kwargs): + response = make_response(content="streamed assistant text") + + messages = logger._extract_all_messages( + kwargs, response, response_model="gpt-4", vendor="openai" + ) + + # All extracted messages must carry no content payload + for msg in messages: + assert ( + "content" not in msg + ), f"content leaked despite redaction signal: {msg}" + # And there must actually be at least one user + one assistant entry, + # otherwise the test would pass vacuously. + assert any(not m.get("is_response") for m in messages) + assert any(m.get("is_response") for m in messages) + + def test_global_turn_off_message_logging_blocks_content(self, monkeypatch): + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + logger = make_logger() + assert logger.record_content is True + + kwargs = make_kwargs(messages=[{"role": "user", "content": "user prompt"}]) + self._assert_no_content(logger, kwargs) + + def test_dynamic_param_turn_off_message_logging_blocks_content(self): + logger = make_logger() + assert logger.record_content is True + + kwargs = make_kwargs(messages=[{"role": "user", "content": "user prompt"}]) + kwargs["standard_callback_dynamic_params"] = { + "turn_off_message_logging": True, + } + self._assert_no_content(logger, kwargs) + + def test_enable_redaction_header_blocks_content(self): + logger = make_logger() + assert logger.record_content is True + + kwargs = make_kwargs(messages=[{"role": "user", "content": "user prompt"}]) + kwargs["litellm_params"]["metadata"]["headers"] = { + "x-litellm-enable-message-redaction": True, + } + self._assert_no_content(logger, kwargs) + + def test_dynamic_param_explicit_false_overrides_global_redaction(self, monkeypatch): + """The dynamic param has higher priority than the global flag (see + should_redact_message_logging). When a caller explicitly opts back into + message logging per-request, NR must record content again.""" + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + logger = make_logger() + + kwargs = make_kwargs(messages=[{"role": "user", "content": "ok to log"}]) + kwargs["standard_callback_dynamic_params"] = { + "turn_off_message_logging": False, + } + response = make_response(content="response text") + + messages = logger._extract_all_messages( + kwargs, response, response_model="gpt-4", vendor="openai" + ) + + request_msg = next(m for m in messages if not m.get("is_response")) + response_msg = next(m for m in messages if m.get("is_response")) + assert request_msg["content"] == "ok to log" + assert response_msg["content"] == "response text" + + +class TestExtractAllMessagesTimestamps: + def setup_method(self): + self.logger = make_logger() + + def test_input_messages_get_start_time_timestamp(self): + kwargs = make_kwargs(messages=[{"role": "user", "content": "Hi"}]) + # make_kwargs sets start_time=1_000_000.0 and end_time=1_000_001.5 + response = make_response() + + messages = self.logger._extract_all_messages( + kwargs, response, response_model="gpt-4", vendor="openai" + ) + + input_msg = next(m for m in messages if not m.get("is_response")) + assert input_msg["timestamp"] == int(1_000_000.0 * 1000.0) + + def test_output_messages_get_end_time_timestamp(self): + kwargs = make_kwargs(messages=[{"role": "user", "content": "Hi"}]) + response = make_response() + + messages = self.logger._extract_all_messages( + kwargs, response, response_model="gpt-4", vendor="openai" + ) + + output_msg = next(m for m in messages if m.get("is_response")) + assert output_msg["timestamp"] == int(1_000_001.5 * 1000.0) + + def test_timestamp_forwarded_to_event_data(self): + logger = make_logger() + mock_app = MagicMock() + mock_app.enabled = True + + kwargs = make_kwargs( + traceparent="00-aabbccddeeff00112233445566778899-0011223344556677-01", + messages=[{"role": "user", "content": "Hi"}], + ) + response = make_response() + + with patch("newrelic.agent.application", return_value=mock_app): + logger._process_success(kwargs, response, start_time=1.0, end_time=2.5) + + calls = mock_app.record_custom_event.call_args_list + message_events = [ + c[0][1] for c in calls if c[0][0] == "LlmChatCompletionMessage" + ] + for event in message_events: + assert "timestamp" in event + + +# --------------------------------------------------------------------------- +# Streaming response handling +# --------------------------------------------------------------------------- + + +def make_streaming_response( + model="gpt-4", + response_id="chatcmpl-stream123", + content="Hello from streaming!", + finish_reason="stop", + prompt_tokens=8, + completion_tokens=15, +): + """Build a streaming-assembled response dict using 'delta' instead of 'message'.""" + return { + "id": response_id, + "model": model, + "choices": [ + { + "delta": {"role": "assistant", "content": content}, + "finish_reason": finish_reason, + } + ], + "usage": { + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "total_tokens": prompt_tokens + completion_tokens, + }, + } + + +class TestStreamingResponse: + """Verify graceful handling of streaming-assembled responses. + + When LiteLLM assembles a streaming response, some providers produce a + final choice dict with a 'delta' key instead of 'message'. The integration + must extract content from either key without raising. + """ + + def setup_method(self): + self.logger = make_logger() + + def test_extracts_content_from_delta_key(self): + kwargs = make_kwargs(messages=[{"role": "user", "content": "Hi"}]) + response = make_streaming_response(content="Streamed reply") + + messages = self.logger._extract_all_messages( + kwargs, response, response_model="gpt-4", vendor="openai" + ) + + response_msgs = [m for m in messages if m.get("is_response")] + assert len(response_msgs) == 1 + assert response_msgs[0]["content"] == "Streamed reply" + assert response_msgs[0]["role"] == "assistant" + + def test_streaming_response_records_summary_and_message_events(self): + mock_app = MagicMock() + mock_app.enabled = True + + kwargs = make_kwargs( + traceparent="00-aabbccddeeff00112233445566778899-0011223344556677-01", + messages=[{"role": "user", "content": "Hi"}], + ) + response = make_streaming_response( + response_id="chatcmpl-stream123", + content="Streamed reply", + finish_reason="stop", + prompt_tokens=8, + completion_tokens=15, + ) + + with patch("newrelic.agent.application", return_value=mock_app): + self.logger._process_success(kwargs, response, start_time=1.0, end_time=2.0) + + calls = mock_app.record_custom_event.call_args_list + event_types = [c[0][0] for c in calls] + assert "LlmChatCompletionSummary" in event_types + assert "LlmChatCompletionMessage" in event_types + + message_events = [ + c[0][1] for c in calls if c[0][0] == "LlmChatCompletionMessage" + ] + response_msg = next((e for e in message_events if e.get("is_response")), None) + assert response_msg is not None + assert response_msg["content"] == "Streamed reply" + + @pytest.mark.asyncio + async def test_async_log_success_event_streaming(self): + """async_log_success_event is the primary entry point for streaming calls.""" + mock_app = MagicMock() + mock_app.enabled = True + + kwargs = make_kwargs(messages=[{"role": "user", "content": "Hi"}]) + response = make_streaming_response() + + with patch("newrelic.agent.application", return_value=mock_app): + await self.logger.async_log_success_event( + kwargs, response, start_time=1.0, end_time=2.0 + ) + + calls = mock_app.record_custom_event.call_args_list + event_types = [c[0][0] for c in calls] + assert "LlmChatCompletionSummary" in event_types + assert "LlmChatCompletionMessage" in event_types + + def test_no_content_when_recording_disabled_streaming(self): + logger = make_logger(turn_off_message_logging=True) + kwargs = make_kwargs(messages=[{"role": "user", "content": "secret"}]) + response = make_streaming_response(content="also secret") + + messages = logger._extract_all_messages( + kwargs, response, response_model="gpt-4", vendor="openai" + ) + + for msg in messages: + assert "content" not in msg + + +# --------------------------------------------------------------------------- +# Explicit-None defensive tests +# --------------------------------------------------------------------------- + + +class TestExplicitNoneValues: + """Verify that explicitly None values in kwargs/response don't raise or silently drop events.""" + + def setup_method(self): + self.logger = make_logger() + + # _get_trace_context — chained dict lookups + def test_trace_context_litellm_params_none(self): + kwargs = make_kwargs() + kwargs["litellm_params"] = None + trace_id = self.logger._get_trace_context(kwargs) + assert trace_id is not None # falls back to UUID + + def test_trace_context_metadata_none(self): + kwargs = make_kwargs() + kwargs["litellm_params"] = {"metadata": None} + trace_id = self.logger._get_trace_context(kwargs) + assert trace_id is not None + + def test_trace_context_headers_none(self): + kwargs = make_kwargs() + kwargs["litellm_params"] = {"metadata": {"headers": None}} + trace_id = self.logger._get_trace_context(kwargs) + assert trace_id is not None + + # _get_request_params + def test_request_params_optional_params_none(self): + assert self.logger._get_request_params({"optional_params": None}) == {} + + # _get_model_names + def test_model_names_model_none_in_kwargs(self): + request_model, _ = self.logger._get_model_names( + {"model": None}, make_response() + ) + assert request_model == "unknown" + + def test_model_names_model_none_in_response(self): + response = make_response() + response["model"] = None + _, response_model = self.logger._get_model_names(make_kwargs(), response) + assert response_model == "gpt-4" # falls back to request_model from kwargs + + # _extract_all_messages + def test_extract_messages_messages_none(self): + kwargs = make_kwargs() + kwargs["messages"] = None + response = make_response() + messages = self.logger._extract_all_messages( + kwargs, response, response_model="gpt-4", vendor="openai" + ) + # No request messages, but response message should still be extracted + assert any(m.get("is_response") for m in messages) + + def test_extract_messages_choices_none(self): + kwargs = make_kwargs(messages=[{"role": "user", "content": "Hi"}]) + response = make_response() + response["choices"] = None + messages = self.logger._extract_all_messages( + kwargs, response, response_model="gpt-4", vendor="openai" + ) + # No response messages, but request message should still be extracted + assert any(not m.get("is_response") for m in messages) + + +# --------------------------------------------------------------------------- +# Helper edge cases +# --------------------------------------------------------------------------- + + +class TestExtractUsage: + def setup_method(self): + self.logger = make_logger() + + def test_missing_usage_returns_zeros(self): + response = {"id": "r1", "model": "gpt-4", "choices": []} + usage = self.logger._extract_usage(response) + assert usage == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0} + + def test_explicit_none_token_fields_return_zeros(self): + response = { + "usage": { + "prompt_tokens": None, + "completion_tokens": None, + "total_tokens": None, + } + } + usage = self.logger._extract_usage(response) + assert usage == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0} + + +class TestGetFinishReason: + def setup_method(self): + self.logger = make_logger() + + def test_returns_unknown_when_no_choices(self): + response = {"choices": []} + assert self.logger._get_finish_reason(response) == "unknown" + + def test_returns_unknown_when_choices_missing(self): + assert self.logger._get_finish_reason({}) == "unknown" + + def test_returns_unknown_when_finish_reason_explicitly_none(self): + response = {"choices": [{"finish_reason": None}]} + assert self.logger._get_finish_reason(response) == "unknown" + + +class TestToEpochMs: + def setup_method(self): + self.logger = make_logger() + + def test_float_passthrough(self): + assert self.logger._to_epoch_ms(1.0) == pytest.approx(1000.0) + + def test_datetime_converted(self): + dt = datetime(2024, 1, 1, 0, 0, 0, tzinfo=timezone.utc) + assert self.logger._to_epoch_ms(dt) == pytest.approx(dt.timestamp() * 1000.0) + + +class TestGetDuration: + def setup_method(self): + self.logger = make_logger() + + def test_uses_kwargs_value_when_present(self): + kwargs = {"llm_api_duration_ms": 750.0} + assert self.logger._get_duration(kwargs, 0.0, 1.0) == 750.0 + + def test_calculates_from_float_timestamps(self): + kwargs = {} + result = self.logger._get_duration(kwargs, 1.0, 2.5) + assert result == pytest.approx(1500.0) + + def test_calculates_from_datetime_timestamps(self): + kwargs = {} + start = datetime(2024, 1, 1, 0, 0, 0, tzinfo=timezone.utc) + end = datetime(2024, 1, 1, 0, 0, 1, 500000, tzinfo=timezone.utc) # +1.5s + result = self.logger._get_duration(kwargs, start, end) + assert result == pytest.approx(1500.0) + + def test_returns_none_when_nothing_available(self): + assert self.logger._get_duration({}, None, None) is None + + +class TestGetRequestParams: + def setup_method(self): + self.logger = make_logger() + + def test_includes_only_present_params(self): + kwargs = {"optional_params": {"temperature": 0.7}} + params = self.logger._get_request_params(kwargs) + assert params == {"temperature": 0.7} + assert "max_tokens" not in params + + def test_empty_when_no_optional_params(self): + assert self.logger._get_request_params({}) == {} + + +# --------------------------------------------------------------------------- +# _process_success — comprehensive happy-path +# --------------------------------------------------------------------------- + + +class TestProcessSuccess: + def test_records_summary_and_message_events(self): + logger = make_logger() + mock_app = MagicMock() + mock_app.enabled = True + + kwargs = make_kwargs( + traceparent="00-aabbccddeeff00112233445566778899-0011223344556677-01", + messages=[{"role": "user", "content": "Hello"}], + optional_params={"temperature": 0.5, "max_tokens": 100}, + ) + response = make_response( + response_id="chatcmpl-xyz", + content="Hi there!", + finish_reason="stop", + prompt_tokens=5, + completion_tokens=10, + ) + + with patch("newrelic.agent.application", return_value=mock_app): + logger._process_success(kwargs, response, start_time=1.0, end_time=2.5) + + calls = mock_app.record_custom_event.call_args_list + event_types = [c[0][0] for c in calls] + + assert "LlmChatCompletionSummary" in event_types + assert "LlmChatCompletionMessage" in event_types + + # Verify summary event fields + summary_data = next( + c[0][1] for c in calls if c[0][0] == "LlmChatCompletionSummary" + ) + assert summary_data["vendor"] == "openai" + assert summary_data["request.model"] == "gpt-4" + assert summary_data["response.model"] == "gpt-4" + assert summary_data["response.choices.finish_reason"] == "stop" + assert summary_data["response.usage.prompt_tokens"] == 5 + assert summary_data["response.usage.completion_tokens"] == 10 + assert summary_data["response.usage.total_tokens"] == 15 + assert summary_data["request.temperature"] == 0.5 + assert summary_data["request.max_tokens"] == 100 + assert summary_data["ingest_source"] == "litellm" + assert summary_data["trace_id"] == "aabbccddeeff00112233445566778899" + + # Verify message event id format: "{llm_response_id}-{sequence}" + message_events = [ + c[0][1] for c in calls if c[0][0] == "LlmChatCompletionMessage" + ] + assert any(e["id"].startswith("chatcmpl-xyz-") for e in message_events) + response_msg = next(e for e in message_events if e.get("is_response")) + assert response_msg["content"] == "Hi there!" + assert response_msg["role"] == "assistant" + + def test_skips_when_disabled(self): + logger = make_logger() + logger.enabled = False + + with patch("newrelic.agent.application") as mock_app: + logger._process_success(make_kwargs(), make_response()) + + mock_app.assert_not_called() + + +# --------------------------------------------------------------------------- +# _record_error_metric +# --------------------------------------------------------------------------- + + +class TestRecordErrorMetric: + def setup_method(self): + self.logger = make_logger() + + def test_calls_record_custom_metric(self): + mock_app = MagicMock() + mock_app.enabled = True + + with patch.object(self.logger, "_check_and_emit_periodic_metric"): + with patch("newrelic.agent.application", return_value=mock_app): + self.logger._record_error_metric() + + mock_app.record_custom_metric.assert_called_once_with("LLM/LiteLLM/Error", 1) + + def test_skips_when_app_disabled(self): + mock_app = MagicMock() + mock_app.enabled = False + + with patch.object(self.logger, "_check_and_emit_periodic_metric"): + with patch("newrelic.agent.application", return_value=mock_app): + self.logger._record_error_metric() + + mock_app.record_custom_metric.assert_not_called() + + def test_calls_check_and_emit_periodic_metric(self): + with patch.object( + self.logger, "_check_and_emit_periodic_metric" + ) as mock_periodic: + with patch("newrelic.agent.application", return_value=MagicMock()): + self.logger._record_error_metric() + + mock_periodic.assert_called_once() + + def test_skips_when_logger_disabled(self): + self.logger.enabled = False + with patch("newrelic.agent.application") as mock_app: + self.logger._record_error_metric() + mock_app.assert_not_called() + + def test_handles_exception(self): + with patch( + "newrelic.agent.application", side_effect=RuntimeError("agent down") + ): + self.logger._record_error_metric() # must not raise + + +# --------------------------------------------------------------------------- +# _emit_supportability_metric +# --------------------------------------------------------------------------- + + +class TestEmitSupportabilityMetric: + def setup_method(self): + self.logger = make_logger() + NewRelicLogger._last_metric_emission_time = 0.0 + + def test_records_metric_with_correct_name_and_value(self): + mock_app = MagicMock() + mock_app.enabled = True + with patch("newrelic.agent.application", return_value=mock_app): + with patch.object( + self.logger, "_get_litellm_version", return_value="1.80.0" + ): + self.logger._emit_supportability_metric() + mock_app.record_custom_metric.assert_called_once_with( + "Supportability/Python/ML/LiteLLM/1.80.0", 1 + ) + + def test_updates_last_emission_time(self): + mock_app = MagicMock() + mock_app.enabled = True + fake_now = 9_999_999.0 + with patch("newrelic.agent.application", return_value=mock_app): + with patch( + "litellm.integrations.newrelic.newrelic.time.time", + return_value=fake_now, + ): + self.logger._emit_supportability_metric() + assert NewRelicLogger._last_metric_emission_time == fake_now + + def test_skips_when_app_disabled(self): + mock_app = MagicMock() + mock_app.enabled = False + with patch("newrelic.agent.application", return_value=mock_app): + self.logger._emit_supportability_metric() + mock_app.record_custom_metric.assert_not_called() + # Timestamp is still updated to back off lock contention during registration. + assert NewRelicLogger._last_metric_emission_time != 0.0 + + def test_skips_when_no_app(self): + with patch("newrelic.agent.application", return_value=None): + self.logger._emit_supportability_metric() + # Timestamp is updated even when app is None to back off lock contention + # if the agent never starts or is slow to initialise. + assert NewRelicLogger._last_metric_emission_time != 0.0 + + def test_handles_exception(self): + with patch( + "newrelic.agent.application", side_effect=RuntimeError("agent down") + ): + self.logger._emit_supportability_metric() # must not raise + + +# --------------------------------------------------------------------------- +# _check_and_emit_periodic_metric +# --------------------------------------------------------------------------- + + +class TestCheckAndEmitPeriodicMetric: + def setup_method(self): + self.logger = make_logger() + NewRelicLogger._last_metric_emission_time = 0.0 + + def test_emits_on_first_call(self): + """_last_metric_emission_time starts at 0.0; any real time satisfies 27-hour window.""" + with patch.object(self.logger, "_emit_supportability_metric") as mock_emit: + with patch( + "litellm.integrations.newrelic.newrelic.time.time", + return_value=100_000.0, + ): + self.logger._check_and_emit_periodic_metric() + mock_emit.assert_called_once() + + def test_does_not_re_emit_within_27_hours(self): + recent = 1_000_000.0 + NewRelicLogger._last_metric_emission_time = recent + with patch.object(self.logger, "_emit_supportability_metric") as mock_emit: + with patch( + "litellm.integrations.newrelic.newrelic.time.time", + return_value=recent + 3600, # 1 hour later + ): + self.logger._check_and_emit_periodic_metric() + mock_emit.assert_not_called() + + def test_re_emits_after_27_hours(self): + old = 1_000_000.0 + NewRelicLogger._last_metric_emission_time = old + with patch.object(self.logger, "_emit_supportability_metric") as mock_emit: + with patch( + "litellm.integrations.newrelic.newrelic.time.time", + return_value=old + 97201, # 27 hours + 1 second + ): + self.logger._check_and_emit_periodic_metric() + mock_emit.assert_called_once() + + def test_boundary_exactly_27_hours_triggers_emission(self): + old = 1_000_000.0 + NewRelicLogger._last_metric_emission_time = old + with patch.object(self.logger, "_emit_supportability_metric") as mock_emit: + with patch( + "litellm.integrations.newrelic.newrelic.time.time", + return_value=old + 97200, + ): + self.logger._check_and_emit_periodic_metric() + mock_emit.assert_called_once() + + +# --------------------------------------------------------------------------- +# _get_litellm_version +# --------------------------------------------------------------------------- + + +class TestGetLitellmVersion: + def setup_method(self): + self.logger = make_logger() + + def test_returns_unknown_on_exception(self): + with patch("importlib.metadata.version", side_effect=Exception("no package")): + result = self.logger._get_litellm_version() + assert result == "unknown" + + +# --------------------------------------------------------------------------- +# _record_summary_event — disabled-app and exception paths +# --------------------------------------------------------------------------- + +_USAGE = {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15} + + +class TestRecordSummaryEvent: + def setup_method(self): + self.logger = make_logger() + + def _call(self, **kwargs): + self.logger._record_summary_event( + request_id="req-1", + trace_id="trace-abc", + request_model="gpt-4", + response_model="gpt-4", + vendor="openai", + finish_reason="stop", + num_messages=2, + usage=_USAGE, + **kwargs, + ) + + def test_skips_when_app_disabled(self): + mock_app = MagicMock() + mock_app.enabled = False + with patch("newrelic.agent.application", return_value=mock_app): + self._call() + mock_app.record_custom_event.assert_not_called() + + def test_handles_exception(self): + with patch( + "newrelic.agent.application", side_effect=RuntimeError("agent down") + ): + self._call() # must not raise + + +# --------------------------------------------------------------------------- +# _record_message_events — disabled-app and exception paths +# --------------------------------------------------------------------------- + +_MESSAGES = [ + {"role": "user", "sequence": 0, "response.model": "gpt-4", "vendor": "openai"} +] + + +class TestRecordMessageEvents: + def setup_method(self): + self.logger = make_logger() + + def _call(self): + self.logger._record_message_events( + request_id="req-1", + llm_response_id="resp-1", + trace_id="trace-abc", + messages=_MESSAGES, + ) + + def test_skips_when_app_disabled(self): + mock_app = MagicMock() + mock_app.enabled = False + with patch("newrelic.agent.application", return_value=mock_app): + self._call() + mock_app.record_custom_event.assert_not_called() + + def test_handles_exception(self): + with patch( + "newrelic.agent.application", side_effect=RuntimeError("agent down") + ): + self._call() # must not raise + + +# --------------------------------------------------------------------------- +# CustomLogger interface entry points +# --------------------------------------------------------------------------- + + +class TestLogSuccessEvent: + def test_delegates_to_process_success(self): + logger = make_logger() + with patch.object(logger, "_process_success") as mock_process: + logger.log_success_event(make_kwargs(), make_response(), 1.0, 2.0) + mock_process.assert_called_once() + + def test_exception_is_handled(self): + logger = make_logger() + with patch.object(logger, "_process_success", side_effect=RuntimeError("boom")): + logger.log_success_event(make_kwargs(), make_response(), 1.0, 2.0) + + @pytest.mark.asyncio + async def test_async_delegates_to_process_success(self): + logger = make_logger() + with patch.object(logger, "_process_success") as mock_process: + await logger.async_log_success_event( + make_kwargs(), make_response(), 1.0, 2.0 + ) + mock_process.assert_called_once() + + @pytest.mark.asyncio + async def test_async_exception_is_handled(self): + logger = make_logger() + with patch.object(logger, "_process_success", side_effect=RuntimeError("boom")): + await logger.async_log_success_event( + make_kwargs(), make_response(), 1.0, 2.0 + ) + + +class TestLogFailureEvent: + def test_sync_records_error_metric(self): + logger = make_logger() + with patch.object(logger, "_record_error_metric") as mock_metric: + logger.log_failure_event(make_kwargs(), None, 1.0, 2.0) + mock_metric.assert_called_once() + + def test_sync_exception_is_handled(self): + logger = make_logger() + with patch.object( + logger, "_record_error_metric", side_effect=RuntimeError("boom") + ): + logger.log_failure_event(make_kwargs(), None, 1.0, 2.0) + + @pytest.mark.asyncio + async def test_async_records_error_metric(self): + logger = make_logger() + with patch.object(logger, "_record_error_metric") as mock_metric: + await logger.async_log_failure_event(make_kwargs(), None, 1.0, 2.0) + mock_metric.assert_called_once() + + @pytest.mark.asyncio + async def test_async_exception_is_handled(self): + logger = make_logger() + with patch.object( + logger, "_record_error_metric", side_effect=RuntimeError("boom") + ): + await logger.async_log_failure_event(make_kwargs(), None, 1.0, 2.0) + + +# --------------------------------------------------------------------------- +# async_health_check +# --------------------------------------------------------------------------- + + +class TestAsyncHealthCheck: + @pytest.mark.asyncio + async def test_unhealthy_when_disabled(self): + logger = make_logger() + logger.enabled = False + result = await logger.async_health_check() + assert result["status"] == "unhealthy" + assert result["error_message"] is not None + + @pytest.mark.asyncio + async def test_healthy_when_app_enabled_records_test_event(self): + logger = make_logger() + mock_app = MagicMock() + mock_app.enabled = True + with patch("newrelic.agent.application", return_value=mock_app): + result = await logger.async_health_check() + assert result["status"] == "healthy" + assert result["error_message"] is None + + mock_app.record_custom_event.assert_called_once() + event_type, event_data = mock_app.record_custom_event.call_args[0] + assert event_type == "LiteLLMConnectionTest" + assert event_data["is_test_event"] is True + assert event_data["app_name"] == logger.app_name + assert event_data["source"] == "litellm-proxy" + assert isinstance(event_data["timestamp"], float) + + @pytest.mark.asyncio + async def test_unhealthy_when_app_disabled(self): + logger = make_logger() + mock_app = MagicMock() + mock_app.enabled = False + with patch("newrelic.agent.application", return_value=mock_app): + result = await logger.async_health_check() + assert result["status"] == "unhealthy" + assert result["error_message"] is not None + mock_app.record_custom_event.assert_not_called() + + @pytest.mark.asyncio + async def test_exception_returns_unhealthy(self): + logger = make_logger() + with patch( + "newrelic.agent.application", side_effect=RuntimeError("agent down") + ): + result = await logger.async_health_check() + assert result["status"] == "unhealthy" + assert "agent down" in result["error_message"] + + @pytest.mark.asyncio + async def test_record_custom_event_failure_returns_unhealthy(self): + logger = make_logger() + mock_app = MagicMock() + mock_app.enabled = True + mock_app.record_custom_event.side_effect = RuntimeError("intake unreachable") + with patch("newrelic.agent.application", return_value=mock_app): + result = await logger.async_health_check() + assert result["status"] == "unhealthy" + assert "intake unreachable" in result["error_message"] + + +# --------------------------------------------------------------------------- +# _extract_completion_id fallback chain +# --------------------------------------------------------------------------- + + +class TestExtractCompletionId: + def setup_method(self): + self.logger = make_logger() + + def test_uses_litellm_call_id_when_response_has_no_id(self): + result = self.logger._extract_completion_id( + kwargs={"litellm_call_id": "call-abc-123"}, + response_obj={}, + ) + assert result == "call-abc-123" + + def test_generates_uuid_when_neither_id_present(self): + result = self.logger._extract_completion_id(kwargs={}, response_obj={}) + # UUID4 hex-with-dashes is 36 chars; just confirm shape and uniqueness + assert isinstance(result, str) + assert len(result) == 36 + second = self.logger._extract_completion_id(kwargs={}, response_obj={}) + assert result != second + + +# --------------------------------------------------------------------------- +# StandardLoggingPayload preference across extractors +# --------------------------------------------------------------------------- + + +class TestStandardLoggingPayloadPreference: + """Each extractor that accepts a StandardLoggingPayload must prefer its + values over the raw kwargs/response fallbacks.""" + + def setup_method(self): + self.logger = make_logger() + + def test_trace_context_uses_slo_trace_id_when_no_traceparent(self): + kwargs = {"litellm_params": {"metadata": {"headers": {}}}} + trace_id = self.logger._get_trace_context( + kwargs, standard_logging_object=make_slo() + ) + assert trace_id == "slo-trace-abc" + + def test_vendor_from_slo(self): + # kwargs carries a different provider; SLO must win. + kwargs = {"litellm_params": {"custom_llm_provider": "kwargs-provider"}} + assert ( + self.logger._get_vendor(kwargs, standard_logging_object=make_slo()) + == "slo-provider" + ) + + def test_model_names_uses_slo_model(self): + request_model, _ = self.logger._get_model_names( + {"model": "kwargs-model"}, + make_response(model="response-model"), + standard_logging_object=make_slo(), + ) + assert request_model == "slo-model" + + def test_usage_from_slo_when_any_token_field_present(self): + # make_response defaults to 10/20/30 tokens; SLO sentinels are 100/200/300. + usage = self.logger._extract_usage( + make_response(), standard_logging_object=make_slo() + ) + assert usage == { + "prompt_tokens": 100, + "completion_tokens": 200, + "total_tokens": 300, + } + + def test_duration_from_slo_response_time_converted_to_ms(self): + # SLO response_time is 1.5 seconds; expected 1500.0 ms. + # Pass start/end that would compute a different value to prove SLO won. + duration = self.logger._get_duration( + kwargs={"llm_api_duration_ms": 9999.0}, + start_time=1.0, + end_time=2.0, + standard_logging_object=make_slo(), + ) + assert duration == 1500.0 + + def test_request_params_from_slo_model_parameters(self): + params = self.logger._get_request_params( + {"optional_params": {"temperature": 0.1}}, + standard_logging_object=make_slo(), + ) + assert params == {"temperature": 0.7, "max_tokens": 500} + + def test_extract_all_messages_sources_timestamps_and_messages_from_slo(self): + """Covers three SLO branches at once: startTime, endTime, and messages list.""" + kwargs = make_kwargs(messages=[{"role": "user", "content": "from-kwargs"}]) + messages = self.logger._extract_all_messages( + kwargs, + make_response(), + response_model="gpt-4", + vendor="openai", + standard_logging_object=make_slo(), + ) + + request = next(m for m in messages if not m.get("is_response")) + assert request["content"] == "from-slo" # SLO messages list wins + assert request["timestamp"] == int(2_000_000.0 * 1000.0) # SLO startTime + + response = next(m for m in messages if m.get("is_response")) + assert response["timestamp"] == int(2_000_001.5 * 1000.0) # SLO endTime diff --git a/tests/test_litellm/litellm_core_utils/test_duration_parser.py b/tests/test_litellm/litellm_core_utils/test_duration_parser.py index d95503665ec..3e4446c6672 100644 --- a/tests/test_litellm/litellm_core_utils/test_duration_parser.py +++ b/tests/test_litellm/litellm_core_utils/test_duration_parser.py @@ -34,6 +34,23 @@ class TestStandardizedResetTime(unittest.TestCase): custom_day_result = get_next_standardized_reset_time("3d", base_time, "UTC") self.assertEqual(custom_day_result, custom_day_expected) + def test_week_based_resets(self): + """Test week-based reset durations (1w, 2w). + 1w snaps to the next Monday at midnight (same as 7d). + 2w advances exactly 14 days from the current date at midnight. + """ + # 1w from a Wednesday -> next Monday (5 days away, not 7) + wednesday = datetime(2023, 5, 17, 15, 45, 0, tzinfo=timezone.utc) + weekly_expected = datetime(2023, 5, 22, 0, 0, 0, tzinfo=timezone.utc) + weekly_result = get_next_standardized_reset_time("1w", wednesday, "UTC") + self.assertEqual(weekly_result, weekly_expected) + + # 2w from a Wednesday -> exactly 14 days out (lands on a Wednesday, not Monday) + base_time = datetime(2023, 5, 17, 10, 30, 0, tzinfo=timezone.utc) + two_week_expected = datetime(2023, 5, 31, 0, 0, 0, tzinfo=timezone.utc) + two_week_result = get_next_standardized_reset_time("2w", base_time, "UTC") + self.assertEqual(two_week_result, two_week_expected) + def test_hour_minute_second_resets(self): """Test hour, minute, and second based reset durations""" # Base time: 2023-05-15 15:20:30 UTC (3:20:30 PM) diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 34edd6eccf3..228fb2dd984 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -3151,6 +3151,41 @@ def test_get_error_information_prefers_message_attribute_over_empty_str(): assert info["error_code"] == "401" +def _anthropic_messages_logging_obj(): + return LitellmLogging( + model="openai/my-local", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="anthropic_messages", + start_time=time.time(), + litellm_call_id="28595", + function_id="28595", + ) + + +def _responses_api_response_with_text(text="hello world"): + from openai.types.responses import ResponseOutputMessage, ResponseOutputText + + from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse + + return ResponsesAPIResponse( + id="resp-28595", + created_at=1700000000, + output=[ + ResponseOutputMessage( + id="msg-1", + type="message", + role="assistant", + status="completed", + content=[ + ResponseOutputText(annotations=[], text=text, type="output_text") + ], + ) + ], + usage=ResponseAPIUsage(input_tokens=11, output_tokens=7, total_tokens=18), + ) + + @pytest.mark.parametrize( "event_cls, event_type", [ @@ -3159,34 +3194,68 @@ def test_get_error_information_prefers_message_attribute_over_empty_str(): ("ResponseFailedEvent", "response.failed"), ], ) -def test_handle_anthropic_messages_response_logging_with_terminal_responses_api_events( +def test_handle_anthropic_messages_response_logging_translates_terminal_responses_api_event( event_cls, event_type ): - """Regression test for #28943: when anthropic_messages routes to OpenAI Responses - API and stream=True, success_handler receives a terminal ResponsesAPI event instead - of a ModelResponse. The handler must return the inner ResponsesAPIResponse rather - than crashing with AnthropicResponse.model_validate.""" + """Regression for #28595 / #28943. When anthropic_messages routes to the OpenAI + Responses backend and stream=True, success_handler receives a terminal Responses + API event. The handler must translate it to a ModelResponse whose choices carry + the assistant text, so the proxy UI Logs tab (which reads response.choices[0]) + renders the response content instead of "No response data available".""" import importlib openai_types = importlib.import_module("litellm.types.llms.openai") EventClass = getattr(openai_types, event_cls) - from litellm.types.llms.openai import ResponsesAPIResponse - logging_obj = LitellmLogging( - model="gpt-4o", - messages=[{"role": "user", "content": "hello"}], - stream=True, - call_type="anthropic_messages", - start_time=time.time(), - litellm_call_id="test-rce-123", - function_id="test-fn", - ) - - inner_response = ResponsesAPIResponse( - id="resp_test", created_at=1700000000, output=[] - ) + logging_obj = _anthropic_messages_logging_obj() + inner_response = _responses_api_response_with_text("hello world") event = EventClass(type=event_type, response=inner_response) result = logging_obj._handle_anthropic_messages_response_logging(result=event) - assert result is inner_response + assert isinstance(result, ModelResponse) + assert result.choices[0].message.content == "hello world" # type: ignore[union-attr] + assert result.usage.prompt_tokens == 11 # type: ignore[attr-defined] + assert result.usage.completion_tokens == 7 # type: ignore[attr-defined] + + +def test_handle_anthropic_messages_response_logging_translates_bare_responses_api_response(): + """Non-streaming bridge path: result is a bare ResponsesAPIResponse (no event wrap).""" + logging_obj = _anthropic_messages_logging_obj() + result = logging_obj._handle_anthropic_messages_response_logging( + result=_responses_api_response_with_text("hi there") + ) + + assert isinstance(result, ModelResponse) + assert result.choices[0].message.content == "hi there" # type: ignore[union-attr] + assert result.usage.total_tokens == 18 # type: ignore[attr-defined] + + +def test_handle_anthropic_messages_response_logging_passes_model_response_through(): + """Anthropic-native path already yields a ModelResponse; it must be returned unchanged.""" + logging_obj = _anthropic_messages_logging_obj() + model_response = ModelResponse() + assert ( + logging_obj._handle_anthropic_messages_response_logging(result=model_response) + is model_response + ) + + +def test_handle_anthropic_messages_response_logging_degrades_on_unparseable_responses_payload(): + """If the Responses translation raises (eg. empty output on an incomplete response), + the row must still land: a minimal ModelResponse with model + usage is returned.""" + from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse + + logging_obj = _anthropic_messages_logging_obj() + empty = ResponsesAPIResponse( + id="resp-empty", + created_at=1700000000, + output=[], + usage=ResponseAPIUsage(input_tokens=4, output_tokens=0, total_tokens=4), + ) + + result = logging_obj._handle_anthropic_messages_response_logging(result=empty) + + assert isinstance(result, ModelResponse) + assert result.model == "openai/my-local" + assert result.usage.prompt_tokens == 4 # type: ignore[attr-defined] diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index ca709e768e5..2c164f2169c 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -521,6 +521,290 @@ async def test_transcription_captured_in_backend_to_client(): assert logging_obj.model_call_details["messages"] == streaming.input_messages +@pytest.mark.asyncio +async def test_transcription_session_captures_usage_and_skips_response_create(): + """ + For a transcription-only session (session.type == "transcription", e.g. + gpt-realtime-whisper), the completed event's audio-duration usage must be + captured for cost and response.create must NOT be sent to the backend. + """ + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + + session_created = json.dumps( + { + "type": "session.created", + "session": { + "type": "transcription", + "audio": { + "input": {"transcription": {"model": "gpt-realtime-whisper"}} + }, + }, + } + ).encode() + completed = json.dumps( + { + "type": "conversation.item.input_audio_transcription.completed", + "transcript": "hello world", + "item_id": "item_1", + "usage": {"type": "duration", "seconds": 12.0}, + } + ).encode() + + backend_ws = MagicMock() + backend_ws.recv = AsyncMock( + side_effect=[session_created, completed, ConnectionClosed(None, None)] + ) + backend_ws.send = AsyncMock() + + logging_obj = MagicMock() + logging_obj.model_call_details = {} + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) + await streaming.backend_to_client_send_messages() + + assert streaming._is_transcription_session is True + + captured = [ + m + for m in streaming.messages + if m.get("type") == "conversation.item.input_audio_transcription.completed" + ] + assert len(captured) == 1, "completed usage event must be captured for cost" + assert captured[0]["usage"]["seconds"] == 12.0 + + # Transcript still forwarded to the client. + client_ws.send_text.assert_any_call(completed.decode()) + + # No response.create — transcription sessions have no assistant turn. + sent_to_backend = [ + json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args + ] + assert all( + e.get("type") != "response.create" for e in sent_to_backend + ), f"transcription session must not trigger response.create, got: {sent_to_backend}" + + +@pytest.mark.asyncio +async def test_non_transcription_completed_event_still_triggers_response_create(): + """ + Regression guard: a normal (non-transcription) session with no guardrails must + keep triggering response.create on a completed transcription event. + """ + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + + completed = json.dumps( + { + "type": "conversation.item.input_audio_transcription.completed", + "transcript": "hi", + "item_id": "item_1", + } + ).encode() + + backend_ws = MagicMock() + backend_ws.recv = AsyncMock(side_effect=[completed, ConnectionClosed(None, None)]) + backend_ws.send = AsyncMock() + + logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) + await streaming.backend_to_client_send_messages() + + assert streaming._is_transcription_session is False + sent_to_backend = [ + json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args + ] + assert any(e.get("type") == "response.create" for e in sent_to_backend) + + +def test_client_session_update_marks_transcription_session(): + """A client session.update with type=transcription flags the session.""" + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + assert streaming._is_transcription_session is False + streaming._collect_user_input_from_client_event( + json.dumps({"type": "session.update", "session": {"type": "transcription"}}) + ) + assert streaming._is_transcription_session is True + + +@pytest.mark.asyncio +async def test_transcription_session_update_enforces_authorized_flat_model(): + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + streaming = RealTimeStreaming( + MagicMock(), + backend_ws, + MagicMock(), + model="gpt-realtime-whisper", + force_transcription_model="gpt-realtime-whisper", + ) + + await streaming._send_to_backend( + json.dumps( + { + "type": "session.update", + "session": { + "type": "transcription", + "input_audio_transcription": { + "model": "restricted-transcription-model", + "language": "en", + }, + }, + } + ) + ) + + sent = json.loads(backend_ws.send.await_args.args[0]) + assert sent["session"]["input_audio_transcription"] == { + "model": "gpt-realtime-whisper", + "language": "en", + } + assert streaming._is_transcription_session is True + + +@pytest.mark.asyncio +async def test_transcription_session_update_enforces_authorized_nested_model(): + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + streaming = RealTimeStreaming( + MagicMock(), + backend_ws, + MagicMock(), + model="gpt-realtime-whisper", + force_transcription_model="gpt-realtime-whisper", + ) + + await streaming._send_to_backend( + json.dumps( + { + "type": "session.update", + "session": { + "type": "transcription", + "audio": { + "input": { + "transcription": { + "model": "restricted-transcription-model", + "prompt": "domain words", + }, + "format": {"type": "audio/pcm", "rate": 24000}, + } + }, + }, + } + ) + ) + + sent = json.loads(backend_ws.send.await_args.args[0]) + assert sent["session"]["audio"]["input"]["transcription"] == { + "model": "gpt-realtime-whisper", + "prompt": "domain words", + } + assert sent["session"]["audio"]["input"]["format"] == { + "type": "audio/pcm", + "rate": 24000, + } + assert streaming._is_transcription_session is True + + +@pytest.mark.asyncio +async def test_normal_realtime_session_keeps_nested_transcription_model(): + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + streaming = RealTimeStreaming( + MagicMock(), + backend_ws, + MagicMock(), + model="gpt-4o-realtime-preview", + ) + + await streaming._send_to_backend( + json.dumps( + { + "type": "session.update", + "session": { + "type": "realtime", + "audio": { + "input": { + "transcription": { + "model": "whisper-1", + "language": "en", + } + } + }, + }, + } + ) + ) + + sent = json.loads(backend_ws.send.await_args.args[0]) + assert sent["session"]["audio"]["input"]["transcription"] == { + "model": "whisper-1", + "language": "en", + } + assert streaming._is_transcription_session is False + + +def test_detect_transcription_session_from_backend_transcription_session_events(): + """Backend transcription_session.created/updated events flag the session.""" + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + assert streaming._is_transcription_session is False + streaming._detect_transcription_session_from_backend( + {"type": "transcription_session.created"} + ) + assert streaming._is_transcription_session is True + + streaming2 = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + streaming2._detect_transcription_session_from_backend( + {"type": "transcription_session.updated"} + ) + assert streaming2._is_transcription_session is True + + +def test_detect_transcription_session_from_backend_session_created_with_type(): + """Backend session.created with type=transcription flags the session.""" + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + streaming._detect_transcription_session_from_backend( + {"type": "session.created", "session": {"type": "transcription"}} + ) + assert streaming._is_transcription_session is True + + +def test_detect_transcription_session_from_backend_ignores_non_transcription(): + """Backend session.created without type=transcription does not flag the session.""" + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + streaming._detect_transcription_session_from_backend( + {"type": "session.created", "session": {"model": "gpt-4o-realtime-preview"}} + ) + assert streaming._is_transcription_session is False + + +def test_capture_transcription_usage_deduplicates_when_already_stored(): + """ + When the event is already in messages (logged via store_message), it must not + be appended a second time by _capture_transcription_usage. + """ + import litellm + + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + # Add the event type to the default logged list so _should_store_message returns True. + streaming.logged_real_time_event_types = [ + "conversation.item.input_audio_transcription.completed" + ] + event = { + "type": "conversation.item.input_audio_transcription.completed", + "usage": {"type": "duration", "seconds": 5.0}, + } + streaming.store_message(json.dumps(event)) + initial_count = len(streaming.messages) + streaming._capture_transcription_usage(event) + assert len(streaming.messages) == initial_count # no duplicate + + @pytest.mark.asyncio async def test_client_ack_caches_setup_to_prevent_duplicate_session_update_setup(): websocket = MagicMock() diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py new file mode 100644 index 00000000000..8bc39a6d85e --- /dev/null +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py @@ -0,0 +1,312 @@ +""" +Regression tests for issue #30014. + +When LiteLLM proxies ``client -> /v1/messages -> /v1/chat/completions`` and a +streaming chunk both *triggers* a new Anthropic content block (its type differs +from the active block) and *carries* the first delta of that new block, the +trigger chunk's delta must be re-emitted as a ``content_block_delta``. + +The synthesized ``content_block_start`` always carries an empty body, so before +the fix the first non-empty ``text_delta`` of every transitioned block was +silently dropped — e.g. text resuming after a tool call started from the second +token ("The weather is nice." was lost, "Hi" rendered as ""). Bundled +``input_json_delta`` tool arguments were already preserved and must stay +preserved, and empty trigger deltas must not produce spurious events. +""" + +import os +import sys +from typing import List, Optional +from unittest.mock import MagicMock + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import ( + AnthropicStreamWrapper, +) +from litellm.types.utils import ( + ChatCompletionDeltaToolCall, + Delta, + Function, + StreamingChoices, +) + + +def _make_chunk(delta: Delta, finish_reason: Optional[str] = None) -> MagicMock: + chunk = MagicMock() + chunk.choices = [ + StreamingChoices( + finish_reason=finish_reason, + index=0, + delta=delta, + logprobs=None, + ) + ] + chunk.usage = None + chunk._hidden_params = {} + return chunk + + +def _tool_chunk( + call_id: str, name: Optional[str], arguments: Optional[str] +) -> MagicMock: + return _make_chunk( + Delta( + content=None, + tool_calls=[ + ChatCompletionDeltaToolCall( + id=call_id, + function=Function(name=name, arguments=arguments), + type="function", + index=0, + ) + ], + ) + ) + + +class _AsyncStream: + def __init__(self, items: List[MagicMock]): + self._it = iter(items) + + def __aiter__(self): + return self + + async def __anext__(self): + try: + return next(self._it) + except StopIteration: + raise StopAsyncIteration + + +def _drain_sync(wrapper: AnthropicStreamWrapper) -> List[dict]: + return list(wrapper) + + +async def _drain_async(wrapper: AnthropicStreamWrapper) -> List[dict]: + return [event async for event in wrapper] + + +def _text_deltas(events: List[dict]) -> List[str]: + return [ + e["delta"]["text"] + for e in events + if e.get("type") == "content_block_delta" + and e["delta"].get("type") == "text_delta" + ] + + +def _input_json_deltas(events: List[dict]) -> List[str]: + return [ + e["delta"]["partial_json"] + for e in events + if e.get("type") == "content_block_delta" + and e["delta"].get("type") == "input_json_delta" + ] + + +def test_first_text_delta_after_tool_use_is_not_dropped_sync(): + """A tool_use -> text transition (text resuming after a tool call) carries + the resumed text's first token in the trigger chunk. Without the fix it was + dropped, so "The weather is nice." vanished and the answer began at " Bye.". + """ + chunks = [ + _make_chunk(Delta(content="Let me check.")), + _tool_chunk("call_1", "get_weather", '{"city":'), + _tool_chunk("call_1", None, ' "NY"}'), + _make_chunk(Delta(content="The weather is nice.")), + _make_chunk(Delta(content=" Bye.")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x") + events = _drain_sync(wrapper) + + assert _input_json_deltas(events) == ['{"city":', ' "NY"}'] + assert _text_deltas(events) == [ + "Let me check.", + "The weather is nice.", + " Bye.", + ] + + +@pytest.mark.asyncio +async def test_first_text_delta_after_tool_use_is_not_dropped_async(): + """Async path mirrors the sync regression — the proxy serves the async + iterator, so it must preserve the first resumed text delta too. + """ + chunks = [ + _make_chunk(Delta(content="Let me check.")), + _tool_chunk("call_1", "get_weather", '{"city": "NY"}'), + _make_chunk(Delta(content="The weather is nice.")), + _make_chunk(Delta(content=" Bye.")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + wrapper = AnthropicStreamWrapper( + completion_stream=_AsyncStream(chunks), model="claude-x" + ) + events = await _drain_async(wrapper) + + assert _input_json_deltas(events) == ['{"city": "NY"}'] + assert _text_deltas(events) == [ + "Let me check.", + "The weather is nice.", + " Bye.", + ] + + +def test_single_first_text_token_after_tool_use_preserved_sync(): + """Minimal reproduction of the issue's example: a single short text token + ("Hi") resuming after a tool call. Without the fix the whole answer is + dropped because its only delta sits in the transition trigger chunk. + """ + chunks = [ + _tool_chunk("call_1", "get_weather", '{"city": "NY"}'), + _make_chunk(Delta(content="Hi")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x") + events = _drain_sync(wrapper) + + assert _text_deltas(events) == ["Hi"] + + +def test_multiple_text_deltas_after_tool_use_preserved_sync(): + """Multiple-delta edge case: only the *first* text delta sits in the + transition trigger chunk; the rest stream normally. All of them — leading + one included — must reach the client in order. + """ + chunks = [ + _tool_chunk("call_1", "get_weather", '{"city": "NY"}'), + _make_chunk(Delta(content="Hi")), + _make_chunk(Delta(content=", how ")), + _make_chunk(Delta(content="can I help ")), + _make_chunk(Delta(content="you?")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x") + events = _drain_sync(wrapper) + + assert _text_deltas(events) == ["Hi", ", how ", "can I help ", "you?"] + assert "".join(_text_deltas(events)) == "Hi, how can I help you?" + + +def test_empty_trigger_delta_is_not_re_emitted_sync(): + """A transition whose trigger chunk carries no content (empty text) must + NOT produce a spurious empty ``content_block_delta`` — only the synthesized + ``content_block_start`` is emitted for the new block. Here a ``tool_use -> + text`` transition is triggered by an empty-content chunk; the re-emit guard + must reject it so the new text block opens without a leading empty delta. + """ + chunks = [ + _tool_chunk("call_1", "get_weather", '{"city": "NY"}'), + # tool_use -> text transition triggered by an empty content chunk; the + # real text arrives in the following chunk. + _make_chunk(Delta(content="")), + _make_chunk(Delta(content="real text")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x") + events = _drain_sync(wrapper) + + # No empty-string text_delta should be present. + assert "" not in _text_deltas(events) + assert "".join(_text_deltas(events)) == "real text" + + +def test_bundled_tool_args_on_transition_still_preserved_sync(): + """Existing behavior guard: when the trigger chunk that opens a tool_use + block also carries arguments (xAI/Gemini style), the ``input_json_delta`` + must still be emitted after ``content_block_start``. + """ + chunks = [ + _make_chunk(Delta(content="Calling a tool.")), + _tool_chunk("call_1", "get_weather", '{"city": "NY"}'), + _make_chunk(Delta(content=None), finish_reason="tool_calls"), + ] + wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x") + events = _drain_sync(wrapper) + + assert _text_deltas(events) == ["Calling a tool."] + assert _input_json_deltas(events) == ['{"city": "NY"}'] + + +@pytest.mark.parametrize( + "processed_chunk, expected", + [ + # Non-empty deltas of every type must be re-emitted. + ( + { + "type": "content_block_delta", + "delta": {"type": "text_delta", "text": "x"}, + }, + True, + ), + ( + { + "type": "content_block_delta", + "delta": {"type": "input_json_delta", "partial_json": "{}"}, + }, + True, + ), + ( + { + "type": "content_block_delta", + "delta": {"type": "thinking_delta", "thinking": "t"}, + }, + True, + ), + ( + { + "type": "content_block_delta", + "delta": {"type": "signature_delta", "signature": "s"}, + }, + True, + ), + # Empty deltas must NOT be re-emitted (no spurious events). + ( + { + "type": "content_block_delta", + "delta": {"type": "text_delta", "text": ""}, + }, + False, + ), + ( + { + "type": "content_block_delta", + "delta": {"type": "input_json_delta", "partial_json": ""}, + }, + False, + ), + ( + { + "type": "content_block_delta", + "delta": {"type": "thinking_delta", "thinking": ""}, + }, + False, + ), + ( + { + "type": "content_block_delta", + "delta": {"type": "signature_delta", "signature": ""}, + }, + False, + ), + # Unknown delta type / non-content_block_delta / malformed delta. + ( + {"type": "content_block_delta", "delta": {"type": "other_delta"}}, + False, + ), + ({"type": "message_delta", "delta": {"stop_reason": "stop"}}, False), + ({"type": "content_block_delta", "delta": None}, False), + ], +) +def test_trigger_delta_has_content_branches(processed_chunk, expected): + """Directly exercise the re-emit predicate across all delta types and the + empty/malformed guards, so the helper's behavior is pinned independently of + upstream chunk-translation details. + """ + assert ( + AnthropicStreamWrapper._trigger_delta_has_content(processed_chunk) is expected + ) diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py index 1d25d719384..e44413cf837 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py @@ -278,7 +278,8 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text(): "content_block_delta", # {"city": "content_block_delta", # "NY"} "content_block_stop", # End of first tool_use content block - "content_block_start", # "The weather is nice today" + "content_block_start", # "The weather is nice today" text block + "content_block_delta", # "The weather is nice today." text_delta "content_block_stop", "content_block_start", # Start of second tool_use content block "content_block_delta", # {"city": @@ -288,7 +289,8 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text(): "content_block_delta", # {"city": "content_block_delta", # " CHI"} "content_block_stop", # End of third tool_use content block - "content_block_start", # "The weather is not so nice today" + "content_block_start", # "The weather is not so nice today" text block + "content_block_delta", # "The weather is not so nice today." text_delta "content_block_stop", "message_delta", # Stop reason with merged usage "message_stop", # Final message stop @@ -296,6 +298,20 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text(): assert expected_types == chunk_types + # Regression: the first (and only) text delta of each text block sits in + # the chunk that *triggered* the tool_use -> text transition. It must be + # re-emitted as a content_block_delta instead of being silently dropped. + text_deltas = [ + chunk["delta"]["text"] + for chunk in chunks + if chunk.get("type") == "content_block_delta" + and chunk["delta"].get("type") == "text_delta" + ] + assert text_deltas == [ + "The weather is nice today.", + "The weather is not so nice today.", + ] + get_weather_calls = 0 for chunk in chunks: diff --git a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py index 41d301c5d5f..4638bc4df0f 100644 --- a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py +++ b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py @@ -147,6 +147,103 @@ async def test_construct_url_ga_protocol(): assert "deployment" not in url +@pytest.mark.asyncio +async def test_construct_url_forwards_transcription_intent_ga(): + """ + Transcription sessions connect with intent=transcription. The Azure handler + must forward that query param so gpt-realtime-whisper opens a transcription + session instead of a normal realtime session. + """ + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + handler = AzureOpenAIRealtime() + url = handler._construct_url( + api_base="https://my-endpoint.openai.azure.com", + model="gpt-realtime-whisper", + api_version="2025-04-01-preview", + realtime_protocol="GA", + query_params={"model": "gpt-realtime-whisper", "intent": "transcription"}, + ) + + assert "/openai/v1/realtime?" in url + assert "intent=transcription" in url + assert "model=" not in url + + +@pytest.mark.asyncio +async def test_construct_url_forwards_transcription_intent_ga_without_model_query(): + """ + OpenAI-compatible transcription clients may connect with only + intent=transcription and send the transcription model in session.update. + Preserve that query shape instead of forcing model= into the upstream URL. + """ + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + handler = AzureOpenAIRealtime() + url = handler._construct_url( + api_base="https://my-endpoint.openai.azure.com", + model="gpt-realtime-whisper", + api_version="2025-04-01-preview", + realtime_protocol="GA", + query_params={"intent": "transcription"}, + ) + + assert url == ( + "wss://my-endpoint.openai.azure.com/openai/v1/realtime" + "?intent=transcription" + ) + + +@pytest.mark.asyncio +async def test_construct_url_forwards_transcription_intent_beta(): + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + handler = AzureOpenAIRealtime() + url = handler._construct_url( + api_base="https://my-endpoint.openai.azure.com", + model="whisper-deploy", + api_version="2024-10-01-preview", + query_params={"intent": "transcription"}, + ) + + assert "/openai/realtime?" in url + assert "deployment=whisper-deploy" in url + assert "intent=transcription" in url + + +@pytest.mark.asyncio +async def test_construct_url_encodes_intent_value(): + """A crafted intent value must be URL-encoded, not injected as raw query params.""" + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + handler = AzureOpenAIRealtime() + url = handler._construct_url( + api_base="https://my-endpoint.openai.azure.com", + model="gpt-realtime-whisper", + api_version="2025-04-01-preview", + realtime_protocol="GA", + query_params={"intent": "transcription&foo=bar"}, + ) + assert "intent=transcription%26foo%3Dbar" in url + assert "&foo=bar" not in url + + +@pytest.mark.asyncio +async def test_construct_url_no_intent_when_absent(): + """No intent param leaks into the URL when not provided.""" + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + handler = AzureOpenAIRealtime() + url = handler._construct_url( + api_base="https://my-endpoint.openai.azure.com", + model="gpt-4o-realtime-preview", + api_version="2024-10-01-preview", + realtime_protocol="GA", + query_params={"model": "gpt-4o-realtime-preview"}, + ) + assert "intent=" not in url + + @pytest.mark.asyncio async def test_construct_url_v1_protocol(): """ @@ -368,6 +465,45 @@ async def test_realtime_protocol_from_litellm_params(): assert litellm_params.get("realtime_protocol") == "GA" +@pytest.mark.asyncio +async def test_arealtime_transcription_intent_defaults_to_ga(monkeypatch): + """ + Azure gpt-realtime-whisper transcription connects on the GA /openai/v1/realtime + path. If the DB model lacks realtime_protocol, infer GA from intent=transcription. + """ + from litellm.realtime_api import main as realtime_main + + mock_async_realtime = AsyncMock() + monkeypatch.setattr( + realtime_main, + "azure_realtime", + MagicMock(async_realtime=mock_async_realtime), + ) + + def fake_get_llm_provider(model, api_base=None, api_key=None): + return ( + "gpt-realtime-whisper", + "azure", + "test-key", + "https://my-endpoint.openai.azure.com", + ) + + monkeypatch.setattr(realtime_main, "get_llm_provider", fake_get_llm_provider) + + await realtime_main._arealtime( + model="azure/gpt-realtime-whisper", + websocket=MagicMock(), + api_key="test-key", + api_version="2025-04-01-preview", + query_params={"intent": "transcription"}, + litellm_logging_obj=MagicMock(), + ) + + called_kwargs = mock_async_realtime.call_args.kwargs + assert called_kwargs["realtime_protocol"] == "GA" + assert called_kwargs["query_params"] == {"intent": "transcription"} + + @pytest.mark.asyncio async def test_async_realtime_default_maintains_backwards_compatibility(): """ diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index 5c83f8b34f7..d940f9f47a6 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -5386,3 +5386,45 @@ def test_converse_top_k_zero_forwarded_on_models_that_accept_it(): ) assert result["additionalModelRequestFields"]["top_k"] == 0 + + +@pytest.mark.asyncio +async def test_grounding_source_and_query_rendered_as_text(): + """grounding_source / query content blocks must render as plain text on the + generate path (the model needs to see the RAG context + question). The bedrock + converse dispatch silently drops unrecognised content types, so these would + otherwise vanish from the prompt.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + BedrockConverseMessagesProcessor, + _bedrock_converse_messages_pt, + ) + + messages = [ + { + "role": "user", + "content": [ + {"type": "grounding_source", "text": "Tokyo is the capital of Japan."}, + {"type": "query", "text": "What is the capital of Japan?"}, + ], + } + ] + + result = _bedrock_converse_messages_pt( + messages=messages, + model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + llm_provider="bedrock_converse", + ) + async_result = ( + await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + llm_provider="bedrock_converse", + ) + ) + + assert result == async_result + assert len(result) == 1 + assert result[0]["role"] == "user" + user_content = result[0]["content"] + assert {"text": "Tokyo is the capital of Japan."} in user_content + assert {"text": "What is the capital of Japan?"} in user_content diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 7321abcee46..f7d445d0788 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -81,6 +81,77 @@ def test_prepare_fake_stream_request(): assert result_data["messages"] == [{"role": "user", "content": "Hello"}] +def test_response_api_handler_streams_when_provider_transform_adds_stream(): + handler = BaseLLMHTTPHandler() + config = Mock() + config.validate_environment.return_value = {} + config.get_complete_url.return_value = "https://chatgpt.example.com/responses" + config.transform_responses_api_request.return_value = { + "model": "gpt-5.3-codex", + "input": "hi", + "stream": True, + } + config.sign_request.return_value = ({}, None) + client = HTTPHandler(client=httpx.Client()) + client.post = Mock( + return_value=httpx.Response( + 200, + request=httpx.Request("POST", "https://chatgpt.example.com/responses"), + ) + ) + logging_obj = Mock() + + handler.response_api_handler( + model="gpt-5.3-codex", + input="hi", + responses_api_provider_config=config, + response_api_optional_request_params={}, + custom_llm_provider="chatgpt", + litellm_params=GenericLiteLLMParams(), + logging_obj=logging_obj, + client=client, + ) + + assert client.post.call_args.kwargs["stream"] is True + assert client.post.call_args.kwargs["json"]["stream"] is True + + +@pytest.mark.asyncio +async def test_async_response_api_handler_streams_when_provider_transform_adds_stream(): + handler = BaseLLMHTTPHandler() + config = Mock() + config.validate_environment.return_value = {} + config.get_complete_url.return_value = "https://chatgpt.example.com/responses" + config.transform_responses_api_request.return_value = { + "model": "gpt-5.3-codex", + "input": "hi", + "stream": True, + } + config.sign_request.return_value = ({}, None) + client = AsyncHTTPHandler() + client.post = AsyncMock( + return_value=httpx.Response( + 200, + request=httpx.Request("POST", "https://chatgpt.example.com/responses"), + ) + ) + logging_obj = Mock() + + await handler.async_response_api_handler( + model="gpt-5.3-codex", + input="hi", + responses_api_provider_config=config, + response_api_optional_request_params={}, + custom_llm_provider="chatgpt", + litellm_params=GenericLiteLLMParams(), + logging_obj=logging_obj, + client=client, + ) + + assert client.post.call_args.kwargs["stream"] is True + assert client.post.call_args.kwargs["json"]["stream"] is True + + def test_get_agentic_loop_settings_defaults_and_overrides(): handler = BaseLLMHTTPHandler() diff --git a/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py b/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py index 17373f24a97..174efceb499 100644 --- a/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py +++ b/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py @@ -585,3 +585,221 @@ class TestGithubCopilotResponsesAPIRouting: provider=LlmProviders.GITHUB_COPILOT, ) assert isinstance(config, GithubCopilotResponsesAPIConfig) + + +class TestGithubCopilotReasoningStreamItemIdNormalization: + """GitHub Copilot's native /responses stream tags every reasoning-summary + event with a different item_id (and the reasoning output_item.added / + output_item.done ids also differ). Strict clients (Vercel ai-sdk) key + reasoning state by item_id and crash when a summary delta references an + unregistered id. The config normalizes every reasoning event in an + output_index group to the id from its output_item.added.""" + + def _config(self): + with patch( + "litellm.llms.github_copilot.responses.transformation.Authenticator" + ): + return GithubCopilotResponsesAPIConfig() + + def _transform(self, config, chunk): + return config.transform_streaming_response( + model="github_copilot/gpt-5.5", + parsed_chunk=chunk, + logging_obj=MagicMock(), + ) + + def test_summary_events_normalized_to_output_item_added_id(self): + config = self._config() + + self._transform( + config, + { + "type": "response.output_item.added", + "output_index": 0, + "item": {"id": "stable_rs_id", "type": "reasoning"}, + }, + ) + + summary_chunks = [ + { + "type": "response.reasoning_summary_part.added", + "output_index": 0, + "summary_index": 0, + "item_id": "bad_part_added", + "part": {"type": "summary_text", "text": ""}, + }, + { + "type": "response.reasoning_summary_text.delta", + "output_index": 0, + "summary_index": 0, + "item_id": "bad_delta_1", + "delta": "Hello", + }, + { + "type": "response.reasoning_summary_text.delta", + "output_index": 0, + "summary_index": 0, + "item_id": "bad_delta_2", + "delta": " world", + }, + { + "type": "response.reasoning_summary_text.done", + "output_index": 0, + "summary_index": 0, + "item_id": "bad_text_done", + "text": "Hello world", + }, + { + "type": "response.reasoning_summary_part.done", + "output_index": 0, + "summary_index": 0, + "item_id": "bad_part_done", + "part": {"type": "summary_text", "text": "Hello world"}, + }, + ] + for chunk in summary_chunks: + event = self._transform(config, chunk) + assert event.item_id == "stable_rs_id" + + def test_reasoning_output_item_done_normalized_to_added_id(self): + config = self._config() + self._transform( + config, + { + "type": "response.output_item.added", + "output_index": 0, + "item": {"id": "stable_rs_id", "type": "reasoning"}, + }, + ) + event = self._transform( + config, + { + "type": "response.output_item.done", + "output_index": 0, + "item": { + "id": "different_done_id", + "type": "reasoning", + "encrypted_content": "ENC", + }, + }, + ) + assert event.item.id == "stable_rs_id" + + def test_interleaved_message_item_does_not_corrupt_mapping(self): + config = self._config() + self._transform( + config, + { + "type": "response.output_item.added", + "output_index": 0, + "item": {"id": "stable_rs_id", "type": "reasoning"}, + }, + ) + self._transform( + config, + { + "type": "response.output_item.added", + "output_index": 1, + "item": {"id": "msg_id", "type": "message"}, + }, + ) + event = self._transform( + config, + { + "type": "response.reasoning_summary_text.delta", + "output_index": 0, + "summary_index": 0, + "item_id": "bad_delta", + "delta": "x", + }, + ) + assert event.item_id == "stable_rs_id" + + def test_event_without_registered_item_passes_through_unchanged(self): + config = self._config() + event = self._transform( + config, + { + "type": "response.output_text.delta", + "output_index": 0, + "item_id": "msg_native_id", + "delta": "hi", + }, + ) + assert event.item_id == "msg_native_id" + + def test_message_text_events_normalized_to_output_item_added_id(self): + config = self._config() + self._transform( + config, + { + "type": "response.output_item.added", + "output_index": 0, + "item": {"id": "stable_msg_id", "type": "message"}, + }, + ) + for chunk in [ + { + "type": "response.content_part.added", + "output_index": 0, + "content_index": 0, + "item_id": "bad_cp_added", + "part": {"type": "output_text", "text": ""}, + }, + { + "type": "response.output_text.delta", + "output_index": 0, + "content_index": 0, + "item_id": "bad_text_delta", + "delta": "Paris", + }, + { + "type": "response.output_text.done", + "output_index": 0, + "content_index": 0, + "item_id": "bad_text_done", + "text": "Paris", + }, + ]: + event = self._transform(config, chunk) + assert event.item_id == "stable_msg_id" + + def test_event_without_output_index_passes_through_unchanged(self): + config = self._config() + event = self._transform( + config, + { + "type": "response.completed", + "response": {"id": "resp_1", "status": "completed", "output": []}, + }, + ) + assert event.type == "response.completed" + + def test_normalization_continues_after_a_terminal_event(self): + config = self._config() + self._transform( + config, + { + "type": "response.output_item.added", + "output_index": 0, + "item": {"id": "stable_rs_id", "type": "reasoning"}, + }, + ) + self._transform( + config, + { + "type": "response.completed", + "response": {"id": "resp_1", "status": "completed", "output": []}, + }, + ) + event = self._transform( + config, + { + "type": "response.reasoning_summary_text.delta", + "output_index": 0, + "summary_index": 0, + "item_id": "bad_delta", + "delta": "x", + }, + ) + assert event.item_id == "stable_rs_id" diff --git a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py index e0911e1ef31..53c9e4b207c 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py @@ -11,6 +11,7 @@ import litellm sys.path.insert(0, os.path.abspath("../../../../..")) from litellm import ModelResponse +from litellm.constants import DEFAULT_OCI_CHAT_MAX_TOKENS from litellm.llms.oci.chat.transformation import ( OCIChatConfig, OCIRequestWrapper, @@ -104,6 +105,7 @@ class TestOCIChatConfig: "chatRequest": { "apiFormat": "GENERIC", "isStream": False, + "maxTokens": DEFAULT_OCI_CHAT_MAX_TOKENS, "messages": [ { "role": "USER", @@ -362,6 +364,137 @@ class TestOCIChatConfig: rf = transformed_request["chatRequest"]["responseFormat"] assert rf["type"] == "JSON_OBJECT" + def test_transform_request_response_format_json_schema_generic(self): + """A GENERIC json_schema must become OCI's JSON_SCHEMA shape with the + OpenAI ``strict`` key renamed to ``isStrict``. + + OCI's ResponseJsonSchema rejects ``strict`` (and any other extra key) + with HTTP 400 "Please pass in correct format of request", so the raw + OpenAI body must not be forwarded. + """ + config = OCIChatConfig() + optional_params = { + "oci_compartment_id": TEST_COMPARTMENT_ID, + "response_format": { + "type": "json_schema", + "json_schema": { + "name": "judgment", + "description": "a score and rationale", + "strict": True, + "schema": { + "type": "object", + "properties": {"score": {"type": "integer"}}, + "required": ["score"], + }, + }, + }, + } + transformed_request = config.transform_request( + model=TEST_MODEL_NAME, # xai.grok-4 -> GENERIC + messages=TEST_MESSAGES, # type: ignore + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + rf = transformed_request["chatRequest"]["responseFormat"] + assert rf["type"] == "JSON_SCHEMA" + assert "strict" not in rf["jsonSchema"] + assert rf["jsonSchema"]["isStrict"] is True + assert rf["jsonSchema"]["name"] == "judgment" + assert rf["jsonSchema"]["description"] == "a score and rationale" + assert rf["jsonSchema"]["schema"]["properties"]["score"]["type"] == "integer" + + def test_transform_request_response_format_json_schema_generic_no_strict(self): + """A GENERIC json_schema without ``strict`` must omit ``isStrict``.""" + config = OCIChatConfig() + optional_params = { + "oci_compartment_id": TEST_COMPARTMENT_ID, + "response_format": { + "type": "json_schema", + "json_schema": {"name": "j", "schema": {"type": "object"}}, + }, + } + transformed_request = config.transform_request( + model=TEST_MODEL_NAME, + messages=TEST_MESSAGES, # type: ignore + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + rf = transformed_request["chatRequest"]["responseFormat"] + assert rf["type"] == "JSON_SCHEMA" + assert "isStrict" not in rf["jsonSchema"] + + def test_transform_request_response_format_json_schema_cohere(self): + """A Cohere json_schema must fold the schema onto JSON_OBJECT. + + OCI Cohere has no JSON_SCHEMA type; sending one yields HTTP 400. + """ + config = OCIChatConfig() + optional_params = { + "oci_compartment_id": TEST_COMPARTMENT_ID, + "response_format": { + "type": "json_schema", + "json_schema": { + "name": "judgment", + "strict": True, + "schema": { + "type": "object", + "properties": {"score": {"type": "integer"}}, + }, + }, + }, + } + transformed_request = config.transform_request( + model="cohere.command-latest", + messages=TEST_MESSAGES, # type: ignore + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + rf = transformed_request["chatRequest"]["responseFormat"] + assert rf["type"] == "JSON_OBJECT" + assert "jsonSchema" not in rf + assert rf["schema"]["properties"]["score"]["type"] == "integer" + + def test_transform_request_response_format_cohere_json_object(self): + """Cohere json_object without a schema stays a bare JSON_OBJECT.""" + config = OCIChatConfig() + optional_params = { + "oci_compartment_id": TEST_COMPARTMENT_ID, + "response_format": {"type": "json_object"}, + } + transformed_request = config.transform_request( + model="cohere.command-latest", + messages=TEST_MESSAGES, # type: ignore + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + rf = transformed_request["chatRequest"]["responseFormat"] + assert rf == {"type": "JSON_OBJECT"} + + def test_transform_request_json_schema_without_body_raises_generic(self): + """A GENERIC json_schema with no ``json_schema`` body must raise an early + 400, not silently emit {"type": "JSON_SCHEMA"} (which OCI rejects).""" + from litellm.llms.oci.common_utils import OCIError + + config = OCIChatConfig() + optional_params = { + "oci_compartment_id": TEST_COMPARTMENT_ID, + "response_format": {"type": "json_schema"}, + } + with pytest.raises(OCIError) as exc_info: + config.transform_request( + model=TEST_MODEL_NAME, # GENERIC + messages=TEST_MESSAGES, # type: ignore + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + assert exc_info.value.status_code == 400 + assert "json_schema" in str(exc_info.value) + def test_transform_response_without_token_details(self): """ Tests that responses missing completionTokensDetails and promptTokensDetails @@ -956,6 +1089,44 @@ class TestOCICohereParamMapping: assert result.get("temperature") == 0.5 +class TestOCIDefaultMaxTokens: + """Regression for OCI's tiny server-side token cap (~20 tokens), which + silently truncated responses mid-string whenever the caller omitted + max_tokens (MLflow judges never send it, so their JSON came back cut off). + transform_request injects DEFAULT_OCI_CHAT_MAX_TOKENS when no limit is + supplied, and leaves an explicit limit untouched.""" + + def _chat_request(self, model: str, optional_params: dict) -> dict: + config = OCIChatConfig() + body = config.transform_request( + model=model, + messages=[{"role": "user", "content": "hi"}], + optional_params={**BASE_OCI_PARAMS, **optional_params}, + litellm_params={}, + headers={}, + ) + return body["chatRequest"] + + @pytest.mark.parametrize( + "model", ["cohere.command-latest", "meta.llama-3.3-70b-instruct"] + ) + def test_default_injected_when_max_tokens_omitted(self, model): + chat_request = self._chat_request(model, {}) + assert chat_request["maxTokens"] == DEFAULT_OCI_CHAT_MAX_TOKENS + + @pytest.mark.parametrize( + "model", ["cohere.command-latest", "meta.llama-3.3-70b-instruct"] + ) + def test_explicit_max_tokens_not_overridden(self, model): + chat_request = self._chat_request(model, {"max_tokens": 256}) + assert chat_request["maxTokens"] == 256 + + def test_reasoning_model_defaults_max_completion_tokens(self): + chat_request = self._chat_request("openai.gpt-5", {}) + assert chat_request["maxCompletionTokens"] == DEFAULT_OCI_CHAT_MAX_TOKENS + assert "maxTokens" not in chat_request + + class TestOCIReasoningEffort: """ Reasoning-effort handling for GENERIC reasoning models: @@ -1133,8 +1304,7 @@ class TestOCIStreamingSignedBody: When signed_json_body is provided, the POST must use that exact bytes object, not json.dumps(data) — otherwise the RSA-SHA256 signature is invalid. """ - import httpx - from unittest.mock import MagicMock, patch + from unittest.mock import MagicMock config = OCIChatConfig() signed_bytes = b'{"signed": true}' @@ -1293,6 +1463,68 @@ class TestOCIChatConfigErrorPaths: ) assert "audio" not in result + @pytest.mark.parametrize("model", ["cohere.command-latest", "xai.grok-4"]) + def test_map_openai_params_max_retries_dropped_without_drop_params(self, model): + """max_retries is a litellm control param, not a generation param. It + must be dropped silently (no raise) even when drop_params is False, so + the litellm proxy (which injects max_retries on every request) does not + 500 every OCI call. + """ + config = OCIChatConfig() + result = config.map_openai_params( + non_default_params={"max_retries": 3}, + optional_params={}, + model=model, + drop_params=False, + ) + assert "max_retries" not in result + def test_map_openai_params_cohere_n_default_dropped(self): + """Cohere has no numGenerations field, but n=1 (and None) is the OpenAI + default single-generation request. It must be dropped silently rather + than raising, so standard clients that always send n=1 (e.g. the MLflow + gateway) are not rejected.""" + config = OCIChatConfig() + for n in (1, None): + result = config.map_openai_params( + non_default_params={"n": n}, + optional_params={}, + model="cohere.command-latest", + drop_params=False, + ) + assert "n" not in result and "numGenerations" not in result + + def test_map_openai_params_cohere_n_gt_1_raises_without_drop(self): + """n>1 is genuinely unsupported on Cohere and must raise without drop.""" + config = OCIChatConfig() + with pytest.raises(Exception, match="not supported on OCI"): + config.map_openai_params( + non_default_params={"n": 3}, + optional_params={}, + model="cohere.command-latest", + drop_params=False, + ) + + def test_map_openai_params_cohere_n_gt_1_dropped_with_drop(self): + config = OCIChatConfig() + result = config.map_openai_params( + non_default_params={"n": 3}, + optional_params={}, + model="cohere.command-latest", + drop_params=True, + ) + assert "n" not in result and "numGenerations" not in result + + def test_map_openai_params_generic_n_maps_to_num_generations(self): + """Generic models keep numGenerations, including n>1.""" + config = OCIChatConfig() + result = config.map_openai_params( + non_default_params={"n": 2}, + optional_params={}, + model=TEST_MODEL_NAME, + drop_params=False, + ) + assert result["numGenerations"] == 2 + def test_transform_request_tool_choice_string_mapped(self): config = OCIChatConfig() result = config.transform_request( diff --git a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py index cc914a22eeb..5dd44d72d68 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py @@ -5,6 +5,7 @@ import json from unittest.mock import patch, MagicMock from litellm import ModelResponse +from litellm.constants import DEFAULT_OCI_CHAT_MAX_TOKENS from litellm.llms.oci.chat.cohere import ( adapt_messages_to_cohere_standard, adapt_tool_definitions_to_cohere_standard, @@ -236,25 +237,30 @@ class TestOCICohereToolCalls: assert result.usage.completion_tokens == 22 assert result.usage.total_tokens == 48 - def test_cohere_request_preserves_json_schema_response_format(self): - """Ensure Cohere requests retain JSON schema payloads in responseFormat.""" + def test_cohere_request_folds_json_schema_into_json_object(self): + """A Cohere json_schema must fold the schema onto JSON_OBJECT. + + OCI Cohere has no JSON_SCHEMA type; sending {"type": "JSON_SCHEMA", ...} + (or the raw lowercase "json_schema" with a jsonSchema body) is rejected + with HTTP 400. The schema rides on JSON_OBJECT instead. + """ config = OCIChatConfig() messages = [{"role": "user", "content": "Return structured info"}] - response_format = { - "type": "json_schema", - "json_schema": { - "name": "test_schema", - "strict": True, - "schema": { - "type": "object", - "properties": {"foo": {"type": "string"}}, - "required": ["foo"], - }, - }, + schema = { + "type": "object", + "properties": {"foo": {"type": "string"}}, + "required": ["foo"], } optional_params = { "oci_compartment_id": TEST_COMPARTMENT_ID, - "response_format": response_format, + "response_format": { + "type": "json_schema", + "json_schema": { + "name": "test_schema", + "strict": True, + "schema": schema, + }, + }, } transformed_request = config.transform_request( @@ -265,18 +271,14 @@ class TestOCICohereToolCalls: headers={}, ) - chat_request = transformed_request["chatRequest"] - assert chat_request["apiFormat"] == "COHERE" - assert "responseFormat" in chat_request - - cohere_response_format = chat_request["responseFormat"] - assert cohere_response_format["type"] == "json_schema" + cohere_response_format = transformed_request["chatRequest"]["responseFormat"] + assert cohere_response_format["type"] == "JSON_OBJECT" + assert "jsonSchema" not in cohere_response_format assert "json_schema" not in cohere_response_format - assert "jsonSchema" in cohere_response_format - assert cohere_response_format["jsonSchema"] == response_format["json_schema"] + assert cohere_response_format["schema"] == schema - def test_cohere_request_response_format_text_stays_lowercase(self): - """Ensure Cohere keeps response_format type lowercase (e.g. 'text' not 'TEXT').""" + def test_cohere_request_response_format_text_is_uppercased(self): + """Cohere response_format type 'text' maps to OCI's canonical 'TEXT'.""" config = OCIChatConfig() messages = [{"role": "user", "content": "Hello"}] optional_params = { @@ -292,10 +294,7 @@ class TestOCICohereToolCalls: headers={}, ) - chat_request = transformed_request["chatRequest"] - assert chat_request["apiFormat"] == "COHERE" - assert "responseFormat" in chat_request - assert chat_request["responseFormat"]["type"] == "text" + assert transformed_request["chatRequest"]["responseFormat"] == {"type": "TEXT"} def test_cohere_tool_call_only_message_no_text(self): """Test chat history with an assistant message that has tool calls but no text content.""" @@ -462,7 +461,8 @@ class TestOCICohereToolCalls: assert "tool_choice" not in supported_params def test_cohere_default_parameters(self): - """Test that Cohere requests do not inject hardcoded defaults — caller supplies all params.""" + """maxTokens is defaulted (OCI's server default truncates at ~20 tokens); + every other param is still pass-through with no hardcoded default.""" config = OCIChatConfig() messages = [{"role": "user", "content": "Hello"}] optional_params = {"oci_compartment_id": TEST_COMPARTMENT_ID} @@ -477,8 +477,7 @@ class TestOCICohereToolCalls: chat_request = transformed_request["chatRequest"] - # No hardcoded defaults injected — only pass through what the user supplies - assert "maxTokens" not in chat_request + assert chat_request["maxTokens"] == DEFAULT_OCI_CHAT_MAX_TOKENS assert "topK" not in chat_request assert "topP" not in chat_request assert "frequencyPenalty" not in chat_request diff --git a/tests/test_litellm/llms/oci/chat/test_oci_generic_chat.py b/tests/test_litellm/llms/oci/chat/test_oci_generic_chat.py index 7583e3bc183..a4a5f111513 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_generic_chat.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_generic_chat.py @@ -2,7 +2,6 @@ Unit tests for litellm/llms/oci/chat/generic.py — error paths and stream handling. """ -import json import pytest from unittest.mock import MagicMock @@ -16,7 +15,12 @@ from litellm.llms.oci.chat.generic import ( handle_generic_response, handle_generic_stream_chunk, ) -from litellm.llms.oci.chat.transformation import OCIChatConfig, OCIStreamWrapper +from litellm.llms.oci.chat.transformation import ( + OCIChatConfig, + OCIStreamWrapper, + OCIVendors, + _model_uses_max_completion_tokens, +) from litellm.llms.oci.common_utils import OCIError # --------------------------------------------------------------------------- @@ -271,7 +275,6 @@ class TestHandleGenericStreamChunk: assert result.choices[0].index == 0 def test_image_content_in_stream_raises(self): - from litellm.types.llms.oci import OCIImageContentPart, OCIImageUrl, OCIMessage chunk = { "apiFormat": "GENERIC", @@ -368,10 +371,6 @@ def _register_oci_gpt5_in_catalog(): class TestGpt5MaxCompletionTokens: def test_helper_detects_gpt5_family(self, _register_oci_gpt5_in_catalog): - from litellm.llms.oci.chat.transformation import ( - _model_uses_max_completion_tokens, - ) - assert _model_uses_max_completion_tokens("openai.gpt-5") is True assert _model_uses_max_completion_tokens("openai.gpt-5-mini") is True assert _model_uses_max_completion_tokens("openai.gpt-5-nano") is True @@ -382,11 +381,40 @@ class TestGpt5MaxCompletionTokens: assert _model_uses_max_completion_tokens("cohere.command-latest") is False assert _model_uses_max_completion_tokens("") is False + def test_helper_covers_openai_models_absent_from_catalog(self): + """OCI keeps adding OpenAI models (gpt-4.1, gpt-5.1..5.5, o-series) + faster than the litellm catalog tracks them. The vendor-prefix rule + must route them to maxCompletionTokens even with no catalog entry, + since OpenAI accepts max_completion_tokens on every chat model while + the reasoning families hard-reject max_tokens.""" + import litellm + + for name in ( + "openai.gpt-5.2", + "openai.gpt-4.1", + "openai.o3", + "oci/openai.gpt-5.1-codex", + ): + assert f"oci/{name.removeprefix('oci/')}" not in litellm.model_cost + assert _model_uses_max_completion_tokens(name) is True + + assert _model_uses_max_completion_tokens("openai.gpt-oss-20b") is False + + def test_default_injection_uses_max_completion_tokens_for_uncataloged_gpt(self): + """Regression: with the injected default maxTokens, a GPT model absent + from the catalog got "maxTokens" on every request and OCI returned 400 + ("Use 'max_completion_tokens' instead") even when the caller never set + max_tokens.""" + from litellm.constants import DEFAULT_OCI_CHAT_MAX_TOKENS + + cfg = OCIChatConfig() + out = cfg._get_optional_params(OCIVendors.GENERIC, {}, model="openai.gpt-5.2") + assert out.get("maxCompletionTokens") == DEFAULT_OCI_CHAT_MAX_TOKENS + assert "maxTokens" not in out + def test_gpt5_routes_max_tokens_to_max_completion_tokens( self, _register_oci_gpt5_in_catalog ): - from litellm.llms.oci.chat.transformation import OCIChatConfig, OCIVendors - cfg = OCIChatConfig() # Both shapes optional_params can take after upstream map_openai_params: # 1. openai-side key still present @@ -404,8 +432,6 @@ class TestGpt5MaxCompletionTokens: assert "maxTokens" not in out_b def test_non_gpt5_keeps_max_tokens(self): - from litellm.llms.oci.chat.transformation import OCIChatConfig, OCIVendors - cfg = OCIChatConfig() out = cfg._get_optional_params( OCIVendors.GENERIC, @@ -416,8 +442,6 @@ class TestGpt5MaxCompletionTokens: assert "maxCompletionTokens" not in out def test_cohere_reasoning_model_keeps_max_tokens(self): - from litellm.llms.oci.chat.transformation import OCIChatConfig, OCIVendors - cfg = OCIChatConfig() out = cfg._get_optional_params( OCIVendors.COHERE, diff --git a/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py b/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py new file mode 100644 index 00000000000..62fc3a8d0aa --- /dev/null +++ b/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py @@ -0,0 +1,226 @@ +""" +Tests for the Realtime transcription_sessions surface used by gpt-realtime-whisper: + - OpenAI / Azure URL construction (POST /v1/realtime/transcription_sessions) + - RealtimeTranscriptionSessionRequest model-resolution + passthrough + - BaseLLMHTTPHandler.async_realtime_transcription_session_handler targeting +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.azure.realtime.http_transformation import AzureRealtimeHTTPConfig +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig +from litellm.types.realtime import RealtimeTranscriptionSessionRequest + + +def test_openai_transcription_session_url(): + cfg = OpenAIRealtimeHTTPConfig() + assert ( + cfg.get_transcription_session_url( + api_base="https://api.openai.com", model="gpt-realtime-whisper" + ) + == "https://api.openai.com/v1/realtime/transcription_sessions" + ) + + +def test_openai_transcription_session_url_strips_trailing_v1(): + """A /v1 suffix must not be duplicated in the path.""" + cfg = OpenAIRealtimeHTTPConfig() + assert ( + cfg.get_transcription_session_url( + api_base="https://api.openai.com/v1", model="gpt-realtime-whisper" + ) + == "https://api.openai.com/v1/realtime/transcription_sessions" + ) + + +def test_azure_transcription_session_url_uses_deployment_and_api_version(): + cfg = AzureRealtimeHTTPConfig() + url = cfg.get_transcription_session_url( + api_base="https://my.openai.azure.com", + model="whisper-deploy", + api_version="2025-04-01-preview", + ) + assert ( + url + == "https://my.openai.azure.com/openai/realtime/transcription_sessions?api-version=2025-04-01-preview" + ) + + +def test_request_resolves_model_returns_none_when_both_absent(): + req = RealtimeTranscriptionSessionRequest(input_audio_format="pcm16") + assert req.resolved_model() is None + + req = RealtimeTranscriptionSessionRequest( + model="openai/gpt-realtime-whisper", + input_audio_transcription={"model": "gpt-realtime-whisper"}, + ) + assert req.resolved_model() == "openai/gpt-realtime-whisper" + + +def test_request_resolves_model_from_input_audio_transcription(): + req = RealtimeTranscriptionSessionRequest( + input_audio_transcription={"model": "gpt-realtime-whisper", "language": "en"}, + ) + assert req.resolved_model() == "gpt-realtime-whisper" + + +def test_request_passthrough_excludes_routing_hint(): + """Unknown fields pass through; the litellm-only `model` hint is not forwarded.""" + req = RealtimeTranscriptionSessionRequest( + model="openai/gpt-realtime-whisper", + input_audio_format="pcm16", + input_audio_transcription={"model": "gpt-realtime-whisper"}, + turn_detection=None, + ) + forwarded = req.model_dump(exclude_none=True, exclude={"model"}) + assert "model" not in forwarded + assert forwarded["input_audio_format"] == "pcm16" + assert forwarded["input_audio_transcription"] == {"model": "gpt-realtime-whisper"} + + +@pytest.mark.asyncio +async def test_handler_posts_to_transcription_sessions_url(): + handler = BaseLLMHTTPHandler() + + mock_response = MagicMock(spec=httpx.Response) + mock_client = MagicMock(spec=AsyncHTTPHandler) + mock_client.post = AsyncMock(return_value=mock_response) + + logging_obj = MagicMock() + logging_obj.pre_call = MagicMock() + + request_body = {"input_audio_transcription": {"model": "gpt-realtime-whisper"}} + result = await handler.async_realtime_transcription_session_handler( + api_base="https://api.openai.com", + api_key="sk-test", + request_data=request_body, + logging_obj=logging_obj, + timeout=10.0, + provider_config=OpenAIRealtimeHTTPConfig(), + model="gpt-realtime-whisper", + client=mock_client, + ) + + assert result is mock_response + _, kwargs = mock_client.post.call_args + assert kwargs["url"] == "https://api.openai.com/v1/realtime/transcription_sessions" + assert kwargs["json"] == request_body + assert kwargs["headers"]["Authorization"] == "Bearer sk-test" + + +@pytest.mark.asyncio +async def test_client_secret_handler_still_targets_client_secrets_url(): + """Refactor regression: the client_secrets handler must keep its own URL.""" + handler = BaseLLMHTTPHandler() + + mock_response = MagicMock(spec=httpx.Response) + mock_client = MagicMock(spec=AsyncHTTPHandler) + mock_client.post = AsyncMock(return_value=mock_response) + + logging_obj = MagicMock() + logging_obj.pre_call = MagicMock() + + await handler.async_realtime_client_secret_handler( + api_base="https://api.openai.com", + api_key="sk-test", + request_data={"session": {"type": "realtime"}}, + logging_obj=logging_obj, + timeout=10.0, + provider_config=OpenAIRealtimeHTTPConfig(), + model="gpt-4o-realtime-preview", + client=mock_client, + ) + + _, kwargs = mock_client.post.call_args + assert kwargs["url"] == "https://api.openai.com/v1/realtime/client_secrets" + + +@pytest.mark.asyncio +async def test_sdk_fn_routes_openai_transcription_session(monkeypatch): + """ + litellm.acreate_realtime_transcription_session resolves the OpenAI provider + from the transcription model and POSTs to the OpenAI transcription_sessions URL. + """ + import litellm + + monkeypatch.setenv("OPENAI_API_KEY", "sk-unit-test") + + mock_response = MagicMock(spec=httpx.Response) + mock_client = MagicMock(spec=AsyncHTTPHandler) + mock_client.post = AsyncMock(return_value=mock_response) + + result = await litellm.acreate_realtime_transcription_session( + model="openai/gpt-realtime-whisper", + transcription_session={ + "input_audio_format": "pcm16", + "input_audio_transcription": {"model": "gpt-realtime-whisper"}, + }, + client=mock_client, + ) + + assert result is mock_response + _, kwargs = mock_client.post.call_args + assert kwargs["url"].endswith("/v1/realtime/transcription_sessions") + # The litellm-only routing hint must not be forwarded upstream. + assert "model" not in kwargs["json"] + assert kwargs["json"]["input_audio_transcription"] == { + "model": "gpt-realtime-whisper" + } + + +def test_append_query_params_skips_existing_keys(): + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + + url = "wss://example.com/v1/realtime?model=gpt-4o" + result = BaseLLMHTTPHandler._append_query_params( + url, {"model": "ignored", "intent": "transcription"} + ) + assert "model=ignored" not in result + assert "intent=transcription" in result + + +def test_append_query_params_no_params_returns_unchanged(): + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + + url = "wss://example.com/v1/realtime?model=gpt-4o" + assert BaseLLMHTTPHandler._append_query_params(url, None) == url + assert BaseLLMHTTPHandler._append_query_params(url, {}) == url + + +def test_append_query_params_encodes_special_chars(): + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + + url = "wss://example.com/v1/realtime" + result = BaseLLMHTTPHandler._append_query_params(url, {"intent": "a&b=c"}) + assert "intent=a%26b%3Dc" in result + assert "&b=c" not in result + + +def test_azure_construct_url_encodes_model_and_api_version(): + """model and api-version must be URL-encoded to prevent query-string injection.""" + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + h = AzureOpenAIRealtime() + url = h._construct_url( + "https://x.openai.azure.com", + "deploy&evil=1", + "2024-10-01-preview", + ) + assert "evil=1" not in url.split("?", 1)[1] + + url_ga = h._construct_url( + "https://x.openai.azure.com", + "deploy&evil=1", + None, + realtime_protocol="GA", + ) + assert "evil=1" not in url_ga.split("?", 1)[1] diff --git a/tests/test_litellm/llms/parallel_ai/__init__.py b/tests/test_litellm/llms/parallel_ai/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py b/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py new file mode 100644 index 00000000000..b5c1a86205b --- /dev/null +++ b/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py @@ -0,0 +1,324 @@ +""" +Tests for Parallel AI Search API integration (v1 endpoint). +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm + +MOCK_V1_RESPONSE = { + "search_id": "search_abc123", + "session_id": "session_xyz", + "results": [ + { + "url": "https://example.com/1", + "title": "Test Result 1", + "publish_date": "2026-01-15", + "excerpts": ["First excerpt.", "Second excerpt."], + }, + { + "url": "https://example.com/2", + "title": None, + "publish_date": None, + "excerpts": ["Only excerpt."], + }, + ], + "usage": [{"name": "search_advanced", "count": 1}], +} + + +def _mock_response(): + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = MOCK_V1_RESPONSE + return mock_response + + +class TestParallelAISearch: + @pytest.fixture(autouse=True) + def _set_api_key(self, monkeypatch): + monkeypatch.setenv("PARALLEL_API_KEY", "test-api-key") + monkeypatch.delenv("PARALLEL_AI_API_BASE", raising=False) + + @pytest.mark.asyncio + async def test_v1_endpoint_and_headers(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="latest developments in AI", + search_provider="parallel_ai", + ) + + call_args = mock_post.call_args + assert call_args.kwargs["url"] == "https://api.parallel.ai/v1/search" + + headers = call_args.kwargs.get("headers", {}) + assert headers["x-api-key"] == "test-api-key" + assert headers["Content-Type"] == "application/json" + assert "parallel-beta" not in headers + + @pytest.mark.asyncio + async def test_string_query_maps_to_search_queries_and_objective(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="latest developments in AI", + search_provider="parallel_ai", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["search_queries"] == ["latest developments in AI"] + assert json_data["objective"] == "latest developments in AI" + + @pytest.mark.asyncio + async def test_list_query_maps_to_search_queries(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query=["AI developments", "machine learning trends"], + search_provider="parallel_ai", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["search_queries"] == [ + "AI developments", + "machine learning trends", + ] + assert "objective" not in json_data + + @pytest.mark.asyncio + async def test_mode_param_passthrough(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + mode="turbo", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["mode"] == "turbo" + + @pytest.mark.asyncio + async def test_default_mode_is_basic(self): + """v1 defaults to 'advanced' server-side; litellm must send 'basic' to keep v1beta's default tier and cost tracking accurate.""" + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["mode"] == "basic" + + @pytest.mark.parametrize( + "processor,expected_mode", [("base", "basic"), ("pro", "advanced")] + ) + @pytest.mark.asyncio + async def test_legacy_processor_maps_to_mode(self, processor, expected_mode): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + processor=processor, + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["mode"] == expected_mode + assert "processor" not in json_data + + @pytest.mark.asyncio + async def test_explicit_mode_wins_over_processor(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + mode="turbo", + processor="pro", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["mode"] == "turbo" + assert "processor" not in json_data + + @pytest.mark.asyncio + async def test_top_level_v1_params_pass_through(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + session_id="session_123", + max_chars_total=4000, + max_tokens_per_page=1024, + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["session_id"] == "session_123" + assert json_data["max_chars_total"] == 4000 + assert "max_tokens_per_page" not in json_data + + @pytest.mark.asyncio + async def test_optional_params_nest_under_advanced_settings(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + max_results=5, + country="US", + search_domain_filter=["arxiv.org", "nature.com"], + exclude_domains=["reddit.com"], + max_chars_per_result=1500, + ) + + json_data = mock_post.call_args.kwargs.get("json") + advanced_settings = json_data["advanced_settings"] + assert advanced_settings["max_results"] == 5 + assert advanced_settings["location"] == "US" + assert advanced_settings["source_policy"]["include_domains"] == [ + "arxiv.org", + "nature.com", + ] + assert advanced_settings["source_policy"]["exclude_domains"] == [ + "reddit.com" + ] + assert advanced_settings["excerpt_settings"]["max_chars_per_result"] == 1500 + + assert "max_results" not in json_data + assert "source_policy" not in json_data + assert "search_domain_filter" not in json_data + assert "exclude_domains" not in json_data + assert "max_chars_per_result" not in json_data + assert "country" not in json_data + + @pytest.mark.asyncio + async def test_explicit_advanced_settings_take_precedence(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + max_results=5, + advanced_settings={"max_results": 7}, + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["advanced_settings"]["max_results"] == 7 + + @pytest.mark.asyncio + async def test_response_transformation(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + response = await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + ) + + assert response.object == "search" + assert len(response.results) == 2 + + first = response.results[0] + assert first.title == "Test Result 1" + assert first.url == "https://example.com/1" + assert first.snippet == "First excerpt. ... Second excerpt." + assert first.date == "2026-01-15" + + second = response.results[1] + assert second.title == "" + assert second.snippet == "Only excerpt." + assert second.date is None + + @pytest.mark.parametrize( + "api_base", + [ + "https://proxy.internal.example.com", + "https://proxy.internal.example.com/", + "https://proxy.internal.example.com/v1", + "https://proxy.internal.example.com/v1/", + "https://proxy.internal.example.com/v1/search", + ], + ) + @pytest.mark.asyncio + async def test_custom_api_base_appends_v1_search(self, api_base): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + api_base=api_base, + ) + + call_args = mock_post.call_args + assert ( + call_args.kwargs["url"] + == "https://proxy.internal.example.com/v1/search" + ) + + @pytest.mark.asyncio + async def test_missing_api_key_raises(self, monkeypatch): + monkeypatch.delenv("PARALLEL_API_KEY", raising=False) + monkeypatch.delenv("PARALLEL_AI_API_KEY", raising=False) + + with pytest.raises(Exception, match="PARALLEL_API_KEY"): + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + ) diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py b/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py index d4d76ab3079..92464ae2c31 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py @@ -302,3 +302,39 @@ def test_streaming_tool_call_finish_reason_with_empty_content_in_final_chunk(): assert len(response2.choices) == 1 # Must be "tool_calls", NOT "stop" assert response2.choices[0].finish_reason == "tool_calls" + + +def test_streaming_metadata_only_chunk_does_not_yield_empty_choices(): + """ + web_search + reasoning makes Gemini emit mid-stream chunks that carry only + grounding/thought metadata — no content part and no finishReason. + _process_candidates skips content-less candidates, so without a fallback + `choices` is empty and the downstream streaming handler hits + `IndexError: list index out of range` on choices[0]. + + Ref: https://github.com/BerriAI/litellm/issues/28884 + """ + logging_obj = _make_logging_obj() + iterator = ModelResponseIterator( + streaming_response=iter([]), + sync_stream=True, + logging_obj=logging_obj, + ) + + # Grounding-only chunk: a candidate with groundingMetadata but no content + # part and no finishReason (what web_search + reasoning produces mid-stream). + metadata_only_chunk = { + "candidates": [ + { + "index": 0, + "groundingMetadata": {"webSearchQueries": ["weather boston"]}, + } + ] + } + + response = iterator.chunk_parser(metadata_only_chunk) + assert response is not None + # Must expose at least one choice so downstream choices[0] is safe. + assert len(response.choices) == 1 + assert response.choices[0].finish_reason is None + assert response.choices[0].delta.content is None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py index d900f690c57..d51cf8c5b72 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -88,6 +88,76 @@ async def test_fetch_tools_from_passthrough_raises_on_upstream_401(): mock_client.list_tools.assert_awaited_with(raise_on_error=True) +@pytest.mark.asyncio +async def test_fetch_tools_from_delegated_oauth2_raises_on_upstream_401(): + manager = MCPServerManager() + delegated_server = MCPServer( + server_id="oauth1", + name="delegated_docs", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + ) + + response = httpx.Response( + status_code=401, + headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + upstream_error = httpx.HTTPStatusError( + "401", request=response.request, response=response + ) + + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(side_effect=upstream_error) + + with pytest.raises(MCPUpstreamAuthError) as exc_info: + await manager._fetch_tools_with_timeout( + mock_client, delegated_server.name, server=delegated_server + ) + + assert exc_info.value.status_code == 401 + assert exc_info.value.www_authenticate == ( + 'Bearer resource_metadata="https://upstream"' + ) + assert exc_info.value.server_name == "delegated_docs" + mock_client.list_tools.assert_awaited_with(raise_on_error=True) + + +@pytest.mark.asyncio +async def test_fetch_tools_from_client_credentials_oauth2_keeps_swallow_behavior(): + manager = MCPServerManager() + m2m_server = MCPServer( + server_id="oauth-m2m", + name="m2m_docs", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + oauth2_flow="client_credentials", + ) + + response = httpx.Response( + status_code=401, + headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + upstream_error = httpx.HTTPStatusError( + "401", request=response.request, response=response + ) + + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(side_effect=upstream_error) + + tools = await manager._fetch_tools_with_timeout( + mock_client, m2m_server.name, server=m2m_server + ) + + assert tools == [] + mock_client.list_tools.assert_awaited_with(raise_on_error=False) + + @pytest.mark.asyncio async def test_fetch_tools_from_passthrough_returns_tools_on_success(): manager = MCPServerManager() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index b6550fee6b9..c9009ed526d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5489,3 +5489,84 @@ async def test_create_mcp_client_sampling_enabled(): client = await manager._create_mcp_client(server=server) assert client._sampling_callback is not None + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_sets_model_in_model_call_details(): + """Regression test: MCP tools/call spend logs persisted with model="". + + execute_mcp_tool set logging_obj.model only; the spend-log writer reads + model_call_details["model"], which stays None when function_setup builds + the logging object without a "model" kwarg. + """ + import uuid + from datetime import timezone + + from litellm.proxy._experimental.mcp_server import server as mcp_module + from litellm.proxy._types import LitellmUserRoles + from litellm.utils import Rules, function_setup + + user = UserAPIKeyAuth( + api_key="sk-user", + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + + fake_server = MagicMock() + fake_server.name = "openapi-petstore" + fake_server.is_byok = False + fake_server.auth_type = None + fake_server.mcp_info = None + fake_server.server_id = "srv-1" + fake_server.server_name = "openapi-petstore" + + fake_tool = MagicMock() + fake_tool.name = "list_pets" + + start_time = datetime.now(timezone.utc) + litellm_logging_obj, _ = function_setup( + original_function="call_mcp_tool", + rules_obj=Rules(), + start_time=start_time, + litellm_call_id=str(uuid.uuid4()), + name="list_pets", + arguments={"limit": 10}, + ) + assert litellm_logging_obj.model_call_details.get("model") is None + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=fake_server, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "pre_call_tool_check", + new=AsyncMock(return_value={}), + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=fake_tool, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + new=AsyncMock(return_value=[]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ), + ): + await mcp_module.execute_mcp_tool( + name="list_pets", + arguments={"limit": 10}, + allowed_mcp_servers=[fake_server], + start_time=start_time, + user_api_key_auth=user, + litellm_logging_obj=litellm_logging_obj, + ) + + assert litellm_logging_obj.model_call_details["model"] == "MCP: list_pets" + assert litellm_logging_obj.model == "MCP: list_pets" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index d52af94c47f..d80d15e2140 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -673,6 +673,110 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401(): assert "Bearer authorization_uri=" in exc_info.value.headers["www-authenticate"] +@pytest.mark.asyncio +async def test_handle_streamable_http_mcp_delegated_server_surfaces_upstream_challenge(): + """ + OAuth2 server with ``delegate_auth_to_upstream=True`` should let the + upstream MCP server's RFC 9728 challenge reach the client instead of + pre-emptively returning LiteLLM's gateway authorization_uri challenge. + """ + from fastapi import HTTPException + + try: + from litellm.proxy._experimental.mcp_server.exceptions import ( + MCPUpstreamAuthError, + ) + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/delegated_oauth_server", + "scheme": "https", + "query_string": b"", + "root_path": "", + "server": ("litellm.example.com", 443), + "headers": [ + (b"content-type", b"application/json"), + (b"host", b"litellm.example.com"), + ], + } + receive = AsyncMock( + return_value={ + "type": "http.request", + "body": b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}', + "more_body": False, + } + ) + send = AsyncMock() + user_auth = MagicMock() + user_auth.user_id = None + delegated_server = MagicMock() + delegated_server.auth_type = MCPAuth.oauth2 + delegated_server.delegate_auth_to_upstream = True + delegated_server.needs_user_oauth_token = True + delegated_server.server_id = "delegated-oauth-server" + + upstream_challenge = ( + 'Bearer resource_metadata="https://upstream.example.com/.well-known/oauth-protected-resource"' + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=( + user_auth, + None, + ["delegated_oauth_server"], + None, + None, + None, + ), + ), + patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + new_callable=AsyncMock, + return_value=None, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + return_value=delegated_server, + ), + patch.object( + session_manager_stateful, + "handle_request", + new_callable=AsyncMock, + side_effect=MCPUpstreamAuthError( + status_code=401, + www_authenticate=upstream_challenge, + server_name="delegated_oauth_server", + ), + ) as mock_handle_request, + ): + with pytest.raises(HTTPException) as exc_info: + await handle_streamable_http_mcp(scope, receive, send) + + assert mock_handle_request.await_count == 1 + assert exc_info.value.status_code == 401 + assert exc_info.value.headers == {"www-authenticate": upstream_challenge} + + @pytest.mark.asyncio async def test_per_user_oauth_with_stored_token_skips_preemptive_401(): """ @@ -759,19 +863,16 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401(): @pytest.mark.asyncio -async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without_token(): +async def test_handle_streamable_http_mcp_delegated_server_without_token_reaches_session_manager(): """ - OAuth2 server with ``delegate_auth_to_upstream=True`` and no Authorization - header must still emit a pre-emptive 401 with WWW-Authenticate so the - client kicks off PKCE. The 401 points at LiteLLM's discovery shim, which - in turn delegates to the upstream OAuth issuer. + OAuth2 server with ``delegate_auth_to_upstream=True`` and no stored token + should not receive LiteLLM's gateway authorization_uri challenge. The + request continues so the upstream MCP server can emit its RFC 9728 challenge. """ - from fastapi import HTTPException - try: from litellm.proxy._experimental.mcp_server.server import ( handle_streamable_http_mcp, - session_manager, + session_manager_stateless, ) except ImportError: pytest.skip("MCP server not available") @@ -785,7 +886,13 @@ async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without (b"host", b"litellm.example.com"), ], } - receive = AsyncMock() + receive = AsyncMock( + return_value={ + "type": "http.request", + "body": b'{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{}}', + "more_body": False, + } + ) send = AsyncMock() user_auth = MagicMock() user_auth.user_id = None @@ -819,19 +926,22 @@ async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without new_callable=AsyncMock, return_value=False, ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + new_callable=AsyncMock, + return_value=None, + ) as mock_get_stored_token, patch( "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", return_value=delegated_server, ), patch.object( - session_manager, + session_manager_stateless, "handle_request", new_callable=AsyncMock, ) as mock_handle_request, ): - with pytest.raises(HTTPException) as exc_info: - await handle_streamable_http_mcp(scope, receive, send) + await handle_streamable_http_mcp(scope, receive, send) - assert exc_info.value.status_code == 401 - assert "www-authenticate" in exc_info.value.headers - assert mock_handle_request.await_count == 0 + assert mock_get_stored_token.await_count == 1 + assert mock_handle_request.await_count == 1 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 71178c4826c..f43d8e85aca 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -2503,3 +2503,267 @@ async def test_post_call_success_hook_only_runs_output_scan(): mock_make.call_args.kwargs.get("logging_event_type") == GuardrailEventHooks.post_call ) + + +# --------------------------------------------------------------------------- +# Contextual grounding: request-side qualifiers +# --------------------------------------------------------------------------- +# +# Bedrock contextual grounding tags each ApplyGuardrail content block with a +# `qualifiers` array (grounding_source / query / guard_content). A caller marks +# message content blocks `{"type": "grounding_source", ...}` / `{"type": "query", ...}`; +# at post_call the hook assembles one source="OUTPUT" call carrying the source + +# query + the model response (as guard_content). A request without these tags +# produces the plain-text payload with no qualifiers. + +_GROUNDING_SOURCE_TEXT = "Tokyo is the capital of Japan." +_GROUNDING_QUERY_TEXT = "What is the capital of Japan?" +_GROUNDING_RESPONSE_TEXT = "The capital of Japan is Tokyo." + + +def _grounding_guardrail() -> BedrockGuardrail: + return BedrockGuardrail( + guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT" + ) + + +def _grounding_messages() -> list: + return [ + { + "role": "system", + "content": [{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}], + }, + { + "role": "user", + "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}], + }, + ] + + +def _model_response(content: str) -> ModelResponse: + from litellm.types.utils import Choices, Message, ModelResponse + + return ModelResponse( + choices=[ + Choices( + index=0, + message=Message(role="assistant", content=content), + finish_reason="stop", + ) + ] + ) + + +# Expected OUTPUT content blocks, keyed by their grounding qualifier, so the +# per-test assertions read as the block sequence they expect. +_GROUNDING_SOURCE_BLOCK = { + "text": {"text": _GROUNDING_SOURCE_TEXT, "qualifiers": ["grounding_source"]} +} +_QUERY_BLOCK = {"text": {"text": _GROUNDING_QUERY_TEXT, "qualifiers": ["query"]}} +_GUARD_BLOCK = { + "text": {"text": _GROUNDING_RESPONSE_TEXT, "qualifiers": ["guard_content"]} +} + + +def _input_request(messages: list) -> dict: + """Arrange a guardrail and act: build the Bedrock INPUT payload.""" + return _grounding_guardrail().convert_to_bedrock_format( + source="INPUT", messages=messages + ) + + +def _output_request(messages: list, response=None) -> dict: + """Arrange a guardrail and act: build the Bedrock OUTPUT payload.""" + return _grounding_guardrail().convert_to_bedrock_format( + source="OUTPUT", response=response, messages=messages + ) + + +def test_grounding_input_strips_grounding_and_query_qualifiers(): + """Grounding is OUTPUT-only: tagged source/query reach Bedrock as plain text on an + INPUT scan, so a tag cannot change how input-safety policies scan content (no bypass). + """ + expected_request = { + "source": "INPUT", + "content": [ + {"text": {"text": _GROUNDING_SOURCE_TEXT}}, + {"text": {"text": _GROUNDING_QUERY_TEXT}}, + ], + } + + actual_request = _input_request(_grounding_messages()) + + assert actual_request == expected_request + + +def test_grounding_input_leaves_existing_guarded_text_unqualified(): + """An existing guarded_text input block keeps its legacy unqualified payload.""" + expected_request = {"source": "INPUT", "content": [{"text": {"text": "policy"}}]} + + actual_request = _input_request( + [{"role": "user", "content": [{"type": "guarded_text", "text": "policy"}]}] + ) + + assert actual_request == expected_request + + +def test_grounding_output_assembles_source_query_and_response(): + """OUTPUT emits grounding_source + query (from the request) then the response as + guard_content, so Bedrock can grade the response against the source and query.""" + expected_request = { + "source": "OUTPUT", + "content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK, _GUARD_BLOCK], + } + + actual_request = _output_request( + _grounding_messages(), _model_response(_GROUNDING_RESPONSE_TEXT) + ) + + assert actual_request == expected_request + + +def test_grounding_output_keeps_legacy_payload_without_tags(): + """Without grounding tags the OUTPUT payload is the legacy single response block.""" + expected_request = { + "source": "OUTPUT", + "content": [{"text": {"text": "Hi there."}}], + } + + actual_request = _output_request( + [{"role": "user", "content": "hello"}], _model_response("Hi there.") + ) + + assert actual_request == expected_request + + +def test_grounding_output_combines_multiple_sources(): + """Every grounding_source block is emitted; Bedrock combines them into one corpus.""" + uk_source_text = "London is the capital of UK." + uk_source_block = { + "text": {"text": uk_source_text, "qualifiers": ["grounding_source"]} + } + messages = [ + { + "role": "system", + "content": [ + {"type": "grounding_source", "text": uk_source_text}, + {"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}, + ], + }, + {"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]}, + ] + expected_request = { + "source": "OUTPUT", + "content": [ + uk_source_block, + _GROUNDING_SOURCE_BLOCK, + _QUERY_BLOCK, + _GUARD_BLOCK, + ], + } + + actual_request = _output_request( + messages, _model_response(_GROUNDING_RESPONSE_TEXT) + ) + + assert actual_request == expected_request + + +def test_grounding_output_keeps_grounding_for_non_model_response(): + """Harvested grounding blocks survive a non-ModelResponse output instead of being + silently dropped (regression guard for the unconditional content assignment).""" + expected_request = { + "source": "OUTPUT", + "content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK], + } + + actual_request = _output_request(_grounding_messages(), response=None) + + assert actual_request == expected_request + + +@pytest.mark.parametrize( + "role, is_trusted", + [ + ("system", True), + ("developer", True), + ("tool", False), + ("function", False), + ("user", False), + ("assistant", False), + ], +) +def test_grounding_source_trusted_only_from_app_roles(role, is_trusted): + """grounding_source is honored only from app-authored roles (system/developer). A + tag on a user, tool, function or assistant message is ignored, so neither a forwarded + end user nor an externally-influenced tool result can supply fake evidence for the + grounding check to grade the response against; query is always collected.""" + messages = [ + { + "role": role, + "content": [{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}], + }, + {"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]}, + ] + expected_content = [_QUERY_BLOCK, _GUARD_BLOCK] + if is_trusted: + expected_content = [_GROUNDING_SOURCE_BLOCK, *expected_content] + + actual_request = _output_request( + messages, _model_response(_GROUNDING_RESPONSE_TEXT) + ) + + assert actual_request == {"source": "OUTPUT", "content": expected_content} + + +@pytest.mark.asyncio +async def test_grounding_output_blocked_raises_400(): + """A BLOCKED contextualGroundingPolicy filter raises HTTP 400.""" + guardrail = _grounding_guardrail() + + mock_bedrock_response = MagicMock() + mock_bedrock_response.status_code = 200 + mock_bedrock_response.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "assessments": [ + { + "contextualGroundingPolicy": { + "filters": [ + { + "type": "GROUNDING", + "threshold": 0.7, + "score": 0.1, + "action": "BLOCKED", + } + ] + } + } + ], + "outputs": [{"text": "Response blocked: not grounded in the provided source."}], + } + + mock_credentials = MagicMock() + mock_credentials.access_key = "test-access-key" + mock_credentials.secret_key = "test-secret-key" + mock_credentials.token = None + + with ( + patch.object( + guardrail.async_handler, "post", new_callable=AsyncMock + ) as mock_post, + patch.object( + guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1") + ), + patch.object(guardrail, "_prepare_request", return_value=MagicMock()), + ): + mock_post.return_value = mock_bedrock_response + + with pytest.raises(HTTPException) as exc_info: + await guardrail.make_bedrock_api_request( + source="OUTPUT", + response=_model_response("The capital of Japan is Paris."), + messages=_grounding_messages(), + request_data={"messages": _grounding_messages()}, + ) + + assert exc_info.value.status_code == 400 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py new file mode 100644 index 00000000000..4160a835ca4 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py @@ -0,0 +1,653 @@ +""" +Unit tests for Ovalix guardrail: config resolution and apply_guardrail behavior +with mocked Tracker service responses (allow, anonymize, block). +""" + +import os +from typing import Any, List +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix import ( + OvalixGuardrail, + OvalixGuardrailBlockedException, + OvalixGuardrailMissingSecrets, +) +from litellm.types.utils import GenericGuardrailAPIInputs + +# Example Tracker responses (as returned by the checkpoint API) +TRACKER_RESPONSE_ALLOW = { + "action_type": "allow", + "data_type": "TEXT", + "original_data": {"content": "how are you?"}, + "modified_data": {"content": "how are you?"}, + "alerts": [], +} + +TRACKER_RESPONSE_ANONYMIZE = { + "action_type": "anonymize", + "data_type": "TEXT", + "original_data": {"content": "Hello, my name is David."}, + "modified_data": {"content": "Hello, my name is {Name}. How are you?"}, + "alerts": [ + { + "title": "Sensitive Data Alert", + "subtitle": "We've identified that you were trying to share sensitive information", + "alerts": ["Name:\tDavid\nRedacted to:\t{Name}"], + } + ], +} + +TRACKER_RESPONSE_BLOCK = { + "action_type": "block", + "data_type": "TEXT", + "original_data": {"content": "I am 15 YO"}, + "modified_data": {"content": "This message was blocked by Ovalix"}, + "alerts": [ + { + "title": "Sensitive Data Alert", + "subtitle": "We've identified that you were trying to share sensitive information", + "alerts": ["Age:\t15\nBlocked"], + } + ], +} + + +def _ovalix_env(): + return { + "OVALIX_TRACKER_API_BASE": "https://tracker.test", + "OVALIX_TRACKER_API_KEY": "key", + "OVALIX_APPLICATION_ID": "app-1", + "OVALIX_PRE_CHECKPOINT_ID": "pre-1", + "OVALIX_POST_CHECKPOINT_ID": "post-1", + } + + +def _guardrail_kwargs(): + return { + "guardrail_name": "ovalix-test", + "event_hook": "pre_call", + "default_on": True, + } + + +class TestOvalixGuardrailConfigModel: + """Minimal config model tests: wiring only.""" + + def test_get_config_model_returns_ovalix_config_model(self): + """get_config_model returns OvalixGuardrailConfigModel for proxy/config wiring.""" + config_model = OvalixGuardrail.get_config_model() + assert config_model is not None + assert config_model.__name__ == "OvalixGuardrailConfigModel" + assert config_model.ui_friendly_name() == "Ovalix Guardrail" + + +class TestOvalixGuardrail: + """Behavioral tests with mocked Tracker checkpoint API.""" + + def setup_method(self): + for key in list(os.environ.keys()): + if key.startswith("OVALIX_"): + del os.environ[key] + + def teardown_method(self): + for key in list(os.environ.keys()): + if key.startswith("OVALIX_"): + del os.environ[key] + + @pytest.fixture + def guardrail_with_env(self): + """Guardrail with OVALIX_* env set; cleans up in teardown.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + yield OvalixGuardrail(**_guardrail_kwargs()) + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + def test_initialization_requires_secrets(self): + """Initialization raises when required Tracker/application/checkpoint config is missing.""" + with pytest.raises(OvalixGuardrailMissingSecrets): + OvalixGuardrail( + guardrail_name="ovalix-test", + event_hook="pre_call", + default_on=True, + ) + + def test_initialization_with_explicit_params(self): + """Guardrail initializes with explicit tracker base, key, app and checkpoint IDs.""" + guardrail = OvalixGuardrail( + tracker_api_base="https://tracker.example", + tracker_api_key="secret", + application_id="app-x", + pre_checkpoint_id="pre-x", + post_checkpoint_id="post-x", + **_guardrail_kwargs(), + ) + assert guardrail._tracker_api_base == "https://tracker.example" + assert guardrail._application_id == "app-x" + assert guardrail._pre_checkpoint_id == "pre-x" + assert guardrail._post_checkpoint_id == "post-x" + + def test_initialization_with_env_vars(self): + """Guardrail picks up OVALIX_* env vars when params not passed.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + assert guardrail._tracker_api_base == "https://tracker.test" + assert guardrail._tracker_api_key == "key" + assert guardrail._application_id == "app-1" + assert guardrail._pre_checkpoint_id == "pre-1" + assert guardrail._post_checkpoint_id == "post-1" + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_call_checkpoint_sends_correct_payload_and_returns_json(self): + """_call_checkpoint POSTs to tracker with application_id, checkpoint_id, actor, session_id, data.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + mock_response = MagicMock() + mock_response.json.return_value = TRACKER_RESPONSE_ALLOW + mock_response.raise_for_status = MagicMock() + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = mock_response + result = await guardrail._call_checkpoint( + content="hello", + checkpoint_id="pre-1", + actor="a1b2c3d4", + session_id="session-1", + ) + + assert result == TRACKER_RESPONSE_ALLOW + mock_post.assert_called_once() + call_args = mock_post.call_args + assert call_args.args[0] == ( + "https://tracker.test/tracking/custom_application/checkpoint" + ) + body = call_args.kwargs["json"] + assert body["application_id"] == "app-1" + assert body["checkpoint_id"] == "pre-1" + assert body["actor"] == "a1b2c3d4" + assert body["session_id"] == "session-1" + assert body["data_type"] == "TEXT" + assert body["data"] == {"content": "hello"} + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_allow_passes_through(self): + """When Tracker returns allow, apply_guardrail returns inputs with texts set to modified_data content.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "how are you?"}], + texts=["how are you?"], + ) + request_data = {} + + mock_response = MagicMock() + mock_response.json.return_value = TRACKER_RESPONSE_ALLOW + mock_response.raise_for_status = MagicMock() + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = mock_response + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result.get("texts") == ["how are you?"] + assert mock_post.call_count == 1 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_anonymize_returns_modified_text(self): + """When Tracker returns anonymize, apply_guardrail returns texts with modified_data content.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[ + {"role": "user", "content": "Hello, my name is David."} + ], + texts=["Hello, my name is David."], + ) + request_data = {} + + mock_response = MagicMock() + mock_response.json.return_value = TRACKER_RESPONSE_ANONYMIZE + mock_response.raise_for_status = MagicMock() + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = mock_response + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result.get("texts") == ["Hello, my name is {Name}. How are you?"] + assert mock_post.call_count == 1 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_block_raises_with_tracker_message(self): + """When Tracker returns block on the (chronologically) last user message, OvalixGuardrailBlockedException is raised.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "I am 15 YO"}], + texts=["I am 15 YO"], + ) + request_data = {} + + mock_response = MagicMock() + mock_response.json.return_value = TRACKER_RESPONSE_BLOCK + mock_response.raise_for_status = MagicMock() + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = mock_response + with pytest.raises(OvalixGuardrailBlockedException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert "This message was blocked by Ovalix" in str(exc_info.value.message) + assert exc_info.value.status_code == 400 + assert mock_post.call_count == 1 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_block_non_last_replaced_in_texts(self): + """When Tracker returns block on a non-last user message, that message is replaced in texts and no exception is raised.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[ + {"role": "user", "content": "I am 15 YO"}, + {"role": "user", "content": "how are you?"}, + ], + texts=["I am 15 YO", "how are you?"], + ) + request_data = {} + + def side_effect(*args, **kwargs): + body = kwargs.get("json", {}) + content = (body.get("data") or {}).get("content", "") + resp = MagicMock() + if "15" in content: + resp.json.return_value = TRACKER_RESPONSE_BLOCK + else: + resp.json.return_value = TRACKER_RESPONSE_ALLOW + resp.raise_for_status = MagicMock() + return resp + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.side_effect = side_effect + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result.get("texts") == [ + "This message was blocked by Ovalix", + "how are you?", + ] + assert mock_post.call_count == 2 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_response_allow_returns_inputs(self): + """When input_type is response and Tracker allows, apply_guardrail returns inputs with texts updated from Tracker.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[ + {"role": "assistant", "content": "Safe assistant reply"} + ], + texts=["Safe assistant reply"], + ) + request_data = {} + + mock_response = MagicMock() + mock_response.json.return_value = TRACKER_RESPONSE_ALLOW + mock_response.raise_for_status = MagicMock() + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = mock_response + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="response", + logging_obj=None, + ) + + assert result.get("texts") == ["how are you?"] + assert mock_post.call_count == 1 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_response_block_raises(self, guardrail_with_env): + """When Tracker blocks on response, apply_guardrail raises OvalixGuardrailBlockedException.""" + guardrail = guardrail_with_env + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "I am 15 YO"}], + texts=["I am 15 YO"], + ) + request_data = {} + + mock_response = MagicMock() + mock_response.json.return_value = TRACKER_RESPONSE_BLOCK + mock_response.raise_for_status = MagicMock() + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = mock_response + with pytest.raises(OvalixGuardrailBlockedException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="response", + logging_obj=None, + ) + + assert "This message was blocked by Ovalix" in str(exc_info.value.message) + assert exc_info.value.status_code == 400 + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + async def test_apply_guardrail_request_missing_modified_data_uses_original_content( + self, guardrail_with_env + ): + """When Tracker response has no modified_data.content, original content is used.""" + guardrail = guardrail_with_env + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "original text"}], + texts=["original text"], + ) + request_data = {} + + mock_response = MagicMock() + mock_response.json.return_value = { + "action_type": "allow", + "data_type": "TEXT", + "original_data": {"content": "original text"}, + "modified_data": {}, + "alerts": [], + } + mock_response.raise_for_status = MagicMock() + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = mock_response + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result.get("texts") == ["original text"] + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + async def test_apply_guardrail_tracker_http_error_raises_guardrail_exception( + self, guardrail_with_env + ): + """When Tracker returns HTTP error (e.g. 400), GuardrailRaisedException is raised.""" + guardrail = guardrail_with_env + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "hello"}], + texts=["hello"], + ) + request_data = {} + + mock_response = MagicMock() + mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( + "Bad Request", + request=MagicMock(), + response=MagicMock(status_code=400), + ) + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = mock_response + with pytest.raises(GuardrailRaisedException): + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + async def test_apply_guardrail_checkpoint_error_raises_guardrail_exception(self): + """When Tracker checkpoint call fails, GuardrailRaisedException is raised.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "hello"}], + texts=["hello"], + ) + request_data = {} + + with patch.object( + guardrail._async_handler, + "post", + new_callable=AsyncMock, + side_effect=httpx.ConnectError("Connection refused"), + ): + with pytest.raises(GuardrailRaisedException): + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_empty_messages_returns_inputs(self): + """When request has no messages, apply_guardrail returns inputs without calling Tracker.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs(structured_messages=[], texts=[]) + request_data = {} + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result == inputs + mock_post.assert_not_called() + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + def test_get_actor_from_metadata(self): + """Actor is taken from metadata.user_api_key_user_email or user_api_key_user_id.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + assert ( + guardrail._get_actor( + {"metadata": {"user_api_key_user_email": "a@b.com"}} + ) + == "a@b.com" + ) + assert ( + guardrail._get_actor({"metadata": {"user_api_key_user_id": "uid-1"}}) + == "uid-1" + ) + assert ( + guardrail._get_actor( + {"litellm_metadata": {"user_api_key_user_id": "uid-2"}} + ) + == "uid-2" + ) + assert guardrail._get_actor({}) == "unknown" + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + def test_get_actor_prefers_email_over_id(self, guardrail_with_env): + """When both user_api_key_user_email and user_api_key_user_id exist, email is used.""" + guardrail = guardrail_with_env + data = { + "metadata": { + "user_api_key_user_email": "primary@test.com", + "user_api_key_user_id": "uid-99", + } + } + assert guardrail._get_actor(data) == "primary@test.com" + + def test_get_tracker_actor_id_is_hash_not_raw_pii(self, guardrail_with_env): + """Tracker API actor field uses a short hash of _get_actor, not email/user id.""" + guardrail = guardrail_with_env + data = {"metadata": {"user_api_key_user_email": "user@example.com"}} + raw = guardrail._get_actor(data) + hashed = guardrail._get_tracker_actor_id(data) + assert raw == "user@example.com" + assert hashed != raw + assert len(hashed) == 8 + assert all(c in "0123456789abcdef" for c in hashed) + + def test_get_session_id_deterministic_and_includes_app_id(self, guardrail_with_env): + """Session ID is stable for same actor/day and includes application_id.""" + guardrail = guardrail_with_env + data = {"metadata": {"user_api_key_user_id": "user-1"}} + session_id_1 = guardrail._get_session_id(data) + session_id_2 = guardrail._get_session_id(data) + assert session_id_1 == session_id_2 + assert "app-1" in session_id_1 + + def test_block_current_message_raises_ovalix_blocked_exception( + self, guardrail_with_env + ): + """_block_current_message raises OvalixGuardrailBlockedException with status_code 400.""" + guardrail = guardrail_with_env + with pytest.raises(OvalixGuardrailBlockedException) as exc_info: + guardrail._block_current_message("Custom block reason") + assert "Custom block reason" in str(exc_info.value.message) + assert exc_info.value.status_code == 400 + + def test_get_trackers_corrected_message(self, guardrail_with_env): + """_get_trackers_corrected_message returns modified_data.content or None.""" + guardrail = guardrail_with_env + assert ( + guardrail._get_trackers_corrected_message( + {"modified_data": {"content": "corrected text"}} + ) + == "corrected text" + ) + assert guardrail._get_trackers_corrected_message({"modified_data": {}}) is None + assert ( + guardrail._get_trackers_corrected_message({"modified_data": "not-a-dict"}) + is None + ) + + @pytest.mark.asyncio + async def test_apply_guardrail_response_no_texts_returns_unchanged(self): + """When input_type is response and inputs have no texts, apply_guardrail returns inputs without calling Tracker.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs() + request_data = {} + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="response", + logging_obj=None, + ) + + assert result == inputs + mock_post.assert_not_called() + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index 64c57ab90e3..d31cfdc39bd 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -756,6 +756,85 @@ async def test_health_services_endpoint_rejects_unknown_service(): await health_services_endpoint(service="totally_unknown_service_xyz") +@pytest.mark.asyncio +@pytest.mark.parametrize( + "role", + [ + None, + LitellmUserRoles.INTERNAL_USER, + LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, + LitellmUserRoles.TEAM, + LitellmUserRoles.CUSTOMER, + ], +) +async def test_health_services_endpoint_newrelic_blocks_non_admin(role): + """ + /health/services?service=newrelic emits a real LiteLLMConnectionTest event + to the configured New Relic account. Only proxy admins (full or view-only) + should be able to trigger it; every other caller must be rejected before + the external event is recorded. + """ + from litellm.proxy._types import ProxyException + + user_api_key_dict = UserAPIKeyAuth( + token="non-admin-token", + user_id="non-admin-user", + user_role=role, + ) + + with patch( + "litellm.integrations.newrelic.newrelic.NewRelicLogger" + ) as MockNewRelicLogger: + mock_instance = MagicMock() + mock_instance.async_health_check = AsyncMock( + return_value={"status": "healthy", "error_message": ""} + ) + MockNewRelicLogger.return_value = mock_instance + + with pytest.raises(ProxyException) as exc_info: + await health_services_endpoint( + user_api_key_dict=user_api_key_dict, + service="newrelic", + ) + + assert str(exc_info.value.code) == "403" + mock_instance.async_health_check.assert_not_awaited() + MockNewRelicLogger.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "admin_role", + [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY], +) +async def test_health_services_endpoint_newrelic_allows_proxy_admin(admin_role): + """ + Proxy admins (full and view-only) can trigger the New Relic test event. + """ + user_api_key_dict = UserAPIKeyAuth( + token="admin-token", + user_id="admin-user", + user_role=admin_role, + ) + + with patch( + "litellm.integrations.newrelic.newrelic.NewRelicLogger" + ) as MockNewRelicLogger: + mock_instance = MagicMock() + mock_instance.async_health_check = AsyncMock( + return_value={"status": "healthy", "error_message": ""} + ) + MockNewRelicLogger.return_value = mock_instance + + result = await health_services_endpoint( + user_api_key_dict=user_api_key_dict, + service="newrelic", + ) + + assert result["status"] == "healthy" + mock_instance.async_health_check.assert_awaited_once() + + @pytest.fixture(scope="function") def proxy_client(monkeypatch): """ diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index ae699ff8e12..50f471721b1 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -6,6 +6,7 @@ import asyncio import os import sys import time +from contextlib import contextmanager from datetime import datetime, timedelta from typing import Any, Dict, List, Optional @@ -20,6 +21,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token +from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.utils import ( EmbeddingResponse, ModelResponse, @@ -3189,6 +3191,246 @@ def test_get_key_mcp_rpm_limit_precedence(): assert get_team_mcp_rpm_limit(none_set) is None +async def _seed_max_parallel_requests_counter( + dual_cache: DualCache, counter_key: str, window_size: int +) -> None: + await dual_cache.async_increment_cache_pipeline( + increment_list=[ + RedisPipelineIncrementOperation( + key=counter_key, increment_value=1, ttl=window_size + ) + ] + ) + + +async def _build_seeded_limiter(): + """Build a v3 limiter whose api-key counter already holds the pre-call +1.""" + api_key = hash_token("sk-disconnect") + cache = DualCache() + limiter = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(cache) + ) + counter_key = f"{{api_key:{api_key}}}:max_parallel_requests" + await _seed_max_parallel_requests_counter(cache, counter_key, limiter.window_size) + user_api_key_dict = UserAPIKeyAuth(api_key=api_key, max_parallel_requests=2) + return limiter, cache, counter_key, user_api_key_dict + + +@contextmanager +def _override_litellm_callbacks(new_callbacks): + """Swap litellm.callbacks so _callback_capabilities recomputes deterministically.""" + saved = litellm.callbacks + litellm.callbacks = new_callbacks + try: + yield + finally: + litellm.callbacks = saved + + +async def _drain_release_task(): + # The disconnect release is scheduled fire-and-forget via create_task. + for _ in range(5): + await asyncio.sleep(0) + + +@pytest.mark.asyncio +async def test_release_max_parallel_requests_on_disconnect_v3(): + """ + Regression for issue #27955: a stream cancelled mid-flight must release the + pre-call +1 reservation. The success/failure logging callbacks never fire + on cancellation, so without an explicit release the api-key counter climbs + by one per cancelled request until the key wedges at its limit. The release + must decrement the api-key max_parallel_requests counter by exactly one. + """ + _api_key = hash_token("sk-12345") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=2) + counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" + + await _seed_max_parallel_requests_counter( + local_cache, counter_key, handler.window_size + ) + assert await local_cache.async_get_cache(key=counter_key) == 1 + + await handler.async_release_max_parallel_requests_on_disconnect(user_api_key_dict) + + assert await local_cache.async_get_cache(key=counter_key) == 0 + + +@pytest.mark.asyncio +async def test_release_max_parallel_requests_on_disconnect_noop_v3(): + """ + The release must be a no-op when the key never reserved a parallel slot + (no api_key, or max_parallel_requests unset). Otherwise a cancelled + no-limit request would drive an unrelated counter negative. + """ + _api_key = hash_token("sk-12345") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" + + await handler.async_release_max_parallel_requests_on_disconnect( + UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None) + ) + assert await local_cache.async_get_cache(key=counter_key) is None + + await handler.async_release_max_parallel_requests_on_disconnect( + UserAPIKeyAuth(api_key=None, max_parallel_requests=5) + ) + assert await local_cache.async_get_cache(key=counter_key) is None + + +@pytest.mark.parametrize("disconnect", ["cancel", "aclose"]) +@pytest.mark.asyncio +async def test_async_streaming_data_generator_releases_counter_on_disconnect_v3( + disconnect, +): + """ + Regression for issue #27955 on the outer SSE generator (used by /v1/messages + and other event-stream routes). A client that disconnects mid-stream raises + GeneratorExit (aclose) or CancelledError into async_streaming_data_generator; + both are BaseException and bypass the success/failure logging callbacks, so + the generator itself must refund the pre-call max_parallel_requests +1. + Releasing inside the nested iterator hook does not work because that + generator is only closed on garbage collection, which is non-deterministic. + """ + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + + limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter() + assert await cache.async_get_cache(key=counter_key) == 1 + + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter + + async def upstream(): + yield ModelResponse() + if disconnect == "cancel": + raise asyncio.CancelledError() + while True: + yield ModelResponse() + + with _override_litellm_callbacks([]): + gen = ProxyBaseLLMRequestProcessing.async_sse_data_generator( + response=upstream(), + user_api_key_dict=user_api_key_dict, + request_data={"model": "claude-test"}, + proxy_logging_obj=proxy_logging_obj, + ) + await gen.__anext__() + if disconnect == "cancel": + with pytest.raises(asyncio.CancelledError): + await gen.__anext__() + else: + await gen.aclose() + await _drain_release_task() + + assert await cache.async_get_cache(key=counter_key) == 0 + + +@pytest.mark.parametrize("disconnect", ["cancel", "aclose"]) +@pytest.mark.asyncio +async def test_async_data_generator_releases_counter_on_disconnect_v3(disconnect): + """ + Regression for issue #27955 on the chat-completions outer generator + (proxy_server.async_data_generator). With only the v3 parallel limiter + enabled, needs_iterator_wrap() is False, so this generator iterates the + upstream response directly and the iterator hook is bypassed entirely -- the + gap that let a disconnect leak the slot in the default limiter-only config. + A mid-stream disconnect must still refund the pre-call +1. + """ + import litellm.proxy.proxy_server as proxy_server + + limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter() + proxy_logging_obj = proxy_server.proxy_logging_obj + saved_hook = proxy_logging_obj.proxy_hook_mapping.get("parallel_request_limiter") + proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter + + async def upstream(): + yield ModelResponse() + if disconnect == "cancel": + raise asyncio.CancelledError() + while True: + yield ModelResponse() + + try: + with _override_litellm_callbacks([]): + assert proxy_logging_obj.needs_iterator_wrap() is False + gen = proxy_server.async_data_generator( + response=upstream(), + user_api_key_dict=user_api_key_dict, + request_data={"model": "gpt-test"}, + ) + await gen.__anext__() + if disconnect == "cancel": + with pytest.raises(asyncio.CancelledError): + await gen.__anext__() + else: + await gen.aclose() + await _drain_release_task() + assert await cache.async_get_cache(key=counter_key) == 0 + finally: + if saved_hook is not None: + proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = ( + saved_hook + ) + else: + proxy_logging_obj.proxy_hook_mapping.pop("parallel_request_limiter", None) + + +@pytest.mark.asyncio +async def test_async_data_generator_releases_counter_when_wrapped_v3(): + """ + Companion to the no-wrap case for issue #27955. With an iterator-override + callback active, needs_iterator_wrap() is True and async_data_generator + drives the chained iterator hook. The refund must still fire exactly once + from the outer generator: the counter returns to 0 (not -1), proving the + nested hook does not also refund and there is no double decrement. + """ + from litellm.integrations.custom_logger import CustomLogger + import litellm.proxy.proxy_server as proxy_server + + class _PassthroughIteratorOverride(CustomLogger): + async def async_post_call_streaming_iterator_hook( + self, user_api_key_dict, response, request_data + ): + async for chunk in response: + yield chunk + + limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter() + proxy_logging_obj = proxy_server.proxy_logging_obj + saved_hook = proxy_logging_obj.proxy_hook_mapping.get("parallel_request_limiter") + proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter + + async def upstream(): + while True: + yield ModelResponse() + + try: + with _override_litellm_callbacks([_PassthroughIteratorOverride()]): + assert proxy_logging_obj.needs_iterator_wrap() is True + gen = proxy_server.async_data_generator( + response=upstream(), + user_api_key_dict=user_api_key_dict, + request_data={"model": "gpt-test"}, + ) + await gen.__anext__() + await gen.aclose() + await _drain_release_task() + assert await cache.async_get_cache(key=counter_key) == 0 + finally: + if saved_hook is not None: + proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = ( + saved_hook + ) + else: + proxy_logging_obj.proxy_hook_mapping.pop("parallel_request_limiter", None) + + def test_tpm_reservation_enabled_by_default(monkeypatch): """Upfront TPM reservation is on unless explicitly disabled via env.""" monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 9ad57b55555..8c26e9e4e1e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -57,6 +57,51 @@ async def test_get_daily_activity_empty_entity_id_list(): assert where_conditions["team_id"] == {"in": []} +@pytest.mark.asyncio +async def test_get_daily_activity_order_has_id_tiebreaker(): + """Regression for #30164. + + ``date`` alone is not a unique sort key for either + ``LiteLLM_DailyUserSpend`` or ``LiteLLM_DailyTeamSpend`` -- a busy + tenant has many rows per date (one per api_key, model, model_group, + provider, endpoint, ...). Offset pagination over a non-unique sort + landed on arbitrary page boundaries between queries, so summing + per-page totals across pages produced non-deterministic results + (sometimes inflated, sometimes deflated). The tiebreaker on the + UUID primary key pins the row order so a client paging through all + results gets the correct total. + """ + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_table = MagicMock() + mock_table.count = AsyncMock(return_value=0) + mock_table.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_verificationtoken = MagicMock() + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_dailyspend = mock_table + + await get_daily_activity( + prisma_client=mock_prisma, + table_name="litellm_dailyspend", + entity_id_field="team_id", + entity_id="team-1", + entity_metadata_field=None, + start_date="2024-01-01", + end_date="2024-01-02", + model=None, + api_key=None, + page=1, + page_size=10, + ) + + mock_table.find_many.assert_called_once() + order = mock_table.find_many.call_args[1]["order"] + assert order == [{"date": "desc"}, {"id": "asc"}], ( + f"order must include the id tiebreaker after date for stable offset " + f"pagination (see #30164); got {order!r}" + ) + + def test_is_user_agent_tag(): """Test _is_user_agent_tag function.""" # Test None and empty string diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 473d61f8a85..046971d033b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -1488,6 +1488,65 @@ async def test_prepare_key_update_data_duration_none_never_expires(): assert result["expires"] is None +@pytest.mark.asyncio +@pytest.mark.parametrize("cleared_value", [[], None]) +async def test_prepare_key_update_data_budget_limits_clears_field(cleared_value): + """budget_limits=[] / None must serialize to JSON null, never reach Prisma raw.""" + from litellm.proxy._types import UpdateKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import ( + prepare_key_update_data, + ) + + existing_key = LiteLLM_VerificationToken( + token="test-token", + key_alias="test-key", + models=["gpt-3.5-turbo"], + user_id="test-user", + team_id=None, + metadata={}, + ) + + update_request = UpdateKeyRequest(key="test-token", budget_limits=cleared_value) + + result = await prepare_key_update_data( + data=update_request, existing_key_row=existing_key + ) + + assert result["budget_limits"] == json.dumps(None) + + +@pytest.mark.asyncio +async def test_prepare_key_update_data_budget_limits_serializes_windows(): + """Non-empty budget_limits stay JSON-encoded with reset_at initialized.""" + from litellm.proxy._types import UpdateKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import ( + prepare_key_update_data, + ) + + existing_key = LiteLLM_VerificationToken( + token="test-token", + key_alias="test-key", + models=["gpt-3.5-turbo"], + user_id="test-user", + team_id=None, + metadata={}, + ) + + update_request = UpdateKeyRequest( + key="test-token", + budget_limits=[{"budget_duration": "1d", "max_budget": 10.0}], + ) + + result = await prepare_key_update_data( + data=update_request, existing_key_row=existing_key + ) + + windows = json.loads(result["budget_limits"]) + assert isinstance(result["budget_limits"], str) + assert windows[0]["max_budget"] == 10.0 + assert windows[0]["reset_at"] is not None + + @pytest.mark.asyncio async def test_validate_team_id_used_in_service_account_request_requires_team_id(): """ @@ -9685,6 +9744,58 @@ class TestKeyOwnerPrivilegeEscalation: ) mock_check.assert_called_once() + @pytest.mark.asyncio + @pytest.mark.parametrize("cleared_value", [[], None]) + async def test_creator_cannot_clear_own_budget_limits(self, cleared_value): + """Clearing budget_limits is a budget change and requires admin.""" + data = UpdateKeyRequest(key="sk-test", budget_limits=cleared_value) + existing = self._make_existing_key(created_by="creator-123") + auth = self._make_auth(user_id="creator-123") + + mock_check = AsyncMock( + side_effect=HTTPException(status_code=403, detail="Not authorized") + ) + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints._check_key_admin_access", + mock_check, + ): + with pytest.raises(HTTPException): + await _validate_update_key_data( + data=data, + existing_key_row=existing, + user_api_key_dict=auth, + llm_router=None, + premium_user=False, + prisma_client=AsyncMock(), + user_api_key_cache=MagicMock(), + ) + mock_check.assert_called_once() + + @pytest.mark.asyncio + async def test_admin_can_clear_budget_limits(self): + data = UpdateKeyRequest(key="sk-test", budget_limits=[]) + existing = self._make_existing_key(created_by="someone-else") + auth = UserAPIKeyAuth( + user_id="admin-user", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + mock_check = AsyncMock() + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints._check_key_admin_access", + mock_check, + ): + await _validate_update_key_data( + data=data, + existing_key_row=existing, + user_api_key_dict=auth, + llm_router=None, + premium_user=False, + prisma_client=AsyncMock(), + user_api_key_cache=MagicMock(), + ) + mock_check.assert_not_called() + @pytest.mark.asyncio async def test_admin_can_update_any_field(self): data = UpdateKeyRequest(key="sk-test", models=["gpt-4"], max_budget=999.0) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py index 3c78a390a45..2d708a3644d 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py @@ -987,6 +987,72 @@ class TestAnthropicBatchPassthroughCostTracking: ) +class TestBuildCompleteStreamingResponseRobustness: + """_build_complete_streaming_response must tolerate non-standard SSE frames.""" + + def _build(self, chunks: List[str]): + return AnthropicPassthroughLoggingHandler._build_complete_streaming_response( + all_chunks=chunks, + litellm_logging_obj=MagicMock(), + model="claude-3-sonnet-20240229", + ) + + def test_done_frame_is_skipped(self): + """A bare 'data: [DONE]' control frame must not break reconstruction.""" + chunks = [ + 'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-3-sonnet-20240229","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":1}}}', + 'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}', + 'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hi"}}', + 'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}', + 'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":2}}', + 'event: message_stop\ndata: {"type":"message_stop"}', + "data: [DONE]", + ] + result = self._build(chunks) + assert result is not None + assert result.choices[0].message.content == "Hi" + + def test_non_json_sse_line_is_skipped(self): + """Non-JSON SSE lines (comments, keep-alive pings) must be skipped.""" + chunks = [ + ": ping", + 'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-3-sonnet-20240229","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":1}}}', + "this is not json at all", + ] + # Must not raise; a malformed stream simply yields no usable response. + result = self._build(chunks) + assert result is None or hasattr(result, "choices") + + def test_mixed_valid_and_invalid_frames(self): + """Valid events are still collected when interleaved with invalid ones.""" + chunks = [ + 'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-3-sonnet-20240229","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":1}}}', + "data: [DONE]", + ": keep-alive", + "not-json", + 'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}', + 'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}', + 'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}', + 'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":2}}', + 'event: message_stop\ndata: {"type":"message_stop"}', + ] + result = self._build(chunks) + assert result is not None + assert result.choices[0].message.content == "Hello" + + def test_done_in_text_payload_is_not_dropped(self): + """A valid event whose text content contains '[DONE]' must NOT be skipped.""" + chunks = [ + 'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-3-sonnet-20240229","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":1}}}', + 'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}', + 'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"The stream ends with [DONE]"}}', + 'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}', + 'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":8}}', + 'event: message_stop\ndata: {"type":"message_stop"}', + ] + result = self._build(chunks) + assert result is not None + assert result.choices[0].message.content == "The stream ends with [DONE]" class TestPureTextFastPathParity: """ The pure-text fast path in _build_complete_streaming_response must produce diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py index 367c89a05f3..65853df392f 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py @@ -241,6 +241,142 @@ async def test_client_secrets_success_with_mock( proxy_app.dependency_overrides.pop(user_api_key_auth, None) +@pytest.mark.asyncio +async def test_client_secrets_transcription_rejects_disallowed_nested_model( + proxy_app, +): + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", + models=["gpt-4o-realtime-preview"], + ) + try: + client = TestClient(proxy_app, raise_server_exceptions=False) + with ( + patch("litellm.proxy.proxy_server.route_request") as mock_route_request, + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + ): + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/client_secrets", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={ + "model": "gpt-4o-realtime-preview", + "session": { + "type": "transcription", + "model": "gpt-4o-realtime-preview", + "audio": { + "input": { + "transcription": { + "model": "gpt-realtime-whisper" + } + } + }, + }, + }, + ) + + assert response.status_code == 403 + assert "Tried to access gpt-realtime-whisper" in response.text + mock_route_request.assert_not_called() + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_client_secrets_transcription_routes_on_nested_model( + proxy_app, + mock_add_litellm_data, + mock_pre_call_hook, +): + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", + models=["gpt-4o-realtime-preview", "gpt-realtime-whisper"], + ) + captured = {} + future_expires_at = int(time.time()) + 3600 + + async def _capturing_route(*args, **kwargs): + captured["data"] = kwargs.get("data") + + async def _inner(): + resp = MagicMock(spec=httpx.Response) + resp.status_code = 200 + resp.text = ( + f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}' + ) + resp.content = ( + f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}' + ).encode() + resp.headers = {} + resp.json.return_value = { + "value": "upstream_ephemeral_key", + "expires_at": future_expires_at, + } + return resp + + return _inner() + + try: + client = TestClient(proxy_app) + with ( + patch( + "litellm.proxy.proxy_server.route_request", + side_effect=_capturing_route, + ), + patch( + "litellm.proxy.proxy_server.add_litellm_data_to_request", + side_effect=mock_add_litellm_data, + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + ): + mock_logging.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook) + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/client_secrets", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={ + "model": "gpt-4o-realtime-preview", + "session": { + "type": "transcription", + "model": "gpt-4o-realtime-preview", + "audio": { + "input": { + "transcription": { + "model": "gpt-realtime-whisper" + } + } + }, + }, + }, + ) + + assert response.status_code == 200 + assert captured["data"]["model"] == "gpt-realtime-whisper" + session = captured["data"]["session"] + assert session["type"] == "transcription" + assert "model" not in session + assert ( + session["audio"]["input"]["transcription"]["model"] + == "gpt-realtime-whisper" + ) + encrypted_value = response.json()["value"] + decoded = _decode_realtime_token_payload( + decrypt_value_helper( + encrypted_value, + key="client_secret.value", + exception_type="debug", + ) + or "" + ) + assert decoded is not None + assert decoded["model_id"] == "gpt-realtime-whisper" + assert decoded["session_type"] == "transcription" + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + def test_realtime_calls_requires_auth(proxy_app): """POST /v1/realtime/calls returns 401 without Authorization. @@ -311,3 +447,547 @@ async def test_realtime_calls_success_with_valid_encrypted_token( assert response.status_code == 201 assert response.content.startswith(b"v=0") assert b"application/sdp" in response.headers.get("content-type", "").encode() + + +def test_token_payload_carries_session_type(): + """The encrypted token records the session kind so /realtime/calls can replay it.""" + payload = _encode_realtime_token_payload( + ephemeral_key="epk", + model_id="gpt-realtime-whisper", + user_id=None, + team_id=None, + expires_at=None, + session_type="transcription", + ) + decoded = _decode_realtime_token_payload(payload) + assert decoded is not None + assert decoded["session_type"] == "transcription" + + +@pytest.mark.asyncio +async def test_realtime_calls_replays_transcription_session_type( + proxy_app, + mock_add_litellm_data, + mock_pre_call_hook, +): + """ + A token minted for a transcription session must drive /realtime/calls to send + session.type == "transcription" upstream, not the default "realtime". + """ + captured = {} + + async def _capturing_route(*args, **kwargs): + captured["session"] = kwargs.get("data", {}).get("session") + + async def _inner(): + resp = MagicMock(spec=httpx.Response) + resp.status_code = 201 + resp.content = b"v=0\r\n" + resp.headers = {"content-type": "application/sdp"} + return resp + + return _inner() + + token_payload = _encode_realtime_token_payload( + ephemeral_key="epk", + model_id="gpt-realtime-whisper", + user_id=None, + team_id=None, + expires_at=int(time.time()) + 3600, + session_type="transcription", + ) + encrypted_token = encrypt_value_helper(token_payload) + + client = TestClient(proxy_app) + with ( + patch( + "litellm.proxy.proxy_server.route_request", + side_effect=_capturing_route, + ), + patch( + "litellm.proxy.proxy_server.add_litellm_data_to_request", + side_effect=mock_add_litellm_data, + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + ): + mock_logging.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook) + mock_logging.post_call_failure_hook = AsyncMock() + + client.post( + "/v1/realtime/calls", + headers={"Authorization": f"Bearer {encrypted_token}"}, + content=b"v=0\r\n", + ) + + assert captured["session"]["type"] == "transcription" + assert ( + captured["session"]["audio"]["input"]["transcription"]["model"] + == "gpt-realtime-whisper" + ) + + +# --- transcription_sessions endpoint --- + + +@pytest.fixture +def mock_route_request_transcription_sessions(): + """Mock route_request to return a fake transcription_sessions upstream response.""" + future_expires_at = int(time.time()) + 3600 + body = { + "id": "sess_abc", + "object": "realtime.transcription_session", + "client_secret": { + "value": "upstream_ephemeral_key", + "expires_at": future_expires_at, + }, + } + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 200 + mock_resp.text = json.dumps(body) + mock_resp.content = json.dumps(body).encode() + mock_resp.headers = {} + mock_resp.json.return_value = body + + async def _mock_route(*args, **kwargs): + async def _inner(): + return mock_resp + + return _inner() + + return _mock_route + + +def test_transcription_sessions_requires_auth(proxy_app): + """POST /v1/realtime/transcription_sessions returns 401 without Authorization.""" + from fastapi import HTTPException + + def _raise_401(): + raise HTTPException(status_code=401, detail="Unauthorized") + + proxy_app.dependency_overrides[user_api_key_auth] = _raise_401 + try: + client = TestClient(proxy_app, raise_server_exceptions=False) + response = client.post( + "/v1/realtime/transcription_sessions", + json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}}, + ) + assert response.status_code == 401 + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_transcription_sessions_rejects_disallowed_resolved_model( + proxy_app, +): + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", + models=["gpt-4o-realtime-preview"], + ) + try: + client = TestClient(proxy_app, raise_server_exceptions=False) + with ( + patch("litellm.proxy.proxy_server.route_request") as mock_route_request, + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + ): + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/transcription_sessions", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={ + "input_audio_transcription": {"model": "gpt-realtime-whisper"} + }, + ) + + assert response.status_code == 403 + assert "Tried to access gpt-realtime-whisper" in response.text + mock_route_request.assert_not_called() + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_transcription_sessions_rejects_disallowed_team_model_scope( + proxy_app, +): + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + + team = LiteLLM_TeamTableCachedObj( + team_id="team-a", + models=["gpt-4o-realtime-preview"], + ) + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", + team_id="team-a", + models=["*"], + ) + try: + client = TestClient(proxy_app, raise_server_exceptions=False) + with ( + patch("litellm.proxy.proxy_server.route_request") as mock_route_request, + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + patch( + "litellm.proxy.auth.auth_checks.get_team_object", + new=AsyncMock(return_value=team), + ), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new=AsyncMock(return_value=None), + ), + ): + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/transcription_sessions", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={ + "input_audio_transcription": {"model": "gpt-realtime-whisper"} + }, + ) + + assert response.status_code == 403 + assert "team" in response.text.lower() + assert "Tried to access gpt-realtime-whisper" in response.text + mock_route_request.assert_not_called() + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_transcription_sessions_rejects_disallowed_project_model_scope( + proxy_app, +): + from litellm.proxy._types import LiteLLM_ProjectTableCachedObj + + project = LiteLLM_ProjectTableCachedObj( + project_id="project-a", + models=["gpt-4o-realtime-preview"], + created_by="test-user", + updated_by="test-user", + ) + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", + project_id="project-a", + models=["*"], + ) + try: + client = TestClient(proxy_app, raise_server_exceptions=False) + with ( + patch("litellm.proxy.proxy_server.route_request") as mock_route_request, + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + patch( + "litellm.proxy.auth.auth_checks.get_project_object", + new=AsyncMock(return_value=project), + ), + ): + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/transcription_sessions", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={ + "input_audio_transcription": {"model": "gpt-realtime-whisper"} + }, + ) + + assert response.status_code == 403 + assert "project" in response.text.lower() + assert "Tried to access gpt-realtime-whisper" in response.text + mock_route_request.assert_not_called() + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_transcription_sessions_rejects_disallowed_team_member_model_scope( + proxy_app, +): + from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_TeamMembership, + LiteLLM_TeamTableCachedObj, + ) + + team = LiteLLM_TeamTableCachedObj(team_id="team-a", models=["*"]) + membership = LiteLLM_TeamMembership( + user_id="test-user", + team_id="team-a", + litellm_budget_table=LiteLLM_BudgetTable( + allowed_models=["gpt-4o-realtime-preview"], + ), + ) + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", + team_id="team-a", + models=["*"], + ) + try: + client = TestClient(proxy_app, raise_server_exceptions=False) + with ( + patch("litellm.proxy.proxy_server.route_request") as mock_route_request, + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + patch( + "litellm.proxy.auth.auth_checks.get_team_object", + new=AsyncMock(return_value=team), + ), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new=AsyncMock(return_value=membership), + ), + ): + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/transcription_sessions", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={ + "input_audio_transcription": {"model": "gpt-realtime-whisper"} + }, + ) + + assert response.status_code == 403 + assert "Team member not allowed to access model" in response.text + mock_route_request.assert_not_called() + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_realtime_transcription_websocket_default_model_checks_key_scope(): + from litellm.proxy import proxy_server + + websocket = MagicMock() + websocket.headers = {} + websocket.close = AsyncMock() + websocket.accept = AsyncMock() + + await proxy_server.realtime_websocket_endpoint( + websocket=websocket, + model=None, + intent="transcription", + user_api_key_dict=UserAPIKeyAuth(models=["gpt-4o-realtime-preview"]), + ) + + websocket.accept.assert_not_awaited() + websocket.close.assert_awaited_once() + _, close_kwargs = websocket.close.call_args + assert close_kwargs["code"] == 1008 + assert "not allowed to access model" in close_kwargs["reason"] + + +@pytest.mark.asyncio +async def test_realtime_transcription_websocket_default_model_checks_team_scope(): + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + + team = LiteLLM_TeamTableCachedObj( + team_id="team-a", + models=["gpt-4o-realtime-preview"], + ) + websocket = MagicMock() + websocket.headers = {} + websocket.close = AsyncMock() + websocket.accept = AsyncMock() + + with ( + patch( + "litellm.proxy.auth.auth_checks.get_team_object", + new=AsyncMock(return_value=team), + ), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new=AsyncMock(return_value=None), + ), + ): + await proxy_server.realtime_websocket_endpoint( + websocket=websocket, + model=None, + intent="transcription", + user_api_key_dict=UserAPIKeyAuth( + user_id="test-user", + team_id="team-a", + models=["*"], + ), + ) + + websocket.accept.assert_not_awaited() + websocket.close.assert_awaited_once() + _, close_kwargs = websocket.close.call_args + assert close_kwargs["code"] == 1008 + assert "not allowed to access model" in close_kwargs["reason"] + + +@pytest.mark.asyncio +async def test_transcription_sessions_encrypts_client_secret( + proxy_app, + mock_route_request_transcription_sessions, + mock_add_litellm_data, + mock_pre_call_hook, +): + """ + POST /v1/realtime/transcription_sessions returns 200 and the ephemeral key + under client_secret.value must be encrypted (never the raw upstream key). + """ + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", team_id="test-team" + ) + captured_route_type = {} + + async def _capturing_route(*args, **kwargs): + captured_route_type["route_type"] = kwargs.get("route_type") + return await mock_route_request_transcription_sessions(*args, **kwargs) + + try: + client = TestClient(proxy_app) + with ( + patch( + "litellm.proxy.proxy_server.route_request", + side_effect=_capturing_route, + ), + patch( + "litellm.proxy.proxy_server.add_litellm_data_to_request", + side_effect=mock_add_litellm_data, + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + ): + mock_logging.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook) + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/transcription_sessions", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={ + "input_audio_format": "pcm16", + "input_audio_transcription": {"model": "gpt-realtime-whisper"}, + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["client_secret"]["value"] != "upstream_ephemeral_key" + # The encrypted value must decrypt back to a payload carrying the raw key. + decrypted = decrypt_value_helper( + data["client_secret"]["value"], + key="client_secret.value", + exception_type="debug", + ) + assert decrypted is not None + assert "upstream_ephemeral_key" in decrypted + # Routed through the dedicated transcription_sessions route type. + assert ( + captured_route_type["route_type"] + == "acreate_realtime_transcription_session" + ) + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + +def test_session_type_coerced_for_unknown_value(): + """An unrecognized session_type in the token falls back to 'realtime'.""" + payload = _encode_realtime_token_payload( + ephemeral_key="epk", + model_id="gpt-4o", + user_id=None, + team_id=None, + expires_at=None, + session_type="INJECTED_TYPE", + ) + # Force-deserialize and check the coercion that happens in proxy_realtime_calls. + decoded = json.loads(payload) + session_type = decoded.get("session_type") or "realtime" + if session_type not in ("realtime", "transcription"): + session_type = "realtime" + assert session_type == "realtime" + + +@pytest.mark.asyncio +async def test_transcription_sessions_returns_upstream_error_verbatim( + proxy_app, + mock_add_litellm_data, + mock_pre_call_hook, +): + """Non-200 upstream response is forwarded unchanged (no encryption attempted).""" + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 400 + mock_resp.content = b'{"error":"bad_request"}' + mock_resp.headers = {} + mock_resp.json.return_value = {"error": "bad_request"} + mock_resp.text = '{"error":"bad_request"}' + + async def _mock_route(*args, **kwargs): + async def _inner(): + return mock_resp + + return _inner() + + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", team_id="test-team" + ) + try: + client = TestClient(proxy_app) + with ( + patch( + "litellm.proxy.proxy_server.route_request", + side_effect=_mock_route, + ), + patch( + "litellm.proxy.proxy_server.add_litellm_data_to_request", + side_effect=mock_add_litellm_data, + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + ): + mock_logging.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook) + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/transcription_sessions", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}}, + ) + assert response.status_code == 400 + assert response.content == b'{"error":"bad_request"}' + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_transcription_sessions_wraps_route_exception( + proxy_app, + mock_add_litellm_data, + mock_pre_call_hook, +): + """A route exception is wrapped in a ProxyException with a human-readable message.""" + from fastapi import HTTPException + + async def _raise_http(*args, **kwargs): + raise HTTPException(status_code=403, detail="Model not allowed") + + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user" + ) + try: + client = TestClient(proxy_app, raise_server_exceptions=False) + with ( + patch( + "litellm.proxy.proxy_server.route_request", + side_effect=_raise_http, + ), + patch( + "litellm.proxy.proxy_server.add_litellm_data_to_request", + side_effect=mock_add_litellm_data, + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + ): + mock_logging.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook) + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/transcription_sessions", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}}, + ) + assert response.status_code == 403 + assert "Model not allowed" in response.text + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 2632d8af4f1..3a1d15ef79c 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -2629,7 +2629,8 @@ async def test_ui_view_spend_logs_with_error_code(client): assert data["total"] == 1 assert len(data["data"]) == 1 assert data["data"][0]["id"] == "log1" - metadata = json.loads(data["data"][0]["metadata"]) + metadata = data["data"][0]["metadata"] + assert isinstance(metadata, dict) assert "error_information" in metadata assert metadata["error_information"]["error_code"] == "404" finally: @@ -2702,7 +2703,8 @@ async def test_ui_view_spend_logs_with_error_message(client): assert data["total"] == 1 assert len(data["data"]) == 1 assert data["data"][0]["id"] == "log1" - metadata = json.loads(data["data"][0]["metadata"]) + metadata = data["data"][0]["metadata"] + assert isinstance(metadata, dict) assert "error_information" in metadata assert ( "Rate limit exceeded" in metadata["error_information"]["error_message"] @@ -2794,7 +2796,8 @@ async def test_ui_view_spend_logs_with_error_code_and_key_alias(client): assert data["total"] == 1 assert len(data["data"]) == 1 assert data["data"][0]["id"] == "log3" - metadata = json.loads(data["data"][0]["metadata"]) + metadata = data["data"][0]["metadata"] + assert isinstance(metadata, dict) assert "user_api_key_alias" in metadata assert metadata["user_api_key_alias"] == "test-key-1" assert "error_information" in metadata @@ -3606,3 +3609,178 @@ async def test_spend_user_fn_strips_password_field(client, monkeypatch): assert "password" not in body[0] finally: app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_spend_logs_rehydrates_metadata_jsonb_text(client, monkeypatch): + """ + Regression for #29674: query_raw returns the JSONB `metadata` column as a + string, so failure rows (status="failure", error_information.error_code=...) + looked like successes at the UI layer because metadata.status was the + string ".status" attribute lookup on a str. The endpoint must re-hydrate + `metadata` to a dict before returning. + """ + failure_metadata = { + "status": "failure", + "error_information": { + "error_code": "403", + "error_message": "Forbidden by upstream", + }, + "user_api_key_alias": "alias-1", + } + + raw_row = { + "request_id": "req-failure-1", + "call_type": "completion", + "api_key": "hashed-key", + "spend": 0.0, + "total_tokens": 0, + "prompt_tokens": 0, + "completion_tokens": 0, + "startTime": "2025-01-01T00:00:00Z", + "endTime": "2025-01-01T00:00:01Z", + "completionStartTime": None, + "model": "gpt-4o", + "model_id": None, + "model_group": None, + "custom_llm_provider": "openai", + "api_base": None, + "user": "u", + "metadata": json.dumps(failure_metadata), # JSONB column comes back as str + "cache_hit": None, + "cache_key": None, + "request_tags": None, + "team_id": None, + "organization_id": None, + "end_user": None, + "requester_ip_address": None, + "session_id": None, + "status": "failure", + "mcp_namespaced_tool_name": None, + "agent_id": None, + "request_duration_ms": 1000, + } + + async def mock_count(*args, **kwargs): + return 1 + + async def mock_query_raw(sql_query, *params): + return [raw_row] + + class MockPrismaClient: + def __init__(self): + self.db = MagicMock() + self.db.litellm_spendlogs = MagicMock() + self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) + self.db.query_raw = AsyncMock(side_effect=mock_query_raw) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + response = client.get( + "/spend/logs/ui", + params={ + "start_date": "2024-12-25 00:00:00", + "end_date": "2025-01-02 23:59:59", + }, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 200, response.text + body = response.json() + assert body["data"], "expected one row in data" + row = body["data"][0] + md = row["metadata"] + # The bug had metadata returned as a JSON string; the fix re-hydrates + # it so the dashboard's metadata.status / metadata.error_information + # accessors work. + assert isinstance(md, dict), f"metadata should be dict, got {type(md)}" + assert md["status"] == "failure" + assert md["error_information"]["error_code"] == "403" + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_spend_logs_metadata_invalid_json_falls_back_to_empty_dict( + client, monkeypatch +): + """ + Defensive: if `metadata` is somehow not valid JSON, fall back to {} rather + than 500-ing the whole UI page. + """ + raw_row = { + "request_id": "req-bad-json", + "call_type": "completion", + "api_key": "hashed-key", + "spend": 0.0, + "total_tokens": 0, + "prompt_tokens": 0, + "completion_tokens": 0, + "startTime": "2025-01-01T00:00:00Z", + "endTime": "2025-01-01T00:00:01Z", + "completionStartTime": None, + "model": "gpt-4o", + "model_id": None, + "model_group": None, + "custom_llm_provider": "openai", + "api_base": None, + "user": "u", + "metadata": "{not-json", + "cache_hit": None, + "cache_key": None, + "request_tags": None, + "team_id": None, + "organization_id": None, + "end_user": None, + "requester_ip_address": None, + "session_id": None, + "status": "success", + "mcp_namespaced_tool_name": None, + "agent_id": None, + "request_duration_ms": 500, + } + + async def mock_count(*args, **kwargs): + return 1 + + async def mock_query_raw(sql_query, *params): + return [raw_row] + + class MockPrismaClient: + def __init__(self): + self.db = MagicMock() + self.db.litellm_spendlogs = MagicMock() + self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) + self.db.query_raw = AsyncMock(side_effect=mock_query_raw) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + response = client.get( + "/spend/logs/ui", + params={ + "start_date": "2024-12-25 00:00:00", + "end_date": "2025-01-02 23:59:59", + }, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 200, response.text + body = response.json() + assert body["data"] + assert body["data"][0]["metadata"] == {} + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) diff --git a/tests/test_litellm/proxy/test_model_list_healthy_only.py b/tests/test_litellm/proxy/test_model_list_healthy_only.py new file mode 100644 index 00000000000..4ab33f3bf50 --- /dev/null +++ b/tests/test_litellm/proxy/test_model_list_healthy_only.py @@ -0,0 +1,92 @@ +""" +Tests for the opt-in `healthy_only` filter on GET /v1/models (`model_list`). +""" + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import proxy_server +from litellm.proxy._types import UserAPIKeyAuth + + +@pytest.fixture +def patched_model_list(monkeypatch): + """Stub router + utility helpers used by `model_list`.""" + from litellm.proxy import utils as proxy_utils + + router = MagicMock() + router.get_fully_blocked_model_names = MagicMock(return_value=set()) + router.async_get_fully_unhealthy_model_names = AsyncMock( + return_value={"claude-sonnet"} + ) + + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "user_model", None) + + async def _fake_get_available_models_for_user(**kwargs): + return ["gpt-4", "claude-sonnet"] + + monkeypatch.setattr( + proxy_utils, + "get_available_models_for_user", + _fake_get_available_models_for_user, + ) + + def _fake_create_model_info_response(model_id, provider="openai", **kwargs): + return {"id": model_id, "object": "model", "created": 0, "owned_by": provider} + + monkeypatch.setattr( + proxy_utils, "create_model_info_response", _fake_create_model_info_response + ) + + return router + + +@pytest.mark.asyncio +async def test_model_list_healthy_only_hides_fully_unhealthy_models( + patched_model_list, +): + response = await proxy_server.model_list( + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + healthy_only=True, + ) + assert [m["id"] for m in response["data"]] == ["gpt-4"] + + +@pytest.mark.asyncio +async def test_model_list_default_keeps_unhealthy_models(patched_model_list): + response = await proxy_server.model_list( + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + assert [m["id"] for m in response["data"]] == ["gpt-4", "claude-sonnet"] + patched_model_list.async_get_fully_unhealthy_model_names.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_model_list_healthy_only_applies_to_scope_expand( + patched_model_list, monkeypatch +): + from litellm.proxy.auth import model_checks + from litellm.proxy.management_endpoints import common_utils + + async def _fake_admin(**kwargs): + return True + + monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _fake_admin) + monkeypatch.setattr( + model_checks, + "get_complete_model_list", + lambda **kwargs: ["gpt-4", "claude-sonnet"], + ) + patched_model_list.get_model_names = MagicMock( + return_value=["gpt-4", "claude-sonnet"] + ) + patched_model_list.get_model_access_groups = MagicMock(return_value={}) + + response = await proxy_server.model_list( + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + scope="expand", + healthy_only=True, + ) + assert [m["id"] for m in response["data"]] == ["gpt-4"] diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index f0015d9df0d..aace9405292 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -118,6 +118,63 @@ async def test_proxy_only_error_log_marks_no_upstream_llm_call(): assert captured.get("flag") is True +@pytest.mark.asyncio +async def test_proxy_only_error_log_keeps_litellm_metadata_in_litellm_params(): + """Responses API requests carry guardrail info under ``litellm_metadata`` + (not ``metadata``). It must land in litellm_params so + ``merge_litellm_metadata`` can surface ``guardrail_information`` in the + spend-log failure row, matching the chat completions path.""" + from litellm.proxy._types import UserAPIKeyAuth + + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + captured = {} + guardrail_info = [{"guardrail_name": "test-guard", "guardrail_status": "blocked"}] + + def fake_update_environment_variables(self, *args, **kwargs): + captured["litellm_params"] = kwargs.get("litellm_params") + captured["optional_params"] = kwargs.get("optional_params") + + from litellm.litellm_core_utils.litellm_logging import Logging + + orig_update_env = Logging.update_environment_variables + orig_pre_call = Logging.pre_call + orig_async_failure = Logging.async_failure_handler + + async def _noop_async_failure(self, *args, **kwargs): + return None + + Logging.update_environment_variables = fake_update_environment_variables + Logging.pre_call = lambda self, *args, **kwargs: None + Logging.async_failure_handler = _noop_async_failure + try: + await proxy_logging_obj._handle_logging_proxy_only_error( + request_data={ + "model": "gpt-4o", + "input": "blocked prompt", + "litellm_metadata": { + "standard_logging_guardrail_information": guardrail_info + }, + }, + user_api_key_dict=UserAPIKeyAuth( + api_key="sk-1234", request_route="/v1/responses" + ), + route="/v1/responses", + original_exception=HTTPException(status_code=400, detail="blocked"), + ) + finally: + Logging.update_environment_variables = orig_update_env + Logging.pre_call = orig_pre_call + Logging.async_failure_handler = orig_async_failure + + assert ( + captured["litellm_params"]["litellm_metadata"][ + "standard_logging_guardrail_information" + ] + == guardrail_info + ) + assert "litellm_metadata" not in captured["optional_params"] + + def test_get_model_group_info_order(): from litellm import Router from litellm.proxy.proxy_server import _get_model_group_info diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 82a4a60bf82..ad08029c2c4 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -463,12 +463,154 @@ def test_realtime_logging_object_allows_null_transcript_in_conversation_item_add usage=usage, results=results, ) - assert logging_result.usage.total_tokens == 18 assert logging_result.results[0]["item"]["content"][0]["transcript"] is None + assert logging_result.results[0]["item"]["content"][0]["transcript"] is None -def test_custom_pricing_with_router_model_id(): +def test_realtime_transcription_duration_cost(monkeypatch): + """ + gpt-realtime-whisper transcription sessions are billed by input audio duration + ($0.017/min). The .completed events carry usage {type: duration, seconds: N}; + cost must equal total_seconds * input_cost_per_second. + """ + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + from litellm.cost_calculator import RealtimeAPITokenUsageProcessor + + results: OpenAIRealtimeStreamList = [ + { + "type": "session.created", + "session": { + "type": "transcription", + "audio": { + "input": {"transcription": {"model": "gpt-realtime-whisper"}} + }, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "transcript": "hello", + "usage": {"type": "duration", "seconds": 60.0}, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "transcript": "world", + "usage": {"type": "duration", "seconds": 30.0}, + }, + ] + + combined = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results( + results=results + ) + cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=combined, + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + ) + + # 90 seconds at $0.017/minute. + expected = 90.0 * (0.017 / 60) + assert abs(cost - expected) < 1e-9 + assert cost > 0 # guards against the duration branch being dropped + + +def test_realtime_transcription_duration_cost_resolves_model_from_litellm_name( + monkeypatch, +): + """When no session event carries the ASR model, the litellm_model_name is used.""" + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + results: OpenAIRealtimeStreamList = [ + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": {"type": "duration", "seconds": 120.0}, + }, + ] + cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="azure", + litellm_model_name="azure/gpt-realtime-whisper", + ) + assert abs(cost - 120.0 * (0.017 / 60)) < 1e-9 + + +def test_realtime_transcription_no_completed_events_is_zero(monkeypatch): + """A realtime stream without transcription completed events adds no extra cost.""" + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + from litellm.cost_calculator import handle_realtime_transcription_cost_calculation + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": "gpt-realtime-whisper"}}, + {"type": "response.done", "response": {"usage": {}}}, + ] + assert ( + handle_realtime_transcription_cost_calculation( + results=results, + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + ) + == 0.0 + ) + + +def test_realtime_transcription_token_billed_fallback(monkeypatch): + """ + Token-billed transcription models price by audio/text tokens. Verify the + fallback path multiplies audio tokens by the model's audio token cost. + """ + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + from litellm.cost_calculator import _transcription_usage_cost + + # gpt-4o-transcribe: input_cost_per_audio_token = 2.5e-06, input_cost_per_token = 2.5e-06, + # output_cost_per_token = 1e-05 + model_info = litellm.get_model_info( + model="gpt-4o-transcribe", custom_llm_provider="openai" + ) + usage = { + "type": "tokens", + "input_tokens": 40, + "output_tokens": 10, + "total_tokens": 50, + "input_token_details": {"audio_tokens": 30, "text_tokens": 10}, + } + cost = _transcription_usage_cost(usage, model_info) + expected = ( + 30 * 2.5e-06 # audio tokens + + 10 * 2.5e-06 # text tokens + + 10 * 1e-05 # output tokens + ) + assert abs(cost - expected) < 1e-12 + + +def test_transcription_usage_cost_returns_zero_for_unknown_type(): + """An unrecognized usage type yields 0 (safe fallback, no exception).""" + from litellm.cost_calculator import _transcription_usage_cost + + assert _transcription_usage_cost({"type": "future_billing_type"}, {}) == 0.0 + assert _transcription_usage_cost({}, {}) == 0.0 + + +def test_get_transcription_model_falls_back_to_session_model(monkeypatch): + """session.model is used when transcription-specific model fields are absent.""" + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + from litellm.cost_calculator import _get_transcription_model_name_from_results + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": "gpt-realtime-whisper"}}, + ] + assert _get_transcription_model_name_from_results(results) == "gpt-realtime-whisper" + from litellm import Router router = Router( diff --git a/tests/test_litellm/test_model_block_unblock.py b/tests/test_litellm/test_model_block_unblock.py new file mode 100644 index 00000000000..318b2f519c4 --- /dev/null +++ b/tests/test_litellm/test_model_block_unblock.py @@ -0,0 +1,199 @@ +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm.proxy._types import ( + BlockModelRequest, + LitellmUserRoles, + ProxyException, + UserAPIKeyAuth, +) +from litellm.types.router import RouterRateLimitError + + +def _setup_model_block_mocks(monkeypatch, *, updated_blocked: bool): + model_id = "model-123" + + existing_row = MagicMock() + existing_row.model_dump.return_value = { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": model_id}, + } + + updated_row = MagicMock() + updated_row.model_id = model_id + updated_row.blocked = updated_blocked + + model_table = MagicMock() + model_table.find_unique = AsyncMock(return_value=existing_row) + model_table.update = AsyncMock(return_value=updated_row) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_proxymodeltable = model_table + + mock_router = MagicMock() + mock_router.get_deployment.return_value = None + + mock_clear_cache = AsyncMock(return_value=None) + mock_audit_log = AsyncMock(return_value=None) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) + monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") + monkeypatch.setattr( + "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", + mock_clear_cache, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.model_management_endpoints.create_object_audit_log", + mock_audit_log, + ) + + return model_id, model_table, updated_row, mock_clear_cache, mock_audit_log + + +def _proxy_admin() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + user_id="admin", + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + ) + + +@pytest.mark.asyncio +async def test_model_block_endpoint_sets_blocked_true(monkeypatch): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + block_model, + ) + + model_id, model_table, updated_row, mock_clear_cache, mock_audit_log = ( + _setup_model_block_mocks(monkeypatch, updated_blocked=True) + ) + + result = await block_model( + data=BlockModelRequest(model_id=model_id), + http_request=MagicMock(), + user_api_key_dict=_proxy_admin(), + litellm_changed_by="operator@example.com", + ) + + assert result == updated_row + model_table.update.assert_awaited_once() + update_kwargs = model_table.update.await_args.kwargs + assert update_kwargs["where"] == {"model_id": model_id} + assert update_kwargs["data"]["blocked"] is True + assert update_kwargs["data"]["updated_by"] == "admin" + assert "updated_at" in update_kwargs["data"] + mock_clear_cache.assert_awaited_once_with() + assert mock_audit_log.call_args.kwargs["action"] == "blocked" + assert ( + mock_audit_log.call_args.kwargs["litellm_changed_by"] == "operator@example.com" + ) + + +@pytest.mark.asyncio +async def test_model_unblock_endpoint_sets_blocked_false(monkeypatch): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + unblock_model, + ) + + model_id, model_table, updated_row, mock_clear_cache, mock_audit_log = ( + _setup_model_block_mocks(monkeypatch, updated_blocked=False) + ) + + result = await unblock_model( + data=BlockModelRequest(model_id=model_id), + http_request=MagicMock(), + user_api_key_dict=_proxy_admin(), + litellm_changed_by=None, + ) + + assert result == updated_row + model_table.update.assert_awaited_once() + assert model_table.update.await_args.kwargs["data"]["blocked"] is False + mock_clear_cache.assert_awaited_once_with() + assert mock_audit_log.call_args.kwargs["action"] == "unblocked" + + +@pytest.mark.asyncio +async def test_model_block_endpoint_requires_proxy_admin(monkeypatch): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + block_model, + ) + + model_id, model_table, _, _, _ = _setup_model_block_mocks( + monkeypatch, updated_blocked=True + ) + non_admin = UserAPIKeyAuth( + user_id="internal-user", + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-user", + ) + + with pytest.raises(ProxyException) as exc_info: + await block_model( + data=BlockModelRequest(model_id=model_id), + http_request=MagicMock(), + user_api_key_dict=non_admin, + litellm_changed_by=None, + ) + + assert exc_info.value.code == "403" + assert "Only proxy admins" in exc_info.value.message + model_table.update.assert_not_awaited() + + +def test_router_returns_no_healthy_deployment_when_model_is_fully_blocked(): + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o-0"}, + "model_info": {"id": "dep-0", "blocked": True}, + }, + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o-1"}, + "model_info": {"id": "dep-1", "blocked": True}, + }, + ] + ) + + with pytest.raises(RouterRateLimitError) as exc_info: + router.get_available_deployment(model="gpt-4o", request_kwargs={}) + + assert "No deployments available for selected model" in str(exc_info.value) + assert "Passed model=gpt-4o" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_route_request_returns_403_when_model_is_fully_blocked(monkeypatch): + from litellm.proxy.route_llm_request import route_request + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": "dep-0", "blocked": True}, + } + ] + ) + monkeypatch.setattr( + "litellm.proxy.route_llm_request.add_shared_session_to_data", + AsyncMock(return_value=None), + ) + + with pytest.raises(litellm.PermissionDeniedError) as exc_info: + await route_request( + data={"model": "gpt-4o"}, + llm_router=router, + user_model=None, + route_type="acreate_eval", + ) + + assert exc_info.value.status_code == 403 + assert "Model is blocked" in exc_info.value.message diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index e681247959f..bd19f951553 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -4468,6 +4468,82 @@ def test_get_fully_blocked_model_names_treats_missing_key_as_unblocked(): assert router.get_fully_blocked_model_names() == set() +def _seed_unhealthy_states(router, unhealthy_ids, timestamp=None): + import time + + ts = timestamp if timestamp is not None else time.time() + router.health_state_cache.set_deployment_health_states( + { + uid: {"is_healthy": False, "timestamp": ts, "reason": "test_unhealthy"} + for uid in unhealthy_ids + } + ) + + +@pytest.mark.asyncio +async def test_async_get_fully_unhealthy_model_names_marks_name_when_all_unhealthy(): + router = _router_with_two_deployments([False, False]) + _seed_unhealthy_states(router, {"dep-0", "dep-1"}) + assert await router.async_get_fully_unhealthy_model_names() == {"gpt-4o"} + + +@pytest.mark.asyncio +async def test_async_get_fully_unhealthy_model_names_keeps_name_when_partial(): + router = _router_with_two_deployments([False, False]) + _seed_unhealthy_states(router, {"dep-0"}) + assert await router.async_get_fully_unhealthy_model_names() == set() + + +@pytest.mark.asyncio +async def test_async_get_fully_unhealthy_model_names_empty_without_health_state(): + router = _router_with_two_deployments([False, False]) + assert await router.async_get_fully_unhealthy_model_names() == set() + + +@pytest.mark.asyncio +async def test_async_get_fully_unhealthy_model_names_ignores_stale_state(): + import time + + router = _router_with_two_deployments([False, False]) + stale_ts = time.time() - (router.health_state_cache.staleness_threshold + 10) + _seed_unhealthy_states(router, {"dep-0", "dep-1"}, timestamp=stale_ts) + assert await router.async_get_fully_unhealthy_model_names() == set() + + +@pytest.mark.asyncio +async def test_async_get_fully_unhealthy_model_names_includes_team_alias(): + import litellm + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": { + "id": "dep-0", + "team_id": "team-1", + "team_public_model_name": "team-gpt", + }, + } + ] + ) + _seed_unhealthy_states(router, {"dep-0"}) + assert await router.async_get_fully_unhealthy_model_names() == { + "gpt-4o", + "team-gpt", + } + + +@pytest.mark.asyncio +async def test_async_get_fully_unhealthy_model_names_noop_with_allowed_fails_policy(): + from litellm.types.router import AllowedFailsPolicy + + router = _router_with_two_deployments([False, False]) + router.allowed_fails_policy = AllowedFailsPolicy(BadRequestErrorAllowedFails=1) + _seed_unhealthy_states(router, {"dep-0", "dep-1"}) + assert await router.async_get_fully_unhealthy_model_names() == set() + + @pytest.mark.asyncio async def test_async_get_healthy_deployments_skips_blocked_deployment(): router = _router_with_two_deployments([True, False]) diff --git a/tests/test_litellm/test_router_block_helpers.py b/tests/test_litellm/test_router_block_helpers.py new file mode 100644 index 00000000000..443209bfe2f --- /dev/null +++ b/tests/test_litellm/test_router_block_helpers.py @@ -0,0 +1,57 @@ +"""Unit tests for Router block helper methods (coverage gate).""" + +from litellm import Router + + +def _make_router(model_name: str, blocked: bool = False) -> Router: + return Router( + model_list=[ + { + "model_name": model_name, + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}, + "model_info": {"blocked": blocked}, + } + ] + ) + + +class TestAreAllDeploymentsBlocked: + def test_all_blocked_returns_true(self): + router = _make_router("gpt-4o", blocked=True) + deployments = router.get_model_list(model_name="gpt-4o") or [] + assert router._are_all_deployments_blocked(deployments) is True + + def test_one_not_blocked_returns_false(self): + router = Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}, + "model_info": {"blocked": True}, + }, + { + "model_name": "gpt-4o", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "fake", + }, + "model_info": {"blocked": False}, + }, + ] + ) + deployments = router.get_model_list(model_name="gpt-4o") or [] + assert router._are_all_deployments_blocked(deployments) is False + + def test_empty_list_returns_false(self): + router = _make_router("gpt-4o") + assert router._are_all_deployments_blocked([]) is False + + +class TestIsModelFullyBlocked: + def test_all_deployments_blocked_returns_true(self): + router = _make_router("gpt-4o", blocked=True) + assert router._is_model_fully_blocked("gpt-4o") is True + + def test_unblocked_deployment_returns_false(self): + router = _make_router("gpt-4o", blocked=False) + assert router._is_model_fully_blocked("gpt-4o") is False diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index b6c9e9c865d..a59f3674da2 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -882,6 +882,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "/v1/completions", "/v1/images/generations", "/v1/realtime", + "/v1/realtime/transcription_sessions", "/v1/images/variations", "/v1/images/edits", "/v1/batch", diff --git a/ui/litellm-dashboard/public/assets/logos/newrelic.png b/ui/litellm-dashboard/public/assets/logos/newrelic.png new file mode 100644 index 00000000000..c841e3e7136 Binary files /dev/null and b/ui/litellm-dashboard/public/assets/logos/newrelic.png differ diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 4aeb1cf3e86..3786b6e728b 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -7250,6 +7250,29 @@ export interface paths { patch: operations["mistral_proxy_route_mistral__endpoint__patch"]; trace?: never; }; + "/model/block": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** + * Block Model + * @description Block a DB-stored model deployment from serving requests. + * + * Parameters: + * - model_id: str - The model deployment id to block. + */ + post: operations["block_model_model_block_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/model/cost_map/source": { parameters: { query?: never; @@ -7468,6 +7491,29 @@ export interface paths { patch?: never; trace?: never; }; + "/model/unblock": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** + * Unblock Model + * @description Unblock a DB-stored model deployment so it can serve requests again. + * + * Parameters: + * - model_id: str - The model deployment id to unblock. + */ + post: operations["unblock_model_model_unblock_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/model/update": { parameters: { query?: never; @@ -7741,6 +7787,15 @@ export interface paths { * - scope: Optional scope parameter. Currently only accepts "expand". * When scope=expand is passed, proxy admins, team admins, and org admins * will receive all proxy models as if they are a proxy admin. + * - healthy_only: When true, hide models whose backing deployments are all marked + * unhealthy by background health checks. Requires + * `background_health_checks: true` in general_settings; without + * health state the listing is returned unfiltered (fail open). + * Models expanded from wildcard routes (e.g. `openai/*`) are not + * filtered, and nothing is hidden when `allowed_fails_policy` is + * configured (cooldown remains the sole exclusion mechanism). + * Hiding is presentation-only: a hidden model can still be + * called directly. */ get: operations["model_list_models_get"]; put?: never; @@ -16563,6 +16618,15 @@ export interface paths { * - scope: Optional scope parameter. Currently only accepts "expand". * When scope=expand is passed, proxy admins, team admins, and org admins * will receive all proxy models as if they are a proxy admin. + * - healthy_only: When true, hide models whose backing deployments are all marked + * unhealthy by background health checks. Requires + * `background_health_checks: true` in general_settings; without + * health state the listing is returned unfiltered (fail open). + * Models expanded from wildcard routes (e.g. `openai/*`) are not + * filtered, and nothing is hidden when `allowed_fails_policy` is + * configured (cooldown remains the sole exclusion mechanism). + * Hiding is presentation-only: a hidden model can still be + * called directly. */ get: operations["model_list_v1_models_get"]; put?: never; @@ -20703,6 +20767,11 @@ export interface components { /** Key */ key: string; }; + /** BlockModelRequest */ + BlockModelRequest: { + /** Model Id */ + model_id: string; + }; /** BlockTeamRequest */ BlockTeamRequest: { /** Team Id */ @@ -25173,6 +25242,34 @@ export interface components { /** Updated By */ updated_by?: string | null; }; + /** LiteLLM_ProxyModelTable */ + LiteLLM_ProxyModelTable: { + /** + * Blocked + * @default false + */ + blocked: boolean; + /** Created At */ + created_at?: string | null; + /** Created By */ + created_by?: string | null; + /** Litellm Params */ + litellm_params: { + [key: string]: unknown; + }; + /** Model Id */ + model_id: string; + /** Model Info */ + model_info?: { + [key: string]: unknown; + } | null; + /** Model Name */ + model_name: string; + /** Updated At */ + updated_at?: string | null; + /** Updated By */ + updated_by?: string | null; + }; /** LiteLLM_SpendLogs */ LiteLLM_SpendLogs: { /** @@ -40448,7 +40545,7 @@ export interface operations { parameters: { query: { /** @description Specify the service being hit. */ - service: ("slack_budget_alerts" | "langfuse" | "langfuse_otel" | "slack" | "openmeter" | "webhook" | "email" | "braintrust" | "datadog" | "datadog_llm_observability" | "generic_api" | "arize" | "galileo" | "sqs") | string; + service: ("slack_budget_alerts" | "langfuse" | "langfuse_otel" | "slack" | "openmeter" | "webhook" | "email" | "braintrust" | "datadog" | "datadog_llm_observability" | "generic_api" | "arize" | "galileo" | "newrelic" | "sqs") | string; }; header?: never; path?: never; @@ -42240,6 +42337,42 @@ export interface operations { }; }; }; + block_model_model_block_post: { + parameters: { + query?: never; + header?: { + /** @description The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability */ + "litellm-changed-by"?: string | null; + }; + path?: never; + cookie?: never; + }; + requestBody: { + content: { + "application/json": components["schemas"]["BlockModelRequest"]; + }; + }; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["LiteLLM_ProxyModelTable"] | null; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; get_model_cost_map_source_model_cost_map_source_get: { parameters: { query?: never; @@ -42515,6 +42648,42 @@ export interface operations { }; }; }; + unblock_model_model_unblock_post: { + parameters: { + query?: never; + header?: { + /** @description The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability */ + "litellm-changed-by"?: string | null; + }; + path?: never; + cookie?: never; + }; + requestBody: { + content: { + "application/json": components["schemas"]["BlockModelRequest"]; + }; + }; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["LiteLLM_ProxyModelTable"] | null; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; update_model_model_update_post: { parameters: { query?: never; @@ -42690,6 +42859,7 @@ export interface operations { include_metadata?: boolean | null; fallback_type?: string | null; scope?: string | null; + healthy_only?: boolean | null; }; header?: never; path?: never; @@ -53562,6 +53732,7 @@ export interface operations { include_metadata?: boolean | null; fallback_type?: string | null; scope?: string | null; + healthy_only?: boolean | null; }; header?: never; path?: never;