mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
Merge branch 'litellm_internal_staging' into litellm_fallback_generalizations_top
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
534a8a3b5b
92 changed files with 11315 additions and 1052 deletions
6
.github/pull_request_template.md
vendored
6
.github/pull_request_template.md
vendored
|
|
@ -83,7 +83,11 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac
|
|||
🚄 Infrastructure
|
||||
✅ Test
|
||||
|
||||
## Changes
|
||||
## Caveats (if any)
|
||||
|
||||
<!-- Short bullet points, just like the TLDR: one line per bullet, roughly 10 words max
|
||||
Call out known limitations, follow-up work, or anything a reviewer should watch out for
|
||||
Leave this section empty if there are none -->
|
||||
|
||||
## QA runbook
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
Do not write comments unless they are:
|
||||
Do not write comments unless they are any of:
|
||||
- absolutely necessary to explain some very complex business logic (in which case, keep it concise and clear)
|
||||
- used as an input for tools to read and act on. For example:
|
||||
- entries in `.git-blame-ignore-revs` saying which commit is excluded from git blame
|
||||
|
|
|
|||
|
|
@ -1,18 +1,18 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 27731
|
||||
"limit": 26391
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2626
|
||||
"limit": 2614
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 329
|
||||
"limit": 327
|
||||
},
|
||||
"reportAttributeAccessIssue": {
|
||||
"limit": 514
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"limit": 116
|
||||
"limit": 114
|
||||
},
|
||||
"reportConstantRedefinition": {
|
||||
"limit": 40
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 8807
|
||||
"limit": 8319
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5835
|
||||
"limit": 5825
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15790
|
||||
"limit": 15695
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -90,7 +90,7 @@
|
|||
"limit": 8
|
||||
},
|
||||
"reportReturnType": {
|
||||
"limit": 217
|
||||
"limit": 213
|
||||
},
|
||||
"reportTypedDictNotRequiredAccess": {
|
||||
"limit": 26
|
||||
|
|
@ -99,22 +99,22 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 45063
|
||||
"limit": 45004
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 113
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 39773
|
||||
"limit": 39649
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 20207
|
||||
"limit": 20132
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 31281
|
||||
"limit": 31156
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 122
|
||||
"limit": 118
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 701
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 862
|
||||
"limit": 857
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_DailyTeamSpend" ADD COLUMN IF NOT EXISTS "ptu_flat_cost" DOUBLE PRECISION NOT NULL DEFAULT 0.0;
|
||||
|
|
@ -30,7 +30,7 @@ model LiteLLM_BudgetTable {
|
|||
end_users LiteLLM_EndUserTable[] // multiple end-users can have the same budget
|
||||
tags LiteLLM_TagTable[] // multiple tags can have the same budget
|
||||
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
|
||||
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
|
||||
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
|
||||
}
|
||||
|
||||
// Models on proxy
|
||||
|
|
@ -893,6 +893,7 @@ model LiteLLM_DailyTeamSpend {
|
|||
api_requests BigInt @default(0)
|
||||
successful_requests BigInt @default(0)
|
||||
failed_requests BigInt @default(0)
|
||||
ptu_flat_cost Float @default(0.0)
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
||||
|
|
|
|||
|
|
@ -197,6 +197,7 @@ standard_logging_payload_excluded_fields: Optional[List[str]] = (
|
|||
None # Fields to exclude from StandardLoggingPayload before callbacks receive it
|
||||
)
|
||||
log_raw_request_response: bool = False
|
||||
request_correlation_in_logs: bool = False
|
||||
redact_messages_in_exceptions: Optional[bool] = False
|
||||
redact_user_api_key_info: Optional[bool] = False
|
||||
# When True (default — preserves historical behavior), the Router appends
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import ast
|
||||
import contextvars
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
|
|
@ -6,12 +7,44 @@ from datetime import datetime
|
|||
from logging import Formatter
|
||||
from typing import Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.litellm_core_utils.secret_redaction import redact_string
|
||||
|
||||
set_verbose = False
|
||||
|
||||
session_id_var: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("session_id", default="")
|
||||
trace_id_var: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("trace_id", default="")
|
||||
|
||||
_MAX_CORRELATION_ID_LENGTH: Final = 256
|
||||
|
||||
|
||||
def _sanitize_correlation_id(value: str) -> str:
|
||||
"""Strip control characters, bound length, and redact credential-shaped
|
||||
content before a caller-controlled trace_id/session_id (e.g.
|
||||
litellm_session_id, x-litellm-trace-id) is stamped into log lines.
|
||||
|
||||
Without the first two, a caller could embed \\r/\\n or terminal escape
|
||||
sequences to forge fake log entries, or submit an oversized value repeated
|
||||
across every log line for the request. Without the redaction, a caller
|
||||
could smuggle a real credential (e.g. an sk-... key) through this field:
|
||||
CorrelationContextFilter stamps trace_id/session_id onto the record after
|
||||
SecretRedactionFilter has already run, so those two fields never otherwise
|
||||
pass through credential redaction.
|
||||
"""
|
||||
stripped: Final = "".join(ch for ch in value if ch.isprintable())
|
||||
return _redact_string(stripped[:_MAX_CORRELATION_ID_LENGTH])
|
||||
|
||||
|
||||
def set_session_id(session_id: str) -> "contextvars.Token[str]":
|
||||
return session_id_var.set(_sanitize_correlation_id(session_id))
|
||||
|
||||
|
||||
def set_trace_id(trace_id: str) -> "contextvars.Token[str]":
|
||||
return trace_id_var.set(_sanitize_correlation_id(trace_id))
|
||||
|
||||
|
||||
if set_verbose is True:
|
||||
logging.warning(
|
||||
"`litellm.set_verbose` is deprecated. Please set `os.environ['LITELLM_LOG'] = 'DEBUG'` for debug logs."
|
||||
|
|
@ -77,6 +110,28 @@ class SecretRedactionFilter(logging.Filter):
|
|||
_secret_filter: Final = SecretRedactionFilter()
|
||||
|
||||
|
||||
class CorrelationContextFilter(logging.Filter):
|
||||
"""Stamps each log record with the current request's trace_id and session_id from contextvars.
|
||||
|
||||
Works in tandem with JsonFormatter: the formatter's record.__dict__ loop picks up these
|
||||
attributes as first-class JSON fields without any formatter-level code.
|
||||
"""
|
||||
|
||||
def filter(self, record: logging.LogRecord) -> bool:
|
||||
if not litellm.request_correlation_in_logs:
|
||||
return True
|
||||
trace_id: Final = trace_id_var.get()
|
||||
if trace_id:
|
||||
record.trace_id = trace_id # rebind-ok: stamping the LogRecord is the Filter interface's contract
|
||||
session_id: Final = session_id_var.get()
|
||||
if session_id:
|
||||
record.session_id = session_id # rebind-ok: stamping the LogRecord is the Filter interface's contract
|
||||
return True
|
||||
|
||||
|
||||
_correlation_filter: Final = CorrelationContextFilter()
|
||||
|
||||
|
||||
json_logs = bool(os.getenv("JSON_LOGS", False))
|
||||
# Create a handler for the logger (you may need to adapt this based on your needs)
|
||||
log_level: Final = os.getenv("LITELLM_LOG", "DEBUG")
|
||||
|
|
@ -84,6 +139,7 @@ numeric_level: Final[str] = getattr(logging, log_level.upper())
|
|||
handler: Final = logging.StreamHandler()
|
||||
handler.setLevel(numeric_level)
|
||||
handler.addFilter(_secret_filter)
|
||||
handler.addFilter(_correlation_filter)
|
||||
|
||||
|
||||
def _try_parse_json_message(message: str) -> dict[str, Any] | None:
|
||||
|
|
@ -146,6 +202,11 @@ def _get_standard_record_attrs() -> frozenset:
|
|||
|
||||
_STANDARD_RECORD_ATTRS: Final = _get_standard_record_attrs()
|
||||
|
||||
# CorrelationContextFilter is the only legitimate source for these two JSON fields;
|
||||
# see JsonFormatter.format() for why they're excluded from the generic message-content
|
||||
# and extra-attribute promotion paths.
|
||||
_RESERVED_CORRELATION_FIELDS: Final = frozenset(("trace_id", "session_id"))
|
||||
|
||||
|
||||
class JsonFormatter(Formatter):
|
||||
def __init__(self):
|
||||
|
|
@ -164,13 +225,18 @@ class JsonFormatter(Formatter):
|
|||
"timestamp": self.formatTime(record),
|
||||
}
|
||||
|
||||
# Parse embedded JSON or Python dict repr in message so sub-fields become first-class properties
|
||||
# Parse embedded JSON or Python dict repr in message so sub-fields become first-class properties.
|
||||
# trace_id/session_id are excluded here unconditionally (not just "if not already
|
||||
# set") - CorrelationContextFilter is the only legitimate source for these two
|
||||
# fields, and a message that merely happens to parse as JSON/dict (e.g. a proxy
|
||||
# log line dumping raw request headers) must never be able to claim them, even on
|
||||
# a record the filter hasn't stamped yet (no correlation context active for it).
|
||||
parsed = _try_parse_json_message(message_str)
|
||||
if parsed is None:
|
||||
parsed = _try_parse_embedded_python_dict(message_str)
|
||||
if parsed is not None:
|
||||
for key, value in parsed.items():
|
||||
if key not in json_record:
|
||||
if key not in json_record and key not in _RESERVED_CORRELATION_FIELDS:
|
||||
json_record[key] = value
|
||||
|
||||
# Include extra attributes passed via logger.debug("msg", extra={...})
|
||||
|
|
@ -178,6 +244,18 @@ class JsonFormatter(Formatter):
|
|||
if key not in _STANDARD_RECORD_ATTRS and key not in json_record:
|
||||
json_record[key] = value
|
||||
|
||||
# trace_id/session_id are reserved: CorrelationContextFilter is the only
|
||||
# legitimate source for these two fields. Without this, a message string
|
||||
# that happens to parse as JSON/dict (e.g. a proxy log line dumping raw
|
||||
# request headers) with a "trace_id"/"session_id" key would have already
|
||||
# claimed the key at the parsed-message step above, and the extra-attributes
|
||||
# loop's "key not in json_record" guard would then skip the real value -
|
||||
# letting a caller-supplied header spoof another request's correlation ids.
|
||||
for reserved_key in _RESERVED_CORRELATION_FIELDS:
|
||||
value = getattr(record, reserved_key, None)
|
||||
if value:
|
||||
json_record[reserved_key] = value
|
||||
|
||||
# Set component/logger only if not already supplied via extra={...}
|
||||
if "component" not in json_record:
|
||||
json_record["component"] = record.name
|
||||
|
|
@ -190,12 +268,34 @@ class JsonFormatter(Formatter):
|
|||
return safe_dumps(json_record)
|
||||
|
||||
|
||||
class CorrelationPlainFormatter(logging.Formatter):
|
||||
"""Appends trace_id/session_id to plain-text log lines stamped by CorrelationContextFilter.
|
||||
|
||||
Mirrors JsonFormatter's handling of these two fields so request_correlation_in_logs
|
||||
behaves the same whether or not json_logs is enabled.
|
||||
"""
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
formatted: Final = super().format(record)
|
||||
trace_id: Final = getattr(record, "trace_id", None)
|
||||
session_id: Final = getattr(record, "session_id", None)
|
||||
if not trace_id and not session_id:
|
||||
return formatted
|
||||
parts: Final = tuple(
|
||||
p
|
||||
for p in (f"trace_id={trace_id}" if trace_id else None, f"session_id={session_id}" if session_id else None)
|
||||
if p
|
||||
)
|
||||
return f"{formatted} [{' '.join(parts)}]"
|
||||
|
||||
|
||||
# Function to set up exception handlers for JSON logging
|
||||
def _setup_json_exception_handlers(formatter):
|
||||
# Create a handler with JSON formatting for exceptions
|
||||
error_handler: Final = logging.StreamHandler()
|
||||
error_handler.setFormatter(formatter)
|
||||
error_handler.addFilter(_secret_filter)
|
||||
error_handler.addFilter(_correlation_filter)
|
||||
|
||||
# Setup excepthook for uncaught exceptions
|
||||
def json_excepthook(exc_type, exc_value, exc_traceback):
|
||||
|
|
@ -243,7 +343,7 @@ if json_logs:
|
|||
handler.setFormatter(JsonFormatter())
|
||||
_setup_json_exception_handlers(JsonFormatter())
|
||||
else:
|
||||
formatter: Final = logging.Formatter(
|
||||
formatter: Final = CorrelationPlainFormatter(
|
||||
"\033[92m%(asctime)s - %(name)s:%(levelname)s\033[0m: %(filename)s:%(lineno)s - %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
|
|
@ -346,6 +446,7 @@ def _initialize_loggers_with_handler(handler: logging.Handler):
|
|||
- Prevents bubbling to parent/root (critical to prevent duplicate JSON logs)
|
||||
"""
|
||||
handler.addFilter(_secret_filter)
|
||||
handler.addFilter(_correlation_filter)
|
||||
for lg in _get_loggers_to_initialize():
|
||||
lg.handlers.clear() # remove any existing handlers
|
||||
lg.addHandler(handler) # add JSON formatter handler
|
||||
|
|
|
|||
|
|
@ -1719,3 +1719,18 @@ BROWSER_SECURITY_HEADERS: Final[frozenset[str]] = frozenset(
|
|||
)
|
||||
|
||||
UNSAFE_PROXY_RESPONSE_HEADERS: Final[frozenset[str]] = HTTP_FRAMING_HEADERS | BROWSER_SECURITY_HEADERS
|
||||
|
||||
# PTU reservation rollup writes rows to LiteLLM_DailyTeamSpend with this
|
||||
# sentinel api_key so PTU flat cost stays distinguishable from real per-request
|
||||
# spend under the table's composite unique constraint.
|
||||
PTU_SENTINEL_API_KEY: Final[str] = "__ptu_flat_cost__"
|
||||
PTU_ROLLUP_JOB_ID: Final[str] = "ptu_flat_cost_rollup_job"
|
||||
PTU_ROLLUP_LOCK_TTL_SECONDS: Final[int] = 900
|
||||
# Furthest back the catch-up pass looks for unpriced PTU days when a deployment
|
||||
# declares no ptu_effective_from, bounding the scan for an open-ended window.
|
||||
PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90
|
||||
# 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
|
||||
# one is seconds old, so a few minutes separates them.
|
||||
PTU_PRUNE_SKEW_GRACE_SECONDS: Final[int] = 300
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import asyncio
|
|||
import contextvars
|
||||
from collections.abc import Coroutine
|
||||
from functools import partial
|
||||
from typing import Any, Final
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -21,8 +21,10 @@ from litellm.types.llms.openai_evals import (
|
|||
CancelRunResponse,
|
||||
CreateEvalRequest,
|
||||
CreateRunRequest,
|
||||
DataSourceConfig,
|
||||
DeleteEvalResponse,
|
||||
Eval,
|
||||
GraderConfig,
|
||||
ListEvalsParams,
|
||||
ListEvalsResponse,
|
||||
ListRunsParams,
|
||||
|
|
@ -41,13 +43,13 @@ DEFAULT_OPENAI_API_BASE: Final = "https://api.openai.com"
|
|||
|
||||
@client
|
||||
async def acreate_eval(
|
||||
data_source_config: dict[str, Any],
|
||||
testing_criteria: list[dict[str, Any]],
|
||||
data_source_config: DataSourceConfig,
|
||||
testing_criteria: list[GraderConfig],
|
||||
name: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -110,17 +112,17 @@ async def acreate_eval(
|
|||
|
||||
@client
|
||||
def create_eval(
|
||||
data_source_config: dict[str, Any],
|
||||
testing_criteria: list[dict[str, Any]],
|
||||
data_source_config: DataSourceConfig,
|
||||
testing_criteria: list[GraderConfig],
|
||||
name: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> Eval | Coroutine[Any, Any, Eval]:
|
||||
) -> Eval | Coroutine[object, object, Eval]:
|
||||
"""
|
||||
Create a new evaluation
|
||||
|
||||
|
|
@ -231,8 +233,8 @@ async def alist_evals(
|
|||
before: str | None = None,
|
||||
order: str | None = None,
|
||||
order_by: str | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -300,12 +302,12 @@ def list_evals(
|
|||
before: str | None = None,
|
||||
order: str | None = None,
|
||||
order_by: str | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> ListEvalsResponse | Coroutine[Any, Any, ListEvalsResponse]:
|
||||
) -> ListEvalsResponse | Coroutine[object, object, ListEvalsResponse]:
|
||||
"""
|
||||
List all evaluations
|
||||
|
||||
|
|
@ -413,8 +415,8 @@ def list_evals(
|
|||
@client
|
||||
async def aget_eval(
|
||||
eval_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -470,12 +472,12 @@ async def aget_eval(
|
|||
@client
|
||||
def get_eval(
|
||||
eval_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> Eval | Coroutine[Any, Any, Eval]:
|
||||
) -> Eval | Coroutine[object, object, Eval]:
|
||||
"""
|
||||
Get an evaluation by ID
|
||||
|
||||
|
|
@ -564,10 +566,10 @@ def get_eval(
|
|||
async def aupdate_eval(
|
||||
eval_id: str,
|
||||
name: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -630,14 +632,14 @@ async def aupdate_eval(
|
|||
def update_eval(
|
||||
eval_id: str,
|
||||
name: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> Eval | Coroutine[Any, Any, Eval]:
|
||||
) -> Eval | Coroutine[object, object, Eval]:
|
||||
"""
|
||||
Update an evaluation
|
||||
|
||||
|
|
@ -783,8 +785,8 @@ def update_eval(
|
|||
@client
|
||||
async def adelete_eval(
|
||||
eval_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -840,12 +842,12 @@ async def adelete_eval(
|
|||
@client
|
||||
def delete_eval(
|
||||
eval_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> DeleteEvalResponse | Coroutine[Any, Any, DeleteEvalResponse]:
|
||||
) -> DeleteEvalResponse | Coroutine[object, object, DeleteEvalResponse]:
|
||||
"""
|
||||
Delete an evaluation
|
||||
|
||||
|
|
@ -933,8 +935,8 @@ def delete_eval(
|
|||
@client
|
||||
async def acancel_eval(
|
||||
eval_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -990,12 +992,12 @@ async def acancel_eval(
|
|||
@client
|
||||
def cancel_eval(
|
||||
eval_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> CancelEvalResponse | Coroutine[Any, Any, CancelEvalResponse]:
|
||||
) -> CancelEvalResponse | Coroutine[object, object, CancelEvalResponse]:
|
||||
"""
|
||||
Cancel a running evaluation
|
||||
|
||||
|
|
@ -1092,12 +1094,12 @@ def cancel_eval(
|
|||
@client
|
||||
async def acreate_run(
|
||||
eval_id: str,
|
||||
data_source: dict[str, Any],
|
||||
data_source: dict[str, object],
|
||||
name: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -1161,16 +1163,16 @@ async def acreate_run(
|
|||
@client
|
||||
def create_run(
|
||||
eval_id: str,
|
||||
data_source: dict[str, Any],
|
||||
data_source: dict[str, object],
|
||||
name: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> Run | Coroutine[Any, Any, Run]:
|
||||
) -> Run | Coroutine[object, object, Run]:
|
||||
"""
|
||||
Create a new run for an evaluation
|
||||
|
||||
|
|
@ -1280,8 +1282,8 @@ async def alist_runs(
|
|||
after: str | None = None,
|
||||
before: str | None = None,
|
||||
order: str | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -1349,12 +1351,12 @@ def list_runs(
|
|||
after: str | None = None,
|
||||
before: str | None = None,
|
||||
order: str | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> ListRunsResponse | Coroutine[Any, Any, ListRunsResponse]:
|
||||
) -> ListRunsResponse | Coroutine[object, object, ListRunsResponse]:
|
||||
"""
|
||||
List all runs for an evaluation
|
||||
|
||||
|
|
@ -1462,8 +1464,8 @@ def list_runs(
|
|||
async def aget_run(
|
||||
eval_id: str,
|
||||
run_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -1522,12 +1524,12 @@ async def aget_run(
|
|||
def get_run(
|
||||
eval_id: str,
|
||||
run_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> Run | Coroutine[Any, Any, Run]:
|
||||
) -> Run | Coroutine[object, object, Run]:
|
||||
"""
|
||||
Get a specific run
|
||||
|
||||
|
|
@ -1618,8 +1620,8 @@ def get_run(
|
|||
async def acancel_run(
|
||||
eval_id: str,
|
||||
run_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -1678,12 +1680,12 @@ async def acancel_run(
|
|||
def cancel_run(
|
||||
eval_id: str,
|
||||
run_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> CancelRunResponse | Coroutine[Any, Any, CancelRunResponse]:
|
||||
) -> CancelRunResponse | Coroutine[object, object, CancelRunResponse]:
|
||||
"""
|
||||
Cancel a running run
|
||||
|
||||
|
|
@ -1783,8 +1785,8 @@ def cancel_run(
|
|||
async def adelete_run(
|
||||
eval_id: str,
|
||||
run_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -1843,12 +1845,12 @@ async def adelete_run(
|
|||
def delete_run(
|
||||
eval_id: str,
|
||||
run_id: str,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> RunDeleteResponse | Coroutine[Any, Any, RunDeleteResponse]:
|
||||
) -> RunDeleteResponse | Coroutine[object, object, RunDeleteResponse]:
|
||||
"""
|
||||
Delete a run
|
||||
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import asyncio
|
|||
import math
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -73,6 +73,18 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b
|
|||
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks"
|
||||
|
||||
|
||||
class _PlanMetadataView(TypedDict):
|
||||
websearch_native_blocks: Sequence[Mapping[str, object]] | None
|
||||
|
||||
|
||||
class _AgenticLoopParamsView(TypedDict):
|
||||
agentic_loop_params: AgenticLoopParams
|
||||
|
||||
|
||||
class _WebSearchSettingsView(TypedDict):
|
||||
websearch_interception_params: WebSearchInterceptionConfig
|
||||
|
||||
|
||||
class WebSearchInterceptionLogger(CustomLogger):
|
||||
"""
|
||||
CustomLogger that intercepts WebSearch tool calls for models that don't
|
||||
|
|
@ -394,7 +406,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
return tool.get("name")
|
||||
|
||||
@classmethod
|
||||
def _sync_forced_tool_choice(cls, tool_choice: Any, converted_tools: list[dict[str, object]]) -> object:
|
||||
def _sync_forced_tool_choice(cls, tool_choice: object, converted_tools: Sequence[Mapping[str, object]]) -> object:
|
||||
"""Repoint a forced ``tool_choice`` at ``litellm_web_search`` when it
|
||||
names a web-search tool that was just converted away.
|
||||
|
||||
|
|
@ -462,7 +474,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True
|
||||
|
||||
# Convert native web search tools to LiteLLM standard
|
||||
converted_tools: Final = []
|
||||
converted_tools: Final[list[dict[str, object]]] = []
|
||||
for tool in tools:
|
||||
if is_web_search_tool(tool):
|
||||
standard_tool = get_litellm_web_search_tool()
|
||||
|
|
@ -833,7 +845,10 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
Anthropic-native clients (Claude Desktop, the Anthropic SDK) can
|
||||
render citations / sources alongside the model's textual reply.
|
||||
"""
|
||||
native_blocks: Final = plan.metadata.get(WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY)
|
||||
metadata_view: Final[_PlanMetadataView] = {
|
||||
"websearch_native_blocks": plan.metadata.get(WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY)
|
||||
}
|
||||
native_blocks: Final = metadata_view["websearch_native_blocks"]
|
||||
if not native_blocks:
|
||||
return response
|
||||
return self._inject_native_blocks(response, native_blocks)
|
||||
|
|
@ -1278,8 +1293,10 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
kwargs_for_followup: Final = self._prepare_followup_kwargs(kwargs)
|
||||
|
||||
if logging_obj is not None:
|
||||
agentic_params: Final[AgenticLoopParams] = logging_obj.model_call_details.get("agentic_loop_params", {})
|
||||
full_model_name = agentic_params.get("model", model)
|
||||
agentic_view: Final[_AgenticLoopParamsView] = {
|
||||
"agentic_loop_params": logging_obj.model_call_details.get("agentic_loop_params", {})
|
||||
}
|
||||
full_model_name = agentic_view["agentic_loop_params"].get("model", model)
|
||||
verbose_logger.debug(
|
||||
"WebSearchInterception: Built anthropic request patch [call_id=%s model=%s messages=%d searches=%d]",
|
||||
_call_id,
|
||||
|
|
@ -1676,7 +1693,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
@staticmethod
|
||||
def initialize_from_proxy_config(
|
||||
litellm_settings: dict[str, Any],
|
||||
callback_specific_params: dict[str, Any],
|
||||
callback_specific_params: Mapping[str, object],
|
||||
) -> "WebSearchInterceptionLogger":
|
||||
"""
|
||||
Static method to initialize WebSearchInterceptionLogger from proxy config.
|
||||
|
|
@ -1700,7 +1717,10 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
# Get websearch_interception_params from litellm_settings or callback_specific_params
|
||||
websearch_params: WebSearchInterceptionConfig = {}
|
||||
if "websearch_interception_params" in litellm_settings:
|
||||
websearch_params = litellm_settings["websearch_interception_params"]
|
||||
settings_view: Final[_WebSearchSettingsView] = {
|
||||
"websearch_interception_params": litellm_settings["websearch_interception_params"]
|
||||
}
|
||||
websearch_params = settings_view["websearch_interception_params"]
|
||||
elif "websearch_interception" in callback_specific_params and isinstance(
|
||||
callback_specific_params["websearch_interception"], dict
|
||||
):
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ This module has no dependencies on proxy code and can be safely imported at the
|
|||
import json
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -71,7 +72,7 @@ def get_litellm_gateway_api_key(
|
|||
return token_data["key"]
|
||||
|
||||
|
||||
def is_cli_token_fresh(token_data: dict, buffer_hours: float = 0.1) -> bool:
|
||||
def is_cli_token_fresh(token_data: Mapping[str, object], buffer_hours: float = 0.1) -> bool:
|
||||
"""Check whether a cached CLI token (as stored in token.json) is still
|
||||
within its expiration window. Used by `lite auth print-token` to fail
|
||||
fast, without a network round trip, once the cached token is past
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import subprocess
|
|||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping
|
||||
from datetime import datetime as dt_object
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
|
||||
|
|
@ -25,7 +25,15 @@ from litellm import (
|
|||
log_raw_request_response,
|
||||
turn_off_message_logging,
|
||||
)
|
||||
from litellm._logging import _is_debugging_on, _redact_string, verbose_logger
|
||||
from litellm._logging import (
|
||||
_is_debugging_on,
|
||||
_redact_string,
|
||||
session_id_var,
|
||||
set_session_id,
|
||||
set_trace_id,
|
||||
trace_id_var,
|
||||
verbose_logger,
|
||||
)
|
||||
from litellm._uuid import uuid
|
||||
from litellm.batches.batch_utils import _handle_completed_batch
|
||||
from litellm.caching.caching import DualCache, InMemoryCache
|
||||
|
|
@ -313,6 +321,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
applied_guardrails: list[str] | None = None,
|
||||
kwargs: dict | None = None,
|
||||
log_raw_request_response: bool = False,
|
||||
supports_correlation_logging: bool = True,
|
||||
):
|
||||
_input: Final[str | None] = messages # save original value of messages
|
||||
if messages is not None:
|
||||
|
|
@ -338,6 +347,36 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.call_type = call_type
|
||||
self.litellm_call_id = litellm_call_id
|
||||
self.litellm_trace_id: str = litellm_trace_id if litellm_trace_id else str(uuid.uuid4())
|
||||
|
||||
# Capture the pre-call *value* (not a contextvars.Token) so restoration works
|
||||
# even if this attempt's own logging ends up dispatched onto a different
|
||||
# asyncio Task/context (e.g. via asyncio.create_task or the logging worker) -
|
||||
# a Token can only be reset in the exact Context where it was created.
|
||||
self._pre_call_trace_id: str = trace_id_var.get()
|
||||
self._pre_call_session_id: str = session_id_var.get()
|
||||
_sid: Final = kwargs.get("litellm_session_id") if kwargs else None
|
||||
self.litellm_session_id: str = str(_sid) if _sid else ""
|
||||
# supports_correlation_logging is False for calls originating from the
|
||||
# sync client entry point (wrapper() in utils.py): a plain OS thread
|
||||
# has no per-call context isolation the way an asyncio Task does, and
|
||||
# a thread pool's worker threads are recycled across unrelated
|
||||
# requests, so stamping trace_id/session_id there risks one request's
|
||||
# ids leaking into a different, later request on the same thread. Sync
|
||||
# support is deferred to a follow-up PR with its own safe-restore
|
||||
# mechanism; async calls (the proxy's only call path) are unaffected.
|
||||
if supports_correlation_logging:
|
||||
set_trace_id(self.litellm_trace_id)
|
||||
set_session_id(self.litellm_session_id)
|
||||
# set_trace_id()/set_session_id() sanitize (strip control chars, bound
|
||||
# length) before storing, so the contextvar's actual value can differ
|
||||
# from self.litellm_trace_id/litellm_session_id. Capture what was
|
||||
# really stored - _restore_correlation_context_if_unclaimed() must
|
||||
# compare against this, not the raw ids, or a caller-supplied id
|
||||
# containing control characters/oversized input would never match
|
||||
# and cleanup would be skipped forever.
|
||||
self._own_trace_id: str = trace_id_var.get()
|
||||
self._own_session_id: str = session_id_var.get()
|
||||
|
||||
self.function_id = function_id
|
||||
self.streaming_chunks: list[Any] = [] # for generating complete stream response
|
||||
self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response
|
||||
|
|
@ -1992,7 +2031,67 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if complete_streaming_response is not None:
|
||||
await self.async_success_handler(result=complete_streaming_response)
|
||||
|
||||
def success_handler(self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs):
|
||||
def _restore_correlation_context(self) -> None:
|
||||
"""Restore trace_id/session_id contextvars to their pre-call value.
|
||||
|
||||
Without this, a nested LiteLLM call sharing the same asyncio Task as an
|
||||
outer request (e.g. a guardrail's own LLM-as-judge call, an MCP sampling
|
||||
call) would leave the outer request's subsequent log lines stamped with
|
||||
the nested call's trace_id/session_id instead of its own.
|
||||
|
||||
Uses a plain set() of the captured pre-call value rather than
|
||||
contextvars.Token-based reset(), since this can end up called from a
|
||||
different asyncio Task/context than __init__ ran in (e.g. the request
|
||||
task's own wrapper() finally block, plus async_success_handler
|
||||
dispatched separately via asyncio.create_task/the logging worker) -
|
||||
reset() only works in the exact Context a Token was created in and
|
||||
raises otherwise. Deliberately NOT idempotent/guarded: each distinct
|
||||
Task that calls this needs its own restore to actually take effect in
|
||||
that Task's view of the contextvars, so calling it multiple times
|
||||
(once per Task involved in this attempt) is required, not just safe.
|
||||
"""
|
||||
set_trace_id(self._pre_call_trace_id)
|
||||
set_session_id(self._pre_call_session_id)
|
||||
|
||||
def _restore_correlation_context_if_unclaimed(self) -> None:
|
||||
"""Guarded variant for __del__-triggered cleanup only.
|
||||
|
||||
__del__ can fire arbitrarily late (delayed by cyclic GC, possibly
|
||||
after the consuming Task/thread has already moved on to a different,
|
||||
still-active call). Unconditionally restoring in that case would
|
||||
stomp the active call's trace_id/session_id with this abandoned
|
||||
stream's stale pre-call snapshot. Only restore if the contextvars
|
||||
still hold the ids *this* call set - i.e. nothing has claimed them
|
||||
since - so an unrelated active call is never overwritten.
|
||||
"""
|
||||
if trace_id_var.get() == self._own_trace_id and session_id_var.get() == self._own_session_id:
|
||||
self._restore_correlation_context()
|
||||
|
||||
def success_handler(
|
||||
self,
|
||||
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
|
||||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
cache_hit: bool | None = None,
|
||||
**kwargs: Any, # kwargs-ok: forwarded to _success_handler_body
|
||||
) -> None:
|
||||
"""Restores trace_id/session_id contextvars once this attempt's own success
|
||||
logging (including any nested calls its callbacks trigger) is fully done."""
|
||||
try:
|
||||
return self._success_handler_body(
|
||||
result=result, start_time=start_time, end_time=end_time, cache_hit=cache_hit, **kwargs
|
||||
)
|
||||
finally:
|
||||
self._restore_correlation_context()
|
||||
|
||||
def _success_handler_body(
|
||||
self,
|
||||
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
|
||||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
cache_hit: bool | None = None,
|
||||
**kwargs: Any, # kwargs-ok: forwarded from success_handler
|
||||
) -> None:
|
||||
verbose_logger.debug("Logging Details LiteLLM-Success Call: Cache_hit=%s", cache_hit)
|
||||
if not self.should_run_logging(event_type="sync_success"): # prevent double logging
|
||||
return
|
||||
|
|
@ -2399,7 +2498,31 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
e,
|
||||
)
|
||||
|
||||
async def async_success_handler(self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs):
|
||||
async def async_success_handler(
|
||||
self,
|
||||
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
|
||||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
cache_hit: bool | None = None,
|
||||
**kwargs: Any, # kwargs-ok: forwarded to _async_success_handler_body
|
||||
) -> None:
|
||||
"""Restores trace_id/session_id contextvars once this attempt's own success
|
||||
logging (including any nested calls its callbacks trigger) is fully done."""
|
||||
try:
|
||||
return await self._async_success_handler_body(
|
||||
result=result, start_time=start_time, end_time=end_time, cache_hit=cache_hit, **kwargs
|
||||
)
|
||||
finally:
|
||||
self._restore_correlation_context()
|
||||
|
||||
async def _async_success_handler_body(
|
||||
self,
|
||||
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
|
||||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
cache_hit: bool | None = None,
|
||||
**kwargs: Any, # kwargs-ok: forwarded from async_success_handler
|
||||
) -> None:
|
||||
"""
|
||||
Implementing async callbacks, to handle asyncio event loop issues when custom integrations need to use async functions.
|
||||
"""
|
||||
|
|
@ -2791,7 +2914,32 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
kwargs=self.model_call_details,
|
||||
)
|
||||
|
||||
def failure_handler(self, exception, traceback_exception, start_time=None, end_time=None):
|
||||
def failure_handler(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
) -> None:
|
||||
"""Restores trace_id/session_id contextvars once this attempt's own failure
|
||||
logging (including any nested calls its callbacks trigger) is fully done."""
|
||||
try:
|
||||
return self._failure_handler_body(
|
||||
exception=exception,
|
||||
traceback_exception=traceback_exception,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
finally:
|
||||
self._restore_correlation_context()
|
||||
|
||||
def _failure_handler_body(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
) -> None:
|
||||
verbose_logger.debug("Logging Details LiteLLM-Failure Call: %s", litellm.failure_callback)
|
||||
if not self.should_run_logging(event_type="sync_failure"): # prevent double logging
|
||||
return
|
||||
|
|
@ -2960,7 +3108,32 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while failure logging %s", e
|
||||
)
|
||||
|
||||
async def async_failure_handler(self, exception, traceback_exception, start_time=None, end_time=None):
|
||||
async def async_failure_handler(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
) -> None:
|
||||
"""Restores trace_id/session_id contextvars once this attempt's own failure
|
||||
logging (including any nested calls its callbacks trigger) is fully done."""
|
||||
try:
|
||||
return await self._async_failure_handler_body(
|
||||
exception=exception,
|
||||
traceback_exception=traceback_exception,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
finally:
|
||||
self._restore_correlation_context()
|
||||
|
||||
async def _async_failure_handler_body(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Implementing async callbacks, to handle asyncio event loop issues when custom integrations need to use async functions.
|
||||
"""
|
||||
|
|
@ -5061,33 +5234,61 @@ class StandardLoggingPayloadSetup:
|
|||
return end_time_float - start_time_float
|
||||
|
||||
@staticmethod
|
||||
def _get_standard_logging_payload_trace_id(
|
||||
def get_standard_logging_payload_trace_id(
|
||||
logging_obj: Logging,
|
||||
litellm_params: dict,
|
||||
litellm_params: Mapping[str, Any],
|
||||
) -> str:
|
||||
"""
|
||||
Returns the `litellm_trace_id` for this request
|
||||
|
||||
This helps link sessions when multiple requests are made in a single session
|
||||
|
||||
Gated behind `litellm.request_correlation_in_logs`:
|
||||
- Off (default): legacy behavior, preserved for backward compatibility -
|
||||
`litellm_session_id` takes priority over `litellm_trace_id` since historically
|
||||
this field doubled as the session-grouping field.
|
||||
- On: `litellm_trace_id` takes priority - trace_id and session_id are independent,
|
||||
see `get_standard_logging_payload_session_id` for session tracking.
|
||||
"""
|
||||
dynamic_litellm_session_id: Final = litellm_params.get("litellm_session_id")
|
||||
dynamic_litellm_trace_id: Final = litellm_params.get("litellm_trace_id")
|
||||
metadata: Final = litellm_params.get("metadata")
|
||||
metadata_session_id: Final = metadata.get("session_id") if metadata else None
|
||||
metadata_trace_id: Final = metadata.get("trace_id") if metadata else None
|
||||
|
||||
# Note: we recommend using `litellm_session_id` for session tracking
|
||||
# `litellm_trace_id` is an internal litellm param
|
||||
ordered_candidates: Final[tuple[Any, Any, Any, Any]] = (
|
||||
(dynamic_litellm_trace_id, dynamic_litellm_session_id, metadata_trace_id, metadata_session_id)
|
||||
if litellm.request_correlation_in_logs
|
||||
else (dynamic_litellm_session_id, dynamic_litellm_trace_id, metadata_session_id, metadata_trace_id)
|
||||
)
|
||||
for candidate in ordered_candidates:
|
||||
if candidate:
|
||||
return str(candidate)
|
||||
return logging_obj.litellm_trace_id
|
||||
|
||||
@staticmethod
|
||||
def get_standard_logging_payload_session_id(
|
||||
logging_obj: Logging,
|
||||
litellm_params: Mapping[str, Any],
|
||||
) -> str:
|
||||
"""
|
||||
Returns the end-user/conversation `litellm_session_id` for this request, independent of trace_id.
|
||||
|
||||
Only populated when `litellm.request_correlation_in_logs` is enabled - off by default
|
||||
to avoid changing existing StandardLoggingPayload shape for callers who haven't opted in.
|
||||
Unlike `get_standard_logging_payload_trace_id`, this never falls back to a generated
|
||||
per-call trace id: it's empty when the caller never supplied a session id.
|
||||
"""
|
||||
if not litellm.request_correlation_in_logs:
|
||||
return ""
|
||||
dynamic_litellm_session_id: Final = litellm_params.get("litellm_session_id")
|
||||
if dynamic_litellm_session_id:
|
||||
return str(dynamic_litellm_session_id)
|
||||
elif dynamic_litellm_trace_id:
|
||||
return str(dynamic_litellm_trace_id)
|
||||
# Fallback: use metadata.session_id or metadata.trace_id for call chaining
|
||||
metadata: Final = litellm_params.get("metadata") or {}
|
||||
metadata_session_id: Final = metadata.get("session_id")
|
||||
metadata_trace_id: Final = metadata.get("trace_id")
|
||||
metadata: Final = litellm_params.get("metadata")
|
||||
metadata_session_id: Final = metadata.get("session_id") if metadata else None
|
||||
if metadata_session_id:
|
||||
return str(metadata_session_id)
|
||||
if metadata_trace_id:
|
||||
return str(metadata_trace_id)
|
||||
return logging_obj.litellm_trace_id
|
||||
return logging_obj.litellm_session_id
|
||||
|
||||
@staticmethod
|
||||
def _get_user_agent_tags(proxy_server_request: dict) -> list[str] | None:
|
||||
|
|
@ -5392,7 +5593,11 @@ def get_standard_logging_object_payload(
|
|||
payload: Final[StandardLoggingPayload] = StandardLoggingPayload(
|
||||
id=str(id),
|
||||
litellm_call_id=kwargs.get("litellm_call_id") or litellm_params.get("litellm_call_id"),
|
||||
trace_id=StandardLoggingPayloadSetup._get_standard_logging_payload_trace_id(
|
||||
trace_id=StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
||||
logging_obj=logging_obj,
|
||||
litellm_params=litellm_params,
|
||||
),
|
||||
session_id=StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
||||
logging_obj=logging_obj,
|
||||
litellm_params=litellm_params,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -49,6 +49,12 @@ _SERVICE_TIER_TO_COST_KEY_SUFFIX: Final[Mapping[str, str]] = MappingProxyType(
|
|||
}
|
||||
)
|
||||
|
||||
_INCLUSIVE_THRESHOLD_PROVIDERS: Final = frozenset({"xai"})
|
||||
|
||||
|
||||
def _uses_inclusive_token_thresholds(custom_llm_provider: str | None) -> bool:
|
||||
return custom_llm_provider in _INCLUSIVE_THRESHOLD_PROVIDERS
|
||||
|
||||
|
||||
def _get_token_detail_value(details: object, key: str) -> int | None:
|
||||
if isinstance(details, dict):
|
||||
|
|
@ -202,7 +208,11 @@ def _parse_above_token_threshold(key: str) -> float:
|
|||
|
||||
|
||||
def _get_token_base_cost(
|
||||
model_info: ModelInfo, usage: Usage, service_tier: str | None = None
|
||||
model_info: ModelInfo,
|
||||
usage: Usage,
|
||||
service_tier: str | None = None,
|
||||
*,
|
||||
threshold_is_inclusive: bool = False,
|
||||
) -> tuple[float, float, float, float, float]:
|
||||
"""
|
||||
Return prompt cost, completion cost, and cache costs for a given model and usage.
|
||||
|
|
@ -210,6 +220,9 @@ def _get_token_base_cost(
|
|||
If input_tokens > threshold and `input_cost_per_token_above_[x]k_tokens` or `input_cost_per_token_above_[x]_tokens` is set,
|
||||
then we use the corresponding threshold cost for all token types.
|
||||
|
||||
`threshold_is_inclusive` switches that comparison to >=, for providers such as xAI
|
||||
that bill the higher tier once the prompt reaches the threshold.
|
||||
|
||||
Returns:
|
||||
Tuple[float, float, float, float] - (prompt_cost, completion_cost, cache_creation_cost, cache_read_cost)
|
||||
"""
|
||||
|
|
@ -262,7 +275,7 @@ def _get_token_base_cost(
|
|||
# Handle both formats: _above_128k_tokens and _above_128_tokens
|
||||
threshold_str = key.split("_above_")[1].split("_tokens")[0]
|
||||
threshold = _parse_above_token_threshold(key)
|
||||
if usage.prompt_tokens > threshold:
|
||||
if usage.prompt_tokens > threshold or (threshold_is_inclusive and usage.prompt_tokens == threshold):
|
||||
# Prefer a service_tier-specific above-threshold key when available,
|
||||
# e.g. input_cost_per_token_priority_above_200k_tokens for Gemini
|
||||
# ON_DEMAND_PRIORITY. Falls back to the standard key automatically
|
||||
|
|
@ -777,7 +790,12 @@ def generic_cost_per_token(
|
|||
cache_creation_cost,
|
||||
cache_creation_cost_above_1hr,
|
||||
cache_read_cost,
|
||||
) = _get_token_base_cost(model_info=model_info, usage=usage, service_tier=service_tier)
|
||||
) = _get_token_base_cost(
|
||||
model_info=model_info,
|
||||
usage=usage,
|
||||
service_tier=service_tier,
|
||||
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
|
||||
)
|
||||
|
||||
prompt_cost = _calculate_input_cost(
|
||||
prompt_tokens_details=prompt_tokens_details,
|
||||
|
|
@ -909,7 +927,12 @@ def get_token_type_cost_breakdown(
|
|||
cache_creation_cost_rate,
|
||||
cache_creation_cost_above_1hr_rate,
|
||||
cache_read_cost_rate,
|
||||
) = _get_token_base_cost(model_info=model_info, usage=usage, service_tier=service_tier)
|
||||
) = _get_token_base_cost(
|
||||
model_info=model_info,
|
||||
usage=usage,
|
||||
service_tier=service_tier,
|
||||
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
|
||||
)
|
||||
|
||||
reasoning_tokens = (
|
||||
_parse_completion_tokens_details(usage)["reasoning_tokens"]
|
||||
|
|
@ -996,9 +1019,13 @@ def calculate_image_response_cost_from_usage(
|
|||
input_tokens_details: Final = getattr(usage, "input_tokens_details", None)
|
||||
prompt_tokens_details: PromptTokensDetailsWrapper | None = None
|
||||
if input_tokens_details is not None:
|
||||
# input_tokens_details may be a dict (e.g. OpenAI image edit responses)
|
||||
# or an object; read it tolerantly like the output side below, so image
|
||||
# input tokens are priced at input_cost_per_image_token instead of
|
||||
# silently falling back to the text rate.
|
||||
prompt_tokens_details = PromptTokensDetailsWrapper(
|
||||
text_tokens=getattr(input_tokens_details, "text_tokens", None),
|
||||
image_tokens=getattr(input_tokens_details, "image_tokens", None),
|
||||
text_tokens=_get_token_detail_value(input_tokens_details, "text_tokens"),
|
||||
image_tokens=_get_token_detail_value(input_tokens_details, "image_tokens"),
|
||||
cached_tokens=0,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -213,7 +213,75 @@ class CustomStreamWrapper:
|
|||
def __aiter__(self) -> AsyncIterator["ModelResponseStream"]:
|
||||
return self
|
||||
|
||||
def _restore_consumer_correlation_context(self, *, guarded: bool = False) -> None:
|
||||
"""Restore trace_id/session_id in the *consuming* thread/task/context.
|
||||
|
||||
wrapper_async() deliberately skips restoring correlation context when
|
||||
it returns a stream, so log lines emitted while the caller iterates it
|
||||
still carry this call's ids (see request_correlation_in_logs).
|
||||
wrapper() (the sync path) never stamps anything in the first place -
|
||||
see Logging.__init__'s supports_correlation_logging - so this method
|
||||
is an inert no-op for sync-created streams, harmless to call anyway
|
||||
since the class is shared between __next__ and __anext__.
|
||||
But the terminal success/failure handlers this stream dispatches to
|
||||
finish the job run on a *different* Task/thread (asyncio.create_task,
|
||||
threading.Thread, or the shared executor) - restoring there fixes up
|
||||
that detached context, not the one actually running the caller's
|
||||
`for`/`async for` loop. Call this at every point control genuinely
|
||||
returns to that consuming context: natural exhaustion (StopIteration/
|
||||
StopAsyncIteration), a raised failure, or explicit aclose(). Never let
|
||||
this raise - it must not break the caller's actual stream handling.
|
||||
|
||||
guarded=True (only __del__ uses this) skips the restore unless the
|
||||
contextvars still hold the ids this stream's own call set, so a
|
||||
delayed finalizer never overwrites a different, still-active call
|
||||
that has since taken over the same Task/thread's context.
|
||||
"""
|
||||
try:
|
||||
logging_obj: Final = getattr(self, "logging_obj", None)
|
||||
if logging_obj is None:
|
||||
return
|
||||
method_name: Final = (
|
||||
"_restore_correlation_context_if_unclaimed" if guarded else "_restore_correlation_context"
|
||||
)
|
||||
restore: Final = getattr(logging_obj, method_name, None)
|
||||
if restore is not None:
|
||||
restore()
|
||||
except Exception as restore_error: # noqa: BLE001 # best-effort cleanup; must not raise into the caller
|
||||
verbose_logger.debug("could not restore correlation context: %s", restore_error)
|
||||
|
||||
def __del__(self) -> None:
|
||||
"""Best-effort correlation-context cleanup for an abandoned async stream.
|
||||
|
||||
Only meaningfully applies to streams created by wrapper_async(): it
|
||||
leaves contextvars "open" across the caller's iteration, so if the
|
||||
caller never fully consumes the stream - stops early, drops the
|
||||
reference, cancels it - none of the exit points
|
||||
_restore_consumer_correlation_context() is called from ever run. For a
|
||||
sync stream (wrapper()), this is a no-op in practice: wrapper() never
|
||||
stamps trace_id/session_id for sync calls in the first place (see
|
||||
Logging.__init__'s supports_correlation_logging), so there is nothing
|
||||
for this to clean up.
|
||||
|
||||
This is a best-effort fallback, not a guarantee: __del__ timing is
|
||||
unpredictable (delayed by cyclic GC, not guaranteed at interpreter
|
||||
shutdown, and may run on a different thread), so this can only reduce
|
||||
how long the leak persists, not eliminate it. That's an acceptable
|
||||
trade specifically because its blast radius is bounded to the one
|
||||
asyncio Task this stream's own call ran in - each async call has its
|
||||
own copy of the contextvars, and Tasks (unlike a thread pool's worker
|
||||
threads) are never recycled across requests, so a delayed or missed
|
||||
cleanup here can never misattribute a *different* request's logs.
|
||||
guarded=True additionally ensures it never clobbers a different,
|
||||
still-active call's context within that same Task if this fires late.
|
||||
"""
|
||||
self._restore_consumer_correlation_context(guarded=True)
|
||||
|
||||
async def aclose(self):
|
||||
# Restore the consumer's outer context only after the underlying
|
||||
# provider stream's own close (and its diagnostic logging below, if
|
||||
# closing fails) completes - not before - so those log lines still
|
||||
# carry this closing stream's own trace_id/session_id.
|
||||
if self.completion_stream is not None:
|
||||
stream_to_close: Final = self.completion_stream
|
||||
self.completion_stream = None
|
||||
|
|
@ -233,6 +301,7 @@ class CustomStreamWrapper:
|
|||
"CustomStreamWrapper.aclose: error closing completion_stream: %s",
|
||||
e,
|
||||
)
|
||||
self._restore_consumer_correlation_context()
|
||||
|
||||
def check_send_stream_usage(self, stream_options: dict | None):
|
||||
return stream_options is not None and stream_options.get("include_usage", False) is True
|
||||
|
|
@ -1839,6 +1908,7 @@ class CustomStreamWrapper:
|
|||
if self.sent_stream_usage is False and self.send_stream_usage is True:
|
||||
self.sent_stream_usage = True
|
||||
return response
|
||||
self._restore_consumer_correlation_context()
|
||||
raise # Re-raise StopIteration
|
||||
else:
|
||||
self.sent_last_chunk = True
|
||||
|
|
@ -1852,6 +1922,19 @@ class CustomStreamWrapper:
|
|||
processed_chunk,
|
||||
cache_hit,
|
||||
) # log response
|
||||
# Deliberately do NOT restore context here even though
|
||||
# completion_stream is already exhausted: this chunk is still
|
||||
# real data belonging to this call, and the caller's own
|
||||
# (application-level) log statements processing it run
|
||||
# immediately after this return, in this same synchronous
|
||||
# frame - restoring first would make those lines carry the
|
||||
# wrong ids, which is exactly what leaving context open during
|
||||
# iteration is meant to prevent (see
|
||||
# _restore_consumer_correlation_context's docstring). A caller
|
||||
# that keeps iterating gets cleaned up on its next __next__()
|
||||
# call (immediate StopIteration, handled above); one that
|
||||
# stops right here relies on aclose() or the best-effort
|
||||
# __del__ guard instead.
|
||||
return processed_chunk
|
||||
except Exception as e:
|
||||
traceback_exception: Final = traceback.format_exc()
|
||||
|
|
@ -1879,8 +1962,12 @@ class CustomStreamWrapper:
|
|||
cache_hit = False
|
||||
if self.custom_llm_provider is not None and self.custom_llm_provider == "cached_response":
|
||||
cache_hit = True
|
||||
self._check_max_streaming_duration()
|
||||
try:
|
||||
# Inside the try (not before it) so a raised litellm.Timeout flows
|
||||
# through the same except Exception -> _handle_stream_fallback_error
|
||||
# path as every other failure, restoring the consumer's correlation
|
||||
# context - a check before the try would bypass that entirely.
|
||||
self._check_max_streaming_duration()
|
||||
if self.completion_stream is None:
|
||||
await self.fetch_stream()
|
||||
|
||||
|
|
@ -2083,10 +2170,17 @@ class CustomStreamWrapper:
|
|||
)
|
||||
)
|
||||
|
||||
self._restore_consumer_correlation_context()
|
||||
raise StopAsyncIteration # Re-raise StopIteration
|
||||
else:
|
||||
self.sent_last_chunk = True
|
||||
processed_chunk: Final = self.finish_reason_handler()
|
||||
# see sync __next__'s sibling branch: deliberately do NOT restore
|
||||
# here - this chunk is still this call's own data, and restoring
|
||||
# before returning it would corrupt the caller's own log
|
||||
# statements processing it. A caller that keeps iterating gets
|
||||
# cleaned up on the next __anext__() call; one that stops here
|
||||
# relies on aclose() or the best-effort __del__ guard.
|
||||
return processed_chunk
|
||||
|
||||
def _log_stream_failure_and_raise(self, e: Exception) -> NoReturn:
|
||||
|
|
@ -2138,7 +2232,12 @@ class CustomStreamWrapper:
|
|||
"""
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
# Map to OpenAI exception format
|
||||
# Map to OpenAI exception format. Some providers' mappers (e.g.
|
||||
# _map_anthropic_exception, _map_aleph_alpha_exception) synchronously
|
||||
# log a debug diagnostic (the raw status code) as part of mapping -
|
||||
# restore the consumer's outer context only after this completes, so
|
||||
# that diagnostic log line still carries the failing stream's own
|
||||
# trace_id/session_id instead of the consumer's (or an empty one).
|
||||
if isinstance(e, OpenAIError):
|
||||
mapped_exception: Exception = e
|
||||
else:
|
||||
|
|
@ -2152,6 +2251,7 @@ class CustomStreamWrapper:
|
|||
)
|
||||
except Exception as mapping_error:
|
||||
mapped_exception = mapping_error
|
||||
self._restore_consumer_correlation_context()
|
||||
|
||||
def _normalize_status_code(exc: Exception) -> int | None:
|
||||
"""Best-effort status_code extraction."""
|
||||
|
|
|
|||
|
|
@ -1,12 +1,13 @@
|
|||
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Final,
|
||||
TypeAlias,
|
||||
cast,
|
||||
)
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.asyncify import run_async_function
|
||||
|
|
@ -39,6 +40,11 @@ _AnthropicSystem: TypeAlias = "str | list[dict[str, object]] | None"
|
|||
_ContextManagementSpec: TypeAlias = "dict[str, object] | list[dict[str, object]] | None"
|
||||
|
||||
|
||||
class _CompletionKwargs(TypedDict, total=False, extra_items=object):
|
||||
model: str
|
||||
custom_llm_provider: str
|
||||
|
||||
|
||||
def _messages_have_compaction_block(messages: _AnthropicMessages) -> bool:
|
||||
"""Return True when any message carries a ``compaction`` content block."""
|
||||
for msg in messages:
|
||||
|
|
@ -312,7 +318,7 @@ ANTHROPIC_ADAPTER: Final = AnthropicAdapter()
|
|||
class LiteLLMMessagesToCompletionTransformationHandler:
|
||||
@staticmethod
|
||||
def _route_openai_thinking_to_responses_api_if_needed(
|
||||
completion_kwargs: dict[str, Any],
|
||||
completion_kwargs: _CompletionKwargs,
|
||||
*,
|
||||
thinking: Mapping[str, object] | None,
|
||||
) -> None:
|
||||
|
|
@ -377,7 +383,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
|
||||
@staticmethod
|
||||
def _normalize_reasoning_effort(
|
||||
completion_kwargs: dict[str, Any],
|
||||
completion_kwargs: _CompletionKwargs,
|
||||
) -> None:
|
||||
"""
|
||||
Normalize reasoning_effort values based on target model capabilities.
|
||||
|
|
@ -393,7 +399,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
if reasoning_effort is None:
|
||||
return
|
||||
|
||||
model: Final = cast(str, completion_kwargs.get("model", ""))
|
||||
model: Final = completion_kwargs.get("model", "")
|
||||
custom_llm_provider: Final = completion_kwargs.get("custom_llm_provider")
|
||||
|
||||
if isinstance(reasoning_effort, str):
|
||||
|
|
@ -417,19 +423,19 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
max_tokens: int,
|
||||
messages: _AnthropicMessages,
|
||||
model: str,
|
||||
metadata: dict | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
stop_sequences: list[str] | None = None,
|
||||
stream: bool | None = False,
|
||||
system: _AnthropicSystem = None,
|
||||
temperature: float | None = None,
|
||||
thinking: dict | None = None,
|
||||
tool_choice: dict | None = None,
|
||||
tools: list[dict] | None = None,
|
||||
thinking: dict[str, object] | None = None,
|
||||
tool_choice: dict[str, object] | None = None,
|
||||
tools: list[dict[str, object]] | None = None,
|
||||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: dict | None = None,
|
||||
output_format: dict[str, object] | None = None,
|
||||
extra_kwargs: Mapping[str, object] | None = None,
|
||||
) -> tuple[dict[str, Any], dict[str, str]]:
|
||||
) -> tuple[_CompletionKwargs, dict[str, str]]:
|
||||
"""Prepare kwargs for litellm.completion/acompletion.
|
||||
|
||||
Returns:
|
||||
|
|
@ -486,7 +492,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
if openai_request is None:
|
||||
raise ValueError("Failed to translate request to OpenAI format")
|
||||
|
||||
completion_kwargs: Final[dict[str, Any]] = dict(openai_request)
|
||||
completion_kwargs: Final[_CompletionKwargs] = {**openai_request}
|
||||
|
||||
if stream:
|
||||
completion_kwargs["stream"] = stream
|
||||
|
|
@ -538,17 +544,17 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
max_tokens: int,
|
||||
messages: _AnthropicMessages,
|
||||
model: str,
|
||||
metadata: dict | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
stop_sequences: list[str] | None = None,
|
||||
stream: bool | None = False,
|
||||
system: str | None = None,
|
||||
temperature: float | None = None,
|
||||
thinking: dict | None = None,
|
||||
tool_choice: dict | None = None,
|
||||
thinking: dict[str, object] | None = None,
|
||||
tool_choice: dict[str, object] | None = None,
|
||||
tools: list[dict[str, object]] | None = None,
|
||||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: dict | None = None,
|
||||
output_format: dict[str, object] | None = None,
|
||||
**kwargs,
|
||||
) -> AnthropicMessagesResponse | AsyncIterator[bytes] | Iterator[bytes]:
|
||||
"""Handle non-Anthropic models asynchronously using the adapter"""
|
||||
|
|
@ -625,17 +631,17 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
max_tokens: int,
|
||||
messages: _AnthropicMessages,
|
||||
model: str,
|
||||
metadata: dict | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
stop_sequences: list[str] | None = None,
|
||||
stream: bool | None = False,
|
||||
system: str | None = None,
|
||||
temperature: float | None = None,
|
||||
thinking: dict | None = None,
|
||||
tool_choice: dict | None = None,
|
||||
thinking: dict[str, object] | None = None,
|
||||
tool_choice: dict[str, object] | None = None,
|
||||
tools: list[dict[str, object]] | None = None,
|
||||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: dict | None = None,
|
||||
output_format: dict[str, object] | None = None,
|
||||
_is_async: bool = False,
|
||||
**kwargs,
|
||||
) -> (
|
||||
|
|
|
|||
|
|
@ -7,8 +7,8 @@ tool through a ``tool_use`` content block, and results are fed back as
|
|||
``tool_result`` blocks in a user message.
|
||||
"""
|
||||
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from typing import Any, Final
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Mapping, Sequence
|
||||
from typing import Any, Final, NamedTuple
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.responses.mcp.request_context import MCPRequestContext
|
||||
|
|
@ -24,14 +24,18 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
|
|||
MAX_MCP_TOOL_USE_ITERATIONS: Final = 10
|
||||
|
||||
|
||||
def _get_response_content(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]:
|
||||
class _AnthropicMessagesCall(NamedTuple):
|
||||
fn: Callable[..., Awaitable[AnthropicMessagesResponse | Iterator[bytes] | AsyncIterator[object]]]
|
||||
|
||||
|
||||
def _get_response_content(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, object]]:
|
||||
content: Final = response.get("content")
|
||||
if not isinstance(content, list):
|
||||
return ()
|
||||
return tuple(block for block in content if isinstance(block, dict))
|
||||
|
||||
|
||||
def _extract_tool_use_blocks(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]:
|
||||
def _extract_tool_use_blocks(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, object]]:
|
||||
"""Return the ``tool_use`` content blocks the model emitted."""
|
||||
return tuple(block for block in _get_response_content(response) if block.get("type") == "tool_use")
|
||||
|
||||
|
|
@ -41,7 +45,7 @@ def _get_stop_reason(response: AnthropicMessagesResponse) -> str | None:
|
|||
return stop_reason if isinstance(stop_reason, str) else None
|
||||
|
||||
|
||||
def _build_tool_result_message(tool_results: Sequence[Mapping[str, Any]]) -> AnthropicMessagesUserMessageParam:
|
||||
def _build_tool_result_message(tool_results: Sequence[Mapping[str, object]]) -> AnthropicMessagesUserMessageParam:
|
||||
"""Turn executed tool results into the user message Anthropic expects."""
|
||||
return AnthropicMessagesUserMessageParam(
|
||||
role="user",
|
||||
|
|
@ -58,11 +62,11 @@ def _build_tool_result_message(tool_results: Sequence[Mapping[str, Any]]) -> Ant
|
|||
|
||||
async def anthropic_messages_with_mcp(
|
||||
max_tokens: int,
|
||||
messages: Sequence[Mapping[str, Any]],
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
model: str,
|
||||
tools: Sequence[Mapping[str, Any]] | None = None,
|
||||
tools: Sequence[Mapping[str, object]] | None = None,
|
||||
**kwargs: Any, # kwargs-ok: forwarded verbatim to litellm.anthropic_messages, which owns the param contract
|
||||
) -> AnthropicMessagesResponse | AsyncIterator[Any]:
|
||||
) -> AnthropicMessagesResponse | Iterator[bytes] | AsyncIterator[object]:
|
||||
"""
|
||||
Expand litellm_proxy MCP references for `/v1/messages` and run the tool loop.
|
||||
|
||||
|
|
@ -81,7 +85,7 @@ async def anthropic_messages_with_mcp(
|
|||
mcp_references, other_tools = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
|
||||
|
||||
if not mcp_references:
|
||||
return await litellm.anthropic_messages(
|
||||
return await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(
|
||||
max_tokens=max_tokens,
|
||||
messages=list(messages),
|
||||
model=model,
|
||||
|
|
@ -114,7 +118,7 @@ async def anthropic_messages_with_mcp(
|
|||
)
|
||||
stream: Final = bool(kwargs.pop("stream", False))
|
||||
|
||||
base_call_args: Final[Mapping[str, Any]] = {
|
||||
base_call_args: Final[Mapping[str, object]] = {
|
||||
"max_tokens": max_tokens,
|
||||
"model": model,
|
||||
"tools": all_tools or None,
|
||||
|
|
@ -123,10 +127,12 @@ async def anthropic_messages_with_mcp(
|
|||
}
|
||||
|
||||
if not should_auto_execute:
|
||||
return await litellm.anthropic_messages(messages=list(messages), stream=stream, **base_call_args)
|
||||
return await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(
|
||||
messages=list(messages), stream=stream, **base_call_args
|
||||
)
|
||||
|
||||
working_messages: Sequence[Mapping[str, Any]] = tuple(messages)
|
||||
response: AnthropicMessagesResponse = await litellm.anthropic_messages(
|
||||
working_messages: Sequence[Mapping[str, object]] = tuple(messages)
|
||||
response: AnthropicMessagesResponse = await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(
|
||||
messages=list(working_messages), stream=False, **base_call_args
|
||||
)
|
||||
|
||||
|
|
@ -161,7 +167,9 @@ async def anthropic_messages_with_mcp(
|
|||
{"role": "assistant", "content": list(_get_response_content(response))},
|
||||
_build_tool_result_message(tool_results),
|
||||
)
|
||||
response = await litellm.anthropic_messages(messages=list(working_messages), stream=False, **base_call_args)
|
||||
response = await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(
|
||||
messages=list(working_messages), stream=False, **base_call_args
|
||||
)
|
||||
else:
|
||||
verbose_logger.warning(
|
||||
"MCP tool loop hit its %s iteration cap for model %s; returning the last response",
|
||||
|
|
|
|||
|
|
@ -8,7 +8,12 @@ from collections.abc import AsyncIterator, Coroutine
|
|||
from typing import Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm.types.llms.anthropic import AnthropicMessagesRequest
|
||||
from litellm.types.llms.anthropic import (
|
||||
AllAnthropicToolsValues,
|
||||
AnthropicMessagesRequest,
|
||||
AnthropicOutputConfig,
|
||||
AnthropicOutputSchema,
|
||||
)
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
|
|
@ -27,24 +32,24 @@ def _build_responses_kwargs(
|
|||
model: str,
|
||||
context_management: dict | None = None,
|
||||
metadata: dict | None = None,
|
||||
output_config: dict | None = None,
|
||||
output_config: AnthropicOutputConfig | None = None,
|
||||
stop_sequences: list[str] | None = None,
|
||||
stream: bool | None = False,
|
||||
system: str | None = None,
|
||||
temperature: float | None = None,
|
||||
thinking: dict | None = None,
|
||||
tool_choice: dict | None = None,
|
||||
tools: list[dict] | None = None,
|
||||
tools: list[AllAnthropicToolsValues | dict] | None = None,
|
||||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: dict | None = None,
|
||||
output_format: AnthropicOutputSchema | None = None,
|
||||
extra_kwargs: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Build the kwargs dict to pass directly to litellm.responses() / litellm.aresponses().
|
||||
"""
|
||||
# Build a typed AnthropicMessagesRequest for the adapter
|
||||
request_data: Final[dict[str, Any]] = {
|
||||
request_data: Final[AnthropicMessagesRequest] = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"max_tokens": max_tokens,
|
||||
|
|
@ -128,19 +133,19 @@ class LiteLLMMessagesToResponsesAPIHandler:
|
|||
model: str,
|
||||
context_management: dict | None = None,
|
||||
metadata: dict | None = None,
|
||||
output_config: dict | None = None,
|
||||
output_config: AnthropicOutputConfig | None = None,
|
||||
stop_sequences: list[str] | None = None,
|
||||
stream: bool | None = False,
|
||||
system: str | None = None,
|
||||
temperature: float | None = None,
|
||||
thinking: dict | None = None,
|
||||
tool_choice: dict | None = None,
|
||||
tools: list[dict] | None = None,
|
||||
tools: list[AllAnthropicToolsValues | dict] | None = None,
|
||||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: dict | None = None,
|
||||
output_format: AnthropicOutputSchema | None = None,
|
||||
**kwargs,
|
||||
) -> AnthropicMessagesResponse | AsyncIterator:
|
||||
) -> AnthropicMessagesResponse | AsyncIterator[bytes]:
|
||||
responses_kwargs: Final = _build_responses_kwargs(
|
||||
max_tokens=max_tokens,
|
||||
messages=messages,
|
||||
|
|
@ -179,23 +184,23 @@ class LiteLLMMessagesToResponsesAPIHandler:
|
|||
model: str,
|
||||
context_management: dict | None = None,
|
||||
metadata: dict | None = None,
|
||||
output_config: dict | None = None,
|
||||
output_config: AnthropicOutputConfig | None = None,
|
||||
stop_sequences: list[str] | None = None,
|
||||
stream: bool | None = False,
|
||||
system: str | None = None,
|
||||
temperature: float | None = None,
|
||||
thinking: dict | None = None,
|
||||
tool_choice: dict | None = None,
|
||||
tools: list[dict] | None = None,
|
||||
tools: list[AllAnthropicToolsValues | dict] | None = None,
|
||||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: dict | None = None,
|
||||
output_format: AnthropicOutputSchema | None = None,
|
||||
_is_async: bool = False,
|
||||
**kwargs,
|
||||
) -> (
|
||||
AnthropicMessagesResponse
|
||||
| AsyncIterator[Any]
|
||||
| Coroutine[Any, Any, AnthropicMessagesResponse | AsyncIterator[Any]]
|
||||
| AsyncIterator[bytes]
|
||||
| Coroutine[None, None, AnthropicMessagesResponse | AsyncIterator[bytes]]
|
||||
):
|
||||
if _is_async:
|
||||
return LiteLLMMessagesToResponsesAPIHandler.async_anthropic_messages_handler(
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -15,6 +15,8 @@ from collections.abc import Mapping, Sequence
|
|||
from typing import Any, Final, NamedTuple, Optional, Protocol, Union, runtime_checkable
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from fastapi import Request
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.shared.context import RequestContext
|
||||
|
|
@ -28,8 +30,9 @@ if typing.TYPE_CHECKING:
|
|||
ToolUseContent,
|
||||
)
|
||||
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter
|
||||
|
|
@ -1016,7 +1019,7 @@ async def _run_budget_checks(
|
|||
general_settings=general_settings or {},
|
||||
route="/chat/completions",
|
||||
llm_router=_llm_router,
|
||||
proxy_logging_obj=typing.cast("ProxyLogging", _proxy_logging_obj),
|
||||
proxy_logging_obj=_proxy_logging_obj,
|
||||
valid_token=user_api_key_auth,
|
||||
request=dummy_request,
|
||||
)
|
||||
|
|
@ -1176,15 +1179,19 @@ async def _build_completion_kwargs(
|
|||
)
|
||||
|
||||
|
||||
class _AcompletionCall(NamedTuple):
|
||||
fn: "Callable[..., Awaitable[ModelResponse | CustomStreamWrapper]]"
|
||||
|
||||
|
||||
async def _run_guardrails_and_call_llm(
|
||||
completion_kwargs: dict[str, Any],
|
||||
completion_kwargs: dict[str, object],
|
||||
user_api_key_auth: "UserAPIKeyAuth",
|
||||
) -> Any:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj as _plo
|
||||
|
||||
if _plo is not None:
|
||||
completion_kwargs = await typing.cast("ProxyLogging", _plo).pre_call_hook(
|
||||
completion_kwargs = await _plo.pre_call_hook(
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
data=completion_kwargs,
|
||||
call_type="acompletion",
|
||||
|
|
@ -1204,10 +1211,10 @@ async def _run_guardrails_and_call_llm(
|
|||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if llm_router is not None:
|
||||
return await llm_router.acompletion(**completion_kwargs)
|
||||
return await litellm.acompletion(**completion_kwargs)
|
||||
return await _AcompletionCall(fn=llm_router.acompletion).fn(**completion_kwargs)
|
||||
return await _AcompletionCall(fn=litellm.acompletion).fn(**completion_kwargs)
|
||||
except ImportError:
|
||||
return await litellm.acompletion(**completion_kwargs)
|
||||
return await _AcompletionCall(fn=litellm.acompletion).fn(**completion_kwargs)
|
||||
|
||||
|
||||
async def handle_sampling_create_message(
|
||||
|
|
|
|||
|
|
@ -4018,6 +4018,12 @@ class LitellmDataForBackendLLMCall(TypedDict, total=False):
|
|||
stream_timeout: float | None
|
||||
user: str | None
|
||||
num_retries: int | None
|
||||
# True when the effective timeout came from a caller-controlled source (the
|
||||
# `x-litellm-timeout`/`x-litellm-stream-timeout` headers, or a `timeout`/`request_timeout`/
|
||||
# `stream_timeout` field in the request body) rather than deployment config, so a
|
||||
# deliberately tiny value isn't treated as a deployment health signal (see
|
||||
# cooldown_handlers._trigger_cooldown_for_failed_deployment).
|
||||
client_side_timeout: bool
|
||||
|
||||
|
||||
class LitellmMetadataFromRequestHeaders(TypedDict, total=False):
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ The A2A SDK can point to LiteLLM's URL and invoke agents registered with LiteLLM
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncGenerator
|
||||
from collections.abc import AsyncGenerator, Mapping
|
||||
from copy import deepcopy
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from urllib.parse import urlparse
|
||||
|
|
@ -36,7 +36,7 @@ from litellm.proxy.agent_endpoints.databricks_oauth import (
|
|||
)
|
||||
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.utils import get_custom_url
|
||||
from litellm.proxy.utils import ProxyLogging, get_custom_url
|
||||
from litellm.types.utils import all_litellm_params
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -46,7 +46,7 @@ if TYPE_CHECKING:
|
|||
|
||||
router: Final = APIRouter()
|
||||
|
||||
_PASCAL_TO_WIRE: Final[dict[str, str]] = {
|
||||
_PASCAL_TO_WIRE: Final[Mapping[str, str]] = {
|
||||
"SendMessage": "message/send",
|
||||
"SendStreamingMessage": "message/stream",
|
||||
"GetTask": "tasks/get",
|
||||
|
|
@ -118,9 +118,9 @@ def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> dict[str, str
|
|||
|
||||
def _forwarding_headers(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_data: dict[str, Any],
|
||||
agent_extra_headers: dict[str, str] | None,
|
||||
) -> dict[str, str] | None:
|
||||
request_data: Mapping[str, object],
|
||||
agent_extra_headers: Mapping[str, str] | None,
|
||||
) -> Mapping[str, str] | None:
|
||||
sanitized: Final = (
|
||||
{k: v for k, v in agent_extra_headers.items() if not k.lower().startswith("x-litellm-")}
|
||||
if agent_extra_headers
|
||||
|
|
@ -136,7 +136,7 @@ def _forwarding_headers(
|
|||
|
||||
|
||||
def _jsonrpc_error(
|
||||
request_id: Any | None,
|
||||
request_id: object,
|
||||
code: int,
|
||||
message: str,
|
||||
status_code: int = 400,
|
||||
|
|
@ -162,7 +162,7 @@ def _get_agent(agent_id: str):
|
|||
return agent
|
||||
|
||||
|
||||
def _enforce_inbound_trace_id(agent: Any, request: Request) -> None:
|
||||
def _enforce_inbound_trace_id(agent: "AgentResponse", request: Request) -> None:
|
||||
"""Raise 400 if agent requires x-litellm-trace-id on inbound calls and it is missing."""
|
||||
agent_litellm_params: Final = agent.litellm_params or {}
|
||||
if not agent_litellm_params.get("require_trace_id_on_calls_to_agent"):
|
||||
|
|
@ -181,8 +181,8 @@ def _enforce_inbound_trace_id(agent: Any, request: Request) -> None:
|
|||
|
||||
async def _forward_jsonrpc(
|
||||
agent_url: str,
|
||||
body: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
body: dict[str, object],
|
||||
extra_headers: Mapping[str, str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
|
@ -205,11 +205,11 @@ async def _forward_jsonrpc(
|
|||
|
||||
async def _a2a_sse_event_source(
|
||||
agent_url: str,
|
||||
body: dict[str, Any],
|
||||
request_id: Any | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
body: Mapping[str, object],
|
||||
request_id: str | int | None = None,
|
||||
extra_headers: Mapping[str, str] | None = None,
|
||||
served_version: A2AVersion = "0.3",
|
||||
) -> AsyncGenerator[dict, None]:
|
||||
) -> AsyncGenerator[Mapping[str, object], None]:
|
||||
"""Stream an upstream A2A SSE response as parsed JSON-RPC event dicts.
|
||||
|
||||
Upstream HTTP/JSON-RPC errors are surfaced as a single JSON-RPC error event
|
||||
|
|
@ -234,7 +234,7 @@ async def _a2a_sse_event_source(
|
|||
try:
|
||||
if not resp.is_success:
|
||||
error_body: Final = await resp.aread()
|
||||
error_event: dict[str, Any] | None = None
|
||||
error_event: Mapping[str, object] | None = None
|
||||
try:
|
||||
parsed: Final = json.loads(error_body)
|
||||
if isinstance(parsed, dict) and "error" in parsed:
|
||||
|
|
@ -267,12 +267,12 @@ async def _a2a_sse_event_source(
|
|||
|
||||
async def _forward_jsonrpc_sse(
|
||||
agent_url: str,
|
||||
body: dict[str, Any],
|
||||
request_id: Any | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
proxy_logging_obj: Any | None = None,
|
||||
user_api_key_dict: Any | None = None,
|
||||
request_data: dict[str, Any] | None = None,
|
||||
body: Mapping[str, object],
|
||||
request_id: str | int | None = None,
|
||||
extra_headers: Mapping[str, str] | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
user_api_key_dict: UserAPIKeyAuth | None = None,
|
||||
request_data: dict[str, object] | None = None,
|
||||
served_version: A2AVersion = "0.3",
|
||||
) -> StreamingResponse:
|
||||
event_source: Final = _a2a_sse_event_source(
|
||||
|
|
@ -283,10 +283,10 @@ async def _forward_jsonrpc_sse(
|
|||
served_version=served_version,
|
||||
)
|
||||
|
||||
def _serialize_chunk(chunk: Any) -> str:
|
||||
def _serialize_chunk(chunk: object) -> str:
|
||||
return f"data: {json.dumps(chunk)}\n\n"
|
||||
|
||||
def _serialize_error(proxy_exc: Any) -> str:
|
||||
def _serialize_error(proxy_exc: object) -> str:
|
||||
return (
|
||||
"data: "
|
||||
+ json.dumps(
|
||||
|
|
@ -331,17 +331,17 @@ async def _forward_jsonrpc_sse(
|
|||
|
||||
async def _handle_stream_message(
|
||||
api_base: str | None,
|
||||
request_id: Any,
|
||||
params: dict[str, Any],
|
||||
litellm_params: dict[str, Any] | None = None,
|
||||
request_id: str | int,
|
||||
params: dict[str, object],
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
agent_id: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
proxy_server_request: dict[str, Any] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
proxy_server_request: dict[str, object] | None = None,
|
||||
*,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
user_api_key_dict: UserAPIKeyAuth | None = None,
|
||||
request_data: dict[str, Any] | None = None,
|
||||
proxy_logging_obj: Any | None = None,
|
||||
request_data: dict[str, object] | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
served_version: A2AVersion = "0.3",
|
||||
) -> StreamingResponse:
|
||||
"""Handle message/stream method via SDK functions.
|
||||
|
|
@ -430,7 +430,7 @@ async def _handle_stream_message(
|
|||
obj = normalize_stream_event(obj, served_version, request_id=request_id)
|
||||
return json.dumps(obj) + "\n"
|
||||
|
||||
def _ndjson_error(proxy_exc: Any) -> str:
|
||||
def _ndjson_error(proxy_exc: object) -> str:
|
||||
return (
|
||||
json.dumps(
|
||||
{
|
||||
|
|
@ -669,7 +669,7 @@ async def invoke_agent_a2a(
|
|||
agent_name: Final = agent_card_params.get("name", agent_id)
|
||||
|
||||
# Get litellm_params (may include custom_llm_provider for completion bridge)
|
||||
litellm_params = agent.litellm_params or {}
|
||||
litellm_params: dict[str, object] = agent.litellm_params or {}
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
|
||||
# Hand the authenticated key hash to the completion bridge so provider
|
||||
|
|
@ -725,7 +725,7 @@ async def invoke_agent_a2a(
|
|||
request_data = data
|
||||
|
||||
# Build merged headers for the backend agent
|
||||
static_headers: Final[dict[str, str]] = dict(agent.static_headers or {})
|
||||
static_headers: Final[Mapping[str, str]] = dict(agent.static_headers or {})
|
||||
|
||||
raw_headers: Final = dict(request.headers)
|
||||
normalized: Final = {k.lower(): v for k, v in raw_headers.items()}
|
||||
|
|
@ -893,7 +893,7 @@ async def invoke_agent_a2a(
|
|||
detail="Push notification URL must be a string",
|
||||
)
|
||||
_validate_push_notification_url(callback_url)
|
||||
forward_body = {
|
||||
forward_body: dict[str, object] = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"method": method,
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import click
|
|||
import requests
|
||||
from rich.console import Console
|
||||
from rich.table import Table
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
|
||||
from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh
|
||||
|
|
@ -18,6 +19,57 @@ from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh
|
|||
from .private_json import write_private_json
|
||||
|
||||
|
||||
class CliTokenData(TypedDict):
|
||||
base_url: str
|
||||
key: str
|
||||
user_id: str
|
||||
user_email: str
|
||||
user_role: str
|
||||
auth_header_name: str
|
||||
jwt_token: str
|
||||
timestamp: float
|
||||
|
||||
|
||||
class CliTeam(TypedDict, total=False):
|
||||
team_id: str | None
|
||||
team_alias: str | None
|
||||
models: list[str]
|
||||
max_budget: float | None
|
||||
|
||||
|
||||
class CliContextObj(TypedDict):
|
||||
base_url: str
|
||||
base_url_explicit: NotRequired[bool]
|
||||
|
||||
|
||||
class CliPollData(TypedDict, total=False):
|
||||
status: str
|
||||
key: str
|
||||
user_id: str
|
||||
teams: list[str]
|
||||
team_details: object
|
||||
requires_team_selection: bool
|
||||
team_id: str
|
||||
|
||||
|
||||
class CliPollRequestKwargs(TypedDict, total=False):
|
||||
timeout: int
|
||||
headers: dict[str, str]
|
||||
|
||||
|
||||
class CliSsoStartData(TypedDict):
|
||||
login_id: str
|
||||
poll_secret: str
|
||||
user_code: str
|
||||
|
||||
|
||||
class CliAuthResult(TypedDict):
|
||||
api_key: str
|
||||
user_id: str | None
|
||||
teams: list[str]
|
||||
team_id: str | None
|
||||
|
||||
|
||||
# Token storage utilities
|
||||
def get_token_file_path() -> str:
|
||||
"""Get the path to store the authentication token"""
|
||||
|
|
@ -27,12 +79,12 @@ def get_token_file_path() -> str:
|
|||
return str(config_dir / "token.json")
|
||||
|
||||
|
||||
def save_token(token_data: dict[str, Any]) -> None:
|
||||
def save_token(token_data: CliTokenData) -> None:
|
||||
"""Save token data to file"""
|
||||
write_private_json(get_token_file_path(), token_data)
|
||||
|
||||
|
||||
def load_token() -> dict[str, Any] | None:
|
||||
def load_token() -> CliTokenData | None:
|
||||
"""Load token data from file"""
|
||||
token_file: Final = get_token_file_path()
|
||||
if not os.path.exists(token_file):
|
||||
|
|
@ -65,7 +117,7 @@ def get_stored_api_key(expected_base_url: str | None = None) -> str | None:
|
|||
|
||||
|
||||
# Team selection utilities
|
||||
def display_teams_table(teams: list[dict[str, Any]]) -> None:
|
||||
def display_teams_table(teams: list[CliTeam]) -> None:
|
||||
"""Display teams in a formatted table"""
|
||||
console: Final = Console()
|
||||
|
||||
|
|
@ -165,7 +217,7 @@ def display_interactive_team_selection(teams: list[dict[str, Any]], selected_ind
|
|||
for i, team in enumerate(teams):
|
||||
team_alias = team.get("team_alias") or "N/A"
|
||||
team_id = team.get("team_id", "N/A")
|
||||
models = team.get("models", [])
|
||||
models: list[str] = team.get("models", [])
|
||||
max_budget = team.get("max_budget")
|
||||
|
||||
# Format models list
|
||||
|
|
@ -249,10 +301,11 @@ def prompt_team_selection_fallback(
|
|||
|
||||
while True:
|
||||
try:
|
||||
choice = click.prompt(
|
||||
prompt_response: str = click.prompt(
|
||||
"\nSelect a team by entering the index number (or 'skip' to continue without a team)",
|
||||
type=str,
|
||||
).strip()
|
||||
)
|
||||
choice = prompt_response.strip()
|
||||
|
||||
if choice.lower() == "skip":
|
||||
return None
|
||||
|
|
@ -275,7 +328,7 @@ def prompt_team_selection_fallback(
|
|||
|
||||
def _response_error_detail(response: requests.Response) -> str | None:
|
||||
try:
|
||||
body: Final = response.json()
|
||||
body: Final[dict[str, object] | list[object] | str | int | float | bool | None] = response.json()
|
||||
except ValueError:
|
||||
return None
|
||||
detail: Final = body.get("detail") if isinstance(body, dict) else None
|
||||
|
|
@ -309,15 +362,15 @@ def _poll_for_ready_data(
|
|||
other_status_log_every: int = 10,
|
||||
http_error_log_every: int = 10,
|
||||
connection_error_log_every: int = 10,
|
||||
) -> dict[str, Any] | None:
|
||||
) -> CliPollData | None:
|
||||
for attempt in range(total_timeout // poll_interval):
|
||||
try:
|
||||
request_kwargs: dict[str, Any] = {"timeout": request_timeout}
|
||||
request_kwargs: CliPollRequestKwargs = {"timeout": request_timeout}
|
||||
if headers is not None:
|
||||
request_kwargs["headers"] = headers
|
||||
response = requests.get(url, **request_kwargs)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
data: CliPollData = response.json()
|
||||
status = data.get("status")
|
||||
if status == "ready":
|
||||
return data
|
||||
|
|
@ -341,7 +394,7 @@ def _poll_for_ready_data(
|
|||
return None
|
||||
|
||||
|
||||
def _normalize_teams(teams, team_details):
|
||||
def _normalize_teams(teams: object, team_details: object) -> list[CliTeam]:
|
||||
"""If team_details are a
|
||||
|
||||
Args:
|
||||
|
|
@ -365,7 +418,7 @@ def _normalize_teams(teams, team_details):
|
|||
return []
|
||||
|
||||
|
||||
def _start_cli_sso_flow(base_url: str) -> dict[str, Any]:
|
||||
def _start_cli_sso_flow(base_url: str) -> CliSsoStartData:
|
||||
start_url: Final = f"{base_url}/sso/cli/start"
|
||||
try:
|
||||
response: Final = requests.post(start_url, timeout=10)
|
||||
|
|
@ -389,7 +442,7 @@ def _start_cli_sso_flow(base_url: str) -> dict[str, Any]:
|
|||
)
|
||||
|
||||
try:
|
||||
data: Final = response.json()
|
||||
data: Final[CliSsoStartData] = response.json()
|
||||
except ValueError:
|
||||
content_type: Final = response.headers.get("content-type", "unknown")
|
||||
raise ValueError(
|
||||
|
|
@ -398,7 +451,7 @@ def _start_cli_sso_flow(base_url: str) -> dict[str, Any]:
|
|||
f"Response starts with: {response.text[:200]!r}"
|
||||
)
|
||||
|
||||
required_fields: Final = ("login_id", "poll_secret", "user_code")
|
||||
required_fields: Final[tuple[str, ...]] = ("login_id", "poll_secret", "user_code")
|
||||
missing_fields: Final = tuple(field for field in required_fields if not isinstance(data.get(field), str))
|
||||
if missing_fields:
|
||||
raise ValueError(
|
||||
|
|
@ -412,7 +465,7 @@ def _get_cli_sso_poll_headers(poll_secret: str) -> dict[str, str]:
|
|||
return {"x-litellm-cli-poll-secret": poll_secret}
|
||||
|
||||
|
||||
def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> dict | None:
|
||||
def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> CliAuthResult | None:
|
||||
"""
|
||||
Poll the server for authentication completion and handle team selection.
|
||||
|
||||
|
|
@ -431,7 +484,7 @@ def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> di
|
|||
teams = data.get("teams", [])
|
||||
team_details: Final = data.get("team_details")
|
||||
user_id = data.get("user_id")
|
||||
normalized_teams: Final[list[dict[str, Any]]] = _normalize_teams(teams, team_details)
|
||||
normalized_teams: Final[list[CliTeam]] = _normalize_teams(teams, team_details)
|
||||
if not normalized_teams:
|
||||
click.echo("Warning: No teams available for selection.")
|
||||
return None
|
||||
|
|
@ -478,7 +531,7 @@ def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> di
|
|||
|
||||
|
||||
def _handle_team_selection_during_polling(
|
||||
base_url: str, key_id: str, poll_secret: str, teams: list[dict[str, Any]]
|
||||
base_url: str, key_id: str, poll_secret: str, teams: list[CliTeam]
|
||||
) -> str | None:
|
||||
"""
|
||||
Handle team selection and re-poll with selected team_id.
|
||||
|
|
@ -522,7 +575,7 @@ def _handle_team_selection_during_polling(
|
|||
return None
|
||||
|
||||
|
||||
def _render_and_prompt_for_team_selection(teams: list[dict[str, Any]]) -> str | None:
|
||||
def _render_and_prompt_for_team_selection(teams: list[CliTeam]) -> str | None:
|
||||
"""Render teams table and prompt user for a team selection.
|
||||
|
||||
Returns the selected team_id as a string, or None if selection was
|
||||
|
|
@ -546,10 +599,11 @@ def _render_and_prompt_for_team_selection(teams: list[dict[str, Any]]) -> str |
|
|||
# Simple selection
|
||||
while True:
|
||||
try:
|
||||
choice = click.prompt(
|
||||
prompt_response: str = click.prompt(
|
||||
"\nSelect a team by entering the index number (or 'skip' to use first team)",
|
||||
type=str,
|
||||
).strip()
|
||||
)
|
||||
choice = prompt_response.strip()
|
||||
|
||||
if choice.lower() == "skip":
|
||||
# Default to the first team's ID if the user skips an
|
||||
|
|
@ -582,7 +636,8 @@ def login(ctx: click.Context):
|
|||
from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
|
||||
from litellm.proxy.client.cli.interface import show_commands
|
||||
|
||||
base_url: Final = ctx.obj["base_url"]
|
||||
ctx_obj: Final[CliContextObj] = ctx.obj
|
||||
base_url: Final = ctx_obj["base_url"]
|
||||
|
||||
try:
|
||||
cli_sso_flow: Final = _start_cli_sso_flow(base_url=base_url)
|
||||
|
|
@ -675,8 +730,9 @@ def print_token(ctx: click.Context):
|
|||
# explicitly pointed us at a server, trust whichever one `lite login`
|
||||
# actually issued this token for -- that's the whole point of not
|
||||
# needing a wrapper command.
|
||||
if ctx.obj.get("base_url_explicit"):
|
||||
base_url: Final = ctx.obj["base_url"]
|
||||
ctx_obj: Final[CliContextObj] = ctx.obj
|
||||
if ctx_obj.get("base_url_explicit"):
|
||||
base_url: Final = ctx_obj["base_url"]
|
||||
if token_data.get("base_url") != base_url.rstrip("/"):
|
||||
click.echo("Not authenticated for this server. Run 'lite login'.", err=True)
|
||||
sys.exit(1)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, TypeVar
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -311,14 +311,17 @@ def _coerce_timeout(value: Any, fallback: float) -> float:
|
|||
return fallback
|
||||
|
||||
|
||||
_ReadResultT: Final = TypeVar("_ReadResultT")
|
||||
|
||||
|
||||
async def call_with_db_reconnect_retry(
|
||||
prisma_client: Any,
|
||||
coro_factory: Callable[[], Awaitable[Any]],
|
||||
coro_factory: Callable[[], Awaitable[_ReadResultT]],
|
||||
*,
|
||||
reason: str,
|
||||
timeout_seconds: float | None = None,
|
||||
lock_timeout_seconds: float | None = None,
|
||||
) -> Any:
|
||||
) -> _ReadResultT:
|
||||
"""Run a Prisma read coroutine with one transport-reconnect-and-retry.
|
||||
|
||||
The canonical "self-heal a transient DB transport blip" wrapper used by
|
||||
|
|
|
|||
|
|
@ -9,10 +9,10 @@ import asyncio
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import AsyncGenerator
|
||||
from collections.abc import AsyncGenerator, Coroutine, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from re import Pattern
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, cast
|
||||
|
||||
import yaml
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -28,6 +28,7 @@ from litellm.types.utils import (
|
|||
GenericGuardrailAPIInputs,
|
||||
GuardrailStatus,
|
||||
GuardrailTracingDetail,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
)
|
||||
|
||||
|
|
@ -83,6 +84,46 @@ WORD_NUMBER_SEQUENCE_PATTERN: Final = re.compile(
|
|||
WORD_NUMBER_TOKEN_FINDER: Final = re.compile(rf"(?:{WORD_NUMBER_TOKEN_REGEX})", re.IGNORECASE)
|
||||
|
||||
|
||||
class ConditionalCategoryConfig(TypedDict):
|
||||
identifier_words: Sequence[str]
|
||||
block_words: Sequence[str]
|
||||
action: ContentFilterAction
|
||||
severity: str
|
||||
|
||||
|
||||
class CompiledPatternEntry(TypedDict):
|
||||
regex: Pattern[str]
|
||||
pattern_name: str
|
||||
action: ContentFilterAction
|
||||
keyword_regex: Pattern[str] | None
|
||||
allow_word_numbers: bool
|
||||
|
||||
|
||||
class _PatternExtraLookup(TypedDict):
|
||||
keyword_pattern: str | None
|
||||
allow_word_numbers: bool
|
||||
|
||||
|
||||
class _CategoryConfigView(TypedDict):
|
||||
category: object
|
||||
enabled: object
|
||||
action: object
|
||||
category_file: str | None
|
||||
|
||||
|
||||
class CategoryFileData(TypedDict, total=False):
|
||||
category_name: str
|
||||
description: str
|
||||
default_action: str
|
||||
keywords: Sequence[Mapping[str, str]]
|
||||
exceptions: Sequence[str]
|
||||
identifier_words: Sequence[str]
|
||||
always_block_keywords: Sequence[Mapping[str, str]]
|
||||
inherit_from: str
|
||||
additional_block_words: Sequence[str]
|
||||
phrase_patterns: Sequence[str]
|
||||
|
||||
|
||||
# Helper data structure for category-based detection
|
||||
class CategoryConfig:
|
||||
"""Configuration for a content category."""
|
||||
|
|
@ -92,13 +133,13 @@ class CategoryConfig:
|
|||
category_name: str,
|
||||
description: str,
|
||||
default_action: ContentFilterAction,
|
||||
keywords: list[dict[str, str]],
|
||||
exceptions: list[str],
|
||||
identifier_words: list[str] | None = None,
|
||||
always_block_keywords: list[dict[str, str]] | None = None,
|
||||
keywords: Sequence[Mapping[str, str]],
|
||||
exceptions: Sequence[str],
|
||||
identifier_words: Sequence[str] | None = None,
|
||||
always_block_keywords: Sequence[Mapping[str, str]] | None = None,
|
||||
inherit_from: str | None = None,
|
||||
additional_block_words: list[str] | None = None,
|
||||
phrase_patterns: list[str] | None = None,
|
||||
additional_block_words: Sequence[str] | None = None,
|
||||
phrase_patterns: Sequence[str] | None = None,
|
||||
):
|
||||
self.category_name = category_name
|
||||
self.description = description
|
||||
|
|
@ -151,7 +192,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
severity_threshold: str = "medium",
|
||||
llm_router: Router | None = None,
|
||||
image_model: str | None = None,
|
||||
competitor_intent_config: dict[str, Any] | None = None,
|
||||
competitor_intent_config: dict[str, object] | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
|
|
@ -194,9 +235,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
# Always-block keywords are checked after exceptions (exceptions take precedence)
|
||||
self.always_block_category_keywords: dict[str, tuple[str, str, ContentFilterAction]] = {}
|
||||
# Store conditional categories (identifier_words + block_words)
|
||||
self.conditional_categories: dict[
|
||||
str, dict[str, Any]
|
||||
] = {} # category_name -> {identifier_words, block_words, action, severity}
|
||||
self.conditional_categories: dict[str, ConditionalCategoryConfig] = {}
|
||||
|
||||
# Competitor intent checker (optional; airline uses major_airlines.json, generic requires competitors)
|
||||
self._competitor_intent_checker: BaseCompetitorIntentChecker | None = None
|
||||
|
|
@ -212,7 +251,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
normalized_blocked_words: Final = self._normalize_blocked_words(blocked_words)
|
||||
|
||||
# Compile regex patterns
|
||||
self.compiled_patterns: list[dict[str, Any]] = []
|
||||
self.compiled_patterns: list[CompiledPatternEntry] = []
|
||||
for pattern_config in normalized_patterns:
|
||||
self._add_pattern(pattern_config)
|
||||
|
||||
|
|
@ -250,7 +289,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
"Loaded %s categories with %s keywords", len(self.loaded_categories), len(self.category_keywords)
|
||||
)
|
||||
|
||||
def _init_competitor_intent_checker(self, competitor_intent_config: dict[str, Any]) -> None:
|
||||
def _init_competitor_intent_checker(self, competitor_intent_config: dict[str, object]) -> None:
|
||||
try:
|
||||
competitor_intent_type: Final = competitor_intent_config.get("competitor_intent_type", "airline")
|
||||
if competitor_intent_type == "generic":
|
||||
|
|
@ -293,6 +332,15 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
result.append(word)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _category_config_view(cat_config: ContentFilterCategoryConfig) -> _CategoryConfigView:
|
||||
return {
|
||||
"category": cat_config.get("category"),
|
||||
"enabled": cat_config.get("enabled", True),
|
||||
"action": cat_config.get("action"),
|
||||
"category_file": cat_config.get("category_file"),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _assert_within_categories_dir(path: str, categories_dir: str) -> None:
|
||||
"""Raise ValueError if path escapes the categories directory."""
|
||||
|
|
@ -395,7 +443,8 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
categories_dir: Final = os.path.join(os.path.dirname(__file__), "categories")
|
||||
|
||||
for cat_config in categories:
|
||||
category_name = cat_config.get("category")
|
||||
view = self._category_config_view(cat_config)
|
||||
category_name = view["category"]
|
||||
if not category_name or not isinstance(category_name, str):
|
||||
verbose_proxy_logger.warning("Category name missing or invalid in config, skipping")
|
||||
continue
|
||||
|
|
@ -405,12 +454,12 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.warning("Category name '%s' contains invalid characters, skipping", category_name)
|
||||
continue
|
||||
|
||||
enabled = cat_config.get("enabled", True)
|
||||
action = cat_config.get("action")
|
||||
enabled = view["enabled"]
|
||||
action = view["action"]
|
||||
severity_threshold = (
|
||||
cat_config.get("severity_threshold", self.severity_threshold) or self.severity_threshold
|
||||
)
|
||||
custom_file = cat_config.get("category_file")
|
||||
custom_file = view["category_file"]
|
||||
|
||||
if not enabled:
|
||||
verbose_proxy_logger.debug("Category %s is disabled, skipping", category_name)
|
||||
|
|
@ -514,7 +563,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
categories_dir: Directory containing category files
|
||||
"""
|
||||
try:
|
||||
block_words: Final = []
|
||||
block_words: Final[list[str]] = []
|
||||
inherit_from = category_config_obj.inherit_from
|
||||
|
||||
# Load inherited block words if specified
|
||||
|
|
@ -605,11 +654,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
"""
|
||||
if file_path.lower().endswith(".json"):
|
||||
return self._load_category_file_json(file_path)
|
||||
with open(file_path, "r") as f:
|
||||
data: Final = yaml.safe_load(f)
|
||||
|
||||
# Handle always_block_keywords if present
|
||||
always_block: Final = data.get("always_block_keywords", [])
|
||||
data: Final = self._read_category_yaml(file_path)
|
||||
|
||||
return CategoryConfig(
|
||||
category_name=data.get("category_name", "unknown"),
|
||||
|
|
@ -618,12 +663,17 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
keywords=data.get("keywords", []),
|
||||
exceptions=data.get("exceptions", []),
|
||||
identifier_words=data.get("identifier_words"),
|
||||
always_block_keywords=always_block,
|
||||
always_block_keywords=data.get("always_block_keywords", []),
|
||||
inherit_from=data.get("inherit_from"),
|
||||
additional_block_words=data.get("additional_block_words"),
|
||||
phrase_patterns=data.get("phrase_patterns"),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _read_category_yaml(file_path: str) -> CategoryFileData:
|
||||
with open(file_path, "r") as f:
|
||||
return yaml.safe_load(f)
|
||||
|
||||
def _load_category_file_json(self, file_path: str) -> CategoryConfig:
|
||||
"""
|
||||
Load a category from the harm_toxic_abuse-style JSON format.
|
||||
|
|
@ -682,13 +732,13 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
pattern_config: ContentFilterPattern configuration
|
||||
"""
|
||||
try:
|
||||
extra_config: dict[str, Any] = {}
|
||||
extra_config: _PatternExtraLookup = {"keyword_pattern": None, "allow_word_numbers": False}
|
||||
if pattern_config.pattern_type == "prebuilt":
|
||||
if not pattern_config.pattern_name:
|
||||
raise ValueError("pattern_name is required for prebuilt patterns")
|
||||
compiled = get_compiled_pattern(pattern_config.pattern_name)
|
||||
pattern_name = pattern_config.pattern_name
|
||||
extra_config = PATTERN_EXTRA_CONFIG.get(pattern_name, {}) or {}
|
||||
extra_config = self._lookup_pattern_extra(pattern_name)
|
||||
elif pattern_config.pattern_type == "regex":
|
||||
if not pattern_config.pattern:
|
||||
raise ValueError("pattern is required for regex patterns")
|
||||
|
|
@ -697,9 +747,8 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
else:
|
||||
raise ValueError(f"Unknown pattern_type: {pattern_config.pattern_type}")
|
||||
|
||||
keyword_regex: Pattern | None = None
|
||||
if extra_config.get("keyword_pattern"):
|
||||
keyword_regex = re.compile(extra_config["keyword_pattern"], re.IGNORECASE)
|
||||
keyword_pattern: Final = extra_config["keyword_pattern"]
|
||||
keyword_regex: Final = re.compile(keyword_pattern, re.IGNORECASE) if keyword_pattern else None
|
||||
|
||||
self.compiled_patterns.append(
|
||||
{
|
||||
|
|
@ -707,7 +756,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
"pattern_name": pattern_name,
|
||||
"action": pattern_config.action,
|
||||
"keyword_regex": keyword_regex,
|
||||
"allow_word_numbers": bool(extra_config.get("allow_word_numbers")),
|
||||
"allow_word_numbers": extra_config["allow_word_numbers"],
|
||||
}
|
||||
)
|
||||
verbose_proxy_logger.debug("Added pattern: %s with action %s", pattern_name, pattern_config.action)
|
||||
|
|
@ -715,6 +764,14 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.error("Error adding pattern %s: %s", pattern_config, e)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def _lookup_pattern_extra(pattern_name: str) -> _PatternExtraLookup:
|
||||
extra: Final = PATTERN_EXTRA_CONFIG.get(pattern_name)
|
||||
return {
|
||||
"keyword_pattern": extra.get("keyword_pattern") if extra is not None else None,
|
||||
"allow_word_numbers": bool(extra.get("allow_word_numbers")) if extra is not None else False,
|
||||
}
|
||||
|
||||
def _load_blocked_words_file(self, file_path: str) -> None:
|
||||
"""
|
||||
Load blocked words from a YAML file.
|
||||
|
|
@ -754,18 +811,16 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
except Exception as e:
|
||||
raise Exception(f"Error loading blocked words file {file_path}: {e}")
|
||||
|
||||
def _find_pattern_spans(self, text: str, pattern_entry: dict[str, Any]) -> list[tuple[int, int]]:
|
||||
def _find_pattern_spans(self, text: str, pattern_entry: CompiledPatternEntry) -> list[tuple[int, int]]:
|
||||
"""Return all match spans for a pattern, applying contextual rules if required."""
|
||||
|
||||
regex: Final[Pattern] = pattern_entry["regex"]
|
||||
keyword_regex: Final[Pattern | None] = pattern_entry.get("keyword_regex")
|
||||
regex: Final[Pattern[str]] = pattern_entry["regex"]
|
||||
keyword_regex: Final[Pattern[str] | None] = pattern_entry.get("keyword_regex")
|
||||
allow_word_numbers: Final[bool] = pattern_entry.get("allow_word_numbers", False)
|
||||
|
||||
keyword_matches: list[re.Match] | None = None
|
||||
if keyword_regex is not None:
|
||||
keyword_matches = list(keyword_regex.finditer(text))
|
||||
if not keyword_matches:
|
||||
return []
|
||||
keyword_matches: Final = list(keyword_regex.finditer(text)) if keyword_regex is not None else None
|
||||
if keyword_matches is not None and not keyword_matches:
|
||||
return []
|
||||
|
||||
match_spans: Final[list[tuple[int, int]]] = []
|
||||
|
||||
|
|
@ -795,7 +850,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
self,
|
||||
value_start: int,
|
||||
value_end: int,
|
||||
keyword_matches: list[re.Match],
|
||||
keyword_matches: Sequence[re.Match[str]],
|
||||
text: str,
|
||||
) -> bool:
|
||||
"""Check if a value is separated from a keyword by an allowed gap."""
|
||||
|
|
@ -861,7 +916,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
def _convert_word_number_sequence(self, sequence: str) -> str | None:
|
||||
"""Convert a spelled-out digit sequence (e.g., 'One-Two') into digits."""
|
||||
|
||||
tokens: Final = WORD_NUMBER_TOKEN_FINDER.findall(sequence)
|
||||
tokens: Final[list[str]] = WORD_NUMBER_TOKEN_FINDER.findall(sequence)
|
||||
if not tokens:
|
||||
return None
|
||||
|
||||
|
|
@ -1328,7 +1383,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
HTTPException: If sensitive content is detected and action is BLOCK
|
||||
"""
|
||||
# Collect all exceptions from loaded categories
|
||||
all_exceptions: Final = []
|
||||
all_exceptions: Final[list[str]] = []
|
||||
for category in self.loaded_categories.values():
|
||||
all_exceptions.extend(category.exceptions)
|
||||
|
||||
|
|
@ -1404,7 +1459,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
if not (images and self.image_model and self.llm_router):
|
||||
return
|
||||
|
||||
tasks: Final = []
|
||||
tasks: Final[list[Coroutine[object, object, ModelResponse]]] = []
|
||||
for image in images:
|
||||
task = self.llm_router.acompletion(
|
||||
model=self.image_model,
|
||||
|
|
@ -1425,12 +1480,10 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
tasks.append(task)
|
||||
|
||||
responses: Final = await asyncio.gather(*tasks)
|
||||
descriptions: Final = []
|
||||
descriptions: Final[list[str]] = []
|
||||
for response in responses:
|
||||
choice = response.choices[0]
|
||||
message = getattr(choice, "message", None)
|
||||
if message and getattr(message, "content", None):
|
||||
image_description = message.content
|
||||
image_description = self._describe_image_response_content(response)
|
||||
if image_description:
|
||||
verbose_proxy_logger.debug("Image description: %s", image_description)
|
||||
descriptions.append(image_description)
|
||||
else:
|
||||
|
|
@ -1447,7 +1500,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
except HTTPException as e:
|
||||
# e.detail can be a string or dict
|
||||
if isinstance(e.detail, dict) and "error" in e.detail:
|
||||
detail_dict = cast(dict[str, Any], e.detail)
|
||||
detail_dict = cast(dict[str, str], e.detail)
|
||||
detail_dict["error"] = detail_dict["error"] + " (Image description): " + description
|
||||
elif isinstance(e.detail, str):
|
||||
e.detail = e.detail + " (Image description): " + description
|
||||
|
|
@ -1455,6 +1508,14 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
e.detail = "Content blocked: Image description detected" + description
|
||||
raise e
|
||||
|
||||
@staticmethod
|
||||
def _describe_image_response_content(response: ModelResponse) -> str | None:
|
||||
choice = response.choices[0]
|
||||
message = getattr(choice, "message", None)
|
||||
if message and getattr(message, "content", None):
|
||||
return message.content
|
||||
return None
|
||||
|
||||
def _count_masked_entities(
|
||||
self,
|
||||
detections: list[ContentFilterDetection],
|
||||
|
|
@ -1484,12 +1545,12 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
category = category_detection["category"]
|
||||
masked_entity_count[category] = masked_entity_count.get(category, 0) + 1
|
||||
|
||||
def _build_match_details(self, detections: list[ContentFilterDetection]) -> list[dict]:
|
||||
def _build_match_details(self, detections: list[ContentFilterDetection]) -> list[dict[str, object]]:
|
||||
"""Build match_details list from content filter detections."""
|
||||
match_details: Final[list[dict]] = []
|
||||
match_details: Final[list[dict[str, object]]] = []
|
||||
for detection in detections:
|
||||
action_taken = detection.get("action", detection.get("action_hint", ""))
|
||||
detail: dict = {"type": detection["type"], "action_taken": action_taken}
|
||||
detail: dict[str, object] = {"type": detection["type"], "action_taken": action_taken}
|
||||
if detection["type"] == "pattern":
|
||||
detail["detection_method"] = "regex"
|
||||
detail["snippet"] = cast(PatternDetection, detection).get("pattern_name", "")
|
||||
|
|
@ -1510,7 +1571,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
|
||||
def _get_detection_methods(self, detections: list[ContentFilterDetection]) -> str:
|
||||
"""Get comma-separated detection methods used."""
|
||||
methods: Final[set] = set()
|
||||
methods: Final[set[str]] = set()
|
||||
for detection in detections:
|
||||
if detection["type"] == "pattern":
|
||||
methods.add("regex")
|
||||
|
|
@ -1659,7 +1720,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
guardrail_json_response = exception_str if exception_str else [dict(detection) for detection in detections]
|
||||
|
||||
# Competitor intent: add confidence and classification to tracing if present
|
||||
tracing_kw: Final[dict[str, Any]] = {
|
||||
tracing_kw: Final[GuardrailTracingDetail] = {
|
||||
"guardrail_id": self.config_guardrail_id or self.guardrail_name,
|
||||
"policy_template": self.config_policy_template or self._get_policy_templates(),
|
||||
"detection_method": (self._get_detection_methods(detections) if detections else None),
|
||||
|
|
|
|||
|
|
@ -55,7 +55,7 @@ for pattern_data in _PATTERNS_DATA["patterns"]:
|
|||
PATTERN_EXTRA_CONFIG[pattern_data["name"]] = extra_config
|
||||
|
||||
|
||||
def get_compiled_pattern(pattern_name: str) -> Pattern:
|
||||
def get_compiled_pattern(pattern_name: str) -> Pattern[str]:
|
||||
"""
|
||||
Get a compiled regex pattern by name.
|
||||
|
||||
|
|
|
|||
|
|
@ -73,12 +73,14 @@ import hashlib
|
|||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Final, Optional
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
import jwt
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey
|
||||
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey, RSAPublicKey
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import DualCache
|
||||
|
|
@ -90,13 +92,28 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from jwt.types import Options
|
||||
|
||||
|
||||
class _OIDCDiscoveryDocument(TypedDict, total=False):
|
||||
jwks_uri: str
|
||||
|
||||
|
||||
class _JWTDecodeKwargs(TypedDict):
|
||||
algorithms: Sequence[str]
|
||||
options: "Options"
|
||||
audience: NotRequired[str]
|
||||
issuer: NotRequired[str]
|
||||
|
||||
|
||||
# Module-level singleton for the JWKS discovery endpoint to access.
|
||||
_mcp_jwt_signer_instance: Optional["MCPJWTSigner"] = None
|
||||
|
||||
_MCP_JWT_CALL_TYPES: Final = frozenset({"call_mcp_tool", "list_mcp_tools"})
|
||||
|
||||
# Simple in-memory JWKS cache: keyed by JWKS URI → (keys_list, fetched_at).
|
||||
_jwks_cache: Final[dict[str, tuple]] = {}
|
||||
_jwks_cache: Final[dict[str, tuple[Sequence[Mapping[str, object]], float]]] = {}
|
||||
_JWKS_CACHE_TTL: Final = 3600 # 1 hour
|
||||
|
||||
|
||||
|
|
@ -133,7 +150,7 @@ def _int_to_base64url(n: int) -> str:
|
|||
return base64.urlsafe_b64encode(n.to_bytes(byte_length, byteorder="big")).rstrip(b"=").decode("ascii")
|
||||
|
||||
|
||||
def _compute_kid(public_key: Any) -> str:
|
||||
def _compute_kid(public_key: RSAPublicKey) -> str:
|
||||
"""Derive a key ID from the public key's DER encoding (SHA-256, first 16 hex chars)."""
|
||||
der_bytes: Final = public_key.public_bytes(
|
||||
encoding=serialization.Encoding.DER,
|
||||
|
|
@ -142,7 +159,7 @@ def _compute_kid(public_key: Any) -> str:
|
|||
return hashlib.sha256(der_bytes).hexdigest()[:16]
|
||||
|
||||
|
||||
async def _fetch_jwks(jwks_uri: str) -> list[dict[str, Any]]:
|
||||
async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]:
|
||||
"""
|
||||
Fetch and cache a JWKS from the given URI.
|
||||
|
||||
|
|
@ -163,12 +180,13 @@ async def _fetch_jwks(jwks_uri: str) -> list[dict[str, Any]]:
|
|||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
resp: Final = await client.get(jwks_uri, headers={"Accept": "application/json"})
|
||||
resp.raise_for_status()
|
||||
keys = resp.json().get("keys", [])
|
||||
_jwks_cache[jwks_uri] = (keys, now)
|
||||
return keys
|
||||
jwks_body: Final[Mapping[str, Sequence[Mapping[str, object]]]] = resp.json()
|
||||
fetched_keys: Final = jwks_body.get("keys", [])
|
||||
_jwks_cache[jwks_uri] = (fetched_keys, now)
|
||||
return fetched_keys
|
||||
|
||||
|
||||
async def _fetch_oidc_discovery(discovery_uri: str) -> dict[str, Any]:
|
||||
async def _fetch_oidc_discovery(discovery_uri: str) -> _OIDCDiscoveryDocument:
|
||||
"""Fetch an OIDC discovery document and return its parsed JSON."""
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
|
|
@ -178,7 +196,8 @@ async def _fetch_oidc_discovery(discovery_uri: str) -> dict[str, Any]:
|
|||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
resp: Final = await client.get(discovery_uri, headers={"Accept": "application/json"})
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
document: Final[_OIDCDiscoveryDocument] = resp.json()
|
||||
return document
|
||||
|
||||
|
||||
class MCPJWTSigner(CustomGuardrail):
|
||||
|
|
@ -230,8 +249,8 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# FR-12: End-user identity mapping
|
||||
end_user_claim_sources: list[str] | None = None,
|
||||
# FR-13: Claim operations
|
||||
add_claims: dict[str, Any] | None = None,
|
||||
set_claims: dict[str, Any] | None = None,
|
||||
add_claims: Mapping[str, object] | None = None,
|
||||
set_claims: Mapping[str, object] | None = None,
|
||||
remove_claims: list[str] | None = None,
|
||||
# FR-14: Two-token model
|
||||
channel_token_audience: str | None = None,
|
||||
|
|
@ -283,7 +302,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
self.verify_issuer: str | None = verify_issuer
|
||||
self.verify_audience: str | None = verify_audience
|
||||
# Cached OIDC discovery document (fetched lazily, TTL = 24 h)
|
||||
self._oidc_discovery_doc: dict[str, Any] | None = None
|
||||
self._oidc_discovery_doc: _OIDCDiscoveryDocument | None = None
|
||||
self._oidc_discovery_fetched_at: float = 0.0
|
||||
|
||||
# --- FR-12: End-user identity mapping ---
|
||||
|
|
@ -294,8 +313,8 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
]
|
||||
|
||||
# --- FR-13: Claim operations ---
|
||||
self.add_claims: dict[str, Any] = add_claims or {}
|
||||
self.set_claims: dict[str, Any] = set_claims or {}
|
||||
self.add_claims: Mapping[str, object] = add_claims or {}
|
||||
self.set_claims: Mapping[str, object] = set_claims or {}
|
||||
self.remove_claims: list[str] = remove_claims or []
|
||||
|
||||
# --- FR-14: Two-token model ---
|
||||
|
|
@ -347,7 +366,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
"""
|
||||
return 3600 if self._persistent_key else 300
|
||||
|
||||
def get_jwks(self) -> dict[str, Any]:
|
||||
def get_jwks(self) -> Mapping[str, Sequence[Mapping[str, str]]]:
|
||||
"""
|
||||
Return the JWKS for the RSA public key.
|
||||
Used by GET /.well-known/jwks.json so MCP servers can verify tokens.
|
||||
|
|
@ -374,7 +393,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# the IdP, short enough to pick up jwks_uri changes after key rotation.
|
||||
_OIDC_DISCOVERY_TTL = 86400
|
||||
|
||||
async def _get_oidc_discovery(self) -> dict[str, Any]:
|
||||
async def _get_oidc_discovery(self) -> _OIDCDiscoveryDocument:
|
||||
"""Fetch and cache the OIDC discovery document with a 24-hour TTL.
|
||||
|
||||
Only caches when the doc contains a 'jwks_uri' so that a transient or
|
||||
|
|
@ -391,7 +410,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
return doc
|
||||
return self._oidc_discovery_doc or {}
|
||||
|
||||
async def _verify_incoming_jwt(self, raw_token: str) -> dict[str, Any]:
|
||||
async def _verify_incoming_jwt(self, raw_token: str) -> dict[str, object]:
|
||||
"""
|
||||
Verify an incoming Bearer JWT against the configured IdP's JWKS.
|
||||
|
||||
|
|
@ -438,8 +457,8 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# it infers from the key type (RSAPublicKey → RS256).
|
||||
alg: Final = getattr(signing_jwk, "algorithm_name", None) or "RS256"
|
||||
|
||||
decode_options: Final[dict[str, Any]] = {"verify_exp": True}
|
||||
decode_kwargs: Final[dict[str, Any]] = {
|
||||
decode_options: Final[Options] = {"verify_exp": True}
|
||||
decode_kwargs: Final[_JWTDecodeKwargs] = {
|
||||
"algorithms": [alg],
|
||||
"options": decode_options,
|
||||
}
|
||||
|
|
@ -451,10 +470,10 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
if self.verify_issuer:
|
||||
decode_kwargs["issuer"] = self.verify_issuer
|
||||
|
||||
payload: Final[dict[str, Any]] = jwt.decode(raw_token, signing_jwk.key, **decode_kwargs)
|
||||
payload: Final[dict[str, object]] = jwt.decode(raw_token, signing_jwk.key, **decode_kwargs)
|
||||
return payload
|
||||
|
||||
async def _introspect_opaque_token(self, token: str) -> dict[str, Any]:
|
||||
async def _introspect_opaque_token(self, token: str) -> dict[str, object]:
|
||||
"""
|
||||
Perform RFC 7662 token introspection for opaque (non-JWT) tokens.
|
||||
|
||||
|
|
@ -479,7 +498,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
headers={"Accept": "application/json"},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
result: Final[dict[str, Any]] = resp.json()
|
||||
result: Final[dict[str, object]] = resp.json()
|
||||
if not result.get("active", False):
|
||||
raise jwt.exceptions.ExpiredSignatureError(
|
||||
"MCPJWTSigner: incoming token is inactive (introspection returned active=false)"
|
||||
|
|
@ -492,7 +511,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
|
||||
def _validate_required_claims(
|
||||
self,
|
||||
jwt_claims: dict[str, Any] | None,
|
||||
jwt_claims: Mapping[str, object] | None,
|
||||
) -> None:
|
||||
"""
|
||||
Raise HTTP 403 if any required_claims are absent from the verified
|
||||
|
|
@ -522,7 +541,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
def _resolve_end_user_identity(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
jwt_claims: dict[str, Any] | None,
|
||||
jwt_claims: Mapping[str, object] | None,
|
||||
) -> str:
|
||||
"""
|
||||
Resolve the outbound JWT 'sub' using the ordered end_user_claim_sources list.
|
||||
|
|
@ -545,19 +564,19 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
value = str(raw) if raw else None
|
||||
|
||||
elif source == "litellm:user_id":
|
||||
uid = getattr(user_api_key_dict, "user_id", None)
|
||||
uid = user_api_key_dict.user_id
|
||||
value = str(uid) if uid else None
|
||||
|
||||
elif source == "litellm:email":
|
||||
email = getattr(user_api_key_dict, "user_email", None)
|
||||
email = user_api_key_dict.user_email
|
||||
value = str(email) if email else None
|
||||
|
||||
elif source == "litellm:end_user_id":
|
||||
eid = getattr(user_api_key_dict, "end_user_id", None)
|
||||
eid = user_api_key_dict.end_user_id
|
||||
value = str(eid) if eid else None
|
||||
|
||||
elif source == "litellm:team_id":
|
||||
tid = getattr(user_api_key_dict, "team_id", None)
|
||||
tid = user_api_key_dict.team_id
|
||||
value = str(tid) if tid else None
|
||||
|
||||
else:
|
||||
|
|
@ -568,7 +587,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
return value
|
||||
|
||||
# Final fallback for service accounts with no user identity
|
||||
token: Final = getattr(user_api_key_dict, "token", None) or getattr(user_api_key_dict, "api_key", None)
|
||||
token: Final = user_api_key_dict.token or user_api_key_dict.api_key
|
||||
if token:
|
||||
return "apikey:" + hashlib.sha256(str(token).encode()).hexdigest()[:16]
|
||||
return "litellm-proxy"
|
||||
|
|
@ -615,7 +634,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# FR-13: Claim operations
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _apply_claim_operations(self, claims: dict[str, Any]) -> dict[str, Any]:
|
||||
def _apply_claim_operations(self, claims: dict[str, object]) -> dict[str, object]:
|
||||
"""Apply add_claims, set_claims, and remove_claims to the claim dict."""
|
||||
# add_claims: insert only when key is absent
|
||||
for k, v in self.add_claims.items():
|
||||
|
|
@ -637,9 +656,9 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
|
||||
def _passthrough_optional_claims(
|
||||
self,
|
||||
claims: dict[str, Any],
|
||||
jwt_claims: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
claims: dict[str, object],
|
||||
jwt_claims: Mapping[str, object] | None,
|
||||
) -> dict[str, object]:
|
||||
"""Forward optional_claims from verified incoming token into the outbound JWT."""
|
||||
if not self.optional_claims or not jwt_claims:
|
||||
return claims
|
||||
|
|
@ -656,7 +675,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
data: dict,
|
||||
jwt_claims: dict[str, Any] | None = None,
|
||||
jwt_claims: Mapping[str, object] | None = None,
|
||||
call_type: CallTypesLiteral | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
|
|
@ -669,7 +688,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
jwt_claims if available. None for pure API-key requests.
|
||||
"""
|
||||
now: Final = int(time.time())
|
||||
claims: dict[str, Any] = {
|
||||
claims: dict[str, object] = {
|
||||
"iss": self.issuer,
|
||||
"aud": self.audience,
|
||||
"iat": now,
|
||||
|
|
@ -681,18 +700,18 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
claims["sub"] = self._resolve_end_user_identity(user_api_key_dict, jwt_claims)
|
||||
|
||||
# email passthrough when available from LiteLLM context
|
||||
user_email: Final = getattr(user_api_key_dict, "user_email", None)
|
||||
user_email: Final = user_api_key_dict.user_email
|
||||
if user_email:
|
||||
claims["email"] = user_email
|
||||
|
||||
# act — RFC 8693 delegation claim (team/org context)
|
||||
team_id: Final = getattr(user_api_key_dict, "team_id", None)
|
||||
org_id: Final = getattr(user_api_key_dict, "org_id", None)
|
||||
team_id: Final = user_api_key_dict.team_id
|
||||
org_id: Final = user_api_key_dict.org_id
|
||||
act_sub: Final = team_id or org_id or "litellm-proxy"
|
||||
claims["act"] = {"sub": act_sub}
|
||||
|
||||
# end_user_id when set separately from user_id
|
||||
end_user_id: Final = getattr(user_api_key_dict, "end_user_id", None)
|
||||
end_user_id: Final = user_api_key_dict.end_user_id
|
||||
if end_user_id:
|
||||
claims["end_user_id"] = end_user_id
|
||||
|
||||
|
|
@ -710,8 +729,8 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
|
||||
def _build_channel_token_claims(
|
||||
self,
|
||||
base_claims: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
base_claims: Mapping[str, object],
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Build claims for the channel token (FR-14 two-token model).
|
||||
|
||||
|
|
@ -776,7 +795,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# ------------------------------------------------------------------
|
||||
# FR-5: Verify incoming token before re-signing
|
||||
# ------------------------------------------------------------------
|
||||
jwt_claims: dict[str, Any] | None = None
|
||||
jwt_claims: dict[str, object] | None = None
|
||||
raw_token: Final[str | None] = hook_data.get("incoming_bearer_token")
|
||||
|
||||
if self.access_token_discovery_uri and raw_token:
|
||||
|
|
@ -810,7 +829,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
|
||||
# Fall back to LiteLLM-decoded JWT claims (available when proxy uses JWT auth).
|
||||
if jwt_claims is None:
|
||||
jwt_claims = getattr(user_api_key_dict, "jwt_claims", None)
|
||||
jwt_claims = user_api_key_dict.jwt_claims
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# FR-15: Validate required claims
|
||||
|
|
@ -896,7 +915,7 @@ async def inject_mcp_jwt_headers_for_upstream(
|
|||
if auth_hdr.lower().startswith("bearer "):
|
||||
incoming_bearer_token = auth_hdr[len("bearer ") :]
|
||||
|
||||
hook_data: Final[dict[str, Any]] = {
|
||||
hook_data: Final = {
|
||||
"mcp_tool_name": "" if for_list_tools else mcp_tool_name,
|
||||
"incoming_bearer_token": incoming_bearer_token,
|
||||
"extra_headers": merged,
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ Provides real-time threat detection, DLP, URL filtering, content masking, and po
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import AsyncIterable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
from urllib.parse import urlparse
|
||||
|
|
@ -166,7 +167,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
GuardrailEventHooks.during_mcp_call: GuardrailEventHooks.during_call,
|
||||
}
|
||||
|
||||
def should_run_guardrail(self, data: Any, event_type: GuardrailEventHooks) -> bool:
|
||||
def should_run_guardrail(self, data: Mapping[str, object], event_type: GuardrailEventHooks) -> bool:
|
||||
if super().should_run_guardrail(data, event_type):
|
||||
return True
|
||||
compat: Final = self._MCP_COMPAT_MAP.get(event_type)
|
||||
|
|
@ -175,7 +176,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return True
|
||||
return False
|
||||
|
||||
def _extract_text_from_messages(self, messages: list[dict[str, Any]]) -> str:
|
||||
def _extract_text_from_messages(self, messages: Sequence[Mapping[str, object]]) -> str:
|
||||
"""Extract text content from messages array."""
|
||||
if not isinstance(messages, list) or not messages:
|
||||
return ""
|
||||
|
|
@ -242,10 +243,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
self,
|
||||
content: str = "",
|
||||
is_response: bool = False,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
call_id: str | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
call_id: object = None,
|
||||
tool_event: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""Call PANW Prisma AIRS API to scan content or a tool_event."""
|
||||
|
||||
if tool_event is None and not content.strip():
|
||||
|
|
@ -275,7 +276,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
else:
|
||||
app_name_value = self.app_name # Defaults to "LiteLLM"
|
||||
|
||||
panw_metadata: Final = {
|
||||
panw_metadata: Final[dict[str, object]] = {
|
||||
"app_user": (
|
||||
(metadata.get("app_user") or metadata.get("user") or "litellm_user") if metadata else "litellm_user"
|
||||
),
|
||||
|
|
@ -295,13 +296,13 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
panw_metadata["litellm_trace_id"] = metadata["litellm_trace_id"]
|
||||
|
||||
# Build contents: tool_event takes priority, else prompt/response text
|
||||
contents: list[dict[str, Any]]
|
||||
contents: Sequence[Mapping[str, object]]
|
||||
if tool_event is not None:
|
||||
contents = [{"tool_event": tool_event}]
|
||||
else:
|
||||
contents = [{"response" if is_response else "prompt": content}]
|
||||
|
||||
payload: Final = {
|
||||
payload: Final[dict[str, object]] = {
|
||||
"metadata": panw_metadata,
|
||||
"contents": contents,
|
||||
}
|
||||
|
|
@ -325,7 +326,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
# If neither profile_name nor profile_id is provided, PANW API will use the
|
||||
# profile linked to the API key (if configured in Strata Cloud Manager)
|
||||
if profile_name or profile_id:
|
||||
ai_profile: Final = {}
|
||||
ai_profile: Final[dict[str, object]] = {}
|
||||
if profile_id:
|
||||
ai_profile["profile_id"] = profile_id
|
||||
if profile_name:
|
||||
|
|
@ -333,7 +334,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
payload["ai_profile"] = ai_profile
|
||||
|
||||
if is_response and tool_event is None:
|
||||
payload["metadata"]["is_response"] = True
|
||||
panw_metadata["is_response"] = True
|
||||
|
||||
headers: Final = {
|
||||
"Content-Type": "application/json",
|
||||
|
|
@ -355,7 +356,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
response.raise_for_status()
|
||||
|
||||
result: Final = response.json()
|
||||
result: Final[dict[str, object]] = response.json()
|
||||
|
||||
# Validate response format
|
||||
if "action" not in result:
|
||||
|
|
@ -489,7 +490,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
return "unknown"
|
||||
|
||||
def _get_masked_text(self, scan_result: dict[str, Any], is_response: bool = False) -> str | None:
|
||||
def _get_masked_text(self, scan_result: Mapping[str, object], is_response: bool = False) -> str | None:
|
||||
"""Extract masked text from PANW scan result."""
|
||||
masked_key: Final = "response_masked_data" if is_response else "prompt_masked_data"
|
||||
masked_data: Final = scan_result.get(masked_key)
|
||||
|
|
@ -511,7 +512,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
@staticmethod
|
||||
def _apply_mcp_masking(
|
||||
request_data: dict,
|
||||
original_args: Any,
|
||||
original_args: object,
|
||||
masked_text: str,
|
||||
*,
|
||||
is_blocked: bool = True,
|
||||
|
|
@ -544,7 +545,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
# If the original args were structured, preserve the type.
|
||||
if isinstance(original_args, (dict, list)):
|
||||
try:
|
||||
parsed: Final = json.loads(masked_text)
|
||||
parsed: Final[object] = json.loads(masked_text)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
@ -556,7 +557,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
}
|
||||
},
|
||||
)
|
||||
masked_value: Any = parsed
|
||||
masked_value: object = parsed
|
||||
else:
|
||||
masked_value = masked_text
|
||||
|
||||
|
|
@ -572,7 +573,9 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
else:
|
||||
verbose_proxy_logger.info("PANW Prisma AIRS: MCP request allowed with PII masking applied")
|
||||
|
||||
def _apply_masking_to_messages(self, messages: list[dict[str, Any]], masked_text: str) -> list[dict[str, Any]]:
|
||||
def _apply_masking_to_messages(
|
||||
self, messages: list[dict[str, object]], masked_text: str
|
||||
) -> Sequence[Mapping[str, object]]:
|
||||
"""Apply masked text to the last user message."""
|
||||
if not messages:
|
||||
return messages
|
||||
|
|
@ -622,7 +625,9 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
if hasattr(choice.message.function_call, "arguments"):
|
||||
choice.message.function_call.arguments = masked_text
|
||||
|
||||
def _build_error_detail(self, scan_result: dict[str, Any], is_response: bool = False) -> dict[str, Any]:
|
||||
def _build_error_detail(
|
||||
self, scan_result: Mapping[str, object], is_response: bool = False
|
||||
) -> Mapping[str, Mapping[str, object]]:
|
||||
"""Build enhanced error detail with scan information."""
|
||||
action_type: Final = "Response" if is_response else "Prompt"
|
||||
code_suffix: Final = "_response_blocked" if is_response else "_blocked"
|
||||
|
|
@ -642,7 +647,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
},
|
||||
)
|
||||
|
||||
error_detail: Final = {
|
||||
error_detail: Final[dict[str, dict[str, object]]] = {
|
||||
"error": {
|
||||
"message": error_msg,
|
||||
"type": "guardrail_violation",
|
||||
|
|
@ -672,12 +677,12 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
def _handle_api_error_with_logging(
|
||||
self,
|
||||
scan_result: dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
scan_result: dict[str, object],
|
||||
data: dict[str, object],
|
||||
start_time: datetime,
|
||||
event_type: GuardrailEventHooks,
|
||||
is_response: bool = False,
|
||||
) -> dict[str, Any] | None:
|
||||
) -> None:
|
||||
"""Handle API errors with fail-open/fail-closed logic."""
|
||||
end_time: Final = datetime.now()
|
||||
duration: Final = (end_time - start_time).total_seconds()
|
||||
|
|
@ -722,7 +727,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=f"{self.guardrail_name}:unscanned"
|
||||
)
|
||||
return None
|
||||
return
|
||||
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
|
|
@ -783,7 +788,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return metadata
|
||||
|
||||
@staticmethod
|
||||
def _extract_text_from_sse_bytes(chunks: list[bytes]) -> str:
|
||||
def _extract_text_from_sse_bytes(chunks: Sequence[bytes]) -> str:
|
||||
"""Extract text from Anthropic SSE byte chunks (content_block_delta → text_delta)."""
|
||||
texts: Final[list[str]] = []
|
||||
raw: Final = b"".join(chunks).decode("utf-8", errors="replace")
|
||||
|
|
@ -804,7 +809,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return "".join(texts)
|
||||
|
||||
@staticmethod
|
||||
def _extract_text_from_streaming_events(chunks: list) -> str:
|
||||
def _extract_text_from_streaming_events(chunks: Sequence[object]) -> str:
|
||||
"""Extract text from /v1/responses streaming events (object or dict)."""
|
||||
|
||||
def _attr(c, key):
|
||||
|
|
@ -960,7 +965,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
cache: DualCache,
|
||||
data: dict[str, Any],
|
||||
call_type: CallTypesLiteral,
|
||||
) -> dict[str, Any] | None:
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Pre-call hook to scan user prompts before sending to LLM.
|
||||
|
||||
|
|
@ -1075,10 +1080,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
@log_guardrail_information
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
data: dict[str, Any],
|
||||
data: dict[str, object],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
) -> Any:
|
||||
response: object,
|
||||
) -> object:
|
||||
"""
|
||||
Post-call hook to scan LLM responses before returning to user.
|
||||
|
||||
|
|
@ -1193,7 +1198,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
assembled_model_response: ModelResponse,
|
||||
request_data: dict,
|
||||
start_time: datetime,
|
||||
) -> tuple[bool, ModelResponse, dict[str, Any]]:
|
||||
) -> tuple[bool, ModelResponse, dict[str, object]]:
|
||||
"""
|
||||
Scan assembled streaming response and apply masking if needed.
|
||||
Returns (content_was_modified, response, scan_result).
|
||||
|
|
@ -1255,8 +1260,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
request_data: dict,
|
||||
response: AsyncIterable[object],
|
||||
request_data: dict[str, object],
|
||||
):
|
||||
"""
|
||||
Process streaming response chunks and scan the assembled response.
|
||||
|
|
@ -1367,7 +1372,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
# returns a proper JSON error response with the correct status code.
|
||||
# (Raising from a generator hits create_response's generic except → 500.)
|
||||
detail: Final = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)}
|
||||
error_obj: Final[dict[str, Any]] = dict(detail.get("error", detail))
|
||||
error_obj: Final[dict[str, object]] = dict(detail.get("error", detail))
|
||||
error_obj["code"] = e.status_code
|
||||
yield f"data: {json.dumps({'error': error_obj})}\n\n"
|
||||
except Exception as e:
|
||||
|
|
@ -1378,8 +1383,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
self,
|
||||
tool_calls: list,
|
||||
is_response: bool,
|
||||
metadata: dict[str, Any],
|
||||
call_id: str,
|
||||
metadata: Mapping[str, object],
|
||||
call_id: object,
|
||||
request_data: dict,
|
||||
start_time: datetime,
|
||||
) -> None:
|
||||
|
|
@ -1416,7 +1421,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
tool_name = func.get("name")
|
||||
|
||||
# --- build tool_event payload (canonical PANW schema) -----------
|
||||
tool_event: dict[str, Any] = {
|
||||
tool_event: dict[str, object] = {
|
||||
"metadata": {
|
||||
"ecosystem": "openai",
|
||||
"method": "tools/call",
|
||||
|
|
@ -1472,7 +1477,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
@staticmethod
|
||||
def _is_anthropic_request(
|
||||
request_data: dict,
|
||||
request_data: Mapping[str, object],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> bool:
|
||||
"""Detect if the current request is an Anthropic /v1/messages call."""
|
||||
|
|
@ -1497,7 +1502,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
def _use_latest_user_only(
|
||||
self,
|
||||
request_data: dict,
|
||||
request_data: Mapping[str, object],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> bool:
|
||||
"""Resolve whether to scan only the latest user message.
|
||||
|
|
@ -1515,8 +1520,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
@staticmethod
|
||||
def _get_latest_user_text_indices(
|
||||
texts: list[str],
|
||||
messages: list,
|
||||
texts: Sequence[str],
|
||||
messages: Sequence[object],
|
||||
) -> set | None:
|
||||
"""Return text indices belonging to only the latest scannable human-authored (user or developer) message.
|
||||
|
||||
|
|
@ -1569,8 +1574,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
@staticmethod
|
||||
def _get_scannable_text_indices(
|
||||
texts: list[str],
|
||||
structured_messages: list,
|
||||
texts: Sequence[str],
|
||||
structured_messages: Sequence[object],
|
||||
) -> set | None:
|
||||
"""Derive which ``texts`` indices originate from user/system messages.
|
||||
|
||||
|
|
@ -1627,7 +1632,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
|
|
@ -1798,7 +1803,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
# "mcp_tool_name"/"mcp_arguments". Check canonical first, then fallback.
|
||||
mcp_tool_name: Final = request_data.get("mcp_tool_name") or self._mcp_name_fallback(request_data)
|
||||
if mcp_tool_name and input_type == "request":
|
||||
mcp_tool_event: Final[dict[str, Any]] = {
|
||||
mcp_tool_event: Final[dict[str, object]] = {
|
||||
"metadata": {
|
||||
"ecosystem": "mcp",
|
||||
"method": "tools/call",
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Pre-call hook that filters MCP tools semantically before LLM inference.
|
|||
Reduces context window size and improves tool selection accuracy.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -164,7 +165,7 @@ class SemanticToolFilterHook(CustomLogger):
|
|||
return [name for name in names if name]
|
||||
|
||||
@staticmethod
|
||||
def _narrow_mcp_references(tools: list[Any], selected_tool_names: list[str]) -> list[Any]:
|
||||
def _narrow_mcp_references(tools: Sequence[Mapping[str, object]], selected_tool_names: list[str]) -> list[object]:
|
||||
"""
|
||||
Restrict each litellm_proxy MCP reference to the semantically selected tools.
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import re
|
|||
import time
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
|
@ -66,6 +67,32 @@ _SESSION_ID_VALUE_RE: Final = re.compile(r"^[a-zA-Z0-9_\-]{8,}$")
|
|||
|
||||
_SHA256_HEX_RE: Final = re.compile(r"^[0-9a-f]{64}$")
|
||||
|
||||
# W3C Trace Context traceparent header: https://www.w3.org/TR/trace-context/
|
||||
# e.g. "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"
|
||||
_TRACEPARENT_RE: Final = re.compile(r"^[0-9a-f]{2}-([0-9a-f]{32})-[0-9a-f]{16}-[0-9a-f]{2}$", re.IGNORECASE)
|
||||
|
||||
|
||||
def _trace_id_from_traceparent(traceparent: str) -> str | None:
|
||||
"""Extract the trace-id from a W3C Trace Context traceparent header, e.g.
|
||||
"00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01" -> the 32-hex
|
||||
trace-id in the middle. An all-zero trace-id is invalid per spec and is
|
||||
rejected, matching how the OpenTelemetry SDK itself treats it."""
|
||||
match: Final = _TRACEPARENT_RE.match(traceparent.strip())
|
||||
if not match:
|
||||
return None
|
||||
trace_id: Final = match.group(1).lower()
|
||||
return trace_id if trace_id != "0" * 32 else None
|
||||
|
||||
|
||||
def _session_id_from_baggage(baggage: str) -> str | None:
|
||||
"""Extract a session.id entry from a W3C Baggage header
|
||||
(https://www.w3.org/TR/baggage/), e.g. "session.id=abc-123,user.id=42"."""
|
||||
for pair in baggage.split(","):
|
||||
key, _, value = pair.strip().partition("=")
|
||||
if key.strip() == "session.id" and value.strip():
|
||||
return value.strip()
|
||||
return None
|
||||
|
||||
|
||||
def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
||||
"""Only proxy-validated keys are stamped, proven by the unforgeable
|
||||
|
|
@ -210,6 +237,11 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = (
|
|||
"_code_interpreter_interception_sandbox_key",
|
||||
"_code_interpreter_interception_session_scoped",
|
||||
"max_agentic_loops",
|
||||
# Recomputed below from the actual caller-controlled timeout sources (headers and
|
||||
# body fields); a client-forged value here would let a request either dodge cooldown
|
||||
# protection on a real deployment failure or force a false "not caller-controlled"
|
||||
# reading that lets its own bad timeout cool down deployments other tenants rely on.
|
||||
"client_side_timeout",
|
||||
)
|
||||
|
||||
_UNTRUSTED_METADATA_CONTROL_FIELDS: Final = (
|
||||
|
|
@ -1035,6 +1067,7 @@ class LiteLLMProxyRequestSetup:
|
|||
def add_litellm_data_for_backend_llm_call(
|
||||
*,
|
||||
headers: dict,
|
||||
request_data: Mapping[str, Any],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
general_settings: dict[str, Any] | None = None,
|
||||
) -> LitellmDataForBackendLLMCall:
|
||||
|
|
@ -1053,13 +1086,29 @@ class LiteLLMProxyRequestSetup:
|
|||
if _organization is not None:
|
||||
data["organization"] = _organization
|
||||
|
||||
timeout: Final = LiteLLMProxyRequestSetup._get_timeout_from_request(headers)
|
||||
if timeout is not None:
|
||||
data["timeout"] = timeout
|
||||
header_timeout: Final = LiteLLMProxyRequestSetup._get_timeout_from_request(headers)
|
||||
if header_timeout is not None:
|
||||
data["timeout"] = header_timeout
|
||||
|
||||
stream_timeout: Final = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers)
|
||||
if stream_timeout is not None:
|
||||
data["stream_timeout"] = stream_timeout
|
||||
header_stream_timeout: Final = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers)
|
||||
if header_stream_timeout is not None:
|
||||
data["stream_timeout"] = header_stream_timeout
|
||||
|
||||
# Router._get_timeout resolves the effective per-attempt timeout from any of
|
||||
# kwargs["timeout"], kwargs["request_timeout"], or kwargs["stream_timeout"], and a
|
||||
# caller can supply any of those via the request body as well as the headers above.
|
||||
# A deliberately tiny value can force a 408 on every deployment in a fallback chain,
|
||||
# so this marker (never trusted verbatim from the client; stripped above) must cover
|
||||
# every source cooldown_handlers._trigger_cooldown_for_failed_deployment needs to
|
||||
# distinguish from a real deployment health signal.
|
||||
if (
|
||||
header_timeout is not None
|
||||
or header_stream_timeout is not None
|
||||
or request_data.get("timeout") is not None
|
||||
or request_data.get("request_timeout") is not None
|
||||
or request_data.get("stream_timeout") is not None
|
||||
):
|
||||
data["client_side_timeout"] = True
|
||||
|
||||
num_retries: Final = LiteLLMProxyRequestSetup._get_num_retries_from_request(headers)
|
||||
if num_retries is not None:
|
||||
|
|
@ -1113,6 +1162,33 @@ class LiteLLMProxyRequestSetup:
|
|||
body_metadata["user_id"] = session_id
|
||||
verbose_proxy_logger.debug("Extracted session_id from Anthropic metadata.user_id")
|
||||
|
||||
# Last-resort fallback: the W3C standards for trace/session propagation
|
||||
# (https://www.w3.org/TR/trace-context/, https://www.w3.org/TR/baggage/).
|
||||
# Lower priority than everything above - only fires when neither the
|
||||
# explicit litellm headers nor the Anthropic-metadata path found
|
||||
# anything - but lets a caller's existing traceparent/baggage headers
|
||||
# (from real OTel instrumentation) correlate with litellm's own logs
|
||||
# instead of generating an unrelated trace_id.
|
||||
normalized_headers: Final = MappingProxyType({k.lower(): v for k, v in headers.items() if isinstance(k, str)})
|
||||
if "litellm_trace_id" not in data:
|
||||
traceparent: Final = normalized_headers.get("traceparent")
|
||||
if isinstance(traceparent, str):
|
||||
trace_id_from_traceparent: Final = _trace_id_from_traceparent(traceparent)
|
||||
if trace_id_from_traceparent:
|
||||
metadata_from_headers["trace_id"] = trace_id_from_traceparent
|
||||
data["litellm_trace_id"] = trace_id_from_traceparent # rebind-ok: data is an out-param
|
||||
verbose_proxy_logger.debug(
|
||||
"Extracted trace_id from W3C traceparent header: %s", trace_id_from_traceparent
|
||||
)
|
||||
if "litellm_session_id" not in data:
|
||||
baggage: Final = normalized_headers.get("baggage")
|
||||
if isinstance(baggage, str):
|
||||
session_id_from_baggage: Final = _session_id_from_baggage(baggage)
|
||||
if session_id_from_baggage:
|
||||
metadata_from_headers["session_id"] = session_id_from_baggage
|
||||
data["litellm_session_id"] = session_id_from_baggage # rebind-ok: data is an out-param
|
||||
verbose_proxy_logger.debug("Extracted session_id from W3C baggage header")
|
||||
|
||||
if isinstance(data[_metadata_variable_name], dict):
|
||||
data[_metadata_variable_name].update(metadata_from_headers)
|
||||
return data
|
||||
|
|
@ -1545,6 +1621,7 @@ async def add_litellm_data_to_request(
|
|||
data.update(
|
||||
LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call(
|
||||
headers=_headers,
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
general_settings=general_settings,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import asyncio
|
|||
import json
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
|
@ -37,8 +37,26 @@ from litellm.types.management_endpoints import (
|
|||
CacheSettingsField,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
class _CacheConfigRow(Protocol):
|
||||
cache_settings: str | Mapping[str, object] | None
|
||||
|
||||
|
||||
class _CacheConfigTable(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, str]) -> _CacheConfigRow | None: ...
|
||||
|
||||
async def upsert(self, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]) -> _CacheConfigRow: ...
|
||||
|
||||
|
||||
def _cache_config_table(prisma_client: "PrismaClient") -> _CacheConfigTable:
|
||||
return CacheConfigRepository(prisma_client).table
|
||||
|
||||
|
||||
# Cache fields holding credentials. Masked on read so plaintext Redis /
|
||||
# Sentinel passwords never leave the server in a GET response. `url` is here
|
||||
# because a Redis/Valkey URL can embed a password inline
|
||||
|
|
@ -197,7 +215,7 @@ def _saved_secret_is_reusable(incoming: Mapping[str, object], saved: Mapping[str
|
|||
return True
|
||||
|
||||
|
||||
def _merge_over_saved(incoming: Mapping[str, object], saved: Mapping[str, object]) -> dict[str, Any]:
|
||||
def _merge_over_saved(incoming: Mapping[str, object], saved: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""Keep the stored secret behind any credential the caller echoed back redacted or omitted.
|
||||
|
||||
GET returns credentials as the marker and the form never re-prefills a
|
||||
|
|
@ -339,7 +357,7 @@ class CacheSettingsManager:
|
|||
return normalized1 == normalized2
|
||||
|
||||
@staticmethod
|
||||
async def init_cache_settings_in_db(prisma_client, proxy_config):
|
||||
async def init_cache_settings_in_db(prisma_client: "PrismaClient", proxy_config):
|
||||
"""
|
||||
Initialize cache settings from database into the router on startup.
|
||||
Only reinitializes if cache params have changed.
|
||||
|
|
@ -349,7 +367,7 @@ class CacheSettingsManager:
|
|||
try:
|
||||
cache_config: Final = await call_with_db_reconnect_retry(
|
||||
prisma_client,
|
||||
lambda: CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"}),
|
||||
lambda: _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"}),
|
||||
reason="init_cache_settings_in_db_lookup_failure",
|
||||
)
|
||||
if cache_config is not None and cache_config.cache_settings:
|
||||
|
|
@ -444,7 +462,7 @@ async def get_cache_settings(
|
|||
# Read the stored settings (decrypted); an env-only cache has none.
|
||||
stored: dict[str, object] = {}
|
||||
if prisma_client is not None:
|
||||
cache_config = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"})
|
||||
cache_config = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"})
|
||||
if cache_config is not None and cache_config.cache_settings:
|
||||
stored = proxy_config._decrypt_db_variables(
|
||||
variables_dict=_parse_stored_settings(cache_config.cache_settings)
|
||||
|
|
@ -511,9 +529,7 @@ async def test_cache_connection(
|
|||
saved_settings: dict[str, object] = {}
|
||||
if prisma_client is not None:
|
||||
try:
|
||||
existing_row: Final = await CacheConfigRepository(prisma_client).table.find_unique(
|
||||
where={"id": "cache_config"}
|
||||
)
|
||||
existing_row: Final = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"})
|
||||
if existing_row is not None and existing_row.cache_settings:
|
||||
saved_settings = proxy_config._decrypt_db_variables(
|
||||
variables_dict=_parse_stored_settings(existing_row.cache_settings)
|
||||
|
|
@ -590,7 +606,7 @@ async def update_cache_settings(
|
|||
try:
|
||||
# Read the stored row first: its decrypted values back any credential the
|
||||
# caller echoed back redacted, and its key set drives the audit diff.
|
||||
existing_row: Final = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"})
|
||||
existing_row: Final = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"})
|
||||
before_settings: dict[str, object] | None = None
|
||||
saved_settings: dict[str, object] = {}
|
||||
if existing_row is not None and existing_row.cache_settings:
|
||||
|
|
@ -606,7 +622,7 @@ async def update_cache_settings(
|
|||
encrypted_settings: Final = proxy_config._encrypt_env_variables(environment_variables=cache_settings)
|
||||
|
||||
# Save to database
|
||||
await CacheConfigRepository(prisma_client).table.upsert(
|
||||
await _cache_config_table(prisma_client).upsert(
|
||||
where={"id": "cache_config"},
|
||||
data={
|
||||
"create": {
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from fastapi import HTTPException, status
|
|||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import PTU_SENTINEL_API_KEY
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.repositories.table_repositories import DeletedVerificationTokenRepository
|
||||
|
|
@ -150,6 +151,7 @@ def update_metrics(existing_metrics: SpendMetrics, record: DailySpendRecord) ->
|
|||
prompt_tokens: Final = record.prompt_tokens or 0
|
||||
completion_tokens: Final = record.completion_tokens or 0
|
||||
existing_metrics.spend += record.spend or 0.0
|
||||
existing_metrics.flat_cost += getattr(record, "ptu_flat_cost", None) or 0.0
|
||||
existing_metrics.prompt_tokens += prompt_tokens
|
||||
existing_metrics.completion_tokens += completion_tokens
|
||||
existing_metrics.total_tokens += prompt_tokens + completion_tokens
|
||||
|
|
@ -208,30 +210,43 @@ def update_breakdown_metrics(
|
|||
entity_id_field: str | None = None,
|
||||
entity_metadata_field: Mapping[str, dict[str, object]] | None = None,
|
||||
) -> BreakdownMetrics:
|
||||
"""Updates breakdown metrics for a single record using the existing update_metrics function"""
|
||||
"""Updates breakdown metrics for a single record using the existing update_metrics function.
|
||||
|
||||
PTU sentinel rows (api_key == PTU_SENTINEL_API_KEY) add their flat cost to every
|
||||
parent bucket but never appear as an api_key row, and are kept out of the
|
||||
per-request provider breakdown."""
|
||||
|
||||
is_ptu_sentinel: Final = record.api_key == PTU_SENTINEL_API_KEY
|
||||
|
||||
# A PTU sentinel row keys on the deployment id so a rename cannot move it, and carries
|
||||
# the operator-facing name in model_group. The breakdown key is rendered directly as a
|
||||
# label, so display the name; two deployments sharing one name merge here, which is
|
||||
# what the write path used to do by collapsing them into a single row.
|
||||
model_key: Final = (record.model_group or record.model) if is_ptu_sentinel else record.model
|
||||
|
||||
# Update model breakdown
|
||||
if record.model and record.model not in breakdown.models:
|
||||
breakdown.models[record.model] = MetricWithMetadata(
|
||||
if model_key and model_key not in breakdown.models:
|
||||
breakdown.models[model_key] = MetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=model_metadata.get(record.model, {}), # Add any model-specific metadata here
|
||||
metadata=model_metadata.get(model_key, {}), # Add any model-specific metadata here
|
||||
)
|
||||
if record.model:
|
||||
breakdown.models[record.model].metrics = update_metrics(breakdown.models[record.model].metrics, record)
|
||||
if model_key:
|
||||
breakdown.models[model_key].metrics = update_metrics(breakdown.models[model_key].metrics, record)
|
||||
|
||||
# Update API key breakdown for this model
|
||||
if record.api_key not in breakdown.models[record.model].api_key_breakdown:
|
||||
breakdown.models[record.model].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
if not is_ptu_sentinel:
|
||||
# Update API key breakdown for this model
|
||||
if record.api_key not in breakdown.models[model_key].api_key_breakdown:
|
||||
breakdown.models[model_key].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
)
|
||||
breakdown.models[model_key].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.models[model_key].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
breakdown.models[record.model].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.models[record.model].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
|
||||
# Update model group breakdown
|
||||
model_group_key: Final = record.model_group or record.model
|
||||
|
|
@ -245,19 +260,20 @@ def update_breakdown_metrics(
|
|||
breakdown.model_groups[model_group_key].metrics, record
|
||||
)
|
||||
|
||||
# Update API key breakdown for this model
|
||||
if record.api_key not in breakdown.model_groups[model_group_key].api_key_breakdown:
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
if not is_ptu_sentinel:
|
||||
# Update API key breakdown for this model
|
||||
if record.api_key not in breakdown.model_groups[model_group_key].api_key_breakdown:
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
)
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
|
||||
if record.mcp_namespaced_tool_name:
|
||||
if record.mcp_namespaced_tool_name not in breakdown.mcp_servers:
|
||||
|
|
@ -288,28 +304,29 @@ def update_breakdown_metrics(
|
|||
record,
|
||||
)
|
||||
|
||||
# Update provider breakdown
|
||||
provider: Final = record.custom_llm_provider or "unknown"
|
||||
if provider not in breakdown.providers:
|
||||
breakdown.providers[provider] = MetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=provider_metadata.get(provider, {}), # Add any provider-specific metadata here
|
||||
)
|
||||
breakdown.providers[provider].metrics = update_metrics(breakdown.providers[provider].metrics, record)
|
||||
if not is_ptu_sentinel:
|
||||
# Update provider breakdown
|
||||
provider: Final = record.custom_llm_provider or "unknown"
|
||||
if provider not in breakdown.providers:
|
||||
breakdown.providers[provider] = MetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=provider_metadata.get(provider, {}), # Add any provider-specific metadata here
|
||||
)
|
||||
breakdown.providers[provider].metrics = update_metrics(breakdown.providers[provider].metrics, record)
|
||||
|
||||
# Update API key breakdown for this provider
|
||||
if record.api_key not in breakdown.providers[provider].api_key_breakdown:
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
# Update API key breakdown for this provider
|
||||
if record.api_key not in breakdown.providers[provider].api_key_breakdown:
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
)
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
|
||||
# Update endpoint breakdown
|
||||
if record.endpoint:
|
||||
|
|
@ -336,16 +353,17 @@ def update_breakdown_metrics(
|
|||
record,
|
||||
)
|
||||
|
||||
# Update api key breakdown
|
||||
if record.api_key not in breakdown.api_keys:
|
||||
breakdown.api_keys[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
), # Add any api_key-specific metadata here
|
||||
)
|
||||
breakdown.api_keys[record.api_key].metrics = update_metrics(breakdown.api_keys[record.api_key].metrics, record)
|
||||
if not is_ptu_sentinel:
|
||||
# Update api key breakdown
|
||||
if record.api_key not in breakdown.api_keys:
|
||||
breakdown.api_keys[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
), # Add any api_key-specific metadata here
|
||||
)
|
||||
breakdown.api_keys[record.api_key].metrics = update_metrics(breakdown.api_keys[record.api_key].metrics, record)
|
||||
|
||||
# Update entity-specific metrics if entity_id_field is provided
|
||||
if entity_id_field:
|
||||
|
|
@ -358,19 +376,20 @@ def update_breakdown_metrics(
|
|||
)
|
||||
breakdown.entities[entity_value].metrics = update_metrics(breakdown.entities[entity_value].metrics, record)
|
||||
|
||||
# Update API key breakdown for this entity
|
||||
if record.api_key not in breakdown.entities[entity_value].api_key_breakdown:
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
if not is_ptu_sentinel:
|
||||
# Update API key breakdown for this entity
|
||||
if record.api_key not in breakdown.entities[entity_value].api_key_breakdown:
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
)
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
|
||||
return breakdown
|
||||
|
||||
|
|
@ -599,6 +618,14 @@ def _build_aggregated_sql_query(
|
|||
# total_successful_requests metadata they feed) once the admin UI reads SGR
|
||||
# only from LiteLLM_DailyGatewayRequests. The remaining spend, token and
|
||||
# api_requests rollups are still served from here.
|
||||
#
|
||||
# Only LiteLLM_DailyTeamSpend carries ptu_flat_cost; other daily tables emit a
|
||||
# constant zero so the SpendMetrics.flat_cost response shape stays uniform.
|
||||
ptu_flat_cost_select: Final = (
|
||||
"SUM(ptu_flat_cost)::float AS ptu_flat_cost"
|
||||
if table_name == "litellm_dailyteamspend"
|
||||
else "0::float AS ptu_flat_cost"
|
||||
)
|
||||
sql_query: Final = f"""
|
||||
SELECT
|
||||
date,
|
||||
|
|
@ -612,6 +639,7 @@ def _build_aggregated_sql_query(
|
|||
custom_llm_provider, mcp_namespaced_tool_name,
|
||||
endpoint) AS group_level,
|
||||
SUM(spend)::float AS spend,
|
||||
{ptu_flat_cost_select},
|
||||
SUM(prompt_tokens)::bigint AS prompt_tokens,
|
||||
SUM(completion_tokens)::bigint AS completion_tokens,
|
||||
SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens,
|
||||
|
|
@ -707,7 +735,9 @@ async def _aggregate_spend_records(
|
|||
The per-row loop is offloaded to a worker thread via asyncio.to_thread so
|
||||
a large result set doesn't peg the event loop.
|
||||
"""
|
||||
api_keys: Final[set[str]] = {record.api_key for record in records if record.api_key}
|
||||
api_keys: Final[set[str]] = {
|
||||
record.api_key for record in records if record.api_key and record.api_key != PTU_SENTINEL_API_KEY
|
||||
}
|
||||
|
||||
api_key_metadata: dict[str, _KeyMetadataDict] = {}
|
||||
if api_keys:
|
||||
|
|
@ -754,6 +784,7 @@ def _record_to_spend_metrics(record: _GroupingSetsRow) -> SpendMetrics:
|
|||
completion_tokens: Final = record.completion_tokens or 0
|
||||
return SpendMetrics(
|
||||
spend=record.spend or 0.0,
|
||||
flat_cost=getattr(record, "ptu_flat_cost", None) or 0.0,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
|
|
@ -820,6 +851,7 @@ def _aggregate_grouping_sets_records_sync(
|
|||
for record in records:
|
||||
level = record.group_level
|
||||
metrics = _record_to_spend_metrics(record)
|
||||
is_ptu_sentinel = record.api_key == PTU_SENTINEL_API_KEY
|
||||
|
||||
if level == _GROUP_GRAND_TOTAL:
|
||||
total_metrics = metrics
|
||||
|
|
@ -832,7 +864,7 @@ def _aggregate_grouping_sets_records_sync(
|
|||
breakdown = ensure_date(record.date)["breakdown"]
|
||||
|
||||
if level == _GROUP_DATE_API_KEY:
|
||||
if record.api_key:
|
||||
if record.api_key and not is_ptu_sentinel:
|
||||
breakdown.api_keys[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=metrics,
|
||||
metadata=_key_metadata(api_key_metadata, record.api_key),
|
||||
|
|
@ -841,13 +873,13 @@ def _aggregate_grouping_sets_records_sync(
|
|||
if record.model:
|
||||
assign_metric_with_metadata(breakdown.models, record.model, metrics)
|
||||
elif level == _GROUP_DATE_MODEL_API_KEY:
|
||||
if record.model and record.api_key:
|
||||
if record.model and record.api_key and not is_ptu_sentinel:
|
||||
assign_api_key_breakdown(breakdown.models, record.model, record.api_key, metrics)
|
||||
elif level == _GROUP_DATE_MODEL_GROUP:
|
||||
if record.model_group:
|
||||
assign_metric_with_metadata(breakdown.model_groups, record.model_group, metrics)
|
||||
elif level == _GROUP_DATE_MODEL_GROUP_API_KEY:
|
||||
if record.model_group and record.api_key:
|
||||
if record.model_group and record.api_key and not is_ptu_sentinel:
|
||||
assign_api_key_breakdown(
|
||||
breakdown.model_groups,
|
||||
record.model_group,
|
||||
|
|
@ -855,10 +887,17 @@ def _aggregate_grouping_sets_records_sync(
|
|||
metrics,
|
||||
)
|
||||
elif level == _GROUP_DATE_PROVIDER:
|
||||
# Only PTU sentinel rows carry ptu_flat_cost and they have no provider, so at
|
||||
# this level the sentinel's cost would land under "unknown". Withholding the
|
||||
# flat cost matches the per-row path, which skips sentinel rows outright. The
|
||||
# bucket itself is still assigned unconditionally: a legacy row predating the
|
||||
# api_requests column backfills to all zeroes, and skipping those would drop a
|
||||
# provider the base build reported.
|
||||
provider_metrics = metrics.model_copy(update={"flat_cost": 0.0}) # mutable-ok: pydantic update payload
|
||||
provider = record.custom_llm_provider or "unknown"
|
||||
assign_metric_with_metadata(breakdown.providers, provider, metrics)
|
||||
assign_metric_with_metadata(breakdown.providers, provider, provider_metrics)
|
||||
elif level == _GROUP_DATE_PROVIDER_API_KEY:
|
||||
if record.api_key:
|
||||
if record.api_key and not is_ptu_sentinel:
|
||||
provider = record.custom_llm_provider or "unknown"
|
||||
assign_api_key_breakdown(breakdown.providers, provider, record.api_key, metrics)
|
||||
elif level == _GROUP_DATE_MCP:
|
||||
|
|
@ -898,7 +937,7 @@ async def _aggregate_grouping_sets_records(
|
|||
records: Sequence[_GroupingSetsRow],
|
||||
) -> _AggregatedSpendData:
|
||||
"""Async wrapper: fetch api_key_metadata, then dispatch on a worker thread."""
|
||||
api_keys: Final[set[str]] = {r.api_key for r in records if r.api_key}
|
||||
api_keys: Final[set[str]] = {r.api_key for r in records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY}
|
||||
|
||||
api_key_metadata: dict[str, _KeyMetadataDict] = {}
|
||||
if api_keys:
|
||||
|
|
@ -1008,6 +1047,7 @@ async def get_daily_activity(
|
|||
results=aggregated["results"],
|
||||
metadata=DailySpendMetadata(
|
||||
total_spend=metadata_metrics.spend,
|
||||
total_flat_cost=metadata_metrics.flat_cost,
|
||||
total_prompt_tokens=metadata_metrics.prompt_tokens,
|
||||
total_completion_tokens=metadata_metrics.completion_tokens,
|
||||
total_tokens=metadata_metrics.total_tokens,
|
||||
|
|
@ -1098,6 +1138,7 @@ async def get_daily_activity_aggregated(
|
|||
results=aggregated["results"],
|
||||
metadata=DailySpendMetadata(
|
||||
total_spend=aggregated["totals"].spend,
|
||||
total_flat_cost=aggregated["totals"].flat_cost,
|
||||
total_prompt_tokens=aggregated["totals"].prompt_tokens,
|
||||
total_completion_tokens=aggregated["totals"].completion_tokens,
|
||||
total_tokens=aggregated["totals"].total_tokens,
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import datetime
|
|||
import json
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from json import JSONDecodeError
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, Protocol, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
||||
|
|
@ -79,6 +80,7 @@ from litellm.types.router import (
|
|||
SPECIAL_MODEL_INFO_PARAMS,
|
||||
Deployment,
|
||||
GenericLiteLLMParams,
|
||||
ModelInfo,
|
||||
updateDeployment,
|
||||
)
|
||||
from litellm.utils import get_utc_datetime
|
||||
|
|
@ -233,6 +235,96 @@ def _raise_on_strategy_router_write_violation(
|
|||
)
|
||||
|
||||
|
||||
_PTU_MODEL_INFO_FIELDS: Final = ("ptu_count", "cost_per_ptu_per_hour", "ptu_effective_from", "ptu_effective_to")
|
||||
|
||||
|
||||
def _explicitly_cleared_ptu_fields(model_info: ModelInfo | None) -> frozenset[str]:
|
||||
"""The PTU fields a patch sends as an explicit null, which update_db_model drops."""
|
||||
if model_info is None:
|
||||
return frozenset()
|
||||
return frozenset(
|
||||
field
|
||||
for field in _PTU_MODEL_INFO_FIELDS
|
||||
if field in model_info.model_fields_set and getattr(model_info, field) is None
|
||||
)
|
||||
|
||||
|
||||
def _merged_ptu_model_info(*, db_model: Deployment, patch_data: updateDeployment) -> Mapping[str, object]:
|
||||
"""The model_info a patch would store, which is the stored blob updated by the patch.
|
||||
|
||||
A PTU invariant holds over the deployment as it will exist, not over whichever subset
|
||||
of fields a caller happened to send.
|
||||
"""
|
||||
empty: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
stored: Final = db_model.model_info.model_dump(exclude_none=True) if db_model.model_info else empty
|
||||
incoming: Final = patch_data.model_info.model_dump(exclude_none=True) if patch_data.model_info else empty
|
||||
cleared: Final = _explicitly_cleared_ptu_fields(patch_data.model_info)
|
||||
return MappingProxyType({k: v for k, v in {**stored, **incoming}.items() if k not in cleared})
|
||||
|
||||
|
||||
def _validate_ptu_model_info(model_info: Mapping[str, object]) -> None:
|
||||
"""Enforce the PTU cross-field invariant on the effective model_info.
|
||||
|
||||
ptu_count and cost_per_ptu_per_hour must be set together, and a team_id and a
|
||||
ptu_effective_from are required when they are. The start is mandatory rather than
|
||||
defaulted because flat cost accrues from it: inferring one would let a deployment
|
||||
configured today be billed for days it did not exist. Per-field bounds (positive
|
||||
count, non-negative rate) are enforced by ModelInfo itself.
|
||||
|
||||
Window ordering is checked before the count/rate gate. A patch that touches only one
|
||||
end of the window carries no count or rate, and ModelInfo sees one field at a time, so
|
||||
leaving it to either would let an inverted window reach the row; the next load then
|
||||
fails to parse it and drops the deployment out of the router, where no further patch
|
||||
can repair it because each one re-parses the stored value first.
|
||||
"""
|
||||
effective_from: Final = _coerce_ptu_datetime(model_info.get("ptu_effective_from"))
|
||||
effective_to: Final = _coerce_ptu_datetime(model_info.get("ptu_effective_to"))
|
||||
if effective_from is not None and effective_to is not None and effective_to <= effective_from:
|
||||
raise HTTPException(status_code=400, detail="ptu_effective_to must be after ptu_effective_from")
|
||||
|
||||
has_count: Final = model_info.get("ptu_count") is not None
|
||||
has_rate: Final = model_info.get("cost_per_ptu_per_hour") is not None
|
||||
if not has_count and not has_rate:
|
||||
return
|
||||
if has_count != has_rate:
|
||||
raise HTTPException(status_code=400, detail="ptu_count and cost_per_ptu_per_hour must be set together")
|
||||
if effective_from is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"ptu_effective_from is required when PTU fields are set. Flat cost accrues from that "
|
||||
"instant, so without it the start would have to be inferred and a deployment configured "
|
||||
"today could be billed for days it did not exist"
|
||||
),
|
||||
)
|
||||
if not model_info.get("team_id"):
|
||||
raise HTTPException(
|
||||
status_code=400, detail="team_id is required when PTU fields are set (one model maps to one team)"
|
||||
)
|
||||
|
||||
|
||||
def _parse_ptu_datetime(value: object) -> datetime.datetime | None:
|
||||
"""``value`` as a datetime, parsing an ISO string, else None."""
|
||||
if isinstance(value, datetime.datetime):
|
||||
return value
|
||||
if not isinstance(value, str):
|
||||
return None
|
||||
try:
|
||||
return datetime.datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def _coerce_ptu_datetime(value: object) -> datetime.datetime | None:
|
||||
"""Coerce a model_info effective-window value (datetime or ISO string) to UTC, else None."""
|
||||
parsed: Final = _parse_ptu_datetime(value)
|
||||
if parsed is None:
|
||||
return None
|
||||
if parsed.tzinfo is None:
|
||||
return parsed.replace(tzinfo=datetime.timezone.utc)
|
||||
return parsed.astimezone(datetime.timezone.utc)
|
||||
|
||||
|
||||
def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel:
|
||||
merged_model_name: Final = updated_patch.model_name or db_model.model_name
|
||||
merged_litellm_params: Final = db_model.litellm_params.model_dump(exclude_none=True)
|
||||
|
|
@ -270,6 +362,10 @@ def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> Pr
|
|||
if field in SPECIAL_MODEL_INFO_PARAMS and getattr(updated_patch.model_info, field) is None:
|
||||
merged_model_info.pop(field, None)
|
||||
merged_litellm_params.pop(field, None)
|
||||
for field in _explicitly_cleared_ptu_fields(updated_patch.model_info):
|
||||
merged_model_info.pop(field, None)
|
||||
|
||||
_validate_ptu_model_info(merged_model_info)
|
||||
|
||||
# convert to prisma compatible format
|
||||
|
||||
|
|
@ -716,6 +812,17 @@ async def _update_team_model_in_db(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
# Validated before any write, beside the premium check the create path already runs
|
||||
# here. The team ACL is updated below and autocommits, so a validator that raises
|
||||
# further down would leave the team mutated and the deployment row never written.
|
||||
#
|
||||
# The merged view is what gets stored, so that is what has to satisfy the invariants.
|
||||
# Validating the patch alone rejected a partial edit of an already valid deployment:
|
||||
# raising the rate on a configured model carries no ptu_effective_from, which the
|
||||
# stored row supplies.
|
||||
if patch_data.model_info is not None:
|
||||
_validate_ptu_model_info(_merged_ptu_model_info(db_model=db_model, patch_data=patch_data))
|
||||
|
||||
patch_team_id: Final = patch_data.model_info.team_id if patch_data.model_info else None
|
||||
|
||||
# No team_id in patch, proceed with standard update
|
||||
|
|
@ -1424,6 +1531,8 @@ async def add_new_model(
|
|||
|
||||
model_response: LiteLLM_ProxyModelTable | None = None
|
||||
# update DB
|
||||
_validate_ptu_model_info(model_params.model_info.model_dump(exclude_none=True))
|
||||
|
||||
if store_model_in_db is True:
|
||||
"""
|
||||
- store model_list in db
|
||||
|
|
|
|||
|
|
@ -18,13 +18,16 @@ Scoping:
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import (
|
||||
CommonProxyErrors,
|
||||
LiteLLM_TeamTable,
|
||||
LitellmUserRoles,
|
||||
UserAPIKeyAuth,
|
||||
user_api_key_has_admin_view,
|
||||
|
|
@ -40,10 +43,56 @@ from litellm.types.memory_management import (
|
|||
MemoryUpdateRequest,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
def _serialize_metadata_for_prisma(metadata: Any) -> str:
|
||||
class _MemoryRecord(Protocol):
|
||||
memory_id: str
|
||||
key: str
|
||||
value: str
|
||||
metadata: object
|
||||
user_id: str | None
|
||||
team_id: str | None
|
||||
created_at: datetime | None
|
||||
created_by: str | None
|
||||
updated_at: datetime | None
|
||||
updated_by: str | None
|
||||
|
||||
|
||||
class _MemoryTableActions(Protocol):
|
||||
async def create(self, data: Mapping[str, object]) -> _MemoryRecord: ...
|
||||
|
||||
async def find_many(
|
||||
self,
|
||||
where: Mapping[str, object] | None = ...,
|
||||
order: Mapping[str, str] | None = ...,
|
||||
skip: int = ...,
|
||||
take: int = ...,
|
||||
) -> Sequence[_MemoryRecord]: ...
|
||||
|
||||
async def count(self, where: Mapping[str, object] | None = ...) -> int: ...
|
||||
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _MemoryRecord: ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> _MemoryRecord | None: ...
|
||||
|
||||
|
||||
def _memory_table(prisma_client: "PrismaClient") -> _MemoryTableActions:
|
||||
return MemoryRepository(prisma_client).table
|
||||
|
||||
|
||||
class _TeamTableActions(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, str]) -> LiteLLM_TeamTable | None: ...
|
||||
|
||||
|
||||
def _team_table(prisma_client: "PrismaClient") -> _TeamTableActions:
|
||||
return TeamRepository(prisma_client).table
|
||||
|
||||
|
||||
def _serialize_metadata_for_prisma(metadata: object) -> str:
|
||||
"""
|
||||
Encode a `metadata` payload for the `Json?` column.
|
||||
|
||||
|
|
@ -62,25 +111,25 @@ def _is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
|||
return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
|
||||
|
||||
def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> dict | None:
|
||||
def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, object] | None:
|
||||
"""
|
||||
Prisma `where` fragment restricting rows to those the caller can see.
|
||||
Returns None for admins (no restriction).
|
||||
"""
|
||||
if user_api_key_has_admin_view(user_api_key_dict):
|
||||
return None
|
||||
ors: Final[list[dict]] = []
|
||||
if user_api_key_dict.user_id:
|
||||
ors.append({"user_id": user_api_key_dict.user_id})
|
||||
if user_api_key_dict.team_id:
|
||||
ors.append({"team_id": user_api_key_dict.team_id})
|
||||
ors: Final = [
|
||||
{field: value}
|
||||
for field, value in (("user_id", user_api_key_dict.user_id), ("team_id", user_api_key_dict.team_id))
|
||||
if value
|
||||
]
|
||||
if not ors:
|
||||
# Caller has neither user_id nor team_id — match nothing.
|
||||
return {"memory_id": "__no_match__"}
|
||||
return {"OR": ors}
|
||||
|
||||
|
||||
def _row_to_model(row: Any) -> LiteLLM_MemoryRow:
|
||||
def _row_to_model(row: _MemoryRecord) -> LiteLLM_MemoryRow:
|
||||
return LiteLLM_MemoryRow(
|
||||
memory_id=row.memory_id,
|
||||
key=row.key,
|
||||
|
|
@ -95,7 +144,7 @@ def _row_to_model(row: Any) -> LiteLLM_MemoryRow:
|
|||
)
|
||||
|
||||
|
||||
def _require_prisma():
|
||||
def _require_prisma() -> "PrismaClient":
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
|
|
@ -113,7 +162,9 @@ def _internal_error(log_message: str, exc: Exception, default_detail: str) -> HT
|
|||
return HTTPException(status_code=500, detail=default_detail)
|
||||
|
||||
|
||||
async def _assert_write_access(prisma_client: Any, row: Any, user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
async def _assert_write_access(
|
||||
prisma_client: "PrismaClient", row: _MemoryRecord, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> None:
|
||||
"""
|
||||
Enforce ownership for mutations (PUT/DELETE).
|
||||
|
||||
|
|
@ -153,7 +204,7 @@ async def _assert_write_access(prisma_client: Any, row: Any, user_api_key_dict:
|
|||
)
|
||||
|
||||
|
||||
async def _is_team_admin_for(prisma_client: Any, user_api_key_dict: UserAPIKeyAuth, team_id: str) -> bool:
|
||||
async def _is_team_admin_for(prisma_client: "PrismaClient", user_api_key_dict: UserAPIKeyAuth, team_id: str) -> bool:
|
||||
"""
|
||||
True if the caller is a team admin of `team_id`, or an org admin for the
|
||||
team's organization. Mirrors the auth pattern used by team-management
|
||||
|
|
@ -168,7 +219,7 @@ async def _is_team_admin_for(prisma_client: Any, user_api_key_dict: UserAPIKeyAu
|
|||
)
|
||||
|
||||
try:
|
||||
team_obj: Final = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
|
||||
team_obj: Final = await _team_table(prisma_client).find_unique(where={"team_id": team_id})
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error loading team for write-auth check (team_id=%s): %s", team_id, e)
|
||||
return False
|
||||
|
|
@ -269,7 +320,7 @@ async def create_memory(
|
|||
# `metadata` is a `Json?` column — prisma-client-python rejects raw
|
||||
# Python values, so JSON-encode any non-null payload and omit the field
|
||||
# entirely when None so the column defaults to SQL NULL.
|
||||
create_data: Final[dict] = {
|
||||
create_data: Final[dict[str, object]] = {
|
||||
"key": body.key,
|
||||
"value": body.value,
|
||||
"user_id": user_id,
|
||||
|
|
@ -281,7 +332,7 @@ async def create_memory(
|
|||
create_data["metadata"] = _serialize_metadata_for_prisma(body.metadata)
|
||||
|
||||
try:
|
||||
row: Final = await MemoryRepository(prisma_client).table.create(data=create_data)
|
||||
row: Final = await _memory_table(prisma_client).create(data=create_data)
|
||||
except Exception as e:
|
||||
# Key is globally unique. Any duplicate → 409.
|
||||
if _is_unique_violation(e):
|
||||
|
|
@ -325,14 +376,14 @@ async def list_memory(
|
|||
# top-level "AND" — safer than `dict.update` since future visibility
|
||||
# filters could grow an "OR" key that would clobber this one if merged
|
||||
# by key.
|
||||
key_filter: Final[dict] = {}
|
||||
key_filter: Final[dict[str, object]] = {}
|
||||
if key_prefix is not None:
|
||||
key_filter["key"] = {"startsWith": key_prefix}
|
||||
elif key is not None:
|
||||
key_filter["key"] = key
|
||||
|
||||
vis: Final = _visibility_filter(user_api_key_dict)
|
||||
where: dict
|
||||
where: Mapping[str, object]
|
||||
if vis is None:
|
||||
where = key_filter
|
||||
elif not key_filter:
|
||||
|
|
@ -341,8 +392,8 @@ async def list_memory(
|
|||
where = {"AND": [key_filter, vis]}
|
||||
|
||||
try:
|
||||
total: Final = await MemoryRepository(prisma_client).table.count(where=where)
|
||||
rows: Final = await MemoryRepository(prisma_client).table.find_many(
|
||||
total: Final = await _memory_table(prisma_client).count(where=where)
|
||||
rows: Final = await _memory_table(prisma_client).find_many(
|
||||
where=where,
|
||||
order={"updated_at": "desc"},
|
||||
skip=(page - 1) * page_size,
|
||||
|
|
@ -354,12 +405,14 @@ async def list_memory(
|
|||
return MemoryListResponse(memories=[_row_to_model(r) for r in rows], total=total)
|
||||
|
||||
|
||||
async def _find_memory_for_caller(prisma_client: Any, key: str, user_api_key_dict: UserAPIKeyAuth) -> Any:
|
||||
async def _find_memory_for_caller(
|
||||
prisma_client: "PrismaClient", key: str, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> _MemoryRecord:
|
||||
"""Look up a memory row by key, scoped to the caller's visibility."""
|
||||
key_filter: Final[dict] = {"key": key}
|
||||
key_filter: Final[Mapping[str, object]] = {"key": key}
|
||||
vis: Final = _visibility_filter(user_api_key_dict)
|
||||
where: Final[dict] = key_filter if vis is None else {"AND": [key_filter, vis]}
|
||||
rows = await MemoryRepository(prisma_client).table.find_many(where=where, take=1, order={"updated_at": "desc"})
|
||||
where: Final[Mapping[str, object]] = key_filter if vis is None else {"AND": [key_filter, vis]}
|
||||
rows = await _memory_table(prisma_client).find_many(where=where, take=1, order={"updated_at": "desc"})
|
||||
if not rows:
|
||||
raise HTTPException(status_code=404, detail=f"Memory with key '{key}' not found")
|
||||
return rows[0]
|
||||
|
|
@ -415,7 +468,7 @@ async def upsert_memory(
|
|||
fields_sent: Final = body.model_fields_set
|
||||
metadata_in_payload: Final = "metadata" in fields_sent
|
||||
|
||||
data: Final[dict] = {}
|
||||
data: Final[dict[str, object]] = {}
|
||||
if body.value is not None:
|
||||
data["value"] = body.value
|
||||
if metadata_in_payload:
|
||||
|
|
@ -427,7 +480,7 @@ async def upsert_memory(
|
|||
)
|
||||
data["updated_by"] = user_api_key_dict.user_id
|
||||
|
||||
async def _find_existing() -> Any:
|
||||
async def _find_existing() -> _MemoryRecord | None:
|
||||
"""Return the caller-visible row for `key`, or None."""
|
||||
try:
|
||||
return await _find_memory_for_caller(prisma_client, key, user_api_key_dict)
|
||||
|
|
@ -444,7 +497,7 @@ async def upsert_memory(
|
|||
# their team) — otherwise a teammate could overwrite a personal
|
||||
# entry through the OR-based visibility filter.
|
||||
await _assert_write_access(prisma_client, existing, user_api_key_dict)
|
||||
row = await MemoryRepository(prisma_client).table.update(
|
||||
row = await _memory_table(prisma_client).update(
|
||||
where={"memory_id": existing.memory_id},
|
||||
data=data,
|
||||
)
|
||||
|
|
@ -459,7 +512,7 @@ async def upsert_memory(
|
|||
# Omit `metadata` when None so the column defaults to SQL NULL;
|
||||
# otherwise JSON-encode for Prisma — same pattern as
|
||||
# `create_memory` above.
|
||||
create_data: Final[dict] = {
|
||||
create_data: Final[dict[str, object]] = {
|
||||
"key": key,
|
||||
"value": body.value,
|
||||
"user_id": user_id,
|
||||
|
|
@ -470,7 +523,7 @@ async def upsert_memory(
|
|||
if body.metadata is not None:
|
||||
create_data["metadata"] = _serialize_metadata_for_prisma(body.metadata)
|
||||
try:
|
||||
row = await MemoryRepository(prisma_client).table.create(data=create_data)
|
||||
row = await _memory_table(prisma_client).create(data=create_data)
|
||||
except Exception as e:
|
||||
# Race: a concurrent PUT/POST created the row after our check.
|
||||
# Re-read and fall back to an update so the PUT stays idempotent
|
||||
|
|
@ -487,7 +540,7 @@ async def upsert_memory(
|
|||
)
|
||||
# Same write-authorization check as the non-race path.
|
||||
await _assert_write_access(prisma_client, existing_after_race, user_api_key_dict)
|
||||
row = await MemoryRepository(prisma_client).table.update(
|
||||
row = await _memory_table(prisma_client).update(
|
||||
where={"memory_id": existing_after_race.memory_id},
|
||||
data=data,
|
||||
)
|
||||
|
|
@ -515,7 +568,7 @@ async def delete_memory(
|
|||
# Visibility != write authority — see the upsert handler for the rationale.
|
||||
await _assert_write_access(prisma_client, row, user_api_key_dict)
|
||||
try:
|
||||
await MemoryRepository(prisma_client).table.delete(where={"memory_id": row.memory_id})
|
||||
await _memory_table(prisma_client).delete(where={"memory_id": row.memory_id})
|
||||
except Exception as e:
|
||||
raise _internal_error("Error deleting memory: %s", e, "Internal error deleting memory entry.")
|
||||
|
||||
|
|
|
|||
|
|
@ -612,7 +612,7 @@ from litellm.secret_managers.main import (
|
|||
normalize_nonempty_secret_str,
|
||||
str_to_bool,
|
||||
)
|
||||
from litellm.types.integrations.slack_alerting import SlackAlertingArgs
|
||||
from litellm.types.integrations.slack_alerting import AlertType, SlackAlertingArgs
|
||||
from litellm.types.llms.anthropic import (
|
||||
AnthropicMessagesRequest,
|
||||
AnthropicResponse,
|
||||
|
|
@ -8470,6 +8470,42 @@ class ProxyStartupEvent:
|
|||
|
||||
await cls._initialize_spend_tracking_background_jobs(scheduler=scheduler)
|
||||
|
||||
### PTU DAILY ROLLUP ###
|
||||
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import (
|
||||
PTU_ROLLUP_JOB_ID,
|
||||
run_scheduled_ptu_rollup,
|
||||
)
|
||||
|
||||
async def _alert_ptu_rollup_failure(message: str) -> None:
|
||||
await proxy_logging_obj.alerting_handler(
|
||||
message=message,
|
||||
level="High",
|
||||
alert_type=AlertType.failed_tracking_spend,
|
||||
)
|
||||
|
||||
async def _scheduled_ptu_rollup() -> None:
|
||||
# Reuse the PodLockManager from db_spend_update_writer so only one pod
|
||||
# reconciles a day; a multi-pod race could prune another pod's fresh rows
|
||||
await run_scheduled_ptu_rollup(
|
||||
prisma_client,
|
||||
pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager,
|
||||
alert=_alert_ptu_rollup_failure,
|
||||
)
|
||||
|
||||
scheduler.add_job(
|
||||
_scheduled_ptu_rollup,
|
||||
"cron",
|
||||
hour=0,
|
||||
minute=15,
|
||||
timezone="UTC",
|
||||
id=PTU_ROLLUP_JOB_ID,
|
||||
replace_existing=True,
|
||||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"PTU rollup job scheduled at 00:15 UTC daily (only models with PTU config accrue flat cost)"
|
||||
)
|
||||
|
||||
### SPEND LOG CLEANUP ###
|
||||
if (
|
||||
general_settings.get("maximum_spend_logs_retention_period") is not None
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ model LiteLLM_BudgetTable {
|
|||
end_users LiteLLM_EndUserTable[] // multiple end-users can have the same budget
|
||||
tags LiteLLM_TagTable[] // multiple tags can have the same budget
|
||||
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
|
||||
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
|
||||
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
|
||||
}
|
||||
|
||||
// Models on proxy
|
||||
|
|
@ -893,6 +893,7 @@ model LiteLLM_DailyTeamSpend {
|
|||
api_requests BigInt @default(0)
|
||||
successful_requests BigInt @default(0)
|
||||
failed_requests BigInt @default(0)
|
||||
ptu_flat_cost Float @default(0.0)
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
||||
|
|
|
|||
654
litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py
Normal file
654
litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py
Normal file
|
|
@ -0,0 +1,654 @@
|
|||
"""
|
||||
Daily rollup for per-model PTU (provisioned throughput) flat cost.
|
||||
|
||||
v1 reads PTU config straight off the model deployment
|
||||
(``LiteLLM_ProxyModelTable.model_info``): a deployment carrying ``ptu_count``
|
||||
and ``cost_per_ptu_per_hour`` accrues flat cost of
|
||||
``ptu_count * cost_per_ptu_per_hour * active_hours`` for a given UTC day, where
|
||||
``active_hours`` is the overlap between the day and the optional
|
||||
``[ptu_effective_from, ptu_effective_to)`` window (a window opening at 23:00
|
||||
charges one hour that day). The amount is written to ``LiteLLM_DailyTeamSpend``
|
||||
under a sentinel api_key so the rows are distinguishable from per-request rows
|
||||
and share the existing unique constraint.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime, time, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
PTU_PRUNE_SKEW_GRACE_SECONDS,
|
||||
PTU_ROLLUP_JOB_ID,
|
||||
PTU_ROLLUP_LOCK_TTL_SECONDS,
|
||||
PTU_ROLLUP_MAX_BACKFILL_DAYS,
|
||||
PTU_SENTINEL_API_KEY,
|
||||
)
|
||||
from litellm.types.router import ModelInfo
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
_HOURS_PER_DAY: Final = 24
|
||||
_UPSERT_ATTEMPTS: Final = 3
|
||||
_UPSERT_RETRY_BACKOFF_SECONDS: Final = 0.5
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RollupResult:
|
||||
day: date
|
||||
models_processed: int
|
||||
rows_written: int
|
||||
rows_failed: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BackfillResult:
|
||||
start: date
|
||||
end: date
|
||||
days_scanned: int
|
||||
rows_written: int
|
||||
rows_failed: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PTUModel:
|
||||
"""A model deployment carrying valid manual PTU config."""
|
||||
|
||||
model_id: str
|
||||
model_name: str
|
||||
team_id: str
|
||||
ptu_count: int
|
||||
cost_per_ptu_per_hour: float
|
||||
effective_from: datetime | None = None
|
||||
effective_to: datetime | None = None
|
||||
|
||||
|
||||
def _parse_utc_datetime(value: object) -> datetime | None:
|
||||
"""Parse a model_info datetime (ISO string or datetime) into a UTC-aware datetime, else None."""
|
||||
parsed: Final = _coerce_datetime(value)
|
||||
if parsed is None:
|
||||
return None
|
||||
if parsed.tzinfo is None:
|
||||
return parsed.replace(tzinfo=timezone.utc)
|
||||
return parsed.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _coerce_datetime(value: object) -> datetime | None:
|
||||
"""``value`` as a datetime, parsing an ISO string, else None."""
|
||||
if isinstance(value, datetime):
|
||||
return value
|
||||
if not isinstance(value, str):
|
||||
return None
|
||||
try:
|
||||
return datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def _public_model_name(row: object, model_info: Mapping[str, object]) -> str:
|
||||
"""The name an operator recognises for this deployment.
|
||||
|
||||
Creating a team-scoped deployment rewrites model_name to a synthetic routing key
|
||||
(``model_name_<team_id>_<uuid4>``) and keeps the chosen name in
|
||||
``model_info.team_public_model_name``. PTU config is only accepted alongside a
|
||||
team_id, so every PTU deployment carries that synthetic name; keying the sentinel
|
||||
row on it would file each charge under a UUID that no usage view can resolve and
|
||||
that never lines up with the same model's request rows.
|
||||
"""
|
||||
public_name: Final = model_info.get("team_public_model_name")
|
||||
if isinstance(public_name, str) and public_name:
|
||||
return public_name
|
||||
return str(getattr(row, "model_name", "") or "")
|
||||
|
||||
|
||||
def _decode_model_info(raw: object) -> "Mapping[str, object] | None":
|
||||
"""A deployment's model_info as a dict, decoding a JSON string, else None."""
|
||||
if isinstance(raw, str):
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
if isinstance(raw, dict):
|
||||
return raw
|
||||
return None
|
||||
|
||||
|
||||
def _parse_ptu_model(row: object) -> PTUModel | None:
|
||||
"""Return a PTUModel when the deployment carries valid manual PTU config, else None.
|
||||
|
||||
Valid means model_info has a positive ptu_count, a non-negative
|
||||
cost_per_ptu_per_hour, and a team_id (1 model -> 1 team).
|
||||
"""
|
||||
raw_model_info: Final = getattr(row, "model_info", None)
|
||||
model_info: Final = _decode_model_info(raw_model_info)
|
||||
if model_info is None:
|
||||
return None
|
||||
ptu_count: Final = model_info.get("ptu_count")
|
||||
cost_per_hour: Final = model_info.get("cost_per_ptu_per_hour")
|
||||
team_id: Final = model_info.get("team_id")
|
||||
if ptu_count is None or cost_per_hour is None or not team_id:
|
||||
return None
|
||||
try:
|
||||
ptu_count_int: Final = int(ptu_count)
|
||||
cost_per_hour_float: Final = float(cost_per_hour)
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
return None
|
||||
if not 0 < ptu_count_int <= ModelInfo.MAX_PTU_COUNT:
|
||||
return None
|
||||
if not 0 <= cost_per_hour_float <= ModelInfo.MAX_COST_PER_PTU_PER_HOUR:
|
||||
return None
|
||||
if model_info.get("ptu_effective_from") is None:
|
||||
# The endpoints require a start; a row without one predates that rule or was
|
||||
# written around them, and inferring one would bill days the deployment did not exist
|
||||
return None
|
||||
raw_from: Final = model_info.get("ptu_effective_from")
|
||||
raw_to: Final = model_info.get("ptu_effective_to")
|
||||
effective_from: Final = _parse_utc_datetime(raw_from)
|
||||
effective_to: Final = _parse_utc_datetime(raw_to)
|
||||
# A present-but-unparseable bound would read as "no bound" and silently widen the
|
||||
# window to the whole day, so the deployment is skipped until the config is fixed
|
||||
if (raw_from is not None and effective_from is None) or (raw_to is not None and effective_to is None):
|
||||
return None
|
||||
if effective_from is not None and effective_to is not None and effective_to <= effective_from:
|
||||
return None
|
||||
return PTUModel(
|
||||
model_id=str(getattr(row, "model_id", "") or ""),
|
||||
model_name=_public_model_name(row, model_info),
|
||||
team_id=str(team_id),
|
||||
ptu_count=ptu_count_int,
|
||||
cost_per_ptu_per_hour=cost_per_hour_float,
|
||||
effective_from=effective_from,
|
||||
effective_to=effective_to,
|
||||
)
|
||||
|
||||
|
||||
def _active_hours_on_day(model: PTUModel, day: date) -> float:
|
||||
"""Hours the model's PTU window overlaps ``day`` (UTC), clamped to [0, 24]."""
|
||||
day_start: Final = datetime.combine(day, time.min, tzinfo=timezone.utc)
|
||||
day_end: Final = day_start + timedelta(days=1)
|
||||
start: Final = max(day_start, model.effective_from) if model.effective_from else day_start
|
||||
end: Final = min(day_end, model.effective_to) if model.effective_to else day_end
|
||||
if end <= start:
|
||||
return 0.0
|
||||
return (end - start).total_seconds() / 3600.0
|
||||
|
||||
|
||||
def _compute_daily_flat_cost(model: PTUModel, day: date) -> float:
|
||||
"""Flat cost for ``day``: ptu_count * cost_per_ptu_per_hour * active_hours."""
|
||||
return float(model.ptu_count) * model.cost_per_ptu_per_hour * _active_hours_on_day(model, day)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PTUCharge:
|
||||
"""One sentinel row's worth of flat cost for a deployment on a day.
|
||||
|
||||
``model_id`` is the row's identity and goes in the unique key; ``model_name`` is what
|
||||
an operator reads and rides alongside it. A deployment can be renamed, so keying on
|
||||
the name would let two runs holding different config views write the same day twice.
|
||||
"""
|
||||
|
||||
team_id: str
|
||||
model_id: str
|
||||
model_name: str
|
||||
flat_cost: float
|
||||
|
||||
|
||||
def _aggregate_charges(ptu_models: tuple[PTUModel, ...], day: date) -> tuple[_PTUCharge, ...]:
|
||||
"""One charge per deployment that accrues cost on ``day``. Zero-cost deployments are
|
||||
dropped, which keeps a day outside a window from writing a row.
|
||||
|
||||
Deployments sharing a public name inside a team no longer need collapsing: each keys
|
||||
its own row on its own id, and the read path merges them back under the shared name.
|
||||
"""
|
||||
return tuple(
|
||||
_PTUCharge(
|
||||
team_id=model.team_id,
|
||||
model_id=model.model_id,
|
||||
model_name=model.model_name,
|
||||
flat_cost=_compute_daily_flat_cost(model, day),
|
||||
)
|
||||
for model in sorted(ptu_models, key=lambda m: (m.team_id, m.model_id))
|
||||
if _compute_daily_flat_cost(model, day) > 0
|
||||
)
|
||||
|
||||
|
||||
async def _upsert_ptu_daily_row(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
team_id: str,
|
||||
model_id: str,
|
||||
model_name: str,
|
||||
date_str: str,
|
||||
flat_cost: float,
|
||||
) -> None:
|
||||
"""Idempotent upsert of a sentinel-api_key row on LiteLLM_DailyTeamSpend.
|
||||
|
||||
``model`` holds the deployment id because it is part of the table's unique key and a
|
||||
rename must not move the row. ``model_group`` carries the operator-facing name, which
|
||||
is outside the key and is what the usage views display.
|
||||
"""
|
||||
where: Final = { # mutable-ok: prisma upsert filter payload
|
||||
"team_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint": { # mutable-ok: prisma composite-key filter
|
||||
"team_id": team_id,
|
||||
"date": date_str,
|
||||
"api_key": PTU_SENTINEL_API_KEY,
|
||||
"model": model_id,
|
||||
"custom_llm_provider": "",
|
||||
"mcp_namespaced_tool_name": "",
|
||||
"endpoint": "",
|
||||
}
|
||||
}
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
await prisma_client.db.litellm_dailyteamspend.upsert(
|
||||
where=where,
|
||||
data={ # mutable-ok: prisma upsert data payload
|
||||
"create": { # mutable-ok: prisma create payload
|
||||
"team_id": team_id,
|
||||
"date": date_str,
|
||||
"api_key": PTU_SENTINEL_API_KEY,
|
||||
"model": model_id,
|
||||
"model_group": model_name,
|
||||
"custom_llm_provider": "",
|
||||
"mcp_namespaced_tool_name": "",
|
||||
"endpoint": "",
|
||||
"ptu_flat_cost": flat_cost,
|
||||
},
|
||||
"update": { # mutable-ok: prisma update payload
|
||||
"model_group": model_name,
|
||||
"ptu_flat_cost": flat_cost,
|
||||
"updated_at": now,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _upsert_charge_with_retry(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
charge: _PTUCharge,
|
||||
date_str: str,
|
||||
) -> bool:
|
||||
"""Write one charge, retrying transient failures. Returns False once attempts are spent.
|
||||
|
||||
The upsert is idempotent on the sentinel unique key, so a retry can only rewrite the
|
||||
same amount for the same day. Retrying in-run matters because the scheduled job moves
|
||||
on to the next date: a write lost here is a day of PTU cost that no later run replays.
|
||||
"""
|
||||
for attempt in range(1, _UPSERT_ATTEMPTS + 1):
|
||||
try:
|
||||
await _upsert_ptu_daily_row(
|
||||
prisma_client,
|
||||
team_id=charge.team_id,
|
||||
model_id=charge.model_id,
|
||||
model_name=charge.model_name,
|
||||
date_str=date_str,
|
||||
flat_cost=charge.flat_cost,
|
||||
)
|
||||
return True
|
||||
except Exception as exc: # noqa: BLE001 # one bad row must not stop the batch
|
||||
if attempt < _UPSERT_ATTEMPTS:
|
||||
verbose_proxy_logger.warning(
|
||||
"PTU rollup: upsert attempt %d/%d failed for team=%s model=%s day=%s: %s",
|
||||
attempt,
|
||||
_UPSERT_ATTEMPTS,
|
||||
charge.team_id,
|
||||
charge.model_name,
|
||||
date_str,
|
||||
exc,
|
||||
)
|
||||
await asyncio.sleep(_UPSERT_RETRY_BACKOFF_SECONDS * attempt)
|
||||
continue
|
||||
verbose_proxy_logger.error(
|
||||
"PTU rollup: upsert failed after %d attempts for team=%s model=%s day=%s "
|
||||
"(rerun the rollup for that date to recover): %s",
|
||||
_UPSERT_ATTEMPTS,
|
||||
charge.team_id,
|
||||
charge.model_name,
|
||||
date_str,
|
||||
exc,
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
async def _load_ptu_models(prisma_client: "PrismaClient") -> tuple[PTUModel, ...]:
|
||||
"""Every model deployment currently carrying valid manual PTU config."""
|
||||
rows: Final = await prisma_client.db.litellm_proxymodeltable.find_many()
|
||||
return tuple(parsed for parsed in (_parse_ptu_model(row) for row in rows) if parsed is not None)
|
||||
|
||||
|
||||
async def run_ptu_flat_cost_rollup(
|
||||
prisma_client: "PrismaClient",
|
||||
target_date: date | None = None,
|
||||
may_prune: bool = True,
|
||||
) -> RollupResult:
|
||||
"""Rollup one UTC day of flat PTU cost across all PTU-configured model deployments.
|
||||
|
||||
Defaults to yesterday UTC. Authoritative for the day: it upserts the current charges
|
||||
first, then deletes the day's sentinel rows this run did not refresh, so a
|
||||
since-removed, invalidated, or now-out-of-window deployment leaves no stale charge.
|
||||
|
||||
The prune predicate is ``updated_at < run_started`` rather than "not in the charge
|
||||
set I computed", which matters under concurrency: whether a row is garbage becomes a
|
||||
property of the row instead of one run's in-memory config snapshot, so a run can
|
||||
never delete a row a concurrent run just wrote. It is still skipped when any charge
|
||||
failed to write, since a row whose replacement never landed would look unrefreshed.
|
||||
"""
|
||||
day: Final = target_date or (datetime.now(timezone.utc).date() - timedelta(days=1))
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_proxy_logger.warning("PTU rollup: prisma_client is None, skipping")
|
||||
return RollupResult(day=day, models_processed=0, rows_written=0)
|
||||
|
||||
date_str: Final = day.isoformat()
|
||||
run_started: Final = datetime.now(timezone.utc)
|
||||
|
||||
ptu_models: Final = await _load_ptu_models(prisma_client)
|
||||
charges: Final = _aggregate_charges(ptu_models, day)
|
||||
|
||||
landed: Final = tuple(
|
||||
[await _upsert_charge_with_retry(prisma_client, charge=charge, date_str=date_str) for charge in charges]
|
||||
)
|
||||
rows_written: Final = sum(landed)
|
||||
rows_failed: Final = len(charges) - rows_written
|
||||
|
||||
if not may_prune:
|
||||
verbose_proxy_logger.info(
|
||||
"PTU rollup for %s: ran without the cross-pod lock, skipping the prune so a "
|
||||
"concurrent pod's charges cannot be swept by this run's cutoff",
|
||||
date_str,
|
||||
)
|
||||
elif rows_failed:
|
||||
# A charge that never landed leaves its row looking unrefreshed, so the prune
|
||||
# would delete the very row the failed write was meant to replace
|
||||
verbose_proxy_logger.warning(
|
||||
"PTU rollup: %d charge(s) failed for %s, skipping the prune so a row whose "
|
||||
"replacement did not land is not deleted; rerun that date to reconcile",
|
||||
rows_failed,
|
||||
date_str,
|
||||
)
|
||||
else:
|
||||
await _prune_unrefreshed_sentinel_rows(prisma_client, date_str=date_str, run_started=run_started)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"PTU rollup for %s: %d PTU models processed, %d rows written, %d rows failed",
|
||||
date_str,
|
||||
len(ptu_models),
|
||||
rows_written,
|
||||
rows_failed,
|
||||
)
|
||||
return RollupResult(
|
||||
day=day,
|
||||
models_processed=len(ptu_models),
|
||||
rows_written=rows_written,
|
||||
rows_failed=rows_failed,
|
||||
)
|
||||
|
||||
|
||||
def _backfill_window(ptu_models: tuple[PTUModel, ...], end: date) -> tuple[date, ...]:
|
||||
"""The UTC days the catch-up pass considers, oldest first, through ``end`` inclusive.
|
||||
|
||||
Starts at the earliest declared ``ptu_effective_from``, floored at
|
||||
``PTU_ROLLUP_MAX_BACKFILL_DAYS`` before ``end``. A start is required alongside the
|
||||
count and rate, so a deployment without one is not priced rather than being given the
|
||||
floor, which would bill it for the whole cap window. Empty when there is no PTU
|
||||
config, or when every declared window opens after ``end``.
|
||||
"""
|
||||
floor: Final = end - timedelta(days=PTU_ROLLUP_MAX_BACKFILL_DAYS)
|
||||
starts: Final = tuple(model.effective_from.date() for model in ptu_models if model.effective_from)
|
||||
if not starts:
|
||||
return ()
|
||||
start: Final = max(min(starts), floor)
|
||||
return tuple(start + timedelta(days=offset) for offset in range((end - start).days + 1))
|
||||
|
||||
|
||||
async def _existing_sentinel_keys(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
start: date,
|
||||
end: date,
|
||||
) -> frozenset[tuple[str, str, str]]:
|
||||
"""``(team_id, deployment id, date)`` of every PTU sentinel row within ``[start, end]``.
|
||||
|
||||
The row's ``model`` column holds the deployment id, so this is an exact identity and
|
||||
survives a rename. Nothing here reads the display name.
|
||||
"""
|
||||
date_range: Final = {"gte": start.isoformat(), "lte": end.isoformat()} # mutable-ok: prisma range filter
|
||||
rows: Final = await prisma_client.db.litellm_dailyteamspend.find_many(
|
||||
where={"api_key": PTU_SENTINEL_API_KEY, "date": date_range} # mutable-ok: prisma find filter
|
||||
)
|
||||
return frozenset(
|
||||
(
|
||||
str(getattr(row, "team_id", "") or ""),
|
||||
str(getattr(row, "model", "") or ""),
|
||||
str(getattr(row, "date", "") or ""),
|
||||
)
|
||||
for row in rows
|
||||
)
|
||||
|
||||
|
||||
async def run_ptu_flat_cost_backfill(
|
||||
prisma_client: "PrismaClient",
|
||||
today: date | None = None,
|
||||
) -> BackfillResult:
|
||||
"""Price the elapsed days of every PTU window that carry no sentinel row yet.
|
||||
|
||||
Writes only the charges that are missing and never rewrites or deletes an existing
|
||||
row, so a day already priced keeps the amount it was billed, whatever the config says
|
||||
now. A day counts as priced when a sentinel row exists for that deployment id, so
|
||||
renaming a deployment neither re-prices its history nor files a second charge beside
|
||||
the row already there. Zero-cost days write nothing, which leaves a day
|
||||
outside a window reconsidered on each run rather than recorded as done.
|
||||
|
||||
It deletes nothing. Removing a deployment stops it accruing new charges and leaves the
|
||||
days it was billed for standing, since those days were incurred.
|
||||
"""
|
||||
end: Final = (today or datetime.now(timezone.utc).date()) - timedelta(days=1)
|
||||
|
||||
if not prisma_client:
|
||||
verbose_proxy_logger.warning("PTU backfill: prisma_client is None, skipping")
|
||||
return BackfillResult(start=end, end=end, days_scanned=0, rows_written=0)
|
||||
|
||||
ptu_models: Final = await _load_ptu_models(prisma_client)
|
||||
days: Final = _backfill_window(ptu_models, end)
|
||||
|
||||
if not days:
|
||||
return BackfillResult(start=end, end=end, days_scanned=0, rows_written=0)
|
||||
|
||||
priced: Final = await _existing_sentinel_keys(prisma_client, start=days[0], end=days[-1])
|
||||
missing: Final = tuple(
|
||||
(day.isoformat(), charge)
|
||||
for day in days
|
||||
for charge in _aggregate_charges(ptu_models, day)
|
||||
if (charge.team_id, charge.model_id, day.isoformat()) not in priced
|
||||
)
|
||||
if not missing:
|
||||
return BackfillResult(start=days[0], end=days[-1], days_scanned=len(days), rows_written=0)
|
||||
|
||||
landed: Final = tuple(
|
||||
[
|
||||
await _upsert_charge_with_retry(prisma_client, charge=charge, date_str=date_str)
|
||||
for date_str, charge in missing
|
||||
]
|
||||
)
|
||||
rows_written: Final = sum(landed)
|
||||
verbose_proxy_logger.info(
|
||||
"PTU backfill for %s to %s: %d unpriced charge(s) found, %d written, %d failed",
|
||||
days[0].isoformat(),
|
||||
days[-1].isoformat(),
|
||||
len(missing),
|
||||
rows_written,
|
||||
len(missing) - rows_written,
|
||||
)
|
||||
return BackfillResult(
|
||||
start=days[0],
|
||||
end=days[-1],
|
||||
days_scanned=len(days),
|
||||
rows_written=rows_written,
|
||||
rows_failed=len(missing) - rows_written,
|
||||
)
|
||||
|
||||
|
||||
async def run_scheduled_ptu_rollup(
|
||||
prisma_client: "PrismaClient",
|
||||
pod_lock_manager: "PodLockManager | None" = None,
|
||||
target_date: date | None = None,
|
||||
alert: Callable[[str], Awaitable[None]] | None = None,
|
||||
) -> RollupResult | None:
|
||||
"""Run the daily rollup under a cross-pod lock so only one proxy reconciles a day.
|
||||
|
||||
Every proxy process schedules this cron, and the read-charge-prune sequence is not
|
||||
atomic: two pods reading different config snapshots can have the loser's prune delete
|
||||
a row the winner just wrote. Returns None when another pod holds the lock, since that
|
||||
pod is doing the work. A deployment without a Redis-backed lock manager runs
|
||||
unguarded, as ``SpendLogCleanup`` does, and so does a run that cannot reach Redis at
|
||||
all: the lock exists to avoid duplicate work, so no lock problem may cost a day.
|
||||
|
||||
The lease is a fixed TTL with no renewal, so a long scan can outlive it. That costs
|
||||
duplicate work rather than correctness: the upserts are idempotent on the sentinel
|
||||
key and the prune reads only the row's own timestamp, so a second pod arriving
|
||||
mid-run cannot corrupt the day.
|
||||
"""
|
||||
if pod_lock_manager is None or pod_lock_manager.redis_cache is None:
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False)
|
||||
|
||||
if not await pod_lock_manager.acquire_lock(cronjob_id=PTU_ROLLUP_JOB_ID, ttl=PTU_ROLLUP_LOCK_TTL_SECONDS):
|
||||
if await _lock_is_held(pod_lock_manager):
|
||||
verbose_proxy_logger.info("PTU rollup: another pod holds the rollup lock, skipping this run")
|
||||
return None
|
||||
# acquire_lock reports contention and a Redis outage the same way, so an
|
||||
# unreachable Redis would otherwise skip the day on every pod at once. The
|
||||
# reconcile is safe to run concurrently, so losing the lock costs duplicate
|
||||
# work; losing the day costs a team's charges
|
||||
verbose_proxy_logger.warning(
|
||||
"PTU rollup: could not take the rollup lock and no other pod holds it, "
|
||||
"running unguarded rather than skipping the day"
|
||||
)
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False)
|
||||
|
||||
try:
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=True)
|
||||
finally:
|
||||
await pod_lock_manager.release_lock(cronjob_id=PTU_ROLLUP_JOB_ID)
|
||||
|
||||
|
||||
async def _lock_is_held(pod_lock_manager: "PodLockManager") -> bool:
|
||||
"""True only when the rollup lock is readable and someone is holding it.
|
||||
|
||||
A Redis that cannot be read is reported as "not held" so the caller runs the day
|
||||
rather than skipping it; the cost of being wrong here is a duplicate reconcile.
|
||||
"""
|
||||
try:
|
||||
lock_key: Final = pod_lock_manager.get_redis_lock_key(PTU_ROLLUP_JOB_ID)
|
||||
return bool(await pod_lock_manager.redis_cache.async_get_cache(lock_key))
|
||||
except Exception as exc: # noqa: BLE001 # an unreadable lock must not skip the day
|
||||
verbose_proxy_logger.warning("PTU rollup: could not read the rollup lock: %s", exc)
|
||||
return False
|
||||
|
||||
|
||||
async def _run_and_alert(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
target_date: date | None,
|
||||
alert: "Callable[[str], Awaitable[None]] | None",
|
||||
may_prune: bool = True,
|
||||
) -> RollupResult:
|
||||
"""Reconcile the day, catch up any days left unpriced, and alert on charges that did not land.
|
||||
|
||||
A charge that exhausts its retries leaves that team showing no PTU cost for the date,
|
||||
and the scheduled job moves on to the next day rather than replaying it. That is a
|
||||
silent underbill unless someone is reading proxy logs, so it is escalated to whatever
|
||||
alerting the deployment has configured.
|
||||
|
||||
The catch-up pass runs only on the scheduled shape, where ``target_date`` is None. An
|
||||
explicit date means reconcile exactly that day, so it stays a single-day operation.
|
||||
Its failure is contained: the day's own result is returned either way.
|
||||
"""
|
||||
result: Final = await run_ptu_flat_cost_rollup(prisma_client, target_date=target_date, may_prune=may_prune)
|
||||
if result.rows_failed:
|
||||
await _deliver_alert(
|
||||
alert,
|
||||
f"PTU flat-cost rollup for {result.day.isoformat()}: {result.rows_failed} of "
|
||||
f"{result.rows_written + result.rows_failed} team charges failed to write. Those teams show no PTU "
|
||||
f"cost for that date until the rollup is rerun for it.",
|
||||
)
|
||||
if target_date is None:
|
||||
await _backfill_and_alert(prisma_client, alert=alert)
|
||||
return result
|
||||
|
||||
|
||||
async def _backfill_and_alert(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
alert: "Callable[[str], Awaitable[None]] | None",
|
||||
) -> None:
|
||||
"""Catch up unpriced PTU days, alerting on charges that did not land.
|
||||
|
||||
Never raises: the day's own rollup has already run and its result must reach the
|
||||
caller whatever the catch-up pass does.
|
||||
"""
|
||||
try:
|
||||
backfill: Final = await run_ptu_flat_cost_backfill(prisma_client)
|
||||
except Exception as exc: # noqa: BLE001 # the catch-up pass must not fail the day's rollup
|
||||
verbose_proxy_logger.error("PTU backfill: catch-up pass failed, the day's rollup still stands: %s", exc)
|
||||
return
|
||||
if backfill.rows_failed:
|
||||
await _deliver_alert(
|
||||
alert,
|
||||
f"PTU flat-cost backfill for {backfill.start.isoformat()} to {backfill.end.isoformat()}: "
|
||||
f"{backfill.rows_failed} of {backfill.rows_written + backfill.rows_failed} previously unpriced charges "
|
||||
f"failed to write. Those days stay unpriced until a later run picks them up.",
|
||||
)
|
||||
|
||||
|
||||
async def _deliver_alert(alert: "Callable[[str], Awaitable[None]] | None", message: str) -> None:
|
||||
"""Send an operator alert when one is configured, swallowing a broken channel."""
|
||||
if alert is None:
|
||||
return
|
||||
try:
|
||||
await alert(message)
|
||||
except Exception as exc: # noqa: BLE001 # a broken alert channel must not fail the rollup
|
||||
verbose_proxy_logger.error("PTU rollup: could not deliver the failed-charge alert: %s", exc)
|
||||
|
||||
|
||||
async def _prune_unrefreshed_sentinel_rows(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
date_str: str,
|
||||
run_started: datetime,
|
||||
) -> None:
|
||||
"""Delete the day's PTU sentinel rows this run did not refresh.
|
||||
|
||||
Every charge the run wrote bumps ``updated_at`` past ``run_started``, so anything
|
||||
left below that mark is a (team, model) the current config no longer prices. The mark
|
||||
is pulled back by ``PTU_PRUNE_SKEW_GRACE_SECONDS`` because the two timestamps come
|
||||
from different hosts: a stale row is hours old, a concurrently written one is seconds
|
||||
old, and the grace separates them without waiting on clocks agreeing. The
|
||||
predicate reads only the row, never the caller's config snapshot, which is what
|
||||
makes it safe to run twice, out of order, or beside another pod: a row written
|
||||
after this run began is out of reach of its delete. Mirrors the retention predicate
|
||||
``SpendLogCleanup`` deletes by."""
|
||||
cutoff: Final = run_started - timedelta(seconds=PTU_PRUNE_SKEW_GRACE_SECONDS)
|
||||
await prisma_client.db.litellm_dailyteamspend.delete_many(
|
||||
where={ # mutable-ok: prisma delete filter
|
||||
"date": date_str,
|
||||
"api_key": PTU_SENTINEL_API_KEY,
|
||||
"updated_at": {"lt": cutoff}, # mutable-ok: prisma comparison filter
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
__all__ = (
|
||||
"PTU_ROLLUP_JOB_ID",
|
||||
"PTU_SENTINEL_API_KEY",
|
||||
"BackfillResult",
|
||||
"PTUModel",
|
||||
"RollupResult",
|
||||
"run_ptu_flat_cost_backfill",
|
||||
"run_ptu_flat_cost_rollup",
|
||||
"run_scheduled_ptu_rollup",
|
||||
)
|
||||
|
|
@ -17,7 +17,8 @@ from __future__ import annotations
|
|||
|
||||
import hashlib
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -35,10 +36,32 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import httpx
|
||||
|
||||
from litellm import Router
|
||||
from litellm.types.rag import RAGIngestOptions
|
||||
|
||||
|
||||
class S3VectorDataPayload(TypedDict):
|
||||
float32: Sequence[float]
|
||||
|
||||
|
||||
class S3VectorEntry(TypedDict):
|
||||
key: str
|
||||
data: S3VectorDataPayload
|
||||
metadata: Mapping[str, str]
|
||||
|
||||
|
||||
class S3VectorsQueryMatch(TypedDict, total=False):
|
||||
key: str
|
||||
distance: float
|
||||
metadata: Mapping[str, str]
|
||||
|
||||
|
||||
class S3VectorsQueryResponse(TypedDict, total=False):
|
||||
vectors: Sequence[S3VectorsQueryMatch]
|
||||
|
||||
|
||||
class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
||||
"""
|
||||
S3 Vectors RAG ingestion using httpx + AWS SigV4 signing.
|
||||
|
|
@ -66,10 +89,10 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
BaseAWSLLM.__init__(self)
|
||||
|
||||
# Extract config
|
||||
self.vector_bucket_name = self.vector_store_config["vector_bucket_name"]
|
||||
self.index_name = self.vector_store_config.get("index_name")
|
||||
self.distance_metric = self.vector_store_config.get("distance_metric", S3_VECTORS_DEFAULT_DISTANCE_METRIC)
|
||||
self.non_filterable_metadata_keys = self.vector_store_config.get(
|
||||
self.vector_bucket_name: str = self.vector_store_config["vector_bucket_name"]
|
||||
self.index_name: str | None = self.vector_store_config.get("index_name")
|
||||
self.distance_metric: str = self.vector_store_config.get("distance_metric", S3_VECTORS_DEFAULT_DISTANCE_METRIC)
|
||||
self.non_filterable_metadata_keys: Sequence[str] = self.vector_store_config.get(
|
||||
"non_filterable_metadata_keys",
|
||||
S3_VECTORS_DEFAULT_NON_FILTERABLE_METADATA_KEYS,
|
||||
)
|
||||
|
|
@ -78,7 +101,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
self.dimension = self._get_dimension_from_config()
|
||||
|
||||
# Get AWS region using BaseAWSLLM method
|
||||
_aws_region: Final = self.vector_store_config.get("aws_region_name")
|
||||
_aws_region: Final[str | None] = self.vector_store_config.get("aws_region_name")
|
||||
self.aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
|
||||
aws_region_name=str(_aws_region) if _aws_region else None
|
||||
)
|
||||
|
|
@ -135,7 +158,8 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
Returns None if dimension should be auto-detected.
|
||||
"""
|
||||
if "dimension" in self.vector_store_config:
|
||||
return int(self.vector_store_config["dimension"])
|
||||
configured_dimension: Final[int] = self.vector_store_config["dimension"]
|
||||
return int(configured_dimension)
|
||||
return None
|
||||
|
||||
async def _ensure_config_initialized(self):
|
||||
|
|
@ -258,7 +282,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
get_body: Final = safe_dumps({"vectorBucketName": self.vector_bucket_name})
|
||||
|
||||
try:
|
||||
response = await self._sign_and_execute_request("POST", get_url, data=get_body)
|
||||
response: httpx.Response = await self._sign_and_execute_request("POST", get_url, data=get_body)
|
||||
if response.status_code == 200:
|
||||
verbose_logger.debug("Vector bucket %s exists", self.vector_bucket_name)
|
||||
return
|
||||
|
|
@ -294,7 +318,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
get_body: Final = safe_dumps({"vectorBucketName": self.vector_bucket_name, "indexName": self.index_name})
|
||||
|
||||
try:
|
||||
response = await self._sign_and_execute_request("POST", get_url, data=get_body)
|
||||
response: httpx.Response = await self._sign_and_execute_request("POST", get_url, data=get_body)
|
||||
if response.status_code == 200:
|
||||
verbose_logger.debug("Vector index %s exists", self.index_name)
|
||||
return
|
||||
|
|
@ -311,7 +335,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
)
|
||||
|
||||
# Prepare index configuration per AWS API docs
|
||||
index_config: Final = {
|
||||
index_config: Final[dict[str, object]] = {
|
||||
"vectorBucketName": self.vector_bucket_name,
|
||||
"indexName": self.index_name,
|
||||
"dataType": "float32",
|
||||
|
|
@ -336,7 +360,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
verbose_logger.exception("Error creating vector index: %s", e)
|
||||
raise
|
||||
|
||||
async def _put_vectors(self, vectors: list[dict[str, Any]]):
|
||||
async def _put_vectors(self, vectors: Sequence[S3VectorEntry]):
|
||||
"""
|
||||
Call PutVectors API to store vectors in S3 Vectors.
|
||||
|
||||
|
|
@ -355,7 +379,9 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
}
|
||||
|
||||
try:
|
||||
response: Final = await self._sign_and_execute_request("POST", url, data=safe_dumps(request_body))
|
||||
response: Final[httpx.Response] = await self._sign_and_execute_request(
|
||||
"POST", url, data=safe_dumps(request_body)
|
||||
)
|
||||
|
||||
if response.status_code in (200, 201):
|
||||
verbose_logger.info("Successfully stored %s vectors in index %s", len(vectors), self.index_name)
|
||||
|
|
@ -442,24 +468,18 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
raise ValueError(error_msg)
|
||||
|
||||
# Prepare vectors for PutVectors API
|
||||
vectors: Final = []
|
||||
for i, (chunk, embedding) in enumerate(zip(chunks, embeddings)):
|
||||
# Build metadata dict
|
||||
metadata: dict[str, str] = {
|
||||
"source_text": chunk, # Non-filterable (for reference)
|
||||
"chunk_index": str(i), # Filterable
|
||||
}
|
||||
|
||||
if filename:
|
||||
metadata["filename"] = filename # Filterable
|
||||
|
||||
vector_obj = {
|
||||
"key": f"{filename}_{i}" if filename else f"chunk_{i}",
|
||||
"data": {"float32": embedding},
|
||||
"metadata": metadata,
|
||||
}
|
||||
|
||||
vectors.append(vector_obj)
|
||||
vectors: Final = [
|
||||
S3VectorEntry(
|
||||
key=f"{filename}_{i}" if filename else f"chunk_{i}",
|
||||
data=S3VectorDataPayload(float32=embedding),
|
||||
metadata=(
|
||||
{"source_text": chunk, "chunk_index": str(i), "filename": filename}
|
||||
if filename
|
||||
else {"source_text": chunk, "chunk_index": str(i)}
|
||||
),
|
||||
)
|
||||
for i, (chunk, embedding) in enumerate(zip(chunks, embeddings))
|
||||
]
|
||||
|
||||
# Call PutVectors API
|
||||
await self._put_vectors(vectors)
|
||||
|
|
@ -468,7 +488,9 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
vector_store_id: Final = f"{self.vector_bucket_name}:{self.index_name}"
|
||||
return vector_store_id, filename
|
||||
|
||||
async def query_vector_store(self, vector_store_id: str, query: str, top_k: int = 5) -> dict[str, Any] | None:
|
||||
async def query_vector_store(
|
||||
self, vector_store_id: str, query: str, top_k: int = 5
|
||||
) -> S3VectorsQueryResponse | None:
|
||||
"""
|
||||
Query S3 Vectors using QueryVectors API.
|
||||
|
||||
|
|
@ -489,7 +511,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
embedding_model: Final = self.embedding_config.get("model", "text-embedding-3-small")
|
||||
|
||||
response = await litellm.aembedding(model=embedding_model, input=[query])
|
||||
query_embedding: Final = response.data[0]["embedding"]
|
||||
query_embedding: Final[Sequence[float]] = response.data[0]["embedding"]
|
||||
|
||||
# Call QueryVectors API
|
||||
url: Final = f"https://s3vectors.{self.aws_region_name}.api.aws/QueryVectors"
|
||||
|
|
@ -504,15 +526,18 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
}
|
||||
|
||||
try:
|
||||
response = await self._sign_and_execute_request("POST", url, data=safe_dumps(request_body))
|
||||
query_response: Final[httpx.Response] = await self._sign_and_execute_request(
|
||||
"POST", url, data=safe_dumps(request_body)
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
results: Final = response.json()
|
||||
if query_response.status_code == 200:
|
||||
results: Final[S3VectorsQueryResponse] = query_response.json()
|
||||
matches: Final = results.get("vectors")
|
||||
verbose_logger.debug("Query returned %s results", len(results.get("vectors", [])))
|
||||
|
||||
# Check if query terms appear in results
|
||||
if results.get("vectors"):
|
||||
for result in results["vectors"]:
|
||||
if matches:
|
||||
for result in matches:
|
||||
metadata = result.get("metadata", {})
|
||||
source_text = metadata.get("source_text", "")
|
||||
if query.lower() in source_text.lower():
|
||||
|
|
@ -521,7 +546,9 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
# Return results even if exact match not found
|
||||
return results
|
||||
else:
|
||||
verbose_logger.error("QueryVectors failed with status %s: %s", response.status_code, response.text)
|
||||
verbose_logger.error(
|
||||
"QueryVectors failed with status %s: %s", query_response.status_code, query_response.text
|
||||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error querying vectors: %s", e)
|
||||
|
|
|
|||
|
|
@ -2,7 +2,10 @@ import re
|
|||
import traceback
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypedDict, overload
|
||||
|
||||
from openai.types.chat import ChatCompletionToolParam
|
||||
from openai.types.responses.function_tool_param import FunctionToolParam
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
|
||||
|
|
@ -18,6 +21,7 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamingResponse,
|
||||
)
|
||||
from litellm.types.llms.openai import ToolParam as ResponsesToolParam
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
Choices,
|
||||
|
|
@ -36,10 +40,14 @@ else:
|
|||
MCPTool = Any
|
||||
|
||||
# NOTE: We intentionally keep ToolParam as a broad type here to avoid tight coupling
|
||||
# to optional OpenAI SDK typing symbols in environments that may not have them available.
|
||||
# `Any` is used to keep mypy compatible with the broader OpenAI tool union types
|
||||
# passed around in Responses API while still allowing dict-style access at runtime.
|
||||
ToolParam = Any
|
||||
ToolParam: TypeAlias = Mapping[str, object]
|
||||
|
||||
|
||||
class MCPToolResult(TypedDict):
|
||||
tool_call_id: str | None
|
||||
result: str
|
||||
name: str | None
|
||||
|
||||
|
||||
LITELLM_PROXY_MCP_SERVER_URL: Final = "litellm_proxy"
|
||||
LITELLM_PROXY_MCP_SERVER_URL_PREFIX: Final = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/"
|
||||
|
|
@ -199,13 +207,12 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
_get_tools_from_mcp_servers,
|
||||
)
|
||||
|
||||
mcp_servers: Final[list[str]] = []
|
||||
if mcp_tools_with_litellm_proxy:
|
||||
for _tool in mcp_tools_with_litellm_proxy:
|
||||
# if user specifies servers as server_url: litellm_proxy/mcp/zapier,github then return zapier,github
|
||||
server_url = _tool.get("server_url", "") if isinstance(_tool, dict) else ""
|
||||
if isinstance(server_url, str) and server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL_PREFIX):
|
||||
mcp_servers.append(server_url.split("/")[-1])
|
||||
mcp_servers: Final = [
|
||||
server_url.split("/")[-1]
|
||||
for _tool in (mcp_tools_with_litellm_proxy or ())
|
||||
for server_url in (_tool.get("server_url", "") if isinstance(_tool, dict) else "",)
|
||||
if isinstance(server_url, str) and server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL_PREFIX)
|
||||
]
|
||||
|
||||
# Resolve toolset names: collect all toolset IDs first, then apply their
|
||||
# combined permissions in a single pass so multiple toolsets are unioned
|
||||
|
|
@ -279,15 +286,15 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
)
|
||||
|
||||
server_names: Final[list[str]] = []
|
||||
for server in allowed_mcp_servers:
|
||||
if server is None:
|
||||
continue
|
||||
server_name = (
|
||||
getattr(server, "server_name", None) or getattr(server, "alias", None) or getattr(server, "name", None)
|
||||
server_names: Final = [
|
||||
server_name
|
||||
for server in allowed_mcp_servers
|
||||
if server is not None
|
||||
for server_name in (
|
||||
getattr(server, "server_name", None) or getattr(server, "alias", None) or getattr(server, "name", None),
|
||||
)
|
||||
if isinstance(server_name, str):
|
||||
server_names.append(server_name)
|
||||
if isinstance(server_name, str)
|
||||
]
|
||||
|
||||
return tools, server_names
|
||||
|
||||
|
|
@ -305,8 +312,8 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
List of deduplicated MCP tools
|
||||
The returned dictionary maps each tool_name to the server_name
|
||||
"""
|
||||
seen_names: Final = set()
|
||||
deduplicated_tools: Final = []
|
||||
seen_names: Final[set[str]] = set()
|
||||
deduplicated_tools: Final[list[MCPTool]] = []
|
||||
tool_server_map: Final[dict[str, str]] = {}
|
||||
|
||||
for tool in mcp_tools:
|
||||
|
|
@ -331,7 +338,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
) -> list[MCPTool]:
|
||||
"""Filter MCP tools based on allowed_tools parameter from the original tool configs."""
|
||||
# Collect all allowed tool names from all MCP tool configs
|
||||
allowed_tool_names: Final = set()
|
||||
allowed_tool_names: Final[set[str]] = set()
|
||||
for tool_config in mcp_tools_with_litellm_proxy:
|
||||
if isinstance(tool_config, dict) and "allowed_tools" in tool_config:
|
||||
allowed_tools = tool_config.get("allowed_tools", [])
|
||||
|
|
@ -343,23 +350,13 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
return mcp_tools
|
||||
|
||||
# Filter tools based on allowed names
|
||||
filtered_tools: Final = []
|
||||
for mcp_tool in mcp_tools:
|
||||
if isinstance(mcp_tool, dict):
|
||||
tool_name = mcp_tool.get("name")
|
||||
else:
|
||||
tool_name = getattr(mcp_tool, "name", None)
|
||||
|
||||
if not tool_name:
|
||||
continue
|
||||
|
||||
if tool_name in allowed_tool_names:
|
||||
filtered_tools.append(mcp_tool)
|
||||
continue
|
||||
|
||||
unprefixed_name, _ = split_server_prefix_from_name(tool_name)
|
||||
if unprefixed_name in allowed_tool_names:
|
||||
filtered_tools.append(mcp_tool)
|
||||
filtered_tools: Final = [
|
||||
mcp_tool
|
||||
for mcp_tool in mcp_tools
|
||||
for tool_name in (mcp_tool.get("name") if isinstance(mcp_tool, dict) else getattr(mcp_tool, "name", None),)
|
||||
if tool_name
|
||||
and (tool_name in allowed_tool_names or split_server_prefix_from_name(tool_name)[0] in allowed_tool_names)
|
||||
]
|
||||
|
||||
return filtered_tools
|
||||
|
||||
|
|
@ -448,24 +445,37 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
return deduplicated_mcp_tools, tool_server_map
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def _transform_mcp_tools_to_openai(
|
||||
mcp_tools: Sequence[MCPTool],
|
||||
target_format: Literal["responses"] = ...,
|
||||
) -> list[FunctionToolParam]: ...
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def _transform_mcp_tools_to_openai(
|
||||
mcp_tools: Sequence[MCPTool],
|
||||
target_format: Literal["chat"],
|
||||
) -> list[ChatCompletionToolParam]: ...
|
||||
|
||||
@staticmethod
|
||||
def _transform_mcp_tools_to_openai(
|
||||
mcp_tools: Sequence[MCPTool],
|
||||
target_format: Literal["responses", "chat"] = "responses",
|
||||
) -> list[Any]:
|
||||
) -> Sequence[FunctionToolParam | ChatCompletionToolParam]:
|
||||
"""Transform MCP tools to OpenAI-compatible format."""
|
||||
from litellm.experimental_mcp_client.tools import (
|
||||
transform_mcp_tool_to_openai_responses_api_tool,
|
||||
transform_mcp_tool_to_openai_tool,
|
||||
)
|
||||
|
||||
openai_tools: Final[list[Any]] = []
|
||||
for mcp_tool in mcp_tools:
|
||||
if target_format == "chat":
|
||||
openai_tool = transform_mcp_tool_to_openai_tool(mcp_tool)
|
||||
else:
|
||||
openai_tool = transform_mcp_tool_to_openai_responses_api_tool(mcp_tool)
|
||||
openai_tools.append(openai_tool)
|
||||
openai_tools: Final = [
|
||||
transform_mcp_tool_to_openai_tool(mcp_tool)
|
||||
if target_format == "chat"
|
||||
else transform_mcp_tool_to_openai_responses_api_tool(mcp_tool)
|
||||
for mcp_tool in mcp_tools
|
||||
]
|
||||
|
||||
return openai_tools
|
||||
|
||||
|
|
@ -496,9 +506,9 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
return True
|
||||
|
||||
@staticmethod
|
||||
def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> list[Any]:
|
||||
def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> list[object]:
|
||||
"""Extract tool calls from the response output."""
|
||||
tool_calls: Final[list[Any]] = []
|
||||
tool_calls: Final[list[object]] = []
|
||||
for output_item in response.output:
|
||||
# Check if this is a function call output item
|
||||
if isinstance(output_item, dict) and output_item.get("type") == "function_call":
|
||||
|
|
@ -533,7 +543,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
def _extract_tool_call_details(
|
||||
tool_call,
|
||||
tool_call: object,
|
||||
) -> tuple[str | None, str | None, str | None]:
|
||||
"""Extract tool name, arguments, and call_id from a tool call."""
|
||||
if isinstance(tool_call, dict):
|
||||
|
|
@ -566,7 +576,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
return tool_name, tool_arguments, tool_call_id
|
||||
|
||||
@staticmethod
|
||||
def _parse_tool_arguments(tool_arguments: Any) -> dict[str, Any]:
|
||||
def _parse_tool_arguments(tool_arguments: str | None) -> dict[str, object]:
|
||||
"""Parse tool arguments, handling both string and dict formats."""
|
||||
import json
|
||||
|
||||
|
|
@ -591,23 +601,18 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
# Fallback to generic handling if MCP types not available
|
||||
return "Tool executed successfully"
|
||||
|
||||
text_parts: Final = []
|
||||
other_content_types: Final = []
|
||||
|
||||
for content_item in result.content:
|
||||
if isinstance(content_item, TextContent):
|
||||
# Text content - extract the text
|
||||
text_parts.append(str(content_item.text))
|
||||
elif isinstance(content_item, ImageContent):
|
||||
# Image content
|
||||
other_content_types.append("Image")
|
||||
elif isinstance(content_item, EmbeddedResource):
|
||||
# Embedded resource
|
||||
other_content_types.append("EmbeddedResource")
|
||||
else:
|
||||
# Other unknown content types
|
||||
content_type = type(content_item).__name__
|
||||
other_content_types.append(content_type)
|
||||
text_parts: Final = [
|
||||
str(content_item.text) for content_item in result.content if isinstance(content_item, TextContent)
|
||||
]
|
||||
other_content_types: Final = [
|
||||
"Image"
|
||||
if isinstance(content_item, ImageContent)
|
||||
else "EmbeddedResource"
|
||||
if isinstance(content_item, EmbeddedResource)
|
||||
else type(content_item).__name__
|
||||
for content_item in result.content
|
||||
if not isinstance(content_item, TextContent)
|
||||
]
|
||||
|
||||
# Combine text parts if any
|
||||
result_text = " ".join(text_parts) if text_parts else ""
|
||||
|
|
@ -631,7 +636,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
litellm_call_id: str | None = None,
|
||||
litellm_trace_id: str | None = None,
|
||||
request_tags: list[str] | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> list[MCPToolResult]:
|
||||
"""Execute tool calls and return results."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -645,11 +650,11 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
)
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
tool_results: Final = []
|
||||
tool_results: Final[list[MCPToolResult]] = []
|
||||
tool_call_id: str | None = None
|
||||
rules_obj: Final = Rules()
|
||||
for tool_call in tool_calls:
|
||||
logging_request_data: dict[str, Any] = {}
|
||||
logging_request_data: dict[str, object] = {}
|
||||
tool_name: str | None = None
|
||||
try:
|
||||
(
|
||||
|
|
@ -678,7 +683,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
sanitized_tool_name = strip_known_server_prefix(resolved_tool_name, mcp_server)
|
||||
|
||||
start_time = datetime.now()
|
||||
logging_input = [
|
||||
logging_input: Sequence[Mapping[str, object]] = [
|
||||
{
|
||||
"role": "tool",
|
||||
"content": {
|
||||
|
|
@ -688,13 +693,14 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
}
|
||||
]
|
||||
tool_logging_call_id = litellm_call_id or str(uuid.uuid4())
|
||||
logging_metadata: dict[str, object] = {
|
||||
"tool_call_id": tool_call_id,
|
||||
"tool_name": sanitized_tool_name,
|
||||
"server_name": server_name,
|
||||
}
|
||||
logging_request_data = {
|
||||
"model": f"MCP: {tool_name}",
|
||||
"metadata": {
|
||||
"tool_call_id": tool_call_id,
|
||||
"tool_name": sanitized_tool_name,
|
||||
"server_name": server_name,
|
||||
},
|
||||
"metadata": logging_metadata,
|
||||
"input": logging_input,
|
||||
"call_type": CallTypes.call_mcp_tool.value,
|
||||
"litellm_call_id": tool_logging_call_id,
|
||||
|
|
@ -712,7 +718,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
if litellm_trace_id:
|
||||
logging_request_data["litellm_trace_id"] = litellm_trace_id
|
||||
if request_tags:
|
||||
logging_request_data["metadata"]["tags"] = request_tags
|
||||
logging_metadata["tags"] = request_tags
|
||||
if user_api_key_auth is not None:
|
||||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
LiteLLMProxyRequestSetup,
|
||||
|
|
@ -902,16 +908,16 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
def _create_follow_up_messages_for_chat(
|
||||
original_messages: list[Any],
|
||||
original_messages: list[object],
|
||||
response: ModelResponse,
|
||||
tool_results: Sequence[Mapping[str, object]],
|
||||
) -> list[Any]:
|
||||
) -> Sequence[Mapping[str, object]]:
|
||||
"""Create follow-up chat messages that include tool execution results."""
|
||||
from copy import deepcopy
|
||||
|
||||
from litellm.utils import convert_list_message_to_dict
|
||||
|
||||
follow_up_messages: list[Any] = convert_list_message_to_dict(deepcopy(original_messages))
|
||||
follow_up_messages: list[dict[str, object]] = convert_list_message_to_dict(deepcopy(original_messages))
|
||||
|
||||
if not follow_up_messages:
|
||||
follow_up_messages = []
|
||||
|
|
@ -950,9 +956,9 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
response: ResponsesAPIResponse,
|
||||
tool_results: Sequence[Mapping[str, object]],
|
||||
original_input: str | ResponseInputParam | None = None,
|
||||
) -> list[Any]:
|
||||
) -> list[object]:
|
||||
"""Create follow-up input with tool results in proper format."""
|
||||
follow_up_input: Final[list[Any]] = []
|
||||
follow_up_input: Final[list[object]] = []
|
||||
|
||||
# Add original user input if available to maintain conversation context
|
||||
if original_input:
|
||||
|
|
@ -964,8 +970,8 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
follow_up_input.append(original_input)
|
||||
|
||||
# Add the assistant message with function calls
|
||||
assistant_message_content: Final[list[Any]] = []
|
||||
function_calls: Final[list[dict[str, Any]]] = []
|
||||
assistant_message_content: Final[list[object]] = []
|
||||
function_calls: Final[list[dict[str, object]]] = []
|
||||
|
||||
for output_item in response.output:
|
||||
if not isinstance(output_item, dict) and hasattr(output_item, "model_dump"):
|
||||
|
|
@ -1027,7 +1033,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
async def _make_follow_up_call(
|
||||
follow_up_input: list[Any],
|
||||
model: str,
|
||||
all_tools: list[Any] | None,
|
||||
all_tools: Sequence[ResponsesToolParam] | None,
|
||||
response_id: str,
|
||||
**call_params: Any,
|
||||
) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator:
|
||||
|
|
@ -1044,7 +1050,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
async def _log_mcp_tool_failure(
|
||||
*,
|
||||
proxy_logging_obj: Optional["ProxyLogging"],
|
||||
user_api_key_auth: Any,
|
||||
user_api_key_auth: "UserAPIKeyAuth | None",
|
||||
request_data: dict[str, object],
|
||||
error: Exception,
|
||||
) -> None:
|
||||
|
|
@ -1072,7 +1078,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
all_tools: Sequence[object] | None,
|
||||
mcp_tools_with_litellm_proxy: list[Mapping[str, object]],
|
||||
mcp_discovery_events: list[ResponsesAPIStreamingResponse],
|
||||
call_params: dict[str, Any],
|
||||
call_params: Mapping[str, object],
|
||||
previous_response_id: str | None,
|
||||
tool_server_map: dict[str, str],
|
||||
**kwargs,
|
||||
|
|
@ -1115,10 +1121,10 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
input: str | ResponseInputParam,
|
||||
model: str,
|
||||
all_tools: Sequence[object] | None,
|
||||
call_params: dict[str, Any],
|
||||
call_params: Mapping[str, object],
|
||||
previous_response_id: str | None,
|
||||
**kwargs,
|
||||
) -> dict[str, Any]:
|
||||
**kwargs: object,
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Build a clean request parameters dictionary for MCP streaming.
|
||||
|
||||
|
|
@ -1126,7 +1132,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
in a clean, maintainable way.
|
||||
"""
|
||||
# Start with the core required parameters
|
||||
request_params: Final = {
|
||||
request_params: Final[dict[str, object]] = {
|
||||
"input": input,
|
||||
"model": model,
|
||||
"tools": all_tools,
|
||||
|
|
@ -1146,7 +1152,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
def _create_tool_execution_events(
|
||||
tool_calls: Sequence[object], tool_results: list[dict[str, Any]]
|
||||
tool_calls: Sequence[object], tool_results: Sequence[MCPToolResult]
|
||||
) -> list[ResponsesAPIStreamingResponse]:
|
||||
"""
|
||||
Create MCP tool execution events for streaming.
|
||||
|
|
|
|||
|
|
@ -19,13 +19,13 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
ResponsesAPIStreamingResponse,
|
||||
ToolParam,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import MCPToolResult
|
||||
else:
|
||||
MCPTool = Any
|
||||
|
||||
|
|
@ -33,7 +33,7 @@ MAX_MCP_TOOL_CALL_ROUNDS: Final = 5
|
|||
|
||||
|
||||
async def create_mcp_list_tools_events(
|
||||
mcp_tools_with_litellm_proxy: list[ToolParam],
|
||||
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]],
|
||||
user_api_key_auth: "UserAPIKeyAuth | None",
|
||||
base_item_id: str,
|
||||
pre_processed_mcp_tools: list[MCPTool],
|
||||
|
|
@ -44,13 +44,14 @@ async def create_mcp_list_tools_events(
|
|||
|
||||
try:
|
||||
# Extract MCP server names
|
||||
mcp_servers: Final = []
|
||||
for tool in mcp_tools_with_litellm_proxy:
|
||||
if isinstance(tool, dict) and "server_url" in tool:
|
||||
server_url = tool.get("server_url")
|
||||
if isinstance(server_url, str) and server_url.startswith("litellm_proxy/mcp/"):
|
||||
server_name = server_url.split("/")[-1]
|
||||
mcp_servers.append(server_name)
|
||||
_mcp_servers: Final = [
|
||||
server_url.split("/")[-1]
|
||||
for tool in mcp_tools_with_litellm_proxy
|
||||
if isinstance(tool, dict)
|
||||
and "server_url" in tool
|
||||
and isinstance(server_url := tool.get("server_url"), str)
|
||||
and server_url.startswith("litellm_proxy/mcp/")
|
||||
]
|
||||
|
||||
# Emit list tools in progress event
|
||||
in_progress_event: Final = MCPListToolsInProgressEvent(
|
||||
|
|
@ -65,15 +66,14 @@ async def create_mcp_list_tools_events(
|
|||
filtered_mcp_tools: Final = pre_processed_mcp_tools
|
||||
|
||||
# Convert tools to dict format for the event
|
||||
mcp_tools_dict: Final = []
|
||||
for tool in filtered_mcp_tools:
|
||||
if hasattr(tool, "model_dump") and callable(getattr(tool, "model_dump")):
|
||||
# Type cast to help mypy understand this is safe after hasattr check
|
||||
mcp_tools_dict.append(cast(Any, tool).model_dump())
|
||||
elif hasattr(tool, "__dict__"):
|
||||
mcp_tools_dict.append(tool.__dict__)
|
||||
else:
|
||||
mcp_tools_dict.append({"name": getattr(tool, "name", str(tool))})
|
||||
_mcp_tools_dict: Final = [
|
||||
tool.model_dump()
|
||||
if hasattr(tool, "model_dump") and callable(getattr(tool, "model_dump"))
|
||||
else tool.__dict__
|
||||
if hasattr(tool, "__dict__")
|
||||
else {"name": getattr(tool, "name", str(tool))}
|
||||
for tool in filtered_mcp_tools
|
||||
]
|
||||
|
||||
# Emit list tools completed event
|
||||
completed_event: Final = MCPListToolsCompletedEvent(
|
||||
|
|
@ -96,21 +96,18 @@ async def create_mcp_list_tools_events(
|
|||
server_label = str(server_label_value) if server_label_value is not None else ""
|
||||
|
||||
# Format tools for OpenAI output_item.done format
|
||||
formatted_tools: Final = []
|
||||
for tool in filtered_mcp_tools:
|
||||
tool_dict = {
|
||||
formatted_tools: Final = [
|
||||
{
|
||||
"name": getattr(tool, "name", "unknown"),
|
||||
"description": getattr(tool, "description", ""),
|
||||
"annotations": {"read_only": False},
|
||||
**dict.fromkeys(
|
||||
("input_schema",) if hasattr(tool, "inputSchema") or hasattr(tool, "input_schema") else (),
|
||||
getattr(tool, "inputSchema", getattr(tool, "input_schema", None)),
|
||||
),
|
||||
}
|
||||
|
||||
# Add input_schema if available
|
||||
if hasattr(tool, "inputSchema"):
|
||||
tool_dict["input_schema"] = getattr(tool, "inputSchema")
|
||||
elif hasattr(tool, "input_schema"):
|
||||
tool_dict["input_schema"] = getattr(tool, "input_schema")
|
||||
|
||||
formatted_tools.append(tool_dict)
|
||||
for tool in filtered_mcp_tools
|
||||
]
|
||||
|
||||
# Create the output_item.done event with MCP tools list
|
||||
output_item_done_event = OutputItemDoneEvent(
|
||||
|
|
@ -166,7 +163,7 @@ async def create_mcp_list_tools_events(
|
|||
|
||||
def create_mcp_call_events(
|
||||
tool_name: str,
|
||||
tool_call_id: str,
|
||||
tool_call_id: str | None,
|
||||
arguments: str,
|
||||
result: str | None = None,
|
||||
base_item_id: str | None = None,
|
||||
|
|
@ -256,9 +253,12 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
4. Emits tool execution events in the stream
|
||||
"""
|
||||
|
||||
model: str
|
||||
tool_results: "Sequence[MCPToolResult]"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_iterator: Any, # Can be None - will be created internally
|
||||
base_iterator: "BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None", # created internally when None
|
||||
mcp_events: list[ResponsesAPIStreamingResponse],
|
||||
tool_server_map: dict[str, str],
|
||||
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] | None = None,
|
||||
|
|
@ -285,7 +285,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
self.tool_server_map = tool_server_map
|
||||
|
||||
# Iterator references
|
||||
self.base_iterator: Any | ResponsesAPIResponse | None = base_iterator # Will be created when needed
|
||||
self.base_iterator: BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None = (
|
||||
base_iterator # Will be created when needed
|
||||
)
|
||||
|
||||
# Response collection for tool execution
|
||||
self.collected_response: ResponsesAPIResponse | None = None
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ import weakref
|
|||
from collections import defaultdict
|
||||
from collections.abc import AsyncGenerator, Callable, Generator, Mapping, Sequence
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeVar, Union, cast
|
||||
|
||||
import anyio
|
||||
|
|
@ -117,6 +118,7 @@ from litellm.router_utils.cooldown_handlers import (
|
|||
DEFAULT_COOLDOWN_TIME_SECONDS,
|
||||
_async_get_cooldown_deployments,
|
||||
_async_get_cooldown_deployments_with_debug_info,
|
||||
_first_present, # pyright: ignore[reportPrivateUsage] - shared internal helper across router_utils submodules, matching the other cooldown_handlers imports on this line
|
||||
_get_cooldown_deployments,
|
||||
_set_cooldown_deployments,
|
||||
is_advisor_orchestration_failure,
|
||||
|
|
@ -1898,7 +1900,7 @@ class Router:
|
|||
# Set per-deployment num_retries on exception for retry logic
|
||||
if deployment is not None:
|
||||
self._set_deployment_num_retries_on_exception(e, deployment)
|
||||
self._set_failed_deployment_id_on_exception(e, deployment)
|
||||
self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs)
|
||||
raise e
|
||||
|
||||
def _get_silent_experiment_kwargs(self, **kwargs) -> dict:
|
||||
|
|
@ -2961,7 +2963,7 @@ class Router:
|
|||
# Set per-deployment num_retries on exception for retry logic
|
||||
if deployment is not None:
|
||||
self._set_deployment_num_retries_on_exception(e, deployment)
|
||||
self._set_failed_deployment_id_on_exception(e, deployment)
|
||||
self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs)
|
||||
raise e
|
||||
except Exception as e:
|
||||
verbose_router_logger.info("litellm.acompletion(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
|
|
@ -2970,7 +2972,7 @@ class Router:
|
|||
# Set per-deployment num_retries on exception for retry logic
|
||||
if deployment is not None:
|
||||
self._set_deployment_num_retries_on_exception(e, deployment)
|
||||
self._set_failed_deployment_id_on_exception(e, deployment)
|
||||
self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs)
|
||||
raise e
|
||||
|
||||
def _update_kwargs_before_fallbacks(
|
||||
|
|
@ -3016,7 +3018,7 @@ class Router:
|
|||
except (ValueError, TypeError):
|
||||
pass # Skip if value can't be converted to int
|
||||
|
||||
def _set_failed_deployment_id_on_exception(self, exception: Exception, deployment: dict) -> None:
|
||||
def _set_failed_deployment_id_on_exception(self, exception: Exception, deployment: Mapping[str, Any]) -> None:
|
||||
"""
|
||||
Stamp the failed deployment's `model_info.id` on the exception so the
|
||||
fallback layer can exclude it from subsequent re-picks within the same
|
||||
|
|
@ -3035,6 +3037,16 @@ class Router:
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
def _stamp_failed_deployment_id_with_effective_model_info(
|
||||
self, exception: Exception, deployment: Mapping[str, Any], kwargs: Mapping[str, Any]
|
||||
) -> None:
|
||||
# A client-side-credential call gets a dynamic deployment id generated inside
|
||||
# _update_kwargs_with_deployment and stamped into kwargs["model_info"]; stamping
|
||||
# the static shared deployment's id instead would let one tenant's bad credentials
|
||||
# cool down the deployment every other tenant sharing this config relies on.
|
||||
effective_model_info: Final = kwargs.get("model_info") or deployment.get("model_info") or MappingProxyType({})
|
||||
self._set_failed_deployment_id_on_exception(exception, MappingProxyType({"model_info": effective_model_info}))
|
||||
|
||||
def _update_kwargs_with_default_litellm_params(
|
||||
self, kwargs: dict, metadata_variable_name: str | None = "metadata"
|
||||
) -> None:
|
||||
|
|
@ -4521,10 +4533,11 @@ class Router:
|
|||
|
||||
passthrough_on_no_deployment: Final = kwargs.pop("passthrough_on_no_deployment", False)
|
||||
function_name: Final = "_ageneric_api_call_with_fallbacks"
|
||||
deployment = None # rebind-ok: pre-init so the except block can stamp a failure with no deployment picked
|
||||
try:
|
||||
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
try:
|
||||
deployment: Final = await self.async_get_available_deployment(
|
||||
deployment = await self.async_get_available_deployment( # rebind-ok: set on success, see pre-init above
|
||||
model=model,
|
||||
request_kwargs=kwargs,
|
||||
messages=kwargs.get("messages", None),
|
||||
|
|
@ -4601,6 +4614,8 @@ class Router:
|
|||
)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
if deployment is not None:
|
||||
self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs)
|
||||
raise e
|
||||
|
||||
async def _aresponses_with_streaming_fallbacks(
|
||||
|
|
@ -7078,7 +7093,9 @@ class Router:
|
|||
)
|
||||
|
||||
# Determine cooldown time with priority: deployment config > response header > router default
|
||||
deployment_cooldown: Final = litellm_params.get("cooldown_time", None)
|
||||
deployment_cooldown: Final = _first_present(
|
||||
_model_info if isinstance(_model_info, dict) else None, litellm_params, key="cooldown_time"
|
||||
)
|
||||
|
||||
header_cooldown = None
|
||||
if exception_headers is not None:
|
||||
|
|
@ -11800,6 +11817,23 @@ class Router:
|
|||
and allowed_fails_policy.BadRequestErrorAllowedFails is not None
|
||||
):
|
||||
return allowed_fails_policy.BadRequestErrorAllowedFails
|
||||
if (
|
||||
isinstance(exception, litellm.InternalServerError)
|
||||
and allowed_fails_policy.InternalServerErrorAllowedFails is not None
|
||||
):
|
||||
return allowed_fails_policy.InternalServerErrorAllowedFails
|
||||
if (
|
||||
isinstance(exception, litellm.ServiceUnavailableError)
|
||||
and allowed_fails_policy.ServiceUnavailableErrorAllowedFails is not None
|
||||
):
|
||||
return allowed_fails_policy.ServiceUnavailableErrorAllowedFails
|
||||
if (
|
||||
isinstance(exception, litellm.BadGatewayError)
|
||||
and allowed_fails_policy.BadGatewayErrorAllowedFails is not None
|
||||
):
|
||||
return allowed_fails_policy.BadGatewayErrorAllowedFails
|
||||
if isinstance(exception, litellm.NotFoundError) and allowed_fails_policy.NotFoundErrorAllowedFails is not None:
|
||||
return allowed_fails_policy.NotFoundErrorAllowedFails
|
||||
|
||||
def _initialize_alerting(self):
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ Wrapper around router cache. Meant to handle model cooldown logic
|
|||
|
||||
import functools
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
|
@ -28,6 +29,12 @@ class CooldownCacheValue(TypedDict):
|
|||
cooldown_time: float
|
||||
|
||||
|
||||
# Cap on the corrected in-memory TTL set in `_corrected_active_cooldown`: re-checks the
|
||||
# real remaining cooldown against Redis at least this often, so an entry that later gets
|
||||
# deleted or extended in Redis before its original deadline is still noticed promptly.
|
||||
_MAX_CORRECTED_IN_MEMORY_TTL_SECONDS: Final = 60.0
|
||||
|
||||
|
||||
class CooldownCache:
|
||||
def __init__(self, cache: DualCache, default_cooldown_time: float):
|
||||
self.cache = cache
|
||||
|
|
@ -100,6 +107,30 @@ class CooldownCache:
|
|||
def get_cooldown_cache_key(model_id: str) -> str:
|
||||
return "deployment:" + model_id + ":cooldown"
|
||||
|
||||
def _corrected_active_cooldown(
|
||||
self,
|
||||
key: str,
|
||||
result: Mapping[str, Any],
|
||||
current_time: float,
|
||||
) -> CooldownCacheValue | None:
|
||||
"""
|
||||
Return a CooldownCacheValue if the cooldown is still active, or None if it has expired.
|
||||
|
||||
Also corrects the in-memory TTL when DualCache promotes a Redis entry using the
|
||||
default 600s TTL instead of the true remaining cooldown time.
|
||||
"""
|
||||
cooldown_cache_value: Final = CooldownCacheValue(**result) # pyright: ignore[reportUnknownArgumentType] - result comes from an untyped cache read, not from our own code
|
||||
remaining: Final = (cooldown_cache_value["timestamp"] + cooldown_cache_value["cooldown_time"]) - current_time
|
||||
if remaining <= 0:
|
||||
self.cache.in_memory_cache.delete_cache(key)
|
||||
return None
|
||||
current_expiry: Final = self.cache.in_memory_cache.ttl_dict.get(key)
|
||||
if current_expiry is not None and current_expiry > current_time + remaining + 5:
|
||||
corrected_ttl: Final = min(remaining, _MAX_CORRECTED_IN_MEMORY_TTL_SECONDS)
|
||||
self.cache.in_memory_cache.delete_cache(key)
|
||||
self.cache.in_memory_cache.set_cache(key, result, ttl=corrected_ttl)
|
||||
return cooldown_cache_value
|
||||
|
||||
async def async_get_active_cooldowns(
|
||||
self, model_ids: list[str], parent_otel_span: Span | None
|
||||
) -> list[tuple[str, CooldownCacheValue]]:
|
||||
|
|
@ -117,11 +148,13 @@ class CooldownCache:
|
|||
if results is None or all(v is None for v in results):
|
||||
return active_cooldowns
|
||||
|
||||
# Process the results
|
||||
current_time: Final = time.time()
|
||||
for model_id, result in zip(model_ids, results):
|
||||
if result and isinstance(result, dict):
|
||||
cooldown_cache_value = CooldownCacheValue(**result)
|
||||
active_cooldowns.append((model_id, cooldown_cache_value))
|
||||
key = CooldownCache.get_cooldown_cache_key(model_id)
|
||||
cooldown_cache_value = self._corrected_active_cooldown(key, result, current_time)
|
||||
if cooldown_cache_value is not None:
|
||||
active_cooldowns.append((model_id, cooldown_cache_value))
|
||||
|
||||
return active_cooldowns
|
||||
|
||||
|
|
@ -134,11 +167,13 @@ class CooldownCache:
|
|||
results: Final = self.cache.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
|
||||
active_cooldowns: Final = []
|
||||
# Process the results
|
||||
current_time: Final = time.time()
|
||||
for model_id, result in zip(model_ids, results):
|
||||
if result and isinstance(result, dict):
|
||||
cooldown_cache_value = CooldownCacheValue(**result)
|
||||
active_cooldowns.append((model_id, cooldown_cache_value))
|
||||
key = CooldownCache.get_cooldown_cache_key(model_id)
|
||||
cooldown_cache_value = self._corrected_active_cooldown(key, result, current_time)
|
||||
if cooldown_cache_value is not None:
|
||||
active_cooldowns.append((model_id, cooldown_cache_value))
|
||||
|
||||
return active_cooldowns
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ Router cooldown handlers
|
|||
|
||||
import asyncio
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
|
|
@ -58,6 +60,148 @@ def is_advisor_orchestration_failure(exception: BaseException | None) -> bool:
|
|||
return bool(getattr(exception, _ADVISOR_ORCHESTRATION_FAILURE_ATTR, False))
|
||||
|
||||
|
||||
_EXCEPTION_POLICY_FIELDS: Final[tuple[tuple[type, str], ...]] = (
|
||||
# ContentPolicyViolationError subclasses BadRequestError, so it must be checked first.
|
||||
(litellm.ContentPolicyViolationError, "ContentPolicyViolationErrorAllowedFails"),
|
||||
(litellm.BadRequestError, "BadRequestErrorAllowedFails"),
|
||||
(litellm.AuthenticationError, "AuthenticationErrorAllowedFails"),
|
||||
(litellm.Timeout, "TimeoutErrorAllowedFails"),
|
||||
(litellm.RateLimitError, "RateLimitErrorAllowedFails"),
|
||||
(litellm.InternalServerError, "InternalServerErrorAllowedFails"),
|
||||
(litellm.ServiceUnavailableError, "ServiceUnavailableErrorAllowedFails"),
|
||||
(litellm.BadGatewayError, "BadGatewayErrorAllowedFails"),
|
||||
(litellm.NotFoundError, "NotFoundErrorAllowedFails"),
|
||||
)
|
||||
|
||||
|
||||
def _first_present(*sources: Mapping[str, Any] | None, key: str) -> int | float | None:
|
||||
"""Return *key* from the first source mapping where it's set, so callers can
|
||||
support a setting living in more than one deployment config location. Sources
|
||||
are checked in order from most to least specific to that setting."""
|
||||
for source in sources:
|
||||
if source is None:
|
||||
continue
|
||||
value = source.get(key)
|
||||
if value is not None:
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _get_deployment_cooldown_policy(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
deployment: str,
|
||||
) -> tuple[Mapping[str, int] | None, int | None]:
|
||||
"""Return (allowed_fails_policy, allowed_fails) from deployment model_info, or (None, None).
|
||||
|
||||
`model_info` is the only supported location for these two fields (unlike
|
||||
`cooldown_time`, they have no pre-existing `litellm_params` precedent): `litellm_params`
|
||||
gets copied wholesale into the actual provider call kwargs (see e.g.
|
||||
Router._image_generation's `data = deployment["litellm_params"].copy()`), so a new
|
||||
field placed there would leak into the outgoing LLM request instead of staying
|
||||
router-internal.
|
||||
"""
|
||||
dep: Final = litellm_router_instance.get_model_info(id=deployment)
|
||||
if dep is None:
|
||||
return None, None
|
||||
mi: Final[Mapping[str, Any]] = dep.get("model_info") or MappingProxyType({})
|
||||
raw: Final = mi.get("allowed_fails_policy")
|
||||
policy: Final[Mapping[str, int] | None] = raw if isinstance(raw, dict) else None
|
||||
allowed: Final[int | None] = mi.get("allowed_fails")
|
||||
return policy, allowed
|
||||
|
||||
|
||||
def _resolve_allowed_fails_from_policy(
|
||||
policy: Mapping[str, int] | None,
|
||||
exception: Exception,
|
||||
) -> int | None:
|
||||
"""Match *exception* against *policy* and return the configured allowed-fail count, or None."""
|
||||
if policy is None:
|
||||
return None
|
||||
for exc_type, field in _EXCEPTION_POLICY_FIELDS:
|
||||
if isinstance(exception, exc_type):
|
||||
value = policy.get(field)
|
||||
if value is not None:
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _should_cooldown_based_on_deployment_policy(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
deployment: str,
|
||||
original_exception: Exception,
|
||||
dep_policy: Mapping[str, int] | None,
|
||||
dep_allowed_fails: int | None,
|
||||
is_single_deployment_model_group: bool,
|
||||
) -> bool:
|
||||
"""Resolve deployment-level allowed-fails and delegate to the shared counting logic.
|
||||
|
||||
When the deployment's policy doesn't cover *original_exception*'s type and no
|
||||
deployment-wide `allowed_fails` is set either, defer to router-level behavior
|
||||
instead of forcing an immediate cooldown.
|
||||
|
||||
A generic, deployment-wide `allowed_fails` predates this feature's per-exception-type
|
||||
policy and is a much less deliberate opt-in, so on a single-deployment model group it
|
||||
still defers to the "avoid cooldowns on single deployment model groups" safety net
|
||||
(see `_should_cooldown_deployment`'s BASE CASE) rather than silently disabling it. An
|
||||
explicit, named-exception-type `allowed_fails_policy` entry is unambiguous enough to
|
||||
override that safety net, matching `_has_explicit_allowed_fails_policy_for_exception`.
|
||||
"""
|
||||
allowed_fails_from_policy: Final = _resolve_allowed_fails_from_policy(dep_policy, original_exception)
|
||||
if allowed_fails_from_policy is None and dep_allowed_fails is not None and is_single_deployment_model_group:
|
||||
return False
|
||||
|
||||
allowed_fails_override: Final[int | None] = (
|
||||
allowed_fails_from_policy if allowed_fails_from_policy is not None else dep_allowed_fails
|
||||
)
|
||||
cache_key_suffix: Final[str | None] = (
|
||||
type(original_exception).__name__
|
||||
if allowed_fails_from_policy is not None
|
||||
else ("generic" if dep_allowed_fails is not None else None)
|
||||
)
|
||||
|
||||
dep: Final = litellm_router_instance.get_model_info(id=deployment)
|
||||
cooldown_time_override: Final = (
|
||||
_first_present(dep.get("model_info"), dep.get("litellm_params"), key="cooldown_time")
|
||||
if dep is not None
|
||||
else None
|
||||
)
|
||||
|
||||
return should_cooldown_based_on_allowed_fails_policy(
|
||||
litellm_router_instance=litellm_router_instance,
|
||||
deployment=deployment,
|
||||
original_exception=original_exception,
|
||||
allowed_fails_override=allowed_fails_override,
|
||||
cooldown_time_override=cooldown_time_override,
|
||||
cache_key_suffix=cache_key_suffix,
|
||||
)
|
||||
|
||||
|
||||
def _has_explicit_allowed_fails_policy_for_exception(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
deployment: str | None,
|
||||
original_exception: Exception,
|
||||
) -> bool:
|
||||
"""True if this deployment has an explicit, deployment-level allowed_fails_policy
|
||||
entry matching *original_exception*'s type.
|
||||
|
||||
`_is_cooldown_required` skips cooldown evaluation for most 4XX errors (BadRequestError,
|
||||
ContentPolicyViolationError) by default, since a generic client error is usually not the
|
||||
deployment's fault. A deployment-level allowed_fails_policy entry naming that exact
|
||||
exception type is this PR's own per-deployment opt-in, so it overrides that default.
|
||||
|
||||
Deliberately scoped to the deployment level only, and to the named-exception-type
|
||||
policy dict rather than a plain `allowed_fails` integer: a pre-existing router-wide
|
||||
`allowed_fails_policy` (or a deployment's generic `allowed_fails`) predates this
|
||||
feature and must keep its existing behavior for 4XX types `_is_cooldown_required`
|
||||
already excludes, rather than silently start cooling down deployments whose configs
|
||||
never opted into this specific override.
|
||||
"""
|
||||
if deployment is None:
|
||||
return False
|
||||
dep_policy, _ = _get_deployment_cooldown_policy(litellm_router_instance, deployment)
|
||||
return _resolve_allowed_fails_from_policy(dep_policy, original_exception) is not None
|
||||
|
||||
|
||||
def _is_cooldown_required(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
model_id: str,
|
||||
|
|
@ -155,6 +299,10 @@ def _should_run_cooldown_logic(
|
|||
model_id=deployment,
|
||||
exception_status=exception_status,
|
||||
exception_str=str(original_exception),
|
||||
) and not _has_explicit_allowed_fails_policy_for_exception(
|
||||
litellm_router_instance=litellm_router_instance,
|
||||
deployment=deployment,
|
||||
original_exception=original_exception,
|
||||
):
|
||||
verbose_router_logger.debug("Should Not Run Cooldown Logic: _is_cooldown_required returned False")
|
||||
return False
|
||||
|
|
@ -190,11 +338,24 @@ def _should_cooldown_deployment(
|
|||
|
||||
- v1 logic (Legacy): if allowed fails or allowed fail policy set, coolsdown if num fails in this minute > allowed fails
|
||||
"""
|
||||
## BASE CASE - single deployment
|
||||
model_group: Final = litellm_router_instance.get_model_group(id=deployment)
|
||||
is_single_deployment_model_group = False
|
||||
if model_group is not None and len(model_group) == 1:
|
||||
is_single_deployment_model_group = True
|
||||
|
||||
## CHECK DEPLOYMENT-LEVEL POLICY FIRST (overrides router-level)
|
||||
dep_policy, dep_allowed_fails = _get_deployment_cooldown_policy(litellm_router_instance, deployment)
|
||||
if dep_policy is not None or dep_allowed_fails is not None:
|
||||
return _should_cooldown_based_on_deployment_policy(
|
||||
litellm_router_instance,
|
||||
deployment,
|
||||
original_exception,
|
||||
dep_policy,
|
||||
dep_allowed_fails,
|
||||
is_single_deployment_model_group,
|
||||
)
|
||||
|
||||
## BASE CASE - single deployment
|
||||
if (
|
||||
litellm_router_instance.allowed_fails_policy is None
|
||||
and _is_allowed_fails_set_on_router(litellm_router_instance=litellm_router_instance) is False
|
||||
|
|
@ -382,29 +543,50 @@ def should_cooldown_based_on_allowed_fails_policy(
|
|||
litellm_router_instance: LitellmRouter,
|
||||
deployment: str,
|
||||
original_exception: Any,
|
||||
allowed_fails_override: int | None = None,
|
||||
cooldown_time_override: float | None = None,
|
||||
cache_key_suffix: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if fails are within the allowed limit and update the number of fails.
|
||||
|
||||
When *allowed_fails_override* / *cooldown_time_override* are supplied they
|
||||
take precedence over the router-level values (used by deployment-level overrides).
|
||||
|
||||
When *cache_key_suffix* is supplied the fail counter is keyed as
|
||||
``{deployment}:{cache_key_suffix}`` so that different exception types are
|
||||
tracked independently per deployment.
|
||||
|
||||
Returns:
|
||||
- True if fails exceed the allowed limit (should cooldown)
|
||||
- False if fails are within the allowed limit (should not cooldown)
|
||||
"""
|
||||
allowed_fails: Final = (
|
||||
litellm_router_instance.get_allowed_fails_from_policy(
|
||||
exception=original_exception,
|
||||
)
|
||||
or litellm_router_instance.allowed_fails
|
||||
allowed_fails_from_policy: Final = litellm_router_instance.get_allowed_fails_from_policy(
|
||||
exception=original_exception
|
||||
)
|
||||
allowed_fails: Final = (
|
||||
allowed_fails_override
|
||||
if allowed_fails_override is not None
|
||||
else (
|
||||
allowed_fails_from_policy
|
||||
if allowed_fails_from_policy is not None
|
||||
else litellm_router_instance.allowed_fails
|
||||
)
|
||||
)
|
||||
cooldown_time: Final = (
|
||||
cooldown_time_override
|
||||
if cooldown_time_override is not None
|
||||
else (litellm_router_instance.cooldown_time or DEFAULT_COOLDOWN_TIME_SECONDS)
|
||||
)
|
||||
cooldown_time: Final = litellm_router_instance.cooldown_time or DEFAULT_COOLDOWN_TIME_SECONDS
|
||||
|
||||
current_fails: Final = litellm_router_instance.failed_calls.get_cache(key=deployment) or 0
|
||||
cache_key: Final = f"{deployment}:{cache_key_suffix}" if cache_key_suffix else deployment
|
||||
current_fails: Final = litellm_router_instance.failed_calls.get_cache(key=cache_key) or 0
|
||||
updated_fails: Final = current_fails + 1
|
||||
|
||||
if updated_fails > allowed_fails:
|
||||
return True
|
||||
else:
|
||||
litellm_router_instance.failed_calls.set_cache(key=deployment, value=updated_fails, ttl=cooldown_time)
|
||||
litellm_router_instance.failed_calls.set_cache(key=cache_key, value=updated_fails, ttl=cooldown_time)
|
||||
|
||||
return False
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import hashlib
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
|
@ -12,6 +13,16 @@ from litellm.router_utils.add_retry_fallback_headers import (
|
|||
add_fallback_headers_to_response,
|
||||
get_fallback_error_info,
|
||||
)
|
||||
from litellm.router_utils.batch_utils import _get_router_metadata_variable_name
|
||||
from litellm.router_utils.cooldown_handlers import (
|
||||
_first_present, # pyright: ignore[reportPrivateUsage] - shared internal helper, used across router_utils
|
||||
_set_cooldown_deployments, # pyright: ignore[reportPrivateUsage] - shared helper, used across router_utils
|
||||
cast_exception_status_to_int,
|
||||
is_advisor_orchestration_failure,
|
||||
)
|
||||
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
|
||||
increment_deployment_failures_for_current_minute,
|
||||
)
|
||||
from litellm.types.router import LiteLLMParamsTypedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -21,6 +32,116 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
LitellmRouter = Any
|
||||
|
||||
# Status codes a generic API call's caller-supplied resource id can trigger on its own
|
||||
# (e.g. a nonexistent file/batch/thread id), independent of the selected deployment's health.
|
||||
_REQUEST_SCOPED_STATUS_CODES: Final = frozenset((404,))
|
||||
|
||||
|
||||
def _trigger_cooldown_for_failed_deployment(
|
||||
litellm_router: LitellmRouter,
|
||||
kwargs: Mapping[str, Any],
|
||||
exception: Exception,
|
||||
) -> None:
|
||||
"""
|
||||
Trigger cooldown for a failed fallback deployment.
|
||||
|
||||
In the fallback path the normal failure-callback cooldown is skipped because the
|
||||
Logging object sets has_logged_async_failure=True after the first failure and
|
||||
blocks all subsequent failure callbacks. This helper ensures every failed
|
||||
fallback deployment is evaluated for cooldown regardless.
|
||||
"""
|
||||
try:
|
||||
if is_advisor_orchestration_failure(exception):
|
||||
verbose_router_logger.debug(
|
||||
"Not triggering cooldown for fallback deployment: failure originated "
|
||||
"from advisor orchestration, not the selected deployment."
|
||||
)
|
||||
return
|
||||
|
||||
exception_status: Final[str | int] = getattr(exception, "status_code", "")
|
||||
|
||||
# Generic API calls (files, batches, threads, rerank, ...) take a caller-supplied
|
||||
# resource id, so a 404 there usually means "that id doesn't exist" rather than
|
||||
# "this deployment is unhealthy". Left unguarded, one bad id would 404 every
|
||||
# deployment in the fallback chain and cool all of them down from a single request.
|
||||
if (
|
||||
kwargs.get("original_generic_function") is not None
|
||||
and cast_exception_status_to_int(exception_status) in _REQUEST_SCOPED_STATUS_CODES
|
||||
):
|
||||
verbose_router_logger.debug(
|
||||
"Not triggering cooldown for fallback deployment: status %s on a generic API "
|
||||
"call is caller-attributable, not a deployment health signal.",
|
||||
exception_status,
|
||||
)
|
||||
return
|
||||
|
||||
# The proxy's `x-litellm-timeout` header lets a caller set an arbitrarily short
|
||||
# timeout, which litellm.Timeout reports as status 408 regardless of the deployment's
|
||||
# actual health. Left unguarded, a caller could force a 408 on every deployment in
|
||||
# the fallback chain from a single request with a near-zero timeout.
|
||||
if kwargs.get("client_side_timeout") and cast_exception_status_to_int(exception_status) == 408:
|
||||
verbose_router_logger.debug(
|
||||
"Not triggering cooldown for fallback deployment: a caller-supplied "
|
||||
"x-litellm-timeout caused this 408, not deployment health."
|
||||
)
|
||||
return
|
||||
|
||||
# Only Router._set_failed_deployment_id_on_exception()'s server-stamped id is
|
||||
# trusted here: a metadata-bucket lookup (e.g. "metadata"/"litellm_metadata")
|
||||
# can't reliably tell a caller-supplied bucket from a router-authored one
|
||||
# without knowing this call's function_name, so a client with permission to
|
||||
# set metadata could otherwise get an arbitrary deployment cooled down.
|
||||
deployment_id: Final[str | None] = getattr(exception, "failed_deployment_id", None)
|
||||
|
||||
if deployment_id is None:
|
||||
verbose_router_logger.debug("Cannot trigger cooldown for fallback: no failed_deployment_id on exception")
|
||||
return
|
||||
|
||||
# Priority: deployment config > response header > router default, matching
|
||||
# Router.deployment_callback_on_failure's precedence for the primary path.
|
||||
deployment_dict: Final = litellm_router.get_model_info(id=deployment_id)
|
||||
deployment_cooldown: Final = (
|
||||
_first_present(
|
||||
deployment_dict.get("model_info"), deployment_dict.get("litellm_params"), key="cooldown_time"
|
||||
)
|
||||
if deployment_dict is not None
|
||||
else None
|
||||
)
|
||||
exception_headers: Final = litellm.litellm_core_utils.exception_mapping_utils._get_response_headers(
|
||||
original_exception=exception
|
||||
)
|
||||
_get_retry_after: Final = (
|
||||
litellm.utils._get_retry_after_from_exception_header # pyright: ignore[reportPrivateUsage] - as router.py
|
||||
)
|
||||
header_cooldown: Final = (
|
||||
_get_retry_after(response_headers=exception_headers) if exception_headers is not None else None
|
||||
)
|
||||
time_to_cooldown: Final = (
|
||||
deployment_cooldown
|
||||
if deployment_cooldown is not None and deployment_cooldown >= 0
|
||||
else (
|
||||
header_cooldown
|
||||
if header_cooldown is not None and header_cooldown >= 0
|
||||
else litellm_router.cooldown_time
|
||||
)
|
||||
)
|
||||
|
||||
increment_deployment_failures_for_current_minute(
|
||||
litellm_router_instance=litellm_router,
|
||||
deployment_id=deployment_id,
|
||||
)
|
||||
_set_cooldown_deployments(
|
||||
litellm_router_instance=litellm_router,
|
||||
exception_status=exception_status,
|
||||
original_exception=exception,
|
||||
deployment=deployment_id,
|
||||
time_to_cooldown=time_to_cooldown,
|
||||
)
|
||||
|
||||
verbose_router_logger.debug("Triggered cooldown for fallback deployment %s", deployment_id)
|
||||
except Exception as e: # noqa: BLE001 - best-effort cooldown trigger must never break the fallback response itself
|
||||
verbose_router_logger.debug("Error triggering cooldown for fallback deployment: %s", e)
|
||||
|
||||
|
||||
def fallback_attempt_key(fallback_target: object) -> str | None:
|
||||
"""
|
||||
|
|
@ -131,6 +252,28 @@ def get_fallback_model_group(fallbacks: list[Any], model_group: str) -> tuple[li
|
|||
return fallback_model_group, generic_fallback_idx
|
||||
|
||||
|
||||
PROVIDER_SCOPED_RESOURCE_KEYS: Final = ("input_file_id", "training_file")
|
||||
|
||||
|
||||
def _get_fallback_target_model_group(fallback_entry: str | Mapping[str, object]) -> str | None:
|
||||
if isinstance(fallback_entry, str):
|
||||
return fallback_entry
|
||||
target: Final = fallback_entry.get("model")
|
||||
return target if isinstance(target, str) else None
|
||||
|
||||
|
||||
def references_provider_scoped_resource(kwargs: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
True when the request names a file that only exists under one provider's credentials.
|
||||
|
||||
Batch and fine-tuning jobs are created from a file the caller already uploaded, and
|
||||
that file lives in the account of the deployment that stored it. Handing the id to a
|
||||
different model group can only fail, and the second provider's error replaces the
|
||||
error the caller actually needs to see.
|
||||
"""
|
||||
return any(kwargs.get(key) for key in PROVIDER_SCOPED_RESOURCE_KEYS)
|
||||
|
||||
|
||||
async def run_async_fallback(
|
||||
*args: tuple[Any],
|
||||
litellm_router: LitellmRouter,
|
||||
|
|
@ -176,6 +319,10 @@ async def run_async_fallback(
|
|||
|
||||
error_from_fallbacks = original_exception
|
||||
fallback_errors = (get_fallback_error_info(original_exception),)
|
||||
metadata_variable_name: Final = _get_router_metadata_variable_name(
|
||||
function_name=getattr(kwargs.get("original_function"), "__name__", None)
|
||||
)
|
||||
same_model_group_only: Final = references_provider_scoped_resource(kwargs)
|
||||
# Read out of kwargs and narrowed here rather than declared as a parameter: every caller
|
||||
# reaches this function by spreading a loosely-typed kwargs dict, so a declared parameter
|
||||
# would carry an annotation that no call site can actually be checked against.
|
||||
|
|
@ -188,6 +335,13 @@ async def run_async_fallback(
|
|||
for mg in fallback_model_group:
|
||||
if mg == original_model_group:
|
||||
continue
|
||||
if same_model_group_only and _get_fallback_target_model_group(mg) != original_model_group:
|
||||
verbose_router_logger.info(
|
||||
"Skipping fallback to model_group = %s: request is pinned to model_group = %s by its uploaded file",
|
||||
mask_sensitive_structure(mg),
|
||||
original_model_group,
|
||||
)
|
||||
continue
|
||||
attempt_key = fallback_attempt_key(mg)
|
||||
if attempt_key is not None:
|
||||
if attempt_key in attempted:
|
||||
|
|
@ -205,9 +359,10 @@ async def run_async_fallback(
|
|||
kwargs["model"] = mg
|
||||
elif isinstance(mg, dict):
|
||||
kwargs.update(mg)
|
||||
kwargs.setdefault("metadata", {}).update(
|
||||
{"model_group": kwargs.get("model", None)}
|
||||
) # update model_group used, if fallbacks are done
|
||||
kwargs[metadata_variable_name] = {
|
||||
**(kwargs.get(metadata_variable_name) or {}),
|
||||
"model_group": kwargs.get("model", None),
|
||||
}
|
||||
fallback_depth = fallback_depth + 1
|
||||
kwargs["fallback_depth"] = fallback_depth
|
||||
kwargs["max_fallbacks"] = max_fallbacks
|
||||
|
|
@ -236,6 +391,13 @@ async def run_async_fallback(
|
|||
kwargs=kwargs,
|
||||
original_exception=original_exception,
|
||||
)
|
||||
logging_obj = kwargs.get("litellm_logging_obj")
|
||||
if logging_obj is not None and logging_obj.model_call_details.get("has_logged_async_failure", False):
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=litellm_router,
|
||||
kwargs=kwargs,
|
||||
exception=e,
|
||||
)
|
||||
raise error_from_fallbacks
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ Pydantic models for Memory management endpoints.
|
|||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
|
@ -12,7 +11,7 @@ class LiteLLM_MemoryRow(BaseModel):
|
|||
memory_id: str
|
||||
key: str
|
||||
value: str
|
||||
metadata: Any | None = None
|
||||
metadata: object | None = None
|
||||
user_id: str | None = None
|
||||
team_id: str | None = None
|
||||
created_at: datetime | None = None
|
||||
|
|
@ -24,7 +23,7 @@ class LiteLLM_MemoryRow(BaseModel):
|
|||
class MemoryCreateRequest(BaseModel):
|
||||
key: str = Field(..., description="Memory key (acts as the namespace in the URL).")
|
||||
value: str = Field(..., description="Memory content. Typically markdown/text for LLM context.")
|
||||
metadata: Any | None = Field(
|
||||
metadata: object | None = Field(
|
||||
default=None,
|
||||
description="Optional JSON metadata (tags, structured fields).",
|
||||
)
|
||||
|
|
@ -40,7 +39,7 @@ class MemoryCreateRequest(BaseModel):
|
|||
|
||||
class MemoryUpdateRequest(BaseModel):
|
||||
value: str | None = None
|
||||
metadata: Any | None = None
|
||||
metadata: object | None = None
|
||||
# Only honored on create (when the row doesn't yet exist) and only for
|
||||
# PROXY_ADMIN callers — mirrors MemoryCreateRequest so admins can bootstrap
|
||||
# rows scoped to another user/team via PUT, not just POST.
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ class GroupByDimension(str, Enum):
|
|||
|
||||
class SpendMetrics(BaseModel):
|
||||
spend: float = Field(default=0.0)
|
||||
flat_cost: float = Field(default=0.0)
|
||||
prompt_tokens: int = Field(default=0)
|
||||
completion_tokens: int = Field(default=0)
|
||||
cache_read_input_tokens: int = Field(default=0)
|
||||
|
|
@ -75,6 +76,7 @@ class DailySpendData(BaseModel):
|
|||
|
||||
class DailySpendMetadata(BaseModel):
|
||||
total_spend: float = Field(default=0.0)
|
||||
total_flat_cost: float = Field(default=0.0)
|
||||
total_prompt_tokens: int = Field(default=0)
|
||||
total_completion_tokens: int = Field(default=0)
|
||||
total_tokens: int = Field(default=0)
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ litellm.Router Types - includes RouterConfig, UpdateRouterConfig, ModelInfo etc
|
|||
import datetime
|
||||
import enum
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Final, Generic, Literal, TypeVar, get_type_hints
|
||||
from typing import Any, ClassVar, Final, Generic, Literal, TypeVar, get_type_hints
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
|
@ -127,6 +127,14 @@ class UpdateRouterConfig(BaseModel):
|
|||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
|
||||
def _as_utc(value: datetime.datetime | None) -> datetime.datetime | None:
|
||||
if value is None:
|
||||
return None
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=datetime.timezone.utc)
|
||||
return value.astimezone(datetime.timezone.utc)
|
||||
|
||||
|
||||
class ModelInfo(MirroredPricingParams):
|
||||
id: str | None # Allow id to be optional on input, but it will always be present as a str in the model instance
|
||||
db_model: bool = False # used for proxy - to separate models which are stored in the db vs. config.
|
||||
|
|
@ -151,6 +159,17 @@ class ModelInfo(MirroredPricingParams):
|
|||
# admin-toggled pause flag; mirrors LiteLLM_ProxyModelTable.blocked
|
||||
blocked: bool | None = None
|
||||
|
||||
# Bounds live on the model rather than litellm.constants: names there reach
|
||||
# litellm/__init__ through several modules' star re-exports, and a Final rebound that
|
||||
# way trips the basedpyright gate.
|
||||
MAX_PTU_COUNT: ClassVar[int] = 1_000_000
|
||||
MAX_COST_PER_PTU_PER_HOUR: ClassVar[float] = 1_000_000.0
|
||||
|
||||
ptu_count: int | None = None
|
||||
cost_per_ptu_per_hour: float | None = None
|
||||
ptu_effective_from: datetime.datetime | None = None
|
||||
ptu_effective_to: datetime.datetime | None = None
|
||||
|
||||
def __init__(self, id: str | int | None = None, **params) -> None:
|
||||
if id is None:
|
||||
id = str(uuid.uuid4()) # Generate a UUID if id is None or not provided
|
||||
|
|
@ -158,6 +177,23 @@ class ModelInfo(MirroredPricingParams):
|
|||
id = str(id)
|
||||
super().__init__(id=id, **params)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_ptu_bounds(self) -> "ModelInfo":
|
||||
if self.ptu_count is not None and not 0 < self.ptu_count <= self.MAX_PTU_COUNT:
|
||||
raise ValueError(f"ptu_count must be a positive integer no greater than {self.MAX_PTU_COUNT}")
|
||||
if (
|
||||
self.cost_per_ptu_per_hour is not None
|
||||
and not 0 <= self.cost_per_ptu_per_hour <= self.MAX_COST_PER_PTU_PER_HOUR
|
||||
):
|
||||
raise ValueError(
|
||||
f"cost_per_ptu_per_hour must be a finite number between 0 and {self.MAX_COST_PER_PTU_PER_HOUR}"
|
||||
)
|
||||
start: Final = _as_utc(self.ptu_effective_from)
|
||||
end: Final = _as_utc(self.ptu_effective_to)
|
||||
if start is not None and end is not None and end <= start:
|
||||
raise ValueError("ptu_effective_to must be after ptu_effective_from")
|
||||
return self
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
def __contains__(self, key) -> bool:
|
||||
|
|
@ -422,6 +458,9 @@ class LiteLLMParamsTypedDict(TypedDict, total=False):
|
|||
max_budget: float | None
|
||||
budget_duration: str | None
|
||||
|
||||
# per-deployment cooldown override
|
||||
cooldown_time: float | None
|
||||
|
||||
|
||||
class DeploymentTypedDict(TypedDict, total=False):
|
||||
model_name: Required[str]
|
||||
|
|
@ -513,6 +552,9 @@ class AllowedFailsPolicy(BaseModel):
|
|||
RateLimitErrorAllowedFails: int | None = None
|
||||
ContentPolicyViolationErrorAllowedFails: int | None = None
|
||||
InternalServerErrorAllowedFails: int | None = None
|
||||
ServiceUnavailableErrorAllowedFails: int | None = None
|
||||
BadGatewayErrorAllowedFails: int | None = None
|
||||
NotFoundErrorAllowedFails: int | None = None
|
||||
|
||||
|
||||
class AlertingConfig(BaseModel):
|
||||
|
|
|
|||
|
|
@ -3129,6 +3129,7 @@ class StandardAuditLogPayload(TypedDict):
|
|||
class StandardLoggingPayload(TypedDict):
|
||||
id: str
|
||||
trace_id: str # Trace multiple LLM calls belonging to same overall request (e.g. fallbacks/retries)
|
||||
session_id: str # End-user/conversation session id (litellm_session_id), independent of trace_id
|
||||
litellm_call_id: str | None # UUID returned in x-litellm-call-id response header
|
||||
call_type: str
|
||||
stream: bool | None
|
||||
|
|
|
|||
104
litellm/utils.py
104
litellm/utils.py
|
|
@ -711,14 +711,71 @@ def _remove_thought_signatures_from_messages(messages: list, thought_signature_s
|
|||
return processed_messages
|
||||
|
||||
|
||||
def _restore_correlation_context_if_supported(logging_obj: object) -> None:
|
||||
"""Call logging_obj._restore_correlation_context() if it's actually there.
|
||||
|
||||
Some call sites (tests, narrow unit paths) inject a minimal stand-in
|
||||
object as litellm_logging_obj instead of a real Logging instance - this
|
||||
method is new plumbing specific to request_correlation_in_logs, not part
|
||||
of any pre-existing stand-in's expected interface. `object` (not `Any`)
|
||||
is deliberate: the getattr() below is exactly how this stays type-safe
|
||||
while still tolerating a stand-in that lacks the method.
|
||||
"""
|
||||
restore: Final = getattr(logging_obj, "_restore_correlation_context", None)
|
||||
if restore is not None:
|
||||
restore()
|
||||
|
||||
|
||||
def _is_streaming_response_for_correlation(result: object) -> bool:
|
||||
"""True if `result` is a lazy stream wrapper rather than an already-complete response.
|
||||
|
||||
Only wrapper_async() consults this - it must NOT restore the originating
|
||||
Task's trace_id/session_id as soon as a streaming call returns this: the
|
||||
caller is about to iterate it over however many subsequent lines of their
|
||||
own code, and those log lines should still show this call's ids, not the
|
||||
pre-call ones. This is safe specifically because each async call already
|
||||
runs in its own asyncio Task with its own copy of the contextvars, so
|
||||
leaving it "open" can only affect that one Task, never a different,
|
||||
unrelated future request - Tasks, unlike a thread pool's worker threads,
|
||||
are never recycled across requests. The corresponding terminal handler
|
||||
(async_success_handler, dispatched once the full stream is actually
|
||||
assembled) is what restores it once streaming genuinely finishes.
|
||||
|
||||
wrapper() (the sync path) does NOT consult this at all: sync calls pass
|
||||
supports_correlation_logging=False into function_setup()/Logging(), so
|
||||
they never stamp trace_id/session_id in the first place - a plain OS
|
||||
thread has no per-call isolation the way an asyncio Task does, and a
|
||||
thread pool's worker threads *are* recycled across unrelated requests, so
|
||||
stamping ids there without a safe restore mechanism could permanently
|
||||
misattribute a later, unrelated request's logs. Full sync support is
|
||||
deferred to a follow-up PR with its own restore mechanism; see
|
||||
Logging.__init__'s supports_correlation_logging parameter.
|
||||
|
||||
Genuinely circular otherwise: utils.py -> streaming_handler.py ->
|
||||
redact_messages.py -> llms/vertex_ai/common_utils.py -> utils.py, which
|
||||
needs names (supports_response_schema, etc.) this module hasn't finished
|
||||
defining yet at that point in its own top-to-bottom execution.
|
||||
"""
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
return isinstance(result, CustomStreamWrapper)
|
||||
|
||||
|
||||
# Runs once per call to check if the user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
|
||||
def function_setup(
|
||||
original_function: str, rules_obj, start_time, *args, **kwargs
|
||||
): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
|
||||
original_function: str,
|
||||
rules_obj: Rules,
|
||||
start_time: datetime.datetime,
|
||||
*args: Any, # positional passthrough to the wrapped LLM call (ANN401 ignored, see ruff-strict.toml)
|
||||
is_async_call: bool = True,
|
||||
**kwargs: Any, # kwargs-ok: forwarded to Logging()/callbacks, varies per call_type
|
||||
) -> tuple[LiteLLMLoggingObject, dict[str, Any]]:
|
||||
### NOTICES ###
|
||||
if litellm.set_verbose is True:
|
||||
verbose_logger.warning(
|
||||
"`litellm.set_verbose` is deprecated. Please set `os.environ['LITELLM_LOG'] = 'DEBUG'` for debug logs."
|
||||
)
|
||||
logging_obj: LiteLLMLoggingObject | None = None # rebind-ok: set to the real object further down on success
|
||||
try:
|
||||
global callback_list, add_breadcrumb, user_logger_fn, Logging
|
||||
|
||||
|
|
@ -1001,7 +1058,8 @@ def function_setup(
|
|||
):
|
||||
stream = True
|
||||
get_litellm_logging_class: Final = getattr(sys.modules[__name__], "get_litellm_logging_class")
|
||||
logging_obj: Final = get_litellm_logging_class()( # Victim for object pool
|
||||
# Victim for object pool
|
||||
logging_obj = get_litellm_logging_class()( # rebind-ok: 2nd assignment to logging_obj (see initial None above)
|
||||
model=model,
|
||||
messages=messages,
|
||||
stream=stream,
|
||||
|
|
@ -1016,6 +1074,7 @@ def function_setup(
|
|||
dynamic_async_failure_callbacks=dynamic_async_failure_callbacks,
|
||||
kwargs=kwargs,
|
||||
applied_guardrails=applied_guardrails,
|
||||
supports_correlation_logging=is_async_call,
|
||||
)
|
||||
|
||||
## check if metadata is passed in
|
||||
|
|
@ -1040,6 +1099,15 @@ def function_setup(
|
|||
)
|
||||
return logging_obj, kwargs
|
||||
except Exception as e:
|
||||
# If Logging() was constructed above before this failed, its __init__ already
|
||||
# mutated trace_id_var/session_id_var - restore them *before* logging the
|
||||
# exception below, since we're about to raise without ever returning
|
||||
# logging_obj to the caller's wrapper()/wrapper_async() (which would
|
||||
# otherwise be the one doing this restore). Restoring first means this
|
||||
# diagnostic log line itself doesn't get stamped with a call's ids when
|
||||
# that call never actually produced a usable logging object.
|
||||
if logging_obj is not None:
|
||||
_restore_correlation_context_if_supported(logging_obj)
|
||||
verbose_logger.exception("litellm.utils.py::function_setup() - [Non-Blocking] Error in function_setup")
|
||||
raise e
|
||||
|
||||
|
|
@ -1296,7 +1364,9 @@ def client(original_function):
|
|||
|
||||
try:
|
||||
if logging_obj is None:
|
||||
logging_obj, kwargs = function_setup(original_function.__name__, rules_obj, start_time, *args, **kwargs)
|
||||
logging_obj, kwargs = function_setup(
|
||||
original_function.__name__, rules_obj, start_time, *args, is_async_call=False, **kwargs
|
||||
)
|
||||
|
||||
# Type assertion: logging_obj is guaranteed to be non-None after function_setup
|
||||
assert logging_obj is not None, "logging_obj should not be None after function_setup"
|
||||
|
|
@ -1807,9 +1877,11 @@ def client(original_function):
|
|||
kwargs["retry_strategy"] = "exponential_backoff_retry"
|
||||
elif isinstance(e, openai.APIError): # generic api error
|
||||
kwargs["retry_strategy"] = "constant_retry"
|
||||
return await litellm.acompletion_with_retries(*args, **kwargs)
|
||||
result = await litellm.acompletion_with_retries(*args, **kwargs)
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
return result
|
||||
elif (
|
||||
isinstance(e, litellm.exceptions.ContextWindowExceededError)
|
||||
and context_window_fallback_dict
|
||||
|
|
@ -1820,7 +1892,8 @@ def client(original_function):
|
|||
args[0] = context_window_fallback_dict[model]
|
||||
else:
|
||||
kwargs["model"] = context_window_fallback_dict[model]
|
||||
return await original_function(*args, **kwargs)
|
||||
result = await original_function(*args, **kwargs)
|
||||
return result
|
||||
elif call_type == CallTypes.aresponses.value:
|
||||
_is_litellm_router_call = "model_group" in (
|
||||
kwargs.get("metadata") or {}
|
||||
|
|
@ -1837,9 +1910,11 @@ def client(original_function):
|
|||
kwargs["retry_strategy"] = "exponential_backoff_retry"
|
||||
elif isinstance(e, openai.APIError): # generic api error
|
||||
kwargs["retry_strategy"] = "constant_retry"
|
||||
return await litellm.aresponses_with_retries(*args, **kwargs)
|
||||
result = await litellm.aresponses_with_retries(*args, **kwargs)
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
return result
|
||||
|
||||
deployment_num_retries: Final = kwargs.get("num_retries")
|
||||
if deployment_num_retries is not None:
|
||||
|
|
@ -1849,6 +1924,21 @@ def client(original_function):
|
|||
setattr(e, "timeout", timeout)
|
||||
raise e
|
||||
|
||||
finally:
|
||||
# Restore trace_id/session_id contextvars to their pre-call value once
|
||||
# this call (in this asyncio Task) is fully done - see
|
||||
# request_correlation_in_logs. Unlike wrapper()'s sync path, it's safe to
|
||||
# skip restoring when returning a stream: each async call already runs in
|
||||
# its own Task with its own copy of the contextvars (asyncio.create_task
|
||||
# copies context at creation), so leaving this Task's own view "open"
|
||||
# while the caller iterates the stream can only affect that one Task -
|
||||
# never a different, unrelated future request, since Tasks (unlike a
|
||||
# thread pool's worker threads) are never recycled across requests. The
|
||||
# corresponding terminal handler (async_success_handler) restores it once
|
||||
# streaming genuinely finishes; aclose()/__del__ cover early termination.
|
||||
if not _is_streaming_response_for_correlation(result):
|
||||
_restore_correlation_context_if_supported(logging_obj)
|
||||
|
||||
get_coroutine_checker: Final = getattr(sys.modules[__name__], "get_coroutine_checker")
|
||||
is_coroutine: Final = get_coroutine_checker().is_async_callable(original_function)
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,18 +1,18 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 3114
|
||||
"limit": 3106
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 71
|
||||
},
|
||||
"ANN003": {
|
||||
"limit": 834
|
||||
"limit": 832
|
||||
},
|
||||
"ANN201": {
|
||||
"limit": 2031
|
||||
},
|
||||
"ANN202": {
|
||||
"limit": 865
|
||||
"limit": 861
|
||||
},
|
||||
"ANN204": {
|
||||
"limit": 713
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 133
|
||||
},
|
||||
"ANN401": {
|
||||
"limit": 1555
|
||||
"limit": 1495
|
||||
},
|
||||
"ASYNC230": {
|
||||
"limit": 11
|
||||
|
|
@ -39,7 +39,7 @@
|
|||
"limit": 505
|
||||
},
|
||||
"B009": {
|
||||
"limit": 81
|
||||
"limit": 79
|
||||
},
|
||||
"B010": {
|
||||
"limit": 190
|
||||
|
|
@ -78,7 +78,7 @@
|
|||
"limit": 1
|
||||
},
|
||||
"C901": {
|
||||
"limit": 314
|
||||
"limit": 313
|
||||
},
|
||||
"D419": {
|
||||
"limit": 6
|
||||
|
|
@ -234,7 +234,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"TID251": {
|
||||
"limit": 1238
|
||||
"limit": 1226
|
||||
},
|
||||
"TRY002": {
|
||||
"limit": 528
|
||||
|
|
|
|||
|
|
@ -16,6 +16,17 @@ external = [
|
|||
"PLC0415", "E402", "BLE001", "ARG002", "S102", "S324", "S606", "D401", "F403", "F405",
|
||||
]
|
||||
|
||||
[lint.per-file-ignores]
|
||||
# ANN401 (explicit `Any` disallowed) has no per-line/function-level ignore mechanism
|
||||
# in ruff, only file-level. These two files each have a handful of parameters that
|
||||
# are genuinely heterogeneous with no fitting concrete type: a response object that
|
||||
# varies across every LLM call type (completion/embedding/transcription/etc. each
|
||||
# return a different shape), and *args/**kwargs forwarded verbatim with no fixed
|
||||
# shape. Tried the closest existing union (CostResponseTypes) first; basedpyright
|
||||
# caught a real mismatch, confirming Any is correct here, not a shortcut.
|
||||
"litellm/litellm_core_utils/litellm_logging.py" = ["ANN401"]
|
||||
"litellm/utils.py" = ["ANN401"]
|
||||
|
||||
[lint.mccabe]
|
||||
max-complexity = 15
|
||||
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ model LiteLLM_BudgetTable {
|
|||
end_users LiteLLM_EndUserTable[] // multiple end-users can have the same budget
|
||||
tags LiteLLM_TagTable[] // multiple tags can have the same budget
|
||||
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
|
||||
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
|
||||
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
|
||||
}
|
||||
|
||||
// Models on proxy
|
||||
|
|
@ -893,6 +893,7 @@ model LiteLLM_DailyTeamSpend {
|
|||
api_requests BigInt @default(0)
|
||||
successful_requests BigInt @default(0)
|
||||
failed_requests BigInt @default(0)
|
||||
ptu_flat_cost Float @default(0.0)
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
||||
|
|
|
|||
|
|
@ -471,8 +471,8 @@ def test_get_final_response_obj():
|
|||
litellm.turn_off_message_logging = False
|
||||
|
||||
|
||||
def test_get_standard_logging_payload_trace_id():
|
||||
"""Test _get_standard_logging_payload_trace_id with different input scenarios"""
|
||||
def testget_standard_logging_payload_trace_id():
|
||||
"""Test get_standard_logging_payload_trace_id with different input scenarios"""
|
||||
# Test case 1: When litellm_trace_id is provided in litellm_params
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
|
@ -482,33 +482,134 @@ def test_get_standard_logging_payload_trace_id():
|
|||
|
||||
# Test when litellm_trace_id is in litellm_params
|
||||
litellm_params = {"litellm_trace_id": "dynamic-trace-id"}
|
||||
result = StandardLoggingPayloadSetup._get_standard_logging_payload_trace_id(
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == "dynamic-trace-id"
|
||||
|
||||
# Test case 2: When litellm_trace_id is not provided in litellm_params
|
||||
litellm_params = {}
|
||||
result = StandardLoggingPayloadSetup._get_standard_logging_payload_trace_id(
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == "default-trace-id"
|
||||
|
||||
# Test case 3: When litellm_params is None
|
||||
result = StandardLoggingPayloadSetup._get_standard_logging_payload_trace_id(
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
||||
logging_obj=mock_logging_obj, litellm_params={}
|
||||
)
|
||||
assert result == "default-trace-id"
|
||||
|
||||
# Test case 4: When litellm_trace_id in params is not a string
|
||||
litellm_params = {"litellm_trace_id": 12345}
|
||||
result = StandardLoggingPayloadSetup._get_standard_logging_payload_trace_id(
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == "12345"
|
||||
assert isinstance(result, str)
|
||||
|
||||
|
||||
def testget_standard_logging_payload_trace_id_prioritizes_trace_id_when_flag_on(monkeypatch):
|
||||
"""With request_correlation_in_logs on, an explicit litellm_trace_id wins over litellm_session_id."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.litellm_trace_id = "default-trace-id"
|
||||
|
||||
litellm_params = {"litellm_trace_id": "the-trace-id", "litellm_session_id": "the-session-id"}
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == "the-trace-id"
|
||||
|
||||
|
||||
def testget_standard_logging_payload_trace_id_prioritizes_session_id_when_flag_off(monkeypatch):
|
||||
"""With request_correlation_in_logs off (default), legacy behavior is preserved:
|
||||
litellm_session_id still wins over litellm_trace_id."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", False)
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.litellm_trace_id = "default-trace-id"
|
||||
|
||||
litellm_params = {"litellm_trace_id": "the-trace-id", "litellm_session_id": "the-session-id"}
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == "the-session-id"
|
||||
|
||||
|
||||
def testget_standard_logging_payload_session_id_when_flag_on(monkeypatch):
|
||||
"""Test get_standard_logging_payload_session_id with different input scenarios, flag enabled"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.litellm_session_id = ""
|
||||
|
||||
# Test case 1: litellm_session_id provided directly in litellm_params
|
||||
litellm_params = {"litellm_session_id": "dynamic-session-id"}
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == "dynamic-session-id"
|
||||
|
||||
# Test case 2: falls back to metadata.session_id when not in litellm_params directly
|
||||
litellm_params = {"metadata": {"session_id": "metadata-session-id"}}
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == "metadata-session-id"
|
||||
|
||||
# Test case 3: falls back to logging_obj.litellm_session_id when nothing else is set
|
||||
mock_logging_obj.litellm_session_id = "obj-session-id"
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
||||
logging_obj=mock_logging_obj, litellm_params={}
|
||||
)
|
||||
assert result == "obj-session-id"
|
||||
|
||||
# Test case 4: empty string when no session id was supplied anywhere
|
||||
mock_logging_obj.litellm_session_id = ""
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
||||
logging_obj=mock_logging_obj, litellm_params={}
|
||||
)
|
||||
assert result == ""
|
||||
|
||||
# Test case 5: non-string session id in params is coerced to str
|
||||
litellm_params = {"litellm_session_id": 98765}
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == "98765"
|
||||
assert isinstance(result, str)
|
||||
|
||||
# Test case 6: trace_id and session_id are independent - passing only a trace id
|
||||
# must not populate session_id
|
||||
litellm_params = {"litellm_trace_id": "some-trace-id"}
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == ""
|
||||
|
||||
|
||||
def testget_standard_logging_payload_session_id_empty_when_flag_off(monkeypatch):
|
||||
"""When request_correlation_in_logs is off (default), session_id is always empty,
|
||||
even if litellm_session_id was explicitly supplied - preserves the pre-existing
|
||||
StandardLoggingPayload shape for callers who haven't opted in."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", False)
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.litellm_session_id = "obj-session-id"
|
||||
|
||||
litellm_params = {"litellm_session_id": "dynamic-session-id"}
|
||||
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
||||
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
||||
)
|
||||
assert result == ""
|
||||
|
||||
|
||||
def test_truncate_standard_logging_payload():
|
||||
"""
|
||||
1. original messages, response, and error_str should NOT BE MODIFIED, since these are from kwargs
|
||||
|
|
|
|||
|
|
@ -2697,6 +2697,79 @@ def test_get_timeout_from_request():
|
|||
assert timeout == 90.5
|
||||
|
||||
|
||||
def test_add_litellm_data_for_backend_llm_call_marks_client_side_timeout():
|
||||
"""A caller-supplied x-litellm-timeout must be marked with client_side_timeout=True,
|
||||
so the router's fallback-cooldown trigger can tell it apart from a deployment
|
||||
actually timing out (a caller could otherwise force every deployment in a fallback
|
||||
chain to look unhealthy with a single near-zero timeout request)."""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test_api_key")
|
||||
|
||||
data = LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call(
|
||||
headers={"x-litellm-timeout": "0.001"},
|
||||
request_data={},
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
assert data["timeout"] == 0.001
|
||||
assert data["client_side_timeout"] is True
|
||||
|
||||
data_without_header = LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call(
|
||||
headers={},
|
||||
request_data={},
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
assert "client_side_timeout" not in data_without_header
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"request_data",
|
||||
[
|
||||
{"timeout": 0.001},
|
||||
{"request_timeout": 0.001},
|
||||
{"stream_timeout": 0.001},
|
||||
],
|
||||
)
|
||||
def test_add_litellm_data_for_backend_llm_call_marks_client_side_timeout_from_body(
|
||||
request_data,
|
||||
):
|
||||
"""Router._get_timeout resolves the effective timeout from kwargs["timeout"],
|
||||
kwargs["request_timeout"], or kwargs["stream_timeout"], and a caller can supply any
|
||||
of those directly in the request body, not just via the x-litellm-timeout header.
|
||||
Missing this would let a caller force a 408 on every deployment in a fallback chain
|
||||
without it being recognized as caller-controlled, cooling down deployments other
|
||||
tenants rely on."""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test_api_key")
|
||||
|
||||
data = LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call(
|
||||
headers={},
|
||||
request_data=request_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
assert data["client_side_timeout"] is True
|
||||
|
||||
|
||||
def test_add_litellm_data_for_backend_llm_call_ignores_forged_client_side_timeout():
|
||||
"""The caller-supplied client_side_timeout key itself must never be trusted verbatim:
|
||||
the marker is always recomputed from the actual timeout sources, so a caller can't
|
||||
forge client_side_timeout=True to dodge cooldown on a real deployment failure."""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test_api_key")
|
||||
|
||||
data = LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call(
|
||||
headers={},
|
||||
request_data={"client_side_timeout": True},
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
assert "client_side_timeout" not in data
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"ui_exists, ui_has_content",
|
||||
[
|
||||
|
|
|
|||
779
tests/router_unit_tests/test_router_cooldown_per_deployment.py
Normal file
779
tests/router_unit_tests/test_router_cooldown_per_deployment.py
Normal file
|
|
@ -0,0 +1,779 @@
|
|||
"""
|
||||
Tests for per-deployment cooldown policy overrides, DualCache TTL correction,
|
||||
and fallback-path cooldown gap fix.
|
||||
"""
|
||||
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.router_utils.cooldown_cache import CooldownCache, CooldownCacheValue
|
||||
from litellm.router_utils.cooldown_handlers import (
|
||||
_get_deployment_cooldown_policy,
|
||||
_has_explicit_allowed_fails_policy_for_exception,
|
||||
_resolve_allowed_fails_from_policy,
|
||||
_should_cooldown_deployment,
|
||||
mark_advisor_orchestration_failure,
|
||||
should_cooldown_based_on_allowed_fails_policy,
|
||||
)
|
||||
from litellm.router_utils.fallback_event_handlers import _trigger_cooldown_for_failed_deployment
|
||||
from litellm.types.router import AllowedFailsPolicy
|
||||
|
||||
|
||||
def _make_router(model_list: list, **kwargs) -> Router:
|
||||
return Router(model_list=model_list, **kwargs)
|
||||
|
||||
|
||||
class TestDeploymentLevelAllowedFails:
|
||||
def test_deployment_level_allowed_fails_overrides_router_level(self):
|
||||
"""
|
||||
A deployment with model_info.allowed_fails=0 must enter cooldown after 1
|
||||
failure even when the router-level allowed_fails=10.
|
||||
"""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {
|
||||
"id": "primary",
|
||||
"allowed_fails": 0,
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {"id": "secondary"},
|
||||
},
|
||||
],
|
||||
allowed_fails=10,
|
||||
)
|
||||
|
||||
_exception = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
should_cooldown = _should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="primary",
|
||||
exception_status=429,
|
||||
original_exception=_exception,
|
||||
)
|
||||
|
||||
assert should_cooldown is True, "Deployment-level allowed_fails=0 should force cooldown after first failure"
|
||||
|
||||
def test_deployment_level_allowed_fails_does_not_affect_other_deployments(self):
|
||||
"""
|
||||
A deployment without model_info.allowed_fails must still use the router-level
|
||||
allowed_fails and not be pulled into cooldown prematurely.
|
||||
"""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {
|
||||
"id": "primary",
|
||||
"allowed_fails": 0,
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {"id": "secondary"},
|
||||
},
|
||||
],
|
||||
allowed_fails=10,
|
||||
)
|
||||
|
||||
_exception = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
should_cooldown = _should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="secondary",
|
||||
exception_status=429,
|
||||
original_exception=_exception,
|
||||
)
|
||||
|
||||
assert should_cooldown is False, (
|
||||
"secondary has no deployment-level policy; with allowed_fails=10 it should not cool down on first failure"
|
||||
)
|
||||
|
||||
|
||||
class TestDeploymentLevelAllowedFailsPolicyByExceptionType:
|
||||
def test_rate_limit_error_triggers_cooldown_with_zero_threshold(self):
|
||||
"""
|
||||
RateLimitErrorAllowedFails=0 must trigger cooldown after 1 RateLimitError
|
||||
even when allowed_fails=5 for other exception types.
|
||||
"""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {
|
||||
"id": "primary",
|
||||
"allowed_fails_policy": {
|
||||
"RateLimitErrorAllowedFails": 0,
|
||||
"InternalServerErrorAllowedFails": 5,
|
||||
},
|
||||
},
|
||||
},
|
||||
],
|
||||
allowed_fails=10,
|
||||
)
|
||||
|
||||
rate_limit_exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
should_cooldown = _should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="primary",
|
||||
exception_status=429,
|
||||
original_exception=rate_limit_exc,
|
||||
)
|
||||
|
||||
assert should_cooldown is True, "RateLimitErrorAllowedFails=0 must trigger cooldown on first rate limit error"
|
||||
|
||||
def test_internal_server_error_respects_per_exception_threshold(self):
|
||||
"""
|
||||
InternalServerErrorAllowedFails=5 must allow 5 InternalServerErrors before cooldown.
|
||||
"""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {
|
||||
"id": "primary",
|
||||
"allowed_fails_policy": {
|
||||
"RateLimitErrorAllowedFails": 0,
|
||||
"InternalServerErrorAllowedFails": 5,
|
||||
},
|
||||
},
|
||||
},
|
||||
],
|
||||
allowed_fails=10,
|
||||
)
|
||||
|
||||
ise = litellm.InternalServerError("Internal error", "openai", "gpt-4")
|
||||
|
||||
for _ in range(5):
|
||||
should_cooldown = _should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="primary",
|
||||
exception_status=500,
|
||||
original_exception=ise,
|
||||
)
|
||||
assert should_cooldown is False, "Should not cooldown within the allowed_fails threshold"
|
||||
|
||||
should_cooldown = _should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="primary",
|
||||
exception_status=500,
|
||||
original_exception=ise,
|
||||
)
|
||||
assert should_cooldown is True, "Should cooldown after exceeding InternalServerErrorAllowedFails=5"
|
||||
|
||||
|
||||
class TestExceptionTypeCountersTrackedIndependently:
|
||||
def test_cache_key_suffix_separates_exception_type_counters(self):
|
||||
"""
|
||||
When cache_key_suffix is provided, fail counters for different exception types
|
||||
must be independent; RateLimitError fails must not bleed into generic counters.
|
||||
"""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {"id": "primary"},
|
||||
},
|
||||
],
|
||||
allowed_fails=10,
|
||||
)
|
||||
|
||||
rate_limit_exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
ise = litellm.InternalServerError("Internal error", "openai", "gpt-4")
|
||||
|
||||
for _ in range(3):
|
||||
should_cooldown_based_on_allowed_fails_policy(
|
||||
litellm_router_instance=router,
|
||||
deployment="primary",
|
||||
original_exception=rate_limit_exc,
|
||||
allowed_fails_override=5,
|
||||
cache_key_suffix="RateLimitError",
|
||||
)
|
||||
|
||||
rl_counter = router.failed_calls.get_cache(key="primary:RateLimitError") or 0
|
||||
generic_counter = router.failed_calls.get_cache(key="primary:generic") or 0
|
||||
|
||||
assert rl_counter == 3, "RateLimitError counter should be 3"
|
||||
assert generic_counter == 0, "generic counter must be untouched by RateLimitError increments"
|
||||
|
||||
should_cooldown_based_on_allowed_fails_policy(
|
||||
litellm_router_instance=router,
|
||||
deployment="primary",
|
||||
original_exception=ise,
|
||||
allowed_fails_override=5,
|
||||
cache_key_suffix="generic",
|
||||
)
|
||||
|
||||
generic_counter_after = router.failed_calls.get_cache(key="primary:generic") or 0
|
||||
rl_counter_after = router.failed_calls.get_cache(key="primary:RateLimitError") or 0
|
||||
|
||||
assert generic_counter_after == 1, "generic counter should now be 1"
|
||||
assert rl_counter_after == 3, "RateLimitError counter must remain unchanged after InternalServerError"
|
||||
|
||||
|
||||
class TestCooldownCacheTTLCorrection:
|
||||
def _make_cooldown_cache(self) -> CooldownCache:
|
||||
in_memory = InMemoryCache()
|
||||
dual_cache = DualCache(in_memory_cache=in_memory)
|
||||
return CooldownCache(cache=dual_cache, default_cooldown_time=60.0)
|
||||
|
||||
def test_expired_entry_evicted_and_not_returned(self):
|
||||
"""
|
||||
An entry with timestamp+cooldown_time in the past must be evicted from
|
||||
in-memory cache and excluded from the active cooldown list.
|
||||
"""
|
||||
cc = self._make_cooldown_cache()
|
||||
model_id = "expired-deployment"
|
||||
key = CooldownCache.get_cooldown_cache_key(model_id)
|
||||
|
||||
expired_value: CooldownCacheValue = {
|
||||
"exception_received": "Rate limit",
|
||||
"status_code": "429",
|
||||
"timestamp": time.time() - 120.0,
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
|
||||
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
assert active == [], "Expired cooldown entry must not appear in active cooldowns"
|
||||
assert cc.cache.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache"
|
||||
|
||||
def test_active_entry_is_returned(self):
|
||||
"""
|
||||
An entry whose cooldown window has not elapsed must appear in the active list.
|
||||
"""
|
||||
cc = self._make_cooldown_cache()
|
||||
model_id = "active-deployment"
|
||||
key = CooldownCache.get_cooldown_cache_key(model_id)
|
||||
|
||||
active_value: CooldownCacheValue = {
|
||||
"exception_received": "Rate limit",
|
||||
"status_code": "429",
|
||||
"timestamp": time.time(),
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, active_value, ttl=60)
|
||||
|
||||
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
assert len(active) == 1
|
||||
assert active[0][0] == model_id
|
||||
|
||||
def test_ttl_corrected_when_in_memory_expiry_far_exceeds_remaining(self):
|
||||
"""
|
||||
When DualCache backfills from Redis using the default 600s TTL, the in-memory
|
||||
TTL must be corrected to min(remaining, 60) seconds.
|
||||
"""
|
||||
cc = self._make_cooldown_cache()
|
||||
model_id = "backfilled-deployment"
|
||||
key = CooldownCache.get_cooldown_cache_key(model_id)
|
||||
|
||||
remaining = 30.0
|
||||
value: CooldownCacheValue = {
|
||||
"exception_received": "Rate limit",
|
||||
"status_code": "429",
|
||||
"timestamp": time.time() - (60.0 - remaining),
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, value, ttl=600)
|
||||
|
||||
before_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
assert before_expiry is not None
|
||||
|
||||
cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
assert after_expiry is not None
|
||||
corrected_remaining = after_expiry - time.time()
|
||||
assert corrected_remaining <= 60.0, "Corrected TTL must not exceed 60s"
|
||||
assert corrected_remaining > 0, "Corrected TTL must be positive (cooldown still active)"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_expired_entry_evicted(self):
|
||||
"""
|
||||
Async path must also evict expired entries.
|
||||
"""
|
||||
cc = self._make_cooldown_cache()
|
||||
model_id = "async-expired"
|
||||
key = CooldownCache.get_cooldown_cache_key(model_id)
|
||||
|
||||
expired_value: CooldownCacheValue = {
|
||||
"exception_received": "Rate limit",
|
||||
"status_code": "429",
|
||||
"timestamp": time.time() - 120.0,
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
|
||||
active = await cc.async_get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
assert active == [], "Expired entry must not appear in async active cooldowns"
|
||||
assert cc.cache.in_memory_cache.get_cache(key) is None
|
||||
|
||||
|
||||
class TestFallbackDeploymentCooldown:
|
||||
def test_trigger_cooldown_for_failed_deployment_calls_set_cooldown(self):
|
||||
"""
|
||||
_trigger_cooldown_for_failed_deployment must call _set_cooldown_deployments
|
||||
with the deployment ID stamped on the exception.
|
||||
"""
|
||||
mock_router = MagicMock()
|
||||
mock_router.cooldown_time = 60.0
|
||||
mock_router.get_model_info.return_value = None
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
exc.failed_deployment_id = "fallback-deployment"
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown:
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=mock_router,
|
||||
kwargs={},
|
||||
exception=exc,
|
||||
)
|
||||
|
||||
mock_set_cooldown.assert_called_once()
|
||||
call_kwargs = mock_set_cooldown.call_args[1]
|
||||
assert call_kwargs["deployment"] == "fallback-deployment"
|
||||
assert call_kwargs["original_exception"] is exc
|
||||
|
||||
def test_trigger_cooldown_no_op_when_deployment_id_missing(self):
|
||||
"""
|
||||
_trigger_cooldown_for_failed_deployment must not raise and must skip
|
||||
_set_cooldown_deployments when the exception has no failed_deployment_id.
|
||||
"""
|
||||
mock_router = MagicMock()
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown:
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=mock_router,
|
||||
kwargs={},
|
||||
exception=RuntimeError("no stamped deployment id"),
|
||||
)
|
||||
|
||||
mock_set_cooldown.assert_not_called()
|
||||
|
||||
def test_trigger_cooldown_does_not_trust_caller_supplied_metadata_bucket(self):
|
||||
"""
|
||||
A metadata bucket can't reliably be told apart from a caller-supplied one
|
||||
without knowing the call's function_name, so a client with permission to
|
||||
set metadata must not be able to get an arbitrary deployment cooled down
|
||||
by forging a deployment_model_name marker.
|
||||
"""
|
||||
mock_router = MagicMock()
|
||||
mock_router.cooldown_time = 60.0
|
||||
mock_router.get_model_info.return_value = None
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
kwargs = {
|
||||
"metadata": {
|
||||
"model_info": {"id": "attacker-chosen-deployment"},
|
||||
"deployment_model_name": "gpt-4",
|
||||
}
|
||||
}
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown:
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=mock_router,
|
||||
kwargs=kwargs,
|
||||
exception=exc,
|
||||
)
|
||||
|
||||
mock_set_cooldown.assert_not_called()
|
||||
|
||||
def test_trigger_cooldown_increments_failure_counter_before_cooldown_check(self):
|
||||
"""
|
||||
The fallback path must feed the same per-minute failure counter the
|
||||
primary path uses, or repeated fallback failures never accumulate toward
|
||||
the default percent-fail-rate cooldown threshold.
|
||||
"""
|
||||
mock_router = MagicMock()
|
||||
mock_router.cooldown_time = 60.0
|
||||
mock_router.get_model_info.return_value = None
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
exc.failed_deployment_id = "fallback-deployment"
|
||||
|
||||
with (
|
||||
patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown,
|
||||
patch(
|
||||
"litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute"
|
||||
) as mock_increment,
|
||||
):
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc)
|
||||
|
||||
mock_increment.assert_called_once_with(
|
||||
litellm_router_instance=mock_router, deployment_id="fallback-deployment"
|
||||
)
|
||||
mock_set_cooldown.assert_called_once()
|
||||
|
||||
def test_trigger_cooldown_uses_deployment_cooldown_time_override(self):
|
||||
"""
|
||||
When the deployment has a model_info.cooldown_time, that value must be
|
||||
passed as time_to_cooldown rather than the router-level cooldown_time.
|
||||
"""
|
||||
mock_router = MagicMock()
|
||||
mock_router.cooldown_time = 300.0
|
||||
mock_router.get_model_info.return_value = {"model_info": {"cooldown_time": 30.0}}
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
exc.failed_deployment_id = "fallback-deployment"
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown:
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=mock_router,
|
||||
kwargs={},
|
||||
exception=exc,
|
||||
)
|
||||
|
||||
call_kwargs = mock_set_cooldown.call_args[1]
|
||||
assert call_kwargs["time_to_cooldown"] == 30.0, (
|
||||
"Deployment-level cooldown_time must override router-level value"
|
||||
)
|
||||
|
||||
def test_trigger_cooldown_skipped_for_advisor_orchestration_failure(self):
|
||||
"""
|
||||
A failure tagged as originating from advisor orchestration (not the selected
|
||||
deployment) must not cool down the fallback deployment, matching the same
|
||||
guard already applied in Router.deployment_callback_on_failure.
|
||||
"""
|
||||
mock_router = MagicMock()
|
||||
mock_router.cooldown_time = 60.0
|
||||
mock_router.get_model_info.return_value = None
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
exc.failed_deployment_id = "fallback-deployment"
|
||||
mark_advisor_orchestration_failure(exc)
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown:
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=mock_router,
|
||||
kwargs={},
|
||||
exception=exc,
|
||||
)
|
||||
|
||||
mock_set_cooldown.assert_not_called()
|
||||
|
||||
def test_trigger_cooldown_falls_back_to_litellm_params_cooldown_time(self):
|
||||
"""
|
||||
cooldown_time has pre-existing litellm_params support on the primary
|
||||
failure path (Router.deployment_callback_on_failure), so it must still be
|
||||
honored as a fallback when model_info doesn't set it, unlike the new
|
||||
allowed_fails/allowed_fails_policy fields which are model_info-only.
|
||||
"""
|
||||
mock_router = MagicMock()
|
||||
mock_router.cooldown_time = 300.0
|
||||
mock_router.get_model_info.return_value = {"litellm_params": {"cooldown_time": 30.0}}
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
exc.failed_deployment_id = "fallback-deployment"
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown:
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=mock_router,
|
||||
kwargs={},
|
||||
exception=exc,
|
||||
)
|
||||
|
||||
call_kwargs = mock_set_cooldown.call_args[1]
|
||||
assert call_kwargs["time_to_cooldown"] == 30.0, (
|
||||
"litellm_params.cooldown_time must still be honored as a fallback"
|
||||
)
|
||||
|
||||
def test_trigger_cooldown_prefers_model_info_cooldown_time_over_litellm_params(self):
|
||||
mock_router = MagicMock()
|
||||
mock_router.cooldown_time = 300.0
|
||||
mock_router.get_model_info.return_value = {
|
||||
"model_info": {"cooldown_time": 15.0},
|
||||
"litellm_params": {"cooldown_time": 30.0},
|
||||
}
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
exc.failed_deployment_id = "fallback-deployment"
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown:
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=mock_router,
|
||||
kwargs={},
|
||||
exception=exc,
|
||||
)
|
||||
|
||||
call_kwargs = mock_set_cooldown.call_args[1]
|
||||
assert call_kwargs["time_to_cooldown"] == 15.0, "model_info.cooldown_time must take priority"
|
||||
|
||||
|
||||
class TestSingleDeploymentModelGroupProtection:
|
||||
def test_generic_allowed_fails_does_not_bypass_single_deployment_protection(self):
|
||||
"""
|
||||
Setting only a generic model_info.allowed_fails on a single-deployment model
|
||||
group must not disable the "avoid cooldowns on single deployment model groups"
|
||||
safety net; before this feature existed the field had no effect at all here,
|
||||
so a plain 500 error must behave the same as the no-policy control.
|
||||
"""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {"id": "solo", "allowed_fails": 1},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
exc = Exception("Internal error")
|
||||
for _ in range(2):
|
||||
should_cooldown = _should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="solo",
|
||||
exception_status=500,
|
||||
original_exception=exc,
|
||||
)
|
||||
assert should_cooldown is False, (
|
||||
"single-deployment model group must stay protected from a generic allowed_fails override"
|
||||
)
|
||||
|
||||
def test_named_exception_policy_still_overrides_single_deployment_protection(self):
|
||||
"""
|
||||
Unlike a generic allowed_fails, an explicit per-exception-type allowed_fails_policy
|
||||
entry is a deliberate, unambiguous opt-in and must still apply even on a
|
||||
single-deployment model group.
|
||||
"""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {
|
||||
"id": "solo",
|
||||
"allowed_fails_policy": {"RateLimitErrorAllowedFails": 0},
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
should_cooldown = _should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="solo",
|
||||
exception_status=429,
|
||||
original_exception=exc,
|
||||
)
|
||||
assert should_cooldown is True, "explicit per-exception-type policy must still cool down a solo deployment"
|
||||
|
||||
|
||||
class TestShouldCooldownBasedOnAllowedFailsPolicyFalsyZero:
|
||||
def test_router_level_policy_of_zero_is_not_swallowed_by_allowed_fails(self):
|
||||
"""
|
||||
Router.get_allowed_fails_from_policy returning 0 (a legitimate "cooldown after
|
||||
the very first failure" policy) must not be treated as falsy and replaced by
|
||||
router.allowed_fails.
|
||||
"""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {"id": "primary"},
|
||||
},
|
||||
],
|
||||
allowed_fails=10,
|
||||
allowed_fails_policy=AllowedFailsPolicy(RateLimitErrorAllowedFails=0),
|
||||
)
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
should_cooldown = should_cooldown_based_on_allowed_fails_policy(
|
||||
litellm_router_instance=router,
|
||||
deployment="primary",
|
||||
original_exception=exc,
|
||||
)
|
||||
assert should_cooldown is True, "RateLimitErrorAllowedFails=0 must cool down after the first failure"
|
||||
|
||||
|
||||
class TestResolveAllowedFailsFromPolicyFallsThrough:
|
||||
def test_none_value_on_first_match_falls_through_to_next_type(self):
|
||||
"""
|
||||
ContentPolicyViolationError is also a BadRequestError; if the policy names
|
||||
ContentPolicyViolationError but leaves its value unset (None) while setting
|
||||
BadRequestErrorAllowedFails, resolution must fall through to the
|
||||
BadRequestError entry rather than stopping at the first isinstance match.
|
||||
"""
|
||||
policy = {
|
||||
"ContentPolicyViolationErrorAllowedFails": None,
|
||||
"BadRequestErrorAllowedFails": 3,
|
||||
}
|
||||
exc = litellm.ContentPolicyViolationError("flagged", "openai", "gpt-4")
|
||||
result = _resolve_allowed_fails_from_policy(policy=policy, exception=exc)
|
||||
assert result == 3, "must fall through to BadRequestErrorAllowedFails when the more specific field is unset"
|
||||
|
||||
|
||||
class TestDeploymentCallbackOnFailureCooldownTimePrecedence:
|
||||
def test_model_info_cooldown_time_used_in_primary_sync_path(self):
|
||||
"""
|
||||
Router.deployment_callback_on_failure (the primary sync failure-callback path,
|
||||
as opposed to the fallback path covered by TestFallbackDeploymentCooldown) must
|
||||
also honor a model_info.cooldown_time, not just litellm_params.cooldown_time.
|
||||
"""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4"},
|
||||
"model_info": {"id": "primary", "cooldown_time": 15.0},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
kwargs = {
|
||||
"exception": exc,
|
||||
"litellm_params": {
|
||||
"model_info": {"id": "primary", "cooldown_time": 15.0},
|
||||
},
|
||||
}
|
||||
|
||||
with patch("litellm.router._set_cooldown_deployments") as mock_set_cooldown:
|
||||
router.deployment_callback_on_failure(
|
||||
kwargs=kwargs,
|
||||
completion_response=None,
|
||||
start_time=0,
|
||||
end_time=1,
|
||||
)
|
||||
|
||||
mock_set_cooldown.assert_called_once()
|
||||
call_kwargs = mock_set_cooldown.call_args[1]
|
||||
assert call_kwargs["time_to_cooldown"] == 15.0, (
|
||||
"model_info.cooldown_time must be honored in the primary sync failure-callback path"
|
||||
)
|
||||
|
||||
def test_litellm_params_cooldown_time_still_honored_as_fallback(self):
|
||||
"""cooldown_time has pre-existing litellm_params support on this primary
|
||||
path; it must keep working when model_info doesn't set it."""
|
||||
router = _make_router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4", "cooldown_time": 20.0},
|
||||
"model_info": {"id": "primary"},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4")
|
||||
kwargs = {
|
||||
"exception": exc,
|
||||
"litellm_params": {
|
||||
"model_info": {"id": "primary"},
|
||||
"cooldown_time": 20.0,
|
||||
},
|
||||
}
|
||||
|
||||
with patch("litellm.router._set_cooldown_deployments") as mock_set_cooldown:
|
||||
router.deployment_callback_on_failure(
|
||||
kwargs=kwargs,
|
||||
completion_response=None,
|
||||
start_time=0,
|
||||
end_time=1,
|
||||
)
|
||||
|
||||
call_kwargs = mock_set_cooldown.call_args[1]
|
||||
assert call_kwargs["time_to_cooldown"] == 20.0, "litellm_params.cooldown_time must still be honored"
|
||||
|
||||
|
||||
class TestNewAllowedFailsPolicyFields:
|
||||
def test_service_unavailable_error_matched_by_policy(self):
|
||||
"""
|
||||
ServiceUnavailableError must be matched against ServiceUnavailableErrorAllowedFails.
|
||||
"""
|
||||
policy = {"ServiceUnavailableErrorAllowedFails": 0}
|
||||
exc = litellm.ServiceUnavailableError("Service unavailable", "openai", "gpt-4")
|
||||
result = _resolve_allowed_fails_from_policy(policy=policy, exception=exc)
|
||||
assert result == 0
|
||||
|
||||
def test_bad_gateway_error_matched_by_policy(self):
|
||||
"""
|
||||
BadGatewayError must be matched against BadGatewayErrorAllowedFails.
|
||||
"""
|
||||
policy = {"BadGatewayErrorAllowedFails": 2}
|
||||
exc = litellm.BadGatewayError("Bad gateway", "openai", "gpt-4")
|
||||
result = _resolve_allowed_fails_from_policy(policy=policy, exception=exc)
|
||||
assert result == 2
|
||||
|
||||
def test_not_found_error_matched_by_policy(self):
|
||||
"""
|
||||
NotFoundError must be matched against NotFoundErrorAllowedFails.
|
||||
"""
|
||||
policy = {"NotFoundErrorAllowedFails": 1}
|
||||
exc = litellm.NotFoundError("Not found", "openai", "gpt-4")
|
||||
result = _resolve_allowed_fails_from_policy(policy=policy, exception=exc)
|
||||
assert result == 1
|
||||
|
||||
def test_unknown_exception_type_returns_none(self):
|
||||
"""
|
||||
An exception type not in the policy mapping must return None.
|
||||
"""
|
||||
policy = {"RateLimitErrorAllowedFails": 0}
|
||||
exc = ValueError("unexpected error")
|
||||
result = _resolve_allowed_fails_from_policy(policy=policy, exception=exc)
|
||||
assert result is None
|
||||
|
||||
def test_allowed_fails_policy_model_accepts_new_fields(self):
|
||||
"""
|
||||
AllowedFailsPolicy Pydantic model must accept the three new fields.
|
||||
"""
|
||||
policy = AllowedFailsPolicy(
|
||||
ServiceUnavailableErrorAllowedFails=3,
|
||||
BadGatewayErrorAllowedFails=2,
|
||||
NotFoundErrorAllowedFails=1,
|
||||
)
|
||||
assert policy.ServiceUnavailableErrorAllowedFails == 3
|
||||
assert policy.BadGatewayErrorAllowedFails == 2
|
||||
assert policy.NotFoundErrorAllowedFails == 1
|
||||
|
||||
|
||||
class TestRouterLevelGetAllowedFailsFromPolicy:
|
||||
"""Router.get_allowed_fails_from_policy must handle all AllowedFailsPolicy fields."""
|
||||
|
||||
def _make_router(self, **policy_kwargs):
|
||||
return Router(
|
||||
model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "fake"}}],
|
||||
allowed_fails_policy=AllowedFailsPolicy(**policy_kwargs),
|
||||
)
|
||||
|
||||
def test_internal_server_error_returned(self):
|
||||
router = self._make_router(InternalServerErrorAllowedFails=7)
|
||||
exc = litellm.InternalServerError("500 error", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) == 7
|
||||
|
||||
def test_service_unavailable_error_returned(self):
|
||||
router = self._make_router(ServiceUnavailableErrorAllowedFails=4)
|
||||
exc = litellm.ServiceUnavailableError("503 error", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) == 4
|
||||
|
||||
def test_bad_gateway_error_returned(self):
|
||||
router = self._make_router(BadGatewayErrorAllowedFails=2)
|
||||
exc = litellm.BadGatewayError("502 error", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) == 2
|
||||
|
||||
def test_not_found_error_returned(self):
|
||||
router = self._make_router(NotFoundErrorAllowedFails=1)
|
||||
exc = litellm.NotFoundError("404 error", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) == 1
|
||||
|
||||
def test_unmatched_exception_returns_none(self):
|
||||
router = self._make_router(InternalServerErrorAllowedFails=5)
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) is None
|
||||
|
|
@ -19,7 +19,9 @@ from litellm.router_utils.cooldown_handlers import (
|
|||
_should_cooldown_deployment,
|
||||
cast_exception_status_to_int,
|
||||
_is_cooldown_required,
|
||||
_has_explicit_allowed_fails_policy_for_exception,
|
||||
)
|
||||
from litellm.types.router import AllowedFailsPolicy
|
||||
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
|
||||
increment_deployment_failures_for_current_minute,
|
||||
increment_deployment_successes_for_current_minute,
|
||||
|
|
@ -107,6 +109,137 @@ def test_should_run_cooldown_logic(testing_litellm_router):
|
|||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def single_deployment_router():
|
||||
"""A router with one deployment whose model_info.id is the lookup-able
|
||||
"dep-1" (unlike `testing_litellm_router`'s top-level "model_id" key, which
|
||||
is not absorbed into model_info.id and so never resolves via
|
||||
get_model_info/get_model_group)."""
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-5-mini",
|
||||
"litellm_params": {"model": "gpt-5-mini"},
|
||||
"model_info": {"id": "dep-1"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_should_run_cooldown_logic_generic_bad_request_excluded_by_default(
|
||||
single_deployment_router,
|
||||
):
|
||||
"""A generic BadRequestError/ContentPolicyViolationError (400) is excluded from
|
||||
cooldown evaluation by _is_cooldown_required when no allowed_fails_policy is
|
||||
configured for that exception type. This is the pre-existing, intentional
|
||||
default: a client error is usually not the deployment's fault."""
|
||||
exc = litellm.BadRequestError("bad request", "openai", "gpt-5-mini")
|
||||
assert (
|
||||
_should_run_cooldown_logic(single_deployment_router, "dep-1", 400, exc) is False
|
||||
)
|
||||
|
||||
|
||||
def test_should_run_cooldown_logic_router_level_policy_does_not_override_bad_request_exclusion(
|
||||
single_deployment_router,
|
||||
):
|
||||
"""A router-level allowed_fails_policy is a pre-existing, router-wide setting that
|
||||
predates the per-deployment override feature, so it must keep its existing behavior
|
||||
and stay subject to the generic 4XX exclusion. Only an explicit deployment-level
|
||||
policy (an unambiguous per-exception opt-in for that one deployment) overrides it;
|
||||
see test_should_run_cooldown_logic_explicit_deployment_level_policy_overrides_content_policy_exclusion."""
|
||||
exc = litellm.BadRequestError("bad request", "openai", "gpt-5-mini")
|
||||
single_deployment_router.allowed_fails_policy = AllowedFailsPolicy(
|
||||
BadRequestErrorAllowedFails=5
|
||||
)
|
||||
assert (
|
||||
_should_run_cooldown_logic(single_deployment_router, "dep-1", 400, exc) is False
|
||||
)
|
||||
|
||||
|
||||
def test_should_run_cooldown_logic_explicit_deployment_level_policy_overrides_content_policy_exclusion(
|
||||
single_deployment_router,
|
||||
):
|
||||
"""Same as the router-level case, but for a deployment-level allowed_fails_policy
|
||||
entry (this PR's per-deployment feature) targeting ContentPolicyViolationError."""
|
||||
exc = litellm.ContentPolicyViolationError("flagged content", "openai", "gpt-5-mini")
|
||||
deployment_dict = single_deployment_router.get_model_info(id="dep-1")
|
||||
deployment_dict["model_info"]["allowed_fails_policy"] = {
|
||||
"ContentPolicyViolationErrorAllowedFails": 0
|
||||
}
|
||||
assert (
|
||||
_should_run_cooldown_logic(single_deployment_router, "dep-1", 400, exc) is True
|
||||
)
|
||||
|
||||
|
||||
class TestHasExplicitAllowedFailsPolicyForException:
|
||||
def test_no_policy_anywhere_returns_false(self, single_deployment_router):
|
||||
exc = litellm.BadRequestError("bad request", "openai", "gpt-5-mini")
|
||||
assert (
|
||||
_has_explicit_allowed_fails_policy_for_exception(
|
||||
single_deployment_router, "dep-1", exc
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_router_level_policy_for_matching_exception_returns_false(
|
||||
self, single_deployment_router
|
||||
):
|
||||
"""Deliberately scoped to deployment-level only: a router-level policy
|
||||
predates this feature and must not be treated as an explicit per-exception
|
||||
opt-in for cooldown-gate purposes."""
|
||||
exc = litellm.RateLimitError("rate limited", "openai", "gpt-5-mini")
|
||||
single_deployment_router.allowed_fails_policy = AllowedFailsPolicy(
|
||||
RateLimitErrorAllowedFails=3
|
||||
)
|
||||
assert (
|
||||
_has_explicit_allowed_fails_policy_for_exception(
|
||||
single_deployment_router, "dep-1", exc
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_router_level_policy_for_different_exception_returns_false(
|
||||
self, single_deployment_router
|
||||
):
|
||||
exc = litellm.BadRequestError("bad request", "openai", "gpt-5-mini")
|
||||
single_deployment_router.allowed_fails_policy = AllowedFailsPolicy(
|
||||
RateLimitErrorAllowedFails=3
|
||||
)
|
||||
assert (
|
||||
_has_explicit_allowed_fails_policy_for_exception(
|
||||
single_deployment_router, "dep-1", exc
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_deployment_level_policy_for_matching_exception_returns_true(
|
||||
self, single_deployment_router
|
||||
):
|
||||
exc = litellm.ContentPolicyViolationError("flagged", "openai", "gpt-5-mini")
|
||||
deployment_dict = single_deployment_router.get_model_info(id="dep-1")
|
||||
deployment_dict["model_info"]["allowed_fails_policy"] = {
|
||||
"ContentPolicyViolationErrorAllowedFails": 0
|
||||
}
|
||||
assert (
|
||||
_has_explicit_allowed_fails_policy_for_exception(
|
||||
single_deployment_router, "dep-1", exc
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_none_deployment_returns_false(self, single_deployment_router):
|
||||
exc = litellm.RateLimitError("rate limited", "openai", "gpt-5-mini")
|
||||
single_deployment_router.allowed_fails_policy = AllowedFailsPolicy(
|
||||
RateLimitErrorAllowedFails=3
|
||||
)
|
||||
assert (
|
||||
_has_explicit_allowed_fails_policy_for_exception(
|
||||
single_deployment_router, None, exc
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
def test_should_cooldown_deployment_rate_limit_error(testing_litellm_router):
|
||||
"""
|
||||
Test the _should_cooldown_deployment function when a rate limit error occurs
|
||||
|
|
|
|||
|
|
@ -2143,6 +2143,54 @@ def test_token_type_cost_breakdown_matches_real_gemini_numbers():
|
|||
assert breakdown.cache_creation_cost == 0.0
|
||||
|
||||
|
||||
def test_token_type_cost_breakdown_xai_at_exactly_200k_uses_higher_tier_rates():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=200_000,
|
||||
completion_tokens=2_000,
|
||||
total_tokens=202_000,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
reasoning_tokens=1_500, text_tokens=500
|
||||
),
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
cached_tokens=50_000, text_tokens=150_000
|
||||
),
|
||||
)
|
||||
|
||||
breakdown = get_token_type_cost_breakdown(
|
||||
model="grok-4.20-0309-reasoning", custom_llm_provider="xai", usage=usage
|
||||
)
|
||||
|
||||
assert breakdown.reasoning_cost == pytest.approx(1_500 * 5e-06)
|
||||
assert breakdown.cache_read_cost == pytest.approx(50_000 * 4e-07)
|
||||
|
||||
|
||||
def test_token_type_cost_breakdown_xai_just_below_200k_uses_base_tier_rates():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=199_999,
|
||||
completion_tokens=2_000,
|
||||
total_tokens=201_999,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
reasoning_tokens=1_500, text_tokens=500
|
||||
),
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
cached_tokens=50_000, text_tokens=149_999
|
||||
),
|
||||
)
|
||||
|
||||
breakdown = get_token_type_cost_breakdown(
|
||||
model="grok-4.20-0309-reasoning", custom_llm_provider="xai", usage=usage
|
||||
)
|
||||
|
||||
assert breakdown.reasoning_cost == pytest.approx(1_500 * 2.5e-06)
|
||||
assert breakdown.cache_read_cost == pytest.approx(50_000 * 2e-07)
|
||||
|
||||
|
||||
def test_token_type_cost_breakdown_includes_cache_creation_from_top_level_usage():
|
||||
"""
|
||||
Bedrock/Anthropic report cache tokens as top-level usage fields; the Usage
|
||||
|
|
@ -2446,6 +2494,60 @@ def test_token_type_cost_breakdown_applies_regional_uplift():
|
|||
assert text_input_cost + eu.cache_read_cost == pytest.approx(prompt_cost)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("details_as_dict", [True, False])
|
||||
def test_image_response_input_image_tokens_priced_at_image_rate(details_as_dict):
|
||||
"""
|
||||
Image input tokens must be priced at input_cost_per_image_token even when
|
||||
input_tokens_details is a plain dict, as in OpenAI image edit responses.
|
||||
|
||||
Regression test: dict-shaped input_tokens_details was read with getattr(),
|
||||
which returns None for dicts, so image input tokens silently fell back to
|
||||
the text input rate (e.g. $5/M instead of $8/M for gpt-image-2).
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
||||
calculate_image_response_cost_from_usage,
|
||||
)
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
mock_model_info = {
|
||||
"input_cost_per_token": 5e-6,
|
||||
"input_cost_per_image_token": 8e-6,
|
||||
"output_cost_per_image_token": 3e-5,
|
||||
}
|
||||
|
||||
input_details = {"text_tokens": 19, "image_tokens": 512}
|
||||
image_response = ImageResponse(data=[ImageObject(b64_json="x")])
|
||||
# Mirror the usage shape of a real OpenAI images.edit response:
|
||||
# a Usage object carrying input_tokens/output_tokens with detail dicts.
|
||||
image_response.usage = Usage(
|
||||
prompt_tokens=0,
|
||||
completion_tokens=0,
|
||||
total_tokens=689,
|
||||
input_tokens=531,
|
||||
input_tokens_details=(
|
||||
input_details
|
||||
if details_as_dict
|
||||
else ImageUsageInputTokensDetails(**input_details)
|
||||
),
|
||||
output_tokens=158,
|
||||
output_tokens_details={"image_tokens": 158, "text_tokens": 0},
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.llm_cost_calc.utils.get_model_info",
|
||||
return_value=mock_model_info,
|
||||
):
|
||||
cost = calculate_image_response_cost_from_usage(
|
||||
model="gpt-image-2",
|
||||
image_response=image_response,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
expected = 19 * 5e-6 + 512 * 8e-6 + 158 * 3e-5
|
||||
assert cost is not None
|
||||
assert round(cost, 12) == round(expected, 12)
|
||||
GEMINI_DAY0_LAUNCH_PRICING = [
|
||||
("gemini-3.6-flash", 1.5e-06, 7.5e-06, 1.5e-07),
|
||||
("gemini/gemini-3.6-flash", 1.5e-06, 7.5e-06, 1.5e-07),
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import httpx
|
|||
from openai._legacy_response import HttpxBinaryResponseContent
|
||||
|
||||
import litellm
|
||||
from litellm._logging import session_id_var, trace_id_var
|
||||
from litellm.constants import SENTRY_DENYLIST, SENTRY_PII_DENYLIST
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
|
||||
|
|
@ -3312,6 +3313,51 @@ def test_failure_handler_runs_sync_callbacks_for_non_pass_through_requests(
|
|||
dummy_logger.log_failure_event.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_failure_handler_runs_callbacks_and_restores_correlation_context(logging_obj):
|
||||
"""await logging_obj.async_failure_handler(...) must dispatch async failure callbacks
|
||||
and, once its own body completes, restore trace_id/session_id contextvars via
|
||||
_restore_correlation_context() (the fix for the nested-call context leak)."""
|
||||
from litellm._logging import session_id_var, trace_id_var
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class DummyLogger(CustomLogger):
|
||||
pass
|
||||
|
||||
logging_obj.call_type = "acompletion"
|
||||
logging_obj.stream = False
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
logging_obj.litellm_params = {}
|
||||
|
||||
dummy_logger = DummyLogger()
|
||||
dummy_logger.async_log_failure_event = AsyncMock()
|
||||
|
||||
# logging_obj is constructed by the fixture (before this line runs), so it
|
||||
# already captured whatever was ambient at that point as its own pre-call
|
||||
# value - assert restoration lands back on THAT captured value, not a
|
||||
# value set here (which would be too late to affect __init__'s snapshot).
|
||||
trace_id_var.set("mutated-during-call")
|
||||
session_id_var.set("mutated-during-call")
|
||||
try:
|
||||
with patch.object(
|
||||
logging_obj,
|
||||
"get_combined_callback_list",
|
||||
return_value=[dummy_logger],
|
||||
):
|
||||
await logging_obj.async_failure_handler(
|
||||
exception=Exception("test error"),
|
||||
traceback_exception="",
|
||||
)
|
||||
|
||||
dummy_logger.async_log_failure_event.assert_called_once()
|
||||
assert trace_id_var.get() == logging_obj._pre_call_trace_id
|
||||
assert session_id_var.get() == logging_obj._pre_call_session_id
|
||||
assert trace_id_var.get() != "mutated-during-call"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_merge_hidden_params_from_response_into_metadata_populates_metadata():
|
||||
"""Streaming completion path should mirror non-stream: metadata.hidden_params from response."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -4230,3 +4276,199 @@ def test_pre_call_does_not_pin_request_in_module_state(logging_obj):
|
|||
logging_obj.post_call(original_response='{"ok": true}', input=big_input, api_key="sk-test")
|
||||
|
||||
assert litellm.error_logs == {}
|
||||
|
||||
|
||||
def test_logging_init_sets_trace_id():
|
||||
"""Logging.__init__() must call set_trace_id with self.litellm_trace_id."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
trace_id_var.set("")
|
||||
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="call-001",
|
||||
function_id="fn-001",
|
||||
kwargs={},
|
||||
)
|
||||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
|
||||
|
||||
def test_logging_init_skips_stamping_when_correlation_logging_unsupported():
|
||||
"""supports_correlation_logging=False (what wrapper(), the sync entry
|
||||
point, always passes) must leave trace_id_var/session_id_var completely
|
||||
untouched, even though self.litellm_trace_id/litellm_session_id (the
|
||||
plain attributes used by StandardLoggingPayload) are still populated as
|
||||
usual - only the ambient contextvar stamping is gated."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="call-sync-excluded",
|
||||
function_id="fn-sync-excluded",
|
||||
kwargs={"litellm_session_id": "should-not-be-stamped"},
|
||||
litellm_trace_id="should-not-be-stamped-either",
|
||||
supports_correlation_logging=False,
|
||||
)
|
||||
|
||||
assert trace_id_var.get() == ""
|
||||
assert session_id_var.get() == ""
|
||||
# The plain attributes are unaffected - only the contextvar stamping is gated.
|
||||
assert log_obj.litellm_trace_id == "should-not-be-stamped-either"
|
||||
assert log_obj.litellm_session_id == "should-not-be-stamped"
|
||||
|
||||
|
||||
def test_logging_init_sets_session_id_when_provided():
|
||||
"""Logging.__init__() must call set_session_id when litellm_session_id is in kwargs."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
session_id_var.set("")
|
||||
|
||||
Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="call-002",
|
||||
function_id="fn-002",
|
||||
kwargs={"litellm_session_id": "my-session-99"},
|
||||
)
|
||||
assert session_id_var.get() == "my-session-99"
|
||||
|
||||
|
||||
def test_logging_init_resets_session_id_to_empty_when_absent():
|
||||
"""When no session_id is in kwargs, Logging.__init__() must reset session_id_var to ""
|
||||
so a prior request's session_id does not leak into subsequent log records."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
session_id_var.set("preexisting-sid")
|
||||
|
||||
Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="call-003",
|
||||
function_id="fn-003",
|
||||
kwargs={},
|
||||
)
|
||||
assert session_id_var.get() == ""
|
||||
|
||||
|
||||
def test_restore_correlation_context_resets_to_pre_call_value():
|
||||
"""_restore_correlation_context() must put trace_id_var/session_id_var back to
|
||||
whatever they were immediately before this Logging instance was constructed.
|
||||
This is the mechanism that prevents a nested call (e.g. a guardrail's own
|
||||
LLM-as-judge call sharing the same asyncio Task) from leaking its trace_id/
|
||||
session_id into the outer call's subsequent log lines."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
trace_id_var.set("outer-trace")
|
||||
session_id_var.set("outer-session")
|
||||
try:
|
||||
inner = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="inner-call",
|
||||
function_id="fn-inner",
|
||||
kwargs={"litellm_session_id": "inner-session"},
|
||||
)
|
||||
assert trace_id_var.get() == inner.litellm_trace_id
|
||||
assert session_id_var.get() == "inner-session"
|
||||
|
||||
inner._restore_correlation_context()
|
||||
|
||||
assert trace_id_var.get() == "outer-trace"
|
||||
assert session_id_var.get() == "outer-session"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_restore_correlation_context_safe_to_call_repeatedly():
|
||||
"""Calling _restore_correlation_context() more than once must not raise.
|
||||
|
||||
It's deliberately NOT guarded against repeat calls: wrapper()'s finally
|
||||
block and a terminal handler (success_handler/failure_handler) can both
|
||||
end up calling it for the same instance, potentially from different
|
||||
asyncio Tasks - each call needs to take effect in its own Task's view of
|
||||
the contextvars, so repeat calls are expected, not just tolerated."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="call-idempotent",
|
||||
function_id="fn-idempotent",
|
||||
kwargs={},
|
||||
)
|
||||
log_obj._restore_correlation_context()
|
||||
log_obj._restore_correlation_context() # must not raise
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_restore_correlation_context_works_across_asyncio_task_boundary():
|
||||
"""_restore_correlation_context() must succeed even when it's called from a
|
||||
different asyncio Task than the one Logging.__init__() ran in - exactly what
|
||||
happens on litellm's real async success path, where async_success_handler is
|
||||
dispatched via asyncio.create_task / the global logging worker rather than
|
||||
awaited directly in the request's own task.
|
||||
|
||||
A contextvars.Token can only be reset in the exact Context it was created in
|
||||
and raises ValueError otherwise (verified separately against raw contextvars,
|
||||
not just this codebase). The fix uses a plain set() of the captured pre-call
|
||||
value instead, which works regardless of which Task calls it. This test
|
||||
fails with a token-based implementation - the child task's reset() would
|
||||
raise, get silently swallowed, and leave the child's view unrestored - and
|
||||
passes with the value-based one.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
trace_id_var.set("outer-trace-cross-task")
|
||||
session_id_var.set("outer-session-cross-task")
|
||||
try:
|
||||
# __init__ runs in THIS (outer) task's context.
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="acompletion",
|
||||
start_time=None,
|
||||
litellm_call_id="cross-task-call",
|
||||
function_id="fn-cross-task",
|
||||
kwargs={"litellm_session_id": "cross-task-session"},
|
||||
)
|
||||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "cross-task-session"
|
||||
|
||||
async def restore_in_new_task():
|
||||
# Simulates async_success_handler running in a task spawned after
|
||||
# __init__ already ran elsewhere - a different Context object.
|
||||
log_obj._restore_correlation_context()
|
||||
return trace_id_var.get(), session_id_var.get()
|
||||
|
||||
trace_in_child, session_in_child = await asyncio.create_task(restore_in_new_task())
|
||||
|
||||
assert trace_in_child == "outer-trace-cross-task"
|
||||
assert session_in_child == "outer-session-cross-task"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ Covers:
|
|||
import os
|
||||
import sys
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -69,8 +69,12 @@ class TestCustomStreamWrapperMaxDuration:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_raise_on_async_anext_when_exceeded(self):
|
||||
"""__anext__ should check the limit before iterating."""
|
||||
"""__anext__ should check the limit before iterating, dispatching the
|
||||
same failure-callback/logging path every other stream failure goes
|
||||
through (dispatch_failure_handlers is async on the real Logging class,
|
||||
so the mock needs to be awaitable too)."""
|
||||
wrapper = _make_custom_stream_wrapper()
|
||||
wrapper.logging_obj.dispatch_failure_handlers = AsyncMock()
|
||||
wrapper._stream_created_time = time.time() - 20
|
||||
with patch("litellm.constants.LITELLM_MAX_STREAMING_DURATION_SECONDS", 10.0):
|
||||
with pytest.raises(litellm.Timeout):
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ import traceback
|
|||
from typing import Optional
|
||||
|
||||
import litellm
|
||||
from litellm import verbose_logger
|
||||
from litellm._logging import session_id_var, trace_id_var
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.streaming_handler import (
|
||||
AUDIO_ATTRIBUTE,
|
||||
|
|
@ -3551,3 +3553,613 @@ def test_openai_custom_tool_call_stream_deltas_survive_conversion(logging_obj: L
|
|||
assert combined_input == "*** Begin Patch\n*** End Patch\n"
|
||||
finish_reasons = [chunk.choices[0].finish_reason for chunk in emitted if chunk.choices]
|
||||
assert "tool_calls" in finish_reasons
|
||||
|
||||
|
||||
def test_sync_completion_never_stamps_correlation_context(monkeypatch):
|
||||
"""wrapper() (the sync entry point) does not participate in
|
||||
request_correlation_in_logs at all: Logging.__init__() is called with
|
||||
supports_correlation_logging=False for every sync call, so
|
||||
trace_id_var/session_id_var are never touched, regardless of whether the
|
||||
caller passes litellm_trace_id/litellm_session_id or the call streams.
|
||||
|
||||
This is a deliberate scoping decision, not an oversight: a plain OS
|
||||
thread has no per-call isolation the way an asyncio Task does, and a
|
||||
thread pool's worker threads are recycled across unrelated requests, so
|
||||
safely supporting this for the sync path needs its own restore mechanism
|
||||
with its own tests - tracked as a separate, follow-up piece of work.
|
||||
Async (acompletion/wrapper_async, the only path the proxy uses) is
|
||||
unaffected - see test_async_streaming_completion_does_not_reset_context_before_iteration."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
# Reset explicitly rather than asserting a clean slate - this must hold
|
||||
# regardless of what any other test left behind in these module-level
|
||||
# contextvars.
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
try:
|
||||
litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="Hello there!",
|
||||
litellm_trace_id="should-never-appear",
|
||||
litellm_session_id="should-never-appear-either",
|
||||
num_retries=0,
|
||||
)
|
||||
assert trace_id_var.get() == ""
|
||||
assert session_id_var.get() == ""
|
||||
|
||||
response = litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="Hello there!",
|
||||
stream=True,
|
||||
litellm_trace_id="should-never-appear-stream",
|
||||
litellm_session_id="should-never-appear-stream-either",
|
||||
num_retries=0,
|
||||
)
|
||||
for _ in response:
|
||||
pass
|
||||
assert trace_id_var.get() == ""
|
||||
assert session_id_var.get() == ""
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_abandoned_sync_stream_cannot_contaminate_a_later_call_on_the_same_thread(monkeypatch):
|
||||
"""The maintainer-reported blocking bug reproduced live in this session -
|
||||
request A starts a sync stream, consumes one chunk, abandons it; request
|
||||
B runs next on the same forced-reuse ThreadPoolExecutor worker - is now
|
||||
structurally impossible rather than merely restored-after-the-fact: since
|
||||
sync calls never stamp trace_id_var/session_id_var at all
|
||||
(supports_correlation_logging=False), there is nothing for request A to
|
||||
leave behind for request B to inherit."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
pool = ThreadPoolExecutor(max_workers=1)
|
||||
try:
|
||||
|
||||
def call_a_abandon_stream():
|
||||
response = litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "call A"}],
|
||||
mock_response="call A response",
|
||||
stream=True,
|
||||
litellm_session_id="SESSION-AAA",
|
||||
litellm_trace_id="TRACE-AAA",
|
||||
num_retries=0,
|
||||
)
|
||||
next(response) # consume exactly one chunk, then abandon it
|
||||
|
||||
def call_b_non_streaming():
|
||||
litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "call B"}],
|
||||
mock_response="call B response",
|
||||
litellm_session_id="SESSION-BBB",
|
||||
litellm_trace_id="TRACE-BBB",
|
||||
num_retries=0,
|
||||
)
|
||||
return trace_id_var.get(), session_id_var.get()
|
||||
|
||||
pool.submit(call_a_abandon_stream).result()
|
||||
ids_after_b = pool.submit(call_b_non_streaming).result()
|
||||
|
||||
assert ids_after_b == ("", "")
|
||||
finally:
|
||||
pool.shutdown(wait=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_streaming_completion_does_not_reset_context_before_iteration(monkeypatch):
|
||||
"""Same as above for wrapper_async()/acompletion()."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
trace_id_var.set("outer-trace-async-stream")
|
||||
session_id_var.set("outer-session-async-stream")
|
||||
try:
|
||||
response = await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="Hello there!",
|
||||
stream=True,
|
||||
litellm_session_id="async-streaming-call-session",
|
||||
num_retries=0,
|
||||
)
|
||||
assert session_id_var.get() == "async-streaming-call-session"
|
||||
|
||||
async for _ in response:
|
||||
pass
|
||||
|
||||
# Once the stream is genuinely exhausted, the *consuming* task's own
|
||||
# context must be restored - async_success_handler's own dispatch (via
|
||||
# asyncio.create_task) only fixes up its own detached task, not this one.
|
||||
assert session_id_var.get() == "outer-session-async-stream"
|
||||
assert trace_id_var.get() == "outer-trace-async-stream"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_stream_wrapper_del_restores_correlation_context():
|
||||
"""CustomStreamWrapper.__del__ is the best-effort fallback for an abandoned
|
||||
stream (caller never exhausts it, so the normal terminal-handler restore
|
||||
never fires). Testing this via real garbage collection is unreliable in
|
||||
practice - CPython's per-chunk logging submits work to a thread pool
|
||||
executor whose worker thread transiently holds its own reference to the
|
||||
wrapper (a bound method argument) until that task completes, so refcount
|
||||
doesn't reliably hit zero on a deterministic schedule even with polling.
|
||||
Call __del__ directly instead: it's a plain method, calling it early
|
||||
doesn't run actual finalization, and this exercises exactly the logic that
|
||||
real garbage collection would eventually trigger.
|
||||
"""
|
||||
trace_id_var.set("outer-trace-abandoned")
|
||||
session_id_var.set("outer-session-abandoned")
|
||||
try:
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="abandoned-stream-call",
|
||||
function_id="fn-abandoned-stream",
|
||||
kwargs={"litellm_session_id": "abandoned-stream-session"},
|
||||
)
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=iter([]),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=log_obj,
|
||||
)
|
||||
wrapper.__del__()
|
||||
|
||||
assert trace_id_var.get() == "outer-trace-abandoned"
|
||||
assert session_id_var.get() == "outer-session-abandoned"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_stream_wrapper_del_never_raises_with_broken_logging_obj():
|
||||
"""__del__ runs during garbage collection, possibly at interpreter
|
||||
shutdown - it must never raise regardless of what's wrong with logging_obj,
|
||||
or Python prints an ignored "exception in __del__" warning and, worse,
|
||||
could mask the real error a caller is in the middle of handling."""
|
||||
|
||||
class ExplodingLogging:
|
||||
model_call_details: dict = {}
|
||||
|
||||
def _restore_correlation_context(self):
|
||||
raise RuntimeError("logging_obj is in a bad state")
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=iter([]),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=ExplodingLogging(),
|
||||
)
|
||||
wrapper.__del__() # must not raise
|
||||
|
||||
|
||||
def test_stream_wrapper_del_does_not_clobber_a_newer_active_call():
|
||||
"""A delayed finalizer must never stomp a different, still-active call's
|
||||
context. If an abandoned stream's __del__ fires late - after a new call
|
||||
has already started in the same Task/thread and claimed the contextvars -
|
||||
unconditionally restoring the abandoned stream's own pre-call snapshot
|
||||
would corrupt the active call's subsequent log lines with stale ids."""
|
||||
trace_id_var.set("outer-trace-before-abandoned-call")
|
||||
session_id_var.set("outer-session-before-abandoned-call")
|
||||
try:
|
||||
abandoned_log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="abandoned-stream-call",
|
||||
function_id="fn-abandoned-stream",
|
||||
kwargs={"litellm_session_id": "abandoned-stream-session"},
|
||||
)
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=iter([]),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=abandoned_log_obj,
|
||||
)
|
||||
|
||||
# A new, unrelated call starts in this same Task/thread before the
|
||||
# abandoned stream's __del__ ever fires, and claims the contextvars.
|
||||
Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="newer-active-call",
|
||||
function_id="fn-newer-active-call",
|
||||
kwargs={"litellm_session_id": "newer-active-session"},
|
||||
)
|
||||
assert trace_id_var.get() != abandoned_log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "newer-active-session"
|
||||
|
||||
# The delayed finalizer for the abandoned stream must not clobber
|
||||
# the newer call's still-active ids.
|
||||
wrapper.__del__()
|
||||
|
||||
assert trace_id_var.get() != abandoned_log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "newer-active-session"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_stream_wrapper_del_restores_when_own_session_id_needed_sanitizing():
|
||||
"""The __del__ guard must compare against the *sanitized* id actually
|
||||
stored in the contextvar, not the raw litellm_session_id/litellm_trace_id
|
||||
- set_session_id()/set_trace_id() strip control characters before
|
||||
storing, so a caller-supplied id containing e.g. a newline would never
|
||||
equal the raw attribute, and the guard would wrongly conclude some other
|
||||
call has claimed the context and skip cleanup forever."""
|
||||
trace_id_var.set("outer-trace-needs-sanitizing")
|
||||
session_id_var.set("outer-session-needs-sanitizing")
|
||||
try:
|
||||
raw_session_id = "abandoned\nsession\rwith-control-chars"
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="abandoned-stream-needs-sanitizing",
|
||||
function_id="fn-abandoned-stream-needs-sanitizing",
|
||||
kwargs={"litellm_session_id": raw_session_id},
|
||||
)
|
||||
# Sanity: the contextvar holds the sanitized value, which differs
|
||||
# from the raw litellm_session_id this test constructed it with.
|
||||
assert session_id_var.get() != raw_session_id
|
||||
assert log_obj.litellm_session_id == raw_session_id
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=iter([]),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=log_obj,
|
||||
)
|
||||
wrapper.__del__()
|
||||
|
||||
assert trace_id_var.get() == "outer-trace-needs-sanitizing"
|
||||
assert session_id_var.get() == "outer-session-needs-sanitizing"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_stream_wrapper_next_keeps_context_active_through_synthesized_finish_reason_chunk():
|
||||
"""When the underlying stream ends without ever emitting an explicit
|
||||
finish_reason chunk, __next__ synthesizes one via finish_reason_handler()
|
||||
and returns it. That chunk is still this call's own data - the caller's
|
||||
own (application-level) log statements processing it run immediately
|
||||
after this return, in the same synchronous frame, so context must NOT be
|
||||
restored yet or those log lines would carry the wrong ids. A caller that
|
||||
keeps iterating (the common, non-early-break pattern) still gets a
|
||||
correct, deterministic restore on the very next __next__() call, since
|
||||
completion_stream is already exhausted and immediately re-raises
|
||||
StopIteration."""
|
||||
trace_id_var.set("outer-trace-finish-reason")
|
||||
session_id_var.set("outer-session-finish-reason")
|
||||
try:
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="finish-reason-call",
|
||||
function_id="fn-finish-reason",
|
||||
kwargs={"litellm_session_id": "finish-reason-session"},
|
||||
)
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=iter([]),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=log_obj,
|
||||
)
|
||||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "finish-reason-session"
|
||||
|
||||
chunk = next(wrapper)
|
||||
|
||||
assert chunk.choices[0].finish_reason is not None
|
||||
# Still this call's own ids - not restored yet.
|
||||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "finish-reason-session"
|
||||
|
||||
# A caller that keeps iterating (doesn't break early) still gets a
|
||||
# deterministic restore right here, on the next real StopIteration.
|
||||
with pytest.raises(StopIteration):
|
||||
next(wrapper)
|
||||
assert trace_id_var.get() == "outer-trace-finish-reason"
|
||||
assert session_id_var.get() == "outer-session-finish-reason"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_stream_wrapper_del_cleans_up_after_synthesized_finish_reason_chunk():
|
||||
"""A caller that breaks immediately after seeing finish_reason (the
|
||||
early-break pattern) never triggers the next()-driven restore above - it
|
||||
relies on the best-effort __del__ guard instead, same as any other
|
||||
abandoned stream. The guard must still recognize this call's own
|
||||
(unrestored) ids as unclaimed and clean them up."""
|
||||
trace_id_var.set("outer-trace-finish-reason-del")
|
||||
session_id_var.set("outer-session-finish-reason-del")
|
||||
try:
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="finish-reason-del-call",
|
||||
function_id="fn-finish-reason-del",
|
||||
kwargs={"litellm_session_id": "finish-reason-del-session"},
|
||||
)
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=iter([]),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=log_obj,
|
||||
)
|
||||
|
||||
chunk = next(wrapper)
|
||||
assert chunk.choices[0].finish_reason is not None
|
||||
|
||||
wrapper.__del__()
|
||||
|
||||
assert trace_id_var.get() == "outer-trace-finish-reason-del"
|
||||
assert session_id_var.get() == "outer-session-finish-reason-del"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_wrapper_anext_keeps_context_active_through_synthesized_finish_reason_chunk():
|
||||
"""Async sibling of test_stream_wrapper_next_keeps_context_active_through_synthesized_finish_reason_chunk -
|
||||
_finalize_completed_stream()'s else branch must not restore before
|
||||
returning the synthesized chunk either."""
|
||||
trace_id_var.set("outer-trace-anext-finish-reason")
|
||||
session_id_var.set("outer-session-anext-finish-reason")
|
||||
try:
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="anext-finish-reason-call",
|
||||
function_id="fn-anext-finish-reason",
|
||||
kwargs={"litellm_session_id": "anext-finish-reason-session"},
|
||||
)
|
||||
|
||||
async def _empty_aiter():
|
||||
return
|
||||
yield # pragma: no cover - makes this an async generator
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=_empty_aiter(),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=log_obj,
|
||||
)
|
||||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "anext-finish-reason-session"
|
||||
|
||||
chunk = await wrapper.__anext__()
|
||||
|
||||
assert chunk.choices[0].finish_reason is not None
|
||||
# Still this call's own ids - not restored yet.
|
||||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "anext-finish-reason-session"
|
||||
|
||||
# A caller that keeps iterating still gets a deterministic restore
|
||||
# right here, on the next real StopAsyncIteration.
|
||||
with pytest.raises(StopAsyncIteration):
|
||||
await wrapper.__anext__()
|
||||
assert trace_id_var.get() == "outer-trace-anext-finish-reason"
|
||||
assert session_id_var.get() == "outer-session-anext-finish-reason"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_wrapper_anext_max_duration_timeout_restores_consumer_correlation_context(monkeypatch):
|
||||
"""_check_max_streaming_duration() raises litellm.Timeout when a client keeps
|
||||
an async stream open past LITELLM_MAX_STREAMING_DURATION_SECONDS. That raise
|
||||
must flow through the same except Exception -> _handle_stream_fallback_error
|
||||
path as every other failure so the consumer's outer correlation context gets
|
||||
restored - calling the check before entering __anext__()'s try block would
|
||||
let the Timeout bypass that restoration entirely."""
|
||||
monkeypatch.setattr(litellm.constants, "LITELLM_MAX_STREAMING_DURATION_SECONDS", 1)
|
||||
trace_id_var.set("outer-trace-max-duration")
|
||||
session_id_var.set("outer-session-max-duration")
|
||||
try:
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="max-duration-call",
|
||||
function_id="fn-max-duration",
|
||||
kwargs={"litellm_session_id": "max-duration-session"},
|
||||
)
|
||||
|
||||
async def _empty_aiter():
|
||||
return
|
||||
yield # pragma: no cover - makes this an async generator
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=_empty_aiter(),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=log_obj,
|
||||
)
|
||||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "max-duration-session"
|
||||
|
||||
wrapper._stream_created_time = time.time() - 10
|
||||
|
||||
with pytest.raises(Exception):
|
||||
await wrapper.__anext__()
|
||||
|
||||
assert trace_id_var.get() == "outer-trace-max-duration"
|
||||
assert session_id_var.get() == "outer-session-max-duration"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_wrapper_aclose_restores_consumer_correlation_context():
|
||||
"""Explicit early termination (aclose(), e.g. on client disconnect or a
|
||||
router fallback aborting an in-progress stream) must restore the caller's
|
||||
correlation context too - not just __del__'s best-effort GC-timed fallback,
|
||||
since aclose() is normally called deterministically by the consumer/
|
||||
framework, unlike __del__."""
|
||||
trace_id_var.set("outer-trace-aclose")
|
||||
session_id_var.set("outer-session-aclose")
|
||||
try:
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="aclose-call",
|
||||
function_id="fn-aclose",
|
||||
kwargs={"litellm_session_id": "aclose-session"},
|
||||
)
|
||||
|
||||
async def _empty_aiter():
|
||||
return
|
||||
yield # pragma: no cover - makes this an async generator
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=_empty_aiter(),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=log_obj,
|
||||
)
|
||||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "aclose-session"
|
||||
|
||||
await wrapper.aclose()
|
||||
|
||||
assert trace_id_var.get() == "outer-trace-aclose"
|
||||
assert session_id_var.get() == "outer-session-aclose"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_wrapper_aclose_keeps_context_active_through_close_failure_diagnostic(monkeypatch):
|
||||
"""If closing the underlying provider stream raises, aclose()'s except
|
||||
branch logs a debug diagnostic. That log line must still carry the
|
||||
closing stream's own trace_id/session_id - the outer context must not be
|
||||
restored until after the close attempt (and its diagnostic) completes."""
|
||||
trace_id_var.set("outer-trace-close-fail")
|
||||
session_id_var.set("outer-session-close-fail")
|
||||
try:
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="close-fail-call",
|
||||
function_id="fn-close-fail",
|
||||
kwargs={"litellm_session_id": "close-fail-session"},
|
||||
)
|
||||
|
||||
class _RaisingAsyncCloseStream:
|
||||
async def aclose(self):
|
||||
raise RuntimeError("boom closing stream")
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
raise StopAsyncIteration
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=_RaisingAsyncCloseStream(),
|
||||
model="gpt-3.5-turbo",
|
||||
logging_obj=log_obj,
|
||||
)
|
||||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
assert session_id_var.get() == "close-fail-session"
|
||||
|
||||
captured_ids = {}
|
||||
real_debug = verbose_logger.debug
|
||||
|
||||
def fake_debug(msg, *args, **kwargs):
|
||||
if "error closing completion_stream" in msg:
|
||||
captured_ids["trace_id"] = trace_id_var.get()
|
||||
captured_ids["session_id"] = session_id_var.get()
|
||||
return real_debug(msg, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(verbose_logger, "debug", fake_debug)
|
||||
|
||||
await wrapper.aclose()
|
||||
|
||||
assert captured_ids["trace_id"] == log_obj.litellm_trace_id
|
||||
assert captured_ids["session_id"] == "close-fail-session"
|
||||
assert trace_id_var.get() == "outer-trace-close-fail"
|
||||
assert session_id_var.get() == "outer-session-close-fail"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_handle_stream_fallback_error_restores_context_only_after_exception_mapping(monkeypatch):
|
||||
"""_map_anthropic_exception/_map_aleph_alpha_exception synchronously log a
|
||||
debug diagnostic (the raw status code) as part of exception_type()'s
|
||||
mapping. The consumer's outer context must not be restored until that
|
||||
mapping call returns, or the diagnostic log line would carry the outer
|
||||
(or empty) trace_id/session_id instead of the failing stream's own."""
|
||||
trace_id_var.set("outer-trace-fallback")
|
||||
session_id_var.set("outer-session-fallback")
|
||||
try:
|
||||
log_obj = Logging(
|
||||
model="claude-3-opus",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="fallback-error-call",
|
||||
function_id="fn-fallback-error",
|
||||
kwargs={"litellm_session_id": "fallback-error-session"},
|
||||
)
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=iter([]),
|
||||
model="claude-3-opus",
|
||||
custom_llm_provider="anthropic",
|
||||
logging_obj=log_obj,
|
||||
)
|
||||
|
||||
captured_ids = {}
|
||||
|
||||
def fake_exception_type(**kwargs):
|
||||
captured_ids["trace_id"] = trace_id_var.get()
|
||||
captured_ids["session_id"] = session_id_var.get()
|
||||
return ValueError("mapped boom")
|
||||
|
||||
monkeypatch.setattr("litellm.litellm_core_utils.streaming_handler.exception_type", fake_exception_type)
|
||||
|
||||
with pytest.raises(Exception):
|
||||
wrapper._handle_stream_fallback_error(RuntimeError("boom"))
|
||||
|
||||
# The mapper ran while the stream's own ids were still active.
|
||||
assert captured_ids["trace_id"] == log_obj.litellm_trace_id
|
||||
assert captured_ids["session_id"] == "fallback-error-session"
|
||||
# Restored to the consumer's outer context once mapping/raise completes.
|
||||
assert trace_id_var.get() == "outer-trace-fallback"
|
||||
assert session_id_var.get() == "outer-session-fallback"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
|
|
|||
|
|
@ -3562,6 +3562,8 @@ def test_supports_native_structured_outputs():
|
|||
assert config._supports_native_structured_outputs("nvidia.nemotron-nano-3-30b")
|
||||
# DeepSeek: old substring "deepseek-v3.1" didn't match real ID
|
||||
assert config._supports_native_structured_outputs("deepseek.v3-v1:0")
|
||||
assert config._supports_native_structured_outputs("deepseek.v3.2")
|
||||
assert config._supports_native_structured_outputs("zai.glm-5")
|
||||
|
||||
# Unsupported models -- should fall back to tool-call approach
|
||||
assert not config._supports_native_structured_outputs(
|
||||
|
|
|
|||
|
|
@ -432,10 +432,10 @@ class TestXAICostCalculator:
|
|||
model="grok-4.20-beta-0309-reasoning", usage=usage
|
||||
)
|
||||
|
||||
# Input: 100 tokens * $2e-6 = $0.0002
|
||||
# Output: 200 tokens * $6e-6 = $0.0012
|
||||
expected_prompt_cost = 100 * 2e-6
|
||||
expected_completion_cost = 200 * 6e-6
|
||||
# Input: 100 tokens * $1.25e-6 = $0.000125
|
||||
# Output: 200 tokens * $2.5e-6 = $0.0005
|
||||
expected_prompt_cost = 100 * 1.25e-6
|
||||
expected_completion_cost = 200 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
|
@ -448,10 +448,38 @@ class TestXAICostCalculator:
|
|||
model="grok-4.20-beta-0309-non-reasoning", usage=usage
|
||||
)
|
||||
|
||||
# Input: 50 tokens * $2e-6 = $0.0001
|
||||
# Output: 100 tokens * $6e-6 = $0.0006
|
||||
expected_prompt_cost = 50 * 2e-6
|
||||
expected_completion_cost = 100 * 6e-6
|
||||
# Input: 50 tokens * $1.25e-6 = $0.0000625
|
||||
# Output: 100 tokens * $2.5e-6 = $0.00025
|
||||
expected_prompt_cost = 50 * 1.25e-6
|
||||
expected_completion_cost = 100 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_grok_4_20_at_exactly_200k_prompt_tokens_uses_higher_tier(self):
|
||||
"""xAI bills the >=200k tier once the prompt reaches 200k, so the boundary is inclusive."""
|
||||
usage = Usage(prompt_tokens=200_000, completion_tokens=1_000, total_tokens=201_000)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="grok-4.20-0309-reasoning", usage=usage
|
||||
)
|
||||
|
||||
expected_prompt_cost = 200_000 * 2.5e-6
|
||||
expected_completion_cost = 1_000 * 5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_grok_4_20_just_below_200k_prompt_tokens_uses_base_tier(self):
|
||||
"""One token under the boundary still bills at the base rates."""
|
||||
usage = Usage(prompt_tokens=199_999, completion_tokens=1_000, total_tokens=200_999)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="grok-4.20-0309-reasoning", usage=usage
|
||||
)
|
||||
|
||||
expected_prompt_cost = 199_999 * 1.25e-6
|
||||
expected_completion_cost = 1_000 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
|
@ -464,10 +492,10 @@ class TestXAICostCalculator:
|
|||
model="grok-4.20-multi-agent-beta-0309", usage=usage
|
||||
)
|
||||
|
||||
# Input: 200 tokens * $2e-6 = $0.0004
|
||||
# Output: 300 tokens * $6e-6 = $0.0018
|
||||
expected_prompt_cost = 200 * 2e-6
|
||||
expected_completion_cost = 300 * 6e-6
|
||||
# Input: 200 tokens * $1.25e-6 = $0.00025
|
||||
# Output: 300 tokens * $2.5e-6 = $0.00075
|
||||
expected_prompt_cost = 200 * 1.25e-6
|
||||
expected_completion_cost = 300 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
|
|
|||
|
|
@ -7,9 +7,7 @@ from unittest.mock import AsyncMock, MagicMock
|
|||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
sys.path.insert(0, os.path.abspath("../../../..")) # Adds the parent directory to the system path
|
||||
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
_adjust_dates_for_timezone,
|
||||
|
|
@ -108,8 +106,7 @@ async def test_get_daily_activity_order_has_id_tiebreaker():
|
|||
mock_table.find_many.assert_called_once()
|
||||
order = mock_table.find_many.call_args[1]["order"]
|
||||
assert order == [{"date": "desc"}, {"id": "asc"}], (
|
||||
f"order must include the id tiebreaker after date for stable offset "
|
||||
f"pagination (see #30164); got {order!r}"
|
||||
f"order must include the id tiebreaker after date for stable offset pagination (see #30164); got {order!r}"
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -301,9 +298,7 @@ async def test_get_api_key_metadata_returns_active_key_metadata():
|
|||
mock_active_key.key_alias = "my-active-key"
|
||||
mock_active_key.team_id = "team-abc"
|
||||
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[mock_active_key]
|
||||
)
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[mock_active_key])
|
||||
|
||||
result = await get_api_key_metadata(
|
||||
prisma_client=mock_prisma,
|
||||
|
|
@ -329,9 +324,7 @@ async def test_get_api_key_metadata_falls_back_to_deleted_keys():
|
|||
mock_deleted_key.key_alias = "toto-test-2"
|
||||
mock_deleted_key.team_id = "team-xyz"
|
||||
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(
|
||||
return_value=[mock_deleted_key]
|
||||
)
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[mock_deleted_key])
|
||||
|
||||
result = await get_api_key_metadata(
|
||||
prisma_client=mock_prisma,
|
||||
|
|
@ -360,9 +353,7 @@ async def test_get_api_key_metadata_mixed_active_and_deleted_keys():
|
|||
mock_active_key.key_alias = "active-alias"
|
||||
mock_active_key.team_id = "team-active"
|
||||
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[mock_active_key]
|
||||
)
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[mock_active_key])
|
||||
|
||||
# One deleted key found
|
||||
mock_deleted_key = MagicMock()
|
||||
|
|
@ -370,9 +361,7 @@ async def test_get_api_key_metadata_mixed_active_and_deleted_keys():
|
|||
mock_deleted_key.key_alias = "deleted-alias"
|
||||
mock_deleted_key.team_id = "team-deleted"
|
||||
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(
|
||||
return_value=[mock_deleted_key]
|
||||
)
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[mock_deleted_key])
|
||||
|
||||
result = await get_api_key_metadata(
|
||||
prisma_client=mock_prisma,
|
||||
|
|
@ -397,13 +386,9 @@ async def test_get_api_key_metadata_deleted_table_not_queried_when_all_keys_foun
|
|||
mock_active_key.key_alias = "alias-1"
|
||||
mock_active_key.team_id = "team-1"
|
||||
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[mock_active_key]
|
||||
)
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[mock_active_key])
|
||||
mock_prisma.db.litellm_deletedverificationtoken = MagicMock()
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
|
||||
|
||||
result = await get_api_key_metadata(
|
||||
prisma_client=mock_prisma,
|
||||
|
|
@ -425,9 +410,7 @@ async def test_get_api_key_metadata_deleted_table_error_handled_gracefully():
|
|||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
|
||||
# Deleted table raises an error (e.g., table doesn't exist in older schema)
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(
|
||||
side_effect=Exception("Table not found")
|
||||
)
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(side_effect=Exception("Table not found"))
|
||||
|
||||
result = await get_api_key_metadata(
|
||||
prisma_client=mock_prisma,
|
||||
|
|
@ -458,9 +441,7 @@ async def test_get_api_key_metadata_regenerated_key_uses_most_recent_deleted_rec
|
|||
mock_deleted_2.team_id = "older-team"
|
||||
|
||||
# Ordered by deleted_at desc, so first record is the most recent
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(
|
||||
return_value=[mock_deleted_1, mock_deleted_2]
|
||||
)
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[mock_deleted_1, mock_deleted_2])
|
||||
|
||||
result = await get_api_key_metadata(
|
||||
prisma_client=mock_prisma,
|
||||
|
|
@ -633,9 +614,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
|
|||
mock_deleted_key.team_id = "69cd4b77-b095-4489-8c46-4f2f31d840a2"
|
||||
|
||||
mock_prisma.db.litellm_deletedverificationtoken = MagicMock()
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(
|
||||
return_value=[mock_deleted_key]
|
||||
)
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[mock_deleted_key])
|
||||
|
||||
result = await get_daily_activity_aggregated(
|
||||
prisma_client=mock_prisma,
|
||||
|
|
@ -754,15 +733,9 @@ async def test_model_groups_breakdown_keys_by_public_name_with_model_fallback():
|
|||
mock_prisma.db = MagicMock()
|
||||
|
||||
records = [
|
||||
_daily_user_spend_record(
|
||||
user_id="u1", api_key="key-1", spend=7.0, model="gpt-5.2", model_group="gpt-5.2-eu"
|
||||
),
|
||||
_daily_user_spend_record(
|
||||
user_id="u1", api_key="key-1", spend=3.0, model="gpt-5.2", model_group=None
|
||||
),
|
||||
_daily_user_spend_record(
|
||||
user_id="u1", api_key="key-1", spend=2.0, model="claude-x", model_group=""
|
||||
),
|
||||
_daily_user_spend_record(user_id="u1", api_key="key-1", spend=7.0, model="gpt-5.2", model_group="gpt-5.2-eu"),
|
||||
_daily_user_spend_record(user_id="u1", api_key="key-1", spend=3.0, model="gpt-5.2", model_group=None),
|
||||
_daily_user_spend_record(user_id="u1", api_key="key-1", spend=2.0, model="claude-x", model_group=""),
|
||||
]
|
||||
|
||||
mock_table = MagicMock()
|
||||
|
|
@ -829,9 +802,7 @@ class TestAdjustDatesForTimezone:
|
|||
],
|
||||
)
|
||||
def test_returns_input_dates_unchanged_for_any_offset(self, offset_minutes):
|
||||
start, end = _adjust_dates_for_timezone(
|
||||
"2026-05-29", "2026-05-29", offset_minutes
|
||||
)
|
||||
start, end = _adjust_dates_for_timezone("2026-05-29", "2026-05-29", offset_minutes)
|
||||
assert start == "2026-05-29"
|
||||
assert end == "2026-05-29"
|
||||
|
||||
|
|
@ -859,9 +830,7 @@ class TestAdjustDatesForTimezone:
|
|||
exceeded the multi-day total by ~50% over a 5-day IST window.
|
||||
"""
|
||||
days = ["2026-05-29", "2026-05-30", "2026-05-31", "2026-06-01", "2026-06-02"]
|
||||
single_day_ranges = [
|
||||
_adjust_dates_for_timezone(d, d, offset_minutes) for d in days
|
||||
]
|
||||
single_day_ranges = [_adjust_dates_for_timezone(d, d, offset_minutes) for d in days]
|
||||
multi_day_range = _adjust_dates_for_timezone(days[0], days[-1], offset_minutes)
|
||||
|
||||
per_day_starts = [r[0] for r in single_day_ranges]
|
||||
|
|
@ -894,9 +863,7 @@ class TestAdjustDatesForTimezoneLiveEnd:
|
|||
assert (start, end) == ("2026-07-06", "2026-08-06")
|
||||
|
||||
def test_without_opt_in_live_range_keeps_pass_through(self):
|
||||
start, end = _adjust_dates_for_timezone(
|
||||
"2026-07-06", "2026-08-05", 420, utc_now=self.PT_EVENING_UTC
|
||||
)
|
||||
start, end = _adjust_dates_for_timezone("2026-07-06", "2026-08-05", 420, utc_now=self.PT_EVENING_UTC)
|
||||
assert (start, end) == ("2026-07-06", "2026-08-05")
|
||||
|
||||
def test_pt_historical_range_is_untouched(self):
|
||||
|
|
@ -1173,3 +1140,348 @@ class TestEverySavingsDriverSurvivesTheReadPath:
|
|||
assert f"total_{driver}" in DailySpendMetadata.model_fields, (
|
||||
f"total_{driver} is missing, so the range summary omits the driver"
|
||||
)
|
||||
|
||||
|
||||
def _spend_record(api_key, *, model="gpt-4o-mini-ptu", spend=0.0, ptu_flat_cost=0.0):
|
||||
return SimpleNamespace(
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
model_group=None,
|
||||
mcp_namespaced_tool_name=None,
|
||||
custom_llm_provider="openai",
|
||||
endpoint=None,
|
||||
spend=spend,
|
||||
prompt_tokens=0,
|
||||
completion_tokens=0,
|
||||
cache_read_input_tokens=0,
|
||||
cache_creation_input_tokens=0,
|
||||
compression_saved_tokens=0,
|
||||
compression_savings_spend=0,
|
||||
prompt_caching_savings_spend=0,
|
||||
autorouter_savings_spend=0,
|
||||
total_tokens=0,
|
||||
api_requests=0,
|
||||
successful_requests=0,
|
||||
failed_requests=0,
|
||||
ptu_flat_cost=ptu_flat_cost,
|
||||
)
|
||||
|
||||
|
||||
def test_update_metrics_accumulates_ptu_flat_cost():
|
||||
metrics = update_metrics(SpendMetrics(), _spend_record("real-key", spend=1.0, ptu_flat_cost=240.0))
|
||||
assert metrics.flat_cost == 240.0
|
||||
assert metrics.spend == 1.0
|
||||
|
||||
|
||||
def test_ptu_sentinel_excluded_from_key_breakdown_but_flat_cost_aggregates():
|
||||
from litellm.constants import PTU_SENTINEL_API_KEY
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import update_breakdown_metrics
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import BreakdownMetrics
|
||||
|
||||
breakdown = BreakdownMetrics()
|
||||
update_breakdown_metrics(breakdown, _spend_record("real-key", spend=5.0, ptu_flat_cost=0.0), {}, {}, {})
|
||||
update_breakdown_metrics(breakdown, _spend_record(PTU_SENTINEL_API_KEY, spend=0.0, ptu_flat_cost=240.0), {}, {}, {})
|
||||
|
||||
model_bucket = breakdown.models["gpt-4o-mini-ptu"]
|
||||
# flat cost aggregates into the parent model metrics
|
||||
assert model_bucket.metrics.flat_cost == 240.0
|
||||
assert model_bucket.metrics.spend == 5.0
|
||||
# the sentinel never appears as an api_key row; only the real key does
|
||||
assert PTU_SENTINEL_API_KEY not in model_bucket.api_key_breakdown
|
||||
assert "real-key" in model_bucket.api_key_breakdown
|
||||
|
||||
|
||||
def _grouping_row(
|
||||
group_level,
|
||||
*,
|
||||
api_key=None,
|
||||
model=None,
|
||||
model_group=None,
|
||||
custom_llm_provider="openai",
|
||||
mcp_namespaced_tool_name=None,
|
||||
endpoint=None,
|
||||
spend=0.0,
|
||||
ptu_flat_cost=0.0,
|
||||
):
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import _GroupingSetsRow
|
||||
|
||||
return _GroupingSetsRow(
|
||||
date="2024-01-01",
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
model_group=model_group,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
mcp_namespaced_tool_name=mcp_namespaced_tool_name,
|
||||
endpoint=endpoint,
|
||||
group_level=group_level,
|
||||
spend=spend,
|
||||
ptu_flat_cost=ptu_flat_cost,
|
||||
prompt_tokens=0,
|
||||
completion_tokens=0,
|
||||
cache_read_input_tokens=0,
|
||||
cache_creation_input_tokens=0,
|
||||
compression_saved_tokens=0,
|
||||
compression_savings_spend=0.0,
|
||||
prompt_caching_savings_spend=0.0,
|
||||
autorouter_savings_spend=0.0,
|
||||
api_requests=0,
|
||||
successful_requests=0,
|
||||
failed_requests=0,
|
||||
)
|
||||
|
||||
|
||||
def test_grouping_sets_dispatcher_excludes_ptu_sentinel_from_key_breakdowns():
|
||||
"""The GROUPING SETS path must mirror the per-row path: the flat-cost sentinel
|
||||
aggregates into the date/model/total metrics but never surfaces as an api_key."""
|
||||
from litellm.constants import PTU_SENTINEL_API_KEY
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
_GROUP_DATE_API_KEY,
|
||||
_GROUP_DATE_MODEL,
|
||||
_GROUP_DATE_MODEL_API_KEY,
|
||||
_GROUP_GRAND_TOTAL,
|
||||
_aggregate_grouping_sets_records_sync,
|
||||
)
|
||||
|
||||
records = [
|
||||
_grouping_row(_GROUP_DATE_API_KEY, api_key="real-key", spend=5.0),
|
||||
_grouping_row(_GROUP_DATE_API_KEY, api_key=PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0),
|
||||
_grouping_row(_GROUP_DATE_MODEL, model="gpt-4o-mini-ptu", spend=5.0, ptu_flat_cost=240.0),
|
||||
_grouping_row(_GROUP_DATE_MODEL_API_KEY, model="gpt-4o-mini-ptu", api_key="real-key", spend=5.0),
|
||||
_grouping_row(
|
||||
_GROUP_DATE_MODEL_API_KEY, model="gpt-4o-mini-ptu", api_key=PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0
|
||||
),
|
||||
_grouping_row(_GROUP_GRAND_TOTAL, spend=5.0, ptu_flat_cost=240.0),
|
||||
]
|
||||
|
||||
aggregated = _aggregate_grouping_sets_records_sync(records=records, api_key_metadata={})
|
||||
|
||||
assert aggregated["totals"].flat_cost == 240.0
|
||||
day = aggregated["results"][0]
|
||||
assert PTU_SENTINEL_API_KEY not in day.breakdown.api_keys
|
||||
assert "real-key" in day.breakdown.api_keys
|
||||
|
||||
model_bucket = day.breakdown.models["gpt-4o-mini-ptu"]
|
||||
assert model_bucket.metrics.flat_cost == 240.0
|
||||
assert model_bucket.metrics.spend == 5.0
|
||||
assert PTU_SENTINEL_API_KEY not in model_bucket.api_key_breakdown
|
||||
assert "real-key" in model_bucket.api_key_breakdown
|
||||
|
||||
|
||||
def test_grouping_sets_dispatcher_populates_every_breakdown_level():
|
||||
"""Every GROUPING SETS level lands in its bucket, and the flat-cost sentinel
|
||||
is kept out of the model_group and provider api_key sub-breakdowns too."""
|
||||
from litellm.constants import PTU_SENTINEL_API_KEY
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
_GROUP_DATE_ENDPOINT,
|
||||
_GROUP_DATE_ENDPOINT_API_KEY,
|
||||
_GROUP_DATE_MCP,
|
||||
_GROUP_DATE_MCP_API_KEY,
|
||||
_GROUP_DATE_MODEL_GROUP,
|
||||
_GROUP_DATE_MODEL_GROUP_API_KEY,
|
||||
_GROUP_DATE_PROVIDER,
|
||||
_GROUP_DATE_PROVIDER_API_KEY,
|
||||
_aggregate_grouping_sets_records_sync,
|
||||
)
|
||||
|
||||
records = [
|
||||
_grouping_row(_GROUP_DATE_MODEL_GROUP, model_group="grp", spend=4.0, ptu_flat_cost=240.0),
|
||||
_grouping_row(_GROUP_DATE_MODEL_GROUP_API_KEY, model_group="grp", api_key="real-key", spend=4.0),
|
||||
_grouping_row(
|
||||
_GROUP_DATE_MODEL_GROUP_API_KEY, model_group="grp", api_key=PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0
|
||||
),
|
||||
_grouping_row(_GROUP_DATE_PROVIDER, custom_llm_provider="azure", spend=4.0),
|
||||
_grouping_row(_GROUP_DATE_PROVIDER_API_KEY, custom_llm_provider="azure", api_key="real-key", spend=4.0),
|
||||
_grouping_row(
|
||||
_GROUP_DATE_PROVIDER_API_KEY,
|
||||
custom_llm_provider="azure",
|
||||
api_key=PTU_SENTINEL_API_KEY,
|
||||
ptu_flat_cost=240.0,
|
||||
),
|
||||
_grouping_row(_GROUP_DATE_MCP, mcp_namespaced_tool_name="srv/tool", spend=2.0),
|
||||
_grouping_row(_GROUP_DATE_MCP_API_KEY, mcp_namespaced_tool_name="srv/tool", api_key="real-key", spend=2.0),
|
||||
_grouping_row(_GROUP_DATE_ENDPOINT, endpoint="/v1/chat/completions", spend=3.0),
|
||||
_grouping_row(_GROUP_DATE_ENDPOINT_API_KEY, endpoint="/v1/chat/completions", api_key="real-key", spend=3.0),
|
||||
]
|
||||
|
||||
aggregated = _aggregate_grouping_sets_records_sync(records=records, api_key_metadata={})
|
||||
day = aggregated["results"][0]
|
||||
|
||||
group_bucket = day.breakdown.model_groups["grp"]
|
||||
assert group_bucket.metrics.flat_cost == 240.0
|
||||
assert PTU_SENTINEL_API_KEY not in group_bucket.api_key_breakdown
|
||||
assert "real-key" in group_bucket.api_key_breakdown
|
||||
|
||||
provider_bucket = day.breakdown.providers["azure"]
|
||||
assert PTU_SENTINEL_API_KEY not in provider_bucket.api_key_breakdown
|
||||
assert "real-key" in provider_bucket.api_key_breakdown
|
||||
|
||||
assert "real-key" in day.breakdown.mcp_servers["srv/tool"].api_key_breakdown
|
||||
assert "real-key" in day.breakdown.endpoints["/v1/chat/completions"].api_key_breakdown
|
||||
|
||||
|
||||
def test_grouping_sets_dispatcher_keeps_ptu_flat_cost_out_of_the_provider_breakdown():
|
||||
"""Sentinel rows carry no provider, so their flat cost must not surface under the
|
||||
"unknown" provider - the per-row path skips them for exactly the same reason."""
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
_GROUP_DATE_PROVIDER,
|
||||
_aggregate_grouping_sets_records_sync,
|
||||
)
|
||||
|
||||
records = [
|
||||
_grouping_row(_GROUP_DATE_PROVIDER, custom_llm_provider="azure", spend=4.0),
|
||||
# the sentinel's own provider-level row: empty provider, flat cost only
|
||||
_grouping_row(_GROUP_DATE_PROVIDER, custom_llm_provider="", ptu_flat_cost=240.0),
|
||||
]
|
||||
|
||||
aggregated = _aggregate_grouping_sets_records_sync(records=records, api_key_metadata={})
|
||||
providers = aggregated["results"][0].breakdown.providers
|
||||
|
||||
# the bucket is still reported (a legacy all-zero row must not vanish); only the
|
||||
# flat cost is withheld, so no provider is credited with PTU capacity cost
|
||||
assert providers["azure"].metrics.spend == 4.0
|
||||
assert sum(bucket.metrics.flat_cost for bucket in providers.values()) == 0.0
|
||||
|
||||
|
||||
def test_grouping_sets_dispatcher_keeps_a_real_provider_row_that_shares_the_sentinel_shape():
|
||||
"""A request row whose provider is empty still gets its "unknown" bucket - only the
|
||||
flat cost is withheld, so provider attribution of real spend is unchanged."""
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
_GROUP_DATE_PROVIDER,
|
||||
_aggregate_grouping_sets_records_sync,
|
||||
)
|
||||
|
||||
records = [_grouping_row(_GROUP_DATE_PROVIDER, custom_llm_provider="", spend=4.0, ptu_flat_cost=240.0)]
|
||||
|
||||
aggregated = _aggregate_grouping_sets_records_sync(records=records, api_key_metadata={})
|
||||
unknown = aggregated["results"][0].breakdown.providers["unknown"]
|
||||
|
||||
assert unknown.metrics.spend == 4.0
|
||||
assert unknown.metrics.flat_cost == 0.0
|
||||
|
||||
|
||||
def test_update_breakdown_metrics_covers_mcp_endpoint_and_entity():
|
||||
"""A full request record fans out into the mcp, endpoint, provider and entity
|
||||
breakdowns, while the flat-cost sentinel stays out of the entity api_key sub-map."""
|
||||
from litellm.constants import PTU_SENTINEL_API_KEY
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import update_breakdown_metrics
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import BreakdownMetrics
|
||||
|
||||
breakdown = BreakdownMetrics()
|
||||
record = SimpleNamespace(
|
||||
api_key="real-key",
|
||||
model="gpt-4o-mini-ptu",
|
||||
model_group="grp",
|
||||
mcp_namespaced_tool_name="srv/tool",
|
||||
custom_llm_provider="azure",
|
||||
endpoint="/v1/chat/completions",
|
||||
spend=5.0,
|
||||
prompt_tokens=0,
|
||||
completion_tokens=0,
|
||||
cache_read_input_tokens=0,
|
||||
cache_creation_input_tokens=0,
|
||||
compression_saved_tokens=0,
|
||||
compression_savings_spend=0,
|
||||
prompt_caching_savings_spend=0,
|
||||
autorouter_savings_spend=0,
|
||||
total_tokens=0,
|
||||
api_requests=0,
|
||||
successful_requests=0,
|
||||
failed_requests=0,
|
||||
ptu_flat_cost=0.0,
|
||||
team_id="team-1",
|
||||
)
|
||||
update_breakdown_metrics(breakdown, record, {}, {}, {}, entity_id_field="team_id")
|
||||
|
||||
assert "srv/tool" in breakdown.mcp_servers
|
||||
assert "real-key" in breakdown.mcp_servers["srv/tool"].api_key_breakdown
|
||||
assert "/v1/chat/completions" in breakdown.endpoints
|
||||
assert "azure" in breakdown.providers
|
||||
assert "team-1" in breakdown.entities
|
||||
assert "real-key" in breakdown.entities["team-1"].api_key_breakdown
|
||||
|
||||
sentinel = SimpleNamespace(**{**record.__dict__, "api_key": PTU_SENTINEL_API_KEY, "ptu_flat_cost": 240.0})
|
||||
update_breakdown_metrics(breakdown, sentinel, {}, {}, {}, entity_id_field="team_id")
|
||||
assert PTU_SENTINEL_API_KEY not in breakdown.entities["team-1"].api_key_breakdown
|
||||
assert breakdown.entities["team-1"].metrics.flat_cost == 240.0
|
||||
|
||||
|
||||
def test_grouping_sets_dispatcher_keeps_an_all_zero_legacy_provider_bucket():
|
||||
"""LiteLLM_DailyTeamSpend predates its api_requests column; the migration that added it
|
||||
backfilled NOT NULL DEFAULT 0, so a legacy keyless row is all zeroes. Dropping those
|
||||
would silently remove a provider the base build reported."""
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
_GROUP_DATE_PROVIDER,
|
||||
_aggregate_grouping_sets_records_sync,
|
||||
)
|
||||
|
||||
records = [
|
||||
_grouping_row(_GROUP_DATE_PROVIDER, custom_llm_provider="ollama"), # spend/tokens/requests all 0
|
||||
_grouping_row(_GROUP_DATE_PROVIDER, custom_llm_provider="openai", spend=0.25),
|
||||
]
|
||||
|
||||
providers = _aggregate_grouping_sets_records_sync(records=records, api_key_metadata={})["results"][
|
||||
0
|
||||
].breakdown.providers
|
||||
|
||||
assert set(providers) == {"ollama", "openai"}
|
||||
assert providers["ollama"].metrics.spend == 0.0
|
||||
assert providers["ollama"].metrics.flat_cost == 0.0
|
||||
|
||||
|
||||
class TestSentinelRowsDisplayTheirModelName:
|
||||
"""A sentinel row keys on the deployment id so a rename cannot move it. The usage views
|
||||
render the breakdown key directly as a label, so the read path has to show the name."""
|
||||
|
||||
@staticmethod
|
||||
def _breakdown(records):
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import update_breakdown_metrics
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import BreakdownMetrics
|
||||
|
||||
breakdown = BreakdownMetrics()
|
||||
for record in records:
|
||||
update_breakdown_metrics(breakdown, record, {}, {}, {})
|
||||
return breakdown
|
||||
|
||||
@staticmethod
|
||||
def _sentinel(*, model_id, model_group, flat_cost=480.0):
|
||||
from litellm.constants import PTU_SENTINEL_API_KEY
|
||||
|
||||
record = _spend_record(PTU_SENTINEL_API_KEY, model=model_id, spend=0.0, ptu_flat_cost=flat_cost)
|
||||
record.model_group = model_group
|
||||
return record
|
||||
|
||||
def test_models_breakdown_keys_a_sentinel_row_on_its_public_name(self):
|
||||
models = self._breakdown([self._sentinel(model_id="dep-1", model_group="gpt-4o-ptu")]).models
|
||||
|
||||
assert "gpt-4o-ptu" in models, f"the UI would label this row a UUID: {list(models)}"
|
||||
assert "dep-1" not in models
|
||||
assert models["gpt-4o-ptu"].metrics.flat_cost == pytest.approx(480.0)
|
||||
|
||||
def test_two_deployments_sharing_a_name_merge_under_it(self):
|
||||
"""The write path stopped collapsing them, so the read path has to."""
|
||||
models = self._breakdown(
|
||||
[
|
||||
self._sentinel(model_id="dep-a", model_group="gpt-4o-ptu", flat_cost=240.0),
|
||||
self._sentinel(model_id="dep-b", model_group="gpt-4o-ptu", flat_cost=120.0),
|
||||
]
|
||||
).models
|
||||
|
||||
assert list(models) == ["gpt-4o-ptu"]
|
||||
assert models["gpt-4o-ptu"].metrics.flat_cost == pytest.approx(360.0)
|
||||
|
||||
def test_a_request_row_still_keys_on_its_model(self):
|
||||
"""Scoped to sentinel rows: a request row keys on model as it always has, even
|
||||
though it also carries a model_group."""
|
||||
record = _spend_record("real-key", model="gemini/gemini-2.5-flash", spend=1.25)
|
||||
record.model_group = "gemini-live"
|
||||
|
||||
models = self._breakdown([record]).models
|
||||
|
||||
assert "gemini/gemini-2.5-flash" in models
|
||||
assert "gemini-live" not in models
|
||||
|
||||
def test_a_sentinel_row_without_a_model_group_falls_back_to_the_id(self):
|
||||
"""Never drop the charge: an unexpected row with no display name still reports."""
|
||||
models = self._breakdown([self._sentinel(model_id="dep-1", model_group=None)]).models
|
||||
|
||||
assert models["dep-1"].metrics.flat_cost == pytest.approx(480.0)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,382 @@
|
|||
import datetime
|
||||
import json
|
||||
|
||||
"""Tests for PTU config on the model deployment (v1 model-settings design)."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_merged_ptu_model_info,
|
||||
_validate_ptu_model_info,
|
||||
)
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo, updateDeployment
|
||||
|
||||
|
||||
def test_model_info_accepts_valid_ptu_fields():
|
||||
info = ModelInfo(id="x", team_id="t", ptu_count=5, cost_per_ptu_per_hour=2.0)
|
||||
assert info.ptu_count == 5
|
||||
assert info.cost_per_ptu_per_hour == 2.0
|
||||
|
||||
|
||||
def test_model_info_rejects_non_positive_count():
|
||||
with pytest.raises(ValueError):
|
||||
ModelInfo(id="x", team_id="t", ptu_count=0, cost_per_ptu_per_hour=2.0)
|
||||
|
||||
|
||||
def test_model_info_rejects_negative_rate():
|
||||
with pytest.raises(ValueError):
|
||||
ModelInfo(id="x", team_id="t", ptu_count=5, cost_per_ptu_per_hour=-1.0)
|
||||
|
||||
|
||||
def test_model_info_rejects_a_count_beyond_the_cap():
|
||||
"""flat cost multiplies the count by a float, and an unbounded int overflows that
|
||||
conversion, which aborted the rollup for every team rather than skipping one model."""
|
||||
with pytest.raises(ValueError):
|
||||
ModelInfo(id="x", team_id="t", ptu_count=10**400, cost_per_ptu_per_hour=2.0)
|
||||
|
||||
|
||||
def test_model_info_accepts_a_count_at_the_cap():
|
||||
info = ModelInfo(id="x", team_id="t", ptu_count=ModelInfo.MAX_PTU_COUNT, cost_per_ptu_per_hour=2.0)
|
||||
assert info.ptu_count == ModelInfo.MAX_PTU_COUNT
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rate", [float("nan"), float("inf"), float("-inf")])
|
||||
def test_model_info_rejects_a_non_finite_rate(rate):
|
||||
"""NaN compares False against every bound, so a bare `< 0` check let it through and the
|
||||
deployment then accrued a flat cost of nan."""
|
||||
with pytest.raises(ValueError):
|
||||
ModelInfo(id="x", team_id="t", ptu_count=5, cost_per_ptu_per_hour=rate)
|
||||
|
||||
|
||||
def test_model_info_rejects_a_rate_beyond_the_cap():
|
||||
with pytest.raises(ValueError):
|
||||
ModelInfo(id="x", team_id="t", ptu_count=5, cost_per_ptu_per_hour=ModelInfo.MAX_COST_PER_PTU_PER_HOUR * 2)
|
||||
|
||||
|
||||
def test_model_info_allows_partial_delta_for_patch():
|
||||
# A PATCH delta may carry only one field; bounds-only validation must not reject it.
|
||||
info = ModelInfo(id="x", ptu_count=5)
|
||||
assert info.ptu_count == 5
|
||||
assert info.cost_per_ptu_per_hour is None
|
||||
|
||||
|
||||
def test_validate_helper_no_ptu_is_noop():
|
||||
_validate_ptu_model_info({"team_id": "t"})
|
||||
|
||||
|
||||
def test_validate_helper_requires_both_fields():
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
_validate_ptu_model_info({"team_id": "t", "ptu_count": 5})
|
||||
assert exc.value.status_code == 400
|
||||
assert "set together" in exc.value.detail
|
||||
|
||||
|
||||
def test_validate_helper_requires_team_id():
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
_validate_ptu_model_info(
|
||||
{"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "ptu_effective_from": "2026-08-01T00:00:00Z"}
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
assert "team_id" in exc.value.detail
|
||||
|
||||
|
||||
def test_validate_helper_requires_an_effective_start():
|
||||
"""Flat cost accrues from the start, so it cannot be inferred."""
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
_validate_ptu_model_info({"team_id": "t", "ptu_count": 5, "cost_per_ptu_per_hour": 2.0})
|
||||
assert exc.value.status_code == 400
|
||||
assert "ptu_effective_from is required" in exc.value.detail
|
||||
|
||||
|
||||
def test_validate_helper_passes_full_config():
|
||||
_validate_ptu_model_info(
|
||||
{"team_id": "t", "ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "ptu_effective_from": "2026-08-01T00:00:00Z"}
|
||||
)
|
||||
|
||||
|
||||
def test_model_info_rejects_effective_to_before_from():
|
||||
import datetime
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
ModelInfo(
|
||||
id="x",
|
||||
team_id="t",
|
||||
ptu_count=5,
|
||||
cost_per_ptu_per_hour=2.0,
|
||||
ptu_effective_from=datetime.datetime(2026, 7, 30, tzinfo=datetime.timezone.utc),
|
||||
ptu_effective_to=datetime.datetime(2026, 7, 29, tzinfo=datetime.timezone.utc),
|
||||
)
|
||||
|
||||
|
||||
def test_model_info_accepts_valid_effective_window():
|
||||
import datetime
|
||||
|
||||
info = ModelInfo(
|
||||
id="x",
|
||||
team_id="t",
|
||||
ptu_count=5,
|
||||
cost_per_ptu_per_hour=2.0,
|
||||
ptu_effective_from=datetime.datetime(2026, 7, 30, tzinfo=datetime.timezone.utc),
|
||||
ptu_effective_to=datetime.datetime(2026, 8, 30, tzinfo=datetime.timezone.utc),
|
||||
)
|
||||
assert info.ptu_effective_from is not None
|
||||
|
||||
|
||||
def test_model_info_compares_mixed_naive_and_aware_timestamps():
|
||||
import datetime
|
||||
|
||||
info = ModelInfo(
|
||||
id="x",
|
||||
team_id="t",
|
||||
ptu_count=5,
|
||||
cost_per_ptu_per_hour=2.0,
|
||||
ptu_effective_from=datetime.datetime(2026, 7, 30, 23, 0),
|
||||
ptu_effective_to=datetime.datetime(2026, 7, 31, 0, 0, tzinfo=datetime.timezone.utc),
|
||||
)
|
||||
assert info.ptu_effective_to is not None
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
ModelInfo(
|
||||
id="x",
|
||||
team_id="t",
|
||||
ptu_count=5,
|
||||
cost_per_ptu_per_hour=2.0,
|
||||
ptu_effective_from=datetime.datetime(2026, 7, 31, 2, 0),
|
||||
ptu_effective_to=datetime.datetime(2026, 7, 31, 0, 0, tzinfo=datetime.timezone.utc),
|
||||
)
|
||||
|
||||
|
||||
def test_validate_helper_rejects_effective_to_before_from():
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
_validate_ptu_model_info(
|
||||
{
|
||||
"team_id": "t",
|
||||
"ptu_count": 5,
|
||||
"cost_per_ptu_per_hour": 2.0,
|
||||
"ptu_effective_from": "2026-07-30T00:00:00Z",
|
||||
"ptu_effective_to": "2026-07-29T00:00:00Z",
|
||||
}
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
assert "ptu_effective_to" in exc.value.detail
|
||||
|
||||
|
||||
def test_validate_helper_accepts_valid_window_on_merged_info():
|
||||
_validate_ptu_model_info(
|
||||
{
|
||||
"team_id": "t",
|
||||
"ptu_count": 5,
|
||||
"cost_per_ptu_per_hour": 2.0,
|
||||
"ptu_effective_from": "2026-07-30T00:00:00Z",
|
||||
"ptu_effective_to": "2026-08-30T00:00:00Z",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_validate_helper_rejects_inverted_window_without_count_or_rate():
|
||||
"""A patch that touches only one end of the window merges to a model_info with no count
|
||||
or rate. Returning early on that shape let an inverted window reach the row, and the next
|
||||
load then failed to parse it and dropped the deployment out of the router."""
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
_validate_ptu_model_info(
|
||||
{
|
||||
"team_id": "t",
|
||||
"ptu_effective_from": "2026-08-02T00:00:00Z",
|
||||
"ptu_effective_to": "2026-08-01T00:00:00Z",
|
||||
}
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
assert "ptu_effective_to" in exc.value.detail
|
||||
|
||||
|
||||
def test_validate_helper_rejects_equal_window_bounds_without_count_or_rate():
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
_validate_ptu_model_info(
|
||||
{
|
||||
"ptu_effective_from": "2026-08-01T00:00:00Z",
|
||||
"ptu_effective_to": "2026-08-01T00:00:00Z",
|
||||
}
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
|
||||
|
||||
def test_validate_helper_accepts_ordered_window_without_count_or_rate():
|
||||
"""Window-only edits stay legal; only the ordering is enforced, and no team_id is
|
||||
demanded while the deployment carries no priced PTU config."""
|
||||
_validate_ptu_model_info(
|
||||
{
|
||||
"ptu_effective_from": "2026-08-01T00:00:00Z",
|
||||
"ptu_effective_to": "2026-08-02T00:00:00Z",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_validate_helper_accepts_a_single_open_ended_bound():
|
||||
_validate_ptu_model_info({"ptu_effective_from": "2026-08-01T00:00:00Z"})
|
||||
_validate_ptu_model_info({"ptu_effective_to": "2026-08-02T00:00:00Z"})
|
||||
|
||||
|
||||
class TestPartialPtuEditsUseTheMergedView:
|
||||
"""A PTU invariant holds over the deployment as it will exist, not over whichever
|
||||
subset of fields a caller sent. Validating the patch alone rejected an ordinary edit."""
|
||||
|
||||
@staticmethod
|
||||
def _configured():
|
||||
return Deployment(
|
||||
model_name="gpt-4o",
|
||||
litellm_params=LiteLLM_Params(model="openai/gpt-4o"),
|
||||
model_info=ModelInfo(
|
||||
id="dep-0",
|
||||
team_id="t",
|
||||
ptu_count=10,
|
||||
cost_per_ptu_per_hour=2.0,
|
||||
ptu_effective_from=datetime.datetime(2026, 7, 1, tzinfo=datetime.timezone.utc),
|
||||
),
|
||||
)
|
||||
|
||||
def test_raising_the_rate_on_a_configured_model_is_allowed(self):
|
||||
"""The patch carries no start; the stored row supplies it."""
|
||||
merged = _merged_ptu_model_info(
|
||||
db_model=self._configured(),
|
||||
patch_data=updateDeployment(model_info=ModelInfo(id="dep-0", ptu_count=10, cost_per_ptu_per_hour=3.0)),
|
||||
)
|
||||
_validate_ptu_model_info(merged)
|
||||
assert merged["cost_per_ptu_per_hour"] == 3.0
|
||||
assert merged["ptu_effective_from"] is not None
|
||||
|
||||
def test_a_genuinely_startless_configuration_is_still_rejected(self):
|
||||
"""Merging must not become a way to smuggle PTU config in without a start."""
|
||||
bare = Deployment(model_name="gpt-4o", litellm_params=LiteLLM_Params(model="openai/gpt-4o"))
|
||||
merged = _merged_ptu_model_info(
|
||||
db_model=bare,
|
||||
patch_data=updateDeployment(
|
||||
model_info=ModelInfo(id="dep-0", team_id="t", ptu_count=10, cost_per_ptu_per_hour=2.0)
|
||||
),
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
_validate_ptu_model_info(merged)
|
||||
assert "ptu_effective_from is required" in exc.value.detail
|
||||
|
||||
def test_the_patch_still_wins_over_the_stored_value(self):
|
||||
merged = _merged_ptu_model_info(
|
||||
db_model=self._configured(),
|
||||
patch_data=updateDeployment(model_info=ModelInfo(id="dep-0", ptu_count=25)),
|
||||
)
|
||||
assert merged["ptu_count"] == 25
|
||||
|
||||
def test_an_explicit_null_clears_the_stored_field(self):
|
||||
"""update_db_model drops a PTU field a patch sends as null, so the merged view has to
|
||||
drop it too. Carrying the stored value forward validated a deployment that never
|
||||
existed."""
|
||||
merged = _merged_ptu_model_info(
|
||||
db_model=self._configured(),
|
||||
patch_data=updateDeployment(model_info=ModelInfo(id="dep-0", ptu_count=None)),
|
||||
)
|
||||
assert "ptu_count" not in merged
|
||||
|
||||
def test_clearing_one_half_of_the_pair_is_rejected(self):
|
||||
"""The write leaves a rate with no count. Merging on the stored count hid that."""
|
||||
merged = _merged_ptu_model_info(
|
||||
db_model=self._configured(),
|
||||
patch_data=updateDeployment(model_info=ModelInfo(id="dep-0", ptu_count=None)),
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
_validate_ptu_model_info(merged)
|
||||
assert "must be set together" in exc.value.detail
|
||||
|
||||
def test_clearing_the_whole_pair_is_allowed(self):
|
||||
"""Turning PTU off on a deployment is a legitimate edit."""
|
||||
merged = _merged_ptu_model_info(
|
||||
db_model=self._configured(),
|
||||
patch_data=updateDeployment(model_info=ModelInfo(id="dep-0", ptu_count=None, cost_per_ptu_per_hour=None)),
|
||||
)
|
||||
_validate_ptu_model_info(merged)
|
||||
assert "ptu_count" not in merged
|
||||
assert "cost_per_ptu_per_hour" not in merged
|
||||
|
||||
def test_an_omitted_field_is_not_a_clear(self):
|
||||
"""A partial edit that never mentions the count keeps it. Only an explicit null clears."""
|
||||
merged = _merged_ptu_model_info(
|
||||
db_model=self._configured(),
|
||||
patch_data=updateDeployment(model_info=ModelInfo(id="dep-0", cost_per_ptu_per_hour=3.0)),
|
||||
)
|
||||
assert merged["ptu_count"] == 10
|
||||
|
||||
|
||||
class TestTeamModelUpdateValidatesBeforeWriting:
|
||||
"""Drives the endpoint path itself, not the helpers. The validator sits above the team
|
||||
ACL write, which autocommits, so what it validates has to be right at that call site."""
|
||||
|
||||
@staticmethod
|
||||
async def _run(db_model, patch_data, monkeypatch, touched=None):
|
||||
import litellm.proxy.management_endpoints.model_management_endpoints as mme
|
||||
|
||||
touched = [] if touched is None else touched
|
||||
|
||||
async def _never(*args, **kwargs):
|
||||
touched.append("team_write")
|
||||
|
||||
monkeypatch.setattr(mme, "_setup_new_team_model_assignment", _never)
|
||||
monkeypatch.setattr(mme, "_update_existing_team_model_assignment", _never)
|
||||
monkeypatch.setattr(mme.ModelManagementAuthChecks, "allow_team_model_action", AsyncMock(return_value=True))
|
||||
result = await mme._update_team_model_in_db(
|
||||
db_model=db_model,
|
||||
patch_data=patch_data,
|
||||
user_api_key_dict=MagicMock(),
|
||||
prisma_client=MagicMock(),
|
||||
)
|
||||
return result, touched
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raising_the_rate_on_a_configured_model_reaches_the_write(self, monkeypatch):
|
||||
"""The patch carries no start. Validating it alone rejected this ordinary edit."""
|
||||
db_model = TestPartialPtuEditsUseTheMergedView._configured()
|
||||
patch = updateDeployment(model_info=ModelInfo(id="dep-0", team_id="t", ptu_count=10, cost_per_ptu_per_hour=3.0))
|
||||
|
||||
result, touched = await self._run(db_model, patch, monkeypatch)
|
||||
|
||||
assert touched == ["team_write"]
|
||||
assert json.loads(result["model_info"])["cost_per_ptu_per_hour"] == 3.0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_startless_configuration_is_refused_before_the_team_write(self, monkeypatch):
|
||||
"""And the refusal still lands before anything is committed."""
|
||||
bare = Deployment(model_name="gpt-4o", litellm_params=LiteLLM_Params(model="openai/gpt-4o"))
|
||||
patch = updateDeployment(model_info=ModelInfo(id="dep-0", team_id="t", ptu_count=10, cost_per_ptu_per_hour=2.0))
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await self._run(bare, patch, monkeypatch)
|
||||
|
||||
assert "ptu_effective_from is required" in exc.value.detail
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clearing_half_the_pair_is_refused_before_the_team_write(self, monkeypatch):
|
||||
"""The write drops the nulled field, so validating against the stored one let a
|
||||
deployment with a rate and no count commit."""
|
||||
db_model = TestPartialPtuEditsUseTheMergedView._configured()
|
||||
patch = updateDeployment(model_info=ModelInfo(id="dep-0", team_id="t", ptu_count=None))
|
||||
touched = []
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await self._run(db_model, patch, monkeypatch, touched)
|
||||
|
||||
assert "must be set together" in exc.value.detail
|
||||
assert touched == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clearing_the_whole_pair_reaches_the_write_and_stores_neither_field(self, monkeypatch):
|
||||
"""What the validator approved is what the write persists."""
|
||||
db_model = TestPartialPtuEditsUseTheMergedView._configured()
|
||||
patch = updateDeployment(
|
||||
model_info=ModelInfo(id="dep-0", team_id="t", ptu_count=None, cost_per_ptu_per_hour=None)
|
||||
)
|
||||
|
||||
result, touched = await self._run(db_model, patch, monkeypatch)
|
||||
|
||||
assert touched == ["team_write"]
|
||||
stored = json.loads(result["model_info"])
|
||||
assert "ptu_count" not in stored
|
||||
assert "cost_per_ptu_per_hour" not in stored
|
||||
1481
tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py
Normal file
1481
tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -812,6 +812,71 @@ async def test_add_litellm_data_to_request_strips_callback_control_fields(
|
|||
assert control_field not in snapshot_body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("timeout_field", ["timeout", "request_timeout", "stream_timeout"])
|
||||
async def test_add_litellm_data_to_request_marks_body_timeout_as_client_side(timeout_field):
|
||||
"""Router._get_timeout resolves the effective timeout from any of kwargs["timeout"],
|
||||
kwargs["request_timeout"], or kwargs["stream_timeout"], all settable directly in the
|
||||
request body. Without recognizing all three, a caller could force a 408 on every
|
||||
deployment in a fallback chain without it being flagged as caller-controlled, cooling
|
||||
down deployments other tenants rely on (see cooldown_handlers._trigger_cooldown_for_failed_deployment)."""
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data={
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
timeout_field: 0.001,
|
||||
},
|
||||
request=request_mock,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
assert updated["client_side_timeout"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_ignores_forged_client_side_timeout():
|
||||
"""The client_side_timeout marker itself must never be trusted verbatim from the
|
||||
request body: a caller forging client_side_timeout=True without a real timeout
|
||||
override could dodge cooldown protection on an actual deployment failure."""
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data={
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"client_side_timeout": True,
|
||||
},
|
||||
request=request_mock,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
assert not updated.get("client_side_timeout")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_allows_client_mock_response_with_admin_opt_in():
|
||||
request_mock = MagicMock(spec=Request)
|
||||
|
|
@ -2713,6 +2778,149 @@ def test_get_chain_id_from_headers_generic_vendor_session_id():
|
|||
)
|
||||
|
||||
|
||||
def test_trace_id_from_traceparent_valid():
|
||||
from litellm.proxy.litellm_pre_call_utils import _trace_id_from_traceparent
|
||||
|
||||
assert (
|
||||
_trace_id_from_traceparent("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01")
|
||||
== "4bf92f3577b34da6a3ce929d0e0e4736"
|
||||
)
|
||||
# Case-insensitive, normalized to lowercase
|
||||
assert (
|
||||
_trace_id_from_traceparent("00-4BF92F3577B34DA6A3CE929D0E0E4736-00f067aa0ba902b7-01")
|
||||
== "4bf92f3577b34da6a3ce929d0e0e4736"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"traceparent",
|
||||
[
|
||||
"not-a-traceparent",
|
||||
"00-tooshort-00f067aa0ba902b7-01",
|
||||
"00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7", # missing flags segment
|
||||
"00-4bf92f3577g34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", # non-hex char
|
||||
"00-00000000000000000000000000000000-00f067aa0ba902b7-01", # all-zero trace-id, invalid per spec
|
||||
"",
|
||||
],
|
||||
)
|
||||
def test_trace_id_from_traceparent_rejects_malformed(traceparent: str):
|
||||
from litellm.proxy.litellm_pre_call_utils import _trace_id_from_traceparent
|
||||
|
||||
assert _trace_id_from_traceparent(traceparent) is None
|
||||
|
||||
|
||||
def test_session_id_from_baggage_valid():
|
||||
from litellm.proxy.litellm_pre_call_utils import _session_id_from_baggage
|
||||
|
||||
assert _session_id_from_baggage("session.id=abc-123,user.id=42") == "abc-123"
|
||||
assert _session_id_from_baggage("user.id=42, session.id=xyz-789") == "xyz-789"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"baggage",
|
||||
[
|
||||
"user.id=42",
|
||||
"",
|
||||
"session.id=",
|
||||
],
|
||||
)
|
||||
def test_session_id_from_baggage_absent_or_empty(baggage: str):
|
||||
from litellm.proxy.litellm_pre_call_utils import _session_id_from_baggage
|
||||
|
||||
assert _session_id_from_baggage(baggage) is None
|
||||
|
||||
|
||||
def test_add_litellm_metadata_from_request_headers_traceparent_sets_trace_id_only():
|
||||
"""A bare traceparent header (no litellm-specific headers) sets litellm_trace_id
|
||||
from its trace-id component and leaves litellm_session_id unset."""
|
||||
headers = {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"}
|
||||
data = {"metadata": {}}
|
||||
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
|
||||
headers=headers, data=data, _metadata_variable_name="metadata"
|
||||
)
|
||||
assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
|
||||
assert data["metadata"]["trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
|
||||
assert "litellm_session_id" not in data
|
||||
|
||||
|
||||
def test_add_litellm_metadata_from_request_headers_baggage_sets_session_id_only():
|
||||
"""A bare baggage header (no litellm-specific headers) sets litellm_session_id
|
||||
from its session.id entry and leaves litellm_trace_id unset."""
|
||||
headers = {"baggage": "session.id=baggage-session-42,user.id=7"}
|
||||
data = {"metadata": {}}
|
||||
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
|
||||
headers=headers, data=data, _metadata_variable_name="metadata"
|
||||
)
|
||||
assert data["litellm_session_id"] == "baggage-session-42"
|
||||
assert data["metadata"]["session_id"] == "baggage-session-42"
|
||||
assert "litellm_trace_id" not in data
|
||||
|
||||
|
||||
def test_add_litellm_metadata_from_request_headers_baggage_session_id_not_logged_raw(caplog):
|
||||
"""The raw baggage session.id value must never reach the debug log line -
|
||||
it isn't sanitized until set_session_id() runs much later in
|
||||
Logging.__init__(), so logging it here would let a caller with control
|
||||
characters or terminal escape sequences forge plaintext log output."""
|
||||
import logging
|
||||
|
||||
poisoned = "poisoned\x1b[31mFAKE_RED_TEXT\x1b[0m"
|
||||
headers = {"baggage": f"session.id={poisoned}"}
|
||||
data = {"metadata": {}}
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
|
||||
headers=headers, data=data, _metadata_variable_name="metadata"
|
||||
)
|
||||
assert data["litellm_session_id"] == poisoned
|
||||
assert not any(poisoned in record.getMessage() for record in caplog.records)
|
||||
|
||||
|
||||
def test_add_litellm_metadata_from_request_headers_traceparent_and_baggage_together():
|
||||
"""traceparent and baggage are resolved independently - trace_id and
|
||||
session_id do not have to be the same value, unlike the chain_id path."""
|
||||
headers = {
|
||||
"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
|
||||
"baggage": "session.id=baggage-session-42",
|
||||
}
|
||||
data = {"metadata": {}}
|
||||
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
|
||||
headers=headers, data=data, _metadata_variable_name="metadata"
|
||||
)
|
||||
assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
|
||||
assert data["litellm_session_id"] == "baggage-session-42"
|
||||
|
||||
|
||||
def test_add_litellm_metadata_from_request_headers_explicit_trace_id_beats_traceparent():
|
||||
"""x-litellm-trace-id must win over a traceparent header carrying a
|
||||
different trace-id - explicit litellm headers are always highest priority."""
|
||||
headers = {
|
||||
"x-litellm-trace-id": "explicit-trace-id-value",
|
||||
"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
|
||||
}
|
||||
data = {"metadata": {}}
|
||||
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
|
||||
headers=headers, data=data, _metadata_variable_name="metadata"
|
||||
)
|
||||
assert data["litellm_trace_id"] == "explicit-trace-id-value"
|
||||
assert data["litellm_session_id"] == "explicit-trace-id-value"
|
||||
|
||||
|
||||
def test_add_litellm_metadata_from_request_headers_anthropic_metadata_beats_baggage():
|
||||
"""The existing Anthropic metadata.user_id session_id path must win over a
|
||||
baggage session.id fallback."""
|
||||
data = {
|
||||
"metadata": {
|
||||
"user_id": "user_abc123_account__session_e96634a3-fa28-4083-b354-55542e2dca01",
|
||||
}
|
||||
}
|
||||
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
|
||||
headers={"baggage": "session.id=baggage-session-42"},
|
||||
data=data,
|
||||
_metadata_variable_name="metadata",
|
||||
)
|
||||
assert data["litellm_session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01"
|
||||
assert "litellm_trace_id" not in data
|
||||
|
||||
|
||||
def test_get_internal_user_header_from_mapping_returns_expected_header():
|
||||
mappings = [
|
||||
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"},
|
||||
|
|
|
|||
|
|
@ -11278,3 +11278,40 @@ async def test_setup_prisma_client_returns_none_when_connect_itself_fails(monkey
|
|||
assert result is None
|
||||
assert mock_client.start_db_health_watchdog_task.await_count == 0
|
||||
assert mock_client.health_check.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ptu_rollup_job_registered_at_startup(monkeypatch):
|
||||
"""The PTU rollup cron is registered at startup; only models with PTU config accrue flat cost (asserted in test_ptu_flat_cost_rollup.py)."""
|
||||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import (
|
||||
PTU_ROLLUP_JOB_ID,
|
||||
)
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
|
||||
|
||||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||||
mock_proxy_config = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
|
||||
):
|
||||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||||
general_settings={},
|
||||
prisma_client=mock_prisma_client,
|
||||
proxy_budget_rescheduler_min_time=1,
|
||||
proxy_budget_rescheduler_max_time=2,
|
||||
proxy_batch_write_at=5,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
assert ps.scheduler is not None
|
||||
assert ps.scheduler.get_job(PTU_ROLLUP_JOB_ID) is not None
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ Unit tests for CooldownCache exception masking functionality
|
|||
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -255,3 +256,66 @@ class TestCooldownCacheExceptionMasking:
|
|||
# Should show first 50 characters, then all asterisks
|
||||
expected = "A" * 50 + "*" * 50
|
||||
assert masked == expected
|
||||
|
||||
|
||||
class TestCorrectedActiveCooldown:
|
||||
def _make_cooldown_cache(self) -> CooldownCache:
|
||||
in_memory = InMemoryCache()
|
||||
dual_cache = DualCache(in_memory_cache=in_memory)
|
||||
return CooldownCache(cache=dual_cache, default_cooldown_time=60.0)
|
||||
|
||||
def _entry(self, timestamp: float, cooldown_time: float) -> CooldownCacheValue:
|
||||
return CooldownCacheValue(
|
||||
exception_received="Rate limit",
|
||||
status_code="429",
|
||||
timestamp=timestamp,
|
||||
cooldown_time=cooldown_time,
|
||||
)
|
||||
|
||||
def test_expired_entry_returns_none_and_evicts(self):
|
||||
cc = self._make_cooldown_cache()
|
||||
key = "deployment:expired-dep:cooldown"
|
||||
entry = self._entry(timestamp=time.time() - 120.0, cooldown_time=60.0)
|
||||
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=600)
|
||||
|
||||
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
|
||||
|
||||
assert result is None
|
||||
assert cc.cache.in_memory_cache.get_cache(key) is None
|
||||
|
||||
def test_active_entry_within_window_returns_value(self):
|
||||
cc = self._make_cooldown_cache()
|
||||
key = "deployment:active-dep:cooldown"
|
||||
entry = self._entry(timestamp=time.time(), cooldown_time=60.0)
|
||||
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=60)
|
||||
|
||||
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
|
||||
|
||||
assert result is not None
|
||||
assert result["status_code"] == "429"
|
||||
|
||||
def test_inflated_ttl_is_corrected(self):
|
||||
cc = self._make_cooldown_cache()
|
||||
key = "deployment:backfilled-dep:cooldown"
|
||||
remaining = 30.0
|
||||
entry = self._entry(timestamp=time.time() - (60.0 - remaining), cooldown_time=60.0)
|
||||
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=600)
|
||||
|
||||
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
|
||||
|
||||
assert result is not None
|
||||
corrected_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
assert corrected_expiry is not None
|
||||
assert corrected_expiry - time.time() <= 60.0
|
||||
|
||||
def test_normal_ttl_not_modified(self):
|
||||
cc = self._make_cooldown_cache()
|
||||
key = "deployment:normal-dep:cooldown"
|
||||
entry = self._entry(timestamp=time.time(), cooldown_time=60.0)
|
||||
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=60)
|
||||
original_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
|
||||
cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
|
||||
|
||||
after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
assert after_expiry == original_expiry
|
||||
|
|
|
|||
298
tests/test_litellm/router_utils/test_cooldown_handlers.py
Normal file
298
tests/test_litellm/router_utils/test_cooldown_handlers.py
Normal file
|
|
@ -0,0 +1,298 @@
|
|||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import litellm
|
||||
from litellm.router_utils.cooldown_handlers import (
|
||||
_get_deployment_cooldown_policy,
|
||||
_resolve_allowed_fails_from_policy,
|
||||
_should_cooldown_based_on_deployment_policy,
|
||||
should_cooldown_based_on_allowed_fails_policy,
|
||||
)
|
||||
|
||||
|
||||
class TestGetDeploymentCooldownPolicy:
|
||||
def _make_router(self, deployment_id: str, model_info: dict | None = None):
|
||||
router = MagicMock()
|
||||
if model_info is None:
|
||||
router.get_model_info.return_value = None
|
||||
else:
|
||||
router.get_model_info.return_value = {"model_info": model_info}
|
||||
return router
|
||||
|
||||
def test_deployment_not_found_returns_none_none(self):
|
||||
router = self._make_router("dep-1")
|
||||
policy, allowed = _get_deployment_cooldown_policy(router, "dep-1")
|
||||
assert policy is None
|
||||
assert allowed is None
|
||||
|
||||
def test_no_model_info_returns_none_none(self):
|
||||
router = MagicMock()
|
||||
router.get_model_info.return_value = {"model_info": {}}
|
||||
policy, allowed = _get_deployment_cooldown_policy(router, "dep-1")
|
||||
assert policy is None
|
||||
assert allowed is None
|
||||
|
||||
def test_returns_policy_dict_and_allowed_fails(self):
|
||||
router = self._make_router(
|
||||
"dep-1",
|
||||
{"allowed_fails_policy": {"RateLimitErrorAllowedFails": 2}, "allowed_fails": 3},
|
||||
)
|
||||
policy, allowed = _get_deployment_cooldown_policy(router, "dep-1")
|
||||
assert policy == {"RateLimitErrorAllowedFails": 2}
|
||||
assert allowed == 3
|
||||
|
||||
def test_non_dict_policy_treated_as_none(self):
|
||||
router = self._make_router("dep-1", {"allowed_fails_policy": "invalid", "allowed_fails": 5})
|
||||
policy, allowed = _get_deployment_cooldown_policy(router, "dep-1")
|
||||
assert policy is None
|
||||
assert allowed == 5
|
||||
|
||||
def test_allowed_fails_only(self):
|
||||
router = self._make_router("dep-1", {"allowed_fails": 1})
|
||||
policy, allowed = _get_deployment_cooldown_policy(router, "dep-1")
|
||||
assert policy is None
|
||||
assert allowed == 1
|
||||
|
||||
|
||||
class TestResolveAllowedFailsFromPolicy:
|
||||
def test_none_policy_returns_none(self):
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
assert _resolve_allowed_fails_from_policy(None, exc) is None
|
||||
|
||||
def test_matching_rate_limit_error(self):
|
||||
policy = {"RateLimitErrorAllowedFails": 3}
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
assert _resolve_allowed_fails_from_policy(policy, exc) == 3
|
||||
|
||||
def test_matching_internal_server_error(self):
|
||||
policy = {"InternalServerErrorAllowedFails": 5}
|
||||
exc = litellm.InternalServerError("500", "openai", "gpt-4")
|
||||
assert _resolve_allowed_fails_from_policy(policy, exc) == 5
|
||||
|
||||
def test_matching_service_unavailable_error(self):
|
||||
policy = {"ServiceUnavailableErrorAllowedFails": 4}
|
||||
exc = litellm.ServiceUnavailableError("503", "openai", "gpt-4")
|
||||
assert _resolve_allowed_fails_from_policy(policy, exc) == 4
|
||||
|
||||
def test_matching_bad_gateway_error(self):
|
||||
policy = {"BadGatewayErrorAllowedFails": 2}
|
||||
exc = litellm.BadGatewayError("502", "openai", "gpt-4")
|
||||
assert _resolve_allowed_fails_from_policy(policy, exc) == 2
|
||||
|
||||
def test_matching_not_found_error(self):
|
||||
policy = {"NotFoundErrorAllowedFails": 1}
|
||||
exc = litellm.NotFoundError("404", "openai", "gpt-4")
|
||||
assert _resolve_allowed_fails_from_policy(policy, exc) == 1
|
||||
|
||||
def test_unmatched_exception_returns_none(self):
|
||||
policy = {"RateLimitErrorAllowedFails": 3}
|
||||
exc = litellm.InternalServerError("500", "openai", "gpt-4")
|
||||
assert _resolve_allowed_fails_from_policy(policy, exc) is None
|
||||
|
||||
def test_field_absent_from_policy_returns_none(self):
|
||||
policy: dict[str, int] = {}
|
||||
exc = litellm.InternalServerError("500", "openai", "gpt-4")
|
||||
assert _resolve_allowed_fails_from_policy(policy, exc) is None
|
||||
|
||||
def test_content_policy_violation_not_shadowed_by_bad_request_error(self):
|
||||
"""ContentPolicyViolationError subclasses BadRequestError, so if
|
||||
BadRequestError were checked first, this would incorrectly resolve to
|
||||
BadRequestErrorAllowedFails (10) instead of
|
||||
ContentPolicyViolationErrorAllowedFails (2)."""
|
||||
policy = {"BadRequestErrorAllowedFails": 10, "ContentPolicyViolationErrorAllowedFails": 2}
|
||||
exc = litellm.ContentPolicyViolationError("flagged content", "openai", "gpt-4")
|
||||
assert _resolve_allowed_fails_from_policy(policy, exc) == 2
|
||||
|
||||
|
||||
class TestShouldCooldownBasedOnDeploymentPolicy:
|
||||
def _make_router(self, model_info: dict | None = None):
|
||||
router = MagicMock()
|
||||
if model_info is None:
|
||||
router.get_model_info.return_value = None
|
||||
else:
|
||||
router.get_model_info.return_value = model_info
|
||||
return router
|
||||
|
||||
def test_policy_match_uses_exception_type_as_cache_key_suffix(self):
|
||||
policy = {"RateLimitErrorAllowedFails": 0}
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
router = self._make_router({"litellm_params": {}, "model_info": {}})
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
mock_sc.return_value = True
|
||||
result = _should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, policy, None, is_single_deployment_model_group=False
|
||||
)
|
||||
|
||||
assert result is True
|
||||
call_kwargs = mock_sc.call_args[1]
|
||||
assert call_kwargs["allowed_fails_override"] == 0
|
||||
assert call_kwargs["cache_key_suffix"] == "RateLimitError"
|
||||
|
||||
def test_no_policy_match_uses_dep_allowed_fails_and_generic_suffix(self):
|
||||
policy: dict[str, int] = {}
|
||||
exc = litellm.InternalServerError("500", "openai", "gpt-4")
|
||||
router = self._make_router({"litellm_params": {}, "model_info": {}})
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
mock_sc.return_value = False
|
||||
result = _should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, policy, dep_allowed_fails=3, is_single_deployment_model_group=False
|
||||
)
|
||||
|
||||
assert result is False
|
||||
call_kwargs = mock_sc.call_args[1]
|
||||
assert call_kwargs["allowed_fails_override"] == 3
|
||||
assert call_kwargs["cache_key_suffix"] == "generic"
|
||||
|
||||
def test_dep_allowed_fails_on_single_deployment_group_does_not_cooldown(self):
|
||||
"""A generic, deployment-wide allowed_fails predates the per-exception-type
|
||||
policy and is a less deliberate opt-in, so on a single-deployment model group
|
||||
it must not silently disable the "avoid cooldowns on single deployment model
|
||||
groups" safety net."""
|
||||
exc = litellm.InternalServerError("500", "openai", "gpt-4")
|
||||
router = self._make_router({"litellm_params": {}, "model_info": {}})
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
result = _should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, None, dep_allowed_fails=3, is_single_deployment_model_group=True
|
||||
)
|
||||
|
||||
assert result is False
|
||||
mock_sc.assert_not_called()
|
||||
|
||||
def test_named_policy_on_single_deployment_group_still_cools_down(self):
|
||||
"""Unlike a generic allowed_fails, an explicit per-exception-type policy entry
|
||||
is a deliberate opt-in and must still apply on a single-deployment group."""
|
||||
policy = {"RateLimitErrorAllowedFails": 0}
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
router = self._make_router({"litellm_params": {}, "model_info": {}})
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
mock_sc.return_value = True
|
||||
result = _should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, policy, None, is_single_deployment_model_group=True
|
||||
)
|
||||
|
||||
assert result is True
|
||||
mock_sc.assert_called_once()
|
||||
|
||||
def test_no_policy_and_no_dep_allowed_fails_defers_to_router_level(self):
|
||||
"""When neither a deployment policy nor a deployment-wide allowed_fails covers
|
||||
this exception, defer to router-level behavior instead of forcing an
|
||||
immediate cooldown (allowed_fails_override=0 would trip on the first failure
|
||||
of any exception type the deployment's config doesn't mention)."""
|
||||
exc = litellm.InternalServerError("500", "openai", "gpt-4")
|
||||
router = self._make_router({"litellm_params": {}, "model_info": {}})
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
mock_sc.return_value = True
|
||||
_should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, None, None, is_single_deployment_model_group=False
|
||||
)
|
||||
|
||||
call_kwargs = mock_sc.call_args[1]
|
||||
assert call_kwargs["allowed_fails_override"] is None
|
||||
assert call_kwargs["cache_key_suffix"] is None
|
||||
|
||||
def test_partial_policy_without_dep_allowed_fails_defers_for_uncovered_exception(self):
|
||||
"""A deployment that only sets RateLimitErrorAllowedFails must not force a
|
||||
zero-fail threshold on an unrelated TimeoutError; it should defer to
|
||||
router-level behavior for exception types its policy doesn't mention."""
|
||||
policy = {"RateLimitErrorAllowedFails": 0}
|
||||
exc = litellm.Timeout("timed out", "openai", "gpt-4")
|
||||
router = self._make_router({"litellm_params": {}, "model_info": {}})
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
mock_sc.return_value = False
|
||||
_should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, policy, dep_allowed_fails=None, is_single_deployment_model_group=False
|
||||
)
|
||||
|
||||
call_kwargs = mock_sc.call_args[1]
|
||||
assert call_kwargs["allowed_fails_override"] is None
|
||||
assert call_kwargs["cache_key_suffix"] is None
|
||||
|
||||
def test_cooldown_time_from_model_info_passed_through(self):
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
router = self._make_router({"litellm_params": {}, "model_info": {"cooldown_time": 120.0}})
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
mock_sc.return_value = True
|
||||
_should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, None, None, is_single_deployment_model_group=False
|
||||
)
|
||||
|
||||
call_kwargs = mock_sc.call_args[1]
|
||||
assert call_kwargs["cooldown_time_override"] == 120.0
|
||||
|
||||
def test_cooldown_time_from_litellm_params_used_as_fallback(self):
|
||||
"""cooldown_time has pre-existing litellm_params support on the primary
|
||||
failure path, so it must still be honored here when model_info doesn't
|
||||
set it."""
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
router = self._make_router({"litellm_params": {"cooldown_time": 120.0}, "model_info": {}})
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
mock_sc.return_value = True
|
||||
_should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, None, None, is_single_deployment_model_group=False
|
||||
)
|
||||
|
||||
call_kwargs = mock_sc.call_args[1]
|
||||
assert call_kwargs["cooldown_time_override"] == 120.0
|
||||
|
||||
def test_cooldown_time_from_model_info_takes_priority_over_litellm_params(self):
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
router = self._make_router({"litellm_params": {"cooldown_time": 120.0}, "model_info": {"cooldown_time": 15.0}})
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
mock_sc.return_value = True
|
||||
_should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, None, None, is_single_deployment_model_group=False
|
||||
)
|
||||
|
||||
call_kwargs = mock_sc.call_args[1]
|
||||
assert call_kwargs["cooldown_time_override"] == 15.0
|
||||
|
||||
def test_model_info_none_passes_none_cooldown_time(self):
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
router = self._make_router(None)
|
||||
|
||||
with patch("litellm.router_utils.cooldown_handlers.should_cooldown_based_on_allowed_fails_policy") as mock_sc:
|
||||
mock_sc.return_value = False
|
||||
_should_cooldown_based_on_deployment_policy(
|
||||
router, "dep-1", exc, None, None, is_single_deployment_model_group=False
|
||||
)
|
||||
|
||||
call_kwargs = mock_sc.call_args[1]
|
||||
assert call_kwargs["cooldown_time_override"] is None
|
||||
|
||||
|
||||
class TestShouldCooldownBasedOnAllowedFailsPolicy:
|
||||
def _make_router(self, cooldown_time: float = 60.0) -> MagicMock:
|
||||
router = MagicMock()
|
||||
router.cooldown_time = cooldown_time
|
||||
router.allowed_fails = 0
|
||||
router.allowed_fails_policy = None
|
||||
router.get_allowed_fails_from_policy.return_value = None
|
||||
router.failed_calls.get_cache.return_value = None
|
||||
return router
|
||||
|
||||
def test_cooldown_time_override_zero_is_not_falsy(self):
|
||||
"""cooldown_time_override=0 must be honored; it must not fall through to the router-level value."""
|
||||
router = self._make_router(cooldown_time=60.0)
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
|
||||
should_cooldown_based_on_allowed_fails_policy(
|
||||
litellm_router_instance=router,
|
||||
deployment="dep-1",
|
||||
original_exception=exc,
|
||||
allowed_fails_override=5,
|
||||
cooldown_time_override=0.0,
|
||||
)
|
||||
|
||||
set_cache_call = router.failed_calls.set_cache.call_args
|
||||
assert set_cache_call is not None
|
||||
assert set_cache_call[1]["ttl"] == 0.0, (
|
||||
"cooldown_time_override=0 should be used as TTL, not the router-level 60.0"
|
||||
)
|
||||
|
|
@ -1,9 +1,12 @@
|
|||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
AttemptedFallbackTargets,
|
||||
_trigger_cooldown_for_failed_deployment,
|
||||
fallback_attempt_key,
|
||||
get_fallback_model_group,
|
||||
run_async_fallback,
|
||||
|
|
@ -144,6 +147,452 @@ async def test_run_async_fallback_skips_original_model_group():
|
|||
assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1
|
||||
|
||||
|
||||
def test_trigger_cooldown_calls_set_cooldown_when_deployment_id_present():
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = None
|
||||
|
||||
exc = RuntimeError("upstream error")
|
||||
exc.status_code = 429
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc)
|
||||
|
||||
mock_set.assert_called_once()
|
||||
_, call_kwargs = mock_set.call_args
|
||||
assert call_kwargs["deployment"] == "deployment-abc"
|
||||
assert call_kwargs["exception_status"] == 429
|
||||
|
||||
|
||||
def test_trigger_cooldown_skips_when_no_deployment_id():
|
||||
router = MagicMock()
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=RuntimeError("err"))
|
||||
|
||||
mock_set.assert_not_called()
|
||||
|
||||
|
||||
def test_trigger_cooldown_does_not_trust_caller_supplied_metadata_bucket():
|
||||
"""A metadata bucket can't reliably be told apart from a caller-supplied one
|
||||
without knowing the call's function_name, so a client with permission to set
|
||||
metadata must not be able to get an arbitrary deployment cooled down by
|
||||
forging a deployment_model_name marker."""
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = None
|
||||
|
||||
exc = RuntimeError("err")
|
||||
kwargs = {"metadata": {"model_info": {"id": "attacker-chosen-deployment"}, "deployment_model_name": "gpt-4"}}
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs=kwargs, exception=exc)
|
||||
|
||||
mock_set.assert_not_called()
|
||||
|
||||
|
||||
def test_trigger_cooldown_increments_failure_counter_before_cooldown_check():
|
||||
"""The fallback path must feed the same per-minute failure counter the
|
||||
primary path uses, or repeated fallback failures never accumulate toward
|
||||
the default percent-fail-rate cooldown threshold."""
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = None
|
||||
|
||||
exc = RuntimeError("err")
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
|
||||
with (
|
||||
patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set,
|
||||
patch(
|
||||
"litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute"
|
||||
) as mock_increment,
|
||||
):
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc)
|
||||
|
||||
mock_increment.assert_called_once_with(litellm_router_instance=router, deployment_id="deployment-abc")
|
||||
mock_set.assert_called_once()
|
||||
|
||||
|
||||
def test_trigger_cooldown_uses_deployment_cooldown_time_when_present():
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = {"model_info": {"cooldown_time": 30}}
|
||||
|
||||
exc = RuntimeError("upstream error")
|
||||
exc.status_code = 429
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc)
|
||||
|
||||
_, call_kwargs = mock_set.call_args
|
||||
assert call_kwargs["time_to_cooldown"] == 30
|
||||
|
||||
|
||||
def test_trigger_cooldown_falls_back_to_litellm_params_cooldown_time():
|
||||
"""cooldown_time has pre-existing litellm_params support on the primary
|
||||
failure path, so it must still be honored here when model_info doesn't set
|
||||
it, unlike the new allowed_fails/allowed_fails_policy fields."""
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = {"litellm_params": {"cooldown_time": 30}}
|
||||
|
||||
exc = RuntimeError("upstream error")
|
||||
exc.status_code = 429
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc)
|
||||
|
||||
_, call_kwargs = mock_set.call_args
|
||||
assert call_kwargs["time_to_cooldown"] == 30
|
||||
|
||||
|
||||
def test_trigger_cooldown_uses_response_header_when_no_deployment_config():
|
||||
"""Precedence must match Router.deployment_callback_on_failure's primary path:
|
||||
deployment config, then the response's Retry-After header, then the router
|
||||
default. Without this, the fallback path always skips straight to the router
|
||||
default whenever no deployment-level cooldown_time is configured."""
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = {"model_info": {}}
|
||||
|
||||
exc = RuntimeError("upstream error")
|
||||
exc.status_code = 429
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
exc.litellm_response_headers = httpx.Headers({"retry-after": "45"})
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc)
|
||||
|
||||
_, call_kwargs = mock_set.call_args
|
||||
assert call_kwargs["time_to_cooldown"] == 45
|
||||
|
||||
|
||||
def test_trigger_cooldown_silently_catches_exceptions():
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = None
|
||||
|
||||
exc = RuntimeError("upstream error")
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
|
||||
with patch(
|
||||
"litellm.router_utils.fallback_event_handlers._set_cooldown_deployments",
|
||||
side_effect=RuntimeError("cooldown error"),
|
||||
):
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc)
|
||||
|
||||
|
||||
def test_trigger_cooldown_skips_request_scoped_404_on_generic_api_call():
|
||||
"""A generic API call (files/batches/threads/rerank/...) forwards a caller-supplied
|
||||
resource id, so a 404 there means "that id doesn't exist", not "this deployment is
|
||||
unhealthy". Without this guard, a single bad id would 404 every deployment in the
|
||||
fallback chain and cool all of them down from one request."""
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = None
|
||||
|
||||
exc = RuntimeError("not found")
|
||||
exc.status_code = 404
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
|
||||
with (
|
||||
patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set,
|
||||
patch(
|
||||
"litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute"
|
||||
) as mock_increment,
|
||||
):
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=router,
|
||||
kwargs={"original_generic_function": MagicMock()},
|
||||
exception=exc,
|
||||
)
|
||||
|
||||
mock_set.assert_not_called()
|
||||
mock_increment.assert_not_called()
|
||||
|
||||
|
||||
def test_trigger_cooldown_still_cools_down_404_outside_generic_api_call():
|
||||
"""The request-scoped-404 guard is scoped to generic API calls only: a 404 on a
|
||||
regular completion fallback (no original_generic_function in kwargs) must still
|
||||
cool down the deployment as before."""
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = None
|
||||
|
||||
exc = RuntimeError("not found")
|
||||
exc.status_code = 404
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc)
|
||||
|
||||
mock_set.assert_called_once()
|
||||
|
||||
|
||||
def test_trigger_cooldown_skips_client_side_timeout_408():
|
||||
"""The proxy's x-litellm-timeout header lets a caller set an arbitrarily short
|
||||
timeout, which litellm.Timeout reports as status 408 regardless of the
|
||||
deployment's actual health. Without this guard, a caller could force a 408 on
|
||||
every deployment in the fallback chain from a single request."""
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = None
|
||||
|
||||
exc = RuntimeError("timeout")
|
||||
exc.status_code = 408
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
|
||||
with (
|
||||
patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set,
|
||||
patch(
|
||||
"litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute"
|
||||
) as mock_increment,
|
||||
):
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=router,
|
||||
kwargs={"client_side_timeout": True},
|
||||
exception=exc,
|
||||
)
|
||||
|
||||
mock_set.assert_not_called()
|
||||
mock_increment.assert_not_called()
|
||||
|
||||
|
||||
def test_trigger_cooldown_still_cools_down_408_without_client_side_timeout_flag():
|
||||
"""The client-side-timeout guard is scoped to caller-supplied timeouts only: a
|
||||
408 that did not come from x-litellm-timeout (no client_side_timeout in kwargs)
|
||||
must still cool down the deployment as before."""
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = None
|
||||
|
||||
exc = RuntimeError("timeout")
|
||||
exc.status_code = 408
|
||||
exc.failed_deployment_id = "deployment-abc"
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
_trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc)
|
||||
|
||||
mock_set.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_triggers_cooldown_when_logging_obj_has_logged():
|
||||
router = MagicMock()
|
||||
router.cooldown_time = 60
|
||||
router.get_model_info.return_value = None
|
||||
router.log_retry = MagicMock(side_effect=lambda kwargs, e: kwargs)
|
||||
|
||||
exc = RuntimeError("fallback failed")
|
||||
exc.failed_deployment_id = "dep-xyz"
|
||||
|
||||
async def _always_fail(*args, **kwargs):
|
||||
raise exc
|
||||
|
||||
router.async_function_with_fallbacks = _always_fail
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {"has_logged_async_failure": True}
|
||||
|
||||
kwargs = {
|
||||
"litellm_logging_obj": logging_obj,
|
||||
}
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
with pytest.raises(RuntimeError):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["fallback-model"],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
mock_set.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_skips_cooldown_when_logging_obj_not_logged():
|
||||
router = MagicMock()
|
||||
router.log_retry = MagicMock(side_effect=lambda kwargs, e: kwargs)
|
||||
|
||||
exc = RuntimeError("fallback failed")
|
||||
exc.failed_deployment_id = "dep-xyz"
|
||||
|
||||
async def _always_fail(*args, **kwargs):
|
||||
raise exc
|
||||
|
||||
router.async_function_with_fallbacks = _always_fail
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {"has_logged_async_failure": False}
|
||||
|
||||
kwargs = {
|
||||
"litellm_logging_obj": logging_obj,
|
||||
}
|
||||
|
||||
with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set:
|
||||
with pytest.raises(RuntimeError):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["fallback-model"],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
mock_set.assert_not_called()
|
||||
|
||||
|
||||
class AttemptRecordingRouter:
|
||||
def __init__(self):
|
||||
self.attempted_model_groups = []
|
||||
self.received_kwargs = None
|
||||
|
||||
def log_retry(self, kwargs, e):
|
||||
return kwargs
|
||||
|
||||
async def async_function_with_fallbacks(self, *args, **kwargs):
|
||||
self.attempted_model_groups.append(kwargs.get("model"))
|
||||
self.received_kwargs = kwargs
|
||||
return StreamingWrapper()
|
||||
|
||||
|
||||
async def _acreate_batch(*args, **kwargs):
|
||||
raise AssertionError("only used for its __name__")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_keeps_uploaded_file_requests_in_their_model_group():
|
||||
"""An input_file_id only exists under the credentials of the group it was uploaded
|
||||
to, so a cross-group fallback can only fail with the wrong provider's error."""
|
||||
router = AttemptRecordingRouter()
|
||||
owning_provider_error = RuntimeError("openai connection error")
|
||||
|
||||
with pytest.raises(RuntimeError, match="openai connection error"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["azure-group"],
|
||||
original_model_group="openai-group",
|
||||
original_exception=owning_provider_error,
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="openai-group",
|
||||
input_file_id="file-owned-by-openai",
|
||||
original_function=_acreate_batch,
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_keeps_fine_tuning_requests_in_their_model_group():
|
||||
router = AttemptRecordingRouter()
|
||||
|
||||
with pytest.raises(RuntimeError, match="openai connection error"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["azure-group"],
|
||||
original_model_group="openai-group",
|
||||
original_exception=RuntimeError("openai connection error"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="openai-group",
|
||||
training_file="file-owned-by-openai",
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_allows_same_model_group_retry_for_uploaded_file_requests():
|
||||
"""Order-based fallbacks stay inside the owning group, so they must still run."""
|
||||
router = AttemptRecordingRouter()
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=[{"model": "openai-group", "_target_order": 2}],
|
||||
original_model_group="openai-group",
|
||||
original_exception=RuntimeError("first deployment failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="openai-group",
|
||||
input_file_id="file-owned-by-openai",
|
||||
original_function=_acreate_batch,
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == ["openai-group"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_still_crosses_model_groups_without_an_uploaded_file():
|
||||
router = AttemptRecordingRouter()
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["azure-group"],
|
||||
original_model_group="openai-group",
|
||||
original_exception=RuntimeError("openai connection error"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="openai-group",
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == ["azure-group"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_handles_explicitly_none_metadata():
|
||||
"""/v1/batches always sets `metadata`, and sets it to None when the caller sent
|
||||
none, so setdefault() on it hands back None instead of a dict."""
|
||||
router = AttemptRecordingRouter()
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["azure-group"],
|
||||
original_model_group="openai-group",
|
||||
original_exception=RuntimeError("openai connection error"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="openai-group",
|
||||
metadata=None,
|
||||
)
|
||||
|
||||
assert router.received_kwargs["metadata"] == {"model_group": "azure-group"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_records_batch_model_group_outside_provider_metadata():
|
||||
"""`metadata` on a batch request is forwarded to the provider and stored on the
|
||||
batch, so the router's own model_group belongs in litellm_metadata."""
|
||||
router = AttemptRecordingRouter()
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=[{"model": "openai-group", "_target_order": 2}],
|
||||
original_model_group="openai-group",
|
||||
original_exception=RuntimeError("first deployment failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="openai-group",
|
||||
input_file_id="file-owned-by-openai",
|
||||
metadata={"caller": "nightly-job"},
|
||||
litellm_metadata={"model_group": "openai-group"},
|
||||
original_function=_acreate_batch,
|
||||
)
|
||||
|
||||
assert router.received_kwargs["metadata"] == {"caller": "nightly-job"}
|
||||
assert router.received_kwargs["litellm_metadata"]["model_group"] == "openai-group"
|
||||
|
||||
|
||||
class RecordingFailRouter:
|
||||
def __init__(self):
|
||||
self.attempted_models = []
|
||||
|
|
@ -351,9 +800,7 @@ def test_get_fallback_model_group_does_not_mutate_fallbacks():
|
|||
fallbacks list, which is the live router config shared across requests."""
|
||||
fallbacks = [{"gpt-3.5-turbo": ["claude-3-haiku"]}, "gpt-4o-mini"]
|
||||
|
||||
fallback_model_group, _ = get_fallback_model_group(
|
||||
fallbacks=fallbacks, model_group="unmatched-model"
|
||||
)
|
||||
fallback_model_group, _ = get_fallback_model_group(fallbacks=fallbacks, model_group="unmatched-model")
|
||||
|
||||
assert fallback_model_group == ["gpt-4o-mini"]
|
||||
assert fallbacks == [{"gpt-3.5-turbo": ["claude-3-haiku"]}, "gpt-4o-mini"]
|
||||
|
|
|
|||
|
|
@ -17,9 +17,15 @@ import sys
|
|||
import litellm
|
||||
from litellm._logging import (
|
||||
ALL_LOGGERS,
|
||||
CorrelationContextFilter,
|
||||
CorrelationPlainFormatter,
|
||||
JsonFormatter,
|
||||
_initialize_loggers_with_handler,
|
||||
_turn_on_json,
|
||||
session_id_var,
|
||||
set_session_id,
|
||||
set_trace_id,
|
||||
trace_id_var,
|
||||
verbose_logger,
|
||||
verbose_proxy_logger,
|
||||
verbose_router_logger,
|
||||
|
|
@ -393,3 +399,244 @@ def test_logging_calls_do_not_build_their_message_eagerly():
|
|||
"these logging calls build their message eagerly; pass the values as %-style arguments instead:\n"
|
||||
+ "\n".join(offenders)
|
||||
)
|
||||
|
||||
|
||||
class _JsonCapture(logging.Handler):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.formatter = JsonFormatter()
|
||||
self.records: list[dict] = []
|
||||
self.addFilter(CorrelationContextFilter())
|
||||
|
||||
def emit(self, record):
|
||||
self.records.append(json.loads(self.formatter.format(record)))
|
||||
|
||||
|
||||
def _make_capture_logger(name: str) -> tuple[logging.Logger, _JsonCapture]:
|
||||
lg = logging.getLogger(name)
|
||||
cap = _JsonCapture()
|
||||
lg.addHandler(cap)
|
||||
lg.setLevel(logging.DEBUG)
|
||||
return lg, cap
|
||||
|
||||
|
||||
def test_trace_id_injected_into_json_record(monkeypatch):
|
||||
"""trace_id set via set_trace_id() appears in every JSON record in that context."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
lg, cap = _make_capture_logger("test.trace_inject")
|
||||
set_trace_id("trace-abc-123")
|
||||
try:
|
||||
lg.info("test message")
|
||||
assert len(cap.records) == 1
|
||||
assert cap.records[0]["trace_id"] == "trace-abc-123"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
|
||||
|
||||
def test_session_id_injected_when_set(monkeypatch):
|
||||
"""session_id set via set_session_id() appears in JSON record."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
lg, cap = _make_capture_logger("test.session_inject")
|
||||
set_session_id("sess-xyz-456")
|
||||
try:
|
||||
lg.info("another message")
|
||||
assert cap.records[0]["session_id"] == "sess-xyz-456"
|
||||
finally:
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_trace_id_and_session_id_cannot_be_spoofed_by_message_content(monkeypatch):
|
||||
"""A log message that happens to parse as JSON/dict with "trace_id"/"session_id"
|
||||
keys (e.g. the proxy logging a raw request-header dict) must not override the
|
||||
real correlation ids set via set_trace_id()/set_session_id()."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
lg, cap = _make_capture_logger("test.spoof_attempt")
|
||||
set_trace_id("real-trace-id")
|
||||
set_session_id("real-session-id")
|
||||
try:
|
||||
lg.info('{"trace_id": "attacker-supplied-trace", "session_id": "attacker-supplied-session"}')
|
||||
assert cap.records[0]["trace_id"] == "real-trace-id"
|
||||
assert cap.records[0]["session_id"] == "real-session-id"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_trace_id_and_session_id_cannot_be_injected_with_no_active_context(monkeypatch):
|
||||
"""A message that happens to parse as JSON/dict with "trace_id"/"session_id" keys
|
||||
must not surface those fields at all when CorrelationContextFilter hasn't stamped
|
||||
this record - e.g. a log line emitted before Logging.__init__() runs for a request
|
||||
(request_correlation_in_logs on, but no genuine trace/session id active yet)."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
lg, cap = _make_capture_logger("test.no_context_spoof_attempt")
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
lg.info('{"trace_id": "attacker-supplied-trace", "session_id": "attacker-supplied-session"}')
|
||||
assert "trace_id" not in cap.records[0]
|
||||
assert "session_id" not in cap.records[0]
|
||||
|
||||
|
||||
def test_trace_id_and_session_id_are_redacted_when_credential_shaped(monkeypatch):
|
||||
"""A caller-controlled trace_id/session_id (e.g. from x-litellm-trace-id or a W3C
|
||||
baggage header) that happens to look like a real credential must not reach log
|
||||
records unredacted. CorrelationContextFilter stamps trace_id/session_id onto the
|
||||
record after SecretRedactionFilter has already run, so those two fields would
|
||||
otherwise bypass credential redaction entirely - the fix redacts at set_trace_id()/
|
||||
set_session_id() time instead, before the value ever reaches a log record."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
lg, cap = _make_capture_logger("test.credential_shaped_correlation_id")
|
||||
poisoned_trace_id = "sk-ant-api03-" + "A" * 40
|
||||
poisoned_session_id = "AKIA" + "B" * 16
|
||||
set_trace_id(poisoned_trace_id)
|
||||
set_session_id(poisoned_session_id)
|
||||
try:
|
||||
lg.info("some benign log line")
|
||||
assert cap.records[0]["trace_id"] == "REDACTED"
|
||||
assert cap.records[0]["session_id"] == "REDACTED"
|
||||
assert poisoned_trace_id not in json.dumps(cap.records[0])
|
||||
assert poisoned_session_id not in json.dumps(cap.records[0])
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_session_id_absent_when_not_set():
|
||||
"""session_id must NOT appear in JSON record when not set for this context."""
|
||||
lg, cap = _make_capture_logger("test.no_session")
|
||||
session_id_var.set("")
|
||||
lg.info("no session message")
|
||||
assert "session_id" not in cap.records[0]
|
||||
|
||||
|
||||
def test_trace_id_absent_when_not_set():
|
||||
"""trace_id must NOT appear when not set."""
|
||||
lg, cap = _make_capture_logger("test.no_trace")
|
||||
trace_id_var.set("")
|
||||
lg.info("no trace message")
|
||||
assert "trace_id" not in cap.records[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_contextvar_isolation_between_tasks():
|
||||
"""Two concurrent async tasks each see only their own trace_id."""
|
||||
results: dict[str, str] = {}
|
||||
|
||||
async def task(task_id: str, trace_id: str) -> None:
|
||||
set_trace_id(trace_id)
|
||||
await asyncio.sleep(0)
|
||||
results[task_id] = trace_id_var.get()
|
||||
|
||||
await asyncio.gather(
|
||||
task("A", "trace-for-A"),
|
||||
task("B", "trace-for-B"),
|
||||
)
|
||||
|
||||
assert results["A"] == "trace-for-A"
|
||||
assert results["B"] == "trace-for-B"
|
||||
|
||||
|
||||
def test_trace_id_not_in_log_when_flag_disabled(monkeypatch):
|
||||
"""When request_correlation_in_logs is False (default), trace_id must not appear in JSON records even when set."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", False)
|
||||
lg, cap = _make_capture_logger("test.no_trace_gated")
|
||||
set_trace_id("trace-should-not-appear")
|
||||
try:
|
||||
lg.info("message")
|
||||
assert "trace_id" not in cap.records[0]
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
|
||||
|
||||
def test_session_id_not_in_log_when_flag_disabled(monkeypatch):
|
||||
"""When request_correlation_in_logs is False (default), session_id must not appear in JSON records even when set."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", False)
|
||||
lg, cap = _make_capture_logger("test.no_session_gated")
|
||||
set_session_id("sess-should-not-appear")
|
||||
try:
|
||||
lg.info("message")
|
||||
assert "session_id" not in cap.records[0]
|
||||
finally:
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
class _PlainCapture(logging.Handler):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.formatter = CorrelationPlainFormatter("%(message)s")
|
||||
self.records: list[str] = []
|
||||
self.addFilter(CorrelationContextFilter())
|
||||
|
||||
def emit(self, record):
|
||||
self.records.append(self.formatter.format(record))
|
||||
|
||||
|
||||
def _make_plain_capture_logger(name: str) -> tuple[logging.Logger, _PlainCapture]:
|
||||
lg = logging.getLogger(name)
|
||||
cap = _PlainCapture()
|
||||
lg.addHandler(cap)
|
||||
lg.setLevel(logging.DEBUG)
|
||||
return lg, cap
|
||||
|
||||
|
||||
def test_plain_formatter_appends_trace_id_and_session_id(monkeypatch):
|
||||
"""CorrelationPlainFormatter must append trace_id/session_id to non-JSON log lines too."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
lg, cap = _make_plain_capture_logger("test.plain_trace_session")
|
||||
set_trace_id("plain-trace-1")
|
||||
set_session_id("plain-session-1")
|
||||
try:
|
||||
lg.info("plaintext message")
|
||||
assert cap.records[0] == "plaintext message [trace_id=plain-trace-1 session_id=plain-session-1]"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_plain_formatter_appends_only_trace_id_when_session_id_absent(monkeypatch):
|
||||
"""Only trace_id is appended when session_id was never set."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
lg, cap = _make_plain_capture_logger("test.plain_trace_only")
|
||||
set_trace_id("plain-trace-2")
|
||||
session_id_var.set("")
|
||||
try:
|
||||
lg.info("plaintext message")
|
||||
assert cap.records[0] == "plaintext message [trace_id=plain-trace-2]"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
|
||||
|
||||
def test_plain_formatter_unchanged_when_flag_disabled(monkeypatch):
|
||||
"""When request_correlation_in_logs is False, plain log lines are unmodified even if the contextvars are set."""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", False)
|
||||
lg, cap = _make_plain_capture_logger("test.plain_flag_off")
|
||||
set_trace_id("should-not-appear")
|
||||
set_session_id("should-not-appear")
|
||||
try:
|
||||
lg.info("plaintext message")
|
||||
assert cap.records[0] == "plaintext message"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_set_trace_id_strips_control_characters():
|
||||
"""set_trace_id() must strip \\r/\\n/escape sequences so a caller-controlled
|
||||
trace id can't forge fake log entries when interpolated into plain-text logs."""
|
||||
token = set_trace_id('evil\r\n{"level": "CRITICAL", "message": "forged"}')
|
||||
try:
|
||||
value = trace_id_var.get()
|
||||
assert "\r" not in value
|
||||
assert "\n" not in value
|
||||
finally:
|
||||
trace_id_var.reset(token)
|
||||
|
||||
|
||||
def test_set_session_id_bounds_length():
|
||||
"""set_session_id() must bound length so an oversized caller-supplied value
|
||||
isn't repeated across every log line for the request."""
|
||||
token = set_session_id("a" * 1000)
|
||||
try:
|
||||
assert len(session_id_var.get()) == 256
|
||||
finally:
|
||||
session_id_var.reset(token)
|
||||
|
||||
|
|
|
|||
|
|
@ -6757,6 +6757,68 @@ async def test_acreate_batch_disable_fallbacks_surfaces_owning_provider_error():
|
|||
assert mock_create.call_args.kwargs["model"] == "owning-model"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acreate_batch_surfaces_owning_provider_error_without_disable_fallbacks():
|
||||
"""The router itself has to keep a batch inside the group that owns the input file:
|
||||
the proxy only sets disable_fallbacks on the managed-files route, so the caller
|
||||
otherwise gets the fallback provider's error for a file it never received."""
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "owning-model",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "sk-owning",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "fallback-model",
|
||||
"litellm_params": {
|
||||
"model": "azure/gpt-4o-mini",
|
||||
"api_key": "sk-fallback",
|
||||
"api_base": "https://fallback.openai.azure.com",
|
||||
"api_version": "2024-08-01-preview",
|
||||
},
|
||||
},
|
||||
],
|
||||
fallbacks=[{"owning-model": ["fallback-model"]}],
|
||||
num_retries=0,
|
||||
)
|
||||
attempted_models = []
|
||||
|
||||
async def _acreate_batch(model, **kwargs):
|
||||
attempted_models.append(model)
|
||||
if model == "owning-model":
|
||||
raise litellm.APIConnectionError(
|
||||
message="Connection error - openai is unreachable",
|
||||
model="openai/gpt-4o-mini",
|
||||
llm_provider="openai",
|
||||
)
|
||||
return LiteLLMBatch(
|
||||
id="batch-created-on-the-wrong-provider",
|
||||
completion_window="24h",
|
||||
created_at=0,
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="file-owned-by-openai",
|
||||
object="batch",
|
||||
status="validating",
|
||||
)
|
||||
|
||||
with patch.object(router, "_acreate_batch", _acreate_batch):
|
||||
with pytest.raises(litellm.APIConnectionError, match="openai is unreachable"):
|
||||
await router.acreate_batch(
|
||||
model="owning-model",
|
||||
input_file_id="file-owned-by-openai",
|
||||
endpoint="/v1/chat/completions",
|
||||
completion_window="24h",
|
||||
metadata={"team": "batch-jobs"},
|
||||
)
|
||||
|
||||
assert attempted_models == ["owning-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acreate_batch_request_bedrock_tags_override_deployment_tags():
|
||||
import httpx
|
||||
|
|
@ -7419,6 +7481,47 @@ class TestAutoRouterMaxInputCharsWiring:
|
|||
assert self._registered_auto_router(router).max_input_chars == DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
|
||||
|
||||
|
||||
class TestGetAllowedFailsFromPolicy:
|
||||
def _make_router(self, **policy_kwargs) -> litellm.Router:
|
||||
from litellm.types.router import AllowedFailsPolicy
|
||||
|
||||
return litellm.Router(
|
||||
model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "fake"}}],
|
||||
allowed_fails_policy=AllowedFailsPolicy(**policy_kwargs),
|
||||
)
|
||||
|
||||
def test_no_policy_returns_none(self):
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "fake"}}],
|
||||
)
|
||||
assert router.get_allowed_fails_from_policy(litellm.RateLimitError("429", "openai", "gpt-4")) is None
|
||||
|
||||
def test_internal_server_error_allowed_fails(self):
|
||||
router = self._make_router(InternalServerErrorAllowedFails=7)
|
||||
exc = litellm.InternalServerError("500", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) == 7
|
||||
|
||||
def test_service_unavailable_error_allowed_fails(self):
|
||||
router = self._make_router(ServiceUnavailableErrorAllowedFails=4)
|
||||
exc = litellm.ServiceUnavailableError("503", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) == 4
|
||||
|
||||
def test_bad_gateway_error_allowed_fails(self):
|
||||
router = self._make_router(BadGatewayErrorAllowedFails=2)
|
||||
exc = litellm.BadGatewayError("502", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) == 2
|
||||
|
||||
def test_not_found_error_allowed_fails(self):
|
||||
router = self._make_router(NotFoundErrorAllowedFails=1)
|
||||
exc = litellm.NotFoundError("404", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) == 1
|
||||
|
||||
def test_unmatched_exception_returns_none(self):
|
||||
router = self._make_router(InternalServerErrorAllowedFails=5)
|
||||
exc = litellm.RateLimitError("429", "openai", "gpt-4")
|
||||
assert router.get_allowed_fails_from_policy(exc) is None
|
||||
|
||||
|
||||
class _LogCapture(logging.Handler):
|
||||
def __init__(self, level):
|
||||
super().__init__(level=level)
|
||||
|
|
@ -7550,6 +7653,8 @@ async def test_fallback_failure_detail_from_upstream_is_bounded():
|
|||
assert capture.messages, "the fallback failure path did not log at ERROR"
|
||||
assert huge_message not in "".join(capture.messages)
|
||||
assert max(len(message) for message in capture.messages) < 5_000
|
||||
|
||||
|
||||
def test_stamp_or_clear_metadata_key_writes_and_clears_both_buckets():
|
||||
request_kwargs = {"metadata": {}}
|
||||
litellm.Router._stamp_or_clear_metadata_key(request_kwargs=request_kwargs, key="probe", value=7)
|
||||
|
|
|
|||
|
|
@ -56,16 +56,12 @@ class TestGetExcludedFilteredDeployments:
|
|||
# error. Returning the original list here would re-include the
|
||||
# just-failed deployment and let weighted failover re-pick it.
|
||||
deps = [_make_dep("a"), _make_dep("b")]
|
||||
result = _get_excluded_filtered_deployments(
|
||||
deps, excluded_deployment_ids=["a", "b"]
|
||||
)
|
||||
result = _get_excluded_filtered_deployments(deps, excluded_deployment_ids=["a", "b"])
|
||||
assert result == []
|
||||
|
||||
def test_excluded_set_with_unknown_ids(self):
|
||||
deps = [_make_dep("a"), _make_dep("b")]
|
||||
result = _get_excluded_filtered_deployments(
|
||||
deps, excluded_deployment_ids=["zzz"]
|
||||
)
|
||||
result = _get_excluded_filtered_deployments(deps, excluded_deployment_ids=["zzz"])
|
||||
assert len(result) == 2
|
||||
|
||||
def test_handles_missing_model_info(self):
|
||||
|
|
@ -100,6 +96,159 @@ def test_set_failed_deployment_id_on_exception():
|
|||
assert exc.failed_deployment_id == "dep-a"
|
||||
|
||||
|
||||
def test_stamp_failed_deployment_id_with_effective_model_info_prefers_kwargs():
|
||||
"""kwargs["model_info"] (the dynamic client-side-credential id, when present) must win
|
||||
over the static deployment's model_info, so a bad-credential tenant's failures are
|
||||
attributed to their own dynamic deployment id, not the shared static one."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "test-model",
|
||||
"litellm_params": {"model": "gpt-4o", "api_key": "key"},
|
||||
"model_info": {"id": "dep-a"},
|
||||
}
|
||||
],
|
||||
)
|
||||
exc = Exception("fail")
|
||||
router._stamp_failed_deployment_id_with_effective_model_info(
|
||||
exc, _make_dep("dep-a"), {"model_info": {"id": "dynamic-dep"}}
|
||||
)
|
||||
assert exc.failed_deployment_id == "dynamic-dep"
|
||||
|
||||
|
||||
def test_stamp_failed_deployment_id_with_effective_model_info_falls_back_to_deployment():
|
||||
"""With no dynamic id in kwargs (the common, non-client-side-credential case), the
|
||||
static deployment's own model_info.id must still be stamped."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "test-model",
|
||||
"litellm_params": {"model": "gpt-4o", "api_key": "key"},
|
||||
"model_info": {"id": "dep-a"},
|
||||
}
|
||||
],
|
||||
)
|
||||
exc = Exception("fail")
|
||||
router._stamp_failed_deployment_id_with_effective_model_info(exc, _make_dep("dep-a"), {})
|
||||
assert exc.failed_deployment_id == "dep-a"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ageneric_api_call_with_fallbacks_helper_stamps_failed_deployment_id():
|
||||
"""_ageneric_api_call_with_fallbacks_helper must stamp failed_deployment_id on a
|
||||
failure, same as _completion/_acompletion, so callers identifying the failed
|
||||
deployment (cooldown, weighted failover) work for this call type too instead of
|
||||
depending on which metadata bucket ("metadata" vs "litellm_metadata") it uses."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "test-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "test-key"},
|
||||
"model_info": {"id": "dep-a"},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
async def _failing_original_function(**kwargs):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
await router._ageneric_api_call_with_fallbacks_helper(
|
||||
model="test-model",
|
||||
original_generic_function=_failing_original_function,
|
||||
)
|
||||
|
||||
assert getattr(exc_info.value, "failed_deployment_id", None) == "dep-a"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ageneric_api_call_with_fallbacks_helper_stamps_dynamic_id_for_clientside_credentials():
|
||||
"""A client-side-credential call (tenant-supplied api_key) generates a dynamic
|
||||
deployment id distinct from the shared static deployment. Stamping the static id
|
||||
instead would let one tenant's bad credentials cool down the deployment every
|
||||
other tenant sharing this config relies on."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "test-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "test-key"},
|
||||
"model_info": {"id": "dep-a"},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
async def _failing_original_function(**kwargs):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
await router._ageneric_api_call_with_fallbacks_helper(
|
||||
model="test-model",
|
||||
original_generic_function=_failing_original_function,
|
||||
api_key="tenant-supplied-key",
|
||||
litellm_metadata={"model_group": "test-model"},
|
||||
)
|
||||
|
||||
failed_deployment_id = getattr(exc_info.value, "failed_deployment_id", None)
|
||||
assert failed_deployment_id is not None
|
||||
assert failed_deployment_id != "dep-a"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_stamps_dynamic_id_for_clientside_credentials():
|
||||
"""Same bug as the generic-API-call helper above, but in the regular completion
|
||||
path: _acompletion's exception handlers must stamp the dynamic client-side-credential
|
||||
deployment id, not the shared static deployment's id."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "test-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "test-key"},
|
||||
"model_info": {"id": "dep-a"},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
with patch("litellm.acompletion", new_callable=AsyncMock, side_effect=RuntimeError("boom")):
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
await router._acompletion(
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
api_key="tenant-supplied-key",
|
||||
metadata={"model_group": "test-model"},
|
||||
)
|
||||
|
||||
failed_deployment_id = getattr(exc_info.value, "failed_deployment_id", None)
|
||||
assert failed_deployment_id is not None
|
||||
assert failed_deployment_id != "dep-a"
|
||||
|
||||
|
||||
def test_completion_stamps_dynamic_id_for_clientside_credentials():
|
||||
"""Sync counterpart: _completion's exception handler must stamp the dynamic
|
||||
client-side-credential deployment id, not the shared static deployment's id."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "test-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "test-key"},
|
||||
"model_info": {"id": "dep-a"},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
with patch("litellm.completion", side_effect=RuntimeError("boom")):
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
router._completion(
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
api_key="tenant-supplied-key",
|
||||
metadata={"model_group": "test-model"},
|
||||
)
|
||||
|
||||
failed_deployment_id = getattr(exc_info.value, "failed_deployment_id", None)
|
||||
assert failed_deployment_id is not None
|
||||
assert failed_deployment_id != "dep-a"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_maybe_run_weighted_failover_returns_none_without_failed_id():
|
||||
router = Router(
|
||||
|
|
@ -641,12 +790,8 @@ async def test_maybe_run_weighted_failover_skips_when_remaining_all_in_cooldown(
|
|||
input_kwargs={},
|
||||
)
|
||||
|
||||
assert (
|
||||
result is None
|
||||
), "Should return None when all remaining deployments are in cooldown"
|
||||
assert (
|
||||
not run_async_fallback_called
|
||||
), "run_async_fallback must NOT be called when no healthy deployments remain"
|
||||
assert result is None, "Should return None when all remaining deployments are in cooldown"
|
||||
assert not run_async_fallback_called, "run_async_fallback must NOT be called when no healthy deployments remain"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -705,9 +850,7 @@ async def test_maybe_run_weighted_failover_proceeds_when_one_healthy_remains(
|
|||
)
|
||||
|
||||
assert result == "ok from C"
|
||||
assert (
|
||||
run_async_fallback_called
|
||||
), "run_async_fallback must be called when a healthy deployment remains"
|
||||
assert run_async_fallback_called, "run_async_fallback must be called when a healthy deployment remains"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -11,6 +12,13 @@ sys.path.insert(
|
|||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
from litellm._logging import (
|
||||
CorrelationContextFilter,
|
||||
JsonFormatter,
|
||||
session_id_var,
|
||||
trace_id_var,
|
||||
verbose_logger,
|
||||
)
|
||||
from litellm.proxy.utils import is_valid_api_key
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
|
|
@ -5125,3 +5133,124 @@ def test_ai21_api_key_is_resolved_from_the_documented_env_var(monkeypatch: pytes
|
|||
monkeypatch.setenv("AI21_API_KEY", "sk-ai21-resolved-from-env")
|
||||
|
||||
assert get_api_key(llm_provider="ai21", dynamic_api_key=None) == "sk-ai21-resolved-from-env"
|
||||
|
||||
|
||||
class _JsonCapture(logging.Handler):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.formatter = JsonFormatter()
|
||||
self.records: list[dict] = []
|
||||
self.addFilter(CorrelationContextFilter())
|
||||
|
||||
def emit(self, record):
|
||||
self.records.append(json.loads(self.formatter.format(record)))
|
||||
|
||||
|
||||
def _make_capture_logger(name: str) -> tuple[logging.Logger, _JsonCapture]:
|
||||
lg = logging.getLogger(name)
|
||||
cap = _JsonCapture()
|
||||
lg.addHandler(cap)
|
||||
lg.setLevel(logging.DEBUG)
|
||||
return lg, cap
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapper_async_restores_originating_task_context_after_success(monkeypatch):
|
||||
"""A successful acompletion() dispatches async_success_handler via
|
||||
asyncio.create_task + the global logging worker - a different Task than the
|
||||
one running acompletion() itself (this test's own task). That handler's own
|
||||
restore only fixes up the detached child task it runs in; wrapper_async's own
|
||||
finally block (in litellm/utils.py) must separately restore the *originating*
|
||||
task's trace_id/session_id, since nothing else does.
|
||||
"""
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
trace_id_var.set("outer-trace-wrapper-test")
|
||||
session_id_var.set("outer-session-wrapper-test")
|
||||
try:
|
||||
await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="Hello there!",
|
||||
litellm_session_id="mock-call-session",
|
||||
num_retries=0,
|
||||
)
|
||||
assert trace_id_var.get() == "outer-trace-wrapper-test"
|
||||
assert session_id_var.get() == "outer-session-wrapper-test"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_function_setup_failure_after_logging_construction_restores_context(monkeypatch):
|
||||
"""If function_setup() constructs Logging() (which already mutated
|
||||
trace_id_var/session_id_var in __init__) but then raises before returning,
|
||||
the caller's wrapper() never gets a logging_obj reference to restore from.
|
||||
function_setup()'s own except block must restore the correlation context
|
||||
itself in that case, or it leaks into every subsequent log line in this
|
||||
thread/task until something unrelated happens to reset it."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
|
||||
def _boom(self, *args, **kwargs):
|
||||
raise RuntimeError("simulated failure after Logging() construction")
|
||||
|
||||
monkeypatch.setattr(Logging, "update_environment_variables", _boom)
|
||||
|
||||
trace_id_var.set("pre-setup-failure-trace")
|
||||
session_id_var.set("pre-setup-failure-session")
|
||||
try:
|
||||
with pytest.raises(RuntimeError, match="simulated failure"):
|
||||
litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="Hello there!",
|
||||
litellm_session_id="doomed-call-session",
|
||||
num_retries=0,
|
||||
)
|
||||
assert trace_id_var.get() == "pre-setup-failure-trace"
|
||||
assert session_id_var.get() == "pre-setup-failure-session"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_function_setup_failure_log_line_shows_outer_not_doomed_ids(monkeypatch):
|
||||
"""The 'Error in function_setup' diagnostic log line itself must be stamped
|
||||
with the outer/pre-call correlation ids, not the doomed call's own ids -
|
||||
restoring context must happen *before* logging the exception, not after,
|
||||
since the failed call never produces a usable logging object for anything
|
||||
else to be attributed to."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
|
||||
def _boom(self, *args, **kwargs):
|
||||
raise RuntimeError("simulated failure after Logging() construction")
|
||||
|
||||
monkeypatch.setattr(Logging, "update_environment_variables", _boom)
|
||||
|
||||
lg, cap = _make_capture_logger("test.function_setup_failure_log_order")
|
||||
# verbose_logger is a distinct, module-level logger from our throwaway one -
|
||||
# temporarily attach the same capture handler so we see its own emitted record.
|
||||
verbose_logger.addHandler(cap)
|
||||
try:
|
||||
trace_id_var.set("outer-trace")
|
||||
session_id_var.set("outer-session")
|
||||
with pytest.raises(RuntimeError, match="simulated failure"):
|
||||
litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="Hello there!",
|
||||
litellm_session_id="doomed-call-session",
|
||||
num_retries=0,
|
||||
)
|
||||
setup_failure_records = [r for r in cap.records if "Error in function_setup" in r.get("message", "")]
|
||||
assert len(setup_failure_records) == 1
|
||||
record = setup_failure_records[0]
|
||||
assert record.get("session_id") == "outer-session"
|
||||
assert record.get("trace_id") == "outer-trace"
|
||||
finally:
|
||||
verbose_logger.removeHandler(cap)
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 23149
|
||||
"limit": 23064
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 27166
|
||||
|
|
@ -15,7 +15,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT006": {
|
||||
"limit": 1086
|
||||
"limit": 1078
|
||||
},
|
||||
"LIT007": {
|
||||
"limit": 0
|
||||
|
|
@ -27,7 +27,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16758
|
||||
"limit": 16753
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5598
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import useTeams from "@/app/(dashboard)/hooks/useTeams";
|
||||
import { BarChart, DonutChart } from "@/components/shared/charts";
|
||||
import { buildCostBreakdownTiles, buildSummaryTiles, hasFlatCost, type SummaryTile } from "./entityUsageSummary";
|
||||
import { MoneyCell } from "@/components/shared/table_cells";
|
||||
import { Card as ShadcnCard, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
|
||||
import { formatNumberWithCommas } from "@/utils/dataUtils";
|
||||
|
|
@ -23,8 +24,8 @@ import {
|
|||
Text,
|
||||
Title,
|
||||
} from "@tremor/react";
|
||||
import { ExportOutlined, LoadingOutlined } from "@ant-design/icons";
|
||||
import { Alert, Button } from "antd";
|
||||
import { DownOutlined, ExportOutlined, InfoCircleOutlined, LoadingOutlined, RightOutlined } from "@ant-design/icons";
|
||||
import { Alert, Button, Tooltip } from "antd";
|
||||
import React, { type ReactNode, useMemo, useState } from "react";
|
||||
import TeamMultiSelect from "@/components/common_components/team_multi_select";
|
||||
import { ActivityMetrics, processActivityData } from "@/components/activity_metrics";
|
||||
|
|
@ -76,6 +77,7 @@ interface EntitySpendData {
|
|||
results: ExtendedDailyData[];
|
||||
metadata: {
|
||||
total_spend: number;
|
||||
total_flat_cost?: number;
|
||||
total_api_requests: number;
|
||||
total_successful_requests: number;
|
||||
total_failed_requests: number;
|
||||
|
|
@ -115,6 +117,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
const [topKeysLimit, setTopKeysLimit] = useState<number>(5);
|
||||
const [topModelsLimit, setTopModelsLimit] = useState<number>(5);
|
||||
const [topAgentsLimit, setTopAgentsLimit] = useState<number>(5);
|
||||
const [showCostBreakdown, setShowCostBreakdown] = useState(false);
|
||||
|
||||
const startTime = useMemo(() => (dateValue.from ? new Date(dateValue.from) : null), [dateValue.from]);
|
||||
const endTime = useMemo(() => (dateValue.to ? new Date(dateValue.to) : null), [dateValue.to]);
|
||||
|
|
@ -408,42 +411,39 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
};
|
||||
|
||||
const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1);
|
||||
const showFlatCost = entityType === "team" && hasFlatCost(spendData.metadata);
|
||||
|
||||
const chev = "text-gray-400 text-xs";
|
||||
const expandIcon = showCostBreakdown ? <DownOutlined className={chev} /> : <RightOutlined className={chev} />;
|
||||
const infoIcon = <InfoCircleOutlined className="text-gray-400 hover:text-gray-600" />;
|
||||
|
||||
const renderSummaryTile = ({ title, value, className, tooltip, expandable }: SummaryTile) => (
|
||||
<Card
|
||||
key={title}
|
||||
className={expandable ? "cursor-pointer hover:bg-gray-50 transition-colors" : undefined}
|
||||
onClick={expandable ? () => setShowCostBreakdown(!showCostBreakdown) : undefined}
|
||||
>
|
||||
<div className="flex items-center gap-2">
|
||||
<Title>{title}</Title>
|
||||
{tooltip ? <Tooltip title={tooltip}>{infoIcon}</Tooltip> : null}
|
||||
{expandable ? expandIcon : null}
|
||||
</div>
|
||||
<Text className={`text-2xl font-bold mt-2 ${className ?? ""}`}>{value}</Text>
|
||||
</Card>
|
||||
);
|
||||
|
||||
const breakdownTiles = showFlatCost && showCostBreakdown ? buildCostBreakdownTiles(spendData.metadata) : [];
|
||||
const summaryTiles = [...buildSummaryTiles(spendData.metadata, showFlatCost), ...breakdownTiles];
|
||||
|
||||
const modelViewTitle = modelViewType === "groups" ? "Top Public Model Names" : "Top Litellm Models";
|
||||
|
||||
const costPanel = (
|
||||
<Grid numItems={2} className="gap-2 w-full">
|
||||
{/* Total Spend Card */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<Title>{capitalizedEntityLabel} Spend Overview</Title>
|
||||
<Grid numItems={5} className="gap-4 mt-4">
|
||||
<Card>
|
||||
<Title>Total Spend</Title>
|
||||
<Text className="text-2xl font-bold mt-2">
|
||||
${formatNumberWithCommas(spendData.metadata.total_spend, 2)}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2">{spendData.metadata.total_api_requests.toLocaleString()}</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Successful Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-green-600">
|
||||
{spendData.metadata.total_successful_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Failed Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-red-600">
|
||||
{spendData.metadata.total_failed_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Tokens</Title>
|
||||
<Text className="text-2xl font-bold mt-2">{spendData.metadata.total_tokens.toLocaleString()}</Text>
|
||||
</Card>
|
||||
{summaryTiles.map(renderSummaryTile)}
|
||||
</Grid>
|
||||
</Card>
|
||||
</Col>
|
||||
|
|
@ -456,21 +456,40 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
</CardHeader>
|
||||
<CardContent>
|
||||
<BarChart
|
||||
data={[...spendData.results].sort((a, b) => new Date(a.date).getTime() - new Date(b.date).getTime())}
|
||||
data={[...spendData.results]
|
||||
.sort((a, b) => new Date(a.date).getTime() - new Date(b.date).getTime())
|
||||
.map((row) => ({
|
||||
...row,
|
||||
"Request cost": row.metrics.spend ?? 0,
|
||||
"Flat cost": row.metrics.flat_cost ?? 0,
|
||||
}))}
|
||||
index="date"
|
||||
categories={["metrics.spend"]}
|
||||
colors={["cyan"]}
|
||||
categories={showFlatCost ? ["Request cost", "Flat cost"] : ["metrics.spend"]}
|
||||
colors={showFlatCost ? ["cyan", "violet"] : ["cyan"]}
|
||||
stack={showFlatCost}
|
||||
valueFormatter={valueFormatterSpend}
|
||||
yAxisWidth={100}
|
||||
showLegend={false}
|
||||
showLegend={showFlatCost}
|
||||
customTooltip={({ payload, active }) => {
|
||||
if (!active || !payload?.[0]) return null;
|
||||
const data = payload[0].payload;
|
||||
const entityCount = Object.keys(data.breakdown.entities || {}).length;
|
||||
const requestSpend = data.metrics.spend ?? 0;
|
||||
const flatCost = data.metrics.flat_cost ?? 0;
|
||||
return (
|
||||
<div className="bg-white p-4 shadow-lg rounded-lg border">
|
||||
<p className="font-bold">{data.date}</p>
|
||||
<p className="text-cyan-500">Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}</p>
|
||||
{showFlatCost ? (
|
||||
<>
|
||||
<p className="text-cyan-500">Request cost: ${formatNumberWithCommas(requestSpend, 2)}</p>
|
||||
<p className="text-violet-500">Flat cost: ${formatNumberWithCommas(flatCost, 2)}</p>
|
||||
<p className="font-semibold">
|
||||
Total cost: ${formatNumberWithCommas(requestSpend + flatCost, 2)}
|
||||
</p>
|
||||
</>
|
||||
) : (
|
||||
<p className="text-cyan-500">Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}</p>
|
||||
)}
|
||||
<p className="text-gray-600">Total Requests: {data.metrics.api_requests}</p>
|
||||
<p className="text-gray-600">Successful: {data.metrics.successful_requests}</p>
|
||||
<p className="text-gray-600">Failed: {data.metrics.failed_requests}</p>
|
||||
|
|
|
|||
|
|
@ -0,0 +1,82 @@
|
|||
import { describe, expect, it } from "vitest";
|
||||
import { buildCostBreakdownTiles, buildSummaryTiles, hasFlatCost } from "./entityUsageSummary";
|
||||
|
||||
const metadata = {
|
||||
total_spend: 100,
|
||||
total_flat_cost: 40,
|
||||
total_api_requests: 12,
|
||||
total_successful_requests: 10,
|
||||
total_failed_requests: 2,
|
||||
total_tokens: 3456,
|
||||
};
|
||||
|
||||
describe("hasFlatCost", () => {
|
||||
it("is false when there is no flat cost to report", () => {
|
||||
expect(hasFlatCost({ ...metadata, total_flat_cost: 0 })).toBe(false);
|
||||
const { total_flat_cost, ...noFlat } = metadata;
|
||||
expect(hasFlatCost(noFlat)).toBe(false);
|
||||
});
|
||||
|
||||
it("is true once a flat cost has accrued", () => {
|
||||
expect(hasFlatCost(metadata)).toBe(true);
|
||||
});
|
||||
});
|
||||
|
||||
describe("buildSummaryTiles", () => {
|
||||
it("keeps the row at five tiles either way so adding flat cost never narrows the cards", () => {
|
||||
expect(buildSummaryTiles(metadata, false)).toHaveLength(5);
|
||||
expect(buildSummaryTiles(metadata, true)).toHaveLength(5);
|
||||
});
|
||||
|
||||
it("shows request-only spend under the original title when there is no flat cost", () => {
|
||||
const [first] = buildSummaryTiles(metadata, false);
|
||||
expect(first.title).toBe("Total Spend");
|
||||
expect(first.value).toBe("$100.00");
|
||||
expect(first.expandable).toBeUndefined();
|
||||
});
|
||||
|
||||
it("rolls flat cost into a single expandable Total Cost tile", () => {
|
||||
const [first] = buildSummaryTiles(metadata, true);
|
||||
expect(first.title).toBe("Total Cost");
|
||||
expect(first.value).toBe("$140.00");
|
||||
expect(first.expandable).toBe(true);
|
||||
expect(first.tooltip).toBeTruthy();
|
||||
});
|
||||
|
||||
it("never renders the breakdown titles in the top row", () => {
|
||||
const titles = buildSummaryTiles(metadata, true).map((t) => t.title);
|
||||
expect(titles).not.toContain("Flat Cost");
|
||||
expect(titles).not.toContain("Request Cost");
|
||||
});
|
||||
|
||||
it("treats a missing flat cost as zero", () => {
|
||||
const { total_flat_cost, ...noFlat } = metadata;
|
||||
expect(buildSummaryTiles(noFlat, true)[0].value).toBe("$100.00");
|
||||
});
|
||||
});
|
||||
|
||||
describe("buildCostBreakdownTiles", () => {
|
||||
it("splits the total into request cost and flat cost", () => {
|
||||
const byTitle = Object.fromEntries(buildCostBreakdownTiles(metadata).map((t) => [t.title, t.value]));
|
||||
expect(byTitle["Request Cost"]).toBe("$100.00");
|
||||
expect(byTitle["Flat Cost"]).toBe("$40.00");
|
||||
});
|
||||
|
||||
it("adds up to the Total Cost tile so the expanded view reconciles", () => {
|
||||
const parse = (v: string) => Number(v.replace(/[$,]/g, ""));
|
||||
const parts = buildCostBreakdownTiles(metadata).map((t) => parse(t.value));
|
||||
expect(parts[0] + parts[1]).toBe(parse(buildSummaryTiles(metadata, true)[0].value));
|
||||
});
|
||||
|
||||
it("explains each part, including that flat cost is outside budgets", () => {
|
||||
const byTitle = Object.fromEntries(buildCostBreakdownTiles(metadata).map((t) => [t.title, t.tooltip]));
|
||||
expect(byTitle["Request Cost"]).toBeTruthy();
|
||||
expect(byTitle["Flat Cost"]).toContain("budget");
|
||||
});
|
||||
|
||||
it("treats a missing flat cost as zero", () => {
|
||||
const { total_flat_cost, ...noFlat } = metadata;
|
||||
const byTitle = Object.fromEntries(buildCostBreakdownTiles(noFlat).map((t) => [t.title, t.value]));
|
||||
expect(byTitle["Flat Cost"]).toBe("$0.00");
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,66 @@
|
|||
import { formatNumberWithCommas } from "@/utils/dataUtils";
|
||||
|
||||
export interface SummaryTile {
|
||||
title: string;
|
||||
value: string;
|
||||
className?: string;
|
||||
tooltip?: string;
|
||||
expandable?: boolean;
|
||||
}
|
||||
|
||||
interface SpendSummaryMetadata {
|
||||
total_spend: number;
|
||||
total_flat_cost?: number;
|
||||
total_api_requests: number;
|
||||
total_successful_requests: number;
|
||||
total_failed_requests: number;
|
||||
total_tokens: number;
|
||||
}
|
||||
|
||||
export const TOTAL_COST_TOOLTIP =
|
||||
"Request cost plus flat cost for reserved capacity. Select this tile to see the breakdown.";
|
||||
|
||||
export const REQUEST_COST_TOOLTIP =
|
||||
"Usage-based cost of the requests this entity sent during the selected period, priced per token.";
|
||||
|
||||
export const FLAT_COST_TOOLTIP =
|
||||
"Reserved provisioned throughput, billed per hour whether or not requests are sent. Reported here only; it does not count toward team, key, user, or organization budgets.";
|
||||
|
||||
export const hasFlatCost = (metadata: SpendSummaryMetadata): boolean => (metadata.total_flat_cost ?? 0) > 0;
|
||||
|
||||
export const buildSummaryTiles = (metadata: SpendSummaryMetadata, showFlatCost: boolean): SummaryTile[] => {
|
||||
const flatCost = metadata.total_flat_cost ?? 0;
|
||||
return [
|
||||
showFlatCost
|
||||
? {
|
||||
title: "Total Cost",
|
||||
value: `$${formatNumberWithCommas(metadata.total_spend + flatCost, 2)}`,
|
||||
tooltip: TOTAL_COST_TOOLTIP,
|
||||
expandable: true,
|
||||
}
|
||||
: { title: "Total Spend", value: `$${formatNumberWithCommas(metadata.total_spend, 2)}` },
|
||||
{ title: "Total Requests", value: metadata.total_api_requests.toLocaleString() },
|
||||
{
|
||||
title: "Successful Requests",
|
||||
value: metadata.total_successful_requests.toLocaleString(),
|
||||
className: "text-green-600",
|
||||
},
|
||||
{ title: "Failed Requests", value: metadata.total_failed_requests.toLocaleString(), className: "text-red-600" },
|
||||
{ title: "Total Tokens", value: metadata.total_tokens.toLocaleString() },
|
||||
];
|
||||
};
|
||||
|
||||
export const buildCostBreakdownTiles = (metadata: SpendSummaryMetadata): SummaryTile[] => [
|
||||
{
|
||||
title: "Request Cost",
|
||||
value: `$${formatNumberWithCommas(metadata.total_spend, 2)}`,
|
||||
className: "text-cyan-600",
|
||||
tooltip: REQUEST_COST_TOOLTIP,
|
||||
},
|
||||
{
|
||||
title: "Flat Cost",
|
||||
value: `$${formatNumberWithCommas(metadata.total_flat_cost ?? 0, 2)}`,
|
||||
className: "text-violet-600",
|
||||
tooltip: FLAT_COST_TOOLTIP,
|
||||
},
|
||||
];
|
||||
|
|
@ -0,0 +1,51 @@
|
|||
import { describe, expect, it } from "vitest";
|
||||
import { sumMetadata } from "./usePaginatedDailyActivity";
|
||||
|
||||
describe("sumMetadata", () => {
|
||||
it("sums flat cost across pages instead of keeping the first page's value", () => {
|
||||
// A team whose activity spans more than one page accrues flat cost on each of them.
|
||||
// Keeping page 1's value under-reports the Flat Cost and Total Cost tiles.
|
||||
const merged = sumMetadata({ total_spend: 1, total_flat_cost: 174.5 }, { total_spend: 2, total_flat_cost: 777 });
|
||||
|
||||
expect(merged.total_flat_cost).toBe(951.5);
|
||||
expect(merged.total_spend).toBe(3);
|
||||
});
|
||||
|
||||
it("treats a page missing the field as zero rather than dropping the running total", () => {
|
||||
expect(sumMetadata({ total_flat_cost: 480 }, {}).total_flat_cost).toBe(480);
|
||||
expect(sumMetadata({}, { total_flat_cost: 480 }).total_flat_cost).toBe(480);
|
||||
});
|
||||
|
||||
it("carries non-summable keys through from the first page", () => {
|
||||
const merged = sumMetadata(
|
||||
{ page: 1, total_pages: 3, total_spend: 1 },
|
||||
{ page: 2, total_pages: 3, total_spend: 2 },
|
||||
);
|
||||
|
||||
expect(merged.page).toBe(1);
|
||||
expect(merged.total_pages).toBe(3);
|
||||
});
|
||||
|
||||
it("sums every total_* metric the daily activity metadata exposes", () => {
|
||||
// Guards the class of bug rather than one field: a new backend total that nobody adds
|
||||
// to SUMMABLE_METADATA_KEYS freezes at page 1, and spend still looks right so it reads
|
||||
// as trustworthy.
|
||||
const page = {
|
||||
total_spend: 1,
|
||||
total_prompt_tokens: 1,
|
||||
total_completion_tokens: 1,
|
||||
total_tokens: 1,
|
||||
total_api_requests: 1,
|
||||
total_successful_requests: 1,
|
||||
total_failed_requests: 1,
|
||||
total_cache_read_input_tokens: 1,
|
||||
total_cache_creation_input_tokens: 1,
|
||||
total_flat_cost: 1,
|
||||
};
|
||||
const merged = sumMetadata(page, page);
|
||||
|
||||
for (const key of Object.keys(page)) {
|
||||
expect(merged[key], `${key} must be summed across pages`).toBe(2);
|
||||
}
|
||||
});
|
||||
});
|
||||
|
|
@ -23,6 +23,7 @@ const SUMMABLE_METADATA_KEYS = [
|
|||
"total_failed_requests",
|
||||
"total_cache_read_input_tokens",
|
||||
"total_cache_creation_input_tokens",
|
||||
"total_flat_cost",
|
||||
] as const;
|
||||
|
||||
interface DailyActivityResponse {
|
||||
|
|
@ -68,7 +69,12 @@ const EMPTY_DATA: DailyActivityResponse = {
|
|||
},
|
||||
};
|
||||
|
||||
function sumMetadata(a: Record<string, any>, b: Record<string, any>): Record<string, any> {
|
||||
/**
|
||||
* Combine two pages of metadata. Only keys in SUMMABLE_METADATA_KEYS are added; anything
|
||||
* else keeps the first page's value, so a total the backend adds later is silently frozen
|
||||
* at page 1 until it is listed above. Exported so that contract can be tested directly.
|
||||
*/
|
||||
export function sumMetadata(a: Record<string, any>, b: Record<string, any>): Record<string, any> {
|
||||
const result = { ...a };
|
||||
for (const key of SUMMABLE_METADATA_KEYS) {
|
||||
result[key] = (a[key] || 0) + (b[key] || 0);
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ export interface EntitySpendData {
|
|||
results: any[];
|
||||
metadata: {
|
||||
total_spend: number;
|
||||
total_flat_cost?: number;
|
||||
total_api_requests: number;
|
||||
total_successful_requests: number;
|
||||
total_failed_requests: number;
|
||||
|
|
@ -38,6 +39,8 @@ export interface ExportMetadata {
|
|||
export_scope: ExportScope;
|
||||
summary: {
|
||||
total_spend: number;
|
||||
total_flat_cost?: number;
|
||||
total_cost?: number;
|
||||
total_requests: number;
|
||||
successful_requests: number;
|
||||
failed_requests: number;
|
||||
|
|
|
|||
|
|
@ -1861,6 +1861,87 @@ describe("EntityUsageExport utils", () => {
|
|||
expect(result.summary.failed_requests).toBe(20);
|
||||
expect(result.summary.total_tokens).toBe(4500);
|
||||
});
|
||||
|
||||
it("should include total_flat_cost and total_cost in summary when total_flat_cost is present", () => {
|
||||
const spendWithFlat: EntitySpendData = {
|
||||
...mockSpendData,
|
||||
metadata: { ...mockSpendData.metadata, total_flat_cost: 6.45 },
|
||||
};
|
||||
const result = generateMetadata("team", mockDateRange, [], "daily", spendWithFlat);
|
||||
expect(result.summary.total_flat_cost).toBeCloseTo(6.45, 4);
|
||||
expect(result.summary.total_cost).toBeCloseTo(46.0 + 6.45, 4);
|
||||
});
|
||||
|
||||
it("should omit total_flat_cost and total_cost when total_flat_cost is zero", () => {
|
||||
const zeroFlat = { ...mockSpendData, metadata: { ...mockSpendData.metadata, total_flat_cost: 0 } };
|
||||
const result = generateMetadata("team", mockDateRange, [], "daily", zeroFlat);
|
||||
expect(result.summary.total_flat_cost).toBeUndefined();
|
||||
expect(result.summary.total_cost).toBeUndefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("generateDailyData PTU flat cost", () => {
|
||||
const dayWithFlat: EntitySpendData = {
|
||||
results: [
|
||||
{
|
||||
date: "2025-01-01",
|
||||
breakdown: {
|
||||
entities: {
|
||||
"team-1": {
|
||||
metrics: {
|
||||
spend: 10,
|
||||
flat_cost: 6.45,
|
||||
api_requests: 50,
|
||||
successful_requests: 50,
|
||||
failed_requests: 0,
|
||||
total_tokens: 500,
|
||||
prompt_tokens: 300,
|
||||
completion_tokens: 200,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
},
|
||||
api_key_breakdown: {},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
],
|
||||
metadata: {
|
||||
total_spend: 10,
|
||||
total_flat_cost: 6.45,
|
||||
total_api_requests: 50,
|
||||
total_successful_requests: 50,
|
||||
total_failed_requests: 0,
|
||||
total_tokens: 500,
|
||||
},
|
||||
};
|
||||
|
||||
it("includes Flat Cost ($) and Total Cost ($) columns when total_flat_cost is present", () => {
|
||||
const rows = generateDailyData(dayWithFlat, "Team", {});
|
||||
expect(rows).toHaveLength(1);
|
||||
expect(rows[0]).toHaveProperty("Flat Cost ($)");
|
||||
expect(rows[0]).toHaveProperty("Total Cost ($)");
|
||||
expect(rows[0]["Flat Cost ($)"]).toBe("6.4500");
|
||||
expect(rows[0]["Total Cost ($)"]).toBe("16.4500");
|
||||
});
|
||||
|
||||
it("does not include Flat Cost / Total Cost columns when total_flat_cost is zero", () => {
|
||||
const spendWithoutFlat: EntitySpendData = {
|
||||
...dayWithFlat,
|
||||
metadata: {
|
||||
total_spend: 10,
|
||||
total_api_requests: 50,
|
||||
total_successful_requests: 50,
|
||||
total_failed_requests: 0,
|
||||
total_tokens: 500,
|
||||
total_flat_cost: 0,
|
||||
},
|
||||
};
|
||||
const rows = generateDailyData(spendWithoutFlat, "User", {});
|
||||
expect(rows).toHaveLength(1);
|
||||
expect(rows[0]).not.toHaveProperty("Flat Cost ($)");
|
||||
expect(rows[0]).not.toHaveProperty("Total Cost ($)");
|
||||
});
|
||||
});
|
||||
|
||||
describe("handleExportCSV", () => {
|
||||
|
|
|
|||
|
|
@ -110,31 +110,42 @@ export const getEntityBreakdown = (
|
|||
return Object.values(entitySpend).sort((a, b) => b.metrics.spend - a.metrics.spend);
|
||||
};
|
||||
|
||||
// total_flat_cost defaults to 0 on every entity response, so only a non-zero value
|
||||
// means a PTU-configured team actually accrued flat cost worth exporting.
|
||||
const hasFlatCost = (spendData: EntitySpendData): boolean => (spendData.metadata.total_flat_cost ?? 0) > 0;
|
||||
|
||||
export const generateDailyData = (
|
||||
spendData: EntitySpendData,
|
||||
entityLabel: string,
|
||||
teamAliasMap: Record<string, string> = {},
|
||||
): any[] => {
|
||||
const dailyBreakdown: any[] = [];
|
||||
const includeFlatCost = hasFlatCost(spendData);
|
||||
|
||||
spendData.results.forEach((day) => {
|
||||
Object.entries(resolveEntities(day.breakdown)).forEach(([entity, data]: [string, any]) => {
|
||||
const { id, alias } = resolveEntityDisplay(entity, teamAliasMap, data.metadata);
|
||||
|
||||
dailyBreakdown.push({
|
||||
const row: Record<string, any> = {
|
||||
Date: day.date,
|
||||
[entityLabel]: alias,
|
||||
[`${entityLabel} ID`]: id,
|
||||
"Spend ($)": formatNumberWithCommas(data.metrics.spend, 4),
|
||||
Requests: data.metrics.api_requests,
|
||||
"Successful Requests": data.metrics.successful_requests,
|
||||
"Failed Requests": data.metrics.failed_requests,
|
||||
"Total Tokens": data.metrics.total_tokens,
|
||||
"Prompt Tokens": data.metrics.prompt_tokens || 0,
|
||||
"Completion Tokens": data.metrics.completion_tokens || 0,
|
||||
"Cache Read Input Tokens": data.metrics.cache_read_input_tokens || 0,
|
||||
"Cache Creation Input Tokens": data.metrics.cache_creation_input_tokens || 0,
|
||||
});
|
||||
};
|
||||
if (includeFlatCost) {
|
||||
const flatCost = data.metrics.flat_cost || 0;
|
||||
row["Flat Cost ($)"] = formatNumberWithCommas(flatCost, 4);
|
||||
row["Total Cost ($)"] = formatNumberWithCommas((data.metrics.spend || 0) + flatCost, 4);
|
||||
}
|
||||
row.Requests = data.metrics.api_requests;
|
||||
row["Successful Requests"] = data.metrics.successful_requests;
|
||||
row["Failed Requests"] = data.metrics.failed_requests;
|
||||
row["Total Tokens"] = data.metrics.total_tokens;
|
||||
row["Prompt Tokens"] = data.metrics.prompt_tokens || 0;
|
||||
row["Completion Tokens"] = data.metrics.completion_tokens || 0;
|
||||
row["Cache Read Input Tokens"] = data.metrics.cache_read_input_tokens || 0;
|
||||
row["Cache Creation Input Tokens"] = data.metrics.cache_creation_input_tokens || 0;
|
||||
dailyBreakdown.push(row);
|
||||
});
|
||||
});
|
||||
|
||||
|
|
@ -339,23 +350,31 @@ export const generateMetadata = (
|
|||
selectedFilters: string[],
|
||||
exportScope: ExportScope,
|
||||
spendData: EntitySpendData,
|
||||
): ExportMetadata => ({
|
||||
export_date: new Date().toISOString(),
|
||||
entity_type: entityType,
|
||||
date_range: {
|
||||
from: dateRange.from?.toISOString(),
|
||||
to: dateRange.to?.toISOString(),
|
||||
},
|
||||
filters_applied: selectedFilters.length > 0 ? selectedFilters : "None",
|
||||
export_scope: exportScope,
|
||||
summary: {
|
||||
): ExportMetadata => {
|
||||
const summary: ExportMetadata["summary"] = {
|
||||
total_spend: spendData.metadata.total_spend,
|
||||
total_requests: spendData.metadata.total_api_requests,
|
||||
successful_requests: spendData.metadata.total_successful_requests,
|
||||
failed_requests: spendData.metadata.total_failed_requests,
|
||||
total_tokens: spendData.metadata.total_tokens,
|
||||
},
|
||||
});
|
||||
};
|
||||
if (hasFlatCost(spendData)) {
|
||||
const flatCost = spendData.metadata.total_flat_cost ?? 0;
|
||||
summary.total_flat_cost = flatCost;
|
||||
summary.total_cost = spendData.metadata.total_spend + flatCost;
|
||||
}
|
||||
return {
|
||||
export_date: new Date().toISOString(),
|
||||
entity_type: entityType,
|
||||
date_range: {
|
||||
from: dateRange.from?.toISOString(),
|
||||
to: dateRange.to?.toISOString(),
|
||||
},
|
||||
filters_applied: selectedFilters.length > 0 ? selectedFilters : "None",
|
||||
export_scope: exportScope,
|
||||
summary,
|
||||
};
|
||||
};
|
||||
|
||||
export const handleExportCSV = (
|
||||
spendData: EntitySpendData,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
export interface SpendMetrics {
|
||||
spend: number;
|
||||
flat_cost?: number;
|
||||
prompt_tokens: number;
|
||||
completion_tokens: number;
|
||||
total_tokens: number;
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import React from "react";
|
||||
import { Form, Switch, Select, Tooltip } from "antd";
|
||||
import { Form, Switch, Select, Tooltip, DatePicker } from "antd";
|
||||
import { Text, Accordion, AccordionHeader, AccordionBody, TextInput } from "@tremor/react";
|
||||
import { Row, Col, Typography } from "antd";
|
||||
import TextArea from "antd/es/input/TextArea";
|
||||
|
|
@ -9,6 +9,17 @@ import CacheControlSettings from "./cache_control_settings";
|
|||
import VectorStoreSelector from "../vector_store_management/VectorStoreSelector";
|
||||
import { Tag } from "../tag_management/types";
|
||||
import { formItemValidateJSON } from "../../utils/textUtils";
|
||||
import {
|
||||
PTU_COUNT_FIELD,
|
||||
PTU_RATE_FIELD,
|
||||
PTU_START_FIELD,
|
||||
ptuCountRules,
|
||||
ptuPairRule,
|
||||
ptuRateRules,
|
||||
ptuStartRequiredRule,
|
||||
ptuWindowOrderRule,
|
||||
PTU_END_FIELD,
|
||||
} from "../../utils/ptuValidation";
|
||||
const { Link } = Typography;
|
||||
|
||||
interface AdvancedSettingsProps {
|
||||
|
|
@ -182,6 +193,50 @@ const AdvancedSettings: React.FC<AdvancedSettingsProps> = ({
|
|||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label="PTU Count"
|
||||
name={PTU_COUNT_FIELD}
|
||||
dependencies={[PTU_RATE_FIELD]}
|
||||
rules={[{ validator: validateNumber }, ...ptuCountRules, ptuPairRule(PTU_RATE_FIELD)]}
|
||||
tooltip="Provisioned throughput units for this deployment. Set together with Cost per PTU / Hour and a Team to attribute a flat daily cost."
|
||||
className="mb-4"
|
||||
>
|
||||
<TextInput placeholder="e.g. 15" />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label="Calculated Cost per PTU / Hour (USD)"
|
||||
name={PTU_RATE_FIELD}
|
||||
dependencies={[PTU_COUNT_FIELD]}
|
||||
rules={[{ validator: validateNumber }, ...ptuRateRules, ptuPairRule(PTU_COUNT_FIELD)]}
|
||||
tooltip="Flat cost = PTU count * this rate * active hours, attributed to the deployment's team."
|
||||
className="mb-4"
|
||||
>
|
||||
<TextInput placeholder="e.g. 2.00" />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label="PTU Effective From (UTC)"
|
||||
name={PTU_START_FIELD}
|
||||
dependencies={[PTU_COUNT_FIELD, PTU_END_FIELD]}
|
||||
rules={[ptuStartRequiredRule(PTU_COUNT_FIELD), ptuWindowOrderRule(PTU_END_FIELD, "start")]}
|
||||
tooltip="Start of the PTU window, required when PTU Count is set. Flat cost accrues by the hour within the window; a window opening at 23:00 charges one hour that day."
|
||||
className="mb-4"
|
||||
>
|
||||
<DatePicker showTime style={{ width: "100%" }} />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label="PTU Effective To (UTC)"
|
||||
name={PTU_END_FIELD}
|
||||
dependencies={[PTU_START_FIELD]}
|
||||
rules={[ptuWindowOrderRule(PTU_START_FIELD, "end")]}
|
||||
tooltip="Optional end of the PTU window (exclusive). Leave blank for open-ended."
|
||||
className="mb-4"
|
||||
>
|
||||
<DatePicker showTime style={{ width: "100%" }} />
|
||||
</Form.Item>
|
||||
|
||||
{customPricing && (
|
||||
<div className="ml-6 pl-4 border-l-2 border-gray-200">
|
||||
<Form.Item label="Pricing Model" name="pricing_model" className="mb-4">
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import NotificationManager from "../molecules/notifications_manager";
|
||||
import { Model, modelCreateCall } from "../networking";
|
||||
import { provider_map } from "../provider_info_helpers";
|
||||
import { ptuPickerToUtcIso } from "../../utils/ptuDatetime";
|
||||
|
||||
export const prepareModelAddRequest = async (formValues: Record<string, any>, accessToken: string, form: any) => {
|
||||
try {
|
||||
|
|
@ -163,6 +164,23 @@ export const prepareModelAddRequest = async (formValues: Record<string, any>, ac
|
|||
continue;
|
||||
}
|
||||
|
||||
// Handle the PTU flat-cost fields (attributed to the team via model_info)
|
||||
else if (key === "ptu_count" || key === "cost_per_ptu_per_hour") {
|
||||
if (value !== undefined && value !== null && value !== "") {
|
||||
modelInfoObj[key] = Number(value);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
// Handle the PTU effective window (DatePicker dayjs value -> ISO 8601 UTC string)
|
||||
else if (key === "ptu_effective_from" || key === "ptu_effective_to") {
|
||||
const iso = ptuPickerToUtcIso(value);
|
||||
if (iso !== null) {
|
||||
modelInfoObj[key] = iso;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
// Check if key is any of the specified API related keys
|
||||
else {
|
||||
// Add key-value pair to litellm_params dictionary
|
||||
|
|
|
|||
|
|
@ -17,7 +17,19 @@ import {
|
|||
Title,
|
||||
Button as TremorButton,
|
||||
} from "@tremor/react";
|
||||
import { Button, Form, Input, Modal, Select, Tooltip } from "antd";
|
||||
import { Button, DatePicker, Form, Input, Modal, Select, Tooltip } from "antd";
|
||||
import { formatPtuUtcDisplay, ptuPickerToUtcIso, utcIsoToPickerValue } from "../utils/ptuDatetime";
|
||||
import {
|
||||
PTU_COUNT_FIELD,
|
||||
PTU_RATE_FIELD,
|
||||
ptuCountRules,
|
||||
ptuPairRule,
|
||||
ptuRateRules,
|
||||
ptuStartRequiredRule,
|
||||
ptuWindowOrderRule,
|
||||
PTU_END_FIELD,
|
||||
PTU_START_FIELD,
|
||||
} from "../utils/ptuValidation";
|
||||
import VectorStoreSelector from "./vector_store_management/VectorStoreSelector";
|
||||
import { CheckIcon, CopyIcon } from "lucide-react";
|
||||
import { useEffect, useMemo, useState } from "react";
|
||||
|
|
@ -67,6 +79,62 @@ interface ModelInfoViewProps {
|
|||
modelAccessGroups: string[] | null;
|
||||
}
|
||||
|
||||
interface PtuEditField {
|
||||
name: string;
|
||||
label: string;
|
||||
input: "number" | "datetime";
|
||||
placeholder?: string;
|
||||
isCount?: boolean;
|
||||
isRate?: boolean;
|
||||
isStart?: boolean;
|
||||
pairedWith?: string;
|
||||
windowPeer?: string;
|
||||
bound?: "start" | "end";
|
||||
}
|
||||
|
||||
const PTU_EDIT_FIELDS: PtuEditField[] = [
|
||||
{
|
||||
name: PTU_COUNT_FIELD,
|
||||
label: "PTU Count",
|
||||
input: "number",
|
||||
placeholder: "e.g. 15",
|
||||
isCount: true,
|
||||
pairedWith: PTU_RATE_FIELD,
|
||||
},
|
||||
{
|
||||
name: PTU_RATE_FIELD,
|
||||
label: "Cost per PTU / Hour (USD)",
|
||||
input: "number",
|
||||
placeholder: "e.g. 2.00",
|
||||
isRate: true,
|
||||
pairedWith: PTU_COUNT_FIELD,
|
||||
},
|
||||
{
|
||||
name: PTU_START_FIELD,
|
||||
label: "PTU Effective From (UTC)",
|
||||
input: "datetime",
|
||||
isStart: true,
|
||||
windowPeer: PTU_END_FIELD,
|
||||
bound: "start",
|
||||
},
|
||||
{
|
||||
name: PTU_END_FIELD,
|
||||
label: "PTU Effective To (UTC)",
|
||||
input: "datetime",
|
||||
windowPeer: PTU_START_FIELD,
|
||||
bound: "end",
|
||||
},
|
||||
];
|
||||
|
||||
const ptuFieldDependencies = ({ isStart, pairedWith, windowPeer }: PtuEditField): string[] | undefined => {
|
||||
const deps = [
|
||||
...(isStart ? [PTU_COUNT_FIELD] : []),
|
||||
...(pairedWith ? [pairedWith] : []),
|
||||
...(windowPeer ? [windowPeer] : []),
|
||||
];
|
||||
return deps.length ? deps : undefined;
|
||||
};
|
||||
|
||||
interface ComplexityRouterTierConfig {
|
||||
tiers?: {
|
||||
SIMPLE?: unknown;
|
||||
|
|
@ -427,6 +495,15 @@ export default function ModelInfoView({
|
|||
health_check_model: values.health_check_model,
|
||||
};
|
||||
}
|
||||
const ptuNumber = (val: string | number | null | undefined): number | null =>
|
||||
val !== undefined && val !== null && val !== "" ? Number(val) : null;
|
||||
updatedModelInfo = {
|
||||
...updatedModelInfo,
|
||||
ptu_count: ptuNumber(values.ptu_count),
|
||||
cost_per_ptu_per_hour: ptuNumber(values.cost_per_ptu_per_hour),
|
||||
ptu_effective_from: ptuPickerToUtcIso(values.ptu_effective_from),
|
||||
ptu_effective_to: ptuPickerToUtcIso(values.ptu_effective_to),
|
||||
};
|
||||
} catch (e) {
|
||||
NotificationsManager.fromBackend("Invalid JSON in Model Info");
|
||||
return;
|
||||
|
|
@ -769,6 +846,10 @@ export default function ModelInfoView({
|
|||
output_cost: localModelData.litellm_params?.output_cost_per_token
|
||||
? localModelData.litellm_params.output_cost_per_token * 1_000_000
|
||||
: localModelData.model_info?.output_cost_per_token * 1_000_000 || null,
|
||||
ptu_count: localModelData.model_info?.ptu_count ?? null,
|
||||
cost_per_ptu_per_hour: localModelData.model_info?.cost_per_ptu_per_hour ?? null,
|
||||
ptu_effective_from: utcIsoToPickerValue(localModelData.model_info?.ptu_effective_from),
|
||||
ptu_effective_to: utcIsoToPickerValue(localModelData.model_info?.ptu_effective_to),
|
||||
cache_read_cost:
|
||||
localModelData.litellm_params?.cache_read_input_token_cost !== undefined &&
|
||||
localModelData.litellm_params?.cache_read_input_token_cost !== null
|
||||
|
|
@ -872,6 +953,46 @@ export default function ModelInfoView({
|
|||
)}
|
||||
</div>
|
||||
|
||||
{PTU_EDIT_FIELDS.map((ptuField) => {
|
||||
const { name, label, input, placeholder, isCount, isRate, isStart, pairedWith } = ptuField;
|
||||
const { windowPeer, bound } = ptuField;
|
||||
return (
|
||||
<div key={name}>
|
||||
<Text className="font-medium">{label}</Text>
|
||||
{isEditing ? (
|
||||
<Form.Item
|
||||
name={name}
|
||||
className="mb-0"
|
||||
dependencies={ptuFieldDependencies(ptuField)}
|
||||
rules={[
|
||||
...(isCount ? ptuCountRules : []),
|
||||
...(isRate ? ptuRateRules : []),
|
||||
...(isStart ? [ptuStartRequiredRule(PTU_COUNT_FIELD)] : []),
|
||||
...(pairedWith ? [ptuPairRule(pairedWith)] : []),
|
||||
...(windowPeer && bound ? [ptuWindowOrderRule(windowPeer, bound)] : []),
|
||||
]}
|
||||
>
|
||||
{input === "number" ? (
|
||||
<NumericalInput
|
||||
placeholder={placeholder}
|
||||
step={isCount ? 1 : undefined}
|
||||
min={isCount ? 1 : 0}
|
||||
/>
|
||||
) : (
|
||||
<DatePicker showTime style={{ width: "100%" }} />
|
||||
)}
|
||||
</Form.Item>
|
||||
) : (
|
||||
<div className="mt-1 p-2 bg-gray-50 rounded-sm">
|
||||
{(input === "datetime"
|
||||
? formatPtuUtcDisplay(localModelData?.model_info?.[name])
|
||||
: localModelData?.model_info?.[name]) ?? "Not Set"}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
|
||||
<div>
|
||||
<Text className="font-medium">Cache Read Cost (per 1M tokens)</Text>
|
||||
{isEditing ? (
|
||||
|
|
|
|||
18
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
18
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -24252,6 +24252,11 @@ export interface components {
|
|||
* @default 0
|
||||
*/
|
||||
total_failed_requests: number;
|
||||
/**
|
||||
* Total Flat Cost
|
||||
* @default 0
|
||||
*/
|
||||
total_flat_cost: number;
|
||||
/**
|
||||
* Total Pages
|
||||
* @default 1
|
||||
|
|
@ -32512,6 +32517,11 @@ export interface components {
|
|||
* @default 0
|
||||
*/
|
||||
failed_requests: number;
|
||||
/**
|
||||
* Flat Cost
|
||||
* @default 0
|
||||
*/
|
||||
flat_cost: number;
|
||||
/**
|
||||
* Prompt Caching Savings Spend
|
||||
* @default 0
|
||||
|
|
@ -35304,6 +35314,8 @@ export interface components {
|
|||
cache_creation_input_token_cost?: number | null;
|
||||
/** Cache Read Input Token Cost */
|
||||
cache_read_input_token_cost?: number | null;
|
||||
/** Cost Per Ptu Per Hour */
|
||||
cost_per_ptu_per_hour?: number | null;
|
||||
/** Created At */
|
||||
created_at?: string | null;
|
||||
/** Created By */
|
||||
|
|
@ -35323,6 +35335,12 @@ export interface components {
|
|||
output_cost_per_character?: number | null;
|
||||
/** Output Cost Per Token */
|
||||
output_cost_per_token?: number | null;
|
||||
/** Ptu Count */
|
||||
ptu_count?: number | null;
|
||||
/** Ptu Effective From */
|
||||
ptu_effective_from?: string | null;
|
||||
/** Ptu Effective To */
|
||||
ptu_effective_to?: string | null;
|
||||
/** Team Id */
|
||||
team_id?: string | null;
|
||||
/** Team Public Model Name */
|
||||
|
|
|
|||
96
ui/litellm-dashboard/src/utils/ptuDatetime.test.ts
Normal file
96
ui/litellm-dashboard/src/utils/ptuDatetime.test.ts
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
import dayjs from "dayjs";
|
||||
import { describe, expect, it } from "vitest";
|
||||
import { formatPtuUtcDisplay, ptuPickerToUtcIso, utcIsoToPickerValue } from "./ptuDatetime";
|
||||
|
||||
describe("ptuDatetime", () => {
|
||||
it("stores the picked wall-clock time as UTC instead of shifting across zones", () => {
|
||||
const picked = dayjs("2024-03-10T23:00:00");
|
||||
expect(ptuPickerToUtcIso(picked)).toBe("2024-03-10T23:00:00.000Z");
|
||||
});
|
||||
|
||||
it("returns null for empty picker values", () => {
|
||||
expect(ptuPickerToUtcIso(null)).toBeNull();
|
||||
expect(ptuPickerToUtcIso(undefined)).toBeNull();
|
||||
});
|
||||
|
||||
it("round-trips a UTC ISO string back to the same wall-clock in the picker", () => {
|
||||
const value = utcIsoToPickerValue("2024-03-10T23:00:00.000Z");
|
||||
expect(value).not.toBeNull();
|
||||
expect(value!.format("YYYY-MM-DDTHH:mm:ss")).toBe("2024-03-10T23:00:00");
|
||||
expect(ptuPickerToUtcIso(value)).toBe("2024-03-10T23:00:00.000Z");
|
||||
});
|
||||
|
||||
it("returns null for empty ISO strings", () => {
|
||||
expect(utcIsoToPickerValue(null)).toBeNull();
|
||||
expect(utcIsoToPickerValue(undefined)).toBeNull();
|
||||
expect(utcIsoToPickerValue("")).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
describe("formatPtuUtcDisplay", () => {
|
||||
it("renders the two stored serialisations identically", () => {
|
||||
// the backend writes +00:00, a just-saved form holds the picker's .000Z
|
||||
expect(formatPtuUtcDisplay("2026-08-01T23:00:00+00:00")).toBe("2026-08-01 23:00:00 UTC");
|
||||
expect(formatPtuUtcDisplay("2026-08-01T23:00:00.000Z")).toBe("2026-08-01 23:00:00 UTC");
|
||||
});
|
||||
|
||||
it("shows the UTC instant regardless of the offset it was written with", () => {
|
||||
expect(formatPtuUtcDisplay("2026-08-01T16:00:00-07:00")).toBe("2026-08-01 23:00:00 UTC");
|
||||
});
|
||||
|
||||
it("returns null for empty values so the caller can fall back to Not Set", () => {
|
||||
expect(formatPtuUtcDisplay(null)).toBeNull();
|
||||
expect(formatPtuUtcDisplay(undefined)).toBeNull();
|
||||
expect(formatPtuUtcDisplay("")).toBeNull();
|
||||
});
|
||||
|
||||
it("passes an unparseable value through rather than hiding it", () => {
|
||||
expect(formatPtuUtcDisplay("not-a-date")).toBe("not-a-date");
|
||||
});
|
||||
});
|
||||
|
||||
describe("DST spring-forward gap", () => {
|
||||
// 2027-03-14 02:30 does not exist in America/Los_Angeles: the clock jumps 02:00 -> 03:00.
|
||||
const GAP_ISO = "2027-03-14T02:30:00+00:00";
|
||||
|
||||
it("keeps the stored wall clock when it falls in the local DST gap", () => {
|
||||
const picked = utcIsoToPickerValue(GAP_ISO);
|
||||
expect(picked).not.toBeNull();
|
||||
expect(picked!.format("YYYY-MM-DDTHH:mm:ss")).toBe("2027-03-14T02:30:00");
|
||||
});
|
||||
|
||||
it("round-trips the gap instant back out unchanged, so a save cannot shift it", () => {
|
||||
expect(ptuPickerToUtcIso(utcIsoToPickerValue(GAP_ISO))).toBe("2027-03-14T02:30:00.000Z");
|
||||
});
|
||||
|
||||
it("round-trips a fall-back ambiguous instant unchanged too", () => {
|
||||
// 2027-11-07 01:30 occurs twice in America/Los_Angeles
|
||||
const AMBIGUOUS = "2027-11-07T01:30:00+00:00";
|
||||
expect(ptuPickerToUtcIso(utcIsoToPickerValue(AMBIGUOUS))).toBe("2027-11-07T01:30:00.000Z");
|
||||
});
|
||||
|
||||
it("returns null for an unparseable stored value instead of an Invalid Date picker", () => {
|
||||
expect(utcIsoToPickerValue("not-a-date")).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
describe("sub-second precision", () => {
|
||||
// The backend persists value.isoformat() verbatim, so a window set out of band (curl,
|
||||
// which is this repo's documented setup path) can carry microseconds. Every save re-sends
|
||||
// both window fields, so a lossy round-trip would rewrite an untouched billing window.
|
||||
it("preserves a sub-second stored instant through a save", () => {
|
||||
const stored = "2026-08-01T23:00:00.500000+00:00";
|
||||
expect(ptuPickerToUtcIso(utcIsoToPickerValue(stored))).toBe("2026-08-01T23:00:00.500Z");
|
||||
});
|
||||
|
||||
it("keeps the sub-second component visible on the picker value", () => {
|
||||
expect(utcIsoToPickerValue("2026-08-01T23:00:00.500000+00:00")!.millisecond()).toBe(500);
|
||||
});
|
||||
|
||||
it("still reinterprets a freshly picked local-mode value as UTC", () => {
|
||||
// a value the operator picks has no sub-second part and must not be zone-converted
|
||||
const localPick = dayjs("2026-08-01T23:00:00");
|
||||
expect(localPick.isUTC()).toBe(false);
|
||||
expect(ptuPickerToUtcIso(localPick)).toBe("2026-08-01T23:00:00.000Z");
|
||||
});
|
||||
});
|
||||
61
ui/litellm-dashboard/src/utils/ptuDatetime.ts
Normal file
61
ui/litellm-dashboard/src/utils/ptuDatetime.ts
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
import dayjs, { Dayjs } from "dayjs";
|
||||
import utc from "dayjs/plugin/utc";
|
||||
|
||||
dayjs.extend(utc);
|
||||
|
||||
const WALL_CLOCK_FORMAT = "YYYY-MM-DDTHH:mm:ss";
|
||||
|
||||
/**
|
||||
* Read the wall clock the picker is showing and stamp it as UTC.
|
||||
*
|
||||
* The picker value stays in UTC mode end to end (see `utcIsoToPickerValue`), so `.format()`
|
||||
* returns the digits the operator sees and no zone conversion happens on the way out.
|
||||
*/
|
||||
export const ptuPickerToUtcIso = (value: Dayjs | null | undefined): string | null => {
|
||||
if (!value || typeof value.format !== "function") {
|
||||
return null;
|
||||
}
|
||||
// A value that came from storage is already in UTC mode and holds the exact stored
|
||||
// instant. Every save re-sends both window fields, so routing it back through a
|
||||
// second-granularity wall clock would silently drop any sub-second component of a
|
||||
// window that was set out of band, turning an unrelated edit into a quiet rewrite.
|
||||
if (typeof value.isUTC === "function" && value.isUTC()) {
|
||||
return value.toISOString();
|
||||
}
|
||||
// A freshly picked value is in the browser's zone; its wall clock is what the operator
|
||||
// chose against a UTC-labelled field, so it is reinterpreted rather than converted.
|
||||
return dayjs.utc(value.format(WALL_CLOCK_FORMAT)).toISOString();
|
||||
};
|
||||
|
||||
/**
|
||||
* Hand the picker a UTC-mode Dayjs so it displays the stored wall clock verbatim.
|
||||
*
|
||||
* Re-parsing the wall clock in the browser's zone looks equivalent but is not: a clock reading
|
||||
* that does not exist locally, the hour a DST spring-forward skips, gets advanced by the engine.
|
||||
* `dayjs("2027-03-14T02:30:00")` is 03:30 in America/Los_Angeles, and because the save path
|
||||
* re-stamps whatever the picker holds as UTC, that shift would be written back to the stored
|
||||
* window rather than cancelled out.
|
||||
*/
|
||||
export const utcIsoToPickerValue = (iso: string | null | undefined): Dayjs | null => {
|
||||
if (!iso) {
|
||||
return null;
|
||||
}
|
||||
const parsed = dayjs.utc(iso);
|
||||
return parsed.isValid() ? parsed : null;
|
||||
};
|
||||
|
||||
const DISPLAY_FORMAT = "YYYY-MM-DD HH:mm:ss";
|
||||
|
||||
/**
|
||||
* Render a stored PTU timestamp for the read view. The backend serialises as `+00:00` while a
|
||||
* just-saved form holds the `.000Z` the picker produced, so the same instant would otherwise be
|
||||
* shown two different ways depending on whether the page has been reloaded since the edit. An
|
||||
* unparseable value is passed through rather than hidden.
|
||||
*/
|
||||
export const formatPtuUtcDisplay = (iso: string | null | undefined): string | null => {
|
||||
if (!iso) {
|
||||
return null;
|
||||
}
|
||||
const parsed = dayjs.utc(iso);
|
||||
return parsed.isValid() ? `${parsed.format(DISPLAY_FORMAT)} UTC` : String(iso);
|
||||
};
|
||||
170
ui/litellm-dashboard/src/utils/ptuValidation.test.ts
Normal file
170
ui/litellm-dashboard/src/utils/ptuValidation.test.ts
Normal file
|
|
@ -0,0 +1,170 @@
|
|||
import { describe, expect, it } from "vitest";
|
||||
import {
|
||||
PTU_COUNT_FIELD,
|
||||
PTU_RATE_FIELD,
|
||||
ptuCountRules,
|
||||
ptuPairRule,
|
||||
ptuRateRules,
|
||||
ptuStartRequiredRule,
|
||||
ptuWindowOrderRule,
|
||||
PTU_END_FIELD,
|
||||
PTU_START_FIELD,
|
||||
MAX_PTU_COUNT,
|
||||
MAX_COST_PER_PTU_PER_HOUR,
|
||||
} from "./ptuValidation";
|
||||
|
||||
const validate = (value: unknown) => ptuCountRules[0].validator(null, value);
|
||||
|
||||
describe("ptuCountRules", () => {
|
||||
it("accepts positive whole numbers and empty values", async () => {
|
||||
await expect(validate(5)).resolves.toBeUndefined();
|
||||
await expect(validate("15")).resolves.toBeUndefined();
|
||||
await expect(validate("")).resolves.toBeUndefined();
|
||||
await expect(validate(null)).resolves.toBeUndefined();
|
||||
await expect(validate(undefined)).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("rejects fractional values that the backend integer contract would refuse", async () => {
|
||||
await expect(validate(2.5)).rejects.toThrow("whole number between 1 and");
|
||||
await expect(validate("1.25")).rejects.toThrow("whole number between 1 and");
|
||||
});
|
||||
|
||||
it("rejects zero and negatives, which the backend rejects as a non-positive ptu_count", async () => {
|
||||
await expect(validate(0)).rejects.toThrow("whole number between 1 and");
|
||||
await expect(validate(-1)).rejects.toThrow("whole number between 1 and");
|
||||
await expect(validate("-3")).rejects.toThrow("whole number between 1 and");
|
||||
});
|
||||
|
||||
it("rejects a value that is not a number at all", async () => {
|
||||
await expect(validate("abc")).rejects.toThrow("whole number between 1 and");
|
||||
});
|
||||
});
|
||||
|
||||
describe("ptuPairRule", () => {
|
||||
const rule = (sibling: unknown) => ptuPairRule(PTU_RATE_FIELD)({ getFieldValue: () => sibling });
|
||||
const check = (value: unknown, sibling: unknown) => rule(sibling).validator(null, value);
|
||||
|
||||
it("accepts both set and both cleared, the only shapes the backend stores", async () => {
|
||||
await expect(check(10, 2.0)).resolves.toBeUndefined();
|
||||
await expect(check("", "")).resolves.toBeUndefined();
|
||||
await expect(check(null, undefined)).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("rejects a half-set pair, which the backend answers with a 400", async () => {
|
||||
await expect(check(10, "")).rejects.toThrow("must be set together");
|
||||
await expect(check("", 2.0)).rejects.toThrow("must be set together");
|
||||
await expect(check(null, 2.0)).rejects.toThrow("must be set together");
|
||||
});
|
||||
|
||||
it("reads the sibling by the field name it was given", () => {
|
||||
const seen: string[] = [];
|
||||
ptuPairRule(PTU_COUNT_FIELD)({
|
||||
getFieldValue: (name: string) => {
|
||||
seen.push(name);
|
||||
return 1;
|
||||
},
|
||||
}).validator(null, 1);
|
||||
expect(seen).toEqual([PTU_COUNT_FIELD]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("ptuRateRules", () => {
|
||||
const validate = (value: unknown) => ptuRateRules[0].validator(null, value);
|
||||
|
||||
it("rejects a negative rate, which the backend answers with a 400", async () => {
|
||||
await expect(validate(-1)).rejects.toThrow("must be between 0 and");
|
||||
});
|
||||
|
||||
it("rejects a negative rate typed as a string, which is what an input yields", async () => {
|
||||
await expect(validate("-0.5")).rejects.toThrow("must be between 0 and");
|
||||
});
|
||||
|
||||
it("allows zero, which the backend accepts", async () => {
|
||||
await expect(validate(0)).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("allows a fractional rate", async () => {
|
||||
await expect(validate(2.5)).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("leaves an empty field to the pair rule", async () => {
|
||||
await expect(validate("")).resolves.toBeUndefined();
|
||||
await expect(validate(null)).resolves.toBeUndefined();
|
||||
await expect(validate(undefined)).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("rejects a value that is not a number at all", async () => {
|
||||
await expect(validate("abc")).rejects.toThrow("must be between 0 and");
|
||||
});
|
||||
});
|
||||
|
||||
describe("ptuStartRequiredRule", () => {
|
||||
const rule = (count: unknown, start: unknown) =>
|
||||
ptuStartRequiredRule(PTU_COUNT_FIELD)({ getFieldValue: () => count }).validator(null, start);
|
||||
|
||||
it("rejects PTU config with no effective start, which the backend answers with a 400", async () => {
|
||||
await expect(rule(10, undefined)).rejects.toThrow("PTU Effective From is required when PTU Count is set");
|
||||
await expect(rule(10, "")).rejects.toThrow("PTU Effective From is required when PTU Count is set");
|
||||
});
|
||||
|
||||
it("allows a start once given", async () => {
|
||||
await expect(rule(10, "2026-08-01T00:00:00Z")).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("leaves a deployment with no PTU config alone", async () => {
|
||||
await expect(rule(undefined, undefined)).resolves.toBeUndefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("ptuWindowOrderRule", () => {
|
||||
const form = (values: Record<string, unknown>) => ({ getFieldValue: (name: string) => values[name] });
|
||||
const start = new Date("2026-08-01T00:00:00Z");
|
||||
const end = new Date("2026-09-01T00:00:00Z");
|
||||
|
||||
it("accepts an ordered window from either bound", async () => {
|
||||
await expect(
|
||||
ptuWindowOrderRule(PTU_END_FIELD, "start")(form({ [PTU_END_FIELD]: end })).validator(null, start),
|
||||
).resolves.toBeUndefined();
|
||||
await expect(
|
||||
ptuWindowOrderRule(PTU_START_FIELD, "end")(form({ [PTU_START_FIELD]: start })).validator(null, end),
|
||||
).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("rejects an inverted window from either bound", async () => {
|
||||
await expect(
|
||||
ptuWindowOrderRule(PTU_END_FIELD, "start")(form({ [PTU_END_FIELD]: start })).validator(null, end),
|
||||
).rejects.toThrow("PTU Effective To must be after PTU Effective From");
|
||||
await expect(
|
||||
ptuWindowOrderRule(PTU_START_FIELD, "end")(form({ [PTU_START_FIELD]: end })).validator(null, start),
|
||||
).rejects.toThrow("PTU Effective To must be after PTU Effective From");
|
||||
});
|
||||
|
||||
it("rejects a zero-length window, which the backend also refuses", async () => {
|
||||
await expect(
|
||||
ptuWindowOrderRule(PTU_END_FIELD, "start")(form({ [PTU_END_FIELD]: start })).validator(null, start),
|
||||
).rejects.toThrow("must be after");
|
||||
});
|
||||
|
||||
it("stays silent while either bound is empty, since the window is optional", async () => {
|
||||
await expect(ptuWindowOrderRule(PTU_END_FIELD, "start")(form({})).validator(null, start)).resolves.toBeUndefined();
|
||||
await expect(
|
||||
ptuWindowOrderRule(PTU_START_FIELD, "end")(form({ [PTU_START_FIELD]: start })).validator(null, ""),
|
||||
).resolves.toBeUndefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("backend maximums are mirrored in the form", () => {
|
||||
it("rejects a count above the cap and accepts one at it", async () => {
|
||||
await expect(ptuCountRules[0].validator(null, String(MAX_PTU_COUNT + 1))).rejects.toThrow("1,000,000");
|
||||
await expect(ptuCountRules[0].validator(null, String(MAX_PTU_COUNT))).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("rejects a rate above the cap and accepts one at it", async () => {
|
||||
await expect(ptuRateRules[0].validator(null, String(MAX_COST_PER_PTU_PER_HOUR + 1))).rejects.toThrow("1,000,000");
|
||||
await expect(ptuRateRules[0].validator(null, String(MAX_COST_PER_PTU_PER_HOUR))).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("still accepts a zero rate, which the backend allows", async () => {
|
||||
await expect(ptuRateRules[0].validator(null, "0")).resolves.toBeUndefined();
|
||||
});
|
||||
});
|
||||
116
ui/litellm-dashboard/src/utils/ptuValidation.ts
Normal file
116
ui/litellm-dashboard/src/utils/ptuValidation.ts
Normal file
|
|
@ -0,0 +1,116 @@
|
|||
interface ValidatorRule {
|
||||
validator: (rule: unknown, value: unknown) => Promise<void>;
|
||||
}
|
||||
|
||||
interface FormInstance {
|
||||
getFieldValue: (name: string) => unknown;
|
||||
}
|
||||
|
||||
export const PTU_COUNT_FIELD = "ptu_count";
|
||||
export const PTU_RATE_FIELD = "cost_per_ptu_per_hour";
|
||||
export const PTU_START_FIELD = "ptu_effective_from";
|
||||
export const PTU_END_FIELD = "ptu_effective_to";
|
||||
|
||||
// Mirrors ModelInfo.MAX_PTU_COUNT / MAX_COST_PER_PTU_PER_HOUR. Flat cost multiplies the
|
||||
// count by a float, so the backend caps both; without the same ceiling here the form
|
||||
// reports valid input and the save then fails with a 422 the operator cannot anticipate.
|
||||
export const MAX_PTU_COUNT = 1_000_000;
|
||||
export const MAX_COST_PER_PTU_PER_HOUR = 1_000_000;
|
||||
|
||||
const isFilled = (value: unknown): boolean => value !== undefined && value !== null && value !== "";
|
||||
|
||||
const isPositiveWholeNumber = (value: unknown): boolean => {
|
||||
if (!isFilled(value)) {
|
||||
return true;
|
||||
}
|
||||
const parsed = Number(value);
|
||||
return Number.isInteger(parsed) && parsed > 0 && parsed <= MAX_PTU_COUNT;
|
||||
};
|
||||
|
||||
/** Mirrors the backend contract, which rejects a ptu_count that is not a positive integer. */
|
||||
export const ptuCountRules: ValidatorRule[] = [
|
||||
{
|
||||
validator: (_, value) =>
|
||||
isPositiveWholeNumber(value)
|
||||
? Promise.resolve()
|
||||
: Promise.reject(new Error(`PTU Count must be a whole number between 1 and ${MAX_PTU_COUNT.toLocaleString()}`)),
|
||||
},
|
||||
];
|
||||
|
||||
const isNonNegativeNumber = (value: unknown): boolean => {
|
||||
if (!isFilled(value)) {
|
||||
return true;
|
||||
}
|
||||
const parsed = Number(value);
|
||||
return Number.isFinite(parsed) && parsed >= 0 && parsed <= MAX_COST_PER_PTU_PER_HOUR;
|
||||
};
|
||||
|
||||
/** Mirrors the backend contract, which rejects a negative cost_per_ptu_per_hour. */
|
||||
export const ptuRateRules: ValidatorRule[] = [
|
||||
{
|
||||
validator: (_, value) =>
|
||||
isNonNegativeNumber(value)
|
||||
? Promise.resolve()
|
||||
: Promise.reject(
|
||||
new Error(`Cost per PTU / Hour must be between 0 and ${MAX_COST_PER_PTU_PER_HOUR.toLocaleString()}`),
|
||||
),
|
||||
},
|
||||
];
|
||||
|
||||
/**
|
||||
* The backend rejects a half-set pair with "ptu_count and cost_per_ptu_per_hour must be set
|
||||
* together", so filling or clearing one field without the other is caught in the form. Pair
|
||||
* this with `dependencies` on the sibling field so its error clears when the pair is resolved.
|
||||
*/
|
||||
export const ptuPairRule =
|
||||
(siblingField: string) =>
|
||||
({ getFieldValue }: FormInstance): ValidatorRule => ({
|
||||
validator: (_, value) =>
|
||||
isFilled(value) === isFilled(getFieldValue(siblingField))
|
||||
? Promise.resolve()
|
||||
: Promise.reject(new Error("PTU Count and Cost per PTU / Hour must be set together")),
|
||||
});
|
||||
|
||||
/**
|
||||
* The backend requires an effective start whenever PTU is configured, since flat cost
|
||||
* accrues from that instant and an inferred start would bill days a deployment did not
|
||||
* exist. Pair this with `dependencies` on the count so the error clears when both resolve.
|
||||
*/
|
||||
export const ptuStartRequiredRule =
|
||||
(countField: string) =>
|
||||
({ getFieldValue }: FormInstance): ValidatorRule => ({
|
||||
validator: (_, value) =>
|
||||
isFilled(value) || !isFilled(getFieldValue(countField))
|
||||
? Promise.resolve()
|
||||
: Promise.reject(new Error("PTU Effective From is required when PTU Count is set")),
|
||||
});
|
||||
|
||||
/** Milliseconds for a picker value, which arrives as a Dayjs (or a Date/ISO string in tests). */
|
||||
const toEpochMs = (value: unknown): number => {
|
||||
const raw = (value as { valueOf?: () => unknown } | null)?.valueOf?.();
|
||||
const asNumber = Number(raw);
|
||||
return Number.isFinite(asNumber) ? asNumber : new Date(String(value)).getTime();
|
||||
};
|
||||
|
||||
/**
|
||||
* The backend rejects a window whose end is not strictly after its start, so an inverted or
|
||||
* zero-length window is caught in the form rather than answered with a 422 the operator
|
||||
* cannot anticipate. Pair this with `dependencies` on the sibling bound so the error clears
|
||||
* once the pair is ordered.
|
||||
*/
|
||||
export const ptuWindowOrderRule =
|
||||
(siblingField: string, thisBound: "start" | "end") =>
|
||||
({ getFieldValue }: FormInstance): ValidatorRule => ({
|
||||
validator: (_, value) => {
|
||||
const sibling = getFieldValue(siblingField);
|
||||
if (!isFilled(value) || !isFilled(sibling)) {
|
||||
return Promise.resolve();
|
||||
}
|
||||
const startMs = toEpochMs(thisBound === "start" ? value : sibling);
|
||||
const endMs = toEpochMs(thisBound === "start" ? sibling : value);
|
||||
if (Number.isNaN(startMs) || Number.isNaN(endMs) || endMs > startMs) {
|
||||
return Promise.resolve();
|
||||
}
|
||||
return Promise.reject(new Error("PTU Effective To must be after PTU Effective From"));
|
||||
},
|
||||
});
|
||||
Loading…
Add table
Reference in a new issue