mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
merge: preserve realtime authentication with upstream main
This commit is contained in:
commit
906d5997bf
61 changed files with 4083 additions and 638 deletions
|
|
@ -95,6 +95,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/vertex-ai/",
|
||||
"/assemblyai/",
|
||||
"/eu.assemblyai/",
|
||||
"/deepgram/",
|
||||
"/langfuse/",
|
||||
"/vllm/",
|
||||
"/mistral/",
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -201,6 +201,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
|
|||
"/cohere/",
|
||||
"/comprehendmedical",
|
||||
"/cursor/",
|
||||
"/deepgram/",
|
||||
"/eu.assemblyai/",
|
||||
"/gemini/",
|
||||
"/gigachat/",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
293
litellm/proxy/spend_tracking/daily_global_spend_rollup.py
Normal file
293
litellm/proxy/spend_tracking/daily_global_spend_rollup.py
Normal 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
|
@ -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"),
|
||||
}
|
||||
|
|
@ -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"}
|
||||
|
||||
|
|
|
|||
317
tests/ocr_tests/test_ocr_matrix.py
Normal file
317
tests/ocr_tests/test_ocr_matrix.py
Normal 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)
|
||||
|
|
@ -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)}")
|
||||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
326
tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py
Normal file
326
tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py
Normal 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
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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]
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
},
|
||||
|
|
|
|||
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue