merge: preserve realtime authentication with upstream main

This commit is contained in:
jibanez-staticduo 2026-09-19 00:07:31 +02:00
commit 906d5997bf
No known key found for this signature in database
61 changed files with 4083 additions and 638 deletions

View file

@ -95,6 +95,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/vertex-ai/",
"/assemblyai/",
"/eu.assemblyai/",
"/deepgram/",
"/langfuse/",
"/vllm/",
"/mistral/",

View file

@ -0,0 +1,35 @@
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_DailyGlobalSpend" (
"id" TEXT NOT NULL,
"date" TEXT NOT NULL,
"model" TEXT,
"model_group" TEXT,
"custom_llm_provider" TEXT,
"mcp_namespaced_tool_name" TEXT,
"endpoint" TEXT,
"prompt_tokens" BIGINT NOT NULL DEFAULT 0,
"completion_tokens" BIGINT NOT NULL DEFAULT 0,
"cache_read_input_tokens" BIGINT NOT NULL DEFAULT 0,
"cache_creation_input_tokens" BIGINT NOT NULL DEFAULT 0,
"compression_saved_tokens" BIGINT NOT NULL DEFAULT 0,
"compression_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"prompt_caching_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"gateway_injected_caching_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"autorouter_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"api_requests" BIGINT NOT NULL DEFAULT 0,
"successful_requests" BIGINT NOT NULL DEFAULT 0,
"failed_requests" BIGINT NOT NULL DEFAULT 0,
"total_response_time_ms" BIGINT NOT NULL DEFAULT 0,
"timed_requests" BIGINT NOT NULL DEFAULT 0,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
CONSTRAINT "LiteLLM_DailyGlobalSpend_pkey" PRIMARY KEY ("id")
);
-- CreateIndex
CREATE INDEX IF NOT EXISTS "LiteLLM_DailyGlobalSpend_date_idx" ON "LiteLLM_DailyGlobalSpend"("date");
-- CreateIndex
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_DailyGlobalSpend_date_model_model_group_custom_llm__key" ON "LiteLLM_DailyGlobalSpend"("date", "model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint");

View file

@ -820,6 +820,37 @@ model LiteLLM_DailyUserSpend {
@@index([endpoint])
}
// Key-free daily rollup of LiteLLM_DailyUserSpend, read by the global usage view
model LiteLLM_DailyGlobalSpend {
id String @id @default(uuid())
date String
model String?
model_group String?
custom_llm_provider String?
mcp_namespaced_tool_name String?
endpoint String?
prompt_tokens BigInt @default(0)
completion_tokens BigInt @default(0)
cache_read_input_tokens BigInt @default(0)
cache_creation_input_tokens BigInt @default(0)
compression_saved_tokens BigInt @default(0)
compression_savings_spend Float @default(0.0)
prompt_caching_savings_spend Float @default(0.0)
gateway_injected_caching_savings_spend Float @default(0.0)
autorouter_savings_spend Float @default(0.0)
spend Float @default(0.0)
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@unique([date, model, model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
}
// Track daily organization spend metrics per model and key
model LiteLLM_DailyOrganizationSpend {
id String @id @default(uuid())

View file

@ -1700,6 +1700,9 @@ if TYPE_CHECKING:
from .llms.vertex_ai.vertex_ai_partner_models.ai21.transformation import (
VertexAIAi21Config as VertexAIAi21Config,
)
from .llms.vertex_ai.vertex_ai_partner_models.mistral.transformation import (
VertexAIMistralConfig as VertexAIMistralConfig,
)
from .llms.bedrock.chat.invoke_handler import (
AmazonCohereChatConfig as AmazonCohereChatConfig,
)

View file

@ -184,6 +184,7 @@ LLM_CONFIG_NAMES: Final = (
"VertexAIAnthropicConfig",
"VertexAILlama3Config",
"VertexAIAi21Config",
"VertexAIMistralConfig",
"AmazonCohereChatConfig",
"AmazonBedrockGlobalConfig",
"AmazonAI21Config",
@ -771,6 +772,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
".llms.vertex_ai.vertex_ai_partner_models.ai21.transformation",
"VertexAIAi21Config",
),
"VertexAIMistralConfig": (
".llms.vertex_ai.vertex_ai_partner_models.mistral.transformation",
"VertexAIMistralConfig",
),
"AmazonCohereChatConfig": (
".llms.bedrock.chat.invoke_handler",
"AmazonCohereChatConfig",

View file

@ -533,6 +533,8 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
"vertex_credentials",
"gcs_bucket_name",
"bucket_name",
"s3_endpoint_url",
"s3_region_name",
"timeout",
"max_retries",
"_litellm_internal_model_credentials",

View file

@ -320,6 +320,9 @@ REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float(
# RFC 6455 caps the close frame payload at 125 bytes, 2 of which carry the status code
WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
DEEPGRAM_DEFAULT_API_BASE: Final = "https://api.deepgram.com/v1"
DEEPGRAM_LISTEN_DEFAULT_MODEL: Final = "nova-3"
BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY: Final = "litellm.bedrock_realtime.pending_session_update"
BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY: Final = "litellm.bedrock_realtime.session_committed"
BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY: Final = "litellm.bedrock_realtime.committed_failure"
@ -2079,6 +2082,9 @@ PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90
# Deployments named in the lapsed-window alert before it is truncated, so a fleet-wide
# expiry cannot produce an alert too large for the channel delivering it.
PTU_LAPSED_ALERT_LIMIT: Final[int] = 10
DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID: Final[str] = "daily_global_spend_reconcile_job"
DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS: Final[int] = 3600
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM: Final[str] = "daily_global_spend_reconciled_through"
# Slack allowed when deciding a sentinel row is stale. The row's updated_at and the
# run's cutoff are stamped by different hosts, so clock skew between them must not let
# one run delete a charge another just wrote. A stale row is hours old and a concurrent

View file

@ -44,6 +44,8 @@ OPTIONAL_KWARGS_KEYS: Final = (
"client_side_timeout",
"gcs_bucket_name",
"bucket_name",
"s3_endpoint_url",
"s3_region_name",
"vertex_credentials",
"vertex_project",
"vertex_location",

View file

@ -190,7 +190,7 @@ def get_supported_openai_params(
elif custom_llm_provider == "vertex_ai" or custom_llm_provider == "vertex_ai_beta":
if request_type == "chat_completion":
if model.startswith("mistral"):
return litellm.MistralConfig().get_supported_openai_params(model=model)
return litellm.VertexAIMistralConfig().get_supported_openai_params(model=model)
elif model.startswith("codestral"):
return litellm.CodestralTextCompletionConfig().get_supported_openai_params(model=model)
elif model.startswith("claude"):

View file

@ -1,5 +1,200 @@
import math
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final
from urllib.parse import parse_qs, urlparse
import httpx
import litellm
from litellm.constants import DEEPGRAM_DEFAULT_API_BASE, DEEPGRAM_LISTEN_DEFAULT_MODEL
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.utils import LlmProviders
_WEBSOCKET_SCHEMES: Final = MappingProxyType({"https": "wss", "http": "ws", "wss": "wss", "ws": "ws"})
DEEPGRAM_LISTEN_CALLBACK_PARAMS: Final = frozenset({"callback", "callback_method"})
DEEPGRAM_LISTEN_STREAMING_PRICING_PREFIX: Final = "streaming/"
DEEPGRAM_LISTEN_MULTILINGUAL_LANGUAGE: Final = "multi"
DEEPGRAM_LISTEN_MULTILINGUAL_PRICING_SUFFIX: Final = "-multilingual"
DEEPGRAM_LISTEN_ADDON_PRICING_PARAMS: Final = MappingProxyType(
{
"redact": "redact",
"keyterm": "keyterm",
"detect_entities": "detect_entities",
"diarize": "diarize",
"diarize_model": "diarize",
}
)
_DISABLED_PARAM_VALUES: Final = frozenset({"", "false"})
_SINGLE_VALUED_PARAMS: Final = frozenset({"model", "language"})
class DeepgramException(BaseLLMException):
pass
def deepgram_listen_requested_model(query_string: str) -> str:
return httpx.QueryParams(query_string).get("model") or DEEPGRAM_LISTEN_DEFAULT_MODEL
def _first_occurrences(query_string: str) -> httpx.QueryParams:
"""Authorization and pricing read the first ``model`` and ``language`` value; Deepgram must not see a second one."""
items: Final = httpx.QueryParams(query_string).multi_items()
return httpx.QueryParams(
tuple(
(key, value)
for index, (key, value) in enumerate(items)
if key not in _SINGLE_VALUED_PARAMS or all(earlier != key for earlier, _ in items[:index])
)
)
def deepgram_listen_websocket_target(api_base: str | None, query_string: str) -> str:
listen_url: Final = httpx.URL(f"{(api_base or DEEPGRAM_DEFAULT_API_BASE).rstrip('/')}/listen")
websocket_url: Final = listen_url.copy_with(scheme=_WEBSOCKET_SCHEMES.get(listen_url.scheme, listen_url.scheme))
params: Final = _first_occurrences(query_string)
query: Final = params if params.get("model") else params.remove("model").add("model", DEEPGRAM_LISTEN_DEFAULT_MODEL)
return f"{websocket_url}?{query}"
def deepgram_listen_callback_params(query_string: str) -> tuple[str, ...]:
return tuple(sorted(DEEPGRAM_LISTEN_CALLBACK_PARAMS.intersection(httpx.QueryParams(query_string).keys())))
def deepgram_listen_model(upstream_url: str) -> str:
models: Final = parse_qs(urlparse(upstream_url).query).get("model")
return models[0] if models else DEEPGRAM_LISTEN_DEFAULT_MODEL
def _param_enabled(values: Sequence[str]) -> bool:
return any(value.strip().lower() not in _DISABLED_PARAM_VALUES for value in values)
def deepgram_listen_pricing_model(upstream_url: str) -> str:
"""Registry key, without the provider prefix, for the per-second base rate Deepgram bills a streaming session at:
the multilingual streaming entry when ``language=multi``, otherwise the model's own streaming entry. Pre-recorded
entries are never a substitute: Deepgram prices the two products differently."""
streaming: Final = f"{DEEPGRAM_LISTEN_STREAMING_PRICING_PREFIX}{deepgram_listen_model(upstream_url)}"
language: Final = parse_qs(urlparse(upstream_url).query).get("language", ("",))[0]
if language.strip().lower() == DEEPGRAM_LISTEN_MULTILINGUAL_LANGUAGE:
return f"{streaming}{DEEPGRAM_LISTEN_MULTILINGUAL_PRICING_SUFFIX}"
return streaming
def deepgram_listen_registry_key(upstream_url: str) -> str:
return f"{LlmProviders.DEEPGRAM.value}/{deepgram_listen_pricing_model(upstream_url)}"
def deepgram_listen_is_priced(upstream_url: str) -> bool:
"""Only an exact registry hit counts: the cost calculator resolves a missing ``streaming/<model>`` row to the
pre-recorded ``<model>`` row, which is not the rate Deepgram bills a WebSocket session at."""
registry_key: Final = deepgram_listen_registry_key(upstream_url)
try:
model_info: Final = litellm.get_model_info(model=registry_key, custom_llm_provider=LlmProviders.DEEPGRAM.value)
except Exception:
return False
return model_info["key"] == registry_key
def deepgram_listen_addon_pricing_models(upstream_url: str) -> tuple[str, ...]:
params: Final = parse_qs(urlparse(upstream_url).query)
return tuple(
sorted(
frozenset(
f"{DEEPGRAM_LISTEN_STREAMING_PRICING_PREFIX}{addon}"
for param, addon in DEEPGRAM_LISTEN_ADDON_PRICING_PARAMS.items()
if _param_enabled(params.get(param, ()))
)
)
)
def _channel_count(value: object) -> int | None:
if isinstance(value, bool) or not isinstance(value, int):
return None
return value if value >= 1 else None
def _results_channel_count(frame: Mapping[str, object]) -> int | None:
channel_index: Final = frame.get("channel_index")
if not isinstance(channel_index, list) or len(channel_index) != 2:
return None
return _channel_count(channel_index[1])
def _declared_channel_count(upstream_url: str) -> int | None:
declared: Final = parse_qs(urlparse(upstream_url).query).get("channels")
if not declared or not declared[0].isdigit():
return None
return _channel_count(int(declared[0]))
def deepgram_listen_channel_count(websocket_messages: Sequence[Mapping[str, object]], upstream_url: str) -> int:
metadata_channels: Final = tuple(
channels
for frame in websocket_messages
if frame.get("type") == "Metadata"
if (channels := _channel_count(frame.get("channels"))) is not None
)
if metadata_channels:
return metadata_channels[-1]
results_channels: Final = tuple(
channels
for frame in websocket_messages
if frame.get("type") == "Results"
if (channels := _results_channel_count(frame)) is not None
)
if results_channels:
return max(results_channels)
return _declared_channel_count(upstream_url) or 1
def _seconds(value: object) -> float | None:
if isinstance(value, bool) or not isinstance(value, (int, float)):
return None
return float(value) if math.isfinite(value) and value >= 0 else None
def _results_frame_end(frame: Mapping[str, object]) -> float | None:
start: Final = _seconds(frame.get("start"))
duration: Final = _seconds(frame.get("duration"))
return None if start is None or duration is None else start + duration
def _final_transcript(frame: Mapping[str, object]) -> str | None:
if frame.get("is_final") is not True:
return None
channel: Final = frame.get("channel")
alternatives: Final = channel.get("alternatives") if isinstance(channel, Mapping) else None
first: Final = alternatives[0] if isinstance(alternatives, list) and alternatives else None
transcript: Final = first.get("transcript") if isinstance(first, Mapping) else None
return transcript if isinstance(transcript, str) and transcript else None
def deepgram_listen_audio_seconds(websocket_messages: Sequence[Mapping[str, object]]) -> float:
metadata_durations: Final = tuple(
duration
for frame in websocket_messages
if frame.get("type") == "Metadata"
if (duration := _seconds(frame.get("duration"))) is not None and duration > 0
)
if metadata_durations:
return metadata_durations[-1]
return max(
(
end
for frame in websocket_messages
if frame.get("type") == "Results"
if (end := _results_frame_end(frame)) is not None
),
default=0.0,
)
def deepgram_listen_transcript(websocket_messages: Sequence[Mapping[str, object]]) -> str:
return " ".join(
transcript
for frame in websocket_messages
if frame.get("type") == "Results"
if (transcript := _final_transcript(frame)) is not None
)

View file

@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, cast, get_type_hints, ove
import httpx
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.prompt_templates.common_utils import (
handle_messages_with_content_list_to_str_conversion,
@ -20,16 +21,37 @@ from litellm.llms.openai.chat.gpt_transformation import (
OpenAIChatCompletionStreamingHandler,
OpenAIGPTConfig,
)
from litellm.router_utils.reasoning_effort_capability import (
declared_reasoning_efforts_for_model,
nearest_declared_reasoning_effort,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.mistral import MistralThinkingBlock, MistralToolCallMessage
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse, ModelResponseStream
from litellm.utils import convert_to_model_response_object
from litellm.utils import convert_to_model_response_object, supports_reasoning
if TYPE_CHECKING:
import tiktoken
def _accepted_reasoning_effort(model: str, requested: str, custom_llm_provider: str) -> str:
declared: Final = declared_reasoning_efforts_for_model(model, custom_llm_provider)
if declared is None:
return requested
accepted: Final = nearest_declared_reasoning_effort(requested, declared)
if accepted != requested:
verbose_logger.debug(
"%s: %s takes reasoning_effort %s, sending %s in place of %s",
custom_llm_provider,
model,
declared,
accepted,
requested,
)
return accepted
class MistralConfig(OpenAIGPTConfig):
"""
Reference: https://docs.mistral.ai/api/
@ -86,8 +108,16 @@ class MistralConfig(OpenAIGPTConfig):
def get_config(cls):
return super().get_config()
@property
def custom_llm_provider(self) -> str:
return "mistral"
def get_supported_openai_params(self, model: str) -> list[str]:
supported_params: Final = [
is_magistral: Final = "magistral" in model.lower()
accepts_reasoning_effort: Final = is_magistral or supports_reasoning(
model=model, custom_llm_provider=self.custom_llm_provider
)
return [
"stream",
"temperature",
"top_p",
@ -99,14 +129,10 @@ class MistralConfig(OpenAIGPTConfig):
"stop",
"response_format",
"parallel_tool_calls",
*(("thinking",) if is_magistral else ()),
*(("reasoning_effort",) if accepts_reasoning_effort else ()),
]
# Add reasoning support for magistral models
if "magistral" in model.lower():
supported_params.extend(["thinking", "reasoning_effort"])
return supported_params
def _map_tool_choice(self, tool_choice: str) -> str:
if tool_choice == "auto" or tool_choice == "none":
return tool_choice
@ -171,10 +197,9 @@ class MistralConfig(OpenAIGPTConfig):
optional_params["extra_body"] = {"random_seed": value}
if param == "response_format":
optional_params["response_format"] = value
if param == "reasoning_effort" and "magistral" in model.lower():
# Flag that we need to add reasoning system prompt
optional_params["_add_reasoning_prompt"] = True
if param == "thinking" and "magistral" in model.lower():
if param == "reasoning_effort" and "magistral" not in model.lower():
optional_params["reasoning_effort"] = _accepted_reasoning_effort(model, value, self.custom_llm_provider)
if param in ("reasoning_effort", "thinking") and "magistral" in model.lower():
# Flag that we need to add reasoning system prompt
optional_params["_add_reasoning_prompt"] = True
if param == "parallel_tool_calls":
@ -534,11 +559,13 @@ class MistralConfig(OpenAIGPTConfig):
if "magistral" in model.lower() and optional_params.get("_add_reasoning_prompt", False):
messages = self._add_reasoning_system_prompt_if_needed(messages, optional_params)
upstream_params: Final = {key: value for key, value in optional_params.items() if key != "client_metadata"}
# Call parent transform_request which handles _transform_messages
return super().transform_request(
model=model,
messages=messages,
optional_params=optional_params,
optional_params=upstream_params,
litellm_params=litellm_params,
headers=headers,
)

View file

@ -0,0 +1,7 @@
from litellm.llms.mistral.chat.transformation import MistralConfig
class VertexAIMistralConfig(MistralConfig):
@property
def custom_llm_provider(self) -> str:
return "vertex_ai"

View file

@ -20853,6 +20853,96 @@
"/v1/audio/transcriptions"
]
},
"deepgram/streaming/nova-3": {
"input_cost_per_second": 8e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0048/60 seconds = $0.00008000 per second",
"note": "Nova-3 monolingual streaming, pay as you go",
"original_pricing_per_minute": 0.0048
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/nova-3-multilingual": {
"input_cost_per_second": 9.667e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0058/60 seconds = $0.00009667 per second",
"note": "Nova-3 multilingual (language=multi) streaming, pay as you go",
"original_pricing_per_minute": 0.0058
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/redact": {
"input_cost_per_second": 3.333e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0020/60 seconds = $0.00003333 per second",
"note": "Redaction add-on (redact query param), streaming, pay as you go",
"original_pricing_per_minute": 0.002
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/keyterm": {
"input_cost_per_second": 2.167e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0013/60 seconds = $0.00002167 per second",
"note": "Keyterm Prompting add-on (keyterm query param), streaming, pay as you go",
"original_pricing_per_minute": 0.0013
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/detect_entities": {
"input_cost_per_second": 2.833e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0017/60 seconds = $0.00002833 per second",
"note": "Entity Detection add-on (detect_entities query param), streaming, pay as you go",
"original_pricing_per_minute": 0.0017
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/diarize": {
"input_cost_per_second": 3.333e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0020/60 seconds = $0.00003333 per second",
"note": "Speaker Diarization add-on (diarize / diarize_model query params), streaming, pay as you go",
"original_pricing_per_minute": 0.002
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/whisper": {
"input_cost_per_second": 0.0001,
"litellm_provider": "deepgram",
@ -37107,6 +37197,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37188,6 +37282,15 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"none",
"minimal",
"low",
"medium",
"high",
"xhigh",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-2",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37205,6 +37308,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-3",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37222,6 +37330,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-3",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37239,6 +37352,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-3",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37256,6 +37374,15 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"none",
"minimal",
"low",
"medium",
"high",
"xhigh",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-2",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37551,6 +37678,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37611,6 +37742,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37628,6 +37763,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37661,6 +37800,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37692,6 +37835,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6e-07,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -60037,6 +60184,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6e-07,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -63362,6 +63513,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -63379,6 +63534,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -63396,6 +63555,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -63413,6 +63576,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6e-07,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03",
"supports_assistant_prefill": true,
"supports_function_calling": true,

View file

@ -201,6 +201,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
"/cohere/",
"/comprehendmedical",
"/cursor/",
"/deepgram/",
"/eu.assemblyai/",
"/gemini/",
"/gigachat/",

View file

@ -71,6 +71,7 @@ from litellm.types.utils import (
StandardLoggingVectorStoreRequest,
StandardPassThroughResponseObject,
TextCompletionResponse,
TranscriptionResponse,
)
from litellm.types.videos.main import VideoObject
@ -525,6 +526,7 @@ class LiteLLMRoutes(enum.Enum):
"/gigachat",
"/watsonx",
"/nvidia_nim",
"/deepgram",
]
#########################################################
@ -4806,6 +4808,7 @@ PassThroughEndpointLoggingResultValues = (
| VideoObject
| StandardPassThroughResponseObject
| ResponsesAPIResponse
| TranscriptionResponse
)

View file

@ -664,9 +664,11 @@ def get_websocket_api_key(websocket: WebSocket) -> str | None:
)
async def user_api_key_auth_websocket(websocket: WebSocket):
# Accept the WebSocket connection
async def user_api_key_auth_websocket(websocket: WebSocket) -> UserAPIKeyAuth:
return await user_api_key_auth_websocket_for_model(websocket, model=websocket.query_params.get("model"))
async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str | None) -> UserAPIKeyAuth:
ws_scope: Final = websocket.scope or {}
scope_headers: Final = list(ws_scope.get("headers") or [])
# ``get_request_route`` falls back to ``request.url.path`` when
@ -701,16 +703,12 @@ async def user_api_key_auth_websocket(websocket: WebSocket):
from litellm.proxy.realtime_endpoints.call_sessions import decode_call
call_token: Final = websocket.path_params.get("call_id") or websocket.query_params.get("call_id")
model: Final = (
decode_call(call_token, f"Bearer {api_key}").alias
if call_token is not None
else websocket.query_params.get("model")
)
resolved_model: Final = decode_call(call_token, f"Bearer {api_key}").alias if call_token is not None else model
if call_token is not None:
request.scope["litellm_pinned_realtime_model"] = model
request.scope["litellm_pinned_realtime_model"] = resolved_model
async def return_body():
return _realtime_request_body(model)
return _realtime_request_body(resolved_model)
request.body = return_body
return await user_api_key_auth(request=request, api_key=f"Bearer {api_key}")

View file

@ -1934,6 +1934,7 @@ class ProxyBaseLLMRequestProcessing:
user_api_base: str | None = None,
model: str | None = None,
llm_router: Router | None = None,
rate_limited_model: str | None = None,
*,
internal_realtime_observer: bool = False,
) -> tuple[dict, LiteLLMLoggingObj]:
@ -2099,8 +2100,15 @@ class ProxyBaseLLMRequestProcessing:
# model_info when allow_client_pricing_override is set, so a caller
# could otherwise spoof an unguarded model_info.id while requesting
# a guarded alias and bypass guardrails (veria-ai HIGH on #29654).
merged_for_requested: Final = (
self.data
if rate_limited_model is None
else _check_and_merge_model_level_guardrails(
data=self.data, llm_router=llm_router, trust_client_model_info=False, model_alias=rate_limited_model
)
)
self.data = _check_and_merge_model_level_guardrails(
data=self.data,
data=merged_for_requested,
llm_router=llm_router,
trust_client_model_info=False,
)
@ -2170,7 +2178,7 @@ class ProxyBaseLLMRequestProcessing:
configured_fallbacks: Final = (
self._configured_fallbacks(llm_router=llm_router, user_api_key_dict=user_api_key_dict)
if llm_router is not None and not self.data.get("disable_fallbacks")
if llm_router is not None
else None
)
pristine: Final = independent_snapshot(self.data) if configured_fallbacks else None
@ -2215,7 +2223,6 @@ class ProxyBaseLLMRequestProcessing:
original_model,
fallback_models,
)
try:
for fallback_model in fallback_models:
if fallback_model == original_model:
@ -2238,6 +2245,7 @@ class ProxyBaseLLMRequestProcessing:
model=fallback_model,
route_type=route_type,
llm_router=llm_router,
rate_limited_model=original_model,
)
except ProxyRateLimitError:
continue

View file

@ -11,6 +11,7 @@ from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT
from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.spend_tracking.daily_global_spend_rollup import GLOBAL_SPEND_TABLE_NAME, reconciled_through
from litellm.proxy.spend_tracking.key_metadata_recovery import (
attach_user_emails,
recover_double_hashed_key_metadata,
@ -750,6 +751,66 @@ def _rollup_metric_select(table_name: str) -> str:
_MODEL_GROUP_EXPR: Final = "COALESCE(NULLIF(model_group, ''), model)"
_KEY_FREE_SOURCE_COLUMNS: Final = (
"date",
"model",
"model_group",
"custom_llm_provider",
"mcp_namespaced_tool_name",
"endpoint",
"spend",
"prompt_tokens",
"completion_tokens",
"cache_read_input_tokens",
"cache_creation_input_tokens",
"compression_saved_tokens",
"compression_savings_spend",
"prompt_caching_savings_spend",
"gateway_injected_caching_savings_spend",
"autorouter_savings_spend",
"api_requests",
"successful_requests",
"failed_requests",
"total_response_time_ms",
"timed_requests",
)
async def global_rollup_reconciled_through(prisma_client: PrismaClient, query: _AggregatedQueryKwargs) -> str | None:
"""The last day ``LiteLLM_DailyGlobalSpend`` can answer the key-free arm for, or None to
read it all from the per-key table.
Only an unfiltered read of the user table sums to the same rows as the global table. The
marker read is served from the config cache, so this is not a database round trip per request.
"""
if query["table_name"] != "litellm_dailyuserspend":
return None
if query["entity_id"] is not None or query["api_key"] is not None or query["exclude_entity_ids"]:
return None
try:
return await reconciled_through(prisma_client)
except Exception as exc: # noqa: BLE001 # the per-key table is always a correct answer, so never fail the read
verbose_proxy_logger.warning("Could not read the daily global spend marker, using the per-key table: %s", exc)
return None
def _key_free_source(pg_table: str, where_clause: str, marker_param: str | None) -> str:
"""The relation the key-free arm aggregates: the per-key table alone, or the global rollup
for days through the marker plus the per-key table for the days still open after it."""
if marker_param is None:
return f'"{pg_table}"\n WHERE {where_clause}'
columns: Final = ", ".join(_KEY_FREE_SOURCE_COLUMNS)
return f"""(
SELECT {columns}
FROM "{GLOBAL_SPEND_TABLE_NAME}"
WHERE {where_clause} AND date <= {marker_param}
UNION ALL
SELECT {columns}
FROM "{pg_table}"
WHERE {where_clause} AND date > {marker_param}
) AS key_free_source"""
def _build_aggregated_sql_query(
*,
table_name: str,
@ -762,6 +823,7 @@ def _build_aggregated_sql_query(
exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path
timezone_offset_minutes: int | None = None,
include_current_utc_day: bool = False,
global_rollup_through: str | None = None,
) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params
"""Build the GROUPING SETS query for aggregated daily activity.
@ -786,6 +848,7 @@ def _build_aggregated_sql_query(
exclude_entity_ids=exclude_entity_ids,
)
sentinel_param: Final = f"${len(where_params) + 1}"
marker_param: Final = None if global_rollup_through is None else f"${len(where_params) + 2}"
metric_select: Final = _rollup_metric_select(table_name)
# TODO: drop the successful_requests/failed_requests aggregates (and the
@ -806,8 +869,7 @@ def _build_aggregated_sql_query(
custom_llm_provider, mcp_namespaced_tool_name,
endpoint) AS group_level,
NULL::bigint AS distinct_api_keys,{metric_select}
FROM "{pg_table}"
WHERE {where_clause}
FROM {_key_free_source(pg_table, where_clause, marker_param)}
GROUP BY GROUPING SETS (
(date),
(date, model),
@ -850,7 +912,8 @@ def _build_aggregated_sql_query(
))
"""
return sql_query, [*where_params, PTU_SENTINEL_API_KEY]
marker_params: Final = () if global_rollup_through is None else (global_rollup_through,)
return sql_query, [*where_params, PTU_SENTINEL_API_KEY, *marker_params]
def _build_entity_rollup_sql_query(
@ -1395,7 +1458,10 @@ async def get_daily_activity_aggregated(
timezone_offset_minutes=timezone_offset_minutes,
include_current_utc_day=include_current_utc_day,
)
sql_query, sql_params = _build_aggregated_sql_query(**query_kwargs)
sql_query, sql_params = _build_aggregated_sql_query(
**query_kwargs,
global_rollup_through=await global_rollup_reconciled_through(prisma_client, query_kwargs),
)
entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None
raw_rows, raw_entity_rows = await asyncio.gather(

View file

@ -45,6 +45,13 @@ from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.azure.passthrough.transformation import foreign_azure_deployment
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.deepgram.common_utils import (
deepgram_listen_callback_params,
deepgram_listen_is_priced,
deepgram_listen_registry_key,
deepgram_listen_requested_model,
deepgram_listen_websocket_target,
)
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
@ -57,6 +64,7 @@ from litellm.proxy.auth.user_api_key_auth import (
is_no_auth_dev_mode,
user_api_key_auth,
user_api_key_auth_websocket,
user_api_key_auth_websocket_for_model,
)
from litellm.proxy.common_request_processing import open_sse_before_first_byte
from litellm.proxy.common_utils.http_parsing_utils import (
@ -2925,7 +2933,7 @@ async def _openai_websocket_refusal(
return None
class _OpenAIWebsocketRelay(Protocol):
class _WebsocketRelay(Protocol):
async def __call__(
self,
*,
@ -2939,7 +2947,7 @@ class _OpenAIWebsocketRelay(Protocol):
) -> None: ...
def _openai_websocket_relay() -> _OpenAIWebsocketRelay:
def _websocket_relay() -> _WebsocketRelay:
return websocket_passthrough_request
@ -2957,6 +2965,15 @@ def _proxy_model_allowlists() -> _OpenAIWebsocketModelAllowlists:
return resolve
def _negotiated_websocket_subprotocol(websocket: WebSocket) -> str | None:
requested_subprotocols: Final = tuple(
protocol.strip()
for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",")
if protocol.strip()
)
return requested_subprotocols[0] if requested_subprotocols else None
@router.websocket("/openai_passthrough/{endpoint:path}")
@router.websocket("/openai/{endpoint:path}")
async def openai_websocket_proxy_route(
@ -2964,16 +2981,11 @@ async def openai_websocket_proxy_route(
endpoint: str,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)],
general_settings: Annotated[Mapping[str, object], Depends(_proxy_general_settings)],
relay: Annotated[_OpenAIWebsocketRelay, Depends(_openai_websocket_relay)],
relay: Annotated[_WebsocketRelay, Depends(_websocket_relay)],
model_allowlists: Annotated[_OpenAIWebsocketModelAllowlists, Depends(_proxy_model_allowlists)],
) -> None:
"""WebSocket passthrough for OpenAI prefixes (realtime / responses.connect)."""
requested_subprotocols: Final = tuple(
protocol.strip()
for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",")
if protocol.strip()
)
negotiated_subprotocol: Final = requested_subprotocols[0] if requested_subprotocols else None
negotiated_subprotocol: Final = _negotiated_websocket_subprotocol(websocket)
refusal: Final = await _openai_websocket_refusal(user_api_key_dict, general_settings, model_allowlists)
if refusal is not None:
@ -3032,6 +3044,69 @@ async def openai_websocket_proxy_route(
)
_DEEPGRAM_WS_MISSING_KEY_REASON: Final = (
"Required 'DEEPGRAM_API_KEY' in environment to make pass-through calls to Deepgram."
)
_DEEPGRAM_WS_CALLBACK_REASON: Final = "Deepgram callback delivery is not supported through the proxy: remove {params}"
_DEEPGRAM_WS_UNPRICED_REASON: Final = (
"No streaming price for '{registry_key}': add it to the model cost map to enable it"
)
async def deepgram_listen_user_api_key_auth(websocket: WebSocket) -> UserAPIKeyAuth:
return await user_api_key_auth_websocket_for_model(
websocket, model=deepgram_listen_requested_model(websocket.url.query)
)
@router.websocket("/deepgram/v1/listen")
@router.websocket("/deepgram/listen")
async def deepgram_listen_websocket_route(
websocket: WebSocket,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(deepgram_listen_user_api_key_auth)],
relay: Annotated[_WebsocketRelay, Depends(_websocket_relay)],
) -> None:
deepgram_api_key: Final = passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.DEEPGRAM.value,
region_name=None,
)
if deepgram_api_key is None:
await websocket.close(code=1011, reason=_DEEPGRAM_WS_MISSING_KEY_REASON)
return
await websocket.accept(subprotocol=_negotiated_websocket_subprotocol(websocket))
callback_params: Final = deepgram_listen_callback_params(websocket.url.query)
if callback_params:
await websocket.close(
code=1008,
reason=_DEEPGRAM_WS_CALLBACK_REASON.format(params=", ".join(callback_params)),
)
return
target: Final = deepgram_listen_websocket_target(
api_base=get_secret_str("DEEPGRAM_API_BASE"),
query_string=websocket.url.query,
)
if not deepgram_listen_is_priced(target):
await websocket.close(
code=1008,
reason=_DEEPGRAM_WS_UNPRICED_REASON.format(registry_key=deepgram_listen_registry_key(target)),
)
return
await relay(
websocket=websocket,
target=target,
custom_headers={ # mutable-ok: websocket_passthrough_request requires a plain dict of upstream headers
"Authorization": f"Token {deepgram_api_key}"
},
user_api_key_dict=user_api_key_dict,
forward_headers=False,
endpoint=websocket.url.path,
accept_websocket=False,
)
class BaseOpenAIPassThroughHandler:
@staticmethod
async def _base_openai_pass_through_handler(

View file

@ -8,6 +8,7 @@ import httpx
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
AZURE_SPEECH_BATCH_MODEL,
AZURE_SPEECH_BATCH_PATH_PREFIX,
AZURE_SPEECH_CUSTOM_LLM_PROVIDER,
AZURE_SPEECH_FAST_TRANSCRIPTION_MODEL,
AZURE_SPEECH_FAST_TRANSCRIPTION_PATH,
@ -30,7 +31,8 @@ from litellm.types.utils import StandardPassThroughResponseObject
class AzureSpeechPassthroughLoggingHandler:
@staticmethod
def _is_short_audio_route(url_route: str) -> bool:
return urlparse(url_route).path.startswith(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX)
path: Final = urlparse(url_route).path
return path.rfind(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX) > path.rfind(AZURE_SPEECH_BATCH_PATH_PREFIX)
@staticmethod
def _is_fast_transcription_route(url_route: str) -> bool:

View file

@ -0,0 +1,96 @@
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final
from urllib.parse import urlparse
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.deepgram.common_utils import (
deepgram_listen_addon_pricing_models,
deepgram_listen_audio_seconds,
deepgram_listen_channel_count,
deepgram_listen_is_priced,
deepgram_listen_model,
deepgram_listen_pricing_model,
deepgram_listen_registry_key,
deepgram_listen_transcript,
)
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
from litellm.types.utils import TranscriptionResponse
DEEPGRAM_LISTEN_ROUTE_SUFFIX: Final = "/listen"
def _registry_cost(response: TranscriptionResponse, pricing_model: str) -> float | None:
try:
return litellm.completion_cost(
completion_response=response,
model=pricing_model,
custom_llm_provider=litellm.LlmProviders.DEEPGRAM.value,
call_type="transcription",
)
except Exception as e: # noqa: BLE001 # an unpriced entry must not lose the spend row, only its cost
verbose_proxy_logger.debug("Deepgram listen passthrough: no registry price for '%s': %s", pricing_model, e)
return None
def _audio_cost(response: TranscriptionResponse, upstream_url: str) -> float | None:
if not deepgram_listen_is_priced(upstream_url):
verbose_proxy_logger.warning(
"Deepgram listen passthrough: no registry entry '%s'", deepgram_listen_registry_key(upstream_url)
)
return None
base_cost: Final = _registry_cost(response, deepgram_listen_pricing_model(upstream_url))
if base_cost is None:
return None
addon_costs: Final = tuple(
_registry_cost(response, pricing_model) for pricing_model in deepgram_listen_addon_pricing_models(upstream_url)
)
return base_cost + sum(cost for cost in addon_costs if cost is not None)
class DeepgramListenPassthroughLoggingHandler:
@staticmethod
def is_deepgram_listen_route(url_route: str) -> bool:
path: Final = urlparse(url_route).path
return "/deepgram/" in path and path.endswith(DEEPGRAM_LISTEN_ROUTE_SUFFIX)
def deepgram_listen_passthrough_handler(
self,
websocket_messages: Sequence[Mapping[str, object]],
logging_obj: LiteLLMLoggingObj,
upstream_url: str,
kwargs: Mapping[str, object] = MappingProxyType({}),
) -> PassThroughEndpointLoggingTypedDict:
model: Final = deepgram_listen_model(upstream_url)
audio_seconds: Final = deepgram_listen_audio_seconds(websocket_messages)
channels: Final = deepgram_listen_channel_count(websocket_messages, upstream_url)
billed_seconds: Final = audio_seconds * channels
response: Final = TranscriptionResponse(text=deepgram_listen_transcript(websocket_messages))
response._hidden_params["audio_transcription_duration"] = billed_seconds # pyright: ignore[reportPrivateUsage] # the cost calculator reads the billed duration off the response's hidden params
response_cost: Final = _audio_cost(response, upstream_url)
response._hidden_params["response_cost"] = response_cost # pyright: ignore[reportPrivateUsage] # the logger reads a precomputed cost off the response's hidden params
provider: Final = litellm.LlmProviders.DEEPGRAM.value
logging_obj.model = model # rebind-ok: the spend logger reads model and cost off the shared logging object
logging_obj.model_call_details["model"] = model # rebind-ok: same shared logging object
logging_obj.model_call_details["custom_llm_provider"] = provider # rebind-ok: same shared logging object
logging_obj.model_call_details["response_cost"] = response_cost # rebind-ok: same shared logging object
verbose_proxy_logger.debug(
"Deepgram listen passthrough cost tracking: model %s, audio seconds %s, channels %s, cost %s",
model,
audio_seconds,
channels,
response_cost,
)
logging_result: Final[PassThroughEndpointLoggingTypedDict] = {
"result": response,
"kwargs": {
**kwargs,
"model": model,
"custom_llm_provider": provider,
"response_cost": response_cost,
},
}
return logging_result

View file

@ -8,7 +8,7 @@ from base64 import b64encode
from collections.abc import AsyncGenerator, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from itertools import groupby
from itertools import count, groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
from urllib.parse import urlencode, urlparse
@ -2122,6 +2122,14 @@ def _resolved_vertex_live_setup(
return {**setup_data, "model": setup_model_rewriter(setup_model)}
def _json_object_frame(frame: str | bytes) -> dict[str, object] | None:
try:
decoded: Final = json.loads(frame if isinstance(frame, str) else frame.decode("utf-8"))
except (json.JSONDecodeError, UnicodeDecodeError):
return None
return decoded if isinstance(decoded, dict) else None
def _truncated_close_reason(reason: str) -> str:
"""
Fit a close reason inside the byte budget a WebSocket close frame allows, without splitting a character
@ -2403,70 +2411,41 @@ async def websocket_passthrough_request(
)
await upstream_ws.close()
def _extract_vertex_live_model_from_setup_response(setup_response: Mapping[str, object]) -> None:
extracted_model: Final = _extract_model_from_vertex_ai_setup(setup_response)
if not extracted_model:
verbose_proxy_logger.warning(
"WebSocket passthrough (%s): Failed to extract model from server setup response: %s",
endpoint,
setup_response,
)
return
kwargs["model"] = extracted_model
kwargs["custom_llm_provider"] = "vertex_ai_language_models"
logging_obj.model = extracted_model
logging_obj.model_call_details["model"] = extracted_model
logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai_language_models"
is_vertex_live: Final = bool(endpoint and "/vertex_ai/live" in endpoint)
json_frame_ordinal: Final = count()
async def relay_upstream_frame(upstream_message: str | bytes) -> None:
if isinstance(upstream_message, bytes):
await websocket.send_bytes(upstream_message)
else:
await websocket.send_text(upstream_message)
message_data: Final = _json_object_frame(upstream_message)
if message_data is None:
return
if is_vertex_live and next(json_frame_ordinal) == 0:
_extract_vertex_live_model_from_setup_response(message_data)
return
websocket_messages.append(message_data)
async def forward_upstream_to_client() -> Close | None:
"""Forward messages from upstream to client WebSocket, returning the upstream's close frame"""
try:
# Wait for the first response from upstream
raw_response = await upstream_ws.recv(decode=False)
# Ensure raw_response is bytes before decoding
if isinstance(raw_response, str):
raw_response = raw_response.encode("utf-8")
setup_response: Final[Mapping[str, object]] = json.loads(raw_response.decode("utf-8"))
verbose_proxy_logger.debug("Setup response: %s", setup_response)
# Extract model and provider from setup response for Vertex AI Live
if endpoint and "/vertex_ai/live" in endpoint:
verbose_proxy_logger.debug(
"WebSocket passthrough (%s): Processing server setup response for model extraction",
endpoint,
)
extracted_model: Final = _extract_model_from_vertex_ai_setup(setup_response)
if extracted_model:
kwargs["model"] = extracted_model
kwargs["custom_llm_provider"] = "vertex_ai_language_models"
# Update logging object with correct model
logging_obj.model = extracted_model
logging_obj.model_call_details["model"] = extracted_model
logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai_language_models"
verbose_proxy_logger.debug(
"WebSocket passthrough (%s): Successfully extracted model '%s' and set provider to 'vertex_ai' from server setup response",
endpoint,
extracted_model,
)
else:
verbose_proxy_logger.warning(
"WebSocket passthrough (%s): Failed to extract model from server setup response: %s",
endpoint,
setup_response,
)
else:
verbose_proxy_logger.debug(
"WebSocket passthrough (%s): Not a Vertex AI Live endpoint, skipping model extraction",
endpoint,
)
# Send the setup response to the client
await websocket.send_text(json.dumps(setup_response))
# Now continuously forward messages from upstream to client
async for upstream_message in upstream_ws:
if isinstance(upstream_message, bytes):
await websocket.send_bytes(upstream_message)
# Parse and collect for cost tracking
try:
message_data: dict[str, object] = json.loads(upstream_message.decode())
websocket_messages.append(message_data)
except (json.JSONDecodeError, UnicodeDecodeError):
pass
else:
await websocket.send_text(upstream_message)
# Parse and collect for cost tracking
try:
message_data = json.loads(upstream_message)
websocket_messages.append(message_data)
except json.JSONDecodeError:
pass
while True:
await relay_upstream_frame(await upstream_ws.recv())
except (ConnectionClosedOK, ConnectionClosedError) as e:
verbose_proxy_logger.debug("Upstream WebSocket connection closed: %s", e)
return e.rcvd

View file

@ -26,6 +26,9 @@ from .llm_provider_handlers.cohere_passthrough_logging_handler import (
from .llm_provider_handlers.cursor_passthrough_logging_handler import (
CursorPassthroughLoggingHandler,
)
from .llm_provider_handlers.deepgram_listen_passthrough_logging_handler import (
DeepgramListenPassthroughLoggingHandler,
)
from .llm_provider_handlers.gemini_passthrough_logging_handler import (
GeminiPassthroughLoggingHandler,
)
@ -349,6 +352,21 @@ class PassThroughEndpointLogging:
standard_logging_response_object = vertex_ai_live_handler_result["result"]
kwargs = vertex_ai_live_handler_result["kwargs"]
elif DeepgramListenPassthroughLoggingHandler.is_deepgram_listen_route(url_route):
deepgram_handler_result: Final = (
DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=tuple(
message
for message in (response_body if isinstance(response_body, list) else ())
if isinstance(message, dict)
),
logging_obj=logging_obj,
upstream_url=str(httpx_response.request.url),
kwargs=kwargs,
)
)
standard_logging_response_object = deepgram_handler_result["result"] # rebind-ok: elif-chain
kwargs = deepgram_handler_result["kwargs"] # rebind-ok: elif-chain contract
return_dict["standard_logging_response_object"] = standard_logging_response_object
return_dict["kwargs"] = kwargs

View file

@ -262,6 +262,7 @@ from litellm.constants import (
APSCHEDULER_MISFIRE_GRACE_TIME,
APSCHEDULER_REPLACE_EXISTING,
CLI_SSO_SESSION_TTL_SECONDS,
DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID,
DAYS_IN_A_MONTH,
DEFAULT_HEALTH_CHECK_INTERVAL,
DEFAULT_MODEL_CREATED_AT_TIME,
@ -688,6 +689,9 @@ from litellm.proxy.route_priority import hot_routes_first
from litellm.proxy.search_endpoints.endpoints import router as search_router
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
from litellm.proxy.spend_tracking.daily_global_spend_rollup import (
run_scheduled_daily_global_spend_reconcile,
)
from litellm.proxy.spend_tracking.spend_counter_batch import (
PendingSpendIncrement,
active_spend_counter_batch,
@ -10021,6 +10025,12 @@ class ProxyStartupEvent:
await cls._initialize_spend_tracking_background_jobs(scheduler=scheduler)
cls._initialize_daily_global_spend_reconcile_job(
scheduler=scheduler,
proxy_logging_obj=proxy_logging_obj,
prisma_client=prisma_client,
)
### PTU DAILY ROLLUP ###
from litellm.proxy.spend_tracking.ptu_feature_flag import (
is_ptu_cost_attribution_enabled,
@ -10362,6 +10372,39 @@ class ProxyStartupEvent:
"LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_ENABLED=true to enable)"
)
@classmethod
def _initialize_daily_global_spend_reconcile_job(
cls,
scheduler: AsyncIOScheduler,
proxy_logging_obj: ProxyLogging,
prisma_client: PrismaClient,
) -> None:
async def alert(message: str) -> None:
await proxy_logging_obj.alerting_handler(
message=message,
level="High",
alert_type=AlertType.failed_tracking_spend,
)
async def reconcile() -> None:
await run_scheduled_daily_global_spend_reconcile(
prisma_client,
pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager,
alert=alert,
)
scheduler.add_job(
reconcile,
"cron",
hour=0,
minute=30,
timezone="UTC",
id=DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID,
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
next_run_time=datetime.now(timezone.utc) + timedelta(minutes=2),
)
@classmethod
async def _initialize_slack_alerting_jobs(
cls,

View file

@ -820,6 +820,37 @@ model LiteLLM_DailyUserSpend {
@@index([endpoint])
}
// Key-free daily rollup of LiteLLM_DailyUserSpend, read by the global usage view
model LiteLLM_DailyGlobalSpend {
id String @id @default(uuid())
date String
model String?
model_group String?
custom_llm_provider String?
mcp_namespaced_tool_name String?
endpoint String?
prompt_tokens BigInt @default(0)
completion_tokens BigInt @default(0)
cache_read_input_tokens BigInt @default(0)
cache_creation_input_tokens BigInt @default(0)
compression_saved_tokens BigInt @default(0)
compression_savings_spend Float @default(0.0)
prompt_caching_savings_spend Float @default(0.0)
gateway_injected_caching_savings_spend Float @default(0.0)
autorouter_savings_spend Float @default(0.0)
spend Float @default(0.0)
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@unique([date, model, model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
}
// Track daily organization spend metrics per model and key
model LiteLLM_DailyOrganizationSpend {
id String @id @default(uuid())

View file

@ -0,0 +1,293 @@
"""Roll closed UTC days of ``LiteLLM_DailyUserSpend`` up into ``LiteLLM_DailyGlobalSpend``.
Only days that are over get rolled up, so a pod still flushing per-key spend for the current
day can never leave the global table short; usage reads serve days through the recorded
marker from the global table and later days live from the per-key table. Per-key rows are
dated by request start, so spend can land on a day that was already rolled up (a flush
straddling midnight, a retry after an outage). Each run therefore also rewrites every closed
day that has rows touched since the previous run's scan, whatever the date. The marker lives
in ``LiteLLM_Config``. This runs as a background cron, never in a Prisma migration, since on
a large deployment the first backfill is minutes of work.
"""
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import date, timedelta
from typing import TYPE_CHECKING, Final
from pydantic import BaseModel, ConfigDict, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID,
DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS,
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM,
)
if TYPE_CHECKING:
from litellm.caching.redis_cache import RedisCache
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.utils import PrismaClient
GLOBAL_SPEND_TABLE_NAME: Final = "LiteLLM_DailyGlobalSpend"
# The unique constraint, in constraint order. NULL never matches itself in a unique index, so
# every column is normalized to '' or the same group would be inserted again on every run.
_KEY_COLUMNS: Final = ("date", "model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint")
_METRIC_COLUMNS: Final = (
"prompt_tokens",
"completion_tokens",
"cache_read_input_tokens",
"cache_creation_input_tokens",
"compression_saved_tokens",
"api_requests",
"successful_requests",
"failed_requests",
"total_response_time_ms",
"timed_requests",
"compression_savings_spend",
"prompt_caching_savings_spend",
"gateway_injected_caching_savings_spend",
"autorouter_savings_spend",
"spend",
)
def _quoted(columns: tuple[str, ...]) -> str:
return ", ".join(f'"{column}"' for column in columns)
def _reconcile_day_sql() -> str:
normalized_keys: Final = ", ".join(f"COALESCE(\"{column}\", '')" for column in _KEY_COLUMNS)
sums: Final = ", ".join(f'SUM("{column}")' for column in _METRIC_COLUMNS)
overwrite: Final = ", ".join(f'"{column}" = EXCLUDED."{column}"' for column in _METRIC_COLUMNS)
return (
f'INSERT INTO "{GLOBAL_SPEND_TABLE_NAME}" ("id", {_quoted(_KEY_COLUMNS)}, {_quoted(_METRIC_COLUMNS)}, '
'"updated_at")\n'
f"SELECT gen_random_uuid()::text, {normalized_keys}, {sums}, (NOW() AT TIME ZONE 'UTC')\n"
'FROM "LiteLLM_DailyUserSpend" WHERE "date" = $1\n'
f"GROUP BY {normalized_keys}\n"
f"ON CONFLICT ({_quoted(_KEY_COLUMNS)}) DO UPDATE SET {overwrite}, "
"\"updated_at\" = (NOW() AT TIME ZONE 'UTC')"
)
RECONCILE_DAY_SQL: Final = _reconcile_day_sql()
_DB_NOW_SQL: Final = "SELECT (NOW() AT TIME ZONE 'UTC')::text AS now, (NOW() AT TIME ZONE 'UTC')::date::text AS today"
_ALL_CLOSED_DAYS_SQL: Final = 'SELECT DISTINCT "date" FROM "LiteLLM_DailyUserSpend" WHERE "date" <= $1 ORDER BY "date"'
# Pod clocks drift from the database clock and from each other, so rows are picked up from a
# little before the previous scan; rewriting a day twice is idempotent.
_PENDING_DAYS_SQL: Final = (
'SELECT DISTINCT "date" FROM "LiteLLM_DailyUserSpend" WHERE "date" <= $1 '
'AND ("date" > $2 OR "updated_at" >= $3::timestamp - INTERVAL \'1 hour\') '
'ORDER BY "date"'
)
# Runs can overlap (Redis unreachable, lock expired on a long backfill), so the database keeps the
# later of the stored and the incoming day and scan time in one statement; GREATEST skips NULL.
_ADVANCE_MARKER_SQL: Final = (
'INSERT INTO "LiteLLM_Config" ("param_name", "param_value") '
"VALUES ($1, jsonb_build_object('reconciled_through', $2::text, 'scanned_at', $3::text)) "
'ON CONFLICT ("param_name") DO UPDATE SET "param_value" = jsonb_build_object('
"'reconciled_through', GREATEST(\"LiteLLM_Config\".\"param_value\" ->> 'reconciled_through', "
"EXCLUDED.\"param_value\" ->> 'reconciled_through'), "
"'scanned_at', GREATEST(\"LiteLLM_Config\".\"param_value\" ->> 'scanned_at', "
"EXCLUDED.\"param_value\" ->> 'scanned_at'))"
)
class ReconciledThrough(BaseModel):
"""``reconciled_through`` is the last closed UTC day the global table covers. ``scanned_at`` is
the database clock when the scan behind the last fully successful run started: every per-key
row written before it, on any day through the marker, is in the global table."""
model_config = ConfigDict(frozen=True, extra="ignore")
reconciled_through: str
scanned_at: str | None = None
class _MarkerRow(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore", from_attributes=True)
param_value: object = None
class _DateRow(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
date: str
class _NowRow(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
now: str
today: str
@dataclass(frozen=True, slots=True)
class ReconcileResult:
days_reconciled: tuple[str, ...]
reconciled_through: str | None
failed_day: str | None = None
@dataclass(frozen=True, slots=True)
class _PendingScan:
marker: ReconciledThrough | None
scanned_at: str
days: tuple[str, ...]
def _marker_from_param_value(value: object) -> ReconciledThrough | None:
try:
return (
ReconciledThrough.model_validate_json(value)
if isinstance(value, str)
else ReconciledThrough.model_validate(value)
)
except ValidationError:
return None
async def read_marker(prisma_client: "PrismaClient") -> ReconciledThrough | None:
from litellm.proxy.utils import get_config_param
row: Final = await get_config_param(prisma_client, DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
return None if row is None else _marker_from_param_value(_MarkerRow.model_validate(row).param_value)
async def reconciled_through(prisma_client: "PrismaClient") -> str | None:
"""The last UTC day ``LiteLLM_DailyGlobalSpend`` is known to cover, or None before the first run."""
marker: Final = await read_marker(prisma_client)
return None if marker is None else marker.reconciled_through
async def _advance_marker(prisma_client: "PrismaClient", days: tuple[str, ...], *, scanned_at: str | None) -> None:
"""Move the stored marker to the last of ``days`` and to ``scanned_at`` where those are later
than what is stored, so a slower overlapping run can only add to a faster run's marker."""
from litellm.proxy.utils import invalidate_config_param
await prisma_client.db.execute_raw(
_ADVANCE_MARKER_SQL,
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM,
max(days) if days else None,
scanned_at,
)
await invalidate_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
async def _db_now(prisma_client: "PrismaClient") -> _NowRow:
rows: Final = await prisma_client.db.query_raw(_DB_NOW_SQL)
return _NowRow.model_validate(rows[0])
async def _scan_pending(prisma_client: "PrismaClient") -> _PendingScan:
"""Every closed UTC day (strictly before the database's today) still to roll up, oldest first:
days past the marker, plus any day with per-key rows written since the scan behind the marker.
Before a run has fully succeeded there is no such scan, so every closed day is rolled up."""
marker: Final = await read_marker(prisma_client)
db_now: Final = await _db_now(prisma_client)
last_closed_day: Final = (date.fromisoformat(db_now.today) - timedelta(days=1)).isoformat()
rows: Final = (
await prisma_client.db.query_raw(_ALL_CLOSED_DAYS_SQL, last_closed_day)
if marker is None or marker.scanned_at is None
else await prisma_client.db.query_raw(
_PENDING_DAYS_SQL, last_closed_day, marker.reconciled_through, marker.scanned_at
)
)
return _PendingScan(marker, db_now.now, tuple(_DateRow.model_validate(row).date for row in rows))
async def pending_days(prisma_client: "PrismaClient") -> tuple[str, ...]:
return (await _scan_pending(prisma_client)).days
async def reconcile_day(prisma_client: "PrismaClient", day: str) -> None:
"""Rewrite one day of the global table from the per-key sums. Idempotent: a rerun
overwrites every group with the same totals."""
await prisma_client.db.execute_raw(RECONCILE_DAY_SQL, day)
async def run_daily_global_spend_reconcile(prisma_client: "PrismaClient") -> ReconcileResult:
"""Roll up every pending day, advancing the marker after each; a failing day stops the run
with the marker on the last good day so the next run resumes there. The scan time is only
recorded once every pending day is done, so late rows a failed run saw are found again."""
scan: Final = await _scan_pending(prisma_client)
done: Final = await _reconcile_until_failure(prisma_client, scan)
if len(done) < len(scan.days):
marker: Final = await reconciled_through(prisma_client)
return ReconcileResult(days_reconciled=done, reconciled_through=marker, failed_day=scan.days[len(done)])
if scan.marker is not None or done:
await _advance_marker(prisma_client, done, scanned_at=scan.scanned_at)
return ReconcileResult(days_reconciled=done, reconciled_through=await reconciled_through(prisma_client))
async def _reconcile_until_failure(prisma_client: "PrismaClient", scan: _PendingScan) -> tuple[str, ...]:
for index, day in enumerate(scan.days):
if not await _reconcile_and_record(prisma_client, scan.days[: index + 1]):
return scan.days[:index]
return scan.days
async def _reconcile_and_record(prisma_client: "PrismaClient", done_with_this: tuple[str, ...]) -> bool:
day: Final = done_with_this[-1]
try:
await reconcile_day(prisma_client, day)
await _advance_marker(prisma_client, done_with_this, scanned_at=None)
except Exception as exc: # noqa: BLE001 # one bad day must not lose the days already done
verbose_proxy_logger.exception("Daily global spend reconcile: day %s failed: %s", day, exc)
return False
return True
async def run_scheduled_daily_global_spend_reconcile(
prisma_client: "PrismaClient",
pod_lock_manager: "PodLockManager | None" = None,
alert: Callable[[str], Awaitable[None]] | None = None,
) -> ReconcileResult | None:
"""Run the reconcile under a cross-pod lock so one proxy does the work; the lock only saves
effort (each day is an idempotent rewrite), so an unreachable Redis runs unguarded rather than skipping."""
redis_cache: Final = None if pod_lock_manager is None else pod_lock_manager.redis_cache
if pod_lock_manager is None or redis_cache is None:
return await _run_and_alert(prisma_client, alert=alert)
acquired: Final = await pod_lock_manager.acquire_lock(
cronjob_id=DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID, ttl=DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS
)
if not acquired and await _lock_is_held(pod_lock_manager, redis_cache):
verbose_proxy_logger.info("Daily global spend reconcile: another pod holds the lock, skipping this run")
return None
try:
return await _run_and_alert(prisma_client, alert=alert)
finally:
if acquired:
await pod_lock_manager.release_lock(cronjob_id=DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID)
async def _lock_is_held(pod_lock_manager: "PodLockManager", redis_cache: "RedisCache") -> bool:
try:
lock_key: Final = pod_lock_manager.get_redis_lock_key(DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID)
return bool(await redis_cache.async_get_cache(lock_key))
except Exception as exc: # noqa: BLE001 # an unreadable lock must not skip the run
verbose_proxy_logger.warning("Daily global spend reconcile: could not read the lock: %s", exc)
return False
async def _run_and_alert(
prisma_client: "PrismaClient",
*,
alert: Callable[[str], Awaitable[None]] | None,
) -> ReconcileResult:
result: Final = await run_daily_global_spend_reconcile(prisma_client)
if result.days_reconciled:
verbose_proxy_logger.info(
"Daily global spend reconcile: rolled up %d day(s), reconciled through %s",
len(result.days_reconciled),
result.reconciled_through,
)
if result.failed_day is not None and alert is not None:
await alert(
f"Daily global spend reconcile stopped at {result.failed_day}; usage totals keep reading the per-key "
f"table for ranges past {result.reconciled_through or 'the beginning'} until the next run succeeds."
)
return result

View file

@ -7507,6 +7507,7 @@ def _check_and_merge_model_level_guardrails(
data: dict,
llm_router: Router | None,
trust_client_model_info: bool = True,
model_alias: str | None = None,
) -> dict:
"""
Check if the model has guardrails defined and merge them with existing guardrails in the request data.
@ -7514,6 +7515,7 @@ def _check_and_merge_model_level_guardrails(
Args:
data: The request data dict
llm_router: The LLM router instance to get deployment info from
model_alias: Resolve guardrails for this model group instead of data["model"]
trust_client_model_info: If False, ignore metadata.model_info.id and
resolve guardrails by alias-union only. Set to False on the
pre_call path because add_litellm_data_to_request preserves
@ -7558,13 +7560,13 @@ def _check_and_merge_model_level_guardrails(
# set on ANY eligible deployment still fires (#29652; addresses
# veria-ai HIGH on the single-deployment fallback that would skip
# non-first deployments).
model_alias: Final = data.get("model")
if not isinstance(model_alias, str) or not model_alias:
alias: Final = model_alias if model_alias is not None else data.get("model")
if not isinstance(alias, str) or not alias:
return data
# Pass team_id so team-scoped public model names resolve the same way
# route_request resolves them; otherwise team-scoped deployments are
# invisible to this lookup and their guardrails are silently dropped.
deployments: Final = llm_router.get_model_list(model_name=model_alias, team_id=team_id) or []
deployments: Final = llm_router.get_model_list(model_name=alias, team_id=team_id) or []
seen: Final[set] = set()
union: Final[list] = []
for dep in deployments:

View file

@ -103,6 +103,25 @@ def declared_reasoning_efforts_for_model(model: str, custom_llm_provider: str) -
return declared_reasoning_efforts(entry)
REASONING_EFFORT_STRENGTH_ORDER: Final = ("minimal", "low", "medium", "high", "xhigh", "max")
_STRENGTH_RANK: Final = MappingProxyType({effort: rank for rank, effort in enumerate(REASONING_EFFORT_STRENGTH_ORDER)})
def nearest_declared_reasoning_effort(requested: str, declared: Sequence[str]) -> str:
"""Rounds a request up to the weakest declared level at least as strong as it, and down to the
strongest declared level when it asks for more than the model has, so the caller gets no less
reasoning than it asked for instead of a rejected call. none is the off switch rather than a
strength, so it is never rounded onto the ladder and no level is rounded down to it: a caller
who turned reasoning off must not be billed for it, and a model that cannot turn it off says so
itself. A level outside the strength order is likewise returned as is for upstream to judge."""
ranked: Final = sorted(
(effort for effort in declared if effort in _STRENGTH_RANK), key=lambda effort: _STRENGTH_RANK[effort]
)
if requested in ranked or requested not in _STRENGTH_RANK or not ranked:
return requested
return next((effort for effort in ranked if _STRENGTH_RANK[effort] >= _STRENGTH_RANK[requested]), ranked[-1])
def _supports_none_reasoning_effort(model_info: Mapping[str, object], flag: object) -> bool:
"""Opt-in only where a request path refuses the level. AzureOpenAIGPT5Config raises
UnsupportedParamsError on reasoning_effort='none' without an explicit true, and it is selected

View file

@ -27,13 +27,17 @@ class UpstreamFailure(Exception):
self.__cause__ = cause
def _upstream_failure(error: Exception) -> Exception:
def _upstream_failure(error: Exception, request: LiteLLMOcrRequest) -> Exception:
try:
status, body = _UPSTREAM_ARGS.validate_python(error.args)
headers: Final = _UPSTREAM_HEADERS.validate_python(getattr(error, "headers", None))
except ValidationError:
return error
return UpstreamFailure(httpx.Response(status, content=body.encode(), headers=headers), error)
http_request: Final = httpx.Request("POST", request.api_base or "https://docs.litellm.ai/docs")
return UpstreamFailure(
httpx.Response(status, content=body.encode(), headers=headers, request=http_request),
error,
)
def response(value: Mapping[str, object]) -> OCRResponse:
@ -57,7 +61,7 @@ def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider:
model=request.model.removeprefix(f"{request_provider}/"),
llm_provider=request_provider,
)
original: Final = _upstream_failure(error)
original: Final = _upstream_failure(error, request)
public_error: Final = failures.map_failure(original, request.model, request_provider, arguments(request))
if isinstance(original, UpstreamFailure) and public_error.__context__ is original:
public_error.__context__ = error

View file

@ -302,6 +302,7 @@ class CredentialLiteLLMParams(BaseModel):
aws_bedrock_runtime_endpoint: str | None = None
aws_bedrock_project_id: str | None = None
s3_bucket_name: str | None = None
s3_endpoint_url: str | None = None
s3_region_name: str | None = None
s3_encryption_key_id: str | None = None
aws_batch_role_arn: str | None = None

View file

@ -3778,6 +3778,7 @@ bedrock_batch_litellm_params: Final = (
"aws_batch_role_arn",
"s3_bucket_name",
"s3_region_name",
"s3_endpoint_url",
"s3_output_bucket_name",
"bedrock_tags",
)

View file

@ -4527,7 +4527,7 @@ def get_optional_params(
drop_params=bool(drop_params),
)
else:
optional_params = litellm.MistralConfig().map_openai_params(
optional_params = litellm.VertexAIMistralConfig().map_openai_params(
model=model,
non_default_params=non_default_params,
optional_params=optional_params,
@ -8398,7 +8398,7 @@ class ProviderConfigManager:
elif model in litellm.vertex_mistral_models:
if "codestral" in model:
return litellm.CodestralTextCompletionConfig()
return litellm.MistralConfig()
return litellm.VertexAIMistralConfig()
elif model in litellm.vertex_ai_ai21_models:
return litellm.VertexAIAi21Config()
else:

View file

@ -20853,6 +20853,96 @@
"/v1/audio/transcriptions"
]
},
"deepgram/streaming/nova-3": {
"input_cost_per_second": 8e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0048/60 seconds = $0.00008000 per second",
"note": "Nova-3 monolingual streaming, pay as you go",
"original_pricing_per_minute": 0.0048
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/nova-3-multilingual": {
"input_cost_per_second": 9.667e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0058/60 seconds = $0.00009667 per second",
"note": "Nova-3 multilingual (language=multi) streaming, pay as you go",
"original_pricing_per_minute": 0.0058
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/redact": {
"input_cost_per_second": 3.333e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0020/60 seconds = $0.00003333 per second",
"note": "Redaction add-on (redact query param), streaming, pay as you go",
"original_pricing_per_minute": 0.002
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/keyterm": {
"input_cost_per_second": 2.167e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0013/60 seconds = $0.00002167 per second",
"note": "Keyterm Prompting add-on (keyterm query param), streaming, pay as you go",
"original_pricing_per_minute": 0.0013
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/detect_entities": {
"input_cost_per_second": 2.833e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0017/60 seconds = $0.00002833 per second",
"note": "Entity Detection add-on (detect_entities query param), streaming, pay as you go",
"original_pricing_per_minute": 0.0017
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/streaming/diarize": {
"input_cost_per_second": 3.333e-05,
"litellm_provider": "deepgram",
"metadata": {
"calculation": "$0.0020/60 seconds = $0.00003333 per second",
"note": "Speaker Diarization add-on (diarize / diarize_model query params), streaming, pay as you go",
"original_pricing_per_minute": 0.002
},
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://deepgram.com/pricing",
"supported_endpoints": [
"/v1/listen"
]
},
"deepgram/whisper": {
"input_cost_per_second": 0.0001,
"litellm_provider": "deepgram",
@ -37107,6 +37197,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37188,6 +37282,15 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"none",
"minimal",
"low",
"medium",
"high",
"xhigh",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-2",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37205,6 +37308,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-3",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37222,6 +37330,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-3",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37239,6 +37352,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-3",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37256,6 +37374,15 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"reasoning_effort_levels": [
"none",
"minimal",
"low",
"medium",
"high",
"xhigh",
"max"
],
"source": "https://docs.mistral.ai/models/zai-glm-5-2",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37551,6 +37678,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37611,6 +37742,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37628,6 +37763,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37661,6 +37800,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -37692,6 +37835,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6e-07,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -60037,6 +60184,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6e-07,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -63362,6 +63513,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -63379,6 +63534,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -63396,6 +63555,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -63413,6 +63576,10 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6e-07,
"reasoning_effort_levels": [
"none",
"high"
],
"source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03",
"supports_assistant_prefill": true,
"supports_function_calling": true,

View file

@ -820,6 +820,37 @@ model LiteLLM_DailyUserSpend {
@@index([endpoint])
}
// Key-free daily rollup of LiteLLM_DailyUserSpend, read by the global usage view
model LiteLLM_DailyGlobalSpend {
id String @id @default(uuid())
date String
model String?
model_group String?
custom_llm_provider String?
mcp_namespaced_tool_name String?
endpoint String?
prompt_tokens BigInt @default(0)
completion_tokens BigInt @default(0)
cache_read_input_tokens BigInt @default(0)
cache_creation_input_tokens BigInt @default(0)
compression_saved_tokens BigInt @default(0)
compression_savings_spend Float @default(0.0)
prompt_caching_savings_spend Float @default(0.0)
gateway_injected_caching_savings_spend Float @default(0.0)
autorouter_savings_spend Float @default(0.0)
spend Float @default(0.0)
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@unique([date, model, model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
}
// Track daily organization spend metrics per model and key
model LiteLLM_DailyOrganizationSpend {
id String @id @default(uuid())

View file

@ -1,203 +0,0 @@
"""
Base test class for OCR functionality across different providers.
This follows the same pattern as BaseLLMChatTest in tests/llm_translation/base_llm_unit_tests.py
"""
import pytest
import litellm
import os
from abc import ABC, abstractmethod
# Test resources
TEST_IMAGE_PATH = "test_image_edit.png"
# Tiny in-repo PDF served via jsdelivr (sha-pinned, immutable). The arxiv
# PDF previously used here was several MB — once base64-encoded into the
# Vertex OCR request it ballooned cassettes past 100 MB per test. Keep
# the URL stable across runs so cassettes don't churn.
TEST_PDF_URL = (
"https://cdn.jsdelivr.net/gh/BerriAI/litellm"
"@d769e81c90d453240c61fc572cdb27fae06a89d0"
"/tests/llm_translation/fixtures/dummy.pdf"
)
class BaseOCRTest(ABC):
"""
Abstract base test class that enforces common OCR tests across all providers.
Each provider-specific test class should inherit from this and implement
get_base_ocr_call_args() to return provider-specific configuration.
"""
@abstractmethod
def get_base_ocr_call_args(self) -> dict:
"""Must return the base OCR call args for the specific provider"""
pass
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_basic_ocr_with_url(self, sync_mode):
"""
Test basic OCR with a public URL.
"""
litellm._turn_on_debug()
base_ocr_call_args = self.get_base_ocr_call_args()
print("BASE OCR Call args=", base_ocr_call_args)
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
if sync_mode:
response = litellm.ocr(
document={"type": "document_url", "document_url": TEST_PDF_URL},
**base_ocr_call_args,
)
else:
response = await litellm.aocr(
document={"type": "document_url", "document_url": TEST_PDF_URL},
**base_ocr_call_args,
)
print(f"\n{'='*80}")
print(f"Sync Mode: {sync_mode}")
print(f"Response type: {type(response)}")
print(
f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}"
)
# Check if response has expected OCR format
assert hasattr(response, "pages"), "Response should have 'pages' attribute"
assert hasattr(response, "model"), "Response should have 'model' attribute"
assert hasattr(
response, "object"
), "Response should have 'object' attribute"
assert (
response.object == "ocr"
), f"Expected object='ocr', got '{response.object}'"
# Validate pages structure
assert isinstance(response.pages, list), "pages should be a list"
assert len(response.pages) > 0, "Should have at least one page"
# Check first page structure
first_page = response.pages[0]
assert hasattr(first_page, "index"), "Page should have 'index' attribute"
assert hasattr(
first_page, "markdown"
), "Page should have 'markdown' attribute"
# Extract text from all pages for validation
total_text = "\n\n".join(
page.markdown for page in response.pages if page.markdown
)
print(f"Total pages: {len(response.pages)}")
print(f"Total extracted text length: {len(total_text)} characters")
print(f"First 200 chars: {total_text[:200]}")
print(f"Model: {response.model}")
if response.usage_info:
print(f"Pages processed: {response.usage_info.pages_processed}")
print(f"{'='*80}\n")
assert len(total_text) > 0, "Should extract some text from the document"
#########################################################
# validate we get a response cost in hidden parameters
#########################################################
hidden_params = response._hidden_params
assert isinstance(
hidden_params, dict
), "Hidden parameters should be a dictionary"
print("response usage_info:", response.usage_info)
response_cost = hidden_params.get("response_cost")
assert (
response_cost is not None
), "Response cost should be in hidden parameters"
assert response_cost > 0, "Response cost should be greater than 0"
print("response_cost=", response_cost)
except litellm.RateLimitError as e:
error_msg = str(e)
if "Quota exceeded" in error_msg or "RESOURCE_EXHAUSTED" in error_msg:
pytest.skip(f"Quota exceeded - {error_msg}")
else:
pytest.skip(f"Rate limit exceeded - {error_msg}")
except litellm.InternalServerError:
pytest.skip("Model is overloaded")
except litellm.BadRequestError as e:
error_msg = str(e)
if (
"URL_REJECTED" in error_msg
or "Cannot fetch content from the provided URL" in error_msg
):
pytest.skip(f"URL rejected by provider - {error_msg}")
else:
pytest.fail(f"OCR call failed: {str(e)}")
except Exception as e:
pytest.fail(f"OCR call failed: {str(e)}")
def test_ocr_response_structure(self):
"""
Test that the OCR response has the correct structure.
"""
litellm.set_verbose = True
base_ocr_call_args = self.get_base_ocr_call_args()
try:
response = litellm.ocr(
document={"type": "document_url", "document_url": TEST_PDF_URL},
**base_ocr_call_args,
)
# Validate response structure
assert hasattr(response, "pages"), "Response should have 'pages' attribute"
assert hasattr(response, "model"), "Response should have 'model' attribute"
assert hasattr(
response, "object"
), "Response should have 'object' attribute"
assert hasattr(
response, "usage_info"
), "Response should have 'usage_info' attribute"
assert isinstance(response.pages, list), "pages should be a list"
assert len(response.pages) > 0, "Should have at least one page"
assert response.object == "ocr", "object should be 'ocr'"
# Validate first page structure
first_page = response.pages[0]
assert hasattr(first_page, "index"), "Page should have 'index' attribute"
assert hasattr(
first_page, "markdown"
), "Page should have 'markdown' attribute"
assert isinstance(first_page.markdown, str), "markdown should be a string"
print(f"\nResponse structure validated:")
print(f" - object: {response.object}")
print(f" - model: {response.model}")
print(f" - pages: {len(response.pages)}")
if response.usage_info:
print(f" - pages_processed: {response.usage_info.pages_processed}")
print(f" - doc_size_bytes: {response.usage_info.doc_size_bytes}")
except litellm.RateLimitError as e:
error_msg = str(e)
if "Quota exceeded" in error_msg or "RESOURCE_EXHAUSTED" in error_msg:
pytest.skip(f"Quota exceeded - {error_msg}")
else:
pytest.skip(f"Rate limit exceeded - {error_msg}")
except litellm.InternalServerError:
pytest.skip("Model is overloaded")
except litellm.BadRequestError as e:
error_msg = str(e)
if (
"URL_REJECTED" in error_msg
or "Cannot fetch content from the provided URL" in error_msg
):
pytest.skip(f"URL rejected by provider - {error_msg}")
else:
pytest.fail(f"OCR response structure test failed: {str(e)}")
except Exception as e:
pytest.fail(f"OCR response structure test failed: {str(e)}")

View file

@ -1,29 +0,0 @@
"""
Test OCR functionality with Azure AI API.
Note: Azure AI OCR automatically converts URLs to base64 data URIs since
the Azure AI endpoint doesn't have internet access.
"""
import os
from base_ocr_unit_tests import BaseOCRTest
class TestAzureAIOCR(BaseOCRTest):
"""
Test class for Azure AI OCR functionality.
Inherits from BaseOCRTest and provides Azure AI-specific configuration.
Note: For Azure AI, LiteLLM will automatically convert URLs to base64 data URIs before
sending to the API, since Azure AI OCR endpoint doesn't have internet access.
"""
def get_base_ocr_call_args(self) -> dict:
"""
Return the base OCR call args for Azure AI.
"""
return {
"model": "azure_ai/mistral-document-ai-2512",
"api_key": os.getenv("AZURE_API_KEY"),
"api_base": os.getenv("AZURE_API_BASE"),
}

View file

@ -1,53 +1,13 @@
"""
Test OCR functionality with Azure Document Intelligence API.
Azure Document Intelligence provides advanced document analysis capabilities
using the v4.0 (2024-11-30) API.
"""
import os
"""Azure Document Intelligence request transformation: Mistral-shaped `pages` to Azure's query string."""
import pytest
from base_ocr_unit_tests import BaseOCRTest
from litellm.constants import AZURE_DOCUMENT_INTELLIGENCE_API_VERSION
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
class TestAzureDocumentIntelligenceOCR(BaseOCRTest):
"""
Test class for Azure Document Intelligence OCR functionality.
Inherits from BaseOCRTest and provides Azure Document Intelligence-specific configuration.
Tests the azure_ai/doc-intelligence/<model> provider route.
"""
def get_base_ocr_call_args(self) -> dict:
"""
Return the base OCR call args for Azure Document Intelligence.
Uses prebuilt-layout model which is closest to Mistral OCR format.
"""
# Check for required environment variables
api_key = os.environ.get("AZURE_DOCUMENT_INTELLIGENCE_API_KEY")
endpoint = os.environ.get("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
if not api_key or not endpoint:
pytest.skip(
"AZURE_DOCUMENT_INTELLIGENCE_API_KEY and AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT "
"environment variables are required for Azure Document Intelligence tests"
)
return {
"model": "azure_ai/doc-intelligence/prebuilt-layout",
"api_key": api_key,
"api_base": endpoint,
}
class TestAzureDocumentIntelligencePagesParam:
"""
Unit tests for the Mistral-compatible `pages` parameter translation to
@ -101,7 +61,7 @@ class TestAzureDocumentIntelligencePagesParam:
cfg.map_ocr_params({"pages": [True, False]}, {}, "prebuilt-layout")
def test_map_ocr_params_unsupported_type_raises(self, cfg):
with pytest.raises(ValueError, match='based, Mistral-style\\) or a string like'):
with pytest.raises(ValueError, match="based, Mistral-style\\) or a string like"):
cfg.map_ocr_params({"pages": 5}, {}, "prebuilt-layout")
def test_get_complete_url_appends_pages_query(self, cfg):
@ -110,9 +70,7 @@ class TestAzureDocumentIntelligencePagesParam:
model="azure_ai/doc-intelligence/prebuilt-layout",
optional_params={"pages": "1-3,5"},
)
assert (
f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url
), url
assert f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url, url
assert "pages=1-3,5" in url, url
assert "/documentintelligence/documentModels/prebuilt-layout:analyze" in url
@ -168,4 +126,3 @@ class TestAzureDocumentIntelligencePagesParam:
assert "pages=3,4,5,6,7,8,9" in url
assert req.data == {"urlSource": "https://example.com/x.pdf"}

View file

@ -0,0 +1,317 @@
"""Live provider x auth x input coverage for ``litellm.ocr`` / ``litellm.aocr``.
Each ``Case`` is one hand-picked cell, not the full cross product: every provider
exercises each of its credential kinds in both ``explicit`` (kwargs) and ``env``
(monkeypatched environment) mode at least once, and every input kind a provider
accepts is exercised at least once. Sync and async are spread across the cells.
Every cell also checks the success callback saw the same response and cost.
"""
from __future__ import annotations
import asyncio
import base64
import io
import os
import re
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from types import MappingProxyType
from typing import Final, Literal
import pytest
import litellm
from litellm import Router
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.base_llm.ocr.transformation import OCRResponse
Document = Mapping[str, object]
AuthMode = Literal["explicit", "env"]
CallStyle = Literal["sync", "async"]
@dataclass(frozen=True, slots=True)
class LoggedCall:
payload: Mapping[str, object]
response: object
class RecordingLogger(CustomLogger):
def __init__(self) -> None:
super().__init__() # pyright: ignore[reportUnknownMemberType] # CustomLogger.__init__ is untyped
self.calls: Final[list[LoggedCall]] = [] # mutable-ok: append-only sink the callback hooks write into
def _record(self, kwargs: Mapping[str, object], response_obj: object) -> None:
payload: Final = _string_keyed(kwargs.get("standard_logging_object"))
self.calls.append(LoggedCall(payload, response_obj))
def log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self._record(kwargs, response_obj)
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self._record(kwargs, response_obj)
async def wait_for_call(self, timeout: float = 10.0) -> LoggedCall:
deadline: Final = asyncio.get_running_loop().time() + timeout
while not self.calls:
assert asyncio.get_running_loop().time() < deadline, "success callback never fired"
await asyncio.sleep(0.05)
assert len(self.calls) == 1, self.calls
return self.calls[0]
@pytest.fixture
def logger(monkeypatch: pytest.MonkeyPatch) -> RecordingLogger:
recorder: Final = RecordingLogger()
monkeypatch.setattr(litellm, "callbacks", [recorder])
for registry in ("success_callback", "_async_success_callback", "failure_callback", "_async_failure_callback"):
monkeypatch.setattr(litellm, registry, [])
return recorder
TESTS_DIR: Final = Path(__file__).resolve().parents[1]
PDF_PATH: Final = TESTS_DIR / "llm_translation" / "fixtures" / "dummy.pdf"
PNG_PATH: Final = TESTS_DIR / "image_gen_tests" / "test_image.png"
PINNED_CDN: Final = "https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0"
PDF_URL: Final = f"{PINNED_CDN}/tests/llm_translation/fixtures/dummy.pdf"
PNG_URL: Final = f"{PINNED_CDN}/tests/image_gen_tests/test_image.png"
PDF_TEXT: Final = "Test PDF File"
PNG_TEXT: Final = "LiteLLM"
class _NamedReader(io.BytesIO):
def __init__(self, path: Path) -> None:
super().__init__(path.read_bytes())
self.name: Final = path.name
def _data_uri(path: Path, mime: str) -> str:
return f"data:{mime};base64,{base64.b64encode(path.read_bytes()).decode()}"
@dataclass(frozen=True, slots=True)
class Input:
id: str
build: Callable[[], Document]
expected_text: str
PDF_BY_URL: Final = Input("pdf_url", lambda: {"type": "document_url", "document_url": PDF_URL}, PDF_TEXT)
PNG_BY_URL: Final = Input("image_url", lambda: {"type": "image_url", "image_url": PNG_URL}, PNG_TEXT)
PDF_DATA_URI: Final = Input(
"pdf_data_uri",
lambda: {"type": "document_url", "document_url": _data_uri(PDF_PATH, "application/pdf")},
PDF_TEXT,
)
PNG_DATA_URI: Final = Input(
"image_data_uri", lambda: {"type": "image_url", "image_url": _data_uri(PNG_PATH, "image/png")}, PNG_TEXT
)
PDF_AS_PATH: Final = Input("pdf_path", lambda: {"type": "file", "file": PDF_PATH}, PDF_TEXT)
PDF_AS_BYTES: Final = Input(
"pdf_bytes", lambda: {"type": "file", "file": PDF_PATH.read_bytes(), "mime_type": "application/pdf"}, PDF_TEXT
)
PNG_AS_BYTES: Final = Input(
"image_bytes", lambda: {"type": "file", "file": PNG_PATH.read_bytes(), "mime_type": "image/png"}, PNG_TEXT
)
PNG_AS_FILE_OBJECT: Final = Input(
"image_file_object", lambda: {"type": "file", "file": _NamedReader(PNG_PATH)}, PNG_TEXT
)
@dataclass(frozen=True, slots=True)
class Secret:
"""One credential value: the ``litellm.ocr`` kwarg it travels in, the env var litellm reads
when the kwarg is omitted, and the env var that holds the value in the test process."""
kwarg: str
env: str
source: str | None = None
@property
def source_env(self) -> str:
return self.source or self.env
@dataclass(frozen=True, slots=True)
class Credential:
id: str
secrets: tuple[Secret, ...]
@dataclass(frozen=True, slots=True)
class Provider:
id: str
model: str
credentials: tuple[Credential, ...]
params: Mapping[str, str] = MappingProxyType({})
@property
def env_vars(self) -> frozenset[str]:
return frozenset(secret.env for credential in self.credentials for secret in credential.secrets)
MISTRAL_KEY: Final = Credential("api_key", (Secret("api_key", "MISTRAL_API_KEY"),))
COHERE_KEY: Final = Credential("api_key", (Secret("api_key", "COHERE_API_KEY"),))
REDUCTO_KEY: Final = Credential("api_key", (Secret("api_key", "REDUCTO_API_KEY"),))
AZURE_ENTRA_SECRETS: Final = (
Secret("tenant_id", "AZURE_TENANT_ID", "AZURE_FOUNDRY_TENANT_ID"),
Secret("client_id", "AZURE_CLIENT_ID", "AZURE_FOUNDRY_ADMIN_CLIENT_ID"),
Secret("client_secret", "AZURE_CLIENT_SECRET", "AZURE_FOUNDRY_ADMIN_CLIENT_SECRET"),
)
AZURE_AI_BASE: Final = Secret("api_base", "AZURE_AI_API_BASE")
AZURE_AI_KEY: Final = Credential("api_key", (AZURE_AI_BASE, Secret("api_key", "AZURE_AI_API_KEY")))
AZURE_AI_ENTRA: Final = Credential("entra", (AZURE_AI_BASE, *AZURE_ENTRA_SECRETS))
AZURE_DI_BASE: Final = Secret("api_base", "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
AZURE_DI_KEY: Final = Credential("api_key", (AZURE_DI_BASE, Secret("api_key", "AZURE_DOCUMENT_INTELLIGENCE_API_KEY")))
AZURE_DI_ENTRA: Final = Credential("entra", (AZURE_DI_BASE, *AZURE_ENTRA_SECRETS))
VERTEX_SERVICE_ACCOUNT: Final = Credential(
"service_account",
(Secret("vertex_credentials", "VERTEXAI_CREDENTIALS"), Secret("vertex_project", "VERTEXAI_PROJECT")),
)
MISTRAL: Final = Provider("mistral", "mistral/mistral-ocr-latest", (MISTRAL_KEY,))
AZURE_AI_MISTRAL: Final = Provider(
"azure_ai_mistral", "azure_ai/mistral-document-ai-2512", (AZURE_AI_KEY, AZURE_AI_ENTRA)
)
AZURE_DOC_INTELLIGENCE: Final = Provider(
"azure_doc_intelligence", "azure_ai/doc-intelligence/prebuilt-layout", (AZURE_DI_KEY, AZURE_DI_ENTRA)
)
COHERE: Final = Provider("cohere", "cohere/parse-v5.0", (COHERE_KEY,))
REDUCTO_V3: Final = Provider("reducto_v3", "reducto/parse-v3", (REDUCTO_KEY,))
REDUCTO_LEGACY: Final = Provider("reducto_legacy", "reducto/parse-legacy", (REDUCTO_KEY,))
VERTEX_MISTRAL: Final = Provider(
"vertex_mistral",
"vertex_ai/mistral-ocr-2505",
(VERTEX_SERVICE_ACCOUNT,),
MappingProxyType({"vertex_location": "us-central1"}),
)
@dataclass(frozen=True, slots=True)
class Case:
provider: Provider
credential: Credential
auth: AuthMode
document: Input
call: CallStyle
@property
def id(self) -> str:
return f"{self.provider.id}-{self.credential.id}-{self.auth}-{self.document.id}-{self.call}"
def bind_credentials(self, monkeypatch: pytest.MonkeyPatch) -> Mapping[str, str]:
"""Clear every env var the provider could fall back to, then supply this case's values via kwargs or env."""
values: Final = {secret: os.environ.get(secret.source_env) for secret in self.credential.secrets}
missing: Final = tuple(secret.source_env for secret, value in values.items() if not value)
if missing:
pytest.skip(f"{', '.join(missing)} not set")
for env_var in self.provider.env_vars:
monkeypatch.delenv(env_var, raising=False)
if self.auth == "explicit":
return {secret.kwarg: value for secret, value in values.items() if value}
for secret, value in values.items():
monkeypatch.setenv(secret.env, value or "")
return {}
async def run(self, credentials: Mapping[str, str]) -> OCRResponse:
kwargs: Final = {**self.provider.params, **credentials}
document: Final = self.document.build()
response: Final = (
await litellm.aocr(model=self.provider.model, document=document, **kwargs) # pyright: ignore[reportUnknownMemberType] # @client erases the signature
if self.call == "async"
else litellm.ocr(model=self.provider.model, document=document, **kwargs)
)
assert isinstance(response, OCRResponse)
return response
CASES: Final = (
Case(MISTRAL, MISTRAL_KEY, "explicit", PDF_BY_URL, "sync"),
Case(MISTRAL, MISTRAL_KEY, "env", PNG_BY_URL, "async"),
Case(MISTRAL, MISTRAL_KEY, "explicit", PDF_AS_PATH, "sync"),
Case(MISTRAL, MISTRAL_KEY, "explicit", PNG_AS_BYTES, "async"),
Case(MISTRAL, MISTRAL_KEY, "explicit", PNG_AS_FILE_OBJECT, "sync"),
Case(AZURE_AI_MISTRAL, AZURE_AI_KEY, "explicit", PDF_BY_URL, "sync"),
Case(AZURE_AI_MISTRAL, AZURE_AI_KEY, "env", PNG_BY_URL, "async"),
Case(AZURE_AI_MISTRAL, AZURE_AI_ENTRA, "explicit", PDF_AS_PATH, "sync"),
Case(AZURE_AI_MISTRAL, AZURE_AI_ENTRA, "env", PDF_DATA_URI, "async"),
Case(AZURE_DOC_INTELLIGENCE, AZURE_DI_KEY, "explicit", PDF_BY_URL, "sync"),
Case(AZURE_DOC_INTELLIGENCE, AZURE_DI_KEY, "env", PNG_AS_BYTES, "async"),
Case(AZURE_DOC_INTELLIGENCE, AZURE_DI_ENTRA, "explicit", PNG_BY_URL, "async"),
Case(AZURE_DOC_INTELLIGENCE, AZURE_DI_ENTRA, "env", PDF_AS_PATH, "sync"),
Case(COHERE, COHERE_KEY, "explicit", PNG_BY_URL, "sync"),
Case(COHERE, COHERE_KEY, "env", PNG_DATA_URI, "async"),
Case(REDUCTO_V3, REDUCTO_KEY, "explicit", PDF_AS_PATH, "sync"),
Case(REDUCTO_V3, REDUCTO_KEY, "env", PNG_AS_BYTES, "async"),
Case(REDUCTO_V3, REDUCTO_KEY, "explicit", PDF_DATA_URI, "async"),
Case(REDUCTO_LEGACY, REDUCTO_KEY, "explicit", PDF_AS_BYTES, "sync"),
Case(VERTEX_MISTRAL, VERTEX_SERVICE_ACCOUNT, "explicit", PDF_BY_URL, "sync"),
Case(VERTEX_MISTRAL, VERTEX_SERVICE_ACCOUNT, "env", PNG_BY_URL, "async"),
)
def _response_cost(response: OCRResponse) -> float:
response_cost: Final[object] = response._hidden_params.get("response_cost") # pyright: ignore[reportPrivateUsage, reportUnknownMemberType, reportUnknownVariableType] # response_cost is only surfaced on _hidden_params
assert isinstance(response_cost, float) and response_cost > 0
return response_cost
def _assert_ocr_response(response: OCRResponse, model: str, expected_text: str) -> None:
assert response.object == "ocr"
assert response.model == model.split("/", 1)[1]
assert [page.index for page in response.pages] == list(range(len(response.pages)))
text: Final = re.sub(r"\s+", " ", " ".join(page.markdown for page in response.pages))
assert expected_text.lower() in text.lower(), text
assert response.usage_info is not None
assert response.usage_info.pages_processed == len(response.pages)
_response_cost(response)
def _string_keyed(value: object) -> Mapping[str, object]:
assert isinstance(value, Mapping), type(value)
items: Final = tuple(value.items()) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType, reportUnknownArgumentType] # narrowed from object
return MappingProxyType({str(key): value for key, value in items}) # pyright: ignore[reportUnknownVariableType, reportUnknownArgumentType] # narrowed from object
def _assert_logged(logged: LoggedCall, response: OCRResponse, model: str, logged_model: str, call: CallStyle) -> None:
assert isinstance(logged.response, OCRResponse)
assert logged.response.pages == response.pages
assert logged.payload["status"] == "success"
assert logged.payload["call_type"] == ("aocr" if call == "async" else "ocr")
assert logged.payload["custom_llm_provider"] == model.split("/", 1)[0]
assert logged.payload["model"] == logged_model
assert logged.payload["response_cost"] == _response_cost(response)
@pytest.mark.parametrize("case", CASES, ids=[case.id for case in CASES])
async def test_ocr(case: Case, monkeypatch: pytest.MonkeyPatch, logger: RecordingLogger) -> None:
credentials: Final = case.bind_credentials(monkeypatch)
response: Final = await case.run(credentials)
_assert_ocr_response(response, case.provider.model, case.document.expected_text)
_assert_logged(await logger.wait_for_call(), response, case.provider.model, response.model, case.call)
async def test_router_aocr(monkeypatch: pytest.MonkeyPatch, logger: RecordingLogger) -> None:
case: Final = Case(MISTRAL, MISTRAL_KEY, "explicit", PDF_BY_URL, "async")
router: Final = Router(
model_list=[
{
"model_name": "ocr-alias",
"litellm_params": {"model": MISTRAL.model, **case.bind_credentials(monkeypatch)},
}
]
)
response: Final = await router.aocr(model="ocr-alias", document=PDF_BY_URL.build()) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # Router.aocr is untyped
assert isinstance(response, OCRResponse)
_assert_ocr_response(response, MISTRAL.model, PDF_TEXT)
_assert_logged(await logger.wait_for_call(), response, MISTRAL.model, MISTRAL.model, case.call)

View file

@ -1,92 +0,0 @@
"""
Test OCR functionality with Mistral API.
"""
import os
import sys
import pytest
import litellm
from litellm import Router
from base_ocr_unit_tests import BaseOCRTest, TEST_PDF_URL
class TestMistralOCR(BaseOCRTest):
"""
Test class for Mistral OCR functionality.
"""
def get_base_ocr_call_args(self) -> dict:
"""Return the base OCR call args for Mistral"""
return {
"model": "mistral/mistral-ocr-latest",
"api_key": os.getenv("MISTRAL_API_KEY"),
}
@pytest.mark.asyncio
async def test_router_aocr_with_mistral():
"""
Test OCR with Router using Mistral OCR deployment.
"""
litellm.set_verbose = True
# Create router with Mistral OCR deployment
router = Router(
model_list=[
{
"model_name": "mistral-ocr",
"litellm_params": {
"model": "mistral/mistral-ocr-latest",
"api_key": os.getenv("MISTRAL_API_KEY"),
},
}
]
)
try:
# Call OCR through router
response = await router.aocr(
model="mistral-ocr",
document={"type": "document_url", "document_url": TEST_PDF_URL},
)
print(f"\n{'='*80}")
print("Router OCR Test")
print(f"Response type: {type(response)}")
print(
f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}"
)
# Check if response has expected Mistral OCR format
assert hasattr(response, "pages"), "Response should have 'pages' attribute"
assert hasattr(response, "model"), "Response should have 'model' attribute"
assert hasattr(response, "object"), "Response should have 'object' attribute"
assert (
response.object == "ocr"
), f"Expected object='ocr', got '{response.object}'"
# Validate pages structure
assert isinstance(response.pages, list), "pages should be a list"
assert len(response.pages) > 0, "Should have at least one page"
# Check first page structure
first_page = response.pages[0]
assert hasattr(first_page, "index"), "Page should have 'index' attribute"
assert hasattr(first_page, "markdown"), "Page should have 'markdown' attribute"
# Extract text from all pages for validation
total_text = "\n\n".join(
page.markdown for page in response.pages if page.markdown
)
print(f"Total pages: {len(response.pages)}")
print(f"Total extracted text length: {len(total_text)} characters")
print(f"First 200 chars: {total_text[:200]}")
print(f"Model: {response.model}")
if response.usage_info:
print(f"Pages processed: {response.usage_info.pages_processed}")
print(f"{'='*80}\n")
assert len(total_text) > 0, "Should extract some text from the document"
except Exception as e:
pytest.fail(f"Router OCR call failed: {str(e)}")

View file

@ -1,117 +1,8 @@
"""
Test OCR functionality with Vertex AI OCR APIs (Mistral and DeepSeek).
"""Vertex AI OCR config routing and DeepSeek request shaping (no network)."""
Note: Vertex AI OCR automatically converts URLs to base64 data URIs since
the Vertex AI endpoint doesn't have internet access.
"""
import json
import os
import tempfile
from typing import Final
import pytest
from base_ocr_unit_tests import BaseOCRTest
def load_vertex_ai_credentials():
"""Load Vertex AI credentials for tests"""
# Define the path to the vertex_key.json file
print("loading vertex ai credentials")
filepath = os.path.dirname(os.path.abspath(__file__))
vertex_key_path = filepath + "/vertex_key.json"
# Read the existing content of the file or create an empty dictionary
try:
with open(vertex_key_path, "r") as file:
# Read the file content
print("Read vertexai file path")
content = file.read()
# If the file is empty or not valid JSON, create an empty dictionary
if not content or not content.strip():
service_account_key_data = {}
else:
# Attempt to load the existing JSON content
file.seek(0)
service_account_key_data = json.load(file)
except FileNotFoundError:
# If the file doesn't exist, create an empty dictionary
service_account_key_data = {}
# Update the service_account_key_data with environment variables
private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "")
private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "")
private_key = private_key.replace("\\n", "\n")
service_account_key_data["private_key_id"] = private_key_id
service_account_key_data["private_key"] = private_key
# Create a temporary file
with tempfile.NamedTemporaryFile(mode="w+", delete=False) as temp_file:
# Write the updated content to the temporary files
json.dump(service_account_key_data, temp_file, indent=2)
# Export the temporary file as GOOGLE_APPLICATION_CREDENTIALS
os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name)
class TestVertexAIMistralOCR(BaseOCRTest):
"""
Test class for Vertex AI Mistral OCR functionality.
Inherits from BaseOCRTest and provides Vertex AI-specific configuration.
Note: For Vertex AI, LiteLLM will automatically convert URLs to base64 data URIs before
sending to the API, since Vertex AI OCR endpoint doesn't have internet access.
"""
def setup_method(self):
if os.environ.get("LITELLM_RUN_LIVE_VERTEX_MISTRAL_OCR_TESTS") != "1":
pytest.skip("Live Vertex AI Mistral OCR E2E tests are opt-in")
if os.environ.get("CASSETTE_REDIS_URL"):
pytest.skip(
"Live Vertex AI Mistral OCR E2E tests cannot run under VCR replay"
)
def get_base_ocr_call_args(self) -> dict:
"""
Return the base OCR call args for Vertex AI Mistral OCR.
"""
load_vertex_ai_credentials()
return {
"model": "vertex_ai/mistral-ocr-2505",
"vertex_location": "us-central1",
}
class TestVertexAIDeepSeekOCR(BaseOCRTest):
"""
Test class for Vertex AI DeepSeek OCR functionality.
Inherits from BaseOCRTest and provides Vertex AI-specific configuration.
Note: DeepSeek OCR uses the chat completion API format through the openapi endpoint.
Note: DeepSeek OCR does not support PDF URLs - only image URLs and base64 data.
"""
def get_base_ocr_call_args(self) -> dict:
"""
Return the base OCR call args for Vertex AI DeepSeek OCR.
"""
load_vertex_ai_credentials()
return {
"model": "vertex_ai/deepseek-ocr-maas",
"vertex_location": "us-central1",
}
# Skip PDF URL tests for DeepSeek OCR as it doesn't support PDF URLs
@pytest.mark.skip(reason="DeepSeek OCR does not support PDF URLs")
async def test_basic_ocr_with_url(self, sync_mode):
"""Skip this test for DeepSeek OCR - PDF URLs not supported"""
pass
@pytest.mark.skip(reason="DeepSeek OCR does not support PDF URLs")
def test_ocr_response_structure(self):
"""Skip this test for DeepSeek OCR - PDF URLs not supported"""
pass
def test_vertex_ai_ocr_routing():
@ -126,21 +17,19 @@ def test_vertex_ai_ocr_routing():
# Test DeepSeek OCR routing
deepseek_config = get_vertex_ai_ocr_config("vertex_ai/deepseek-ocr-maas")
assert isinstance(
deepseek_config, VertexAIDeepSeekOCRConfig
), "DeepSeek model should route to VertexAIDeepSeekOCRConfig"
assert isinstance(deepseek_config, VertexAIDeepSeekOCRConfig), (
"DeepSeek model should route to VertexAIDeepSeekOCRConfig"
)
# Test Mistral OCR routing (should use default VertexAIOCRConfig)
mistral_config = get_vertex_ai_ocr_config("vertex_ai/mistral-ocr-2505")
assert isinstance(
mistral_config, VertexAIOCRConfig
), "Mistral model should route to VertexAIOCRConfig"
assert isinstance(mistral_config, VertexAIOCRConfig), "Mistral model should route to VertexAIOCRConfig"
# Test other DeepSeek variants
deepseek_variant = get_vertex_ai_ocr_config("vertex_ai/deepseek-ocr-maas")
assert isinstance(
deepseek_variant, VertexAIDeepSeekOCRConfig
), "DeepSeek variant should route to VertexAIDeepSeekOCRConfig"
assert isinstance(deepseek_variant, VertexAIDeepSeekOCRConfig), (
"DeepSeek variant should route to VertexAIDeepSeekOCRConfig"
)
@pytest.mark.parametrize("model", ("deepseek-ocr-maas", "deepseek-ai/deepseek-ocr-maas"))

View file

@ -1,13 +0,0 @@
{
"type": "service_account",
"project_id": "litellm-ci-cd",
"private_key_id": "",
"private_key": "",
"client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com",
"client_id": "116563532503305622785",
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
"auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs",
"client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com",
"universe_domain": "googleapis.com"
}

View file

@ -278,6 +278,8 @@ def test_extract_credentials_all_supported_keys():
"vertex_credentials",
"gcs_bucket_name",
"bucket_name",
"s3_endpoint_url",
"s3_region_name",
"timeout",
"max_retries",
}

View file

@ -55,6 +55,18 @@ class TestGetLitellmParamsKwargsExtraction:
assert result["timeout"] == 30
assert result["rpm"] == 100
def test_s3_endpoint_kwargs_are_extracted_when_provided(self):
result = get_litellm_params(
s3_endpoint_url="https://bucket.vpce-abc.s3.us-east-1.vpce.amazonaws.com",
s3_region_name="us-east-1",
)
assert result["s3_endpoint_url"] == "https://bucket.vpce-abc.s3.us-east-1.vpce.amazonaws.com"
assert result["s3_region_name"] == "us-east-1"
result_without_s3_kwargs = get_litellm_params()
assert "s3_endpoint_url" not in result_without_s3_kwargs
assert "s3_region_name" not in result_without_s3_kwargs
def test_subset_of_kwargs_only_includes_provided(self):
"""Only provided kwargs appear, others remain absent."""
result = get_litellm_params(azure_ad_token="token123")

View file

@ -2271,6 +2271,19 @@ class TestBedrockFileContentTransformation:
authorization = litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM]["Authorization"]
assert "/eu-west-1/s3/aws4_request" in authorization
def test_s3_request_target_uses_configured_endpoint_url(self):
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
lp = get_litellm_params(
aws_region_name="us-east-1",
s3_endpoint_url="https://bucket.vpce-abc.s3.us-east-1.vpce.amazonaws.com",
)
assert BedrockFilesConfig()._s3_request_target(
optional_params={}, litellm_params=lp
).endpoint_url == "https://bucket.vpce-abc.s3.us-east-1.vpce.amazonaws.com"
def test_validate_environment_merges_and_pops_signed_get_headers(self):
from litellm.llms.bedrock.files.transformation import (
S3_SIGNED_REQUEST_HEADERS_PARAM,

View file

@ -0,0 +1,326 @@
import math
from collections.abc import Mapping, Sequence
from typing import Final
from urllib.parse import parse_qs, urlparse
import pytest
import litellm
from litellm.llms.deepgram.common_utils import (
deepgram_listen_addon_pricing_models,
deepgram_listen_audio_seconds,
deepgram_listen_callback_params,
deepgram_listen_channel_count,
deepgram_listen_is_priced,
deepgram_listen_model,
deepgram_listen_pricing_model,
deepgram_listen_registry_key,
deepgram_listen_requested_model,
deepgram_listen_transcript,
deepgram_listen_websocket_target,
)
NOVA_3_URL: Final = "wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16&sample_rate=16000"
def _results(
start: object,
duration: object,
transcript: str = "",
is_final: object = True,
channel_index: object = (0, 1),
) -> dict[str, object]:
return {
"type": "Results",
"start": start,
"duration": duration,
"is_final": is_final,
"channel_index": list(channel_index) if isinstance(channel_index, tuple) else channel_index,
"channel": {"alternatives": [{"transcript": transcript, "confidence": 0.9}]},
}
def _metadata(duration: object, channels: object = 1) -> dict[str, object]:
return {"type": "Metadata", "request_id": "req-1", "duration": duration, "channels": channels}
@pytest.mark.parametrize(
("api_base", "query_string", "expected"),
[
pytest.param(
None,
"model=nova-3&encoding=linear16",
"wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16",
id="default",
),
pytest.param(
None,
"encoding=linear16&sample_rate=16000",
"wss://api.deepgram.com/v1/listen?encoding=linear16&sample_rate=16000&model=nova-3",
id="model added when missing",
),
pytest.param(
None,
"model=&encoding=linear16",
"wss://api.deepgram.com/v1/listen?encoding=linear16&model=nova-3",
id="empty model replaced",
),
pytest.param(
"http://localhost:9000/v1/",
"model=nova-2",
"ws://localhost:9000/v1/listen?model=nova-2",
id="custom base becomes ws",
),
pytest.param(
"wss://dg.internal/v1",
"model=nova-3&keywords=a&keywords=b",
"wss://dg.internal/v1/listen?model=nova-3&keywords=a&keywords=b",
id="repeated keys preserved",
),
pytest.param(
None,
"model=nova-2&encoding=linear16&model=nova-3",
"wss://api.deepgram.com/v1/listen?model=nova-2&encoding=linear16",
id="only the authorized first model reaches deepgram",
),
pytest.param(
None,
"language=en&model=nova-3&language=multi",
"wss://api.deepgram.com/v1/listen?language=en&model=nova-3",
id="only the priced first language reaches deepgram",
),
pytest.param(
None,
"model=&model=nova-2",
"wss://api.deepgram.com/v1/listen?model=nova-3",
id="blank first model is the default, later models dropped",
),
],
)
def test_deepgram_listen_websocket_target(api_base: str | None, query_string: str, expected: str):
assert deepgram_listen_websocket_target(api_base=api_base, query_string=query_string) == expected
@pytest.mark.parametrize(
("query_string", "expected"),
[
pytest.param("model=nova-3&encoding=linear16", (), id="no callback"),
pytest.param("model=nova-3&callback=https%3A%2F%2Fevil.example%2Fsink", ("callback",), id="callback"),
pytest.param(
"callback_method=put&model=nova-3&callback=wss%3A%2F%2Fevil.example",
("callback", "callback_method"),
id="callback and method",
),
pytest.param("model=nova-3&callback_method=put", ("callback_method",), id="method alone"),
pytest.param("model=nova-3&callbacks=x&my_callback=y", (), id="only exact names match"),
],
)
def test_deepgram_listen_callback_params(query_string: str, expected: tuple[str, ...]):
assert deepgram_listen_callback_params(query_string) == expected
@pytest.mark.parametrize(
("frames", "expected_seconds"),
[
pytest.param((_results(0.0, 2.0), _results(2.0, 3.5), _metadata(6.25)), 6.25, id="metadata wins"),
pytest.param((_metadata(4.0), _results(0.0, 9.0), _metadata(5.5)), 5.5, id="last metadata wins"),
pytest.param((_results(0.0, 2.0), _results(2.0, 3.5), _results(1.0, 1.0)), 5.5, id="furthest results end"),
pytest.param((_results(0.0, 0.0), _metadata(0.0)), 0.0, id="zero metadata is a real zero"),
pytest.param(
(_metadata(0.0), _results(0.0, 2.0), _results(2.0, 3.5)),
5.5,
id="handshake metadata zero does not hide streamed results",
),
pytest.param((_metadata(0.0), _results(0.0, 2.0), _metadata(0.0)), 2.0, id="only zero metadata frames"),
pytest.param((_results(0.0, 1.5), _metadata("6.25")), 1.5, id="string metadata is ignored"),
pytest.param((_results(0.0, 1.5), _metadata(True)), 1.5, id="boolean metadata is ignored"),
pytest.param((_results(0.0, 1.5), _metadata(-3.0)), 1.5, id="negative metadata is ignored"),
pytest.param((_results(0.0, 1.5), _metadata(math.nan), _metadata(math.inf)), 1.5, id="nan/inf ignored"),
pytest.param((_results("0", 2.0), _results(0.0, None), _results(0.0, 0.75)), 0.75, id="malformed results"),
pytest.param(({"type": "SpeechStarted", "timestamp": 3.0}, {"type": "UtteranceEnd"}), 0.0, id="no usage"),
pytest.param((), 0.0, id="no frames"),
],
)
def test_deepgram_listen_audio_seconds(frames: Sequence[Mapping[str, object]], expected_seconds: float):
assert deepgram_listen_audio_seconds(frames) == expected_seconds
@pytest.mark.parametrize(
("frames", "upstream_url", "expected_channels"),
[
pytest.param((_results(0.0, 2.0), _metadata(6.25)), NOVA_3_URL, 1, id="mono"),
pytest.param((_results(0.0, 2.0, channel_index=(0, 2)), _metadata(6.25, 2)), NOVA_3_URL, 2, id="stereo"),
pytest.param((_metadata(1.0, 3), _metadata(1.0, 5)), NOVA_3_URL, 5, id="last metadata wins"),
pytest.param(
(_metadata(1.0, 20), _results(0.0, 1.0, channel_index=(1, 2))),
NOVA_3_URL,
20,
id="metadata beats channel_index",
),
pytest.param(
(_results(0.0, 1.0, channel_index=(0, 2)), _results(0.0, 1.0, channel_index=(3, 4))),
NOVA_3_URL,
4,
id="widest channel_index without metadata",
),
pytest.param(
(_results(0.0, 1.0, channel_index=(0, 2)),),
f"{NOVA_3_URL}&channels=7&multichannel=true",
2,
id="frames beat the declared query",
),
pytest.param((), f"{NOVA_3_URL}&channels=7&multichannel=true", 7, id="declared query when no frames"),
pytest.param((), f"{NOVA_3_URL}&channels=0", 1, id="zero declared channels"),
pytest.param((), f"{NOVA_3_URL}&channels=-2", 1, id="negative declared channels"),
pytest.param((), f"{NOVA_3_URL}&channels=two", 1, id="non numeric declared channels"),
pytest.param((), NOVA_3_URL, 1, id="nothing declared"),
pytest.param((_metadata(1.0, "2"), _metadata(1.0, True), _metadata(1.0, 0)), NOVA_3_URL, 1, id="bad metadata"),
pytest.param((_metadata(1.0, 3), _metadata(1.0, True)), NOVA_3_URL, 3, id="boolean does not shadow a count"),
pytest.param((_metadata(1.0, 2.0), _metadata(1.0, -1)), NOVA_3_URL, 1, id="float and negative metadata"),
pytest.param(
(_metadata(1.0, 2), {**_results(0.0, 1.0), "channels": 9}, {"type": "UtteranceEnd", "channels": 11}),
NOVA_3_URL,
2,
id="channels on non metadata frames ignored",
),
pytest.param(
(
_results(0.0, 1.0, channel_index=[0]),
_results(0.0, 1.0, channel_index=(0, "2")),
_results(0.0, 1.0, channel_index=(0, 0)),
),
NOVA_3_URL,
1,
id="bad channel_index",
),
],
)
def test_deepgram_listen_channel_count(
frames: Sequence[Mapping[str, object]], upstream_url: str, expected_channels: int
):
assert deepgram_listen_channel_count(frames, upstream_url) == expected_channels
def test_deepgram_listen_transcript_joins_final_results_only():
frames = (
_results(0.0, 1.0, "hello wor", is_final=False),
_results(0.0, 1.5, "hello world"),
_results(1.5, 0.5, "", is_final=True),
_results(2.0, 1.0, "how are you", is_final="yes"),
{"type": "Results", "start": 3.0, "duration": 1.0, "is_final": True, "channel": {"alternatives": []}},
_results(4.0, 1.0, "goodbye"),
_metadata(5.0),
)
assert deepgram_listen_transcript(frames) == "hello world goodbye"
@pytest.mark.parametrize(
("upstream_url", "expected_model"),
[
(NOVA_3_URL, "nova-3"),
("wss://api.deepgram.com/v1/listen?encoding=linear16&model=nova-2-medical", "nova-2-medical"),
("wss://api.deepgram.com/v1/listen?model=nova-3&model=nova-2", "nova-3"),
("wss://api.deepgram.com/v1/listen?encoding=linear16", litellm.constants.DEEPGRAM_LISTEN_DEFAULT_MODEL),
],
)
def test_deepgram_listen_model_comes_from_the_upstream_query(upstream_url: str, expected_model: str):
assert deepgram_listen_model(upstream_url) == expected_model
@pytest.mark.parametrize(
"query_string",
[
"model=nova-2&language=en",
"language=en",
"model=&language=en",
"",
"model=nova-3-medical",
"model=nova-2&model=nova-3",
"model=&model=nova-3-medical",
],
)
def test_requested_model_is_the_only_model_the_upstream_target_carries(query_string: str):
"""Authorization runs against ``deepgram_listen_requested_model``; the upstream URL is built separately, so the
two must always agree or a key could be authorized for one model and reach another. Deepgram reads the last
repeated ``model``, so the target must carry exactly one."""
target: Final = deepgram_listen_websocket_target(None, query_string)
assert parse_qs(urlparse(target).query)["model"] == [deepgram_listen_requested_model(query_string)]
assert deepgram_listen_requested_model(query_string) == deepgram_listen_model(target)
@pytest.mark.parametrize(
("upstream_url", "expected"),
[
pytest.param(NOVA_3_URL, "streaming/nova-3", id="monolingual"),
pytest.param(f"{NOVA_3_URL}&language=en", "streaming/nova-3", id="explicit language"),
pytest.param(f"{NOVA_3_URL}&language=multi", "streaming/nova-3-multilingual", id="multilingual"),
pytest.param(f"{NOVA_3_URL}&language=MULTI", "streaming/nova-3-multilingual", id="multilingual any case"),
pytest.param(
"wss://api.deepgram.com/v1/listen?model=nova-2&language=multi",
"streaming/nova-2-multilingual",
id="other model",
),
pytest.param("wss://api.deepgram.com/v1/listen?encoding=linear16", "streaming/nova-3", id="default model"),
],
)
def test_deepgram_listen_pricing_model_is_the_streaming_entry_never_the_prerecorded_one(
upstream_url: str, expected: str
):
assert deepgram_listen_pricing_model(upstream_url) == expected
assert deepgram_listen_registry_key(upstream_url) == f"deepgram/{expected}"
NOVA_2_URL: Final = "wss://api.deepgram.com/v1/listen?model=nova-2"
@pytest.mark.usefixtures("local_model_cost_map")
@pytest.mark.parametrize(
("upstream_url", "extra_rows", "expected"),
[
pytest.param(NOVA_3_URL, (), True, id="streaming entry present"),
pytest.param(f"{NOVA_3_URL}&language=multi", (), True, id="multilingual entry present"),
pytest.param(NOVA_2_URL, (), False, id="only the pre-recorded entry"),
pytest.param(f"{NOVA_2_URL}&language=multi", ("deepgram/streaming/nova-2",), False, id="needs multilingual"),
pytest.param("wss://api.deepgram.com/v1/listen?model=nova-99-unmapped", (), False, id="nothing priced"),
pytest.param(NOVA_2_URL, ("deepgram/streaming/nova-2",), True, id="operator-supplied streaming entry"),
pytest.param(NOVA_2_URL, ("streaming/nova-2",), False, id="a row under another key is not the entry"),
],
)
def test_deepgram_listen_is_priced(
monkeypatch: pytest.MonkeyPatch, upstream_url: str, extra_rows: tuple[str, ...], expected: bool
):
"""The bundled map prices only nova-3 for streaming; nova-2 has a pre-recorded row, which must never count."""
monkeypatch.delitem(litellm.model_cost, "deepgram/streaming/nova-2", raising=False)
assert "deepgram/nova-2" in litellm.model_cost
for row in extra_rows:
monkeypatch.setitem(litellm.model_cost, row, dict(litellm.model_cost["deepgram/streaming/nova-3"]))
assert deepgram_listen_is_priced(upstream_url) is expected
@pytest.mark.parametrize(
("upstream_url", "expected"),
[
pytest.param(NOVA_3_URL, (), id="no add-ons"),
pytest.param(f"{NOVA_3_URL}&redact=pci", ("streaming/redact",), id="redact"),
pytest.param(f"{NOVA_3_URL}&redact=pci&redact=ssn", ("streaming/redact",), id="repeated redact once"),
pytest.param(f"{NOVA_3_URL}&keyterm=a&keyterm=b", ("streaming/keyterm",), id="keyterm"),
pytest.param(f"{NOVA_3_URL}&detect_entities=true", ("streaming/detect_entities",), id="detect_entities"),
pytest.param(f"{NOVA_3_URL}&diarize=true", ("streaming/diarize",), id="diarize"),
pytest.param(f"{NOVA_3_URL}&diarize_model=v1", ("streaming/diarize",), id="diarize_model"),
pytest.param(f"{NOVA_3_URL}&diarize=true&diarize_model=latest", ("streaming/diarize",), id="diarize both once"),
pytest.param(f"{NOVA_3_URL}&detect_entities=false&diarize=FALSE&redact=", (), id="disabled"),
pytest.param(
f"{NOVA_3_URL}&detect_entities=false&detect_entities=true",
("streaming/detect_entities",),
id="any enabling value wins",
),
pytest.param(
f"{NOVA_3_URL}&diarize=true&redact=pci&keyterm=x&detect_entities=true",
("streaming/detect_entities", "streaming/diarize", "streaming/keyterm", "streaming/redact"),
id="all, sorted",
),
],
)
def test_deepgram_listen_addon_pricing_models(upstream_url: str, expected: tuple[str, ...]):
assert deepgram_listen_addon_pricing_models(upstream_url) == expected

View file

@ -51,14 +51,19 @@ class TestMistralReasoningSupport:
assert "reasoning_effort" in supported_params
assert "thinking" in supported_params
# Test non-magistral model doesn't include reasoning parameters
supported_params_reasoning = mistral_config.get_supported_openai_params(
"mistral/mistral-medium-latest"
)
assert "reasoning_effort" in supported_params_reasoning
assert "thinking" not in supported_params_reasoning
supported_params_normal = mistral_config.get_supported_openai_params(
"mistral/mistral-large-latest"
)
assert "reasoning_effort" not in supported_params_normal
assert "thinking" not in supported_params_normal
def test_map_openai_params_reasoning_effort(self):
def test_map_openai_params_reasoning_effort(self, local_model_cost_map):
"""Test that reasoning_effort parameter is properly mapped for magistral models."""
mistral_config = MistralConfig()
@ -73,16 +78,93 @@ class TestMistralReasoningSupport:
assert result.get("_add_reasoning_prompt") is True
# Test reasoning_effort ignored for non-magistral model
optional_params_normal = {}
result_normal = mistral_config.map_openai_params(
non_default_params={"reasoning_effort": "low"},
optional_params=optional_params_normal,
model="mistral/mistral-large-latest",
model="mistral/mistral-medium-latest",
drop_params=False,
)
assert "_add_reasoning_prompt" not in result_normal
assert result_normal["reasoning_effort"] == "high"
@pytest.mark.parametrize(
("model", "requested", "sent"),
[
("mistral-medium-latest", "high", "high"),
("mistral-medium-latest", "none", "none"),
("mistral-medium-latest", "low", "high"),
("mistral-medium-latest", "medium", "high"),
("mistral-medium-latest", "xhigh", "high"),
("mistral-small-latest", "medium", "high"),
("mistral-vibe-cli-latest", "medium", "high"),
("zai-glm-5", "none", "none"),
("zai-glm-5", "minimal", "low"),
("zai-glm-5", "medium", "high"),
("zai-glm-5", "xhigh", "max"),
("zai-glm-5-2", "medium", "medium"),
("zai-glm-5-2", "xhigh", "xhigh"),
],
)
def test_reasoning_effort_is_sent_as_a_level_the_model_accepts(self, local_model_cost_map, model, requested, sent):
import litellm
optional_params = litellm.get_optional_params(
model=model,
custom_llm_provider="mistral",
reasoning_effort=requested,
)
assert optional_params["reasoning_effort"] == sent
def test_reasoning_effort_is_forwarded_verbatim_when_the_map_declares_no_levels(
self, local_model_cost_map, monkeypatch
):
import litellm
monkeypatch.setitem(
litellm.model_cost,
"mistral/undeclared-reasoner",
{"litellm_provider": "mistral", "mode": "chat", "supports_reasoning": True},
)
optional_params = litellm.get_optional_params(
model="undeclared-reasoner",
custom_llm_provider="mistral",
reasoning_effort="medium",
)
assert optional_params["reasoning_effort"] == "medium"
def test_reasoning_effort_stays_unsupported_for_non_reasoning_models(self):
import litellm
with pytest.raises(litellm.UnsupportedParamsError):
litellm.get_optional_params(
model="codestral-latest",
custom_llm_provider="mistral",
reasoning_effort="high",
)
dropped = litellm.get_optional_params(
model="codestral-latest",
custom_llm_provider="mistral",
reasoning_effort="high",
drop_params=True,
)
assert "reasoning_effort" not in dropped
def test_client_metadata_stripped_from_request(self):
mistral_config = MistralConfig()
request = mistral_config.transform_request(
model="mistral-medium-latest",
messages=[{"role": "user", "content": "hi"}],
optional_params={"client_metadata": {"originator": "codex_cli_rs"}, "temperature": 0.2},
litellm_params={},
headers={},
)
assert "client_metadata" not in request
assert request["temperature"] == 0.2
def test_map_openai_params_thinking(self):
"""Test that thinking parameter is properly mapped for magistral models."""

View file

@ -0,0 +1,17 @@
import litellm
def test_reasoning_effort_stays_unsupported_on_vertex_partner_models(local_model_cost_map):
assert "reasoning_effort" in litellm.get_supported_openai_params(
model="mistral-medium-3", custom_llm_provider="mistral"
)
assert "reasoning_effort" not in litellm.get_supported_openai_params(
model="mistral-medium-3", custom_llm_provider="vertex_ai"
)
dropped = litellm.get_optional_params(
model="mistral-medium-3",
custom_llm_provider="vertex_ai",
reasoning_effort="high",
drop_params=True,
)
assert "reasoning_effort" not in dropped

View file

@ -9029,3 +9029,31 @@ async def test_router_settings_model_group_alias_authorizes_target_for_team(monk
await authorize()
assert (await request.json())["model"] == target
assert get_client_requested_model(request) == "AgentX-LLM"
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ["configured-voice", None])
async def test_websocket_auth_explicit_model_overrides_query(monkeypatch, model):
import importlib
from unittest.mock import AsyncMock
from fastapi import WebSocket
from litellm.proxy import proxy_server
auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setattr(proxy_server, "general_settings", {})
seen = []
async def authenticate(request, api_key):
seen.append((await request.json(), api_key))
return "authenticated"
monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate)
websocket = WebSocket({
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
"path": "/v1/realtime", "path_params": {},
"query_string": b"model=untrusted-query",
"headers": [(b"x-litellm-api-key", b"owner")],
}, AsyncMock(), AsyncMock())
assert await auth_module.user_api_key_auth_websocket_for_model(websocket, model) == "authenticated"
assert seen == [({"model": model or ""}, "Bearer owner")]

View file

@ -1,3 +1,4 @@
import pathlib
import re
from collections.abc import Sequence
from datetime import datetime, timedelta, timezone
@ -10,10 +11,11 @@ import pytest
from psycopg.rows import dict_row
from pytest_postgresql import factories
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT
from litellm.constants import (
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM,
PTU_SENTINEL_API_KEY,
USAGE_TOP_API_KEYS_LIMIT,
)
from litellm.proxy.management_endpoints.common_daily_activity import (
_adjust_dates_for_timezone,
_build_aggregated_sql_query,
@ -23,8 +25,12 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
get_api_key_metadata,
get_daily_activity,
get_daily_activity_aggregated,
global_rollup_reconciled_through,
update_metrics,
)
from litellm.proxy.spend_tracking.daily_global_spend_rollup import RECONCILE_DAY_SQL
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.proxy.utils import evict_config_param
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendMetadata,
SpendMetrics,
@ -1578,6 +1584,172 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both
assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == {"key-1"}
def _prisma_with_marker(marker: str | None) -> MagicMock:
prisma = MagicMock()
prisma.db = MagicMock()
prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
row = (
None if marker is None else SimpleNamespace(param_name="m", param_value=f'{{"reconciled_through": "{marker}"}}')
)
prisma.get_generic_data = AsyncMock(return_value=row)
return prisma
def _unfiltered_user_query(**overrides):
return {
"table_name": "litellm_dailyuserspend",
"entity_id_field": "user_id",
"entity_id": None,
"start_date": "2026-06-01",
"end_date": "2026-06-02",
"model": None,
"api_key": None,
"exclude_entity_ids": None,
"timezone_offset_minutes": None,
"include_current_utc_day": False,
**overrides,
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
("marker", "overrides", "expected"),
[
("2026-06-02", {}, "2026-06-02"),
("2026-06-02", {"model": "gpt-5"}, "2026-06-02"),
("2026-05-01", {}, "2026-05-01"),
(None, {}, None),
("2026-06-02", {"api_key": "sk-1"}, None),
("2026-06-02", {"api_key": []}, None),
("2026-06-02", {"entity_id": "u-1"}, None),
("2026-06-02", {"exclude_entity_ids": ["u-1"]}, None),
("2026-06-02", {"table_name": "litellm_dailyteamspend", "entity_id_field": "team_id"}, None),
],
)
async def test_global_rollup_marker_is_used_only_for_unfiltered_user_reads(marker, overrides, expected):
"""Anything that filters by key or entity has no counterpart in the global table; the
SQL splits the range at the marker itself, so the marker passes through unchanged."""
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
prisma = _prisma_with_marker(marker)
assert await global_rollup_reconciled_through(prisma, _unfiltered_user_query(**overrides)) == expected
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
@pytest.mark.asyncio
async def test_global_rollup_marker_read_failure_falls_back_to_the_per_key_table():
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
prisma = _prisma_with_marker(None)
prisma.get_generic_data = AsyncMock(side_effect=RuntimeError("db down"))
assert await global_rollup_reconciled_through(prisma, _unfiltered_user_query()) is None
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
_GLOBAL_SPEND_MIGRATION: Final = (
pathlib.Path(__file__).resolve().parents[4]
/ "litellm-proxy-extras"
/ "litellm_proxy_extras"
/ "migrations"
/ "20260915000000_add_daily_global_spend"
/ "migration.sql"
)
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_serves_closed_days_from_the_global_table_and_open_days_live(
_aggregated_postgresql: psycopg.Connection,
):
"""Day 1 is rolled up and day 2 is still open (never rolled up), so a marker of day 1 must
give the same response as reading everything per-key: day 1 from the global table, day 2
live, one grand total across both. Per-key rows that land after the rollup then tell the
two sources apart: a late day 1 row is invisible to totals until the next reconcile while a
late day 2 row shows up at once, and both keys rank in the key breakdown, which stays
per-key throughout."""
n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 3
rows: Final = [
(
f"row-{day}-{i:03d}",
f"user-{i % 7}",
day,
f"key-{i:03d}",
"gpt-5" if i % 2 else "claude",
"" if i % 3 else "gpt-5",
"openai" if i % 2 else None,
"/v1/chat/completions" if i % 5 else None,
10,
float(i + 1),
1,
1,
)
for day in ("2026-06-01", "2026-06-02")
for i in range(n_keys)
]
_seed_daily_user_spend(_aggregated_postgresql, rows)
with _aggregated_postgresql.cursor() as cur:
cur.execute(
'UPDATE "LiteLLM_DailyUserSpend" SET total_response_time_ms = prompt_tokens * 25, '
"timed_requests = api_requests"
)
cur.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal
cur.execute(
re.sub(r"\$(\d+)", r"%(p\1)s", RECONCILE_DAY_SQL), # pyright: ignore[reportArgumentType] # $N -> psycopg
{"p1": "2026-06-01"},
)
_aggregated_postgresql.commit()
async def read(marker: str | None):
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
prisma = _prisma_with_marker(marker)
prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, [])
return await get_daily_activity_aggregated(
prisma_client=prisma,
entity_metadata_field=None,
**_unfiltered_user_query(),
)
from_per_key = await read(None)
from_global = await read("2026-06-01")
assert from_global.model_dump() == from_per_key.model_dump()
seeded_spend: Final = 2 * sum(float(i + 1) for i in range(n_keys))
assert from_global.metadata.total_spend == pytest.approx(seeded_spend)
assert from_global.metadata.total_response_time_ms == 2 * n_keys * 10 * 25
assert from_global.metadata.total_timed_requests == 2 * n_keys
assert {day.date.isoformat() for day in from_global.results} == {"2026-06-01", "2026-06-02"}
assert len(from_global.results[0].breakdown.api_keys) == USAGE_TOP_API_KEYS_LIMIT
assert set(from_global.results[0].breakdown.model_groups) == {"gpt-5", "claude"}
with _aggregated_postgresql.cursor() as cur:
cur.executemany(
"""
INSERT INTO "LiteLLM_DailyUserSpend"
(id, user_id, date, api_key, model, model_group, custom_llm_provider,
endpoint, prompt_tokens, spend, api_requests, successful_requests)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
""",
[
("late-1", "user-late", "2026-06-01", "key-late-1", "gpt-5", "", "openai", None, 10, 1000.0, 1, 1),
("late-2", "user-late", "2026-06-02", "key-late-2", "gpt-5", "", "openai", None, 10, 500.0, 1, 1),
],
)
_aggregated_postgresql.commit()
late_per_key = await read(None)
late_global = await read("2026-06-01")
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
assert late_per_key.metadata.total_spend == pytest.approx(seeded_spend + 1000.0 + 500.0)
assert late_global.metadata.total_spend == pytest.approx(seeded_spend + 500.0)
by_day: Final = {day.date.isoformat(): day for day in late_global.results}
assert by_day["2026-06-01"].metrics.spend == pytest.approx(seeded_spend / 2)
assert by_day["2026-06-02"].metrics.spend == pytest.approx(seeded_spend / 2 + 500.0)
assert by_day["2026-06-01"].breakdown.api_keys["key-late-1"].metrics.spend == pytest.approx(1000.0)
assert by_day["2026-06-02"].breakdown.api_keys["key-late-2"].metrics.spend == pytest.approx(500.0)
assert late_global.metadata.total_api_keys == n_keys + 2
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_reports_exact_limit_key_count_as_complete(
_aggregated_postgresql: psycopg.Connection,
@ -1634,7 +1806,20 @@ async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_mo
"""Rows stored with an empty or NULL model_group must land in the model_groups
breakdown under their model name instead of vanishing from the usage UI."""
rows: Final = [
("row-0", "user-0", "2026-06-01", "key-0", "gpt-5", "gpt-5-eu", "openai", "/v1/chat/completions", 10, 7.0, 1, 1),
(
"row-0",
"user-0",
"2026-06-01",
"key-0",
"gpt-5",
"gpt-5-eu",
"openai",
"/v1/chat/completions",
10,
7.0,
1,
1,
),
("row-1", "user-1", "2026-06-01", "key-1", "gpt-5", "", "openai", "/v1/chat/completions", 10, 3.0, 1, 1),
("row-2", "user-2", "2026-06-01", "key-2", "claude-x", None, "anthropic", "/v1/messages", 10, 2.0, 1, 1),
]

View file

@ -21,6 +21,11 @@ SHORT_AUDIO_URL = (
)
BATCH_URL = "https://eastus.api.cognitive.microsoft.com/speechtotext/v3.2/transcriptions"
FAST_URL = "https://eastus.api.cognitive.microsoft.com/speechtotext/transcriptions:transcribe?api-version=2024-11-15"
PREFIXED_SHORT_AUDIO_URL = (
"https://apim.example.com/speech-proxy/speech/recognition/conversation/cognitiveservices/v1?language=en-US"
)
PREFIXED_BATCH_URL = "https://apim.example.com/speech/speechtotext/v3.2/transcriptions"
PREFIXED_FAST_URL = "https://apim.example.com/speech/speechtotext/transcriptions:transcribe?api-version=2024-11-15"
FAST_BODY = {"durationMilliseconds": 5061, "combinedPhrases": [{"text": "Hello world."}]}
FAST_AUDIO_SECONDS = 5.061
TRANSCRIPT_BODY = {
@ -87,6 +92,9 @@ class TestAzureSpeechPassthroughHandler:
(FAST_URL, "azure_speech/fast-transcription", FAST_AUDIO_SECONDS * PRICE_PER_SECOND),
(BATCH_URL, "azure_speech/batch-transcription", 0.0),
(f"{BATCH_URL}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files", "azure_speech/batch-transcription", 0.0),
(PREFIXED_SHORT_AUDIO_URL, "azure_speech/short-audio", TRANSCRIPT_AUDIO_SECONDS * PRICE_PER_SECOND),
(PREFIXED_FAST_URL, "azure_speech/fast-transcription", FAST_AUDIO_SECONDS * PRICE_PER_SECOND),
(PREFIXED_BATCH_URL, "azure_speech/batch-transcription", 0.0),
],
)
def test_records_model_provider_and_cost(self, url_route: str, expected_model: str, expected_cost: float):

View file

@ -0,0 +1,302 @@
"""Deepgram ``/v1/listen`` WebSocket passthrough: duration extraction and duration based cost tracking."""
from datetime import datetime
from types import SimpleNamespace
from typing import Final
import pytest
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.deepgram_listen_passthrough_logging_handler import (
DeepgramListenPassthroughLoggingHandler,
)
from litellm.proxy.pass_through_endpoints.success_handler import PassThroughEndpointLogging
from litellm.types.passthrough_endpoints.pass_through_endpoints import PassthroughStandardLoggingPayload
from litellm.types.utils import StandardLoggingPayload, TranscriptionResponse
NOVA_3_URL: Final = "wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16&sample_rate=16000"
pytestmark: Final = pytest.mark.usefixtures("local_model_cost_map")
def _results(start: object, duration: object, transcript: str = "", is_final: object = True) -> dict[str, object]:
return {
"type": "Results",
"start": start,
"duration": duration,
"is_final": is_final,
"channel": {"alternatives": [{"transcript": transcript, "confidence": 0.9}]},
}
def _metadata(duration: object, channels: int = 1) -> dict[str, object]:
return {"type": "Metadata", "request_id": "req-1", "duration": duration, "channels": channels}
@pytest.mark.parametrize(
("url_route", "expected"),
[
("/deepgram/v1/listen", True),
("/deepgram/listen", True),
("/deepgram/v1/listen?model=nova-3", True),
("/litellm/deepgram/v1/listen", True),
("/deepgram/v1/speak", False),
("/deepgram/v1/listen/extra", False),
("/openai/v1/realtime", False),
("/vertex_ai/live", False),
("", False),
],
)
def test_is_deepgram_listen_route(url_route: str, expected: bool):
assert DeepgramListenPassthroughLoggingHandler.is_deepgram_listen_route(url_route) is expected
def _logging_obj(call_id: str = "call-dg") -> LiteLLMLoggingObj:
return LiteLLMLoggingObj(
model="unknown",
messages=[{"role": "user", "content": "WebSocket connection"}],
stream=True,
call_type="pass_through_endpoint",
start_time=datetime.now(),
litellm_call_id=call_id,
function_id="websocket_passthrough",
)
def _registry_cost(pricing_model: str, seconds: float) -> float:
"""Derives the expected charge from the live cost map rather than pinning a vendor price."""
per_second: Final = litellm.model_cost[f"deepgram/{pricing_model}"]["input_cost_per_second"]
assert per_second > 0
return per_second * seconds
def _cost(upstream_url: str, *frames: dict[str, object]) -> float:
handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=frames, logging_obj=_logging_obj(), upstream_url=upstream_url
)
response_cost = handler_result["kwargs"]["response_cost"]
assert isinstance(response_cost, float)
return response_cost
def test_handler_bills_metadata_duration_at_the_registry_rate_and_names_the_model():
frames = (_results(0.0, 5.0, "first sentence"), _results(5.0, 7.5, "second sentence"), _metadata(12.5))
logging_obj = _logging_obj()
handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=frames,
logging_obj=logging_obj,
upstream_url=NOVA_3_URL,
kwargs={"litellm_params": {"metadata": {}}},
)
result = handler_result["result"]
assert isinstance(result, TranscriptionResponse)
assert result.text == "first sentence second sentence"
assert result._hidden_params["audio_transcription_duration"] == 12.5
assert result._hidden_params["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 12.5))
assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 12.5))
assert handler_result["kwargs"]["model"] == "nova-3"
assert handler_result["kwargs"]["custom_llm_provider"] == "deepgram"
assert handler_result["kwargs"]["litellm_params"] == {"metadata": {}}
assert logging_obj.model == "nova-3"
assert logging_obj.model_call_details["model"] == "nova-3"
assert logging_obj.model_call_details["custom_llm_provider"] == "deepgram"
assert logging_obj.model_call_details["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 12.5))
def test_handler_bills_streaming_not_prerecorded_rates():
"""Deepgram prices /v1/listen over a WebSocket separately from pre-recorded transcription, so the streaming entry
must be the one charged; the two registry rows only need to differ for this to matter, whatever their values."""
streaming = litellm.model_cost["deepgram/streaming/nova-3"]["input_cost_per_second"]
prerecorded = litellm.model_cost["deepgram/nova-3"]["input_cost_per_second"]
assert streaming != prerecorded
assert _cost(NOVA_3_URL, _metadata(60.0)) == pytest.approx(60.0 * streaming)
def test_handler_bills_multilingual_streaming_when_language_is_multi():
monolingual = _cost(NOVA_3_URL, _metadata(60.0))
multilingual = _cost(f"{NOVA_3_URL}&language=multi", _metadata(60.0))
assert multilingual == pytest.approx(_registry_cost("streaming/nova-3-multilingual", 60.0))
assert multilingual > monolingual
@pytest.mark.parametrize(
("query", "addons"),
[
pytest.param("redact=pci", ("redact",), id="redaction"),
pytest.param("redact=pci&redact=numbers", ("redact",), id="redaction counted once"),
pytest.param("keyterm=LiteLLM&keyterm=Deepgram", ("keyterm",), id="keyterm prompting"),
pytest.param("detect_entities=true", ("detect_entities",), id="entity detection"),
pytest.param("diarize=true", ("diarize",), id="diarization"),
pytest.param("diarize_model=v1", ("diarize",), id="diarization via diarize_model"),
pytest.param("diarize=true&diarize_model=v1", ("diarize",), id="diarization counted once"),
pytest.param(
"redact=pci&keyterm=x&detect_entities=true&diarize=true",
("redact", "keyterm", "detect_entities", "diarize"),
id="every add-on",
),
pytest.param("detect_entities=false&diarize=False&redact=", (), id="disabled add-ons cost nothing"),
],
)
def test_handler_adds_each_priced_add_on_once_on_top_of_the_base_rate(query: str, addons: tuple[str, ...]):
base = _cost(NOVA_3_URL, _metadata(60.0))
expected = base + sum(_registry_cost(f"streaming/{addon}", 60.0) for addon in addons)
assert _cost(f"{NOVA_3_URL}&{query}", _metadata(60.0)) == pytest.approx(expected)
def test_handler_add_ons_scale_with_channels_like_the_base_rate():
stereo_plain = _cost(f"{NOVA_3_URL}&channels=2", _metadata(60.0, channels=2))
stereo_redacted = _cost(f"{NOVA_3_URL}&channels=2&redact=pci", _metadata(60.0, channels=2))
assert stereo_redacted - stereo_plain == pytest.approx(_registry_cost("streaming/redact", 120.0))
@pytest.mark.parametrize(
"upstream_url",
[
pytest.param("wss://api.deepgram.com/v1/listen?model=nova-2", id="only a pre-recorded entry"),
pytest.param("wss://api.deepgram.com/v1/listen?model=nova-99-not-in-registry", id="no entry at all"),
],
)
def test_handler_never_substitutes_another_rate_for_a_missing_streaming_entry(monkeypatch, upstream_url):
"""The route refuses these sessions up front; should the registry change under a live one, the spend row
keeps the duration and carries no cost, rather than the pre-recorded rate or any other stand-in."""
monkeypatch.delitem(litellm.model_cost, "deepgram/streaming/nova-2", raising=False)
assert "deepgram/nova-2" in litellm.model_cost
handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=(_metadata(60.0),), logging_obj=_logging_obj(), upstream_url=upstream_url
)
assert handler_result["kwargs"]["response_cost"] is None
assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 60.0
def test_handler_falls_back_to_results_frames_when_the_stream_ends_without_metadata():
frames = (_results(0.0, 30.0, "a"), _results(30.0, 30.0, "b"), _results(60.0, 12.5, "c"))
handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=frames, logging_obj=_logging_obj(), upstream_url=NOVA_3_URL
)
assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 72.5
assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 72.5))
def test_handler_charges_more_for_more_audio_on_the_same_model():
short = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=(_metadata(10.0),), logging_obj=_logging_obj(), upstream_url=NOVA_3_URL
)
long = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=(_metadata(30.0),), logging_obj=_logging_obj(), upstream_url=NOVA_3_URL
)
assert long["kwargs"]["response_cost"] == pytest.approx(3 * short["kwargs"]["response_cost"])
assert short["kwargs"]["response_cost"] > 0
def test_handler_bills_every_channel_of_a_multichannel_session():
"""Deepgram bills processed audio per channel (deepgram.com/pricing FAQ, 2026-09-17), so a stereo session must be
charged for twice its wall-clock duration or budgets can be bypassed by requesting more channels."""
mono = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=(_metadata(30.0),), logging_obj=_logging_obj(), upstream_url=NOVA_3_URL
)
stereo = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=(_metadata(30.0, channels=2),),
logging_obj=_logging_obj(),
upstream_url=f"{NOVA_3_URL}&multichannel=true&channels=2",
)
assert stereo["result"]._hidden_params["audio_transcription_duration"] == 60.0
assert stereo["kwargs"]["response_cost"] == pytest.approx(2 * mono["kwargs"]["response_cost"])
assert stereo["kwargs"]["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 60.0))
def test_handler_bills_the_declared_channels_when_the_stream_dies_before_any_frame_reports_them():
handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler(
websocket_messages=(_results(0.0, 10.0, "a"),),
logging_obj=_logging_obj(),
upstream_url=f"{NOVA_3_URL}&multichannel=true&channels=3",
)
assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 30.0
assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 30.0))
class _CapturingLogger(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.payloads: list[StandardLoggingPayload] = []
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
self.payloads.append(kwargs["standard_logging_object"])
@pytest.mark.asyncio
async def test_success_handler_dispatches_deepgram_listen_and_logs_duration_based_spend(monkeypatch):
"""Drives the shared passthrough success handler the way the WebSocket relay does at socket close and reads
what a spend logger receives: Deepgram model and provider, the audio duration billed at the registry rate."""
capturing_logger = _CapturingLogger()
monkeypatch.setattr(litellm, "_async_success_callback", [capturing_logger])
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "callbacks", [])
logging_obj = _logging_obj("call-dg-e2e")
frames = [_results(0.0, 5.0, "hello world", is_final=False), _results(0.0, 5.0, "hello world"), _metadata(20.0)]
user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key", team_id="team-stt", user_id="user-1")
start_time = datetime.now()
passthrough_logging_payload = PassthroughStandardLoggingPayload(
url=NOVA_3_URL, request_body={}, request_method="WEBSOCKET", cost_per_request=None
)
logging_obj.update_environment_variables(
model="unknown",
user="unknown",
optional_params={},
litellm_params={
"metadata": {
"user_api_key": user_api_key_dict.api_key,
"user_api_key_team_id": user_api_key_dict.team_id,
"user_api_key_user_id": user_api_key_dict.user_id,
}
},
call_type="pass_through_endpoint",
)
await PassThroughEndpointLogging().pass_through_async_success_handler(
httpx_response=SimpleNamespace(
status_code=200,
text="WebSocket connection successful",
headers={},
request=SimpleNamespace(method="WEBSOCKET", url=NOVA_3_URL),
),
response_body=frames,
logging_obj=logging_obj,
url_route="/deepgram/v1/listen",
result="websocket_connection_successful",
start_time=start_time,
end_time=datetime.now(),
cache_hit=False,
request_body={},
passthrough_logging_payload=passthrough_logging_payload,
litellm_params={
"metadata": {
"user_api_key": user_api_key_dict.api_key,
"user_api_key_team_id": user_api_key_dict.team_id,
"user_api_key_user_id": user_api_key_dict.user_id,
}
},
)
assert len(capturing_logger.payloads) == 1
payload = capturing_logger.payloads[0]
assert payload["model"] == "nova-3"
assert payload["custom_llm_provider"] == "deepgram"
assert payload["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 20.0))
assert payload["metadata"]["user_api_key_team_id"] == "team-stt"
assert payload["id"] == "call-dg-e2e"

View file

@ -0,0 +1,474 @@
"""Deepgram ``/v1/listen`` passthrough WebSocket route: registration, auth, credential injection, target URL."""
import asyncio
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType, SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from starlette.routing import WebSocketRoute
from starlette.websockets import WebSocketDisconnect
import litellm
from litellm.caching.dual_cache import DualCache
from litellm.proxy._lazy_features import LAZY_FEATURES
from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import _cache_key_object
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_websocket_relay,
deepgram_listen_websocket_route,
router,
)
from litellm.proxy.utils import hash_token
GET_CREDENTIALS: Final = (
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials"
)
USER_API_KEY_AUTH: Final = "litellm.proxy.auth.user_api_key_auth.user_api_key_auth"
LISTEN_PATHS: Final = ("/deepgram/v1/listen", "/deepgram/listen")
NOVA_2_STREAMING_KEY: Final = "deepgram/streaming/nova-2"
pytestmark: Final = pytest.mark.usefixtures("local_model_cost_map")
def _price_nova_2_streaming(monkeypatch: pytest.MonkeyPatch) -> None:
"""An operator-supplied streaming row: the bundled map prices only nova-3 for streaming."""
monkeypatch.setitem(litellm.model_cost, NOVA_2_STREAMING_KEY, dict(litellm.model_cost["deepgram/streaming/nova-3"]))
class _FakeWebSocket:
def __init__(self, path: str, query: str) -> None:
self.url = SimpleNamespace(path=path, query=query)
self.headers = {"authorization": "Bearer sk-litellm-virtual", "x-api-key": "sk-caller-secret"}
self.accepts: list[str | None] = []
self.closed: tuple[int, str] | None = None
async def accept(self, subprotocol: str | None = None) -> None:
self.accepts.append(subprotocol)
async def close(self, code: int = 1000, reason: str = "") -> None:
self.closed = (code, reason)
@dataclass(frozen=True, slots=True)
class _RelayCall:
target: str
custom_headers: Mapping[str, str]
user_api_key_dict: UserAPIKeyAuth
forward_headers: bool
endpoint: str
accept_websocket: bool
class _FakeRelay:
def __init__(self) -> None:
self.calls: list[_RelayCall] = []
async def __call__(
self,
*,
websocket: object,
target: str,
custom_headers: dict[str, str],
user_api_key_dict: UserAPIKeyAuth,
forward_headers: bool,
endpoint: str,
accept_websocket: bool,
) -> None:
self.calls.append(
_RelayCall(
target=target,
custom_headers=MappingProxyType(dict(custom_headers)),
user_api_key_dict=user_api_key_dict,
forward_headers=forward_headers,
endpoint=endpoint,
accept_websocket=accept_websocket,
)
)
async def _serve(websocket: _FakeWebSocket, user_api_key_dict: UserAPIKeyAuth | None = None) -> _FakeRelay:
relay = _FakeRelay()
await deepgram_listen_websocket_route(
websocket=websocket,
user_api_key_dict=user_api_key_dict or UserAPIKeyAuth(),
relay=relay,
)
return relay
def test_deepgram_listen_websocket_routes_registered():
ws_paths = {route.path for route in router.routes if isinstance(route, WebSocketRoute)}
assert set(LISTEN_PATHS) <= ws_paths
@pytest.mark.parametrize("path", LISTEN_PATHS)
def test_deepgram_listen_is_a_lazily_loaded_mapped_pass_through_route(path):
"""The route must be reachable before the passthrough module is imported and must be authed and
billed as a mapped pass-through route like the other provider prefixes."""
feature = next(feature for feature in LAZY_FEATURES if feature.name == "llm_passthrough")
assert feature.matches(path)
assert any(path.startswith(prefix) for prefix in LiteLLMRoutes.mapped_pass_through_routes.value)
@pytest.mark.asyncio
@pytest.mark.parametrize("path", LISTEN_PATHS)
async def test_deepgram_listen_forwards_query_and_injects_only_provider_auth(path, monkeypatch):
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
websocket = _FakeWebSocket(path, "encoding=linear16&sample_rate=16000&keywords=hi%3A2&keywords=there")
caller = UserAPIKeyAuth(api_key="sk-litellm-virtual", team_id="team-stt")
with patch(GET_CREDENTIALS, return_value="dg-provider-key") as get_credentials:
relay = await _serve(websocket, caller)
assert get_credentials.call_args.kwargs == {"custom_llm_provider": "deepgram", "region_name": None}
assert relay.calls == [
_RelayCall(
target=(
"wss://api.deepgram.com/v1/listen"
"?encoding=linear16&sample_rate=16000&keywords=hi%3A2&keywords=there&model=nova-3"
),
custom_headers=MappingProxyType({"Authorization": "Token dg-provider-key"}),
user_api_key_dict=caller,
forward_headers=False,
endpoint=path,
accept_websocket=False,
)
]
assert websocket.accepts == [None]
assert websocket.closed is None
@pytest.mark.asyncio
async def test_deepgram_listen_keeps_caller_chosen_model(monkeypatch):
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
_price_nova_2_streaming(monkeypatch)
websocket = _FakeWebSocket("/deepgram/v1/listen", "model=nova-2&language=en")
with patch(GET_CREDENTIALS, return_value="dg-provider-key"):
relay = await _serve(websocket)
assert [call.target for call in relay.calls] == ["wss://api.deepgram.com/v1/listen?model=nova-2&language=en"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("query", "expected_target"),
[
("", "wss://api.deepgram.com/v1/listen?model=nova-3"),
("model=", "wss://api.deepgram.com/v1/listen?model=nova-3"),
("model=&language=en", "wss://api.deepgram.com/v1/listen?language=en&model=nova-3"),
],
)
async def test_deepgram_listen_defaults_to_nova_3_when_no_model_is_named(query, expected_target, monkeypatch):
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
websocket = _FakeWebSocket("/deepgram/listen", query)
with patch(GET_CREDENTIALS, return_value="dg-provider-key"):
relay = await _serve(websocket)
assert [call.target for call in relay.calls] == [expected_target]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("api_base", "expected_target"),
[
("https://api.eu.deepgram.com/v1/", "wss://api.eu.deepgram.com/v1/listen?model=nova-3"),
("http://localhost:8080/v1", "ws://localhost:8080/v1/listen?model=nova-3"),
("wss://deepgram.internal.example/v1", "wss://deepgram.internal.example/v1/listen?model=nova-3"),
],
)
async def test_deepgram_listen_honours_server_configured_api_base(api_base, expected_target, monkeypatch):
monkeypatch.setenv("DEEPGRAM_API_BASE", api_base)
websocket = _FakeWebSocket("/deepgram/v1/listen", "")
with patch(GET_CREDENTIALS, return_value="dg-provider-key"):
relay = await _serve(websocket)
assert [call.target for call in relay.calls] == [expected_target]
@pytest.mark.asyncio
async def test_deepgram_listen_ignores_caller_supplied_api_base(monkeypatch):
"""V1: the server-configured Deepgram key must only ever go to the server-configured host."""
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
websocket = _FakeWebSocket("/deepgram/v1/listen", "api_base=wss%3A%2F%2Fattacker.example%2Fv1&model=nova-3")
with patch(GET_CREDENTIALS, return_value="dg-provider-key"):
relay = await _serve(websocket)
assert [call.target for call in relay.calls] == [
"wss://api.deepgram.com/v1/listen?api_base=wss%3A%2F%2Fattacker.example%2Fv1&model=nova-3"
]
@pytest.mark.asyncio
async def test_deepgram_listen_closes_cleanly_when_provider_credentials_missing():
websocket = _FakeWebSocket("/deepgram/v1/listen", "model=nova-3")
with patch(GET_CREDENTIALS, return_value=None):
relay = await _serve(websocket)
assert websocket.closed is not None
assert websocket.closed[0] == 1011
assert "DEEPGRAM_API_KEY" in websocket.closed[1]
assert websocket.accepts == []
assert relay.calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize(
"query",
[
pytest.param("model=nova-3&callback=https%3A%2F%2Fsink.example%2Fdg", id="http callback"),
pytest.param("callback=wss%3A%2F%2Fsink.example&callback_method=put&model=nova-3", id="ws callback"),
],
)
async def test_deepgram_listen_rejects_callback_delivery_that_would_go_unbilled(query, monkeypatch):
"""With ``callback`` set, Deepgram sends every Results and Metadata frame to the caller's URL and only a
request id down this socket, so the proxy would meter zero seconds of audio; refuse before contacting Deepgram."""
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
websocket = _FakeWebSocket("/deepgram/v1/listen", query)
with patch(GET_CREDENTIALS, return_value="dg-provider-key"):
relay = await _serve(websocket)
assert relay.calls == []
assert websocket.closed is not None
assert websocket.closed[0] == 1008
assert "callback" in websocket.closed[1]
assert "dg-provider-key" not in websocket.closed[1]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("query", "missing_key"),
[
pytest.param("model=nova-2", "deepgram/streaming/nova-2", id="model with only a pre-recorded price"),
pytest.param("model=nova-99", "deepgram/streaming/nova-99", id="model unknown to the registry"),
pytest.param(
"model=nova-3&language=multi",
"deepgram/streaming/nova-3-multilingual",
id="multilingual session without its own price",
),
],
)
async def test_deepgram_listen_refuses_sessions_it_cannot_price(query, missing_key, monkeypatch):
"""A session with no streaming price would be logged at zero (or at the pre-recorded rate), letting a caller run
up unmetered spend, so the proxy closes it before Deepgram is contacted and names the registry row to add."""
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
monkeypatch.delitem(litellm.model_cost, missing_key, raising=False)
assert "deepgram/nova-2" in litellm.model_cost
websocket = _FakeWebSocket("/deepgram/v1/listen", query)
with patch(GET_CREDENTIALS, return_value="dg-provider-key"):
relay = await _serve(websocket)
assert relay.calls == []
assert websocket.closed is not None
assert websocket.closed[0] == 1008
assert missing_key in websocket.closed[1]
assert "dg-provider-key" not in websocket.closed[1]
@pytest.mark.asyncio
async def test_deepgram_listen_relays_once_the_operator_prices_the_model(monkeypatch):
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
websocket = _FakeWebSocket("/deepgram/v1/listen", "model=nova-2")
with patch(GET_CREDENTIALS, return_value="dg-provider-key"):
assert (await _serve(websocket)).calls == []
_price_nova_2_streaming(monkeypatch)
priced_websocket = _FakeWebSocket("/deepgram/v1/listen", "model=nova-2")
with patch(GET_CREDENTIALS, return_value="dg-provider-key"):
relay = await _serve(priced_websocket)
assert [call.target for call in relay.calls] == ["wss://api.deepgram.com/v1/listen?model=nova-2"]
assert priced_websocket.closed is None
def _app_with_relay(relay: _FakeRelay) -> FastAPI:
app = FastAPI()
app.include_router(router)
app.dependency_overrides[_websocket_relay] = lambda: relay
return app
def test_deepgram_listen_rejects_connections_without_a_litellm_key():
relay = _FakeRelay()
client = TestClient(_app_with_relay(relay))
with patch(GET_CREDENTIALS, return_value="dg-provider-key") as get_credentials:
with pytest.raises(WebSocketDisconnect) as disconnect:
with client.websocket_connect("/deepgram/v1/listen?model=nova-3"):
pass
assert disconnect.value.code == 1008
assert relay.calls == []
get_credentials.assert_not_called()
def test_deepgram_listen_callback_rejection_reaches_the_client_as_a_policy_close(monkeypatch):
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
relay = _FakeRelay()
client = TestClient(_app_with_relay(relay))
with (
patch(GET_CREDENTIALS, return_value="dg-provider-key"),
patch(USER_API_KEY_AUTH, new=AsyncMock(return_value=UserAPIKeyAuth(api_key="hashed"))),
):
with pytest.raises(WebSocketDisconnect) as disconnect:
with client.websocket_connect(
"/deepgram/v1/listen?model=nova-3&callback=https%3A%2F%2Fsink.example%2Fdg",
headers={"Authorization": "Bearer sk-litellm-virtual"},
) as connection:
connection.receive_text()
assert disconnect.value.code == 1008
assert "callback" in disconnect.value.reason
assert relay.calls == []
def test_deepgram_listen_authenticates_the_litellm_key_and_relays_to_deepgram(monkeypatch):
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
relay = _FakeRelay()
client = TestClient(_app_with_relay(relay))
caller = UserAPIKeyAuth(api_key="hashed-sk-litellm", team_id="team-stt")
with (
patch(GET_CREDENTIALS, return_value="dg-provider-key"),
patch(USER_API_KEY_AUTH, new=AsyncMock(return_value=caller)) as auth,
):
with client.websocket_connect(
"/deepgram/v1/listen?model=nova-3&punctuate=true",
headers={"Authorization": "Bearer sk-litellm-virtual"},
):
pass
assert auth.await_args.kwargs["api_key"] == "Bearer sk-litellm-virtual"
assert relay.calls == [
_RelayCall(
target="wss://api.deepgram.com/v1/listen?model=nova-3&punctuate=true",
custom_headers=MappingProxyType({"Authorization": "Token dg-provider-key"}),
user_api_key_dict=caller,
forward_headers=False,
endpoint="/deepgram/v1/listen",
accept_websocket=False,
)
]
async def _cache_restricted_key(virtual_key: str, models: list[str]) -> DualCache:
cache = DualCache()
await _cache_key_object(
hashed_token=hash_token(virtual_key),
user_api_key_obj=UserAPIKeyAuth(token=hash_token(virtual_key), models=models),
user_api_key_cache=cache,
proxy_logging_obj=None,
)
return cache
@pytest.mark.parametrize(
("query", "expect_relay"),
[
pytest.param("model=nova-2", True, id="allowed model named"),
pytest.param("model=nova-3", False, id="denied model named"),
pytest.param("", False, id="model omitted, default denied"),
pytest.param("model=&language=en", False, id="model blank, default denied"),
],
)
def test_deepgram_listen_authorizes_the_model_it_will_actually_send_upstream(query, expect_relay, monkeypatch):
"""A key allowed only ``nova-2`` must not reach ``nova-3`` by leaving ``model`` out and letting the proxy fill
in its default: the real key auth path must see the same model the upstream target will carry."""
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
monkeypatch.setattr(litellm, "max_budget", 0.0)
_price_nova_2_streaming(monkeypatch)
cache = asyncio.run(_cache_restricted_key("sk-only-nova-2", ["nova-2"]))
relay = _FakeRelay()
client = TestClient(_app_with_relay(relay))
with (
patch(GET_CREDENTIALS, return_value="dg-provider-key"),
patch.multiple( # test-quality-ok: the real key auth path reads these proxy_server globals and has no injection seam
"litellm.proxy.proxy_server",
master_key="sk-master",
prisma_client=MagicMock(),
user_api_key_cache=cache,
llm_model_list=None,
llm_router=None,
),
):
if expect_relay:
with client.websocket_connect(
f"/deepgram/v1/listen?{query}", headers={"Authorization": "Bearer sk-only-nova-2"}
):
pass
assert [call.target for call in relay.calls] == [f"wss://api.deepgram.com/v1/listen?{query}"]
return
with pytest.raises(WebSocketDisconnect) as disconnect:
with client.websocket_connect(
f"/deepgram/v1/listen?{query}", headers={"Authorization": "Bearer sk-only-nova-2"}
):
pass
assert disconnect.value.code == 1008
assert relay.calls == []
def test_deepgram_listen_strips_a_second_model_that_would_outrank_the_authorized_one(monkeypatch):
"""Deepgram honours the last repeated ``model``; auth and pricing read the first. A key allowed only ``nova-2``
must not smuggle ``nova-3`` past authorization behind an authorized first value."""
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
monkeypatch.setattr(litellm, "max_budget", 0.0)
_price_nova_2_streaming(monkeypatch)
cache = asyncio.run(_cache_restricted_key("sk-only-nova-2", ["nova-2"]))
relay = _FakeRelay()
client = TestClient(_app_with_relay(relay))
with (
patch(GET_CREDENTIALS, return_value="dg-provider-key"),
patch.multiple( # test-quality-ok: the real key auth path reads these proxy_server globals and has no injection seam
"litellm.proxy.proxy_server",
master_key="sk-master",
prisma_client=MagicMock(),
user_api_key_cache=cache,
llm_model_list=None,
llm_router=None,
),
):
with client.websocket_connect(
"/deepgram/v1/listen?model=nova-2&language=en&model=nova-3&language=multi",
headers={"Authorization": "Bearer sk-only-nova-2"},
):
pass
assert [call.target for call in relay.calls] == ["wss://api.deepgram.com/v1/listen?model=nova-2&language=en"]
def test_deepgram_listen_echoes_the_browser_subprotocol_that_carries_the_litellm_key(monkeypatch):
"""Browsers cannot set headers, so they send the key as a subprotocol and abort the handshake unless the
server echoes that subprotocol back; the key itself must still stay off the upstream connection."""
monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False)
relay = _FakeRelay()
client = TestClient(_app_with_relay(relay))
with (
patch(GET_CREDENTIALS, return_value="dg-provider-key"),
patch(USER_API_KEY_AUTH, new=AsyncMock(return_value=UserAPIKeyAuth(api_key="hashed"))),
):
with client.websocket_connect(
"/deepgram/v1/listen?model=nova-3",
subprotocols=["openai-insecure-api-key.sk-litellm-virtual"],
) as connection:
assert connection.accepted_subprotocol == "openai-insecure-api-key.sk-litellm-virtual"
assert [call.custom_headers for call in relay.calls] == [
MappingProxyType({"Authorization": "Token dg-provider-key"})
]
assert [call.forward_headers for call in relay.calls] == [False]

View file

@ -4942,18 +4942,21 @@ async def test_unusable_upstream_cost_records_zero_not_the_flat_estimate():
class FakeUpstreamWebSocket:
def __init__(self, first_frame: bytes):
self._first_frame = first_frame
"""Serves the given frames in order, then closes normally, the way a real websockets connection does"""
def __init__(self, *frames: str | bytes):
self._frames = iter(frames)
self.close = AsyncMock()
self.send = AsyncMock()
async def recv(self, decode: bool = True):
return self._first_frame
async def recv(self, decode: bool | None = None):
from websockets.exceptions import ConnectionClosedOK
from websockets.frames import Close
def __aiter__(self):
return self
async def __anext__(self):
raise StopAsyncIteration
frame = next(self._frames, None)
if frame is None:
raise ConnectionClosedOK(rcvd=Close(1000, ""), sent=Close(1000, ""), rcvd_then_sent=True)
return frame
class FakeUpstreamConnect:
@ -4974,7 +4977,7 @@ async def test_websocket_passthrough_forwards_non_ascii_first_frame():
first_frame = json.dumps(
{"type": "session.created", "session": {"instructions": "Hablas español, ¿sí?"}},
ensure_ascii=False,
).encode("utf-8")
)
upstream_ws = FakeUpstreamWebSocket(first_frame)
websocket = MagicMock()
@ -5028,7 +5031,7 @@ async def test_websocket_passthrough_propagates_active_trace_context(
from starlette.websockets import WebSocketState
captured: dict[str, dict[str, str]] = {}
upstream_ws = FakeUpstreamWebSocket(b"{}")
upstream_ws = FakeUpstreamWebSocket("{}")
def fake_connect(target, additional_headers):
captured["headers"] = additional_headers
@ -5457,6 +5460,144 @@ async def test_websocket_passthrough_does_not_close_twice_when_success_logging_f
websocket.close.assert_awaited_once_with(code=1008, reason=upstream_reason)
DEEPGRAM_LISTEN_TARGET = "wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16&sample_rate=16000"
DEEPGRAM_INTERIM_FRAME = json.dumps(
{
"type": "Results",
"start": 0.0,
"duration": 1.02,
"is_final": False,
"channel": {"alternatives": [{"transcript": "hello wor", "confidence": 0.71}]},
}
)
DEEPGRAM_FINAL_FRAME = json.dumps(
{
"type": "Results",
"start": 0.0,
"duration": 2.5,
"is_final": True,
"speech_final": True,
"channel": {"alternatives": [{"transcript": "hello world, ¿qué tal?", "confidence": 0.98}]},
},
ensure_ascii=False,
)
DEEPGRAM_METADATA_FRAME = json.dumps({"type": "Metadata", "request_id": "req-1", "duration": 2.5, "channels": 1})
async def _relay_deepgram_listen(upstream_ws, client_receive):
"""Runs the generic relay the way the Deepgram route does and returns (client websocket, success handler mock)"""
websocket = _client_websocket(client_receive)
with (
_patched_websocket_passthrough_environment(upstream_ws),
patch( # test-quality-ok: pass_through_endpoint_logging is a module global read inside websocket_passthrough_request; there is no injection seam
"litellm.proxy.pass_through_endpoints.pass_through_endpoints."
"pass_through_endpoint_logging.pass_through_async_success_handler",
new=AsyncMock(),
) as success_handler,
):
await websocket_passthrough_request(
websocket=websocket,
target=DEEPGRAM_LISTEN_TARGET,
custom_headers={"Authorization": "Token dg-provider-key"},
user_api_key_dict=UserAPIKeyAuth(),
forward_headers=False,
endpoint="/deepgram/v1/listen",
accept_websocket=False,
)
return websocket, success_handler
@pytest.mark.asyncio
async def test_websocket_passthrough_relays_deepgram_transcript_frames_verbatim_and_keeps_them_for_billing():
"""Interim, final and Metadata frames reach the client byte for byte (no JSON round trip, non-ASCII intact,
a binary frame first) and every JSON object frame is what the success handler gets to bill from."""
upstream_ws = FakeUpstreamWebSocket(
b"\x00\x01binary-first",
DEEPGRAM_INTERIM_FRAME,
"not json at all",
DEEPGRAM_FINAL_FRAME,
DEEPGRAM_METADATA_FRAME,
)
websocket, success_handler = await _relay_deepgram_listen(upstream_ws, _pending_receive)
assert [call.args[0] for call in websocket.send_bytes.await_args_list] == [b"\x00\x01binary-first"]
assert [call.args[0] for call in websocket.send_text.await_args_list] == [
DEEPGRAM_INTERIM_FRAME,
"not json at all",
DEEPGRAM_FINAL_FRAME,
DEEPGRAM_METADATA_FRAME,
]
success_call = success_handler.call_args.kwargs
assert success_call["url_route"] == "/deepgram/v1/listen"
assert success_call["response_body"] == [
json.loads(DEEPGRAM_INTERIM_FRAME),
json.loads(DEEPGRAM_FINAL_FRAME),
json.loads(DEEPGRAM_METADATA_FRAME),
]
assert success_call["httpx_response"].request.url == DEEPGRAM_LISTEN_TARGET
assert success_call["logging_obj"].model_call_details.get("custom_llm_provider") is None
websocket.close.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_websocket_passthrough_sends_deepgram_audio_bytes_and_control_text_upstream_unchanged():
upstream_ws = RecordingUpstreamWebSocket()
audio_chunk = bytes(range(256)) * 4
close_stream = json.dumps({"type": "CloseStream"})
await _relay_deepgram_listen(
upstream_ws,
AsyncMock(
side_effect=[
{"type": "websocket.receive", "bytes": audio_chunk},
{"type": "websocket.receive", "text": close_stream},
{"type": "websocket.disconnect"},
]
),
)
assert [call.args[0] for call in upstream_ws.send.await_args_list] == [audio_chunk, close_stream]
assert isinstance(upstream_ws.send.await_args_list[0].args[0], bytes)
upstream_ws.close.assert_awaited_once()
@pytest.mark.asyncio
async def test_websocket_passthrough_vertex_live_setup_ack_names_the_model_but_is_not_billed_as_usage():
"""Vertex Live keeps its special first frame: the setup acknowledgement is forwarded verbatim, read for the
model, and left out of the frames the usage handler sees; later frames are kept as before."""
setup_ack = json.dumps(
{"setupComplete": {}, "model": "projects/p/locations/global/publishers/google/models/gemini-live-2.5-flash"}
)
server_content = json.dumps({"serverContent": {"turnComplete": True}, "usageMetadata": {"totalTokenCount": 12}})
upstream_ws = FakeUpstreamWebSocket(setup_ack, server_content)
websocket = _client_websocket(_pending_receive)
with (
_patched_websocket_passthrough_environment(upstream_ws),
patch( # test-quality-ok: pass_through_endpoint_logging is a module global read inside websocket_passthrough_request; there is no injection seam
"litellm.proxy.pass_through_endpoints.pass_through_endpoints."
"pass_through_endpoint_logging.pass_through_async_success_handler",
new=AsyncMock(),
) as success_handler,
):
await websocket_passthrough_request(
websocket=websocket,
target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent",
custom_headers={"Authorization": "Bearer token"},
user_api_key_dict=UserAPIKeyAuth(),
forward_headers=False,
endpoint="/vertex_ai/live",
accept_websocket=False,
)
assert [call.args[0] for call in websocket.send_text.await_args_list] == [setup_ack, server_content]
success_call = success_handler.call_args.kwargs
assert success_call["response_body"] == [json.loads(server_content)]
assert success_call["logging_obj"].model == "gemini-live-2.5-flash"
assert success_call["logging_obj"].model_call_details["custom_llm_provider"] == "vertex_ai_language_models"
def _passthrough_kwargs_for_reservation(
user_api_key_dict: UserAPIKeyAuth,
parsed_body: dict | None = None,

View file

@ -28,6 +28,7 @@ from typing import List, Optional, Union
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from fastapi import FastAPI
from pydantic import BaseModel
from typing_extensions import TypedDict
@ -1042,6 +1043,57 @@ async def test_spend_report_locks_are_never_released():
proxy_logging_obj.db_spend_update_writer.pod_lock_manager.release_lock.assert_not_awaited()
def _init_daily_global_spend_reconcile_job() -> tuple[AsyncIOScheduler, MagicMock, MagicMock]:
scheduler = AsyncIOScheduler()
proxy_logging_obj = MagicMock()
proxy_logging_obj.alerting_handler = AsyncMock()
prisma_client = MagicMock()
ProxyStartupEvent._initialize_daily_global_spend_reconcile_job(
scheduler=scheduler,
proxy_logging_obj=proxy_logging_obj,
prisma_client=prisma_client,
)
return scheduler, proxy_logging_obj, prisma_client
def test_daily_global_spend_reconcile_job_is_scheduled_nightly_with_an_immediate_catch_up_run():
"""Startup schedules the LiteLLM_DailyGlobalSpend backfill a couple of minutes out, so a
fresh deploy switches usage reads to the global table without waiting for the nightly
run, and after that it fires once a day at 00:30 UTC, when the previous UTC day is closed."""
from datetime import datetime, timedelta, timezone
from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID
scheduler, _, _ = _init_daily_global_spend_reconcile_job()
job = scheduler.get_job(DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID)
assert job is not None
assert timedelta(0) < job.next_run_time - datetime.now(timezone.utc) <= timedelta(minutes=2)
after_catch_up = datetime(2026, 9, 16, 12, 0, tzinfo=timezone.utc)
assert job.trigger.get_next_fire_time(None, after_catch_up) == datetime(2026, 9, 17, 0, 30, tzinfo=timezone.utc)
just_after_a_run = datetime(2026, 9, 17, 0, 30, 1, tzinfo=timezone.utc)
assert job.trigger.get_next_fire_time(None, just_after_a_run) == datetime(2026, 9, 18, 0, 30, tzinfo=timezone.utc)
@pytest.mark.asyncio
async def test_daily_global_spend_reconcile_job_runs_under_the_pod_lock_and_alerts_through_the_proxy(monkeypatch):
from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID
scheduler, proxy_logging_obj, prisma_client = _init_daily_global_spend_reconcile_job()
run = AsyncMock()
monkeypatch.setattr(ps, "run_scheduled_daily_global_spend_reconcile", run)
await scheduler.get_job(DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID).func()
run.assert_awaited_once()
assert run.await_args.args == (prisma_client,)
assert run.await_args.kwargs["pod_lock_manager"] is proxy_logging_obj.db_spend_update_writer.pod_lock_manager
await run.await_args.kwargs["alert"]("day 2026-09-01 failed")
proxy_logging_obj.alerting_handler.assert_awaited_once()
assert proxy_logging_obj.alerting_handler.await_args.kwargs["message"] == "day 2026-09-01 failed"
assert proxy_logging_obj.alerting_handler.await_args.kwargs["level"] == "High"
@pytest.mark.asyncio
async def test_prometheus_fallback_stats_job_skipped_when_another_pod_holds_the_lock(monkeypatch):
"""The boot-time send goes through the same gate, so a losing pod sends nothing at all:

View file

@ -0,0 +1,532 @@
"""Tests for the LiteLLM_DailyGlobalSpend reconcile job (LIT-7818)."""
import json
import pathlib
import re
from datetime import date
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import psycopg
import pytest
from psycopg.rows import dict_row
from pytest_postgresql import factories
from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM
from litellm.proxy.db.daily_spend_bulk_upsert import DAILY_SPEND_TABLES, build_bulk_upsert, merge_by_conflict_key
from litellm.proxy.spend_tracking.daily_global_spend_rollup import (
_ADVANCE_MARKER_SQL,
RECONCILE_DAY_SQL,
read_marker,
reconciled_through,
run_daily_global_spend_reconcile,
run_scheduled_daily_global_spend_reconcile,
)
from litellm.proxy.utils import evict_config_param
USER_TABLE: Final = DAILY_SPEND_TABLES["user"]
TODAY: Final = date(2026, 9, 15)
class _FakeConfigRow:
def __init__(self, param_name: str, param_value: object) -> None:
self.param_name = param_name
self.param_value = param_value
class _FakeConfigTable:
def __init__(self) -> None:
self.rows: dict[str, object] = {}
def advance(self, param_name: str, through: str | None, scanned_at: str | None) -> None:
"""What ``_ADVANCE_MARKER_SQL`` does in Postgres: keep the later of stored and incoming per field."""
stored = self.rows.get(param_name)
current: dict[str, str | None] = json.loads(stored) if isinstance(stored, str) else {}
self.rows[param_name] = json.dumps(
{
"reconciled_through": _greatest(current.get("reconciled_through"), through),
"scanned_at": _greatest(current.get("scanned_at"), scanned_at),
}
)
def _greatest(stored: str | None, incoming: str | None) -> str | None:
present = [value for value in (stored, incoming) if value is not None]
return max(present) if present else None
class _FakeDb:
"""Per-key rows are ``{date: updated_at}`` with a fake database clock that ticks per query,
so "rows written since the last scan" behaves like Postgres would. The database's own
date decides which day is still open, never the pod's clock."""
def __init__(self, prisma: "_FakePrisma") -> None:
self._prisma = prisma
self.litellm_config = _FakeConfigTable()
async def query_raw(self, sql: str, *params: str) -> list[dict[str, str]]:
if sql.startswith("SELECT (NOW()"):
self._prisma.clock += 1
return [{"now": f"clock-{self._prisma.clock:04d}", "today": self._prisma.today.isoformat()}]
rows = self._prisma.user_rows
if len(params) == 1:
(last,) = params
return [{"date": d} for d in sorted(rows) if d <= last]
last, marker, scanned_at = params
return [
{"date": d} for d, written in sorted(rows.items()) if d <= last and (d > marker or written >= scanned_at)
]
async def execute_raw(self, sql: str, *params: str | None) -> int:
if sql == _ADVANCE_MARKER_SQL:
param_name, through, scanned_at = params
assert param_name is not None
self.litellm_config.advance(param_name, through, scanned_at)
return 1
(day,) = params
if day is None or day in self._prisma.failing_days:
raise RuntimeError(f"day {day} exploded")
self._prisma.reconciled.append(day)
landing = self._prisma.marker_landing_on_day.get(day)
if landing is not None:
self.litellm_config.rows[DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM] = landing
return 1
class _FakePrisma:
"""Enough of PrismaClient for the reconcile: per-key dates, a config table, and execute_raw.
``marker_landing_on_day`` stores another pod's marker the moment this run rewrites that day."""
def __init__(
self, user_days: tuple[str, ...], failing_days: frozenset[str] = frozenset(), today: date = TODAY
) -> None:
self.clock = 0
self.today = today
self.user_rows: dict[str, str] = {d: "clock-0000" for d in user_days}
self.failing_days = failing_days
self.marker_landing_on_day: dict[str, str] = {}
self.reconciled: list[str] = []
self.db = _FakeDb(self)
def write_late_row(self, day: str) -> None:
"""A per-key row for ``day`` lands now, after whatever scans already happened."""
self.clock += 1
self.user_rows[day] = f"clock-{self.clock:04d}"
async def get_generic_data(self, key: str, value: str, table_name: str) -> _FakeConfigRow | None:
stored = self.db.litellm_config.rows.get(value)
return None if stored is None else _FakeConfigRow(value, stored)
@pytest.fixture(autouse=True)
async def _fresh_marker_cache():
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
yield
await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM)
@pytest.mark.asyncio
async def test_first_run_rolls_up_every_closed_day_and_never_the_database_s_today():
"""Before any marker exists every closed day with per-key rows is rolled up. Today is left
out: pods are still flushing it, so it is served live from the per-key table until it closes.
The database clock says which day that is; a pod booting with its clock a day ahead must not
roll the open day up and mark it reconciled."""
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-03", "2026-09-14", "2026-09-15"))
result = await run_daily_global_spend_reconcile(prisma)
assert result.days_reconciled == ("2026-09-01", "2026-09-03", "2026-09-14")
assert result.failed_day is None
assert result.reconciled_through == "2026-09-14"
assert await reconciled_through(prisma) == "2026-09-14"
assert "2026-09-15" not in prisma.reconciled
@pytest.mark.asyncio
async def test_later_run_rolls_up_only_new_days_when_nothing_old_changed():
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-12", "2026-09-13", "2026-09-14"), today=date(2026, 9, 14))
await run_daily_global_spend_reconcile(prisma)
prisma.reconciled.clear()
prisma.today = TODAY
result = await run_daily_global_spend_reconcile(prisma)
assert result.days_reconciled == ("2026-09-14",)
assert await reconciled_through(prisma) == "2026-09-14"
@pytest.mark.asyncio
async def test_spend_landing_on_an_old_rolled_up_day_is_folded_in_by_the_next_run():
"""Per-key rows carry the request start date, so a delayed flush or retry can add spend to a
day far behind the marker. That day is rewritten, and the marker never moves back for it."""
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-05", "2026-09-13"), today=date(2026, 9, 14))
await run_daily_global_spend_reconcile(prisma)
prisma.reconciled.clear()
prisma.today = TODAY
prisma.write_late_row("2026-09-01")
prisma.write_late_row("2026-09-03")
result = await run_daily_global_spend_reconcile(prisma)
assert result.days_reconciled == ("2026-09-01", "2026-09-03")
assert "2026-09-05" not in prisma.reconciled
assert await reconciled_through(prisma) == "2026-09-13"
@pytest.mark.asyncio
async def test_a_late_row_seen_by_a_failed_run_is_seen_again_by_the_next_one():
"""The scan time only advances when every pending day was rewritten, otherwise a late row
found by the failed run would be counted as handled."""
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-13"), today=date(2026, 9, 14))
await run_daily_global_spend_reconcile(prisma)
prisma.today = TODAY
prisma.write_late_row("2026-09-01")
prisma.failing_days = frozenset({"2026-09-01"})
failed = await run_daily_global_spend_reconcile(prisma)
prisma.failing_days = frozenset()
prisma.reconciled.clear()
result = await run_daily_global_spend_reconcile(prisma)
assert failed.failed_day == "2026-09-01"
assert failed.reconciled_through == "2026-09-13"
assert result.days_reconciled == ("2026-09-01",)
assert result.failed_day is None
@pytest.mark.asyncio
async def test_a_marker_without_a_scan_time_rolls_every_closed_day_up_again():
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-13"))
prisma.db.litellm_config.rows[DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM] = '{"reconciled_through": "2026-09-13"}'
result = await run_daily_global_spend_reconcile(prisma)
assert result.days_reconciled == ("2026-09-01", "2026-09-13")
marker = await read_marker(prisma)
assert marker is not None and marker.reconciled_through == "2026-09-13" and marker.scanned_at is not None
@pytest.mark.asyncio
async def test_a_run_with_no_new_closed_days_keeps_the_marker():
prisma = _FakePrisma(user_days=("2026-09-13",), today=date(2026, 9, 14))
await run_daily_global_spend_reconcile(prisma)
prisma.reconciled.clear()
result = await run_daily_global_spend_reconcile(prisma)
assert result.days_reconciled == ()
assert result.reconciled_through == "2026-09-13"
@pytest.mark.asyncio
async def test_a_failing_day_stops_the_run_and_leaves_the_marker_on_the_last_good_day():
"""The marker may never claim a day that was not rewritten: reads past it would then trust
a global table missing that day's spend."""
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-02", "2026-09-03"), failing_days=frozenset({"2026-09-02"}))
result = await run_daily_global_spend_reconcile(prisma)
assert result.days_reconciled == ("2026-09-01",)
assert result.failed_day == "2026-09-02"
assert result.reconciled_through == "2026-09-01"
assert prisma.reconciled == ["2026-09-01"]
assert await reconciled_through(prisma) == "2026-09-01"
@pytest.mark.asyncio
async def test_the_next_run_resumes_from_the_failed_day():
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-02", "2026-09-03"), failing_days=frozenset({"2026-09-02"}))
await run_daily_global_spend_reconcile(prisma)
prisma.failing_days = frozenset()
result = await run_daily_global_spend_reconcile(prisma)
assert result.days_reconciled == ("2026-09-01", "2026-09-02", "2026-09-03")
assert await reconciled_through(prisma) == "2026-09-03"
@pytest.mark.asyncio
async def test_a_slower_overlapping_run_never_rewinds_the_marker_a_faster_run_stored():
"""Two pods can reconcile at once (Redis unreachable, or the lock expired on a long backfill).
When the faster one has already stored a later marker, the slower one may only add to it. Putting
its own older prefix back, or dropping the scan time, would send usage reads for every day in
between back to the per-key table until the next run."""
prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-02", "2026-09-03"), failing_days=frozenset({"2026-09-03"}))
prisma.marker_landing_on_day = {
"2026-09-02": '{"reconciled_through": "2026-09-14", "scanned_at": "clock-0009"}',
}
result = await run_daily_global_spend_reconcile(prisma)
assert result.days_reconciled == ("2026-09-01", "2026-09-02")
assert result.reconciled_through == "2026-09-14"
marker = await read_marker(prisma)
assert marker is not None and (marker.reconciled_through, marker.scanned_at) == ("2026-09-14", "clock-0009")
@pytest.mark.asyncio
async def test_a_failure_with_nothing_done_reports_the_previous_marker_and_alerts():
"""When the rewrite of a late day fails the marker must stay put and the operator must hear about it."""
prisma = _FakePrisma(user_days=("2026-09-13",), today=date(2026, 9, 14))
await run_daily_global_spend_reconcile(prisma)
prisma.today = TODAY
prisma.write_late_row("2026-09-12")
prisma.failing_days = frozenset({"2026-09-12"})
alert = AsyncMock()
result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=None, alert=alert)
assert result is not None
assert result.days_reconciled == ()
assert result.failed_day == "2026-09-12"
assert result.reconciled_through == "2026-09-13"
alert.assert_awaited_once()
assert "2026-09-12" in alert.await_args.args[0]
@pytest.mark.asyncio
async def test_a_clean_run_does_not_alert():
prisma = _FakePrisma(user_days=("2026-09-13",))
alert = AsyncMock()
await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=None, alert=alert)
alert.assert_not_awaited()
def _pod_lock(acquired: bool) -> MagicMock:
lock = MagicMock()
lock.redis_cache = MagicMock()
lock.redis_cache.async_get_cache = AsyncMock(return_value="other-pod")
lock.get_redis_lock_key = MagicMock(return_value="lock-key")
lock.acquire_lock = AsyncMock(return_value=acquired)
lock.release_lock = AsyncMock()
return lock
@pytest.mark.asyncio
async def test_scheduled_run_skips_when_another_pod_holds_the_lock():
prisma = _FakePrisma(user_days=("2026-09-13",))
lock = _pod_lock(acquired=False)
result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock)
assert result is None
assert prisma.reconciled == []
lock.release_lock.assert_not_awaited()
@pytest.mark.asyncio
async def test_scheduled_run_runs_and_releases_the_lock_when_it_wins():
prisma = _FakePrisma(user_days=("2026-09-13",))
lock = _pod_lock(acquired=True)
result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock)
assert result is not None and result.days_reconciled == ("2026-09-13",)
lock.release_lock.assert_awaited_once()
@pytest.mark.asyncio
async def test_scheduled_run_proceeds_when_the_lock_cannot_be_acquired_or_read():
"""A Redis outage must not stall the backfill: the day rewrite is idempotent, so running
twice is only wasted effort while skipping forever leaves usage on the slow path."""
prisma = _FakePrisma(user_days=("2026-09-13",))
lock = _pod_lock(acquired=False)
lock.redis_cache.async_get_cache = AsyncMock(side_effect=ConnectionError("redis down"))
result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock)
assert result is not None and result.days_reconciled == ("2026-09-13",)
lock.release_lock.assert_not_awaited()
@pytest.mark.asyncio
async def test_marker_is_read_back_from_the_json_string_the_config_table_stores():
prisma = _FakePrisma(user_days=())
prisma.db.litellm_config.rows[DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM] = '{"reconciled_through": "2026-09-10"}'
assert await reconciled_through(prisma) == "2026-09-10"
@pytest.mark.asyncio
async def test_an_unparseable_marker_reads_as_never_reconciled():
prisma = _FakePrisma(user_days=())
prisma.db.litellm_config.rows[DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM] = '{"something_else": 1}'
assert await reconciled_through(prisma) is None
_rollup_postgresql_proc: Final = factories.postgresql_proc()
_rollup_postgresql: Final = factories.postgresql("_rollup_postgresql_proc")
_MIGRATIONS_DIR: Final = (
pathlib.Path(__file__).resolve().parents[4] / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations"
)
_GLOBAL_SPEND_MIGRATION: Final = _MIGRATIONS_DIR / "20260915000000_add_daily_global_spend" / "migration.sql"
_DAILY_USER_SPEND_DDL: Final = """
CREATE TABLE "LiteLLM_DailyUserSpend" (
id TEXT PRIMARY KEY,
user_id TEXT,
date TEXT NOT NULL,
api_key TEXT NOT NULL,
model TEXT,
model_group TEXT,
custom_llm_provider TEXT,
mcp_namespaced_tool_name TEXT,
endpoint TEXT,
prompt_tokens BIGINT DEFAULT 0,
completion_tokens BIGINT DEFAULT 0,
cache_read_input_tokens BIGINT DEFAULT 0,
cache_creation_input_tokens BIGINT DEFAULT 0,
compression_saved_tokens BIGINT DEFAULT 0,
compression_savings_spend DOUBLE PRECISION DEFAULT 0,
prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
autorouter_savings_spend DOUBLE PRECISION DEFAULT 0,
spend DOUBLE PRECISION DEFAULT 0,
api_requests BIGINT DEFAULT 0,
successful_requests BIGINT DEFAULT 0,
failed_requests BIGINT DEFAULT 0,
total_response_time_ms BIGINT DEFAULT 0,
timed_requests BIGINT DEFAULT 0,
created_at TIMESTAMP DEFAULT now(),
updated_at TIMESTAMP,
UNIQUE (user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint)
)
"""
_PER_KEY_SUMS_SQL: Final = """
SELECT COALESCE(model, '') AS model, COALESCE(model_group, '') AS model_group,
COALESCE(custom_llm_provider, '') AS custom_llm_provider,
SUM(spend) AS spend, SUM(prompt_tokens) AS prompt_tokens, SUM(api_requests) AS api_requests,
SUM(total_response_time_ms) AS total_response_time_ms, SUM(timed_requests) AS timed_requests
FROM "LiteLLM_DailyUserSpend" WHERE date = %s
GROUP BY 1, 2, 3 ORDER BY 1, 2, 3
"""
_GLOBAL_ROWS_SQL: Final = """
SELECT model, model_group, custom_llm_provider, spend, prompt_tokens, api_requests,
total_response_time_ms, timed_requests
FROM "LiteLLM_DailyGlobalSpend" WHERE date = %s ORDER BY 1, 2, 3
"""
def _execute_dollar_sql(conn: psycopg.Connection, sql: str, params: tuple[object, ...]) -> None:
converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql)
conn.execute(
converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query
{f"p{i}": v for i, v in enumerate(params, start=1)},
)
conn.commit()
def _user_txn(**overrides):
return {
"user_id": "u-1",
"date": "2026-09-14",
"api_key": "sk-1",
"model": "gpt-5",
"model_group": "gpt-5",
"custom_llm_provider": "openai",
"mcp_namespaced_tool_name": "",
"endpoint": "/chat/completions",
"prompt_tokens": 10,
"completion_tokens": 20,
"spend": 1.0,
"api_requests": 1,
"successful_requests": 1,
"failed_requests": 0,
"total_response_time_ms": 800,
"timed_requests": 1,
**overrides,
}
def _normalized(rows: list[dict[str, object]]) -> list[tuple[object, ...]]:
return [
(
r["model"],
r["model_group"],
r["custom_llm_provider"],
float(r["spend"]),
int(r["prompt_tokens"]),
int(r["api_requests"]),
int(r["total_response_time_ms"]),
int(r["timed_requests"]),
) # pyright: ignore[reportArgumentType] # dict_row values are untyped
for r in rows
]
def test_reconcile_day_sql_makes_the_global_day_equal_the_per_key_sums(_rollup_postgresql: psycopg.Connection):
"""Against real Postgres and the shipped migration: writer-shaped rows and legacy rows
(NULL and '' dimension spellings) fold into one global day, running the day twice changes
nothing, and other days are left alone."""
conn: Final = _rollup_postgresql
conn.execute(_DAILY_USER_SPEND_DDL) # pyright: ignore[reportArgumentType] # DDL literal
conn.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal
conn.commit()
written_batch = merge_by_conflict_key(
USER_TABLE,
(_user_txn(api_key="sk-1", spend=1.0), _user_txn(api_key="sk-2", user_id="u-2", spend=2.0, prompt_tokens=20)),
)
_execute_dollar_sql(conn, *build_bulk_upsert(USER_TABLE, written_batch))
conn.execute(
"""
INSERT INTO "LiteLLM_DailyUserSpend"
(id, user_id, date, api_key, model, model_group, custom_llm_provider, mcp_namespaced_tool_name,
endpoint, prompt_tokens, spend, api_requests)
VALUES
('legacy-1', 'u-9', '2026-09-14', 'sk-9', 'gpt-5', NULL, 'openai', NULL, NULL, 5, 4.0, 1),
('legacy-2', 'u-9', '2026-09-14', 'sk-9', 'gpt-5', '', 'openai', '', '', 5, 8.0, 1),
('legacy-3', 'u-9', '2026-09-13', 'sk-9', 'claude', '', 'anthropic', '', '', 7, 16.0, 1)
"""
)
conn.commit()
_execute_dollar_sql(conn, RECONCILE_DAY_SQL, ("2026-09-14",))
_execute_dollar_sql(conn, RECONCILE_DAY_SQL, ("2026-09-14",))
with conn.cursor(row_factory=dict_row) as cur:
global_rows = cur.execute(_GLOBAL_ROWS_SQL, ("2026-09-14",)).fetchall()
per_key = cur.execute(_PER_KEY_SUMS_SQL, ("2026-09-14",)).fetchall()
untouched = cur.execute(_GLOBAL_ROWS_SQL, ("2026-09-13",)).fetchall()
assert _normalized(global_rows) == _normalized(per_key)
assert sum(float(r["spend"]) for r in global_rows) == pytest.approx(15.0) # pyright: ignore[reportArgumentType] # dict_row values are untyped
assert sum(int(r["total_response_time_ms"]) for r in global_rows) == 1600 # pyright: ignore[reportArgumentType] # dict_row values are untyped
assert [(r["model"], r["model_group"]) for r in global_rows] == [("gpt-5", ""), ("gpt-5", "gpt-5")]
assert untouched == []
_CONFIG_DDL: Final = 'CREATE TABLE "LiteLLM_Config" (param_name TEXT PRIMARY KEY, param_value JSONB)'
_MARKER_SQL: Final = 'SELECT param_value FROM "LiteLLM_Config" WHERE param_name = %s'
def test_advance_marker_sql_only_ever_moves_the_stored_marker_forward(_rollup_postgresql: psycopg.Connection):
"""Against real Postgres: the statement a slower overlapping run issues after the faster run
already stored a later marker leaves that marker alone, whether it carries an older scan time or
none at all, while a run that is further along moves both fields on."""
conn: Final = _rollup_postgresql
conn.execute(_CONFIG_DDL) # pyright: ignore[reportArgumentType] # DDL literal
conn.commit()
param: Final = DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM
def stored() -> object:
with conn.cursor(row_factory=dict_row) as cur:
row = cur.execute(_MARKER_SQL, (param,)).fetchone()
return None if row is None else row["param_value"]
_execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-01", None))
assert stored() == {"reconciled_through": "2026-09-01", "scanned_at": None}
_execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-14", "2026-09-15 00:30:02.5"))
_execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-02", None))
_execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-03", "2026-09-15 00:30:01.25"))
assert stored() == {"reconciled_through": "2026-09-14", "scanned_at": "2026-09-15 00:30:02.5"}
_execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-15", "2026-09-16 00:30:00.75"))
assert stored() == {"reconciled_through": "2026-09-15", "scanned_at": "2026-09-16 00:30:00.75"}

View file

@ -6711,6 +6711,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
monkeypatch: pytest.MonkeyPatch,
user_api_key_dict: ProxyUserAPIKeyAuth,
fallbacks: list[dict[str, list[str]]],
model_guardrails: dict[str, list[str]] | None = None,
) -> tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]]:
"""Real v3 limiter (the default ``parallel_request_limiter``) wired in through the
``proxy_logging_obj`` seam, so ``common_processing_pre_call_logic`` runs for real:
@ -6738,9 +6739,17 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
proxy_logging_obj = MagicMock(spec=ProxyLogging)
proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=run_limiter)
guardrails_by_group = model_guardrails or {}
router = litellm.Router(
model_list=[
{"model_name": group, "litellm_params": {"model": "openai/gpt-4.1-nano", "api_key": "fake"}}
{
"model_name": group,
"litellm_params": {
"model": "openai/gpt-4.1-nano",
"api_key": "fake",
**({"guardrails": guardrails_by_group[group]} if group in guardrails_by_group else {}),
},
}
for chain in fallbacks
for group in (*chain.keys(), *(m for models in chain.values() for m in models))
],
@ -6752,7 +6761,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
def _otel_key(
rpm_limit: int | None = None,
model_rpm_limit: dict[str, int] | None = None,
disable_fallbacks: bool = False,
disable_fallbacks: bool | None = None,
) -> ProxyUserAPIKeyAuth:
from opentelemetry.sdk.trace import TracerProvider
@ -6763,7 +6772,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
rpm_limit=rpm_limit,
metadata={
**({"model_rpm_limit": model_rpm_limit} if model_rpm_limit else {}),
**({"disable_fallbacks": True} if disable_fallbacks else {}),
**({"disable_fallbacks": disable_fallbacks} if disable_fallbacks is not None else {}),
},
)
@ -6906,6 +6915,81 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
assert exc_info.value.status_code == 429
assert rig[3] == [primary_model, primary_model]
@pytest.mark.asyncio
async def test_key_metadata_disable_fallbacks_false_overrides_request_body(self, monkeypatch: pytest.MonkeyPatch):
primary_model = "gpt-4.1"
fallback_model = "gpt-4.1-mini"
key = self._otel_key(model_rpm_limit={primary_model: 1}, disable_fallbacks=False)
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
request = {
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
"disable_fallbacks": True,
}
await self._pre_call(dict(request), key, rig)
_, (data, _) = await self._pre_call(dict(request), key, rig)
assert data["model"] == fallback_model
assert data["disable_fallbacks"] is False
assert rig[3] == [primary_model, primary_model, fallback_model]
@pytest.mark.asyncio
async def test_fallback_keeps_requested_model_guardrails(self, monkeypatch: pytest.MonkeyPatch):
primary_model = "gpt-4.1"
fallback_model = "gpt-4.1-mini"
guardrail = "pii-guard-for-primary"
key = self._otel_key(model_rpm_limit={primary_model: 1})
rig = self._v3_limiter_rig(
monkeypatch, key, [{primary_model: [fallback_model]}], model_guardrails={primary_model: [guardrail]}
)
run_limiter = rig[0].pre_call_hook
async def limiter_then_guardrail(
user_api_key_dict: ProxyUserAPIKeyAuth, data: dict[str, object], call_type: str
) -> dict[str, object]:
limited = await run_limiter(user_api_key_dict=user_api_key_dict, data=data, call_type=call_type)
if guardrail not in (limited["metadata"].get("guardrails") or []):
return limited
return {
**limited,
"messages": [
{**m, "content": str(m["content"]).replace("123-45-6789", "[REDACTED-SSN]")}
for m in limited["messages"]
],
}
rig[0].pre_call_hook = AsyncMock(side_effect=limiter_then_guardrail)
request = {"model": primary_model, "messages": [{"role": "user", "content": "my ssn is 123-45-6789"}]}
await self._pre_call(dict(request), key, rig)
_, (data, _) = await self._pre_call(dict(request), key, rig)
assert data["model"] == fallback_model
assert guardrail in data["metadata"]["guardrails"]
assert data["messages"] == [{"role": "user", "content": "my ssn is [REDACTED-SSN]"}]
assert rig[3] == [primary_model, primary_model, fallback_model]
@pytest.mark.asyncio
async def test_fallback_keeps_structured_request_guardrails(self, monkeypatch: pytest.MonkeyPatch):
primary_model = "gpt-4.1"
fallback_model = "gpt-4.1-mini"
structured_guardrail = {"pii-guard": {"extra_body": {"threshold": 0.5}}}
key = self._otel_key(model_rpm_limit={primary_model: 1})
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
request = {
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
"guardrails": [structured_guardrail],
}
await self._pre_call(dict(request), key, rig)
_, (data, _) = await self._pre_call(dict(request), key, rig)
assert data["model"] == fallback_model
assert data["metadata"]["guardrails"] == [structured_guardrail]
assert rig[3] == [primary_model, primary_model, fallback_model]
class _RecordingSuccessLogger(CustomLogger):
def __init__(self):

View file

@ -7607,8 +7607,9 @@ async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_servi
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {}) # test-quality-ok: the method reads this module global; no injection seam
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", None) # test-quality-ok: module global holding the YAML endpoints; this case has none
app_routes: Final = patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists") # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker
try:
with settings, yaml_endpoints:
with settings, yaml_endpoints, app_routes:
pc = ProxyConfig()
await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
assert live_routes(), "the stored endpoint should be serving before the row is deleted"
@ -7648,8 +7649,9 @@ async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_rout
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [config_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [config_endpoint]) # test-quality-ok: module global holding the YAML endpoints the reload merges in
app_routes: Final = patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists") # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker
try:
with settings, yaml_endpoints:
with settings, yaml_endpoints, app_routes:
await initialize_pass_through_endpoints(pass_through_endpoints=[config_endpoint])
assert live_paths() == {config_path}

View file

@ -4,6 +4,7 @@ import litellm
from litellm.router_utils.reasoning_effort_capability import (
deployment_is_catalog_mapped,
intersect_supported_reasoning_efforts,
nearest_declared_reasoning_effort,
resolve_supported_reasoning_efforts,
)
@ -415,3 +416,26 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels:
"high",
"xhigh",
)
class TestNearestDeclaredReasoningEffort:
def test_a_declared_level_is_kept(self):
assert nearest_declared_reasoning_effort("high", ("none", "high")) == "high"
assert nearest_declared_reasoning_effort("none", ("none", "high")) == "none"
def test_an_undeclared_level_rounds_up_to_the_next_declared_one(self):
assert nearest_declared_reasoning_effort("medium", ("none", "high")) == "high"
assert nearest_declared_reasoning_effort("minimal", ("low", "high", "max")) == "low"
assert nearest_declared_reasoning_effort("xhigh", ("low", "high", "max")) == "max"
def test_none_is_a_switch_that_is_never_rounded_in_either_direction(self):
assert nearest_declared_reasoning_effort("none", ("low", "high", "max")) == "none"
assert nearest_declared_reasoning_effort("medium", ("none",)) == "medium"
def test_a_level_above_the_ceiling_takes_the_strongest_declared_one(self):
assert nearest_declared_reasoning_effort("max", ("none", "high")) == "high"
assert nearest_declared_reasoning_effort("xhigh", ("none", "low", "medium", "high")) == "high"
def test_a_level_outside_the_strength_order_is_left_for_upstream(self):
assert nearest_declared_reasoning_effort("turbo", ("none", "high")) == "turbo"
assert nearest_declared_reasoning_effort("medium", ()) == "medium"

View file

@ -59,6 +59,17 @@ def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> N
assert public_error.llm_provider == "mistral"
def test_map_failure_maps_upstream_401_to_authentication_error() -> None:
error: Final = RustUpstreamError(401, '{"message": "Unauthorized"}', ())
public_error: Final = map_failure(error, REQUEST, "mistral")
assert isinstance(public_error, litellm.AuthenticationError)
assert public_error.status_code == 401
assert public_error.response.text == '{"message": "Unauthorized"}'
assert public_error.__context__ is error
def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None:
error: Final = RuntimeError("bridge exploded")

View file

@ -906,6 +906,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"/v1/audio/speech",
"/v1/ocr",
"/vertex_ai/live",
"/v1/listen",
"/v1beta/interactions",
],
},

View file

@ -31419,6 +31419,8 @@ export interface components {
s3_bucket_name?: string | null;
/** S3 Encryption Key Id */
s3_encryption_key_id?: string | null;
/** S3 Endpoint Url */
s3_endpoint_url?: string | null;
/** S3 Output Bucket Name */
s3_output_bucket_name?: string | null;
/** S3 Region Name */
@ -42054,6 +42056,8 @@ export interface components {
s3_bucket_name?: string | null;
/** S3 Encryption Key Id */
s3_encryption_key_id?: string | null;
/** S3 Endpoint Url */
s3_endpoint_url?: string | null;
/** S3 Output Bucket Name */
s3_output_bucket_name?: string | null;
/** S3 Region Name */