chore: merge litellm_internal_staging into typing cleanup branch

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-08-10 21:49:03 +00:00
commit 22b7fe6c71
154 changed files with 17428 additions and 2371 deletions

View file

@ -0,0 +1,40 @@
name: "Cache Prisma binaries"
description: >-
Cache the Prisma CLI and engine binaries that `prisma generate` downloads, so
only the first job on a given prisma-client-py version pays for the download.
prisma-client-py shells out to `npm install prisma@<version>` whenever its
binary cache directory has no CLI entrypoint, which pulls ~85 MB of query and
schema engines over the network. That normally takes a few seconds, but it is
unbounded: one shard of a proxy-db run took 5m18s on that single step versus
3.8s on its eleven siblings, which pushed the job past its timeout and got a
fully passing test run cancelled.
Callers must not set PRISMA_BINARY_CACHE_DIR. The prisma-client-py default
(~/.cache/prisma-python/binaries/<prisma-version>/<engine-version>) is already
keyed by both versions, so a cache entry can never be served to a run that
expects different binaries.
runs:
using: composite
steps:
- name: Resolve prisma-client-py version
id: version
shell: bash
run: |
version="$(grep -A1 '^name = "prisma"$' uv.lock | sed -n 's/^version = "\(.*\)"$/\1/p' | head -1)"
if [ -z "${version}" ]; then
echo "could not resolve the prisma package version from uv.lock" >&2
exit 1
fi
echo "version=${version}" >> "$GITHUB_OUTPUT"
- name: Restore Prisma binaries
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
# ~/.cache/prisma-python holds the npm install tree prisma-client-py
# drives; ~/.cache/prisma is where @prisma/engines stages its downloads.
path: |
~/.cache/prisma-python
~/.cache/prisma
key: ${{ runner.os }}-prisma-binaries-${{ steps.version.outputs.version }}

View file

@ -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

View file

@ -18,10 +18,25 @@ on:
type: number
default: 2
timeout-minutes:
description: "Job timeout in minutes"
description: >-
Timeout for the test step alone. Setup (checkout, dependency install,
Prisma client generation) gets its own allowance on top, so a slow
runner or a cold binary download can never cancel passing tests.
required: false
type: number
default: 20
job-timeout-minutes:
description: >-
Backstop for the whole job. Keep it >= `timeout-minutes` plus 35: 30 for
the per-step ceilings on the setup steps below, and 5 for the runner
overhead the job clock charges but no step owns (job init, step
transitions, post-job cleanup). That headroom is what makes the test
budget a floor rather than a hope, since setup cannot overrun into it
without failing its own step first. GitHub expressions have no
arithmetic, so the sum is passed in rather than computed.
required: false
type: number
default: 55
max-failures:
description: "Stop after this many failures"
required: false
@ -44,30 +59,35 @@ jobs:
run:
name: Run tests
runs-on: ubuntu-latest
timeout-minutes: ${{ inputs.timeout-minutes }}
timeout-minutes: ${{ inputs.job-timeout-minutes }}
outputs:
decision: ${{ steps.changes.outputs.decision }}
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
timeout-minutes: 3
with:
persist-credentials: false
- name: Detect backend-relevant changes
id: changes
timeout-minutes: 2
uses: ./.github/actions/detect-backend-changes
- name: Set up Python
timeout-minutes: 3
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
timeout-minutes: 3
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
timeout-minutes: 5
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
@ -79,18 +99,24 @@ jobs:
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 8
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
timeout-minutes: 3
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Run tests
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: ${{ inputs.timeout-minutes }}
env:
TEST_PATH: ${{ inputs.test-path }}
MAX_FAILURES: ${{ inputs.max-failures }}

View file

@ -71,10 +71,12 @@ jobs:
if: steps.changes.outputs.relevant == 'true'
run: .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
if: steps.changes.outputs.relevant == 'true'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.relevant == 'true'
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Set up Node.js

View file

@ -57,9 +57,10 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -43,12 +43,13 @@ jobs:
with:
version: "0.10.9"
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
# The gate provisions its own measurement env (.venv-typecheck: a frozen
# uv sync of its canonical dependency groups plus a generated Prisma
# client), so no install step here can drift from what local runs measure.
- name: Emit basedpyright counts for HEAD
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
python scripts/type_check_gate.py --emit-counts-dir "$RUNNER_TEMP/basedpyright-counts"
counts_file=$(ls "$RUNNER_TEMP"/basedpyright-counts/basedpyright-counts-*.json)

View file

@ -65,6 +65,12 @@ jobs:
- name: check_provider_folders_documented
run: uv run --no-sync python ./tests/code_coverage_tests/check_provider_folders_documented.py
- name: check_prisma_binary_cache
run: uv run --no-sync python ./tests/code_coverage_tests/check_prisma_binary_cache.py
- name: check_workflow_startup_safety
run: uv run --no-sync python ./tests/code_coverage_tests/check_workflow_startup_safety.py
- name: router_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py

View file

@ -71,12 +71,13 @@ jobs:
run: |
uv sync --frozen --group proxy-dev --group e2e-dev
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
# basedpyright resolves Prisma's generated client (litellm/proxy/schema.prisma)
# only after `prisma generate` writes prisma/client.py et al. Without this the
# DB wrappers typed against the generated client would degrade to Unknown.
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
@ -119,7 +120,6 @@ jobs:
- name: Check basedpyright budget (delta vs base)
env:
GH_TOKEN: ${{ github.token }}
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync python scripts/type_check_gate.py --base "$GATE_BASE_SHA"

View file

@ -92,9 +92,10 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -65,10 +65,12 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -28,6 +28,10 @@ concurrency:
# Most of a shard's time is pytest plugin load + xdist worker imports +
# pytest-cov instrumentation, not the tests themselves. Keeping per-shard
# work low and matching worker count to runner cores is what controls it.
# * `timeout` bounds the pytest step only. Checkout, dependency install, and
# Prisma client generation draw on a separate allowance in the base
# workflow, so slow setup shows up as a slow job rather than as a
# cancelled shard whose tests were passing.
# * workers: 4 matches the 4-core ubuntu-latest runner. -n 8 on 4 cores
# oversubscribes 2x and workers fight for CPU during their cold-start
# imports (measured ~441% CPU for -n 8 locally, i.e. ~55% effective).

View file

@ -76,4 +76,5 @@ jobs:
workers: 4
reruns: 2
timeout-minutes: 60
job-timeout-minutes: 95
artifact-name: proxy-server

View file

@ -82,10 +82,12 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -51,9 +51,10 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra proxy
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -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

View file

@ -1,6 +1,6 @@
{
"reportAny": {
"limit": 26163
"limit": 26391
},
"reportArgumentType": {
"limit": 2614
@ -24,7 +24,7 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 8305
"limit": 8319
},
"reportFunctionMemberAccess": {
"limit": 7
@ -57,7 +57,7 @@
"limit": 5825
},
"reportMissingTypeArgument": {
"limit": 15693
"limit": 15695
},
"reportMissingTypeStubs": {
"limit": 40
@ -105,13 +105,13 @@
"limit": 113
},
"reportUnknownMemberType": {
"limit": 39634
"limit": 39643
},
"reportUnknownParameterType": {
"limit": 20132
},
"reportUnknownVariableType": {
"limit": 31144
"limit": 31153
},
"reportUnnecessaryCast": {
"limit": 118

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_DailyTeamSpend" ADD COLUMN IF NOT EXISTS "ptu_flat_cost" DOUBLE PRECISION NOT NULL DEFAULT 0.0;

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -1493,6 +1493,8 @@ SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INT
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500")))
RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN", "100")))
PROXY_BATCH_POLLING_INTERVAL: Final = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600))
MAX_OBJECTS_PER_POLL_CYCLE: Final = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50)))
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS: Final = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7)))
@ -1719,3 +1721,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

View file

@ -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,
),

View file

@ -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,
)

View file

@ -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."""

File diff suppressed because it is too large Load diff

View file

@ -1,7 +1,7 @@
import enum
import json
import os
from collections.abc import Callable
from collections.abc import Callable, Mapping
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal
@ -11,6 +11,7 @@ from pydantic import (
ConfigDict,
Field,
Json,
PositiveInt,
field_validator,
model_validator,
)
@ -1102,6 +1103,8 @@ class AllowedVectorStoreIndexItem(LiteLLMPydanticObjectBase):
class KeyRequestBase(GenerateRequestBase):
key: str | None = None
default_estimated_output_tokens: PositiveInt | None = None
default_estimated_output_tokens_per_model: Mapping[str, PositiveInt] | None = None
budget_id: str | None = None
tags: list[str] | None = None
disable_global_guardrails: bool | None = None
@ -1819,6 +1822,8 @@ class NewTeamRequest(TeamBase):
)
model_tpm_limit: dict[str, int] | None = None
default_estimated_output_tokens: PositiveInt | None = None
default_estimated_output_tokens_per_model: Mapping[str, PositiveInt] | None = None
mcp_rpm_limit: dict[str, int] | None = None
team_member_budget: float | None = None # allow user to set a budget for all team members
team_member_rpm_limit: int | None = None # allow user to set RPM limit for all team members
@ -1883,6 +1888,8 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
prompts: list[str] | None = None
model_rpm_limit: dict[str, int] | None = None
model_tpm_limit: dict[str, int] | None = None
default_estimated_output_tokens: PositiveInt | None = None
default_estimated_output_tokens_per_model: Mapping[str, PositiveInt] | None = None
mcp_rpm_limit: dict[str, int] | None = None
allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None
enforced_batch_output_expires_after: dict | None = None
@ -4018,6 +4025,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):
@ -4097,6 +4110,8 @@ class PassThroughEndpointLoggingTypedDict(TypedDict):
LiteLLM_ManagementEndpoint_MetadataFields: Final = [
"model_rpm_limit",
"model_tpm_limit",
"default_estimated_output_tokens",
"default_estimated_output_tokens_per_model",
"mcp_rpm_limit",
"tag_rpm_limit",
"rpm_limit_type",

View file

@ -1,12 +1,13 @@
import os
import re
import sys
from collections.abc import Iterator, Mapping
from collections.abc import Collection, Iterator, Mapping
from functools import lru_cache
from logging import Logger
from typing import Any, Final
from typing import Any, Final, Protocol
from fastapi import HTTPException, Request, status
from pydantic import PositiveInt, TypeAdapter, ValidationError
import litellm
from litellm import Router, provider_list
@ -999,6 +1000,167 @@ def get_key_model_tpm_limit(
return None
ESTIMATED_OUTPUT_TOKENS_FIELD: Final = "default_estimated_output_tokens"
ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD: Final = "default_estimated_output_tokens_per_model"
ESTIMATED_OUTPUT_TOKENS_METADATA_FIELDS: Final = frozenset(
{ESTIMATED_OUTPUT_TOKENS_FIELD, ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD}
)
_ESTIMATED_OUTPUT_TOKENS_ADAPTER: Final = TypeAdapter(PositiveInt)
_ESTIMATED_OUTPUT_TOKENS_PER_MODEL_ADAPTER: Final = TypeAdapter(Mapping[str, PositiveInt])
def _validated_output_token_estimate(raw: object) -> int | None:
"""Coerce one declared estimate to a positive int, or ignore it."""
if raw is None:
return None
try:
return _ESTIMATED_OUTPUT_TOKENS_ADAPTER.validate_python(raw)
except ValidationError as validation_error:
verbose_proxy_logger.warning(
"Ignoring malformed %s in metadata: %s",
ESTIMATED_OUTPUT_TOKENS_FIELD,
validation_error,
)
return None
def _validated_output_token_estimates_per_model(raw: object) -> Mapping[str, int] | None:
"""Coerce a declared per-model estimate map, or ignore it."""
if raw is None:
return None
try:
return _ESTIMATED_OUTPUT_TOKENS_PER_MODEL_ADAPTER.validate_python(raw)
except ValidationError as validation_error:
verbose_proxy_logger.warning(
"Ignoring malformed %s in metadata: %s",
ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD,
validation_error,
)
return None
def _estimated_output_tokens_from_metadata(
metadata: Mapping[str, Any] | None,
model_name: str | None,
) -> int | None:
"""Resolve the per-model, then global, estimate out of one metadata blob.
The two fields are validated independently so a malformed per-model map
cannot discard a valid global estimate, or the other way round.
"""
if not metadata or ESTIMATED_OUTPUT_TOKENS_METADATA_FIELDS.isdisjoint(metadata):
return None
if model_name is not None:
per_model: Final = _validated_output_token_estimates_per_model(
metadata.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD)
)
per_model_estimate: Final = per_model.get(model_name) if per_model is not None else None
if per_model_estimate is not None:
return per_model_estimate
return _validated_output_token_estimate(metadata.get(ESTIMATED_OUTPUT_TOKENS_FIELD))
def get_estimated_output_tokens(
user_api_key_dict: UserAPIKeyAuth,
model_name: str | None = None,
) -> int | None:
"""Resolve the operator-declared output-token estimate for TPM reservation.
Priority order (returns first found):
1. Key metadata ``default_estimated_output_tokens_per_model[model_name]``
2. Key metadata ``default_estimated_output_tokens``
3. Team metadata ``default_estimated_output_tokens_per_model[model_name]``
4. Team metadata ``default_estimated_output_tokens``
Returns ``None`` when nothing is configured, which leaves the static
heuristic floor in place.
"""
key_estimate: Final = _estimated_output_tokens_from_metadata(user_api_key_dict.metadata, model_name)
if key_estimate is not None:
return key_estimate
return _estimated_output_tokens_from_metadata(user_api_key_dict.team_metadata, model_name)
class OutputTokenEstimateRequest(Protocol):
"""The shape of any management request that can carry an output-token estimate.
Read-only members: the gate inspects a request, it never writes one back.
"""
@property
def metadata(self) -> Mapping[str, object] | None: ...
@property
def default_estimated_output_tokens(self) -> int | None: ...
@property
def default_estimated_output_tokens_per_model(self) -> Mapping[str, int] | None: ...
@property
def model_fields_set(self) -> Collection[str]: ...
def _requested_output_token_estimates(
data: OutputTokenEstimateRequest,
existing_metadata: Mapping[str, object],
) -> tuple[object, object]:
"""The output-token estimates this request would leave stored on the entity.
Mirrors how the management endpoints merge metadata: a supplied ``metadata``
replaces the stored blob wholesale, an omitted one preserves it, and the
dedicated top-level fields overlay whatever survives. Both sources are read
because the same declaration reaches the same stored field either way.
"""
base: Final[Mapping[str, object]] = (
(data.metadata or {}) if "metadata" in data.model_fields_set else existing_metadata
)
return (
data.default_estimated_output_tokens
if data.default_estimated_output_tokens is not None
else base.get(ESTIMATED_OUTPUT_TOKENS_FIELD),
data.default_estimated_output_tokens_per_model
if data.default_estimated_output_tokens_per_model is not None
else base.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD),
)
def enforce_output_token_estimates_are_admin_only(
data: OutputTokenEstimateRequest,
existing_metadata: Mapping[str, object] | None,
user_api_key_dict: UserAPIKeyAuth,
entity: Literal["key", "team"],
) -> None:
"""Only a proxy admin may change what a key or team declares its models emit.
That declaration is what the TPM limiter reserves for a request omitting
``max_tokens``, so lowering or clearing it under-reserves against every
window the request is charged against, including the team and organization
ones the writer may not own. A key's metadata is writable by its holder and
a team's by its team admin, so neither is a trustworthy source for a value
that weakens a limit set above them. Gated on the resulting value rather
than on presence, so a form resending the stored declaration stays a no-op.
"""
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return
stored: Final[Mapping[str, object]] = existing_metadata or {}
if _requested_output_token_estimates(data, stored) == (
stored.get(ESTIMATED_OUTPUT_TOKENS_FIELD),
stored.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD),
):
return
raise HTTPException(
status_code=403,
detail={
"error": f"Only proxy admins can set {ESTIMATED_OUTPUT_TOKENS_FIELD} or "
f"{ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD} on a {entity}. They decide how many output tokens "
"the rate limiter reserves for a request that omits max_tokens."
},
)
def get_model_rate_limit_from_metadata(
user_api_key_dict: UserAPIKeyAuth,
metadata_accessor_key: Literal["team_metadata", "organization_metadata", "project_metadata"],

View file

@ -1,14 +1,21 @@
import asyncio
import json
import time
from collections.abc import Callable, Sequence
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Final, Literal, Protocol, TypeVar
from types import MappingProxyType
from typing import Final, Literal, Protocol, TypeVar, assert_never
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
from litellm.constants import GLOBAL_PROXY_SPEND_CACHE_KEY, LITELLM_PROXY_BUDGET_NAME
from litellm.constants import (
GLOBAL_PROXY_SPEND_CACHE_KEY,
LITELLM_PROXY_BUDGET_NAME,
RESET_BUDGET_JOB_BATCH_SIZE,
RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN,
)
from litellm.proxy._types import (
LiteLLM_BudgetTableFull,
LiteLLM_EndUserTable,
@ -30,7 +37,10 @@ from litellm.repositories.table_repositories import (
TeamMembershipRepository,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.unit_of_work import spend_reset_unit_of_work
from litellm.repositories.unit_of_work import (
budget_cascade_unit_of_work,
spend_reset_unit_of_work,
)
from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
@ -38,6 +48,9 @@ from litellm.types.services import ServiceTypes
_RowT = TypeVar("_RowT")
_LINKED_KEYS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"budget_duration": None, "spend": {"gt": 0}})
_SPENT_ROWS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"spend": {"gt": 0}})
class _TeamMembershipRow(Protocol):
@property
@ -62,39 +75,130 @@ class _TagRow(Protocol):
def tag_name(self) -> str: ...
class _EndUserRow(Protocol):
@property
def user_id(self) -> str: ...
def _team_membership_counter_key(row: _TeamMembershipRow) -> str:
return f"spend:team_member:{row.user_id}:{row.team_id}"
def _team_membership_cache_key(row: _TeamMembershipRow) -> str:
return f"{row.team_id}_{row.user_id}"
def _team_membership_cache_keys(row: _TeamMembershipRow) -> tuple[str, ...]:
return (f"{row.team_id}_{row.user_id}",)
def _key_counter_key(row: _KeyRow) -> str:
return f"spend:key:{row.token}"
def _key_cache_key(row: _KeyRow) -> str:
return row.token
def _key_cache_keys(row: _KeyRow) -> tuple[str, ...]:
return (row.token,)
def _org_counter_key(row: _OrgRow) -> str:
return f"spend:org:{row.organization_id}"
def _org_cache_keys(row: _OrgRow) -> Sequence[str]:
return [
def _org_cache_keys(row: _OrgRow) -> tuple[str, ...]:
return (
f"org_id:{row.organization_id}",
f"org_id:{row.organization_id}:with_budget",
]
)
def _tag_counter_key(row: _TagRow) -> str:
return f"spend:tag:{row.tag_name}"
def _tag_cache_key(row: _TagRow) -> str:
return f"tag:{row.tag_name}"
def _tag_cache_keys(row: _TagRow) -> tuple[str, ...]:
return (f"tag:{row.tag_name}",)
def _budget_link_where(
budget_ids: Sequence[str],
extra: Mapping[str, object] = MappingProxyType({}),
) -> dict[str, object]:
return {"budget_id": {"in": list(budget_ids)}, **extra}
@dataclass(frozen=True, slots=True)
class _BudgetCascade:
"""Everything one budget-tier reset touches, resolved before any write."""
budgets: tuple[LiteLLM_BudgetTableFull, ...] = ()
budget_ids: tuple[str, ...] = ()
budget_resets: tuple[tuple[str, datetime], ...] = ()
endusers: tuple[_EndUserRow, ...] = ()
counter_keys: tuple[str, ...] = ()
cache_keys: tuple[str, ...] = ()
@dataclass(frozen=True, slots=True)
class _BudgetCascadeCommitted:
cascade: _BudgetCascade
advanced: int
@dataclass(frozen=True, slots=True)
class _BudgetCascadeFailed:
cascade: _BudgetCascade
error: Exception
_EMPTY_CASCADE: Final = _BudgetCascade()
@dataclass(frozen=True, slots=True)
class _ChunkOutcome:
"""One chunk of a reset phase: rows read, and rows whose new budget_reset_at
cleared the due cutoff. Anything else is still due and would come straight
back on the next fetch, so it is not progress."""
fetched: int
advanced: int
_NO_PROGRESS: Final = _ChunkOutcome(fetched=0, advanced=0)
def _as_utc(moment: datetime) -> datetime:
return moment if moment.tzinfo is not None else moment.replace(tzinfo=timezone.utc)
def _count_advanced(reset_ats: Iterable[object], cutoff: datetime) -> int:
"""How many rows the write actually moved past the due cutoff.
A budget_duration of "0s" (or one the parser cannot read) resolves to the
current time, so the row is written and stays due. Counting it as progress
would re-read the same chunk until the per-run cap on every tick.
"""
utc_cutoff: Final = _as_utc(cutoff)
return sum(1 for reset_at in reset_ats if isinstance(reset_at, datetime) and _as_utc(reset_at) > utc_cutoff)
def _phase_is_drained(outcome: _ChunkOutcome) -> bool:
"""A short chunk means the due rows ran out. A full chunk that advanced
nothing would be re-read unchanged forever, so it ends the phase too and
those rows wait for the next tick."""
return outcome.fetched < RESET_BUDGET_JOB_BATCH_SIZE or outcome.advanced == 0
async def _run_phase_in_chunks(process_chunk: Callable[[], Awaitable[_ChunkOutcome]]) -> None:
"""Drive one reset phase a chunk at a time, capped so a single run cannot
spin unbounded: leftovers are picked up by the next tick."""
for _ in range(RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN):
if _phase_is_drained(await process_chunk()):
return
def _budget_cascade_event_metadata(cascade: _BudgetCascade) -> dict[str, object]:
return {
"num_budgets_found": len(cascade.budgets),
"budgets_found": json.dumps(cascade.budgets, indent=4, default=str),
"num_endusers_found": len(cascade.endusers),
"endusers_found": json.dumps(cascade.endusers, indent=4, default=str),
}
class ResetBudgetJob:
@ -122,21 +226,14 @@ class ResetBudgetJob:
Updates db
"""
if self.prisma_client is not None:
### RESET KEY BUDGET ###
await self.reset_budget_for_litellm_keys()
if self.prisma_client is None:
return
### RESET USER BUDGET ###
await self.reset_budget_for_litellm_users()
## Reset Team Budget
await self.reset_budget_for_litellm_teams()
### RESET ENDUSER (Customer) BUDGET and corresponding Budget duration ###
await self.reset_budget_for_litellm_budget_table()
### RESET MULTI-WINDOW BUDGETS ###
await self.reset_budget_windows()
await self.reset_budget_for_litellm_keys()
await self.reset_budget_for_litellm_users()
await self.reset_budget_for_litellm_teams()
await self.reset_budget_for_litellm_budget_table()
await self.reset_budget_windows()
@staticmethod
async def _invalidate_spend_counter(counter_key: str) -> None:
@ -194,238 +291,195 @@ class ResetBudgetJob:
e,
)
async def _cascade_reset_spend_for_budget_link(
async def _fetch_linked_rows(
self,
budgets_to_reset: list[LiteLLM_BudgetTableFull],
table: SpendLinkedTable[_RowT],
counter_key_fn: Callable[[_RowT], str],
where: Mapping[str, object],
log_subject: str,
extra_where: dict[str, object] | None = None,
cache_key_fn: Callable[[_RowT], str | Sequence[str]] | None = None,
):
"""
Generic cascade: zero spend on rows whose budget_id is in the reset set.
) -> tuple[_RowT, ...]:
"""Read the rows the cascade will zero, so their counters can be
invalidated once the transaction commits."""
try:
return tuple(await table.find_many(where=where))
except Exception as e:
verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e)
return ()
``cache_key_fn`` is optional: when provided, after the DB update each
matching row's entry or entries in ``user_api_key_cache`` are dropped so
cached spend cannot stay pinned above the zeroed DB row after a reset.
async def _collect_endusers_to_reset(self, budget_ids: Sequence[str]) -> tuple[_EndUserRow, ...]:
linked: Final[Sequence[_EndUserRow] | None] = await self.prisma_client.get_data(
table_name="enduser",
query_type="find_all",
budget_id_list=list(budget_ids),
)
if litellm.max_end_user_budget_id is None or litellm.max_end_user_budget_id not in budget_ids:
return tuple(linked or ())
return (*(linked or ()), *await self._get_endusers_with_no_budget_id())
async def _collect_budget_cascade(self, budgets_to_reset: Sequence[LiteLLM_BudgetTableFull]) -> _BudgetCascade:
"""Resolve every row the expiring budget tiers gate, before any write.
Keys carrying their own budget_duration are left out: they run on their
own schedule via reset_budget_for_litellm_keys(), so sweeping them here
would reset them twice.
"""
budget_ids: Final = [b.budget_id for b in budgets_to_reset if b.budget_id is not None]
budget_ids: Final = tuple(b.budget_id for b in budgets_to_reset if b.budget_id is not None)
if not budget_ids:
return _EMPTY_CASCADE
team_memberships: Final[tuple[_TeamMembershipRow, ...]] = await self._fetch_linked_rows(
table=TeamMembershipRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids),
log_subject="team memberships",
)
keys: Final[tuple[_KeyRow, ...]] = await self._fetch_linked_rows(
table=VerificationTokenRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _LINKED_KEYS_WHERE),
log_subject="keys",
)
orgs: Final[tuple[_OrgRow, ...]] = await self._fetch_linked_rows(
table=OrganizationRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE),
log_subject="orgs",
)
tags: Final[tuple[_TagRow, ...]] = await self._fetch_linked_rows(
table=TagRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE),
log_subject="tags",
)
return _BudgetCascade(
budgets=tuple(budgets_to_reset),
budget_ids=budget_ids,
budget_resets=tuple(
(
b.budget_id,
compute_budget_reset_at(budget_duration=b.budget_duration, settings=self.reset_settings),
)
for b in budgets_to_reset
if b.budget_id is not None and b.budget_duration is not None
),
endusers=await self._collect_endusers_to_reset(budget_ids),
counter_keys=(
*(_team_membership_counter_key(row) for row in team_memberships),
*(_key_counter_key(row) for row in keys),
*(_org_counter_key(row) for row in orgs),
*(_tag_counter_key(row) for row in tags),
),
cache_keys=(
*(key for row in team_memberships for key in _team_membership_cache_keys(row)),
*(key for row in keys for key in _key_cache_keys(row)),
*(key for row in orgs for key in _org_cache_keys(row)),
*(key for row in tags for key in _tag_cache_keys(row)),
),
)
async def _commit_budget_cascade(self, cascade: _BudgetCascade) -> None:
"""Zero the gated spend and advance ``budget_reset_at`` in one transaction.
Advancing the window on its own hides the tier from every later tick
while its dependents stay pinned at the cap for the whole window;
batching both means a mid-cascade failure persists nothing and the rows
stay due for the next run.
"""
if not cascade.budget_ids:
return
where: Final[dict[str, object]] = {"budget_id": {"in": budget_ids}}
if extra_where:
where.update(extra_where)
enduser_ids: Final = tuple(row.user_id for row in cascade.endusers)
async with budget_cascade_unit_of_work(self.prisma_client.db.batch_) as uow:
uow.team_memberships.queue_spend_zero(where=_budget_link_where(cascade.budget_ids))
uow.keys.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _LINKED_KEYS_WHERE))
uow.organizations.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
uow.tags.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
if enduser_ids:
uow.endusers.queue_spend_zero(where={"user_id": {"in": list(enduser_ids)}})
for budget_id, budget_reset_at in cascade.budget_resets:
uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at)
try:
rows: Sequence[_RowT] = await table.find_many(where=where)
except Exception as e:
rows = ()
verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e)
update_result: Final = await table.update_many(where=where, data={"spend": 0})
for row in rows:
await self._invalidate_spend_counter(counter_key_fn(row))
if cache_key_fn is not None:
cache_keys = cache_key_fn(row)
if isinstance(cache_keys, str):
cache_keys = [cache_keys]
for cache_key in cache_keys:
await self._invalidate_user_api_key_cache_entry(cache_key)
return update_result
async def reset_budget_for_litellm_team_members(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the budget for all LiteLLM Team Members if their budget has expired
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=TeamMembershipRepository(self.prisma_client).table,
counter_key_fn=_team_membership_counter_key,
log_subject="team memberships",
cache_key_fn=_team_membership_cache_key,
)
async def reset_budget_for_keys_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the spend for keys linked to budget tiers that are being reset.
Excludes keys with their own budget_duration; those are reset by
reset_budget_for_litellm_keys() to avoid double-resetting.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=VerificationTokenRepository(self.prisma_client).table,
counter_key_fn=_key_counter_key,
log_subject="keys",
extra_where={"budget_duration": None, "spend": {"gt": 0}},
cache_key_fn=_key_cache_key,
)
async def reset_budget_for_orgs_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the spend for orgs linked to budget tiers that are being reset.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=OrganizationRepository(self.prisma_client).table,
counter_key_fn=_org_counter_key,
log_subject="orgs",
extra_where={"spend": {"gt": 0}},
cache_key_fn=_org_cache_keys,
)
async def reset_budget_for_tags_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the spend for tags linked to budget tiers that are being reset.
Also drops each tag's ``user_api_key_cache`` entry so the next
``_tag_max_budget_check`` reloads the zeroed row from the DB.
``SpendCounterReseed.from_db`` intentionally returns ``None`` for
tags, so the budget check falls back to the cached
``LiteLLM_TagTable.spend`` once the spend counter expires; without
this invalidation, that stale ``.spend`` keeps the tag over-budget
indefinitely.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=TagRepository(self.prisma_client).table,
counter_key_fn=_tag_counter_key,
log_subject="tags",
extra_where={"spend": {"gt": 0}},
cache_key_fn=_tag_cache_key,
)
async def reset_budget_for_litellm_budget_table(self):
"""
Resets the budget for all LiteLLM End-Users (Customers), and Team Members if their budget has expired
The corresponding Budget duration is also updated.
"""
async def _invalidate_budget_cascade_caches(self, cascade: _BudgetCascade) -> None:
for counter_key in cascade.counter_keys:
await self._invalidate_spend_counter(counter_key)
for cache_key in cascade.cache_keys:
await self._invalidate_user_api_key_cache_entry(cache_key)
async def _reset_expired_budget_cascade(self) -> _BudgetCascadeCommitted | _BudgetCascadeFailed:
now: Final = datetime.now(timezone.utc)
start_time: Final = time.time()
endusers_to_reset: list[LiteLLM_EndUserTable] | None = None
budgets_to_reset: list[LiteLLM_BudgetTableFull] | None = None
updated_endusers: Final[list[LiteLLM_EndUserTable]] = []
failed_endusers: Final = []
try:
budgets_to_reset = await self.prisma_client.get_data(
table_name="budget", query_type="find_all", reset_at=now
)
if budgets_to_reset is not None and len(budgets_to_reset) > 0:
for budget in budgets_to_reset:
budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now, self.reset_settings)
await self.prisma_client.update_data(
query_type="update_many",
data_list=budgets_to_reset,
table_name="budget",
)
budget_ids_to_reset = [budget.budget_id for budget in budgets_to_reset if budget.budget_id is not None]
endusers_to_reset = await self.prisma_client.get_data(
table_name="enduser",
query_type="find_all",
budget_id_list=budget_ids_to_reset,
)
# Also reset end users with no budget_id (NULL) who use the
# default budget via litellm.max_end_user_budget_id. These
# users are enforced in-memory but never had budget_id
# persisted, so the query above misses them.
if litellm.max_end_user_budget_id is not None and litellm.max_end_user_budget_id in budget_ids_to_reset:
default_budget_endusers: Final = await self._get_endusers_with_no_budget_id()
if default_budget_endusers:
if endusers_to_reset is None:
endusers_to_reset = default_budget_endusers
else:
endusers_to_reset.extend(default_budget_endusers)
await self.reset_budget_for_litellm_team_members(budgets_to_reset=budgets_to_reset)
await self.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset)
await self.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=budgets_to_reset)
await self.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=budgets_to_reset)
if endusers_to_reset is not None and len(endusers_to_reset) > 0:
for enduser in endusers_to_reset:
try:
updated_enduser = await ResetBudgetJob._reset_budget_for_enduser(enduser=enduser)
if updated_enduser is not None:
updated_endusers.append(updated_enduser)
else:
failed_endusers.append(
{
"enduser": enduser,
"error": "Returned None without exception",
}
)
except Exception as e:
failed_endusers.append({"enduser": enduser, "error": str(e)})
verbose_proxy_logger.exception("Failed to reset budget for enduser: %s", enduser)
verbose_proxy_logger.debug(
"Updated users %s",
json.dumps(updated_endusers, indent=4, default=str),
)
await self.prisma_client.update_data(
query_type="update_many",
data_list=updated_endusers,
table_name="enduser",
)
end_time = time.time()
if len(failed_endusers) > 0: # If any endusers failed to reset
raise Exception(
f"Failed to reset {len(failed_endusers)} endusers: {json.dumps(failed_endusers, default=str)}"
)
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
call_type="reset_budget_budget_table",
start_time=start_time,
end_time=end_time,
event_metadata={
"num_budgets_found": (len(budgets_to_reset) if budgets_to_reset else 0),
"budgets_found": json.dumps(budgets_to_reset, indent=4, default=str),
"num_endusers_found": (len(endusers_to_reset) if endusers_to_reset else 0),
"endusers_found": json.dumps(endusers_to_reset, indent=4, default=str),
"num_endusers_updated": len(updated_endusers),
"endusers_updated": json.dumps(updated_endusers, indent=4, default=str),
"num_endusers_failed": len(failed_endusers),
"endusers_failed": json.dumps(failed_endusers, indent=4, default=str),
},
)
budgets_to_reset: Final[Sequence[LiteLLM_BudgetTableFull] | None] = await self.prisma_client.get_data(
table_name="budget",
query_type="find_all",
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
cascade: Final = await self._collect_budget_cascade(budgets_to_reset or ())
except Exception as e:
end_time = time.time()
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
error=e,
call_type="reset_budget_endusers",
start_time=start_time,
end_time=end_time,
event_metadata={
"num_budgets_found": (len(budgets_to_reset) if budgets_to_reset else 0),
"budgets_found": json.dumps(budgets_to_reset, indent=4, default=str),
"num_endusers_found": (len(endusers_to_reset) if endusers_to_reset else 0),
"endusers_found": json.dumps(endusers_to_reset, indent=4, default=str),
},
return _BudgetCascadeFailed(cascade=_EMPTY_CASCADE, error=e)
try:
await self._commit_budget_cascade(cascade)
except Exception as e:
return _BudgetCascadeFailed(cascade=cascade, error=e)
await self._invalidate_budget_cascade_caches(cascade)
return _BudgetCascadeCommitted(
cascade=cascade,
advanced=_count_advanced(
(reset_at for _, reset_at in cascade.budget_resets),
cutoff=datetime.now(timezone.utc),
),
)
async def reset_budget_for_litellm_budget_table(self) -> None:
"""
Resets the spend a budget tier gates (end users, team members, keys,
orgs, tags) and advances the tier's budget_reset_at, atomically.
Caches are invalidated only after the transaction commits, so a failed
run cannot leave a zeroed counter in front of an un-reset DB row.
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_budget_table_chunk)
async def _reset_budget_for_litellm_budget_table_chunk(self) -> _ChunkOutcome:
start_time: Final = time.time()
outcome: Final = await self._reset_expired_budget_cascade()
end_time: Final = time.time()
match outcome:
case _BudgetCascadeCommitted(cascade=cascade, advanced=advanced):
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
call_type="reset_budget_budget_table",
start_time=start_time,
end_time=end_time,
event_metadata={
**_budget_cascade_event_metadata(cascade),
"num_endusers_updated": len(cascade.endusers),
"num_endusers_failed": 0,
},
)
)
)
verbose_proxy_logger.exception("Failed to reset budget for endusers: %s", e)
return _ChunkOutcome(fetched=len(cascade.budgets), advanced=advanced)
case _BudgetCascadeFailed(cascade=cascade, error=error):
verbose_proxy_logger.exception(
"Failed to reset the budget table cascade (team member, enduser, org and tag spend, plus "
"budget_reset_at); nothing was committed and the budgets stay due for the next run: %s",
error,
exc_info=error,
)
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
error=error,
call_type="reset_budget_endusers",
start_time=start_time,
end_time=end_time,
event_metadata=_budget_cascade_event_metadata(cascade),
)
)
return _NO_PROGRESS
case _:
assert_never(outcome)
async def _get_endusers_with_no_budget_id(
self,
@ -486,18 +540,50 @@ class ResetBudgetJob:
for t in updated_teams:
uow.teams.queue_spend_reset(team_id=t.team_id, budget_reset_at=t.budget_reset_at)
async def reset_budget_for_litellm_keys(self):
def _emit_phase_failure(
self,
call_type: str,
error: Exception,
start_time: float,
end_time: float,
event_metadata: dict[str, object],
) -> None:
"""Report rows that could not be reset without failing the chunk: the
rows that did reset are already committed, and raising here would cost
the phase every remaining chunk this tick.
"""
verbose_proxy_logger.error("%s: %s", call_type, error)
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
error=error,
call_type=call_type,
start_time=start_time,
end_time=end_time,
event_metadata=event_metadata,
)
)
async def reset_budget_for_litellm_keys(self) -> None:
"""
Resets the budget for all the litellm keys
Catches Exceptions and logs them
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_keys_chunk)
async def _reset_budget_for_litellm_keys_chunk(self) -> _ChunkOutcome:
now: Final = datetime.utcnow()
start_time: Final = time.time()
keys_to_reset: list[LiteLLM_VerificationToken] | None = None
try:
keys_to_reset = await self.prisma_client.get_data(
table_name="key", query_type="find_all", expires=now, reset_at=now
table_name="key",
query_type="find_all",
expires=now,
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
verbose_proxy_logger.debug("Keys to reset %s", json.dumps(keys_to_reset, indent=4, default=str))
updated_keys: Final[list[LiteLLM_VerificationToken]] = []
@ -528,8 +614,25 @@ class ResetBudgetJob:
await self._invalidate_spend_counter(f"spend:key:{token}")
end_time = time.time()
if len(failed_keys) > 0: # If any keys failed to reset
raise Exception(f"Failed to reset {len(failed_keys)} keys: {json.dumps(failed_keys, default=str)}")
outcome: Final = _ChunkOutcome(
fetched=len(keys_to_reset) if keys_to_reset else 0,
advanced=_count_advanced(
(k.budget_reset_at for k in updated_keys),
cutoff=datetime.now(timezone.utc),
),
)
if len(failed_keys) > 0:
self._emit_phase_failure(
call_type="reset_budget_keys",
error=Exception(f"Failed to reset {len(failed_keys)} keys: {json.dumps(failed_keys, default=str)}"),
start_time=start_time,
end_time=end_time,
event_metadata={
"num_keys_found": len(keys_to_reset) if keys_to_reset else 0,
"keys_found": json.dumps(keys_to_reset, indent=4, default=str),
},
)
return outcome
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
@ -565,16 +668,27 @@ class ResetBudgetJob:
)
)
verbose_proxy_logger.exception("Failed to reset budget for keys: %s", e)
return _NO_PROGRESS
else:
return outcome
async def reset_budget_for_litellm_users(self):
async def reset_budget_for_litellm_users(self) -> None:
"""
Resets the budget for all LiteLLM Internal Users if their budget has expired
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_users_chunk)
async def _reset_budget_for_litellm_users_chunk(self) -> _ChunkOutcome:
now: Final = datetime.utcnow()
start_time: Final = time.time()
users_to_reset: list[LiteLLM_UserTable] | None = None
try:
users_to_reset = await self.prisma_client.get_data(table_name="user", query_type="find_all", reset_at=now)
users_to_reset = await self.prisma_client.get_data(
table_name="user",
query_type="find_all",
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
updated_users: Final[list[LiteLLM_UserTable]] = []
failed_users: Final = []
if users_to_reset is not None and len(users_to_reset) > 0:
@ -609,8 +723,27 @@ class ResetBudgetJob:
await self._invalidate_global_proxy_spend_cache()
end_time = time.time()
if len(failed_users) > 0: # If any users failed to reset
raise Exception(f"Failed to reset {len(failed_users)} users: {json.dumps(failed_users, default=str)}")
outcome: Final = _ChunkOutcome(
fetched=len(users_to_reset) if users_to_reset else 0,
advanced=_count_advanced(
(u.budget_reset_at for u in updated_users),
cutoff=datetime.now(timezone.utc),
),
)
if len(failed_users) > 0:
self._emit_phase_failure(
call_type="reset_budget_users",
error=Exception(
f"Failed to reset {len(failed_users)} users: {json.dumps(failed_users, default=str)}"
),
start_time=start_time,
end_time=end_time,
event_metadata={
"num_users_found": len(users_to_reset) if users_to_reset else 0,
"users_found": json.dumps(users_to_reset, indent=4, default=str),
},
)
return outcome
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
@ -646,16 +779,27 @@ class ResetBudgetJob:
)
)
verbose_proxy_logger.exception("Failed to reset budget for users: %s", e)
return _NO_PROGRESS
else:
return outcome
async def reset_budget_for_litellm_teams(self):
async def reset_budget_for_litellm_teams(self) -> None:
"""
Resets the budget for all LiteLLM Internal Teams if their budget has expired
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_teams_chunk)
async def _reset_budget_for_litellm_teams_chunk(self) -> _ChunkOutcome:
now: Final = datetime.utcnow()
start_time: Final = time.time()
teams_to_reset: list[LiteLLM_TeamTable] | None = None
try:
teams_to_reset = await self.prisma_client.get_data(table_name="team", query_type="find_all", reset_at=now)
teams_to_reset = await self.prisma_client.get_data(
table_name="team",
query_type="find_all",
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
updated_teams: Final[list[LiteLLM_TeamTable]] = []
failed_teams: Final = []
if teams_to_reset is not None and len(teams_to_reset) > 0:
@ -688,8 +832,27 @@ class ResetBudgetJob:
await self._invalidate_spend_counter(f"spend:team:{team_id}")
end_time = time.time()
if len(failed_teams) > 0: # If any teams failed to reset
raise Exception(f"Failed to reset {len(failed_teams)} teams: {json.dumps(failed_teams, default=str)}")
outcome: Final = _ChunkOutcome(
fetched=len(teams_to_reset) if teams_to_reset else 0,
advanced=_count_advanced(
(t.budget_reset_at for t in updated_teams),
cutoff=datetime.now(timezone.utc),
),
)
if len(failed_teams) > 0:
self._emit_phase_failure(
call_type="reset_budget_teams",
error=Exception(
f"Failed to reset {len(failed_teams)} teams: {json.dumps(failed_teams, default=str)}"
),
start_time=start_time,
end_time=end_time,
event_metadata={
"num_teams_found": len(teams_to_reset) if teams_to_reset else 0,
"teams_found": json.dumps(teams_to_reset, indent=4, default=str),
},
)
return outcome
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
@ -725,6 +888,9 @@ class ResetBudgetJob:
)
)
verbose_proxy_logger.exception("Failed to reset budget for teams: %s", e)
return _NO_PROGRESS
else:
return outcome
@staticmethod
async def _reset_expired_window(
@ -882,33 +1048,6 @@ class ResetBudgetJob:
)
return user
@staticmethod
async def _reset_budget_for_enduser(
enduser: LiteLLM_EndUserTable,
) -> LiteLLM_EndUserTable | None:
try:
enduser.spend = 0.0
except Exception as e:
verbose_proxy_logger.exception("Error resetting budget for enduser: %s. Item: %s", e, enduser)
raise e
return enduser
@staticmethod
async def _reset_budget_reset_at_date(
budget: LiteLLM_BudgetTableFull,
current_time: datetime,
reset_settings: BudgetResetSettings,
) -> LiteLLM_BudgetTableFull:
try:
if budget.budget_duration is not None:
budget.budget_reset_at = compute_budget_reset_at(
budget_duration=budget.budget_duration, settings=reset_settings
)
except Exception as e:
verbose_proxy_logger.exception("Error resetting budget_reset_at for budget: %s. Item: %s", e, budget)
raise e
return budget
@staticmethod
async def _reset_budget_for_key(
key: LiteLLM_VerificationToken,

View file

@ -31,6 +31,8 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
ESTIMATED_OUTPUT_TOKENS_FIELD,
get_estimated_output_tokens,
get_key_tag_rpm_limit,
get_model_rate_limit_from_metadata,
)
@ -562,6 +564,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
data: dict,
model: str | None = None,
min_configured_tpm_limit: int | None = None,
configured_output_tokens: int | None = None,
) -> int:
"""
Estimate total tokens this request will consume so we can reserve them
@ -575,6 +578,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
provided, the no-``max_tokens`` output-budget floor is capped at a
fraction of that limit so small TPM caps remain usable. Omit to
preserve the unconstrained floor.
``configured_output_tokens`` is the operator-declared estimate resolved
from key or team metadata. When provided it replaces the heuristic
floor entirely, so the reservation reflects what this tenant's model
actually emits rather than one constant shared by every tenant.
"""
messages = data.get("messages")
prompt: Final = data.get("prompt")
@ -604,7 +612,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
case (_, embeddings_input) if embeddings_input:
# Embeddings have no output tokens
max_tokens_estimate = 0
case _ if total_chars == 0:
case _ if total_chars == 0 and configured_output_tokens is None:
# Fully contentless request (no messages, prompt, or input).
# Don't apply the conservative output-budget floor here — it
# would over-reserve and could push small TPM limits into a
@ -619,7 +627,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# so a small per-tenant TPM cap can't be tripped by the floor
# alone.
output_floor: Final = self._no_max_tokens_output_floor(min_configured_tpm_limit)
max_tokens_estimate = max(estimated_input_tokens, output_floor)
max_tokens_estimate = (
configured_output_tokens
if configured_output_tokens is not None
else max(estimated_input_tokens, output_floor)
)
total_estimated: Final = estimated_input_tokens + max_tokens_estimate
@ -2586,8 +2598,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
data.get("max_tokens") is not None or data.get("max_completion_tokens") is not None
)
is_embedding: Final = data.get("input") is not None
configured_output_tokens: Final = get_estimated_output_tokens(
user_api_key_dict=user_api_key_dict,
model_name=requested_model,
)
if capped_floor < baseline_floor and not has_explicit_max_tokens and not is_embedding:
data["max_tokens"] = capped_floor
data["max_tokens"] = max(capped_floor, configured_output_tokens or 0)
# Floor at 1 token so contentless requests (/responses,
# tool-call continuations, empty messages) still flow
@ -2601,10 +2617,23 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
data=data,
model=requested_model,
min_configured_tpm_limit=min_configured_tpm_limit,
configured_output_tokens=configured_output_tokens,
),
1,
)
if configured_output_tokens is not None and estimated_tokens > min_configured_tpm_limit:
verbose_proxy_logger.debug(
"Reserving %s tokens for model %s (declared %s=%s plus the input estimate) exceeds the "
"smallest TPM limit this request is charged against (%s), so it cannot be admitted even "
"against an empty window. Lower the declared estimate or raise the TPM limit.",
estimated_tokens,
requested_model,
ESTIMATED_OUTPUT_TOKENS_FIELD,
configured_output_tokens,
min_configured_tpm_limit,
)
tpm_response: Final = await self.reserve_tpm_tokens(
descriptors=descriptors,
estimated_tokens=estimated_tokens,

View file

@ -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,
)

View file

@ -20,7 +20,10 @@ from fastapi import APIRouter, Depends, HTTPException
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_view,
validate_budget_duration,
)
from litellm.proxy.utils import jsonify_object
from litellm.repositories.budget_repository import BudgetRepository
@ -72,6 +75,8 @@ async def new_budget(
detail={"error": f"soft_budget must be a non-negative finite number. Received: {budget_obj.soft_budget}"},
)
validate_budget_duration(budget_obj.budget_duration)
# Validate model_max_budget if present
if budget_obj.model_max_budget is not None and len(budget_obj.model_max_budget) > 0:
from litellm.proxy.management_endpoints.key_management_endpoints import (
@ -153,6 +158,8 @@ async def update_budget(
detail={"error": f"soft_budget must be a non-negative finite number. Received: {budget_obj.soft_budget}"},
)
validate_budget_duration(budget_obj.budget_duration)
# Validate model_max_budget if present in update
if budget_obj.model_max_budget is not None and len(budget_obj.model_max_budget) > 0:
from litellm.proxy.management_endpoints.key_management_endpoints import (

View file

@ -8,7 +8,9 @@ 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.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
from litellm.proxy.utils import PrismaClient
from litellm.repositories.table_repositories import DeletedVerificationTokenRepository
from litellm.repositories.verification_token_repository import (
@ -140,6 +142,28 @@ class _GroupingSetsRow(SimpleNamespace):
failed_requests: int | None
def _reported_flat_cost(record: DailySpendRecord | _GroupingSetsRow) -> float:
"""Flat cost a daily row reports, which is zero unless PTU cost attribution is enabled.
Both read paths funnel through here: the paginated path reads the ``ptu_flat_cost``
column straight off the row, and the aggregated path reads the SUM() alias. Rows an
operator accrued during an earlier opt-in stay in the table, so the gate lives on the
read rather than on the query that produced the rows.
The row is checked before the flag because this runs once per metric accumulation, and
a record fans out across roughly a dozen breakdowns. The flag reads through the secret
manager, uncached, so consulting it for every accumulation put thousands of lookups on
a shared endpoint that made none before. Only a row actually carrying flat cost, which
is a sentinel row, reaches it now.
"""
raw: Final = getattr(record, "ptu_flat_cost", None) or 0.0
if not raw:
return 0.0
if not is_ptu_cost_attribution_enabled():
return 0.0
return raw
def update_metrics(existing_metrics: SpendMetrics, record: DailySpendRecord) -> SpendMetrics:
"""Update metrics with new record data.
@ -150,6 +174,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 += _reported_flat_cost(record)
existing_metrics.prompt_tokens += prompt_tokens
existing_metrics.completion_tokens += completion_tokens
existing_metrics.total_tokens += prompt_tokens + completion_tokens
@ -208,30 +233,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 +283,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 +327,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 +376,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 +399,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 +641,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 +662,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 +758,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 +807,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=_reported_flat_cost(record),
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
@ -820,6 +874,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 +887,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 +896,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 +910,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 +960,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 +1070,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 +1161,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,

View file

@ -22,6 +22,35 @@ def validate_finite_spend(spend: float | None) -> None:
)
def validate_budget_duration(budget_duration: str | None) -> None:
"""Reject budget durations that can't be parsed, are non-positive, or
overflow date math, so a bad value can't be persisted and later crash the
budget reset job.
A non-positive duration also resolves to a reset time of "now", which leaves
the row permanently due: the reset job re-reads it every tick and, once
enough of them exist, they fill each batch and starve every other tenant's
reset.
"""
if budget_duration is None:
return
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
try:
if duration_in_seconds(budget_duration) <= 0:
raise ValueError("budget_duration must be positive")
get_budget_reset_time(budget_duration=budget_duration)
except (ValueError, OverflowError):
raise HTTPException(
status_code=400,
detail={
"error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'."
},
)
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
from litellm.proxy._types import (

View file

@ -23,6 +23,7 @@ from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
from litellm.proxy.management_endpoints.common_utils import validate_budget_duration
from litellm.proxy.management_helpers.object_permission_utils import (
_set_object_permission,
handle_update_object_permission_common,
@ -184,6 +185,7 @@ def new_budget_request(data: NewCustomerRequest) -> BudgetNewRequest | None:
if budget_kv_pairs:
budget_request: Final = BudgetNewRequest(**budget_kv_pairs)
validate_budget_duration(budget_request.budget_duration)
if budget_request.budget_reset_at is None and budget_request.budget_duration is not None:
budget_request.budget_reset_at = datetime.utcnow() + timedelta(
seconds=duration_in_seconds(duration=budget_request.budget_duration)

View file

@ -42,6 +42,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_is_user_team_admin,
_user_has_admin_view,
require_caller_user_id_for_non_admin,
validate_budget_duration,
validate_finite_spend,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
@ -508,6 +509,8 @@ async def new_user(
status_code=500,
detail=CommonProxyErrors.db_not_connected_error.value,
)
validate_budget_duration(data.budget_duration)
# Check for duplicate user_id or email
await _check_duplicate_user_id(data.user_id, prisma_client)
await _check_duplicate_user_email(data.user_email, prisma_client)
@ -1187,6 +1190,7 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest | Upda
if "budget_duration" in non_default_values:
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
validate_budget_duration(non_default_values["budget_duration"])
non_default_values["budget_reset_at"] = get_budget_reset_time(
budget_duration=non_default_values["budget_duration"]
)

View file

@ -55,7 +55,10 @@ from litellm.proxy.auth.auth_checks import (
get_project_object,
get_team_object,
)
from litellm.proxy.auth.auth_utils import abbreviate_api_key
from litellm.proxy.auth.auth_utils import (
abbreviate_api_key,
enforce_output_token_estimates_are_admin_only,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.callback_utils import (
decrypt_callback_vars,
@ -79,6 +82,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_set_object_metadata_field,
_team_member_has_permission,
_user_has_admin_view,
validate_budget_duration,
validate_finite_spend,
)
from litellm.proxy.management_endpoints.model_management_endpoints import (
@ -847,12 +851,21 @@ async def _common_key_generation_helper(
premium_user=premium_user,
)
validate_budget_duration(data.budget_duration)
if data.throttle_on_budget_exceeded is True and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
raise HTTPException(
status_code=403,
detail={"error": "Only proxy admins can enable throttle_on_budget_exceeded on a key."},
)
enforce_output_token_estimates_are_admin_only(
data=data,
existing_metadata=None,
user_api_key_dict=user_api_key_dict,
entity="key",
)
if data.metadata is not None and data.metadata.get("service_account_id") is not None and data.team_id is None:
await validate_team_id_used_in_service_account_request(
team_id=data.team_id,
@ -1020,7 +1033,7 @@ async def _common_key_generation_helper(
# Only set budget_duration on key when explicitly provided. Keys with budget_id
# but no explicit budget_duration follow their linked budget tier's schedule;
# reset_budget_for_keys_linked_to_budgets() resets them when the tier resets.
# reset_budget_for_litellm_budget_table() resets them when the tier resets.
# This avoids duplicating budget_duration on keys so tier updates apply automatically.
if "budget_duration" in data_json:
data_json["key_budget_duration"] = data_json.pop("budget_duration", None)
@ -1590,6 +1603,8 @@ async def generate_key_fn(
- budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
- default_estimated_output_tokens: Optional[int] - Proxy admin only. Expected output tokens reserved for TPM limiting when a request omits max_tokens. Positive integer. Falls back to the team setting, then to the built-in estimate.
- default_estimated_output_tokens_per_model: Optional[dict] - Proxy admin only. Per-model override of the above. Example - {"gpt-4": 4096, "gpt-3.5-turbo": 1024}. Takes precedence over the key-wide value.
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
- tag_rpm_limit: Optional[dict] - key-specific per-request-tag rpm limit, keyed by request tag. Example - {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; requests whose tag is absent fall back to the key-level rpm limit.
- tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
@ -1799,6 +1814,8 @@ async def generate_service_account_key_fn(
- budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
- default_estimated_output_tokens: Optional[int] - Proxy admin only. Expected output tokens reserved for TPM limiting when a request omits max_tokens. Positive integer. Falls back to the team setting, then to the built-in estimate.
- default_estimated_output_tokens_per_model: Optional[dict] - Proxy admin only. Per-model override of the above. Example - {"gpt-4": 4096, "gpt-3.5-turbo": 1024}. Takes precedence over the key-wide value.
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
@ -2393,6 +2410,7 @@ async def _validate_update_key_data(
"""Validate permissions and constraints for key update."""
# Reject NaN/±inf spend before it can reach the DB / spend counter.
validate_finite_spend(data.spend)
validate_budget_duration(data.budget_duration)
_is_proxy_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
@ -2479,6 +2497,13 @@ async def _validate_update_key_data(
detail={"error": "Only proxy admins can enable throttle_on_budget_exceeded on a key."},
)
enforce_output_token_estimates_are_admin_only(
data=data,
existing_metadata=_existing_metadata if isinstance(_existing_metadata, dict) else None,
user_api_key_dict=user_api_key_dict,
entity="key",
)
# Personal-key bypass: the caller both created the key AND still owns it
# (user_id == caller). Checking only created_by would let a demoted admin
# who originally created a key for another user continue editing it without
@ -2661,6 +2686,8 @@ async def update_key_fn(
- mcp_rpm_limit: Optional[dict] - Per-MCP-server RPM limits, keyed by MCP server name {"github": 100, "slack": 200}
- tag_rpm_limit: Optional[dict] - Per-request-tag RPM limits, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; absent tags fall back to the key-level rpm limit.
- model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000}
- default_estimated_output_tokens: Optional[int] - Proxy admin only. Expected output tokens reserved for TPM limiting when a request omits max_tokens. Positive integer.
- default_estimated_output_tokens_per_model: Optional[dict] - Proxy admin only. Per-model override of the above {"gpt-4": 4096, "gpt-3.5-turbo": 1024}
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- allowed_cache_controls: Optional[list] - List of allowed cache control values
@ -4635,6 +4662,15 @@ async def _execute_virtual_key_regeneration(
prisma_client=prisma_client,
)
if data is not None:
_existing_key_metadata: Final = getattr(key_in_db, "metadata", None)
enforce_output_token_estimates_are_admin_only(
data=data,
existing_metadata=_existing_key_metadata if isinstance(_existing_key_metadata, dict) else None,
user_api_key_dict=user_api_key_dict,
entity="key",
)
new_token: Final = await get_new_token(data=data)
new_token_hash: Final = hash_token(new_token)
new_token_key_name: Final = abbreviate_api_key(api_key=new_token)

View file

@ -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
@ -55,6 +56,10 @@ from litellm.proxy.management_endpoints.team_endpoints import (
update_team as _legacy_update_team,
)
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
from litellm.proxy.spend_tracking.ptu_feature_flag import (
PTU_COST_ATTRIBUTION_ENV_VAR,
is_ptu_cost_attribution_enabled,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.model_repository import ModelRepository
from litellm.repositories.table_repositories import ModelTableRepository
@ -79,6 +84,7 @@ from litellm.types.router import (
SPECIAL_MODEL_INFO_PARAMS,
Deployment,
GenericLiteLLMParams,
ModelInfo,
updateDeployment,
)
from litellm.utils import get_utc_datetime
@ -233,7 +239,129 @@ 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.
Empty while the feature is off, so disabling pauses PTU rather than letting a client
that round-trips a model_info blob erase a configuration set up during an earlier opt-in.
"""
if model_info is None or not is_ptu_cost_attribution_enabled():
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 _raise_if_ptu_cost_attribution_disabled(incoming_model_info: Mapping[str, object]) -> None:
"""Reject PTU model_info fields unless the operator opted into PTU cost attribution.
Takes the incoming request's model_info rather than the merged deployment, so an
unrelated patch of a model that still stores PTU config from an earlier opt-in is
left alone. The fields are rejected rather than dropped so a caller never believes
a flat cost was configured while the rollup that would price it is not running.
Only a value is rejected. An explicit null reaches the clear loop, which is gated on
the same flag, so a disabled proxy neither writes PTU config nor erases what an
earlier opt-in stored. Disabling pauses the feature rather than discarding its setup.
"""
if is_ptu_cost_attribution_enabled():
return
supplied: Final = tuple(field for field in _PTU_MODEL_INFO_FIELDS if incoming_model_info.get(field) is not None)
if not supplied:
return
raise HTTPException(
status_code=400,
detail=(
f"PTU cost attribution is disabled, so {', '.join(supplied)} cannot be set. "
f"Set {PTU_COST_ATTRIBUTION_ENV_VAR}=true to enable it."
),
)
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:
if updated_patch.model_info is not None:
_raise_if_ptu_cost_attribution_disabled(updated_patch.model_info.model_dump(exclude_none=True))
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)
merged_model_info: Final = db_model.model_info.model_dump(exclude_none=True)
@ -270,6 +398,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 +848,18 @@ 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:
_raise_if_ptu_cost_attribution_disabled(patch_data.model_info.model_dump(exclude_none=True))
_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 +1568,10 @@ async def add_new_model(
model_response: LiteLLM_ProxyModelTable | None = None
# update DB
incoming_model_info: Final = model_params.model_info.model_dump(exclude_none=True)
_raise_if_ptu_cost_attribution_disabled(incoming_model_info)
_validate_ptu_model_info(incoming_model_info)
if store_model_in_db is True:
"""
- store model_list in db

View file

@ -82,6 +82,7 @@ from litellm.proxy.auth.auth_checks import (
get_team_object,
get_user_object,
)
from litellm.proxy.auth.auth_utils import enforce_output_token_estimates_are_admin_only
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch
@ -95,6 +96,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_update_metadata_fields,
_upsert_budget_and_membership,
_user_has_admin_view,
validate_budget_duration,
)
from litellm.proxy.management_endpoints.organization_endpoints import (
add_member_to_organization,
@ -1153,6 +1155,8 @@ async def new_team(
- metadata: Optional[dict] - Metadata for team, store information for team. Example metadata = {"extra_info": "some info"}
- model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit for this team - applied across all keys for this team.
- model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit for this team - applied across all keys for this team.
- default_estimated_output_tokens: Optional[int] - Expected output tokens reserved for TPM limiting when a request omits max_tokens, for keys on this team that do not set their own. Positive integer.
- default_estimated_output_tokens_per_model: Optional[Dict[str, int]] - Per-model override of the above. Example: {"gpt-4": 4096, "gpt-3.5-turbo": 1024}
- mcp_rpm_limit: Optional[Dict[str, int]] - Per-MCP-server RPM limit for this team, keyed by MCP server name (alias if set, else the configured name). Example: {"github": 100, "slack": 200}. Applied across all keys for this team.
- tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for this team - all keys with this team_id will have at max this TPM limit
- rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for this team - all keys associated with this team_id will have at max this RPM limit
@ -1255,6 +1259,9 @@ async def new_team(
detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"},
)
validate_budget_duration(data.budget_duration)
validate_budget_duration(data.team_member_budget_duration)
if data.soft_budget is not None:
if data.max_budget is not None:
# If max_budget is set, soft_budget must be strictly lower than max_budget
@ -1266,6 +1273,13 @@ async def new_team(
},
)
enforce_output_token_estimates_are_admin_only(
data=data,
existing_metadata=None,
user_api_key_dict=user_api_key_dict,
entity="team",
)
# Check if license is over limit
total_teams: Final = await _team_db(prisma_client).count()
if total_teams and _license_check.is_team_count_over_limit(team_count=total_teams):
@ -1863,6 +1877,8 @@ async def update_team(
- allowed_passthrough_routes: Optional[List[str]] - List of allowed pass through routes for the team.
- model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit per model for this team. Example: {"gpt-4": 100, "gpt-3.5-turbo": 200}
- model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit per model for this team. Example: {"gpt-4": 10000, "gpt-3.5-turbo": 20000}
- default_estimated_output_tokens: Optional[int] - Expected output tokens reserved for TPM limiting when a request omits max_tokens, for keys on this team that do not set their own. Positive integer.
- default_estimated_output_tokens_per_model: Optional[Dict[str, int]] - Per-model override of the above. Example: {"gpt-4": 4096, "gpt-3.5-turbo": 1024}
- mcp_rpm_limit: Optional[Dict[str, int]] - Per-MCP-server RPM limit for this team, keyed by MCP server name (alias if set, else the configured name). Example: {"github": 100, "slack": 200}. Applied across all keys for this team.
Example - update team TPM Limit
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
@ -1935,6 +1951,9 @@ async def update_team(
detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"},
)
validate_budget_duration(data.budget_duration)
validate_budget_duration(data.team_member_budget_duration)
existing_team_row = await _team_db(prisma_client).find_unique(where={"team_id": data.team_id})
if existing_team_row is None:
@ -1949,6 +1968,14 @@ async def update_team(
user_api_key_dict=user_api_key_dict,
)
_existing_team_metadata: Final[object] = getattr(existing_team_row, "metadata", None)
enforce_output_token_estimates_are_admin_only(
data=data,
existing_metadata=_existing_team_metadata if isinstance(_existing_team_metadata, dict) else None,
user_api_key_dict=user_api_key_dict,
entity="team",
)
_check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team")
if data.soft_budget is not None:
@ -2959,7 +2986,7 @@ async def team_member_add(
except HTTPException as e:
raise e
_validate_budget_duration(data.budget_duration)
validate_budget_duration(data.budget_duration)
prisma_client = cast(PrismaClient, prisma_client)
@ -3262,29 +3289,6 @@ def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> dict[str, objec
}
def _validate_budget_duration(budget_duration: str | None) -> None:
"""Reject budget durations that can't be parsed, are non-positive, or
overflow date math, so a bad value can't be persisted and later crash the
budget reset job."""
if budget_duration is None:
return
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
try:
if duration_in_seconds(budget_duration) <= 0:
raise ValueError("budget_duration must be positive")
get_budget_reset_time(budget_duration=budget_duration)
except (ValueError, OverflowError):
raise HTTPException(
status_code=400,
detail={
"error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'."
},
)
@router.post(
"/team/member_update",
tags=["team management"],
@ -3322,7 +3326,7 @@ async def team_member_update(
detail={"error": "Either user_id or user_email needs to be passed in"},
)
_validate_budget_duration(data.budget_duration)
validate_budget_duration(data.budget_duration)
_existing_team_row: Final = await _team_db(prisma_client).find_unique(where={"team_id": data.team_id})

View file

@ -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,
@ -6860,9 +6860,19 @@ class ProxyConfig:
guardrail_id = guardrail.get("guardrail_id")
if guardrail_id:
db_guardrail_ids.add(guardrail_id)
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
guardrail=cast(Guardrail, guardrail),
)
try:
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
guardrail=cast(Guardrail, guardrail),
)
except Exception as e: # noqa: BLE001 # one unloadable row must not stop the remaining guardrails
verbose_proxy_logger.error(
"litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - "
"skipping guardrail '%s' (ID: %s): %s: %s",
guardrail.get("guardrail_name"),
guardrail_id,
type(e).__name__,
e,
)
# Drop in-memory DB-backed entries whose row was deleted on another
# pod. Config-loaded entries are never touched.
@ -8470,6 +8480,47 @@ class ProxyStartupEvent:
await cls._initialize_spend_tracking_background_jobs(scheduler=scheduler)
### PTU DAILY ROLLUP ###
from litellm.proxy.spend_tracking.ptu_feature_flag import (
is_ptu_cost_attribution_enabled,
)
if is_ptu_cost_attribution_enabled():
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

View file

@ -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

View file

@ -0,0 +1,18 @@
"""Opt-in flag for PTU (provisioned throughput unit) flat-cost attribution.
The whole feature is inert unless an operator sets
``LITELLM_ENABLE_PTU_COST_ATTRIBUTION``: the daily rollup is not scheduled, the
model endpoints reject PTU config, the daily activity read path reports zero flat
cost, and the model form hides the PTU inputs.
"""
from typing import Final
from litellm.secret_managers.main import get_secret_bool
PTU_COST_ATTRIBUTION_ENV_VAR: Final = "LITELLM_ENABLE_PTU_COST_ATTRIBUTION"
def is_ptu_cost_attribution_enabled() -> bool:
"""Report whether this deployment opted into PTU flat-cost attribution."""
return get_secret_bool(PTU_COST_ATTRIBUTION_ENV_VAR, False) is True

View file

@ -0,0 +1,663 @@
"""
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.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
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.
Returns None without touching the database when PTU cost attribution is off. Proxy
startup already skips scheduling the cron, so this guards the function itself rather
than its one caller, and a deployment that never opted in accrues nothing whatever
reaches it.
"""
if not is_ptu_cost_attribution_enabled():
return None
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",
)

View file

@ -21,6 +21,7 @@ from litellm.proxy.config_resolvers.sso import (
SSO_SECRET_FIELDS,
resolve_sso_config,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
from litellm.proxy.utils import invalidate_config_param
from litellm.repositories.config_repository import ConfigRepository
from litellm.repositories.organization_repository import OrganizationRepository
@ -307,6 +308,27 @@ ALLOWED_UI_SETTINGS_FIELDS: Final = {
"enable_chat_ui",
}
ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: Final = "enable_ptu_cost_attribution"
# UI settings derived from the deployment environment. Deliberately kept out of
# ALLOWED_UI_SETTINGS_FIELDS: they are read-only, never persisted, and PATCH
# rejects them so an admin cannot flip an env-gated feature at runtime.
_DERIVED_UI_SETTINGS_FIELDS: Final[frozenset[str]] = frozenset({ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING})
def _derived_ui_setting_value(key: str) -> object:
"""The environment-derived value GET reports for ``key``.
PATCH compares against this rather than rejecting the key outright, so the body GET
hands back is still a valid PATCH body. Rejecting on presence broke read-modify-write:
a client that edited one setting and sent the rest back unchanged got a 400 and lost
the edit it actually wanted.
"""
if key == ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING:
return is_ptu_cost_attribution_enabled()
return None
# Flags that must be synced from the persisted UISettings into
# general_settings at runtime (on both read and write).
_RUNTIME_GENERAL_SETTINGS_FLAGS: Final = [
@ -1345,21 +1367,15 @@ async def get_ui_settings():
detail={"error": "Database not connected. Please connect a database."},
)
ui_settings: Mapping[str, JsonValue] = {}
db_record: Final = await _ui_settings_db(UISettingsRepository(prisma_client)).find_unique(
where={"id": "ui_settings"}
)
if db_record and db_record.ui_settings:
ui_settings_json: Final = db_record.ui_settings
if isinstance(ui_settings_json, str):
ui_settings = json.loads(ui_settings_json)
else:
ui_settings = dict(ui_settings_json)
stored: Final = (db_record.ui_settings if db_record else None) or "{}"
parsed: Final = json.loads(stored) if isinstance(stored, str) else stored
# Sanitize any unexpected keys from persisted config before returning
ui_settings = {k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS}
ui_settings: Final = {k: v for k, v in parsed.items() if k in ALLOWED_UI_SETTINGS_FIELDS}
# Sync runtime flags into general_settings so the proxy picks them up
# at runtime (covers server restart scenarios).
@ -1377,11 +1393,18 @@ async def get_ui_settings():
# Build config-like object for schema helper
config: Final[dict[str, object]] = {"litellm_settings": {"ui_settings": ui_settings}}
return await _get_settings_with_schema(
settings: Final = await _get_settings_with_schema(
settings_key="ui_settings",
settings_class=_get_effective_ui_settings_class(),
config=config,
)
return UISettingsResponse(
values={
**settings["values"],
ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: is_ptu_cost_attribution_enabled(),
},
field_schema=settings["field_schema"],
)
@router.patch(
@ -1418,6 +1441,20 @@ async def update_ui_settings(
detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."},
)
conflicting_keys: Final = sorted(
key
for key, value in settings_body.items()
if key in _DERIVED_UI_SETTINGS_FIELDS and value != _derived_ui_setting_value(key)
)
if conflicting_keys:
raise HTTPException(
status_code=400,
detail=(
f"Setting(s) {conflicting_keys} are derived from the deployment environment "
"and cannot be changed from the UI."
),
)
# Validate against the same effective class GET advertises, so
# enterprise-registered fields are typed consistently on both sides.
effective_cls: Final = _get_effective_ui_settings_class()

View file

@ -3486,13 +3486,15 @@ class PrismaClient:
r.expires = r.expires.isoformat()
elif query_type == "find_all" and expires is not None and reset_at is not None:
response = await VerificationTokenRepository(self).table.find_many(
take=limit,
where={
"OR": [
{"expires": None},
{"expires": {"gt": expires}},
],
"budget_reset_at": {"lt": reset_at},
}
"NOT": {"budget_duration": None},
},
)
if response is not None and len(response) > 0:
for r in response:
@ -3542,6 +3544,7 @@ class PrismaClient:
response = await UserRepository(self).table.find_many(where=key_val)
elif query_type == "find_all" and reset_at is not None:
response = await UserRepository(self).table.find_many(
take=limit,
where={
# A user seeded from default_internal_user_params
# (or created via /user/new without an explicit
@ -3552,16 +3555,12 @@ class PrismaClient:
# of the row, silently exceeding max_budget. Treat a
# NULL budget_reset_at with a non-NULL budget_duration
# as due, matching the budget-table query below.
"NOT": {"budget_duration": None},
"OR": [
{
"AND": [
{"budget_reset_at": None},
{"NOT": {"budget_duration": None}},
]
},
{"budget_reset_at": None},
{"budget_reset_at": {"lt": reset_at}},
],
}
},
)
elif query_type == "find_all" and user_id_list is not None:
response = await UserRepository(self).table.find_many(where={"user_id": {"in": user_id_list}})
@ -3617,17 +3616,14 @@ class PrismaClient:
elif table_name == "budget" and reset_at is not None:
if query_type == "find_all":
response = await BudgetRepository(self).table.find_many(
take=limit,
where={
"NOT": {"budget_duration": None},
"OR": [
{
"AND": [
{"budget_reset_at": None},
{"NOT": {"budget_duration": None}},
]
},
{"budget_reset_at": None},
{"budget_reset_at": {"lt": reset_at}},
]
}
],
},
)
return response
@ -3645,20 +3641,17 @@ class PrismaClient:
)
elif query_type == "find_all" and reset_at is not None:
response = await TeamRepository(self).table.find_many(
take=limit,
where={
# Same NULL budget_reset_at gap as the user query
# above: a team with a budget_duration but no
# initialized budget_reset_at would never be reset.
"NOT": {"budget_duration": None},
"OR": [
{
"AND": [
{"budget_reset_at": None},
{"NOT": {"budget_duration": None}},
]
},
{"budget_reset_at": None},
{"budget_reset_at": {"lt": reset_at}},
],
}
},
)
elif query_type == "find_all" and user_id is not None:
response = await TeamRepository(self).table.find_many(

View file

@ -70,10 +70,14 @@ from litellm.repositories.table_repositories import (
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.unit_of_work import (
BudgetCascadeUnitOfWork,
BudgetWindowWrites,
KeySpendResetWrites,
LinkedSpendResetWrites,
SpendResetUnitOfWork,
TeamSpendResetWrites,
UserSpendResetWrites,
budget_cascade_unit_of_work,
spend_reset_unit_of_work,
)
from litellm.repositories.user_repository import UserRepository
@ -88,7 +92,9 @@ __all__ = [
"AgentsRepository",
"AuditLogRepository",
"BatchTable",
"BudgetCascadeUnitOfWork",
"BudgetRepository",
"BudgetWindowWrites",
"CacheConfigRepository",
"ClaudeCodePluginRepository",
"ConfigOverridesRepository",
@ -107,6 +113,7 @@ __all__ = [
"InvitationLinkRepository",
"JWTKeyMappingRepository",
"KeySpendResetWrites",
"LinkedSpendResetWrites",
"MCPServerRepository",
"MCPToolsetRepository",
"MCPUserCredentialsRepository",
@ -149,5 +156,6 @@ __all__ = [
"WorkflowEventRepository",
"WorkflowMessageRepository",
"WorkflowRunRepository",
"budget_cascade_unit_of_work",
"spend_reset_unit_of_work",
]

View file

@ -232,6 +232,8 @@ class ObjectPermissionWriteTable(Protocol):
class BatchTable(Protocol):
def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
class PrismaBatch(Protocol):
@property
@ -243,4 +245,19 @@ class PrismaBatch(Protocol):
@property
def litellm_teamtable(self) -> BatchTable: ...
@property
def litellm_budgettable(self) -> BatchTable: ...
@property
def litellm_teammembership(self) -> BatchTable: ...
@property
def litellm_organizationtable(self) -> BatchTable: ...
@property
def litellm_tagtable(self) -> BatchTable: ...
@property
def litellm_endusertable(self) -> BatchTable: ...
async def commit(self) -> None: ...

View file

@ -1,17 +1,21 @@
"""
Unit of work over a single Prisma batch.
Units of work over a single Prisma batch.
``spend_reset_unit_of_work`` opens one ``db.batch_()`` and binds a typed write
Each context manager here opens one ``db.batch_()`` and binds a typed write
repository per table to it, so every update queued through the yielded object
lands in the same transaction. The batch commits when the block exits cleanly
and is abandoned, writing nothing, when the block raises.
Each write repository queues narrow ``{spend, budget_reset_at}`` updates
``spend_reset_unit_of_work`` covers the per-row key/user/team resets;
``budget_cascade_unit_of_work`` covers a budget tier's reset, where the
dependent spend and the tier's next window have to move together.
Each write repository queues narrow ``{spend}`` / ``{budget_reset_at}`` updates
instead of full-model writes, which trip ``prisma.errors.DataError`` on rows
carrying fields the update input type rejects (see #27730).
"""
from collections.abc import AsyncGenerator, Callable
from collections.abc import AsyncGenerator, Callable, Mapping
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime
@ -43,6 +47,24 @@ class TeamSpendResetWrites:
self.table.update(where={"team_id": team_id}, data={"spend": 0, "budget_reset_at": budget_reset_at})
@dataclass(frozen=True, slots=True)
class LinkedSpendResetWrites:
table: BatchTable
def queue_spend_zero(self, where: Mapping[str, object]) -> None:
self.table.update_many(where=where, data={"spend": 0})
@dataclass(frozen=True, slots=True)
class BudgetWindowWrites:
table: BatchTable
def queue_window_advance(self, budget_id: str, budget_reset_at: datetime) -> None:
"""``update_many`` so a tier deleted between the read and the commit is a
no-op row count instead of a P2025 that aborts the whole chunk."""
self.table.update_many(where={"budget_id": budget_id}, data={"budget_reset_at": budget_reset_at})
@dataclass(frozen=True, slots=True)
class SpendResetUnitOfWork:
keys: KeySpendResetWrites
@ -50,6 +72,23 @@ class SpendResetUnitOfWork:
teams: TeamSpendResetWrites
@dataclass(frozen=True, slots=True)
class BudgetCascadeUnitOfWork:
"""Every write a budget-tier reset performs, bound to one batch.
The dependent spend rows and the budget rows' ``budget_reset_at`` advance
must land together: advancing the window without zeroing the spend it
gates leaves the dependents pinned at their cap until the next window.
"""
team_memberships: LinkedSpendResetWrites
keys: LinkedSpendResetWrites
organizations: LinkedSpendResetWrites
tags: LinkedSpendResetWrites
endusers: LinkedSpendResetWrites
budgets: BudgetWindowWrites
@asynccontextmanager
async def spend_reset_unit_of_work(new_batch: Callable[[], PrismaBatch]) -> AsyncGenerator[SpendResetUnitOfWork, None]:
batch = new_batch()
@ -59,3 +98,19 @@ async def spend_reset_unit_of_work(new_batch: Callable[[], PrismaBatch]) -> Asyn
teams=TeamSpendResetWrites(table=batch.litellm_teamtable),
)
await batch.commit()
@asynccontextmanager
async def budget_cascade_unit_of_work(
new_batch: Callable[[], PrismaBatch],
) -> AsyncGenerator[BudgetCascadeUnitOfWork, None]:
batch = new_batch()
yield BudgetCascadeUnitOfWork(
team_memberships=LinkedSpendResetWrites(table=batch.litellm_teammembership),
keys=LinkedSpendResetWrites(table=batch.litellm_verificationtoken),
organizations=LinkedSpendResetWrites(table=batch.litellm_organizationtable),
tags=LinkedSpendResetWrites(table=batch.litellm_tagtable),
endusers=LinkedSpendResetWrites(table=batch.litellm_endusertable),
budgets=BudgetWindowWrites(table=batch.litellm_budgettable),
)
await batch.commit()

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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):

View file

@ -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

View file

@ -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

View file

@ -1,13 +1,3 @@
[[IgnoredVulns]]
id = "GHSA-fwg2-594c-jp42"
ignoreUntil = 2026-08-12
reason = "pypdf 6.15.0 (the fix) published 2026-08-06 and is still inside the P3D exclude-newer window, so uv cannot lock it yet; bump and drop this entry from 2026-08-09"
[[IgnoredVulns]]
id = "GHSA-fp3f-mc75-235c"
ignoreUntil = 2026-08-12
reason = "second pypdf advisory with the same 6.15.0 fix, published 2026-08-07 after the first; drop alongside GHSA-fwg2-594c-jp42 in the same bump"
[[IgnoredVulns]]
id = "GHSA-w8v5-vhqr-4h9v"
ignoreUntil = 2026-09-09

View file

@ -9,10 +9,10 @@
"limit": 832
},
"ANN201": {
"limit": 2031
"limit": 2023
},
"ANN202": {
"limit": 861
"limit": 860
},
"ANN204": {
"limit": 713
@ -237,13 +237,13 @@
"limit": 1225
},
"TRY002": {
"limit": 528
"limit": 524
},
"TRY004": {
"limit": 96
},
"TRY201": {
"limit": 407
"limit": 405
},
"TRY203": {
"limit": 113

View file

@ -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

View file

@ -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

View file

@ -0,0 +1,143 @@
"""Guard the CI cache for Prisma's CLI and engine binaries.
``prisma generate`` shells out to ``npm install prisma@<version>`` whenever the
prisma-client-py binary cache directory has no CLI entrypoint, pulling ~85 MB of
engines over the network. The download is normally seconds and occasionally
minutes, and a job timeout cannot tell the difference from a hung test, so an
uncached job is one slow npm response away from cancelling a passing test run.
Three invariants keep that download off the critical path:
1. No workflow sets ``PRISMA_BINARY_CACHE_DIR``. The prisma-client-py default is
``~/.cache/prisma-python/binaries/<prisma-version>/<engine-version>``, already
keyed by both versions and the only path the cache action restores. Pointing
it elsewhere (``runner.temp`` especially, which is wiped every job) silently
guarantees a cold download.
2. Every job that generates the client also restores the cache.
3. The cache key resolves to a real version from ``uv.lock``. The action fails
the job when it cannot, so a lock format change must break here instead.
"""
import re
import sys
from collections.abc import Iterator, Mapping
from pathlib import Path
from typing import Final
import yaml
from pydantic import BaseModel, Field, ValidationError
REPO_ROOT: Final = Path(__file__).resolve().parent.parent.parent
WORKFLOWS_DIR: Final = REPO_ROOT / ".github" / "workflows"
UV_LOCK: Final = REPO_ROOT / "uv.lock"
CACHE_ACTION: Final = "./.github/actions/cache-prisma-binaries"
# Commands that reach the prisma binary cache: a direct generate, or a script
# that runs one on the caller's behalf.
PRISMA_GENERATE_MARKERS: Final = ("prisma generate", "type_check_gate.py")
class PrismaBinaryCacheError(Exception):
pass
def resolve_prisma_version(lock_text: str) -> str | None:
"""Mirror of the shell lookup in the cache action's version step."""
match: Final = re.search(
r'^name = "prisma"\n^version = "(?P<version>[^"]+)"$',
lock_text,
re.MULTILINE,
)
return match.group("version") if match else None
class WorkflowStep(BaseModel):
"""The two step fields this guard reads; every other key is ignored."""
run: str | None = None
uses: str | None = None
def generates_prisma_client(self) -> bool:
return self.run is not None and any(m in self.run for m in PRISMA_GENERATE_MARKERS)
def restores_cache(self) -> bool:
return self.uses == CACHE_ACTION
class WorkflowJob(BaseModel):
# Absent for jobs that delegate to a reusable workflow via a job-level `uses`.
steps: tuple[WorkflowStep, ...] = ()
class Workflow(BaseModel):
jobs: Mapping[str, WorkflowJob] = Field(default_factory=dict)
def parse_workflow(text: str) -> Workflow | str:
"""Validate untyped YAML at the boundary so the checks below stay typed.
Returns the parsed workflow, or a description of why it could not be read.
"""
parsed: Final = yaml.safe_load(text)
try:
return Workflow.model_validate(parsed if isinstance(parsed, dict) else {})
except ValidationError as exc:
return f"does not parse as a workflow: {exc.error_count()} schema error(s)"
def lock_errors(lock_text: str) -> Iterator[str]:
if not resolve_prisma_version(lock_text):
yield (
"uv.lock has no resolvable `prisma` package version. The version step "
f"in {CACHE_ACTION} greps the same shape and will fail every job that "
"generates the Prisma client."
)
def workflow_errors(rel: Path, text: str) -> Iterator[str]:
if "PRISMA_BINARY_CACHE_DIR" in text:
yield (
f"{rel}: sets PRISMA_BINARY_CACHE_DIR. Leave it unset so the binaries "
f"land in the version-keyed default path the {CACHE_ACTION} action restores."
)
workflow: Final = parse_workflow(text)
if isinstance(workflow, str):
yield f"{rel}: {workflow}"
return
for job_name, job in workflow.jobs.items():
if any(s.generates_prisma_client() for s in job.steps) and not any(
s.restores_cache() for s in job.steps
):
yield (
f"{rel}: job `{job_name}` generates the Prisma client without a "
f"`uses: {CACHE_ACTION}` step, so it downloads ~85 MB of engines "
"on every run."
)
def main() -> None:
errors: Final = (
*lock_errors(UV_LOCK.read_text()),
*(
error
for path in sorted(WORKFLOWS_DIR.glob("*.y*ml"))
for error in workflow_errors(path.relative_to(REPO_ROOT), path.read_text())
),
)
if errors:
raise PrismaBinaryCacheError(
"Prisma binary cache invariants violated:\n - " + "\n - ".join(errors)
)
print("Prisma binary cache invariants hold across .github/workflows/")
if __name__ == "__main__":
try:
main()
except PrismaBinaryCacheError as exc:
print(f"ERROR: {exc}", file=sys.stderr)
sys.exit(1)

View file

@ -0,0 +1,239 @@
"""Catch workflow mistakes that GitHub reports as nothing at all.
A workflow whose YAML is valid but whose expressions are not fails at *startup*:
the run is marked failed, no jobs are created, and no check run is ever posted.
Nothing turns red on the PR, so an entire test suite can silently stop running
while the checks list stays green. These invariants have to be enforced here
because CI cannot enforce them on itself.
1. No arithmetic inside ``${{ }}``. GitHub expressions support grouping, index,
dereference, ``!``, the comparisons, ``&&`` and ``||``, and nothing else. A
``${{ a + b }}`` is a startup failure, not a value. Only ``+`` and ``*`` are
flagged: ``-`` appears in hyphenated input names like ``inputs.timeout-minutes``
and ``/`` inside ref strings, so neither can be told apart from arithmetic by
inspection alone.
2. Callers of the reusable unit-test workflow keep the job timeout at or above
the test budget plus the setup ceilings plus the runner overhead below.
Otherwise the job deadline preempts pytest inside its own advertised budget,
which is the failure the split timeouts exist to prevent, and it shows up as
a cancelled shard whose tests were passing. A budget this check cannot resolve
is reported rather than skipped, so a mistyped input or matrix column surfaces
here instead of leaving the pair silently unchecked.
"""
import re
import sys
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import yaml
from pydantic import BaseModel, Field, ValidationError
REPO_ROOT: Final = Path(__file__).resolve().parent.parent.parent
WORKFLOWS_DIR: Final = REPO_ROOT / ".github" / "workflows"
BASE_WORKFLOW: Final = "./.github/workflows/_test-unit-base.yml"
BASE_WORKFLOW_PATH: Final = WORKFLOWS_DIR / "_test-unit-base.yml"
# Runner time the job clock charges but no step owns: job init, the gaps between
# steps, and post-job cleanup. Without it a job capped at exactly test + setup
# would still preempt pytest inside its own budget.
JOB_OVERHEAD_MINUTES: Final = 5
EXPRESSION: Final = re.compile(r"\$\{\{(?P<body>.*?)\}\}", re.DOTALL)
QUOTED: Final = re.compile(r"'[^']*'")
ARITHMETIC: Final = re.compile(r"[+*]")
MATRIX_REF: Final = re.compile(r"^\$\{\{\s*matrix\.(?P<key>[\w-]+)\s*\}\}$")
class WorkflowStartupError(Exception):
pass
class ReusableCall(BaseModel):
uses: str | None = None
with_: Mapping[str, object] = Field(default_factory=dict, alias="with")
strategy: Mapping[str, object] = Field(default_factory=dict)
steps: tuple[Mapping[str, object], ...] = ()
model_config = {"populate_by_name": True}
class WorkflowFile(BaseModel):
jobs: Mapping[str, ReusableCall] = Field(default_factory=dict)
def parse_workflow(text: str) -> WorkflowFile | str:
parsed: Final = yaml.safe_load(text)
try:
return WorkflowFile.model_validate(parsed if isinstance(parsed, dict) else {})
except ValidationError as exc:
return f"does not parse as a workflow: {exc.error_count()} schema error(s)"
def arithmetic_expressions(text: str) -> Iterator[str]:
for match in EXPRESSION.finditer(text):
body: Final = match.group("body")
if ARITHMETIC.search(QUOTED.sub("", body)):
yield body.strip()
def setup_ceiling_minutes(base_text: str) -> int:
"""Sum the per-step timeouts on everything the base workflow runs before pytest."""
base: Final = yaml.safe_load(base_text)
steps: Final = base["jobs"]["run"]["steps"]
return sum(
s["timeout-minutes"]
for s in steps
if s.get("name") != "Run tests" and isinstance(s.get("timeout-minutes"), int)
)
def base_default(base_text: str, name: str) -> int:
base: Final = yaml.safe_load(base_text)
return base[True]["workflow_call"]["inputs"][name]["default"]
@dataclass(frozen=True, slots=True)
class Column:
"""A budget the caller reads from one column of its own matrix."""
name: str
def budget_source(job: ReusableCall, key: str, fallback: int) -> int | Column | str:
"""A caller passes a literal, or `${{ matrix.x }}` naming a column of its matrix.
Anything else comes back as the reason it could not be read, since a budget
nothing can resolve has to be reported rather than passed over.
"""
value: Final = job.with_.get(key)
if value is None:
return fallback
if isinstance(value, int):
return value
matrix_ref: Final = MATRIX_REF.match(str(value))
if not matrix_ref:
return f"passes `{key}: {value}`, which is neither a number nor a `matrix` reference."
return Column(matrix_ref.group("key"))
def matrix_rows(job: ReusableCall) -> Sequence[Mapping[str, object]]:
matrix: Final = job.strategy.get("matrix", {})
entries: Final = matrix.get("include", ()) if isinstance(matrix, dict) else ()
return tuple(e for e in entries if isinstance(e, dict))
def budget_pairs(job: ReusableCall, test_source: int | Column, job_source: int | Column) -> Iterator[tuple[int, int]]:
"""Pair each shard's test budget with the job budget of that same shard.
Matrix-sourced budgets resolve per `include` row, so two matrix columns are
read off the same row rather than cross-producted across rows.
"""
if isinstance(test_source, int) and isinstance(job_source, int):
yield test_source, job_source
return
for row in matrix_rows(job):
test_budget = row.get(test_source.name) if isinstance(test_source, Column) else test_source
job_budget = row.get(job_source.name) if isinstance(job_source, Column) else job_source
if isinstance(test_budget, int) and isinstance(job_budget, int):
yield test_budget, job_budget
def unresolved_message(where: str, job: ReusableCall, sources: Sequence[int | Column]) -> str:
"""Why no shard yielded a pair of budgets to compare.
Naming only the columns that resolve nowhere keeps the message honest: a
column every row supplies is not what left the pair unchecked.
"""
rows: Final = matrix_rows(job)
missing: Final = tuple(
f"`matrix.{s.name}`"
for s in sources
if isinstance(s, Column) and not any(isinstance(row.get(s.name), int) for row in rows)
)
if missing:
return (
f"{where} reads a budget from {', '.join(missing)}, which no `include` row supplies "
"as a number, so the pair would go unchecked."
)
return (
f"{where} reads both budgets from its matrix, but no single `include` row supplies both "
"as numbers, so the pair would go unchecked."
)
def job_errors(rel: Path, job_name: str, job: ReusableCall, ceiling: int, base_text: str) -> Iterator[str]:
where: Final = f"{rel}: job `{job_name}`"
test_source: Final = budget_source(job, "timeout-minutes", base_default(base_text, "timeout-minutes"))
job_source: Final = budget_source(job, "job-timeout-minutes", base_default(base_text, "job-timeout-minutes"))
sources: Final = (test_source, job_source)
unreadable: Final = tuple(f"{where} {reason}" for reason in sources if isinstance(reason, str))
if unreadable:
yield from unreadable
return
pairs: Final = tuple(budget_pairs(job, test_source, job_source))
if not pairs:
yield unresolved_message(where, job, sources)
return
for test_budget, job_budget in pairs:
required = test_budget + ceiling + JOB_OVERHEAD_MINUTES
if job_budget < required:
yield (
f"{where} gives pytest {test_budget}m but caps the job at "
f"{job_budget}m. Setup can use up to {ceiling}m plus {JOB_OVERHEAD_MINUTES}m of "
f"runner overhead, so the job deadline would preempt pytest; raise "
f"job-timeout-minutes to at least {required}."
)
def timeout_contract_errors(rel: Path, workflow: WorkflowFile, ceiling: int, base_text: str) -> Iterator[str]:
for job_name, job in workflow.jobs.items():
if job.uses == BASE_WORKFLOW:
yield from job_errors(rel, job_name, job, ceiling, base_text)
def workflow_errors(rel: Path, text: str, ceiling: int, base_text: str) -> Iterator[str]:
for expression in arithmetic_expressions(text):
yield (
f"{rel}: `${{{{ {expression} }}}}` uses arithmetic, which GitHub expressions do not "
"support. The workflow will fail at startup with no jobs and no check run."
)
workflow: Final = parse_workflow(text)
if isinstance(workflow, str):
yield f"{rel}: {workflow}"
return
yield from timeout_contract_errors(rel, workflow, ceiling, base_text)
def main() -> None:
base_text: Final = BASE_WORKFLOW_PATH.read_text()
ceiling: Final = setup_ceiling_minutes(base_text)
errors: Final = tuple(
error
for path in sorted(WORKFLOWS_DIR.glob("*.y*ml"))
for error in workflow_errors(path.relative_to(REPO_ROOT), path.read_text(), ceiling, base_text)
)
if errors:
raise WorkflowStartupError(
"Workflow startup invariants violated:\n - " + "\n - ".join(errors)
)
print(f"Workflow startup invariants hold (setup ceiling {ceiling}m)")
if __name__ == "__main__":
try:
main()
except WorkflowStartupError as exc:
print(f"ERROR: {exc}", file=sys.stderr)
sys.exit(1)

View file

@ -44,39 +44,87 @@ def _attrify(d: dict):
return _AttrDict(d)
def _wire_batcher_for_test(prisma_client):
def _wire_batcher_for_test(prisma_client, fail_commit=False):
"""
Wire prisma_client.db.batch_() to return a mock batcher whose .commit() is
awaitable and whose per-table .update() calls get captured. The reset job
writes key/user/team resets via prisma.db.batch_().<table>.update not via
prisma_client.update_data so tests must let that batch path complete.
awaitable and whose per-table .update()/.update_many() calls get captured.
The reset job writes every reset through prisma.db.batch_() key/user/team
rows one by one, and the budget tier's cascade as a single transaction — so
tests must let that batch path complete.
Returns the list that will accumulate {table, where, data} dicts from
each captured update call.
Only committed batches contribute to the returned list, mirroring prisma:
with fail_commit=True the transaction blows up and must persist nothing.
Returns the list that will accumulate {table, op, where, data} dicts from
each captured write.
"""
batch_calls = []
def make_batcher():
queued = []
class _Table:
def __init__(self, table_name):
self._table_name = table_name
def update(self, where=None, data=None):
batch_calls.append(
{"table": self._table_name, "where": where, "data": data}
queued.append(
{
"table": self._table_name,
"op": "update",
"where": where,
"data": data,
}
)
def update_many(self, where=None, data=None):
queued.append(
{
"table": self._table_name,
"op": "update_many",
"where": where,
"data": data,
}
)
async def commit():
if fail_commit:
raise RuntimeError("simulated Postgres failure committing the batch")
batch_calls.extend(queued)
batcher = MagicMock()
batcher.litellm_verificationtoken = _Table("key")
batcher.litellm_usertable = _Table("user")
batcher.litellm_teamtable = _Table("team")
batcher.commit = AsyncMock(return_value=None)
batcher.litellm_budgettable = _Table("budget")
batcher.litellm_teammembership = _Table("team_membership")
batcher.litellm_organizationtable = _Table("org")
batcher.litellm_tagtable = _Table("tag")
batcher.litellm_endusertable = _Table("enduser")
batcher.commit = commit
return batcher
prisma_client.db.batch_ = MagicMock(side_effect=make_batcher)
return batch_calls
def _wire_cascade_reads_for_test(prisma_client):
"""
The budget tier's cascade reads the rows it is about to zero, so their
spend counters can be invalidated after the commit. Give each of those
tables an awaitable find_many so the reads resolve instead of falling into
the job's warn-and-continue path.
"""
for table in (
"litellm_teammembership",
"litellm_verificationtoken",
"litellm_organizationtable",
"litellm_tagtable",
"litellm_endusertable",
):
getattr(prisma_client.db, table).find_many = AsyncMock(return_value=[])
@pytest.mark.asyncio
async def test_reset_budget_keys_partial_failure():
"""
@ -250,41 +298,18 @@ async def test_reset_budget_users_partial_failure():
@pytest.mark.asyncio
async def test_reset_budget_endusers_partial_failure():
async def test_reset_budget_endusers_cascade_failure_is_all_or_nothing():
"""
Test that if one enduser fails to reset, the reset loop still processes the other endusers.
We simulate six endsers where the first fails and the others are updated.
A failure anywhere in the budget-tier cascade must persist nothing, so the
tier stays due and the next scheduler tick retries it. Before the fix the
job committed the new budget_reset_at first and zeroed the dependent spend
afterwards, so a failure here left the tier stamped for the next window
while every end user stayed at the cap.
"""
user1 = {
"user_id": "user1",
"spend": 20.0,
"budget_id": "budget1",
} # Will trigger simulated failure
user2 = {
"user_id": "user2",
"spend": 25.0,
"budget_id": "budget1",
} # Should be updated
user3 = {
"user_id": "user3",
"spend": 30.0,
"budget_id": "budget1",
} # Should be updated
user4 = {
"user_id": "user4",
"spend": 35.0,
"budget_id": "budget1",
} # Should be updated
user5 = {
"user_id": "user5",
"spend": 40.0,
"budget_id": "budget1",
} # Should be updated
user6 = {
"user_id": "user6",
"spend": 45.0,
"budget_id": "budget1",
} # Should be updated
endusers = [
_attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"})
for i in range(1, 7)
]
budget1 = LiteLLM_BudgetTableFull(
**{
@ -301,23 +326,13 @@ async def test_reset_budget_endusers_partial_failure():
if table_name == "budget":
return [budget1]
elif table_name == "enduser":
return [user1, user2, user3, user4, user5, user6]
return endusers
return []
prisma_client.get_data = AsyncMock()
prisma_client.get_data.side_effect = get_data_mock
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
batch_calls = _wire_batcher_for_test(prisma_client, fail_commit=True)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -326,41 +341,13 @@ async def test_reset_budget_endusers_partial_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_enduser(enduser):
if enduser["user_id"] == "user1":
raise Exception("Simulated failure for user1")
enduser["spend"] = 0.0
return enduser
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
async def fake_reset_team_members(budgets_to_reset):
return 1
with (
patch.object(
ResetBudgetJob,
"_reset_budget_for_enduser",
side_effect=fake_reset_enduser,
) as mock_reset_enduser,
patch.object(
ResetBudgetJob,
"reset_budget_for_litellm_team_members",
side_effect=fake_reset_team_members,
) as mock_reset_team_members,
):
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
assert mock_reset_enduser.call_count == 6
assert prisma_client.update_data.await_count == 2
update_call = prisma_client.update_data.call_args
assert update_call.kwargs.get("table_name") == "enduser"
updated_users = update_call.kwargs.get("data_list", [])
assert len(updated_users) == 5
assert updated_users[0]["user_id"] == "user2"
assert updated_users[1]["user_id"] == "user3"
assert updated_users[2]["user_id"] == "user4"
assert updated_users[3]["user_id"] == "user5"
assert updated_users[4]["user_id"] == "user6"
assert batch_calls == [], "a failed cascade must not persist any write"
assert (
prisma_client.update_data.await_count == 0
), "budget_reset_at must not be advanced outside the cascade transaction"
failure_hook_calls = (
proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args_list
@ -369,6 +356,66 @@ async def test_reset_budget_endusers_partial_failure():
call.kwargs.get("call_type") == "reset_budget_endusers"
for call in failure_hook_calls
)
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called()
@pytest.mark.asyncio
async def test_reset_budget_endusers_are_zeroed_with_the_budget_window_advance():
"""
The happy path: every end user the tier gates is zeroed and the tier's
budget_reset_at advances, all inside one transaction.
"""
endusers = [
_attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"})
for i in range(1, 7)
]
budget1 = LiteLLM_BudgetTableFull(
**{
"budget_id": "budget1",
"max_budget": 65.0,
"budget_duration": "2d",
"created_at": datetime.now(timezone.utc) - timedelta(days=3),
}
)
prisma_client = MagicMock()
async def get_data_mock(table_name, *args, **kwargs):
if table_name == "budget":
return [budget1]
elif table_name == "enduser":
return endusers
return []
prisma_client.get_data = AsyncMock()
prisma_client.get_data.side_effect = get_data_mock
prisma_client.update_data = AsyncMock()
batch_calls = _wire_batcher_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
assert prisma_client.db.batch_.call_count == 1, "the cascade must be one transaction"
enduser_writes = [c for c in batch_calls if c["table"] == "enduser"]
assert len(enduser_writes) == 1
assert enduser_writes[0]["where"]["user_id"]["in"] == [f"user{i}" for i in range(1, 7)]
assert enduser_writes[0]["data"] == {"spend": 0}
budget_writes = [c for c in batch_calls if c["table"] == "budget"]
assert len(budget_writes) == 1
assert budget_writes[0]["where"] == {"budget_id": "budget1"}
assert budget_writes[0]["data"]["budget_reset_at"] > datetime.now(timezone.utc)
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called()
@pytest.mark.asyncio
@ -500,16 +547,8 @@ async def test_reset_budget_continues_other_categories_on_failure():
key1, key2 = _attrify(key1), _attrify(key2)
user1, user2 = _attrify(user1), _attrify(user2)
team1, team2 = _attrify(team1), _attrify(team2)
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
enduser1 = _attrify(enduser1)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -541,13 +580,6 @@ async def test_reset_budget_continues_other_categories_on_failure():
).isoformat()
return team
async def fake_reset_enduser(enduser):
enduser["spend"] = 0.0
return enduser
async def fake_reset_team_members(budgets_to_reset):
return 1
with (
patch.object(
ResetBudgetJob, "_reset_budget_for_key", side_effect=fake_reset_key
@ -558,14 +590,6 @@ async def test_reset_budget_continues_other_categories_on_failure():
patch.object(
ResetBudgetJob, "_reset_budget_for_team", side_effect=fake_reset_team
) as mock_reset_team,
patch.object(
ResetBudgetJob, "_reset_budget_for_enduser", side_effect=fake_reset_enduser
) as mock_reset_enduser,
patch.object(
ResetBudgetJob,
"reset_budget_for_litellm_team_members",
side_effect=fake_reset_team_members,
) as mock_reset_team_members,
):
# Call the overall reset_budget method.
await job.reset_budget()
@ -575,29 +599,22 @@ async def test_reset_budget_continues_other_categories_on_failure():
called_tables = {
call.kwargs.get("table_name") for call in prisma_client.get_data.await_args_list
}
if mock_reset_team_members.call_count > 0:
called_tables.add("team_membership")
assert called_tables == {
"key",
"user",
"team",
"budget",
"enduser",
"team_membership",
}
assert called_tables == {"key", "user", "team", "budget", "enduser"}
# After the fix, keys/users/teams write via prisma.db.batch_().<table>.update,
# so only budget + enduser still go through update_data.
calls = prisma_client.update_data.await_args_list
update_data_tables = [c.kwargs.get("table_name") for c in calls]
assert sorted(update_data_tables) == ["budget", "enduser"]
# Every category writes through the batch path now, so update_data is unused.
prisma_client.update_data.assert_not_awaited()
# Check enduser update: enduser succeed.
enduser_call = next(c for c in calls if c.kwargs.get("table_name") == "enduser")
assert len(enduser_call.kwargs.get("data_list", [])) == 1
# The budget tier's cascade still ran despite the failing user category.
assert len([c for c in batch_calls if c["table"] == "team_membership"]) == 1
enduser_writes = [c for c in batch_calls if c["table"] == "enduser"]
assert len(enduser_writes) == 1
assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1"]}}
assert enduser_writes[0]["data"] == {"spend": 0}
# Check the new batch write path: 2 keys + 1 user (user1 failed) + 2 teams.
key_writes = [c for c in batch_calls if c["table"] == "key"]
# `op` separates the per-row resets from the cascade sweep, which also
# targets the key table.
key_writes = [c for c in batch_calls if c["table"] == "key" and c["op"] == "update"]
user_writes = [c for c in batch_calls if c["table"] == "user"]
team_writes = [c for c in batch_calls if c["table"] == "team"]
assert len(key_writes) == 2
@ -974,12 +991,12 @@ async def test_service_logger_teams_failure():
@pytest.mark.asyncio
async def test_service_logger_endusers_success():
"""
Test that when resetting endusers succeeds the service logger success hook is called with
the correct metadata and no exception is logged.
Test that when the budget-tier cascade commits, the service logger success
hook is called with the correct metadata and no exception is logged.
"""
endusers = [
{"user_id": "user1", "spend": 25.0, "budget_id": "budget1"},
{"user_id": "user2", "spend": 25.0, "budget_id": "budget1"},
_attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}),
_attrify({"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}),
]
budgets = [
LiteLLM_BudgetTableFull(
@ -1002,16 +1019,8 @@ async def test_service_logger_endusers_success():
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
batch_calls = _wire_batcher_for_test(prisma_client)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -1020,31 +1029,16 @@ async def test_service_logger_endusers_success():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_enduser(enduser):
enduser["spend"] = 0.0
return enduser
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
mock_verbose_exc.assert_not_called()
async def fake_reset_team_members(budgets_to_reset):
return 1
with (
patch.object(
ResetBudgetJob,
"_reset_budget_for_enduser",
side_effect=fake_reset_enduser,
) as mock_reset_enduser,
patch.object(
ResetBudgetJob,
"reset_budget_for_litellm_team_members",
side_effect=fake_reset_team_members,
) as mock_reset_team_members,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
mock_verbose_exc.assert_not_called()
enduser_writes = [c for c in batch_calls if c["table"] == "enduser"]
assert len(enduser_writes) == 1
assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1", "user2"]}}
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_called_once()
(
@ -1062,12 +1056,12 @@ async def test_service_logger_endusers_success():
@pytest.mark.asyncio
async def test_service_logger_endusers_failure():
"""
Test that a failure during enduser reset calls the failure hook with appropriate metadata,
logs the exception, and does not call the success hook.
Test that a failed cascade calls the failure hook with the rows it had
found, logs the exception, and does not call the success hook.
"""
endusers = [
{"user_id": "user1", "spend": 25.0, "budget_id": "budget1"},
{"user_id": "user2", "spend": 25.0, "budget_id": "budget1"},
_attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}),
_attrify({"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}),
]
budgets = [
LiteLLM_BudgetTableFull(
@ -1090,16 +1084,8 @@ async def test_service_logger_endusers_failure():
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
_wire_batcher_for_test(prisma_client, fail_commit=True)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -1108,39 +1094,16 @@ async def test_service_logger_endusers_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_enduser(enduser):
if enduser["user_id"] == "user1":
raise Exception("Simulated failure for user1")
enduser["spend"] = 0.0
return enduser
async def fake_reset_team_members(budgets_to_reset):
return 1
with (
patch.object(
ResetBudgetJob,
"_reset_budget_for_enduser",
side_effect=fake_reset_enduser,
) as mock_reset_enduser,
patch.object(
ResetBudgetJob,
"reset_budget_for_litellm_team_members",
side_effect=fake_reset_team_members,
) as mock_reset_team_members,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
# Verify exception logging
assert mock_verbose_exc.call_count >= 1
# Verify exception was logged with correct message
assert any(
"Failed to reset budget for enduser" in str(call.args)
for call in mock_verbose_exc.call_args_list
)
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
# The log must name the whole cascade, not just end users: the write
# that failed could have been any of team member / enduser / org / tag
# spend or the budget_reset_at advance.
assert mock_verbose_exc.call_count == 1
assert "budget table cascade" in str(mock_verbose_exc.call_args.args[0])
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_called_once()
(
@ -1158,8 +1121,8 @@ async def test_service_logger_endusers_failure():
@pytest.mark.asyncio
async def test_reset_budget_for_litellm_team_members_called():
"""
Test that when reset_budget_for_litellm_budget_table is called,
team members' budgets are also reset via reset_budget_for_litellm_team_members
Test that when reset_budget_for_litellm_budget_table is called, team
members' spend is zeroed as part of the cascade transaction.
"""
# Arrange
budget1 = LiteLLM_BudgetTableFull(
@ -1171,7 +1134,7 @@ async def test_reset_budget_for_litellm_team_members_called():
}
)
enduser1 = {"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}
enduser1 = _attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"})
prisma_client = MagicMock()
@ -1184,20 +1147,9 @@ async def test_reset_budget_for_litellm_team_members_called():
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
# Mock the db.litellm_teammembership.update_many call
prisma_client.db = MagicMock()
prisma_client.db.litellm_teammembership = MagicMock()
prisma_client.db.litellm_teammembership.update_many = AsyncMock(
return_value={"count": 2}
)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
batch_calls = _wire_batcher_for_test(prisma_client)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -1206,23 +1158,11 @@ async def test_reset_budget_for_litellm_team_members_called():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_enduser(enduser):
enduser["spend"] = 0.0
return enduser
with patch.object(
ResetBudgetJob,
"_reset_budget_for_enduser",
side_effect=fake_reset_enduser,
):
# Act
await job.reset_budget_for_litellm_budget_table()
# Act
await job.reset_budget_for_litellm_budget_table()
# Assert
# Verify that the team membership update was called
prisma_client.db.litellm_teammembership.update_many.assert_called_once()
# Verify the call was made with correct parameters
call_args = prisma_client.db.litellm_teammembership.update_many.call_args
assert call_args.kwargs["where"]["budget_id"]["in"] == ["budget1"]
assert call_args.kwargs["data"]["spend"] == 0
team_member_writes = [c for c in batch_calls if c["table"] == "team_membership"]
assert len(team_member_writes) == 1
assert team_member_writes[0]["where"]["budget_id"]["in"] == ["budget1"]
assert team_member_writes[0]["data"] == {"spend": 0}

View file

@ -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

View file

@ -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",
[

View 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

View file

@ -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

View file

@ -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),

View file

@ -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("")

View file

@ -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):

View file

@ -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("")

View file

@ -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(

View file

@ -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)

File diff suppressed because it is too large Load diff

View file

@ -3,6 +3,7 @@ Unit Tests for the max parallel request limiter v3 for the proxy
"""
import asyncio
import logging
import os
import sys
import time
@ -5100,3 +5101,456 @@ async def test_reserve_tpm_tokens_never_evaluates_the_requests_dimension():
f"reservation pass, got: {response}"
)
assert [s["rate_limit_type"] for s in response["statuses"]] == ["tokens"]
STATIC_OUTPUT_FLOOR = 1024
ONE_TOKEN_PROMPT = [{"role": "user", "content": "hello"}]
ONE_TOKEN_PROMPT_INPUT_ESTIMATE = 1
async def _reserved_tokens_for(
handler,
local_cache,
user_api_key_dict,
data,
call_type="completion",
):
"""Drive the pre-call hook and read back what landed on the :tokens counter."""
tokens_key = handler.create_rate_limit_keys(
key="api_key", value=user_api_key_dict.api_key, rate_limit_type="tokens"
)
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data=data,
call_type=call_type,
)
return int(await local_cache.async_get_cache(key=tokens_key) or 0)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"key_metadata, team_metadata, expected_output_estimate, tier",
[
(
{
"default_estimated_output_tokens_per_model": {"gpt-4o-mini": 3001},
"default_estimated_output_tokens": 2002,
},
{
"default_estimated_output_tokens_per_model": {"gpt-4o-mini": 1503},
"default_estimated_output_tokens": 777,
},
3001,
"key per-model wins over every other tier",
),
(
{"default_estimated_output_tokens": 2002},
{
"default_estimated_output_tokens_per_model": {"gpt-4o-mini": 1503},
"default_estimated_output_tokens": 777,
},
2002,
"key global wins over team config",
),
(
{"default_estimated_output_tokens_per_model": {"some-other-model": 9999}},
{
"default_estimated_output_tokens_per_model": {"gpt-4o-mini": 1503},
"default_estimated_output_tokens": 777,
},
1503,
"team per-model wins when the key has no applicable entry",
),
(
{},
{"default_estimated_output_tokens": 777},
777,
"team global is the last configured tier",
),
({}, {}, STATIC_OUTPUT_FLOOR, "unconfigured falls back to the static floor"),
(
{"unrelated": "value"},
{"unrelated": "value"},
STATIC_OUTPUT_FLOOR,
"unrelated metadata changes nothing",
),
(
{"default_estimated_output_tokens": "not-a-number"},
{},
STATIC_OUTPUT_FLOOR,
"malformed config falls back to the static floor instead of erroring",
),
(
{"default_estimated_output_tokens": 0},
{},
STATIC_OUTPUT_FLOOR,
"a non-positive estimate is rejected, not reserved",
),
],
)
async def test_estimated_output_tokens_resolution_precedence(
monkeypatch, key_metadata, team_metadata, expected_output_estimate, tier
):
"""The no-max_tokens output reservation resolves per key / team / model.
Every configured value here is distinct from the static 1024 floor and
from the input estimate, so the reserved amount identifies which tier the
resolver picked.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(
api_key=hash_token(f"sk-estimate-{expected_output_estimate}-{tier}"),
tpm_limit=1_000_000,
metadata=key_metadata,
team_metadata=team_metadata,
)
reserved = await _reserved_tokens_for(
handler,
local_cache,
user_api_key_dict,
{"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT},
)
assert reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + expected_output_estimate, tier
@pytest.mark.asyncio
async def test_request_max_tokens_outranks_configured_estimate(monkeypatch):
"""An explicit request-level max_tokens stays the top of the precedence order."""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(
api_key=hash_token("sk-estimate-explicit-max-tokens"),
tpm_limit=1_000_000,
metadata={"default_estimated_output_tokens": 2002},
)
reserved = await _reserved_tokens_for(
handler,
local_cache,
user_api_key_dict,
{"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT, "max_tokens": 42},
)
assert reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 42
@pytest.mark.asyncio
async def test_configured_estimate_does_not_apply_to_embeddings(monkeypatch):
"""Embeddings generate no output, so a declared output estimate must not be reserved."""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(
api_key=hash_token("sk-estimate-embeddings"),
tpm_limit=1_000_000,
metadata={"default_estimated_output_tokens": 2002},
)
reserved = await _reserved_tokens_for(
handler,
local_cache,
user_api_key_dict,
{"model": "text-embedding-3-small", "input": "hello"},
call_type="embeddings",
)
assert reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE
@pytest.mark.asyncio
async def test_configured_estimate_applies_to_contentless_requests(monkeypatch):
"""A declared estimate describes generation, so it holds even with no prompt body.
Without config such a request reserves the 1-token floor only; the
declaration is what makes concurrent tool-call continuations countable.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
configured = UserAPIKeyAuth(
api_key=hash_token("sk-estimate-contentless-configured"),
tpm_limit=1_000_000,
metadata={"default_estimated_output_tokens": 2002},
)
unconfigured = UserAPIKeyAuth(
api_key=hash_token("sk-estimate-contentless-plain"),
tpm_limit=1_000_000,
)
assert (
await _reserved_tokens_for(
handler, local_cache, configured, {"model": "gpt-4o-mini", "messages": []}
)
== 2002
)
assert (
await _reserved_tokens_for(
handler, local_cache, unconfigured, {"model": "gpt-4o-mini", "messages": []}
)
== 1
)
@pytest.mark.asyncio
async def test_declared_estimate_never_tightens_the_small_tpm_clamp(monkeypatch):
"""The small-TPM clamp can only be loosened by a declaration, never tightened.
That clamp is the one place the proxy rewrites the caller's generation
budget, and it only fires below a 4096 TPM limit. A declaration above it
raises it, so the tenant is not truncated below what they said their
model emits; a declaration below it changes nothing, because an estimate
describes the typical response and must not become a hard cap that
truncates the tail. The reservation tracks whatever the clamp settles on,
so a small tenant can never generate more than was reserved.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
raised_data: Dict[str, Any] = {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT}
raised_reserved = await _reserved_tokens_for(
handler,
local_cache,
UserAPIKeyAuth(
api_key=hash_token("sk-estimate-hard-cap-raised"),
tpm_limit=2000,
metadata={"default_estimated_output_tokens": 900},
),
raised_data,
)
assert raised_data["max_tokens"] == 900
assert raised_reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 900
lowered_data: Dict[str, Any] = {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT}
lowered_reserved = await _reserved_tokens_for(
handler,
local_cache,
UserAPIKeyAuth(
api_key=hash_token("sk-estimate-hard-cap-lowered"),
tpm_limit=2000,
metadata={"default_estimated_output_tokens": 120},
),
lowered_data,
)
assert lowered_data["max_tokens"] == 500
assert lowered_reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 500
unconfigured_data: Dict[str, Any] = {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT}
unconfigured_reserved = await _reserved_tokens_for(
handler,
local_cache,
UserAPIKeyAuth(
api_key=hash_token("sk-estimate-hard-cap-plain"),
tpm_limit=2000,
),
unconfigured_data,
)
assert unconfigured_data["max_tokens"] == 500
assert unconfigured_reserved == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 500
@pytest.mark.asyncio
async def test_one_malformed_estimate_field_does_not_discard_the_other(monkeypatch):
"""Each declared field is validated on its own.
A per-model map with a bad entry must not take a valid global estimate
down with it, and a bad global must not hide a valid per-model entry.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
broken_map = await _reserved_tokens_for(
handler,
local_cache,
UserAPIKeyAuth(
api_key=hash_token("sk-estimate-broken-map"),
tpm_limit=1_000_000,
metadata={
"default_estimated_output_tokens_per_model": {"gpt-4o-mini": "huge"},
"default_estimated_output_tokens": 2002,
},
),
{"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT},
)
assert broken_map == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 2002
broken_global = await _reserved_tokens_for(
handler,
local_cache,
UserAPIKeyAuth(
api_key=hash_token("sk-estimate-broken-global"),
tpm_limit=1_000_000,
metadata={
"default_estimated_output_tokens_per_model": {"gpt-4o-mini": 3001},
"default_estimated_output_tokens": -5,
},
),
{"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT},
)
assert broken_global == ONE_TOKEN_PROMPT_INPUT_ESTIMATE + 3001
@pytest.mark.asyncio
@pytest.mark.parametrize("declared", [100_000, 5000])
async def test_declared_estimate_over_the_tpm_budget_is_honored_and_explained(monkeypatch, caplog, declared):
"""A declaration bigger than the budget must not be silently shrunk.
Capping it against the TPM limit would re-admit exactly the traffic this
feature exists to hold back, so the request is refused instead and the
reservation is explained rather than leaving an unexplained 429 loop.
``declared == tpm_limit`` is the boundary case: the declaration alone
equals the limit, so only adding the input estimate tips the reservation
over. Comparing the declaration against the limit rather than the
reservation would refuse this request while saying nothing.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(
api_key=hash_token(f"sk-estimate-over-budget-{declared}"),
tpm_limit=5000,
metadata={"default_estimated_output_tokens": declared},
)
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT},
call_type="completion",
)
assert exc_info.value.status_code == 429
explained = [
record.getMessage()
for record in caplog.records
if "cannot be admitted even against an empty window" in record.getMessage()
]
assert len(explained) == 1, f"expected exactly one explanation, got {explained}"
assert str(declared) in explained[0]
assert str(ONE_TOKEN_PROMPT_INPUT_ESTIMATE + declared) in explained[0]
assert "5000" in explained[0]
@pytest.mark.asyncio
async def test_a_key_that_declared_nothing_is_never_blamed_for_a_declaration(monkeypatch, caplog):
"""A request can outgrow its budget on prompt size alone, with no declaration.
The heuristic path reserves input plus the injected clamp, so a long
prompt against a small limit is refused without anyone having declared
anything. Blaming the declared field there would point an operator at a
setting they never set, to fix a 429 whose real cause is prompt size.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
api_key=hash_token("sk-undeclared-long-prompt"),
tpm_limit=1000,
),
cache=local_cache,
data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "x" * 3600}]},
call_type="completion",
)
assert exc_info.value.status_code == 429
assert not [
record for record in caplog.records if "cannot be admitted even against an empty window" in record.getMessage()
]
@pytest.mark.asyncio
async def test_declared_estimate_inside_the_tpm_budget_is_not_explained(monkeypatch, caplog):
"""The explanation is for requests that cannot fit, not for every request.
Without this, a correctly configured key would emit one line per call.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
await handler.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
api_key=hash_token("sk-estimate-within-budget"),
tpm_limit=5000,
metadata={"default_estimated_output_tokens": 1000},
),
cache=local_cache,
data={"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT},
call_type="completion",
)
assert not [
record for record in caplog.records if "cannot be admitted even against an empty window" in record.getMessage()
]
@pytest.mark.asyncio
async def test_configured_estimate_blocks_the_overrun_the_static_floor_admits(monkeypatch):
"""Concurrent unbounded requests must stop at the declared budget.
A key with tpm_limit=8000 whose model really emits ~3000 output tokens
admits 7 concurrent requests under the 1024 floor (7 * 1025 <= 8000), so
once they all report actual usage the window carries ~21000 tokens
against an 8000 limit. Declaring the real output size admits only the two
requests the budget actually covers.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
async def admitted(metadata):
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(
api_key=hash_token(f"sk-overrun-{metadata}"),
tpm_limit=8000,
metadata=metadata,
)
accepted = 0
for _ in range(10):
try:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT},
call_type="completion",
)
except HTTPException:
break
accepted += 1
return accepted
assert await admitted({}) == 7
assert await admitted({"default_estimated_output_tokens": 3000}) == 2

View file

@ -1,11 +1,18 @@
import ast
import importlib.util
from pathlib import Path
from types import ModuleType
from typing import Annotated
import fastapi.dependencies.utils as fastapi_dependency_utils
import pytest
from fastapi import Depends, FastAPI, Header, Query, Request
from fastapi.testclient import TestClient
import litellm.proxy.management_endpoints.management_v1.common as common_module
from litellm.proxy.management_endpoints.management_v1.common import (
ManagementProblem,
PROBLEM_CONTENT_TYPE,
ManagementProblem,
_declared_query_params,
problem_response,
reject_unknown_query_params,
@ -93,3 +100,57 @@ def test_declared_query_params_is_empty_when_the_route_has_no_dependant():
}
)
assert _declared_query_params(request) == frozenset()
# fastapi removed these in 0.140.7, which `pyproject.toml` still allows via
# `fastapi>=0.136.3,<1.0`. Add a name here whenever a supported release drops one.
FASTAPI_NAMES_REMOVED_IN_0_140_7 = frozenset({"get_flat_dependant"})
MANAGEMENT_V1_PACKAGE = Path(str(common_module.__file__)).parent
def _public_names(module: ModuleType) -> frozenset[str]:
return frozenset(name for name in vars(module) if not name.startswith("_"))
def _fastapi_names_imported_by(source_file: Path) -> frozenset[str]:
tree = ast.parse(source_file.read_text())
return frozenset(
alias.name
for node in ast.walk(tree)
if isinstance(node, ast.ImportFrom) and (node.module or "").startswith("fastapi")
for alias in node.names
)
@pytest.mark.parametrize(
"source_file", sorted(MANAGEMENT_V1_PACKAGE.glob("*.py")), ids=lambda path: path.name
)
def test_no_module_imports_a_fastapi_name_removed_in_a_supported_release(source_file: Path):
"""`pyproject.toml` allows fastapi up to <1.0, but CI only ever resolves 0.136.3.
Every other test here passes just as well against a module importing a name
fastapi has since deleted, because the pinned fastapi still has it. On a user's
fastapi>=0.140.7 that import is an ImportError, and `proxy_server` imports this
package unguarded at module level, so it takes the whole proxy down rather than
just these routes. Globbing the package means a new module is covered on sight.
"""
assert not _fastapi_names_imported_by(source_file) & FASTAPI_NAMES_REMOVED_IN_0_140_7
def test_common_still_imports_when_fastapi_has_dropped_those_names(monkeypatch: pytest.MonkeyPatch):
"""The static check above cannot prove the module actually loads; this does.
Behaviour cannot be asserted under the same simulation: on 0.136.3
`get_flat_params` calls `get_flat_dependant` internally, so it raises NameError
once the name is gone. Loading is the part this pins.
"""
for name in FASTAPI_NAMES_REMOVED_IN_0_140_7:
monkeypatch.delattr(fastapi_dependency_utils, name, raising=False)
spec = importlib.util.spec_from_file_location(
"management_v1_common__simulated_fastapi", Path(str(common_module.__file__))
)
assert spec is not None and spec.loader is not None
reimported = importlib.util.module_from_spec(spec)
spec.loader.exec_module(reimported)
assert _public_names(reimported) == _public_names(common_module)

View file

@ -72,6 +72,39 @@ async def test_new_budget_success(client_and_mocks):
mock_table.create.assert_awaited_once()
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
@pytest.mark.asyncio
async def test_new_budget_rejects_a_duration_that_never_advances(
client_and_mocks, bad_duration
):
"""A zero-length window resets to "now", so the row is due again the moment
it is written and the reset job re-reads it on every tick forever."""
client, _, mock_table = client_and_mocks
resp = client.post(
"/budget/new",
json={"budget_id": "budget_bad", "max_budget": 10.0, "budget_duration": bad_duration},
)
assert resp.status_code == 400, resp.text
assert "Invalid budget_duration" in resp.json()["detail"]["error"]
mock_table.create.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_budget_rejects_a_duration_that_never_advances(client_and_mocks):
client, _, mock_table = client_and_mocks
resp = client.post(
"/budget/update",
json={"budget_id": "budget_456", "budget_duration": "0s"},
)
assert resp.status_code == 400, resp.text
assert "Invalid budget_duration" in resp.json()["detail"]["error"]
mock_table.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_new_budget_db_not_connected(client_and_mocks, monkeypatch):
client, mock_prisma, mock_table = client_and_mocks

View file

@ -7,9 +7,9 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
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 +108,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 +300,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 +326,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 +355,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 +363,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 +388,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 +412,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 +443,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 +616,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 +735,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 +804,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 +832,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 +865,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 +1142,580 @@ class TestEverySavingsDriverSurvivesTheReadPath:
assert f"total_{driver}" in DailySpendMetadata.model_fields, (
f"total_{driver} is missing, so the range summary omits the driver"
)
@pytest.fixture
def ptu_cost_attribution_enabled(monkeypatch):
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
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(ptu_cost_attribution_enabled):
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(ptu_cost_attribution_enabled):
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(ptu_cost_attribution_enabled):
"""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(ptu_cost_attribution_enabled):
"""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(ptu_cost_attribution_enabled):
"""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."""
@pytest.fixture(autouse=True)
def _enabled(self, ptu_cost_attribution_enabled):
"""Flat cost is gated off by default, and these assert on the amounts."""
@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)
def _daily_team_row(api_key, *, spend=0.0, ptu_flat_cost=0.0):
"""A LiteLLM_DailyTeamSpend row as the paginated read path receives it from find_many."""
base: Final = _spend_record(api_key, spend=spend, ptu_flat_cost=ptu_flat_cost)
return SimpleNamespace(**{**base.__dict__, "date": "2026-07-01", "team_id": "team-1"})
class TestPtuCostAttributionDisabled:
"""With LITELLM_ENABLE_PTU_COST_ATTRIBUTION unset, both read paths report zero flat
cost, while the sentinel filtering that keeps ``__ptu_flat_cost__`` out of the
breakdowns keeps running.
Filtering is deliberately not gated: an operator can enable the flag, accrue
sentinel rows, then disable it, and those rows stay in LiteLLM_DailyTeamSpend
forever. Gating the filter too would surface the sentinel as a bogus api_key and
mint a provider bucket for its empty provider.
"""
@pytest.fixture(autouse=True)
def _flag_off(self, monkeypatch):
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
def test_paginated_path_reports_zero_flat_cost(self):
metrics = update_metrics(SpendMetrics(), _spend_record("real-key", spend=1.0, ptu_flat_cost=240.0))
assert metrics.flat_cost == 0.0
assert metrics.spend == 1.0
def test_aggregated_path_reports_zero_flat_cost(self):
from litellm.proxy.management_endpoints.common_daily_activity import _GROUP_GRAND_TOTAL
metrics = _record_to_spend_metrics(_grouping_row(_GROUP_GRAND_TOTAL, spend=5.0, ptu_flat_cost=240.0))
assert metrics.flat_cost == 0.0
assert metrics.spend == 5.0
def test_aggregated_totals_and_buckets_report_zero_flat_cost(self):
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_GRAND_TOTAL,
_aggregate_grouping_sets_records_sync,
)
records = [
_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_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 == 0.0
assert aggregated["totals"].spend == 5.0
assert aggregated["results"][0].breakdown.models["gpt-4o-mini-ptu"].metrics.flat_cost == 0.0
def test_sentinel_still_excluded_from_the_api_key_breakdown(self):
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), {}, {}, {})
update_breakdown_metrics(
breakdown, _spend_record(PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0), {}, {}, {}, entity_id_field="team_id"
)
assert PTU_SENTINEL_API_KEY not in breakdown.api_keys
assert PTU_SENTINEL_API_KEY not in breakdown.models["gpt-4o-mini-ptu"].api_key_breakdown
assert "real-key" in breakdown.models["gpt-4o-mini-ptu"].api_key_breakdown
def test_sentinel_still_excluded_from_the_provider_breakdown(self):
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(PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0), {}, {}, {})
assert breakdown.providers == {}
def test_grouping_sets_sentinel_still_excluded_from_breakdowns(self):
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_DATE_PROVIDER,
_aggregate_grouping_sets_records_sync,
)
records = [
_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=PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0
),
_grouping_row(_GROUP_DATE_PROVIDER, custom_llm_provider="", ptu_flat_cost=240.0),
]
day = _aggregate_grouping_sets_records_sync(records=records, api_key_metadata={})["results"][0]
assert PTU_SENTINEL_API_KEY not in day.breakdown.api_keys
assert PTU_SENTINEL_API_KEY not in day.breakdown.models["gpt-4o-mini-ptu"].api_key_breakdown
assert sum(bucket.metrics.flat_cost for bucket in day.breakdown.providers.values()) == 0.0
@pytest.mark.asyncio
async def test_team_daily_activity_endpoint_reports_zero_flat_cost(self):
"""/team/daily/activity reads rows with find_many rather than the aggregated SQL, so
forcing the SQL select to a constant zero would leave this path reporting flat cost."""
from litellm.constants import PTU_SENTINEL_API_KEY
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=2)
mock_table.find_many = AsyncMock(
return_value=[
_daily_team_row("real-key", spend=5.0),
_daily_team_row(PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0),
]
)
mock_prisma.db.litellm_verificationtoken = MagicMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_dailyteamspend = mock_table
result = await get_daily_activity(
prisma_client=mock_prisma,
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id="team-1",
entity_metadata_field=None,
start_date="2026-07-01",
end_date="2026-07-01",
model=None,
api_key=None,
page=1,
page_size=50,
)
assert result.metadata.total_flat_cost == 0.0
assert result.metadata.total_spend == 5.0
assert PTU_SENTINEL_API_KEY not in result.results[0].breakdown.api_keys
@pytest.mark.asyncio
async def test_team_daily_activity_endpoint_reports_flat_cost_once_enabled(self, monkeypatch):
from litellm.constants import PTU_SENTINEL_API_KEY
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=2)
mock_table.find_many = AsyncMock(
return_value=[
_daily_team_row("real-key", spend=5.0),
_daily_team_row(PTU_SENTINEL_API_KEY, ptu_flat_cost=240.0),
]
)
mock_prisma.db.litellm_verificationtoken = MagicMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_dailyteamspend = mock_table
result = await get_daily_activity(
prisma_client=mock_prisma,
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id="team-1",
entity_metadata_field=None,
start_date="2026-07-01",
end_date="2026-07-01",
model=None,
api_key=None,
page=1,
page_size=50,
)
assert result.metadata.total_flat_cost == 240.0
assert PTU_SENTINEL_API_KEY not in result.results[0].breakdown.api_keys
class TestFlagIsNotReadOnTheHotPath:
"""update_metrics runs once per accumulation and a record fans out across roughly a
dozen breakdowns, so a flag that reads through the secret manager must not be consulted
for rows that carry no flat cost at all."""
@staticmethod
def _count_flag_reads(records):
import litellm.proxy.management_endpoints.common_daily_activity as cda
from litellm.types.proxy.management_endpoints.common_daily_activity import BreakdownMetrics
reads = []
real = cda.is_ptu_cost_attribution_enabled
def counted():
reads.append(1)
return real()
cda.is_ptu_cost_attribution_enabled = counted
try:
breakdown = BreakdownMetrics()
for record in records:
cda.update_breakdown_metrics(breakdown, record, {}, {}, {})
finally:
cda.is_ptu_cost_attribution_enabled = real
return len(reads)
def test_a_request_row_never_reads_the_flag(self):
reads = self._count_flag_reads([_spend_record("real-key", spend=5.0, ptu_flat_cost=0.0)])
assert reads == 0, f"{reads} secret-manager lookups for a row with no flat cost"
def test_a_page_of_request_rows_never_reads_the_flag(self):
rows = [_spend_record(f"key-{i}", spend=1.0, ptu_flat_cost=0.0) for i in range(50)]
assert self._count_flag_reads(rows) == 0
def test_a_sentinel_row_still_consults_the_flag(self):
from litellm.constants import PTU_SENTINEL_API_KEY
reads = self._count_flag_reads([_spend_record(PTU_SENTINEL_API_KEY, spend=0.0, ptu_flat_cost=240.0)])
assert reads > 0

View file

@ -628,6 +628,58 @@ class TestValidateFiniteSpendErrorDetail:
}
class TestValidateBudgetDuration:
"""`validate_budget_duration` keeps durations that never advance out of the
database.
A duration of "0s" resolves to a reset time of now, so the row is due again
the instant it is written. The reset job re-reads such rows on every tick
and, once one tenant owns enough of them, they fill each batch and starve
every other tenant's reset.
"""
def test_none_is_allowed(self):
from litellm.proxy.management_endpoints.common_utils import (
validate_budget_duration,
)
assert validate_budget_duration(None) is None
@pytest.mark.parametrize("duration", ["30s", "5m", "1h", "1d", "7d", "30d", "1mo"])
def test_positive_durations_are_allowed(self, duration):
from litellm.proxy.management_endpoints.common_utils import (
validate_budget_duration,
)
assert validate_budget_duration(duration) is None
@pytest.mark.parametrize("duration", ["0s", "0m", "0h", "0d", "-5m", "abc", ""])
def test_non_advancing_durations_are_rejected(self, duration):
from fastapi import HTTPException
from litellm.proxy.management_endpoints.common_utils import (
validate_budget_duration,
)
with pytest.raises(HTTPException) as exc_info:
validate_budget_duration(duration)
assert exc_info.value.status_code == 400
def test_rejection_detail_is_exact(self):
from fastapi import HTTPException
from litellm.proxy.management_endpoints.common_utils import (
validate_budget_duration,
)
with pytest.raises(HTTPException) as exc_info:
validate_budget_duration("0s")
assert exc_info.value.detail == {
"error": "Invalid budget_duration '0s'. Use a format like '1h', '24h', '7d', or '30d'."
}
class TestRequireCallerUserIdErrorDetail:
"""The 403 for a service-account key must carry the exact error body."""

View file

@ -749,6 +749,40 @@ def test_char_new_body(mock_prisma_client, mock_user_api_key_auth):
assert response.json() == _EXPECTED_CUSTOMER
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
def test_customer_new_rejects_a_duration_that_never_advances(
mock_prisma_client, mock_user_api_key_auth, bad_duration
):
"""A zero-length window resets to "now", leaving the customer's budget row
permanently due for the reset job to re-read every tick."""
mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW))
response = client.post(
"/customer/new",
json={"user_id": "c1", "max_budget": 10.0, "budget_duration": bad_duration},
headers={"Authorization": "Bearer k"},
)
assert response.status_code == 400, response.text
assert "Invalid budget_duration" in response.text
mock_prisma_client.db.litellm_endusertable.create.assert_not_awaited()
def test_customer_new_accepts_a_normal_duration(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW))
mock_prisma_client.db.litellm_budgettable.create = AsyncMock(
return_value=_row({"budget_id": "b1", "max_budget": 10.0})
)
response = client.post(
"/customer/new",
json={"user_id": "c1", "max_budget": 10.0, "budget_duration": "30d"},
headers={"Authorization": "Bearer k"},
)
assert response.status_code == 200, response.text
def test_char_update_body(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(
return_value=_row({"user_id": "c1", "blocked": False})

View file

@ -788,6 +788,68 @@ def test_update_internal_user_params_reset_spend_and_max_budget():
assert "budget_duration" not in non_default_values # Should not add default values
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
def test_update_internal_user_params_rejects_a_duration_that_never_advances(bad_duration):
"""A zero-length window resets to "now", so the user row is due again the
moment it is written and the reset job re-reads it on every tick. Enough of
them fill each batch and starve other tenants' resets.
"""
from fastapi import HTTPException
from litellm.proxy._types import UpdateUserRequest
from litellm.proxy.management_endpoints.internal_user_endpoints import (
_update_internal_user_params,
)
data = UpdateUserRequest(user_id="test_user_id", budget_duration=bad_duration)
with pytest.raises(HTTPException) as exc_info:
_update_internal_user_params(data_json=data.model_dump(exclude_unset=True), data=data)
assert exc_info.value.status_code == 400
assert "Invalid budget_duration" in str(exc_info.value.detail)
def test_update_internal_user_params_accepts_a_normal_duration():
from litellm.proxy._types import UpdateUserRequest
from litellm.proxy.management_endpoints.internal_user_endpoints import (
_update_internal_user_params,
)
data = UpdateUserRequest(user_id="test_user_id", budget_duration="30d")
non_default_values = _update_internal_user_params(data_json=data.model_dump(exclude_unset=True), data=data)
assert non_default_values["budget_duration"] == "30d"
assert non_default_values["budget_reset_at"] is not None
@pytest.mark.asyncio
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
async def test_new_user_rejects_a_duration_that_never_advances(mocker, bad_duration):
"""/user/new must reject the same never-advancing durations /user/update does."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
mocker.patch("litellm.proxy.proxy_server.prisma_client", MagicMock())
duplicate_check = mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id",
new=AsyncMock(),
)
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
with pytest.raises(ProxyException) as exc_info:
await new_user(
data=NewUserRequest(budget_duration=bad_duration),
user_api_key_dict=admin,
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
duplicate_check.assert_not_awaited()
@pytest.mark.asyncio
async def test_new_user_license_over_limit(mocker):
"""

View file

@ -2527,6 +2527,72 @@ def _setup_update_key_mocks(monkeypatch, mock_prisma_client):
monkeypatch.setattr("litellm.store_audit_logs", False)
@pytest.mark.asyncio
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
async def test_update_key_rejects_a_duration_that_never_advances(monkeypatch, bad_duration):
"""A zero-length window resets to "now", so the key row is due again the
moment it is written. The reset job re-reads such rows on every tick, and a
tenant with enough of them fills each batch and starves other tenants.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b"
key_in_db = LiteLLM_VerificationToken(token=hashed_token, user_id="test-user")
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=key_in_db
)
mock_prisma_client.update_data = AsyncMock()
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
with pytest.raises(ProxyException) as exc_info:
await update_key_fn(
request=MagicMock(),
data=UpdateKeyRequest(key=hashed_token, budget_duration=bad_duration),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
),
litellm_changed_by=None,
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
mock_prisma_client.update_data.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
async def test_generate_key_rejects_a_duration_that_never_advances(monkeypatch, bad_duration):
"""/key/generate must reject the same never-advancing durations /key/update does."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_fn,
)
mock_prisma_client = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new=AsyncMock(),
) as mock_generate:
with pytest.raises(ProxyException) as exc_info:
await generate_key_fn(
data=GenerateKeyRequest(budget_duration=bad_duration),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234"
),
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
mock_generate.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_key_by_alias_only(monkeypatch):
"""
@ -8081,7 +8147,7 @@ async def test_key_with_budget_id_does_not_store_budget_duration():
budget_duration, the key does NOT get budget_duration stored on it.
Keys with budget_id follow their linked budget tier's reset schedule;
reset_budget_for_keys_linked_to_budgets() resets them when the tier resets.
reset_budget_for_litellm_budget_table() resets them when the tier resets.
This avoids duplicating budget_duration on keys so tier updates apply
automatically to all linked keys.
"""
@ -15442,3 +15508,312 @@ async def test_migrate_encryption_endpoint_rejects_proxy_admin_viewer():
assert exc_info.value.status_code == 403
mock_migrate.assert_not_awaited()
_ESTIMATE = "default_estimated_output_tokens"
_ESTIMATE_PER_MODEL = "default_estimated_output_tokens_per_model"
@pytest.mark.parametrize(
"label, request_body, existing_metadata, allowed",
[
("nothing declared", {}, None, True),
("declared top-level on a key with none stored", {_ESTIMATE: 1}, None, False),
("declared inside metadata on a key with none stored", {"metadata": {_ESTIMATE: 1}}, None, False),
(
"per-model map declared inside metadata",
{"metadata": {_ESTIMATE_PER_MODEL: {"gpt-4": 1}}},
None,
False,
),
("unrelated edit, metadata omitted", {"models": ["gpt-4"]}, {_ESTIMATE: 2000}, True),
("stored value resent unchanged", {_ESTIMATE: 2000}, {_ESTIMATE: 2000}, True),
("stored value lowered", {_ESTIMATE: 1}, {_ESTIMATE: 2000}, False),
("stored value raised", {_ESTIMATE: 9000}, {_ESTIMATE: 2000}, False),
(
"stored value cleared by sending a metadata blob without it",
{"metadata": {"other": "keep"}},
{_ESTIMATE: 2000, "other": "keep"},
False,
),
(
"stored value resent inside the metadata blob",
{"metadata": {_ESTIMATE: 2000, "other": "keep"}},
{_ESTIMATE: 2000, "other": "keep"},
True,
),
(
"per-model map resent unchanged",
{_ESTIMATE_PER_MODEL: {"gpt-4": 4096}},
{_ESTIMATE_PER_MODEL: {"gpt-4": 4096}},
True,
),
(
"one model in the per-model map lowered",
{_ESTIMATE_PER_MODEL: {"gpt-4": 1}},
{_ESTIMATE_PER_MODEL: {"gpt-4": 4096}},
False,
),
],
)
def test_output_token_estimate_admin_gate_matrix(label, request_body, existing_metadata, allowed):
"""A non-admin may only leave a key's stored output-token estimate exactly as it is.
The estimate decides what the TPM limiter reserves for a request that omits
max_tokens, so lowering, raising or clearing it moves a reservation charged
against team and organization windows the key holder does not own. Key
metadata is writable by the key holder, and the declaration can be written
either as a dedicated top-level field or nested in the metadata blob, so
both routes are gated. Resending the stored value is what the edit form
produces on every save and has to stay allowed.
"""
from litellm.proxy.auth.auth_utils import (
enforce_output_token_estimates_are_admin_only,
)
def _call(caller):
enforce_output_token_estimates_are_admin_only(
data=UpdateKeyRequest(key="sk-1", **request_body),
existing_metadata=existing_metadata,
user_api_key_dict=caller,
entity="key",
)
non_admin = UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-non-admin",
user_id="alice",
)
if allowed:
_call(non_admin)
else:
with pytest.raises(HTTPException) as exc:
_call(non_admin)
assert exc.value.status_code == 403
assert "Only proxy admins can set" in str(exc.value.detail)
_call(
UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin",
)
)
@pytest.mark.asyncio
async def test_generate_key_output_token_estimate_rejected_for_non_admin():
"""The /key/update gate does not cover generate, so without its own check a
non-admin could self-mint a key that reserves one output token per
unbounded request and overrun the TPM window it is charged against."""
with patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()):
with pytest.raises(HTTPException) as exc:
await _common_key_generation_helper(
data=GenerateKeyRequest(default_estimated_output_tokens=1, tpm_limit=100000),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-alice",
user_id="alice",
),
litellm_changed_by=None,
team_table=None,
)
assert int(getattr(exc.value, "status_code", 0)) == 403
assert "Only proxy admins can set" in str(exc.value.detail)
@pytest.mark.asyncio
async def test_generate_key_output_token_estimate_in_metadata_rejected_for_non_admin():
"""Writing the declaration into the raw metadata blob lands in the same
stored field, so gating only the dedicated top-level field leaves the
bypass wide open."""
with patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()):
with pytest.raises(HTTPException) as exc:
await _common_key_generation_helper(
data=GenerateKeyRequest(metadata={"default_estimated_output_tokens": 1}),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-alice",
user_id="alice",
),
litellm_changed_by=None,
team_table=None,
)
assert int(getattr(exc.value, "status_code", 0)) == 403
@pytest.mark.asyncio
async def test_generate_key_output_token_estimate_allowed_for_admin():
"""A proxy admin declaring the estimate must reach key creation."""
with (
patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()),
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.premium_user", False),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn"
) as mock_generate_key,
):
mock_generate_key.return_value = {
"key": "sk-test-key",
"expires": None,
"user_id": "admin",
"team_id": None,
}
await _common_key_generation_helper(
data=GenerateKeyRequest(default_estimated_output_tokens=200),
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1"),
litellm_changed_by=None,
team_table=None,
)
assert mock_generate_key.called
def _estimate_key_row(token: str, metadata: dict):
existing_key = MagicMock()
existing_key.token = token
existing_key.user_id = "internal_user"
existing_key.created_by = "internal_user"
existing_key.team_id = None
existing_key.project_id = None
existing_key.max_budget = 10.0
existing_key.key_alias = None
existing_key.models = []
existing_key.metadata = metadata
existing_key.model_dump.return_value = {
"token": token,
"user_id": "internal_user",
"team_id": None,
"max_budget": 10.0,
}
return existing_key
def _wire_update_key_fn(monkeypatch, existing_key):
mock_prisma_client = AsyncMock()
updated_key = MagicMock()
updated_key.token = existing_key.token
updated_key.key_alias = "my-alias"
mock_prisma_client.get_data = AsyncMock(return_value=existing_key)
mock_prisma_client.update_data = AsyncMock(return_value=updated_key)
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
monkeypatch.setattr("litellm.store_audit_logs", False)
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", lambda token: existing_key.token)
async def _noop(**kwargs):
pass
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
_noop,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._enforce_unique_key_alias",
_noop,
)
@pytest.mark.asyncio
async def test_update_key_output_token_estimate_lowered_rejected_for_non_admin(monkeypatch):
"""End-to-end wiring: a key's owner reaches /key/update without any admin
check because metadata is a non-budget field, so the gate has to fire
inside the update path itself rather than only in a helper."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
token = "a1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
_wire_update_key_fn(monkeypatch, _estimate_key_row(token, {_ESTIMATE: 4000}))
mock_request = MagicMock()
mock_request.query_params = {}
with pytest.raises(ProxyException) as exc:
await update_key_fn(
request=mock_request,
data=UpdateKeyRequest(key=token, default_estimated_output_tokens=1),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-internal",
user_id="internal_user",
),
litellm_changed_by=None,
)
assert str(exc.value.code) == "403"
assert "Only proxy admins can set" in str(exc.value.message)
@pytest.mark.asyncio
async def test_update_key_output_token_estimate_unchanged_allows_non_admin_edit(monkeypatch):
"""The edit form resends every field it renders, so gating on presence
would 403 a key owner renaming their own key."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
token = "b1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
_wire_update_key_fn(monkeypatch, _estimate_key_row(token, {_ESTIMATE: 4000}))
mock_request = MagicMock()
mock_request.query_params = {}
result = await update_key_fn(
request=mock_request,
data=UpdateKeyRequest(key=token, key_alias="my-alias", default_estimated_output_tokens=4000),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-internal",
user_id="internal_user",
),
litellm_changed_by=None,
)
assert result is not None
@pytest.mark.asyncio
async def test_regenerate_key_output_token_estimate_lowered_rejected_for_non_admin():
"""/key/regenerate is a third write path into the same stored metadata.
can_modify_verification_token lets a key's own holder regenerate it, and
the request body runs through prepare_key_update_data exactly as an update
does, so gating only generate and update leaves the declaration writable.
"""
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
_execute_virtual_key_regeneration,
)
token = "c1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
key_in_db = LiteLLM_VerificationToken(
token=token,
user_id="internal_user",
metadata={_ESTIMATE: 4000},
)
with pytest.raises(HTTPException) as exc:
await _execute_virtual_key_regeneration(
prisma_client=AsyncMock(),
key_in_db=key_in_db,
hashed_api_key=token,
key="sk-original",
data=RegenerateKeyRequest(key="sk-original", default_estimated_output_tokens=1),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-internal",
user_id="internal_user",
),
litellm_changed_by=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
)
assert exc.value.status_code == 403
assert "Only proxy admins can set" in str(exc.value.detail)

View file

@ -0,0 +1,709 @@
"""Tests for PTU config on the model deployment (v1 model-settings design)."""
import datetime
import json
from contextlib import ExitStack
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm.proxy._types import LiteLLM_ProxyModelTable, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.model_management_endpoints import (
_merged_ptu_model_info,
_raise_if_ptu_cost_attribution_disabled,
_validate_ptu_model_info,
add_new_model,
update_db_model,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
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,
ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc),
)
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,
ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc),
)
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,
ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc),
)
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."""
@pytest.fixture(autouse=True)
def _enabled(self, monkeypatch):
"""PTU writes are gated off by default; these are about the validator, not the gate."""
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
@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."""
@pytest.fixture(autouse=True)
def _enabled(self, monkeypatch):
"""PTU writes are gated off by default; these are about the validator, not the gate."""
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
@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_the_gate_refuses_before_the_team_write(self, monkeypatch):
"""The gate lived inside update_db_model, which runs after the team ACL write, so a
rejected edit still moved the model between teams."""
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
db_model = Deployment(
model_name="gpt-4o",
litellm_params=LiteLLM_Params(model="openai/gpt-4o"),
model_info=ModelInfo(id="dep-0", team_id="team-A"),
)
patch = updateDeployment(
model_info=ModelInfo(
id="dep-0",
team_id="team-B",
ptu_count=15,
cost_per_ptu_per_hour=2.0,
ptu_effective_from=datetime.datetime(2026, 8, 1, tzinfo=datetime.timezone.utc),
)
)
touched = []
with pytest.raises(HTTPException) as exc:
await self._run(db_model, patch, monkeypatch, touched)
assert PTU_COST_ATTRIBUTION_ENV_VAR in exc.value.detail
assert touched == []
@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
class TestPtuCostAttributionGate:
"""PTU config is only writable once an operator sets LITELLM_ENABLE_PTU_COST_ATTRIBUTION.
The fields are rejected rather than dropped: a silent accept-and-drop would let a
caller believe a flat cost was configured while the rollup that prices it is not
even scheduled.
"""
@pytest.fixture(autouse=True)
def _flag_off(self, monkeypatch):
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
@pytest.fixture
def flag_on(self, monkeypatch):
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
@pytest.mark.parametrize(
"model_info",
[
{"team_id": "t", "ptu_count": 5, "cost_per_ptu_per_hour": 2.0},
{"ptu_count": 5},
{"cost_per_ptu_per_hour": 2.0},
{"ptu_effective_from": "2026-08-01T00:00:00Z"},
{"ptu_effective_to": "2026-08-02T00:00:00Z"},
],
)
def test_rejects_any_ptu_field_while_disabled(self, model_info):
with pytest.raises(HTTPException) as exc:
_raise_if_ptu_cost_attribution_disabled(model_info)
assert exc.value.status_code == 400
assert PTU_COST_ATTRIBUTION_ENV_VAR in exc.value.detail
def test_names_every_offending_field(self):
with pytest.raises(HTTPException) as exc:
_raise_if_ptu_cost_attribution_disabled({"ptu_count": 5, "cost_per_ptu_per_hour": 2.0})
assert "ptu_count" in exc.value.detail
assert "cost_per_ptu_per_hour" in exc.value.detail
def test_allows_a_request_without_ptu_fields_while_disabled(self):
_raise_if_ptu_cost_attribution_disabled({"team_id": "t", "access_groups": ["a"]})
def test_allows_every_ptu_field_once_enabled(self, flag_on):
_raise_if_ptu_cost_attribution_disabled(
{
"team_id": "t",
"ptu_count": 5,
"cost_per_ptu_per_hour": 2.0,
"ptu_effective_from": "2026-08-01T00:00:00Z",
"ptu_effective_to": "2026-08-02T00:00:00Z",
}
)
def _deployment_without_ptu() -> Deployment:
return Deployment(
model_name="gpt-4o",
litellm_params=LiteLLM_Params(model="openai/gpt-4o"),
model_info=ModelInfo(id="dep-0", team_id="t"),
)
def _deployment_with_stored_ptu() -> Deployment:
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=15,
cost_per_ptu_per_hour=2.0,
ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc),
),
)
class TestUpdateDbModelPtuGate:
@pytest.fixture(autouse=True)
def _flag_off(self, monkeypatch):
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
def test_patch_carrying_ptu_config_is_rejected(self):
with pytest.raises(HTTPException) as exc:
update_db_model(
db_model=_deployment_without_ptu(),
updated_patch=updateDeployment(model_info=ModelInfo(id="dep-0", team_id="t", ptu_count=15)),
)
assert exc.value.status_code == 400
def test_patch_that_touches_nothing_ptu_still_succeeds(self):
result = update_db_model(
db_model=_deployment_without_ptu(),
updated_patch=updateDeployment(model_info=ModelInfo(id="dep-0", access_groups=["a"])),
)
assert json.loads(result["model_info"])["access_groups"] == ["a"]
def test_unrelated_patch_of_a_model_that_stores_ptu_config_is_not_blocked(self):
"""A deployment configured during an earlier opt-in stays editable: the gate reads the
incoming patch, not the merged deployment, so the stored config is left in place."""
result = update_db_model(
db_model=_deployment_with_stored_ptu(),
updated_patch=updateDeployment(model_name="gpt-4o-renamed"),
)
assert result["model_name"] == "gpt-4o-renamed"
def test_explicit_nulls_do_not_erase_stored_ptu_config_while_disabled(self):
"""A client round-tripping a model_info blob sends the PTU keys as nulls. While the
feature is disabled those nulls must not reach the clear loop: disabling pauses PTU,
it does not silently discard a billing configuration the operator set up earlier."""
result = update_db_model(
db_model=_deployment_with_stored_ptu(),
updated_patch=updateDeployment(
model_info=ModelInfo(id="dep-0", ptu_count=None, cost_per_ptu_per_hour=None)
),
)
stored = json.loads(result["model_info"])
assert stored["ptu_count"] == 15
assert stored["cost_per_ptu_per_hour"] == 2.0
def test_the_merged_view_agrees_with_the_write_while_disabled(self):
"""The validator sees what the write will store. If the merged view honoured a null the
clear loop ignores, a round-tripped blob would 400 on a half-set pair that never forms."""
merged = _merged_ptu_model_info(
db_model=_deployment_with_stored_ptu(),
patch_data=updateDeployment(model_info=ModelInfo(id="dep-0", ptu_count=None)),
)
assert merged["ptu_count"] == 15
_validate_ptu_model_info(merged)
def test_explicit_nulls_still_clear_once_enabled(self, monkeypatch):
"""Clearing remains available to an operator who opted in, which is how PTU config is
removed from a deployment."""
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
result = update_db_model(
db_model=_deployment_with_stored_ptu(),
updated_patch=updateDeployment(
model_info=ModelInfo(id="dep-0", ptu_count=None, cost_per_ptu_per_hour=None)
),
)
stored = json.loads(result["model_info"])
assert "ptu_count" not in stored
assert "cost_per_ptu_per_hour" not in stored
def test_patch_carrying_ptu_config_is_accepted_once_enabled(self, monkeypatch):
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
result = update_db_model(
db_model=_deployment_without_ptu(),
updated_patch=updateDeployment(
model_info=ModelInfo(
id="dep-0",
team_id="t",
ptu_count=15,
cost_per_ptu_per_hour=2.0,
ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc),
)
),
)
stored = json.loads(result["model_info"])
assert stored["ptu_count"] == 15
assert stored["cost_per_ptu_per_hour"] == 2.0
class TestAddNewModelPtuGate:
@pytest.fixture(autouse=True)
def _flag_off(self, monkeypatch):
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
@staticmethod
def _patched_proxy(model_id: str):
"""Patch everything /model/new touches except the PTU gate, and hand back the DB writers."""
db_row = LiteLLM_ProxyModelTable(
model_id=model_id,
model_name="ptu-model",
litellm_params={"model": "openai/gpt-4.1-nano"},
model_info={"id": model_id},
created_by="test-admin",
updated_by="test-admin",
)
add_model_to_db = AsyncMock(return_value=db_row)
add_team_model_to_db = AsyncMock(return_value=db_row)
mock_proxy_config = MagicMock()
mock_proxy_config.add_deployment = AsyncMock(return_value=None)
mock_router = MagicMock()
mock_router.get_model_ids.return_value = [model_id]
proxy_server = "litellm.proxy.proxy_server"
endpoints = "litellm.proxy.management_endpoints.model_management_endpoints"
return (add_model_to_db, add_team_model_to_db), [
patch(f"{proxy_server}.prisma_client", MagicMock()),
patch(f"{proxy_server}.store_model_in_db", True),
patch(f"{proxy_server}.proxy_config", mock_proxy_config),
patch(f"{proxy_server}.proxy_logging_obj", MagicMock()),
patch(f"{proxy_server}.general_settings", {}),
patch(f"{proxy_server}.premium_user", True),
patch(f"{proxy_server}.llm_router", mock_router),
patch(
f"{endpoints}.ModelManagementAuthChecks.can_user_make_model_call",
AsyncMock(return_value=True),
),
patch(f"{endpoints}._add_model_to_db", add_model_to_db),
patch(f"{endpoints}._add_team_model_to_db", add_team_model_to_db),
]
@staticmethod
def _ptu_deployment(model_id: str) -> Deployment:
return Deployment(
model_name="ptu-model",
litellm_params=LiteLLM_Params(model="openai/gpt-4.1-nano", api_key="fake-key"),
model_info=ModelInfo(
id=model_id,
team_id="team-1",
ptu_count=15,
cost_per_ptu_per_hour=2.0,
ptu_effective_from=datetime.datetime(2020, 1, 1, tzinfo=datetime.timezone.utc),
),
)
@pytest.mark.asyncio
async def test_model_new_rejects_ptu_config_while_disabled(self):
(add_model_to_db, add_team_model_to_db), patches = self._patched_proxy("ptu-gate-model")
admin = UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN)
with ExitStack() as stack:
for active_patch in patches:
stack.enter_context(active_patch)
with pytest.raises(Exception) as exc:
await add_new_model(model_params=self._ptu_deployment("ptu-gate-model"), user_api_key_dict=admin)
assert PTU_COST_ATTRIBUTION_ENV_VAR in str(exc.value)
add_model_to_db.assert_not_called()
add_team_model_to_db.assert_not_called()
@pytest.mark.asyncio
async def test_model_new_accepts_a_deployment_without_ptu_config_while_disabled(self):
_, patches = self._patched_proxy("plain-model")
admin = UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN)
with ExitStack() as stack:
for active_patch in patches:
stack.enter_context(active_patch)
result = await add_new_model(
model_params=Deployment(
model_name="ptu-model",
litellm_params=LiteLLM_Params(model="openai/gpt-4.1-nano", api_key="fake-key"),
model_info=ModelInfo(id="plain-model"),
),
user_api_key_dict=admin,
)
assert result.model_id == "plain-model"
@pytest.mark.asyncio
async def test_model_new_accepts_ptu_config_once_enabled(self, monkeypatch):
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
(_, add_team_model_to_db), patches = self._patched_proxy("ptu-gate-model")
admin = UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN)
with ExitStack() as stack:
for active_patch in patches:
stack.enter_context(active_patch)
result = await add_new_model(model_params=self._ptu_deployment("ptu-gate-model"), user_api_key_dict=admin)
assert result.model_id == "ptu-gate-model"
add_team_model_to_db.assert_called_once()

View file

@ -380,6 +380,66 @@ async def test_update_team_permissions_success(mock_db_client, mock_admin_auth):
app.dependency_overrides = {}
@pytest.mark.asyncio
@pytest.mark.parametrize("field", ["budget_duration", "team_member_budget_duration"])
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
async def test_new_team_rejects_a_duration_that_never_advances(
mock_db_client, mock_admin_auth, field, bad_duration
):
"""A zero-length window resets to "now", so the team row is due again the
moment it is written. The reset job re-reads such rows on every tick, and a
tenant with enough of them fills each batch and starves other tenants.
"""
from fastapi import Request
from litellm.proxy._types import NewTeamRequest
from litellm.proxy.management_endpoints.team_endpoints import new_team
mock_db_client.db = MagicMock()
mock_team_create = AsyncMock()
mock_db_client.db.litellm_teamtable = MagicMock()
mock_db_client.db.litellm_teamtable.create = mock_team_create
with pytest.raises(ProxyException) as exc_info:
await new_team(
data=NewTeamRequest(team_alias="my-team", **{field: bad_duration}),
http_request=MagicMock(spec=Request),
user_api_key_dict=mock_admin_auth,
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
mock_team_create.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("field", ["budget_duration", "team_member_budget_duration"])
async def test_update_team_rejects_a_duration_that_never_advances(
mock_db_client, mock_admin_auth, field
):
"""/team/update must reject the same never-advancing durations /team/new does."""
from fastapi import Request
from litellm.proxy._types import UpdateTeamRequest
from litellm.proxy.management_endpoints.team_endpoints import update_team
mock_db_client.db = MagicMock()
mock_find_unique = AsyncMock(return_value=None)
mock_db_client.db.litellm_teamtable = MagicMock()
mock_db_client.db.litellm_teamtable.find_unique = mock_find_unique
with pytest.raises(ProxyException) as exc_info:
await update_team(
data=UpdateTeamRequest(team_id="team-1", **{field: "0s"}),
http_request=MagicMock(spec=Request),
user_api_key_dict=mock_admin_auth,
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
mock_find_unique.assert_not_awaited()
@pytest.mark.asyncio
async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth):
"""
@ -11008,3 +11068,207 @@ def test_validate_member_user_id_provisioning_caps_the_ids_it_echoes_back():
assert f"u{_MAX_REPORTED_UNKNOWN_USER_IDS}" not in detail
assert f"and {500 - _MAX_REPORTED_UNKNOWN_USER_IDS} more" in detail
assert len(detail) < 1000
_TEAM_ESTIMATE = "default_estimated_output_tokens"
_TEAM_ESTIMATE_PER_MODEL = "default_estimated_output_tokens_per_model"
@pytest.mark.parametrize(
"label, request_body, existing_metadata, allowed",
[
("nothing declared", {}, None, True),
("declared top-level with none stored", {_TEAM_ESTIMATE: 1}, None, False),
("declared inside metadata with none stored", {"metadata": {_TEAM_ESTIMATE: 1}}, None, False),
(
"per-model map declared inside metadata",
{"metadata": {_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 1}}},
None,
False,
),
("unrelated edit, metadata omitted", {"tpm_limit": 99}, {_TEAM_ESTIMATE: 2000}, True),
("stored value resent unchanged", {_TEAM_ESTIMATE: 2000}, {_TEAM_ESTIMATE: 2000}, True),
("stored value lowered", {_TEAM_ESTIMATE: 1}, {_TEAM_ESTIMATE: 2000}, False),
("stored value raised", {_TEAM_ESTIMATE: 9000}, {_TEAM_ESTIMATE: 2000}, False),
(
"stored value cleared by sending a metadata blob without it",
{"metadata": {"other": "keep"}},
{_TEAM_ESTIMATE: 2000, "other": "keep"},
False,
),
(
"stored value resent inside the metadata blob",
{"metadata": {_TEAM_ESTIMATE: 2000, "other": "keep"}},
{_TEAM_ESTIMATE: 2000, "other": "keep"},
True,
),
(
"per-model map resent unchanged",
{_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 4096}},
{_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 4096}},
True,
),
(
"one model in the per-model map lowered",
{_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 1}},
{_TEAM_ESTIMATE_PER_MODEL: {"gpt-4": 4096}},
False,
),
],
)
def test_team_output_token_estimate_admin_gate_matrix(label, request_body, existing_metadata, allowed):
"""A team admin may only leave a team's stored output-token estimate exactly as it is.
A team admin can write team metadata, and every key on the team inherits the
team declaration, so without this a team admin could shrink the reservation
for the whole team and under-reserve against an organization TPM window the
organization set above them. Same value-transition rule as the key gate,
including the raw-metadata route and clearing by omission.
"""
from litellm.proxy._types import UpdateTeamRequest
from litellm.proxy.auth.auth_utils import (
enforce_output_token_estimates_are_admin_only,
)
def _call(caller):
enforce_output_token_estimates_are_admin_only(
data=UpdateTeamRequest(team_id="t", **request_body),
existing_metadata=existing_metadata,
user_api_key_dict=caller,
entity="team",
)
team_admin = UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-team-admin",
user_id="team-admin",
)
if allowed:
_call(team_admin)
else:
with pytest.raises(HTTPException) as exc:
_call(team_admin)
assert exc.value.status_code == 403
assert "on a team" in str(exc.value.detail)
_call(
UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin",
)
)
def _wire_update_team(stack, existing_metadata):
"""Mock just enough of update_team to reach (or pass) the estimate gate."""
from unittest.mock import AsyncMock, MagicMock, patch
mock_prisma_client = stack.enter_context(patch("litellm.proxy.proxy_server.prisma_client"))
stack.enter_context(patch("litellm.proxy.proxy_server.llm_router"))
stack.enter_context(patch("litellm.proxy.proxy_server.user_api_key_cache"))
stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj"))
stack.enter_context(patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"))
stack.enter_context(patch("litellm.proxy.management_endpoints.team_endpoints._cache_team_object"))
existing_team = MagicMock()
existing_team.metadata = existing_metadata
existing_team.model_dump.return_value = {
"team_id": "test_team_id",
"team_alias": "test_team",
"metadata": existing_metadata,
"members_with_roles": [{"user_id": "team-admin", "role": "admin"}],
}
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team)
updated_team = MagicMock()
updated_team.team_id = "test_team_id"
updated_team.model_dump.return_value = {"team_id": "test_team_id"}
mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=updated_team)
mock_prisma_client.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data)
return mock_prisma_client
@pytest.mark.asyncio
async def test_update_team_output_token_estimate_lowered_rejected_for_team_admin():
"""End-to-end wiring: _verify_team_access admits a team admin, so the gate
has to fire inside update_team itself."""
import contextlib
from unittest.mock import Mock
from fastapi import Request
from litellm.proxy._types import UpdateTeamRequest
from litellm.proxy.management_endpoints.team_endpoints import update_team
with contextlib.ExitStack() as stack:
_wire_update_team(stack, {_TEAM_ESTIMATE: 4000})
with pytest.raises(ProxyException) as exc:
await update_team(
data=UpdateTeamRequest(team_id="test_team_id", default_estimated_output_tokens=1),
http_request=Mock(spec=Request),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-team-admin",
user_id="team-admin",
),
)
assert str(exc.value.code) == "403"
assert "on a team" in str(exc.value.message)
@pytest.mark.asyncio
async def test_update_team_output_token_estimate_unchanged_allows_team_admin_edit():
"""The team settings form resends every field it renders, so gating on
presence would break a team admin editing an unrelated setting."""
import contextlib
from unittest.mock import Mock
from fastapi import Request
from litellm.proxy._types import UpdateTeamRequest
from litellm.proxy.management_endpoints.team_endpoints import update_team
with contextlib.ExitStack() as stack:
prisma = _wire_update_team(stack, {_TEAM_ESTIMATE: 4000})
await update_team(
data=UpdateTeamRequest(
team_id="test_team_id",
team_alias="renamed",
default_estimated_output_tokens=4000,
),
http_request=Mock(spec=Request),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-team-admin",
user_id="team-admin",
),
)
assert prisma.db.litellm_teamtable.update.called
@pytest.mark.asyncio
async def test_new_team_output_token_estimate_rejected_for_non_admin():
"""/team/new is the other write path into the same stored declaration."""
from unittest.mock import Mock
from fastapi import Request
from litellm.proxy._types import NewTeamRequest
from litellm.proxy.management_endpoints.team_endpoints import new_team
with pytest.raises(ProxyException) as exc:
await new_team(
data=NewTeamRequest(team_alias="t", default_estimated_output_tokens=1),
http_request=Mock(spec=Request),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="sk-alice",
user_id="alice",
),
)
assert str(exc.value.code) == "403"
assert "on a team" in str(exc.value.message)

View file

@ -2639,3 +2639,67 @@ async def test_ProxyConfig__init_non_llm_configs_empty_agents_key_clears_remembe
assert clean_agent_registry.config_agents == ()
clean_agent_registry.load_agents_from_db_and_config(db_agents=None)
assert clean_agent_registry.get_agent_list() == ()
# ---------------------------------------------------------------------------
# _init_guardrails_in_db
# ---------------------------------------------------------------------------
def _db_guardrail_row(guardrail_id: str, guardrail_type: str) -> dict[str, object]:
return {
"guardrail_id": guardrail_id,
"guardrail_name": f"name-{guardrail_id}",
"litellm_params": {"guardrail": guardrail_type, "mode": "pre_call"},
"guardrail_info": None,
"team_id": None,
}
@pytest.mark.asyncio
async def test_ProxyConfig__init_guardrails_in_db_skips_only_the_unloadable_row(monkeypatch):
"""
A single DB row that fails to initialize used to abort the whole loop, so one
typo'd guardrail type left the proxy running with zero guardrails loaded.
The failing row's id must still reach reconcile_db_guardrails so that eviction
pass cannot treat a row that is alive in the DB as one that was deleted.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy.guardrails import guardrail_registry as registry_module
from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams
class _RecordingHandler(registry_module.InMemoryGuardrailHandler):
def __init__(self) -> None:
super().__init__()
self.reconciled_with: list[set[str]] = []
def reconcile_db_guardrails(self, db_guardrail_ids: set[str]) -> list[str]:
self.reconciled_with.append(set(db_guardrail_ids))
return super().reconcile_db_guardrails(db_guardrail_ids)
handler = _RecordingHandler()
monkeypatch.setattr(registry_module, "IN_MEMORY_GUARDRAIL_HANDLER", handler)
def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail:
return CustomGuardrail(
guardrail_name=guardrail["guardrail_name"],
event_hook=GuardrailEventHooks.pre_call,
default_on=False,
)
monkeypatch.setitem(registry_module.guardrail_initializer_registry, "lit5367_ok", _initializer)
prisma_client = MagicMock()
prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(
return_value=[
_db_guardrail_row("first", "lit5367_ok"),
_db_guardrail_row("broken", "litellm_tool_permission"),
_db_guardrail_row("last", "lit5367_ok"),
]
)
await ProxyConfig()._init_guardrails_in_db(prisma_client=prisma_client)
assert sorted(handler.IN_MEMORY_GUARDRAILS) == ["first", "last"]
assert handler.reconciled_with == [{"first", "broken", "last"}]

View file

@ -0,0 +1,33 @@
"""Tests for the opt-in flag that gates PTU flat-cost attribution."""
import pytest
from litellm.proxy.spend_tracking.ptu_feature_flag import (
PTU_COST_ATTRIBUTION_ENV_VAR,
is_ptu_cost_attribution_enabled,
)
def test_disabled_when_env_var_is_unset(monkeypatch):
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
assert is_ptu_cost_attribution_enabled() is False
@pytest.mark.parametrize("value", ["true", "True", "TRUE", " true "])
def test_enabled_for_the_values_the_house_helper_recognises(monkeypatch, value):
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, value)
assert is_ptu_cost_attribution_enabled() is True
@pytest.mark.parametrize("value", ["false", "False", "0", "1", "", "yes", "off", "maybe"])
def test_disabled_for_everything_else(monkeypatch, value):
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, value)
assert is_ptu_cost_attribution_enabled() is False
def test_reads_the_env_var_on_every_call(monkeypatch):
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
assert is_ptu_cost_attribution_enabled() is False
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
assert is_ptu_cost_attribution_enabled() is True

File diff suppressed because it is too large Load diff

View file

@ -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"},

View file

@ -11278,3 +11278,68 @@ 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
async def _run_scheduled_background_jobs():
from litellm.proxy.proxy_server import ProxyStartupEvent
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
return ps.scheduler
@pytest.mark.asyncio
async def test_ptu_rollup_job_registered_at_startup(monkeypatch):
"""The PTU rollup cron is registered once an operator opts in; 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.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import (
PTU_ROLLUP_JOB_ID,
)
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
scheduler = await _run_scheduled_background_jobs()
assert scheduler.get_job(PTU_ROLLUP_JOB_ID) is not None
@pytest.mark.asyncio
async def test_ptu_rollup_job_not_registered_without_opt_in(monkeypatch):
"""Without LITELLM_ENABLE_PTU_COST_ATTRIBUTION the rollup never runs, so no sentinel row
is ever written. This is the gate that keeps the whole feature inert by default."""
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import (
PTU_ROLLUP_JOB_ID,
)
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
scheduler = await _run_scheduled_background_jobs()
assert scheduler.get_job(PTU_ROLLUP_JOB_ID) is None
assert len(scheduler.get_jobs()) > 0

View file

@ -2928,3 +2928,143 @@ def test_update_mcp_semantic_filter_settings_requires_proxy_admin(monkeypatch):
assert "proxy admin" in resp.json()["detail"].lower()
finally:
app.dependency_overrides.pop(user_api_key_auth, None)
class TestPtuCostAttributionUISetting:
"""``enable_ptu_cost_attribution`` is derived from the environment on every GET.
It is deliberately not an allowlisted, persisted setting: the point of gating PTU
flat cost on an env var is that an admin cannot flip it at runtime from the UI.
"""
@staticmethod
def _mock_prisma(monkeypatch, stored=None):
from unittest.mock import AsyncMock, MagicMock
mock_prisma = MagicMock()
mock_record = None
if stored is not None:
mock_record = MagicMock()
mock_record.ui_settings = stored
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=mock_record)
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
return mock_prisma
def test_reported_false_when_the_env_var_is_unset(self, mock_auth, monkeypatch):
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
self._mock_prisma(monkeypatch)
response = client.get("/get/ui_settings")
assert response.status_code == 200
assert response.json()["values"]["enable_ptu_cost_attribution"] is False
def test_reported_true_once_the_env_var_is_set(self, mock_auth, monkeypatch):
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
self._mock_prisma(monkeypatch)
response = client.get("/get/ui_settings")
assert response.status_code == 200
assert response.json()["values"]["enable_ptu_cost_attribution"] is True
def test_a_persisted_true_cannot_forge_the_derived_value(self, mock_auth, monkeypatch):
"""A row written before the allowlist existed must not be able to turn the feature on."""
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
self._mock_prisma(monkeypatch, stored={"enable_ptu_cost_attribution": True})
response = client.get("/get/ui_settings")
assert response.status_code == 200
assert response.json()["values"]["enable_ptu_cost_attribution"] is False
def test_is_not_an_allowlisted_persisted_setting(self):
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
ALLOWED_UI_SETTINGS_FIELDS,
)
assert "enable_ptu_cost_attribution" not in ALLOWED_UI_SETTINGS_FIELDS
def test_the_body_get_returns_is_a_valid_patch_body(self, mock_auth, monkeypatch):
"""Read-modify-write is how a client edits one setting. GET injects the derived key,
so rejecting it on presence made GET's own output an invalid PATCH body: the caller
got a 400 and silently lost the edit it actually wanted."""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
mock_prisma = self._mock_prisma(monkeypatch)
try:
round_tripped = client.get("/get/ui_settings").json()["values"]
assert "enable_ptu_cost_attribution" in round_tripped
response = client.patch("/update/ui_settings", json=round_tripped)
finally:
app.dependency_overrides.clear()
assert response.status_code == 200
assert mock_prisma.db.litellm_uisettings.upsert.called
def test_a_co_submitted_setting_still_applies_alongside_the_derived_key(self, mock_auth, monkeypatch):
"""The derived key riding along must not discard the caller's real edit."""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
mock_prisma = self._mock_prisma(monkeypatch)
try:
response = client.patch(
"/update/ui_settings",
json={"enable_ptu_cost_attribution": False, "enable_chat_ui": True},
)
finally:
app.dependency_overrides.clear()
assert response.status_code == 200
upsert_data = mock_prisma.db.litellm_uisettings.upsert.call_args.kwargs["data"]
persisted = json.loads(upsert_data["create"]["ui_settings"])
assert persisted["enable_chat_ui"] is True
assert "enable_ptu_cost_attribution" not in persisted
def test_patch_rejects_the_derived_setting(self, mock_auth, monkeypatch):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
mock_prisma = self._mock_prisma(monkeypatch)
try:
response = client.patch(
"/update/ui_settings",
json={"enable_ptu_cost_attribution": True},
)
finally:
app.dependency_overrides.clear()
assert response.status_code == 400
assert "enable_ptu_cost_attribution" in str(response.json()["detail"])
assert not mock_prisma.db.litellm_uisettings.upsert.called

View file

@ -3,7 +3,10 @@ from typing import Any, Dict, List, Mapping, Tuple
import pytest
from litellm.repositories.unit_of_work import spend_reset_unit_of_work
from litellm.repositories.unit_of_work import (
budget_cascade_unit_of_work,
spend_reset_unit_of_work,
)
class FakeBatchTable:
@ -14,6 +17,9 @@ class FakeBatchTable:
def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> None:
self._calls.append((self._table_name, dict(where), dict(data)))
def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> None:
self._calls.append((f"{self._table_name}.update_many", dict(where), dict(data)))
class FakeBatch:
def __init__(self):
@ -22,6 +28,11 @@ class FakeBatch:
self.litellm_verificationtoken = FakeBatchTable("litellm_verificationtoken", self.calls)
self.litellm_usertable = FakeBatchTable("litellm_usertable", self.calls)
self.litellm_teamtable = FakeBatchTable("litellm_teamtable", self.calls)
self.litellm_budgettable = FakeBatchTable("litellm_budgettable", self.calls)
self.litellm_teammembership = FakeBatchTable("litellm_teammembership", self.calls)
self.litellm_organizationtable = FakeBatchTable("litellm_organizationtable", self.calls)
self.litellm_tagtable = FakeBatchTable("litellm_tagtable", self.calls)
self.litellm_endusertable = FakeBatchTable("litellm_endusertable", self.calls)
async def commit(self) -> None:
self.commit_count += 1
@ -64,3 +75,53 @@ async def test_empty_block_still_commits_the_batch():
assert batch.commit_count == 1
assert batch.calls == []
async def test_budget_cascade_dependents_and_window_advance_share_one_batch():
batch = FakeBatch()
reset_at = datetime(2026, 8, 3, 12, 0, tzinfo=timezone.utc)
linked = {"budget_id": {"in": ["budget-1"]}}
async with budget_cascade_unit_of_work(lambda: batch) as uow:
uow.team_memberships.queue_spend_zero(where=linked)
uow.keys.queue_spend_zero(where=linked)
uow.organizations.queue_spend_zero(where=linked)
uow.tags.queue_spend_zero(where=linked)
uow.endusers.queue_spend_zero(where={"user_id": {"in": ["enduser-1"]}})
uow.budgets.queue_window_advance(budget_id="budget-1", budget_reset_at=reset_at)
assert batch.commit_count == 0
assert batch.commit_count == 1
assert batch.calls == [
("litellm_teammembership.update_many", linked, {"spend": 0}),
("litellm_verificationtoken.update_many", linked, {"spend": 0}),
("litellm_organizationtable.update_many", linked, {"spend": 0}),
("litellm_tagtable.update_many", linked, {"spend": 0}),
("litellm_endusertable.update_many", {"user_id": {"in": ["enduser-1"]}}, {"spend": 0}),
("litellm_budgettable.update_many", {"budget_id": "budget-1"}, {"budget_reset_at": reset_at}),
]
async def test_budget_window_advance_tolerates_a_tier_deleted_mid_chunk():
"""A tier deleted between the read and the commit must not abort the batch:
``update`` raises P2025 on a missing row and takes every other write in the
chunk down with it, while ``update_many`` just matches nothing."""
batch = FakeBatch()
async with budget_cascade_unit_of_work(lambda: batch) as uow:
uow.budgets.queue_window_advance(budget_id="budget-1", budget_reset_at=datetime.now(timezone.utc))
assert [call[0] for call in batch.calls] == ["litellm_budgettable.update_many"]
async def test_budget_cascade_raising_inside_block_skips_commit():
"""A failure part-way through must leave budget_reset_at where it was, so
the tier is still due on the next tick."""
batch = FakeBatch()
with pytest.raises(RuntimeError, match="boom"):
async with budget_cascade_unit_of_work(lambda: batch) as uow:
uow.team_memberships.queue_spend_zero(where={"budget_id": {"in": ["budget-1"]}})
raise RuntimeError("boom")
assert batch.commit_count == 0

View file

@ -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

View 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"
)

View file

@ -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"]

View file

@ -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)

View file

@ -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)

View file

@ -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

View file

@ -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("")

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 23062
"limit": 23057
},
"LIT002": {
"limit": 27149
"limit": 27156
},
"LIT003": {
"limit": 269
@ -27,9 +27,9 @@
"limit": 0
},
"LIT010": {
"limit": 16735
"limit": 16744
},
"LIT011": {
"limit": 5597
"limit": 5596
}
}

View file

@ -0,0 +1,135 @@
import { getUiSettings } from "@/components/networking";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { renderHook, waitFor } from "@testing-library/react";
import React, { ReactNode } from "react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { PTU_FLAG_REFRESH_MS, usePtuCostAttributionEnabled } from "./usePtuCostAttributionEnabled";
import { useUISettings } from "./useUISettings";
vi.mock("@/components/networking", () => ({
getUiSettings: vi.fn(),
}));
describe("usePtuCostAttributionEnabled", () => {
let queryClient: QueryClient;
beforeEach(() => {
queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
vi.clearAllMocks();
});
const wrapper = ({ children }: { children: ReactNode }) =>
React.createElement(QueryClientProvider, { client: queryClient }, children);
/** Read the flag alongside the query it derives from, so assertions wait for a settled fetch. */
const renderSettledFlag = async (settings: unknown) => {
(getUiSettings as any).mockResolvedValue(settings);
const { result } = renderHook(() => ({ enabled: usePtuCostAttributionEnabled(), query: useUISettings() }), {
wrapper,
});
await waitFor(() => {
expect(result.current.query.isSuccess).toBe(true);
});
return result;
};
it("is true only when the proxy reports the flag as enabled", async () => {
const result = await renderSettledFlag({ values: { enable_ptu_cost_attribution: true } });
expect(result.current.enabled).toBe(true);
});
it("is false when the proxy reports the flag as disabled", async () => {
const result = await renderSettledFlag({ values: { enable_ptu_cost_attribution: false } });
expect(result.current.enabled).toBe(false);
});
it("is false when the proxy omits the flag entirely", async () => {
const result = await renderSettledFlag({ values: { enable_chat_ui: true } });
expect(result.current.enabled).toBe(false);
});
it("is false when the proxy returns no values at all", async () => {
const result = await renderSettledFlag({});
expect(result.current.enabled).toBe(false);
});
it("does not treat a truthy non-boolean as enabled", async () => {
const result = await renderSettledFlag({ values: { enable_ptu_cost_attribution: "false" } });
expect(result.current.enabled).toBe(false);
});
it("does not treat the string 'true' as enabled, since the proxy sends a real boolean", async () => {
const result = await renderSettledFlag({ values: { enable_ptu_cost_attribution: "true" } });
expect(result.current.enabled).toBe(false);
});
it("is false before the settings request resolves", () => {
(getUiSettings as any).mockReturnValue(new Promise(() => {}));
const { result } = renderHook(() => usePtuCostAttributionEnabled(), { wrapper });
expect(result.current).toBe(false);
});
it("is false when the settings request fails", async () => {
(getUiSettings as any).mockRejectedValue(new Error("boom"));
const { result } = renderHook(() => ({ enabled: usePtuCostAttributionEnabled(), query: useUISettings() }), {
wrapper,
});
await waitFor(() => {
expect(result.current.query.isError).toBe(true);
});
expect(result.current.enabled).toBe(false);
});
});
describe("staleness", () => {
let queryClient: QueryClient;
beforeEach(() => {
queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
vi.clearAllMocks();
});
const wrapper = ({ children }: { children: ReactNode }) =>
React.createElement(QueryClientProvider, { client: queryClient }, children);
it("polls the flag once it is on, so an already-open dashboard notices it going off", async () => {
(getUiSettings as any).mockResolvedValue({ values: { enable_ptu_cost_attribution: true } });
const { result } = renderHook(() => usePtuCostAttributionEnabled(), { wrapper });
await waitFor(() => {
expect(result.current).toBe(true);
});
const observers = queryClient.getQueryCache().getAll()[0].observers;
const polling = observers.filter((o: any) => o.options.refetchInterval === PTU_FLAG_REFRESH_MS);
expect(polling.length).toBeGreaterThan(0);
expect(polling[0].options.staleTime).toBe(PTU_FLAG_REFRESH_MS);
expect(PTU_FLAG_REFRESH_MS).toBeLessThan(60 * 60 * 1000);
});
it("does not poll while the flag is off, which is every deployment that never opted in", async () => {
// The hook cannot gate on the flag before reading it, so it starts on the shared
// one-hour cache and only escalates once it has seen the feature enabled. Polling
// unconditionally made a disabled deployment re-fetch settings 120x more often.
(getUiSettings as any).mockResolvedValue({ values: { enable_ptu_cost_attribution: false } });
const { result } = renderHook(() => usePtuCostAttributionEnabled(), { wrapper });
await waitFor(() => {
expect(result.current).toBe(false);
});
const observers = queryClient.getQueryCache().getAll()[0].observers;
expect(observers.every((o: any) => o.options.refetchInterval === undefined)).toBe(true);
expect(observers.every((o: any) => o.options.staleTime === 60 * 60 * 1000)).toBe(true);
});
it("leaves the default alone for every other settings consumer", async () => {
(getUiSettings as any).mockResolvedValue({ values: {} });
const { result } = renderHook(() => useUISettings(), { wrapper });
await waitFor(() => {
expect(result.current.isSuccess).toBe(true);
});
const observers = queryClient.getQueryCache().getAll()[0].observers;
expect(observers[0].options.staleTime).toBe(60 * 60 * 1000);
expect(observers[0].options.refetchInterval).toBeUndefined();
});
});

Some files were not shown because too many files have changed in this diff Show more