Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_fix_search_results_with_guardrails

# Conflicts:
#	tests/test_litellm/test_utils.py
This commit is contained in:
mateo-berri 2026-08-31 21:01:52 -07:00
commit fcbeb2e6a9
222 changed files with 15179 additions and 2528 deletions

View file

@ -4,17 +4,16 @@ description: >-
so only the first job on a given Cargo.lock compiles the bridge from scratch.
litellm builds through maturin, which compiles litellm-rust/crates/python-bridge
in release mode before it can produce a wheel. `uv sync` therefore pays a full
build in every job that installs the workspace: measured at 2m40s per unit shard
on 2026-08-21, more than the whole unit tier spends running tests. Nothing caught
it, because the uv cache holds wheels uv downloads rather than wheels it builds,
and a path dependency whose source moves every commit could never hit that cache
anyway. Cargo rebuilds only what changed when its target directory survives, so a
warm job pays for the bridge crate alone.
in the dev profile for editable installs. `uv sync` therefore pays a full build
in every job that installs the workspace. Nothing caught it, because the uv cache
holds wheels uv downloads rather than wheels it builds, and a path dependency
whose source moves every commit could never hit that cache anyway. Cargo rebuilds
only what changed when its target directory survives, so a warm job pays for the
bridge crate alone.
The key namespace is separate from test-rust.yml's. Both cache the same directory,
but that workflow fills it with debug and clippy artifacts, which a release build
cannot reuse, and a shared key would let whichever ran first deny the other a save.
The key namespace is separate from test-rust.yml's check and release caches. They
cache the same directory for different workloads, and a shared key would let
whichever ran first deny the others a save.
runs:
using: composite
@ -26,6 +25,6 @@ runs:
~/.cargo/registry
~/.cargo/git
litellm-rust/target
key: ${{ runner.os }}-cargo-release-${{ hashFiles('litellm-rust/Cargo.lock') }}
key: ${{ runner.os }}-maturin-dev-${{ hashFiles('litellm-rust/Cargo.lock') }}
restore-keys: |
${{ runner.os }}-cargo-release-
${{ runner.os }}-maturin-dev-

View file

@ -1,230 +0,0 @@
#!/usr/bin/env python3
"""
Detect and close duplicate GitHub issues using title similarity.
Modes:
--scan Compare all open issues against each other (batch)
--issue-number N Check a single issue against older open issues
Requires the `gh` CLI to be authenticated.
"""
import argparse
import difflib
import json
import re
import subprocess
import sys
def normalize_title(title: str) -> str:
"""Strip common prefixes, lowercase, and collapse whitespace."""
title = re.sub(
r"^\[?(bug|feature request|enhancement|question|docs)[:\]]?\s*",
"",
title,
flags=re.IGNORECASE,
)
return " ".join(title.lower().split())
def gh(*args: str) -> str:
"""Run a gh CLI command and return stdout."""
result = subprocess.run(
["gh", *args],
capture_output=True,
text=True,
check=True,
)
return result.stdout
def fetch_open_issues(repo: str | None) -> list[dict]:
"""Fetch all open issues (excluding PRs) via gh api --paginate."""
if repo:
endpoint = (
f"repos/{repo}/issues?state=open&per_page=100&sort=created&direction=asc"
)
else:
endpoint = "repos/{owner}/{repo}/issues?state=open&per_page=100&sort=created&direction=asc"
cmd = ["api", "--paginate", endpoint]
raw = gh(*cmd)
# gh --paginate concatenates JSON arrays, so we may get multiple arrays
issues = []
for line in raw.strip().splitlines():
line = line.strip()
if not line:
continue
parsed = json.loads(line)
if isinstance(parsed, list):
issues.extend(parsed)
else:
issues.append(parsed)
# Filter out pull requests (they also appear in the issues endpoint)
return [i for i in issues if "pull_request" not in i]
def close_as_duplicate(
issue_number: int, duplicate_of: int, repo: str | None, dry_run: bool
) -> None:
"""Close an issue as duplicate of another, adding a comment and label."""
repo_args = ["--repo", repo] if repo else []
if dry_run:
print(
f" [DRY RUN] Would close #{issue_number} as duplicate of #{duplicate_of}"
)
return
# Add comment
comment_body = (
f"Closing as duplicate of #{duplicate_of}.\n\n"
"If you believe this is not a duplicate, please reopen and add context "
"explaining how this differs."
)
gh("issue", "comment", str(issue_number), "--body", comment_body, *repo_args)
# Add label
gh("issue", "edit", str(issue_number), "--add-label", "duplicate", *repo_args)
# Close with not_planned reason
gh(
"api",
f"repos/{repo or '{owner}/{repo}'}/issues/{issue_number}",
"-X",
"PATCH",
"-f",
"state=closed",
"-f",
"state_reason=not_planned",
)
print(f" Closed #{issue_number} as duplicate of #{duplicate_of}")
def find_duplicate(
issue: dict, candidates: list[dict], threshold: float
) -> dict | None:
"""Return the first candidate whose normalized title is above threshold."""
norm = normalize_title(issue["title"])
for candidate in candidates:
if candidate["number"] == issue["number"]:
continue
cand_norm = normalize_title(candidate["title"])
ratio = difflib.SequenceMatcher(None, norm, cand_norm).ratio()
if ratio >= threshold:
return candidate
return None
def scan_all(
issues: list[dict], threshold: float, repo: str | None, dry_run: bool
) -> int:
"""Compare every issue against all older issues. Returns count of duplicates found."""
# Sort oldest first
issues.sort(key=lambda i: i["number"])
closed_count = 0
for idx, issue in enumerate(issues):
older = issues[:idx]
if not older:
continue
dup = find_duplicate(issue, older, threshold)
if dup:
ratio = difflib.SequenceMatcher(
None,
normalize_title(issue["title"]),
normalize_title(dup["title"]),
).ratio()
print(
f"#{issue['number']}: \"{issue['title']}\"\n"
f" -> duplicate of #{dup['number']}: \"{dup['title']}\" "
f"({ratio:.0%} similar)"
)
close_as_duplicate(issue["number"], dup["number"], repo, dry_run)
closed_count += 1
return closed_count
def check_single(
issue_number: int,
issues: list[dict],
threshold: float,
repo: str | None,
dry_run: bool,
) -> bool:
"""Check a single issue against all older open issues. Returns True if duplicate found."""
target = None
for i in issues:
if i["number"] == issue_number:
target = i
break
if target is None:
print(f"Issue #{issue_number} not found among open issues.")
return False
older = [i for i in issues if i["number"] < issue_number]
dup = find_duplicate(target, older, threshold)
if dup:
ratio = difflib.SequenceMatcher(
None,
normalize_title(target["title"]),
normalize_title(dup["title"]),
).ratio()
print(
f"#{target['number']}: \"{target['title']}\"\n"
f" -> duplicate of #{dup['number']}: \"{dup['title']}\" "
f"({ratio:.0%} similar)"
)
close_as_duplicate(issue_number, dup["number"], repo, dry_run)
return True
print(f"#{issue_number}: no duplicate found above threshold {threshold}")
return False
def main() -> None:
parser = argparse.ArgumentParser(
description="Detect and close duplicate GitHub issues"
)
mode = parser.add_mutually_exclusive_group(required=True)
mode.add_argument("--scan", action="store_true", help="Scan all open issues")
mode.add_argument("--issue-number", type=int, help="Check a single issue number")
parser.add_argument(
"--threshold", type=float, default=0.85, help="Similarity threshold (0-1)"
)
parser.add_argument(
"--close",
action="store_true",
help="Actually close duplicates (default is dry-run)",
)
parser.add_argument(
"--repo", type=str, help="Repository (owner/repo). Auto-detected if omitted."
)
args = parser.parse_args()
dry_run = not args.close
if dry_run:
print("=== DRY RUN MODE (pass --close to actually close issues) ===\n")
print("Fetching open issues...")
issues = fetch_open_issues(args.repo)
print(f"Found {len(issues)} open issues.\n")
if args.scan:
count = scan_all(issues, args.threshold, args.repo, dry_run)
print(f"\nTotal duplicates {'found' if dry_run else 'closed'}: {count}")
else:
found = check_single(
args.issue_number, issues, args.threshold, args.repo, dry_run
)
sys.exit(0 if found else 0) # Always exit 0; finding no dup is not an error
if __name__ == "__main__":
main()

View file

@ -0,0 +1,69 @@
name: Auto-close duplicate issues
on:
schedule:
- cron: "0 9 * * *"
workflow_dispatch:
inputs:
dry_run:
description: Log which issues would close without closing anything
type: boolean
default: true
grace_period_days:
description: Days a duplicate notice must go unanswered before the close
type: number
default: 3
pull_request:
paths:
- .github/workflows/auto-close-duplicates.yml
- scripts/auto-close-duplicates.ts
- scripts/auto-close-duplicates.test.ts
permissions: {}
jobs:
test:
if: github.event_name == 'pull_request'
runs-on: ubuntu-latest
timeout-minutes: 5
permissions:
contents: read
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Setup Bun
uses: oven-sh/setup-bun@0c5077e51419868618aeaa5fe8019c62421857d6 # v2.2.0
with:
bun-version: "1.4.0"
- name: Test the sweep
run: bun test scripts/auto-close-duplicates.test.ts
sweep:
if: github.event_name != 'pull_request' && github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
timeout-minutes: 10
permissions:
contents: read
issues: write
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Setup Bun
uses: oven-sh/setup-bun@0c5077e51419868618aeaa5fe8019c62421857d6 # v2.2.0
with:
# Exact version, never latest: the next step holds an issues: write token
bun-version: "1.4.0"
- name: Close unanswered duplicates, reopen ones the reporter answered
run: bun run scripts/auto-close-duplicates.ts
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
DRY_RUN: ${{ inputs.dry_run == true }}
GRACE_PERIOD_DAYS: ${{ inputs.grace_period_days }}

View file

@ -1,12 +1,19 @@
name: Check Duplicate Issues
# Flagging only. "Auto-close duplicate issues" closes a flagged issue 3 days later,
# and only when its title is identical to an older open issue and nobody replied.
# The HTML marker below is the handshake between the two, so keep it in the template.
on:
issues:
types: [opened, edited]
permissions: {}
jobs:
check-duplicate:
runs-on: ubuntu-latest
timeout-minutes: 5
permissions:
issues: write
contents: read
@ -19,35 +26,12 @@ jobs:
threshold: 0.6
reaction: eyes
comment: |
**⚠️ Potential duplicate detected**
<!-- litellm:potential-duplicate candidates={{#issues}}{{number}},{{/issues}} -->
**Potential duplicate detected**
This issue appears similar to existing issue(s):
This looks similar to:
{{#issues}}
- [#{{number}}]({{html_url}}) - {{title}} ({{accuracy}}% similar)
- #{{number}} - {{title}}
{{/issues}}
Please review the linked issue(s) to see if they address your concern. If this is not a duplicate, please provide additional context to help us understand the difference.
- name: Checkout close script
if: github.event.action == 'opened'
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
sparse-checkout: .github/scripts
persist-credentials: false
- name: Set up Python
if: github.event.action == 'opened'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Auto-close if high-confidence duplicate
if: github.event.action == 'opened'
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
python3 .github/scripts/close_duplicate_issues.py \
--issue-number ${{ github.event.issue.number }} \
--repo ${{ github.repository }} \
--threshold 0.85 \
--close
If this is a duplicate, add a thumbs-up reaction to the existing issue and follow along there. When the title is identical to an older open issue, this issue closes automatically in 3 days unless someone responds. If it is not a duplicate, comment here or add a thumbs-down reaction to this comment and it stays open.

77
.github/workflows/test-redis-compat.yml vendored Normal file
View file

@ -0,0 +1,77 @@
name: "Unit Tests: Redis Client Version Compatibility"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "litellm/_redis.py"
- "litellm/_redis_credential_provider.py"
- "tests/test_litellm/test_redis.py"
- "tests/test_litellm/caching/test_redis_connection_pool.py"
- ".github/workflows/test-redis-compat.yml"
- "pyproject.toml"
- "uv.lock"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
redis-compat:
name: "redis-py ${{ matrix.redis-version }}"
runs-on: ubuntu-latest
timeout-minutes: 15
strategy:
fail-fast: false
matrix:
# 5.3.1 is the version pinned in uv.lock (redisvl caps it below 6); the
# newer legs prove the inspect.signature introspection in litellm/_redis.py
# keeps extracting kwargs on the redis-py releases people actually run now.
# Only the exact release 6.0.0 is skipped: rq (pulled by the proxy extra)
# specifies `redis != 6`, which excludes 6.0.0 alone, so 6.4.0 stands in
# for the 6.x line.
redis-version: ["5.3.1", "6.4.0", "7.4.1", "8.0.1"]
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Pin redis-py to the matrix version
env:
REDIS_VERSION: ${{ matrix.redis-version }}
run: |
uv pip install "redis==${REDIS_VERSION:?}"
uv run --no-sync python -c "import redis; assert redis.__version__ == '${REDIS_VERSION:?}', redis.__version__; print('redis-py', redis.__version__)"
- name: Run redis unit tests
run: |
uv run --no-sync pytest \
tests/test_litellm/test_redis.py \
tests/test_litellm/caching/test_redis_connection_pool.py \
--tb=short -vv \
--reruns 2 \
--reruns-delay 1 \
--durations=20

View file

@ -103,6 +103,7 @@ jobs:
tests/test_litellm/completion_extras
tests/test_litellm/compression
tests/test_litellm/containers
tests/test_litellm/endpoints
tests/test_litellm/experimental_mcp_client
tests/test_litellm/models
tests/test_litellm/repositories

View file

@ -23,6 +23,8 @@ When adding new features, add meaningful tests. Don't add tests that don't check
Same thing for bug fixes. The tests should make it so that this specific bug can never happen again without failing tests (i.e., regression)
Never test structure of code only function of it
`tests/test_litellm/` mirrors `litellm/` in a parallel path (see `tests/test_litellm/readme.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `CLAUDE.md`

View file

@ -1,10 +1,10 @@
# syntax=docker/dockerfile:1.7
# Base image for building
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43
@ -40,8 +40,8 @@ COPY --from=uvbin /uvx /usr/local/bin/uvx
RUN apk add --no-cache \
bash \
gcc \
python3 \
python3-dev \
python-3.13 \
python-3.13-dev \
rust \
openssl \
openssl-dev \
@ -51,6 +51,7 @@ RUN apk add --no-cache \
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
UV_PYTHON_DOWNLOADS=0 \
PATH="/app/.venv/bin:${PATH}"
# Copy dependency metadata first for layer caching
@ -65,7 +66,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
--python python3.13
# Copy full source tree
COPY . .
@ -86,7 +87,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
npm_config_cache=/root/.npm \
@ -101,7 +102,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
# node (without npm) is required by the prisma CLI at runtime
RUN apk add --no-cache bash openssl tzdata nodejs python3 libsndfile
RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile
WORKDIR /app
ENV PATH="/app/.venv/bin:${PATH}" \

View file

@ -1,5 +1,5 @@
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
@ -16,7 +16,7 @@ COPY --from=uvbin /uv /uvx /usr/local/bin/
# instead of nodeenv downloading one whose dynamic deps may not be in Wolfi
# (e.g. Node 26.2.0 needs libatomic). Retry for transient apk.cgr.dev flakes.
RUN for i in 1 2 3; do \
apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile nodejs npm && break; \
apk add --no-cache bash gcc python-3.13 python-3.13-dev openssl openssl-dev libsndfile nodejs npm && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done
@ -46,7 +46,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--python python3
--python python3.13
# Stage 2 — copy source and install the project + workspace members.
COPY . .
@ -57,7 +57,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--python python3
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
npm_config_cache=/root/.npm \
@ -71,7 +71,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
RUN for i in 1 2 3; do \
apk add --no-cache bash openssl tzdata python3 libsndfile libatomic && break; \
apk add --no-cache bash openssl tzdata python-3.13 libsndfile libatomic && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done

View file

@ -3,7 +3,7 @@
"limit": 16171
},
"reportArgumentType": {
"limit": 2226
"limit": 2224
},
"reportAssignmentType": {
"limit": 319
@ -57,7 +57,7 @@
"limit": 5611
},
"reportMissingTypeArgument": {
"limit": 15350
"limit": 15348
},
"reportMissingTypeStubs": {
"limit": 40
@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44368
"limit": 44364
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38468
"limit": 38465
},
"reportUnknownParameterType": {
"limit": 19665
"limit": 19663
},
"reportUnknownVariableType": {
"limit": 30066
"limit": 30064
},
"reportUnnecessaryCast": {
"limit": 111

View file

@ -1,10 +1,10 @@
# syntax=docker/dockerfile:1.7
# Base image for building
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43
@ -39,8 +39,8 @@ COPY --from=uvbin /uvx /usr/local/bin/uvx
RUN apk add --no-cache \
bash \
gcc \
python3 \
python3-dev \
python-3.13 \
python-3.13-dev \
openssl \
openssl-dev \
nodejs \
@ -49,6 +49,7 @@ RUN apk add --no-cache \
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
UV_PYTHON_DOWNLOADS=0 \
PATH="/app/.venv/bin:${PATH}"
# Copy dependency metadata first for layer caching
@ -63,7 +64,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
--python python3.13
# Copy full source tree
COPY . .
@ -84,7 +85,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
npm_config_cache=/root/.npm \
@ -98,7 +99,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
# node (without npm) is required by the prisma CLI at runtime
RUN apk add --no-cache bash openssl tzdata nodejs python3 libsndfile
RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile
WORKDIR /app
ENV PATH="/app/.venv/bin:${PATH}" \

View file

@ -1,8 +1,8 @@
# syntax=docker/dockerfile:1.7
# Base images
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG PROXY_EXTRAS_SOURCE=published
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
@ -37,8 +37,8 @@ COPY --from=uvbin /uvx /usr/local/bin/uvx
RUN for i in 1 2 3; do \
apk add --no-cache \
python3 \
python3-dev \
python-3.13 \
python-3.13-dev \
gcc \
rust \
bash \
@ -52,6 +52,7 @@ RUN for i in 1 2 3; do \
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
UV_PYTHON_DOWNLOADS=0 \
PATH="/app/.venv/bin:${PATH}" \
LITELLM_NON_ROOT=true \
XDG_CACHE_HOME=/app/.cache
@ -69,7 +70,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
--python python3.13
# Copy full source tree
COPY . .
@ -96,7 +97,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3 \
--python python3.13 \
--no-sources-package litellm-proxy-extras; \
else \
uv sync --frozen --no-default-groups --no-editable \
@ -105,7 +106,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3; \
--python python3.13; \
fi
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
@ -124,7 +125,7 @@ RUN for i in 1 2 3; do \
apk upgrade --no-cache && break || sleep 5; \
done && \
for i in 1 2 3; do \
apk add --no-cache python3 bash openssl tzdata libsndfile nodejs && break || sleep 5; \
apk add --no-cache python-3.13 bash openssl tzdata libsndfile nodejs && break || sleep 5; \
done
# Copy only what runtime needs. The application is installed inside the venv;

View file

@ -1,5 +1,5 @@
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
@ -16,7 +16,7 @@ COPY --from=uvbin /uv /uvx /usr/local/bin/
# instead of nodeenv downloading one whose dynamic deps may not be in Wolfi
# (e.g. Node 26.2.0 needs libatomic). Retry for transient apk.cgr.dev flakes.
RUN for i in 1 2 3; do \
apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile nodejs npm && break; \
apk add --no-cache bash gcc python-3.13 python-3.13-dev openssl openssl-dev libsndfile nodejs npm && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done
@ -47,7 +47,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra extra_proxy \
--extra semantic-router \
--extra bedrock-realtime \
--python python3
--python python3.13
# Stage 2 — copy source and install the project + workspace members.
COPY . .
@ -59,7 +59,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra extra_proxy \
--extra semantic-router \
--extra bedrock-realtime \
--python python3
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
npm_config_cache=/root/.npm \
@ -73,7 +73,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
RUN for i in 1 2 3; do \
apk add --no-cache bash openssl tzdata python3 libsndfile libatomic && break; \
apk add --no-cache bash openssl tzdata python-3.13 libsndfile libatomic && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done

View file

@ -86,6 +86,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/comprehendmedical",
"/cohere/",
"/gemini/",
"/gigachat/",
"/google/",
"/vertex_ai/",
"/vertex-ai/",

View file

@ -0,0 +1,21 @@
DO $$
BEGIN
IF EXISTS (
SELECT 1 FROM information_schema.columns
WHERE table_name = 'LiteLLM_ShadowEvalJob' AND column_name = 'api_key_id'
) THEN
ALTER TABLE "LiteLLM_ShadowEvalJob" RENAME COLUMN "api_key_id" TO "target_id";
END IF;
END $$;
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN IF NOT EXISTS "target_type" TEXT NOT NULL DEFAULT 'key';
DROP INDEX IF EXISTS "LiteLLM_ShadowEvalJob_one_active_per_key_direction";
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_one_active_per_target_direction"
ON "LiteLLM_ShadowEvalJob"("target_type", "target_id", "direction") WHERE "stopped_at" IS NULL;
DROP INDEX IF EXISTS "LiteLLM_ShadowEvalJob_api_key_id_idx";
CREATE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_target_type_target_id_idx"
ON "LiteLLM_ShadowEvalJob"("target_type", "target_id");

View file

@ -1529,14 +1529,15 @@ model LiteLLM_AutoRouterSession {
model LiteLLM_ShadowEvalJob {
id String @id @default(cuid())
group_id String // legs of one job share this; the API's job id
api_key_id String // hashed virtual key whose traffic this leg shadows
target_type String @default("key") // key | team | user
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
router_name String // the auto-router under evaluation, in either direction
direction String @default("forward") // forward | reverse
baseline_model String? // reverse only: the fixed model the router is judged against
judge_model String
shadow_percentage Float
max_turns Int // sample-count ceiling: the whole budget on pre-max_budget jobs, the error-loop valve otherwise
max_budget Float? // per-key USD cap on the eval's own shadow + judge spend; null on jobs from before spend budgets
max_budget Float? // per-target USD cap on the eval's own shadow + judge spend; null on jobs from before spend budgets
created_at DateTime @default(now())
created_by String?
ends_at DateTime
@ -1544,7 +1545,7 @@ model LiteLLM_ShadowEvalJob {
stopped_by String? // operator who stopped it early; null when it ended on its own
@@index([group_id])
@@index([api_key_id])
@@index([target_type, target_id])
@@index([created_at])
}

View file

@ -512,6 +512,13 @@ class ProxyExtrasDBManager:
try:
import psycopg
except ImportError:
logger.warning(
"psycopg is not installed; skipping the LiteLLM_SpendLogs "
"partition check. If this table is partitioned (see "
"db_scripts/partition_spend_logs.sql), schema reconciliation "
"will try to rewrite its primary key and fail. Install the "
"litellm[extra_proxy] extra, which now includes psycopg."
)
return False
cleaned_url = ProxyExtrasDBManager._strip_prisma_query_params(database_url)

View file

@ -30,3 +30,12 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"]
tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
base64 = "0.22"
[profile.release]
opt-level = 3
lto = "thin"
codegen-units = 1
panic = "unwind"
debug = false
incremental = false
strip = "symbols"

View file

@ -10,7 +10,8 @@ name = "_native"
crate-type = ["cdylib"]
[features]
default = ["extension-module"]
default = ["abi3"]
abi3 = ["pyo3/abi3-py310"]
extension-module = ["pyo3/extension-module"]
[dependencies]

View file

@ -7,6 +7,9 @@ warnings.filterwarnings("ignore", message=".*conflict with protected namespace.*
# Suppress Pydantic 2.11+ deprecation warning about accessing model_fields on instances
# This warning can accumulate during streaming and cause memory leaks
warnings.filterwarnings("ignore", message=".*Accessing the.*attribute on the instance is deprecated.*")
# ReadOnly on TypedDict fields is repo-wide static discipline (LIT012); pydantic warns it
# cannot enforce it at runtime, which floods proxy boot once such a type is schema-walked
warnings.filterwarnings("ignore", message=".*`ReadOnly` qualifier.*")
### INIT VARIABLES #########################
import threading
import os

View file

@ -264,13 +264,17 @@ def _plain_log_format(stdout: TextIO | None, stderr: TextIO | None) -> str:
class LevelRoutingStreamHandler(logging.StreamHandler):
"""Writes records below WARNING to stdout and WARNING and above to stderr.
"""Writes records below WARNING and invalid-key warnings to stdout, others to stderr.
Collectors that derive severity from the stream report every stderr line as an error.
Invalid-key warnings route to stdout so LITELLM_LOG=ERROR can suppress them.
"""
def emit(self, record: logging.LogRecord) -> None:
preferred: Final = sys.stdout if record.levelno < logging.WARNING else sys.stderr
is_stdout_record: Final = record.levelno < logging.WARNING or (
record.levelno == logging.WARNING and record.name == verbose_proxy_stdout_logger.name
)
preferred: Final = sys.stdout if is_stdout_record else sys.stderr
if preferred is None or getattr(preferred, "closed", False):
self.stream = sys.stderr # rebind-ok: fall back to the pre-fix stream rather than raising per record
else:
@ -508,6 +512,9 @@ else:
handler.setFormatter(formatter)
verbose_proxy_logger = logging.getLogger("LiteLLM Proxy")
# Malformed virtual key rejections log through this child; LevelRoutingStreamHandler
# writes its WARNING records to stdout. It has no handler or level of its own.
verbose_proxy_stdout_logger: Final = verbose_proxy_logger.getChild("stdout")
verbose_router_logger = logging.getLogger("LiteLLM Router")
verbose_logger = logging.getLogger("LiteLLM")
@ -520,6 +527,7 @@ verbose_logger.addHandler(handler)
# handlers (JSON mode, uvicorn log config, a host app's root handler).
verbose_router_logger.addFilter(_stdout_truncation_filter)
verbose_proxy_logger.addFilter(_stdout_truncation_filter)
verbose_proxy_stdout_logger.addFilter(_stdout_truncation_filter)
verbose_logger.addFilter(_stdout_truncation_filter)
@ -683,6 +691,7 @@ def _turn_on_json():
- Adds a JSON formatter to all loggers
"""
handler: Final = LevelRoutingStreamHandler()
handler.setLevel(numeric_level)
handler.setFormatter(JsonFormatter())
_initialize_loggers_with_handler(handler)
# Set up exception handlers
@ -700,12 +709,14 @@ def _disable_debugging():
verbose_logger.disabled = True
verbose_router_logger.disabled = True
verbose_proxy_logger.disabled = True
verbose_proxy_stdout_logger.disabled = True
def _enable_debugging():
verbose_logger.disabled = False
verbose_router_logger.disabled = False
verbose_proxy_logger.disabled = False
verbose_proxy_stdout_logger.disabled = False
def print_verbose(print_statement):

View file

@ -13,6 +13,7 @@ import json
# s/o [@Frank Colson](https://www.linkedin.com/in/frank-colson-422b9b183/) for this redis implementation
import os
from collections.abc import Callable, Mapping
from types import MappingProxyType
from typing import Final
from urllib.parse import urlsplit, urlunsplit
@ -38,9 +39,25 @@ from ._logging import verbose_logger
AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default"
def _get_redis_kwargs():
arg_spec: Final = inspect.getfullargspec(redis.Redis)
def _unwrapped_init_args(cls: type) -> frozenset[str]:
"""Every parameter on a single class's own ``__init__``, decorator-unwrapped.
Unlike ``_init_arg_names`` below, this does not walk the MRO: ``redis.Redis``
and ``redis.RedisCluster`` (sync and async) each declare every real
constructor parameter directly on their own ``__init__``, so MRO-walking is
unnecessary — and it actively breaks the several tests here that mock the
class with ``patch(..., autospec=True)``, since ``inspect.getmro`` needs a
real ``__mro__`` that an autospec'd stand-in for a class does not provide.
Still unwraps first: redis-py >= 7.4 decorates these ``__init__``s with
``@deprecated_args`` too, which the same class of bug as ``_init_arg_names``
would otherwise silently empty this allowlist through (see its docstring).
"""
spec: Final = inspect.getfullargspec(inspect.unwrap(cls.__init__))
return frozenset(spec.args + spec.kwonlyargs)
def _get_redis_kwargs():
# Only allow primitive arguments
exclude_args: Final = {
"self",
@ -60,7 +77,7 @@ def _get_redis_kwargs():
"azure_client_secret",
}
available_args: Final = {x for x in arg_spec.args if x not in exclude_args} | include_args
available_args: Final = {x for x in _unwrapped_init_args(redis.Redis) if x not in exclude_args} | include_args
return available_args
@ -120,15 +137,23 @@ def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]:
return tuple(x for x in _init_arg_names(connection_cls) if x not in exclude_args) + include_args
def _get_redis_cluster_kwargs(client=None):
def _get_redis_cluster_kwargs(client: type | None = None):
"""Config kwargs the target cluster client's constructor actually accepts.
Defaults to the sync ``redis.RedisCluster``, but the async cluster client
(``redis.asyncio.cluster.RedisCluster``) declares connection settings such as
``decode_responses`` on its own constructor, where the sync class takes them
through ``**kwargs`` and so never names them in its signature. Introspecting
only the sync class regardless of which client is actually built silently
drops those for every async cluster caller.
"""
if client is None:
client = redis.Redis.from_url
arg_spec: Final = inspect.getfullargspec(redis.RedisCluster)
client = redis.RedisCluster
# Only allow primitive arguments
exclude_args: Final = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"}
available_args = {x for x in arg_spec.args if x not in exclude_args}
available_args = {x for x in _unwrapped_init_args(client) if x not in exclude_args}
available_args |= {
"password",
"username",
@ -161,6 +186,79 @@ def _get_redis_env_kwarg_mapping():
return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs() if x not in exclude_from_environment}
def _str_to_bool(value: str) -> bool:
return value.lower() in ("true", "1", "yes")
def _coerce_redis_kwargs_types(
redis_kwargs: Mapping[str, object],
client: type | tuple[type, ...] = redis.Redis,
) -> dict[str, object]: # mutable-ok: a caller mutates the returned kwargs before constructing its client
"""Coerces string values to the numeric/boolean type ``client``'s constructor
declares for that parameter. ``client`` may be a tuple of client classes; a
parameter's type is taken from the first signature that declares it, which
lets cluster callers coerce cluster-only kwargs such as
``cluster_error_retry_attempts`` alongside the shared connection kwargs.
Environment variables are always strings, and Helm ``--set`` stringifies values
too, so a config value like ``health_check_interval`` or ``socket_timeout``
can arrive as ``"30"``/``"5.5"`` rather than a real number. redis-py's own
connection-health-check arithmetic (``loop.time() + self.health_check_interval``)
then raises ``TypeError`` on every Redis operation instead of connecting.
``max_connections``, ``socket_timeout``, and ``socket_connect_timeout`` use an
explicit target type rather than the parameter's own signature default: redis-py
8.x changed the timeout defaults from ``None`` to int ``5``, so inferring the
type from the default would make a fractional ``"5.5"`` fail ``int()`` and get
silently dropped on 8.x while working on older versions. ``socket_keepalive``
is explicit too: its signature default is ``None``, which carries no type to
infer from, and leaving it a string makes ``"false"`` truthy.
"""
signatures: Final = tuple(inspect.signature(c) for c in (client if isinstance(client, tuple) else (client,)))
explicit_param_types: Final = MappingProxyType(
{
"max_connections": int,
"socket_timeout": float,
"socket_connect_timeout": float,
"socket_keepalive": bool,
}
)
result: Final = dict(redis_kwargs) # mutable-ok: per-key try/except coercion below needs to drop individual keys
for key, value in redis_kwargs.items():
if not isinstance(value, str):
continue
param = next((sig.parameters[key] for sig in signatures if key in sig.parameters), None)
if param is None:
continue
explicit_type = explicit_param_types.get(key)
if explicit_type is bool:
result[key] = _str_to_bool(value)
continue
if explicit_type is not None:
try:
result[key] = explicit_type(value)
except (ValueError, TypeError):
del result[key]
continue
default: object = param.default # pyright: ignore[reportAny] # inspect.Parameter.default is stubbed as Any
if default is inspect.Parameter.empty:
continue
# bool must be checked before int, since bool subclasses int
if isinstance(default, bool):
result[key] = _str_to_bool(value)
elif isinstance(default, int):
try:
result[key] = int(value)
except (ValueError, TypeError):
del result[key]
elif isinstance(default, float):
try:
result[key] = float(value)
except (ValueError, TypeError):
del result[key]
return result
def _redis_kwargs_from_environment():
mapping: Final = _get_redis_env_kwarg_mapping()
@ -505,7 +603,12 @@ def _get_redis_client_logic(**env_overrides):
raise ValueError("Either 'host' or 'url' must be specified for redis.")
# litellm.print_verbose(f"redis_kwargs: {redis_kwargs}")
return redis_kwargs
coercion_client: Final = (
(redis.Redis, redis.RedisCluster, async_redis.RedisCluster)
if redis_kwargs.get("startup_nodes")
else redis.Redis
)
return _coerce_redis_kwargs_types(redis_kwargs, client=coercion_client)
def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
@ -657,7 +760,9 @@ def get_redis_client(**env_overrides):
if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs:
return _init_redis_sentinel(redis_kwargs)
return redis.Redis(**redis_kwargs)
return redis.Redis( # pyright: ignore[reportCallIssue] # object-valued kwargs match no overload statically
**redis_kwargs, # pyright: ignore[reportArgumentType] # allow-listed and coerced against this signature
)
def get_redis_async_client(
@ -669,7 +774,7 @@ def get_redis_async_client(
if "startup_nodes" in redis_kwargs:
from redis.cluster import ClusterNode
args = _get_redis_cluster_kwargs()
args = _get_redis_cluster_kwargs(async_redis.RedisCluster)
cluster_kwargs: Final = {}
for arg in redis_kwargs:
if arg in args:

View file

@ -484,6 +484,22 @@ FIREWORKS_AI_80_B: Final = int(os.getenv("FIREWORKS_AI_80_B", 80))
#### Logging callback constants ####
REDACTED_BY_LITELM_STRING: Final = "REDACTED_BY_LITELM"
MAX_LANGFUSE_INITIALIZED_CLIENTS: Final = int(os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50))
# Backpressure + lifetime bounds for the /v1/messages streaming relay (see
# BaseAnthropicMessagesStreamingIterator.async_sse_wrapper). The relay queue is
# bounded so a slow client throttles the upstream pump instead of letting it
# buffer the whole response in memory; the detached-drain cap bounds how many
# post-disconnect drains may run concurrently so client behavior can't create
# unbounded worker state.
ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE: Final = int(
os.getenv("ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", "1024")
)
# Setting this to 0 disables detached draining entirely: every post-disconnect
# pump bills whatever partial output it has already collected and aborts the
# upstream stream immediately, instead of continuing to drain for the real
# terminal usage.
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS: Final = int(
os.getenv("ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS", "100")
)
LOGGING_WORKER_CONCURRENCY: Final = int(os.getenv("LOGGING_WORKER_CONCURRENCY", 100)) # Must be above 0
LOGGING_WORKER_MAX_QUEUE_SIZE: Final = int(os.getenv("LOGGING_WORKER_MAX_QUEUE_SIZE", 50_000))
LOGGING_WORKER_MAX_TIME_PER_COROUTINE: Final = float(os.getenv("LOGGING_WORKER_MAX_TIME_PER_COROUTINE", 20.0))
@ -806,6 +822,7 @@ openai_compatible_endpoints: Final[list] = [
"https://api.meta.ai/v1",
"https://api.cognition.ai/v1",
"https://api.scx.ai/v1",
"https://gigachat.devices.sberbank.ru/api/v1",
]
@ -1410,6 +1427,12 @@ DEFAULT_SOFT_BUDGET: Final = float(
) # by default all litellm proxy keys have a soft budget of 50.0
# makes it clear this is a rate limit error for a litellm virtual key
RATE_LIMIT_ERROR_MESSAGE_FOR_VIRTUAL_KEY: Final = "LiteLLM Virtual Key user_api_key_hash"
# Prefix of the 401 raised when a submitted virtual key is not shaped like one.
INVALID_VIRTUAL_KEY_ERROR_MESSAGE: Final = "LiteLLM Virtual Key expected"
# Attribute stamped on that 401 at its raise site so log routing recognises it by
# provenance. Message text is caller-influenceable on other 401s, so it must not
# be used to classify.
INVALID_VIRTUAL_KEY_ERROR_MARKER: Final = "_litellm_invalid_virtual_key_error"
# Python garbage collection threshold configuration
# Format: "gen0,gen1,gen2" e.g., "1000,50,50"

View file

@ -1,10 +1,14 @@
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, cast
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS
if TYPE_CHECKING:
from litellm import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import HttpxBinaryResponseContent
from litellm.types.llms.openai import ChatCompletionUserMessage, HttpxBinaryResponseContent
from litellm.types.utils import ModelResponse
@ -16,7 +20,42 @@ def _completion_response_cost(model_response: "ModelResponse") -> float | None:
return response_cost if isinstance(response_cost, float) else None
GEMINI_TTS_CHAT_AUDIO_FORMAT: Final = "pcm16"
class ChatAudioParam(TypedDict):
voice: ReadOnly[str]
format: ReadOnly[NotRequired[str]]
class SpeechToCompletionBridgeTransformationHandler:
def _chat_completion_params(self, optional_params: Mapping[str, object]) -> Mapping[str, object]:
return MappingProxyType(
{
param: value
for param, value in optional_params.items()
if param in OPENAI_CHAT_COMPLETION_PARAMS and param != "response_format"
}
)
def _chat_audio_format(self, model: str, optional_params: Mapping[str, object]) -> str | None:
if self._is_gemini_tts_model(model):
return GEMINI_TTS_CHAT_AUDIO_FORMAT
response_format: Final = optional_params.get("response_format")
return response_format if isinstance(response_format, str) else None
def _chat_audio_param(
self, model: str, voice: str | Mapping[str, object] | None, optional_params: Mapping[str, object]
) -> ChatAudioParam | None:
if not isinstance(voice, str):
return None
audio_format: Final = self._chat_audio_format(model, optional_params)
if audio_format is None:
voice_only: Final[ChatAudioParam] = {"voice": voice}
return voice_only
audio: Final[ChatAudioParam] = {"voice": voice, "format": audio_format}
return audio
def transform_request(
self,
model: str,
@ -28,36 +67,19 @@ class SpeechToCompletionBridgeTransformationHandler:
litellm_logging_obj: "LiteLLMLoggingObj",
custom_llm_provider: str,
) -> dict:
passed_optional_params: Final = {}
for op in optional_params:
if op in OPENAI_CHAT_COMPLETION_PARAMS:
passed_optional_params[op] = optional_params[op]
if voice is not None:
if isinstance(voice, str):
passed_optional_params["audio"] = {"voice": voice}
if "response_format" in optional_params:
passed_optional_params["audio"]["format"] = optional_params["response_format"]
return_kwargs = {
user_message: Final[ChatCompletionUserMessage] = {"role": "user", "content": input}
return_kwargs: Final = {
"model": model,
"messages": [
{
"role": "user",
"content": input,
}
],
"messages": [user_message],
"modalities": ["audio"],
**passed_optional_params,
**self._chat_completion_params(optional_params),
"audio": self._chat_audio_param(model, voice, optional_params),
**litellm_params,
"headers": headers,
"litellm_logging_obj": litellm_logging_obj,
"custom_llm_provider": custom_llm_provider,
}
# filter out None values
return_kwargs = {k: v for k, v in return_kwargs.items() if v is not None}
return return_kwargs
return {k: v for k, v in return_kwargs.items() if v is not None}
def _convert_pcm16_to_wav(self, pcm_data: bytes, sample_rate: int = 24000, channels: int = 1) -> bytes:
"""

View file

@ -108,7 +108,6 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
def __init__(self, completion_stream: object):
self.sent_first_chunk = False
# State tracking for accumulating partial tool calls
self.accumulated_tool_calls = dict[int, _ToolCallAccumulator]()
self._returned_response = False
super().__init__(completion_stream)

View file

@ -62,6 +62,8 @@ class GenAIMapper:
GenAI.RESPONSE_TIME_TO_FIRST_CHUNK: lambda d: d.time_to_first_chunk_seconds,
GenAI.USAGE_INPUT_TOKENS: lambda d: d.usage.input_tokens,
GenAI.USAGE_OUTPUT_TOKENS: lambda d: d.usage.output_tokens,
GenAI.USAGE_CACHE_CREATION_INPUT_TOKENS: lambda d: d.usage.cache_creation_input_tokens,
GenAI.USAGE_CACHE_READ_INPUT_TOKENS: lambda d: d.usage.cache_read_input_tokens,
Error.TYPE: lambda d: d.error.error_type if d.error else None,
Server.ADDRESS: lambda d: d.server.address if d.server else None,
Server.PORT: lambda d: d.server.port if d.server else None,

View file

@ -95,6 +95,22 @@ class LLMUsage:
input_tokens: int | None = None
output_tokens: int | None = None
total_tokens: int | None = None
cache_creation_input_tokens: int | None = None
cache_read_input_tokens: int | None = None
@classmethod
def from_standard_logging_payload(cls, payload: StandardLoggingPayload) -> LLMUsage:
# Cache token counts only exist on the raw provider usage object under metadata
metadata: Final[Mapping[str, object]] = payload.get("metadata") or {}
raw_usage: Final = metadata.get("usage_object")
usage_object: Final[Mapping[str, object]] = raw_usage if isinstance(raw_usage, Mapping) else {}
return cls(
input_tokens=as_int(payload.get("prompt_tokens")),
output_tokens=as_int(payload.get("completion_tokens")),
total_tokens=as_int(payload.get("total_tokens")),
cache_creation_input_tokens=as_int(usage_object.get("cache_creation_input_tokens")),
cache_read_input_tokens=as_int(usage_object.get("cache_read_input_tokens")),
)
@dataclass(frozen=True)
@ -363,11 +379,7 @@ class LLMCallSpanData:
response_model=context.response_model,
response_id=as_str(response.get("id")),
request_params=LLMRequestParams.from_model_parameters(params),
usage=LLMUsage(
input_tokens=as_int(payload.get("prompt_tokens")),
output_tokens=as_int(payload.get("completion_tokens")),
total_tokens=as_int(payload.get("total_tokens")),
),
usage=LLMUsage.from_standard_logging_payload(payload),
finish_reasons=finish_reasons,
error=_parse_error(payload),
response_cost=as_float(payload.get("response_cost")),

View file

@ -110,6 +110,8 @@ class GenAI:
# usage
USAGE_INPUT_TOKENS: Final = "gen_ai.usage.input_tokens"
USAGE_OUTPUT_TOKENS: Final = "gen_ai.usage.output_tokens"
USAGE_CACHE_CREATION_INPUT_TOKENS: Final = "gen_ai.usage.cache_creation.input_tokens"
USAGE_CACHE_READ_INPUT_TOKENS: Final = "gen_ai.usage.cache_read.input_tokens"
# content (opt-in, gated by capture mode)
INPUT_MESSAGES: Final = "gen_ai.input.messages"
OUTPUT_MESSAGES: Final = "gen_ai.output.messages"

View file

@ -592,7 +592,12 @@ _JOBS_CACHE_KEY: Final = "shadow_eval:active_jobs"
class ShadowEvalLogger(CustomLogger):
"""Fires blind pairwise shadow evaluations for keys with an active shadow-eval job."""
"""Fires blind pairwise shadow evaluations for targets with an active shadow-eval job.
A job targets a virtual key, a team, or a user; a request qualifies for a job when
any of its resolved identities (key hash, team id, user id) matches the job's
target, so team and user jobs cover JWT-authenticated traffic, which carries no
key hash at all."""
def __init__(
self,
@ -617,10 +622,10 @@ class ShadowEvalLogger(CustomLogger):
# generation; the refill absorbs written rows and resets.
self._job_starts: dict[str, int] = {} # mutable-ok: per-generation counter
async def _active_jobs(self) -> Mapping[str, tuple[ActiveShadowEvalJob, ...]]:
"""Active jobs by api_key_id, cache-first. A key holds at most one job per
direction, so the value is a collection. A DB fault returns empty without
caching, so sampling pauses for that request and the next one retries."""
async def _active_jobs(self) -> Mapping[tuple[str, str], tuple[ActiveShadowEvalJob, ...]]:
"""Active jobs by (target_type, target_id), cache-first. A target holds at most
one job per direction, so the value is a collection. A DB fault returns empty
without caching, so sampling pauses for that request and the next one retries."""
cached: Final = await self._jobs_cache.async_get_cache(_JOBS_CACHE_KEY)
if cached is not None:
return cached # pyright: ignore[reportReturnType] # cache stores exactly this mapping shape
@ -652,10 +657,10 @@ class ShadowEvalLogger(CustomLogger):
)
for row in grouped or []
}
by_key: Final = tuple(
by_target: Final = tuple(
sorted(
(
(str(record.api_key_id), job)
((str(record.target_type), str(record.target_id)), job)
for record in records or []
if (job := _as_active_job(record, *attempt_stats.get(str(record.id), (0, 0.0)))) is not None
),
@ -663,7 +668,7 @@ class ShadowEvalLogger(CustomLogger):
)
)
jobs: Final = MappingProxyType(
{key: tuple(job for _, job in group) for key, group in groupby(by_key, key=itemgetter(0))}
{target: tuple(job for _, job in group) for target, group in groupby(by_target, key=itemgetter(0))}
)
await self._jobs_cache.async_set_cache(_JOBS_CACHE_KEY, jobs)
self._job_starts = {} # rebind-ok: new generation, counts absorbed into the fill
@ -720,8 +725,18 @@ class ShadowEvalLogger(CustomLogger):
if should_redact_message_logging(dict(kwargs)): # mutable-ok: predicate takes a plain dict
return
metadata: Final = payload.get("metadata") or _EMPTY_METADATA
api_key_hash: Final = metadata.get("user_api_key_hash")
if not api_key_hash:
# Each identity the request resolved to is a candidate target; JWT-auth
# requests carry no key hash but do carry a team and user.
targets: Final = tuple(
(target_type, str(value))
for target_type, value in (
("key", metadata.get("user_api_key_hash")),
("team", metadata.get("user_api_key_team_id")),
("user", metadata.get("user_api_key_user_id")),
)
if value
)
if not targets:
return
request_id: Final = payload.get("id") or ""
if not request_id:
@ -731,8 +746,11 @@ class ShadowEvalLogger(CustomLogger):
return # only surfaces this table can normalize are comparable; unknown types fail closed
if ops.wire_params and _request_mutating_guardrail_ran(request_metadata):
return # the wire-body snapshot predates the rewrite; replaying it would resurrect stripped content
active_jobs: Final = await self._active_jobs()
eligible: Final = self._sampled_jobs(
(await self._active_jobs()).get(str(api_key_hash), ()), request_metadata, request_id
tuple(job for target in targets for job in active_jobs.get(target, ())),
request_metadata,
request_id,
)
if not eligible:
return
@ -1056,7 +1074,7 @@ class ShadowEvalLogger(CustomLogger):
)
_EMPTY_JOBS: Final[Mapping[str, tuple[ActiveShadowEvalJob, ...]]] = MappingProxyType({})
_EMPTY_JOBS: Final[Mapping[tuple[str, str], tuple[ActiveShadowEvalJob, ...]]] = MappingProxyType({})
def _default_prisma_provider() -> "PrismaClient | None":

View file

@ -50,6 +50,9 @@ OPTIONAL_KWARGS_KEYS: Final = (
"vertex_ai_project",
"vertex_ai_location",
"vertex_ai_credentials",
"gigachat_scope",
"gigachat_auth_url",
"gigachat_access_token",
"tpm",
"rpm",
"itpm",

View file

@ -369,6 +369,9 @@ def get_llm_provider(
elif endpoint == "https://api.meta.ai/v1":
custom_llm_provider = "meta"
dynamic_api_key = get_secret_str("META_API_KEY")
elif endpoint == "https://gigachat.devices.sberbank.ru/api/v1":
custom_llm_provider = "gigachat"
dynamic_api_key = get_secret_str("GIGACHAT_API_KEY")
elif (json_provider := JSONProviderRegistry.get_by_base_url(endpoint)) is not None:
custom_llm_provider = json_provider.slug
dynamic_api_key = api_key if api_key is not None else get_secret_str(json_provider.api_key_env)
@ -867,6 +870,9 @@ def _get_openai_compatible_provider_info(
# Manus is OpenAI compatible for responses API
api_base = api_base or get_secret_str("MANUS_API_BASE") or "https://api.manus.im"
dynamic_api_key = api_key or get_secret_str("MANUS_API_KEY")
elif custom_llm_provider == "gigachat":
api_base = api_base or get_secret_str("GIGACHAT_API_BASE") or "https://gigachat.devices.sberbank.ru/api/v1"
dynamic_api_key = api_key or get_secret_str("GIGACHAT_API_KEY")
if api_base is not None and not isinstance(api_base, str):
raise Exception(f"api base needs to be a string. api_base={api_base}")

View file

@ -2141,6 +2141,9 @@ class Logging(LiteLLMLoggingBaseClass):
logging_result: Final = self.normalize_logging_result(result=result)
if isinstance(result, Response) and isinstance(logging_result, (ModelResponse, EmbeddingResponse)):
result = logging_result
if standard_logging_object is None and result is not None and self.stream is not True:
if self._is_recognized_call_type_for_logging(logging_result=logging_result) or isinstance(
logging_result, (dict, list)
@ -6152,7 +6155,10 @@ def get_standard_logging_object_payload(
def emit_standard_logging_payload(payload: StandardLoggingPayload):
if os.getenv("LITELLM_PRINT_STANDARD_LOGGING_PAYLOAD"):
print(json.dumps(payload, indent=4), flush=True) # noqa: T201
try:
print(json.dumps(payload, indent=4, default=str), flush=True) # noqa: T201
except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging
verbose_logger.exception("Error serializing standard logging payload for debug output: %s", e)
def get_standard_logging_metadata(

View file

@ -144,7 +144,6 @@ def _is_choice_non_empty(choice: StreamingChoices) -> bool:
# Check model_extra for dynamically added fields on the choice
choice_extra_fields: Final[Mapping[str, object]] = choice.model_extra or {}
for extra_field_name, extra_field_value in choice_extra_fields.items():
# Skip certain structural fields that are just default/None placeholders
if extra_field_name == "index" and extra_field_value == 0:
continue
if extra_field_name in {"finish_reason", "logprobs"} and extra_field_value is None:
@ -192,7 +191,6 @@ def _is_delta_non_empty(delta: Delta) -> bool:
# Check model_extra for dynamically added fields (this is where Pydantic stores them)
delta_extra_fields: Final[Mapping[str, object]] = delta.model_extra or {}
for extra_field_value in delta_extra_fields.values():
# Even structural fields are meaningful if they have actual content
if _has_meaningful_content(extra_field_value):
return True

View file

@ -205,6 +205,41 @@ def is_non_content_values_set(message: AllMessageValues) -> bool:
return any(message.get(key, None) is not None for key in message if key not in ignore_keys)
_IMAGE_CONTENT_PART_TYPES: Final = frozenset({"image_url", "input_image", "image"})
_IMAGE_SCAN_MAX_DEPTH: Final = 4
def _content_parts_contain_image(parts: Sequence[object]) -> bool:
"""Depth-bounded frontier walk over nested content lists, iterative because the repo bans
recursion; an Anthropic tool_result nests its image parts exactly one level down."""
frontier = parts # rebind-ok: depth-bounded frontier walk
for _ in range(_IMAGE_SCAN_MAX_DEPTH):
if any(isinstance(part, Mapping) and part.get("type") in _IMAGE_CONTENT_PART_TYPES for part in frontier):
return True
frontier = tuple( # rebind-ok: depth-bounded frontier walk
nested
for part in frontier
if isinstance(part, Mapping)
for content in (part.get("content"),)
if isinstance(content, list)
for nested in content
)
if not frontier:
return False
return False
def request_contains_image_content(messages: Sequence[Mapping[str, object]]) -> bool:
"""Whether any message carries an image content part, across the dialects that reach
pre-routing hooks untranslated: chat-completions ``image_url``, Responses ``input_image``,
and Anthropic Messages ``image``, including images nested inside ``tool_result`` blocks."""
return any(
isinstance(content, list) and _content_parts_contain_image(content)
for message in messages
for content in (message.get("content"),)
)
def _audio_or_image_in_message_content(message: AllMessageValues) -> bool:
"""
Checks if message content contains an image or audio

View file

@ -693,6 +693,26 @@ def _count_document_tokens(
)
def _count_file_tokens(
file_value: object,
count_function: TokenCounterFunction,
use_default_image_token_count: bool,
) -> int:
"""An OpenAI `file` block is the chat-completions spelling of a document, so it prices like one."""
if not isinstance(file_value, Mapping):
return 0
filename: Final = file_value.get("filename")
file_data: Final = file_value.get("file_data")
name_tokens: Final = count_function(filename) if isinstance(filename, str) and filename else 0
if not isinstance(file_data, str) or not file_data:
return name_tokens
return name_tokens + calculate_img_tokens(
data=file_data,
mode="auto",
use_default_image_token_count=use_default_image_token_count,
)
def _count_anthropic_content(
content: Mapping[str, Any],
count_function: TokenCounterFunction,
@ -778,6 +798,12 @@ def _count_content_list(
use_default_image_token_count,
default_token_count,
)
elif c["type"] == "file":
num_tokens += _count_file_tokens(
c.get("file"),
count_function,
use_default_image_token_count,
)
elif c["type"] in ("tool_use", "tool_result"):
num_tokens += _count_anthropic_content(
c,
@ -807,7 +833,7 @@ def _count_content_list(
raise ValueError(
f"Invalid content item type: {content_type}. "
f"Expected str or dict with 'type' field "
f"(text, image_url, image, document, tool_use, tool_result, thinking, tool_reference)."
f"(text, image_url, image, document, file, tool_use, tool_result, thinking, tool_reference)."
)
return num_tokens
except Exception as e:

View file

@ -859,12 +859,28 @@ class AnthropicMessagesHandler(BaseTranslation):
@staticmethod
def _image_sources(block: Mapping[str, object]) -> tuple[str, ...]:
"""Normalize an Anthropic image block into strings a guardrail can read.
base64 becomes a data URI so the format travels with the payload, which is what
the OpenAI path already puts in this field. A file source yields nothing: those
bytes live behind the Files API and this extractor has no client to fetch them.
"""
source: Final = block.get("source")
if not isinstance(source, Mapping):
return ()
# Could be base64 or url
source_type: Final = source.get("type")
if source_type == "url":
url: Final = source.get("url")
return (url,) if isinstance(url, str) and url else ()
data: Final = source.get("data")
return (data,) if data else ()
if not isinstance(data, str) or not data:
return ()
media_type: Final = source.get("media_type")
if isinstance(media_type, str) and media_type:
return (f"data:{media_type};base64,{data}",)
return (data,)
async def _apply_guardrail_responses_to_input(
self,

View file

@ -8,6 +8,10 @@ import httpx
from pydantic import TypeAdapter
from typing_extensions import TypedDict
from litellm.constants import (
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS,
ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE,
)
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
@ -21,6 +25,9 @@ from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging()
_UPSTREAM_PUMP_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: stdlib strong-ref set for pump tasks
_DETACHED_STREAM_DRAINS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: bounded strong-ref set, detached drains
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
"Provider stream ended before emitting a message_stop event; "
"the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated."
@ -133,6 +140,34 @@ def _is_terminal_stream_chunk(chunk: object) -> bool:
return _is_message_stop_chunk(chunk) or _is_provider_error_chunk(chunk)
def _try_claim_detached_drain_slot() -> bool:
"""Claim a detached-drain slot for the current task, bounding concurrency.
Returns True if a slot was claimed (the caller may keep draining upstream
for billing) or False if the cap is already reached (the caller should stop
and bill what it has). Only touched from the event loop, so the check +
insert need no lock.
"""
if len(_DETACHED_STREAM_DRAINS) >= ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS:
return False
current_task: Final = asyncio.current_task()
if current_task is not None:
_DETACHED_STREAM_DRAINS.add(current_task)
current_task.add_done_callback(_DETACHED_STREAM_DRAINS.discard)
return True
def _exception_left_unconsumed(queue: "asyncio.Queue[bytes | None | BaseException]", exc: BaseException) -> bool:
"""After client detach the relay never reads the queue again, so drain it here.
The forwarded exception still sitting in the queue means the relay tore
down before re-raising it, so the proxy's failure handling never ran and
the caller must salvage spend itself.
"""
remaining: Final = tuple(queue.get_nowait() for _ in range(queue.qsize()))
return any(item is exc for item in remaining)
def _sse_event(event_type: str, payload: Mapping[str, object]) -> bytes:
return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode()
@ -414,17 +449,167 @@ class BaseAnthropicMessagesStreamingIterator:
async def async_sse_wrapper(
self,
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | dict],
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
) -> AsyncIterator[bytes]:
"""
Generic async SSE wrapper that converts streaming chunks to SSE format
and handles logging.
The upstream read runs in a detached background task (``_pump_upstream``)
so that a client disconnect tears down only this client-facing generator,
never the upstream drain + billing. The provider (e.g. Bedrock) keeps
generating and billing the full response regardless of the client, so
draining it to completion is what lets spend tracking see the real
terminal ``message_delta`` / ``message_stop`` usage instead of a
truncated placeholder count.
Chunks reach the client through a bounded queue. While the client is
connected the pump blocks on a full queue (racing the disconnect
signal), so a slow reader throttles the upstream read exactly as the old
direct ``yield`` did instead of letting the whole response buffer in
memory. Once the client goes away the pump stops enqueueing and only
keeps a single ``collected_chunks`` copy for billing, and the number of
such post-disconnect drains running at once is capped so client behavior
can't create unbounded worker state; over the cap the pump bills what it
has rather than draining further. Detached-drain lifetime is otherwise
bounded by the upstream stream/read timeout.
An upstream failure (Bedrock read / decode / chunk-conversion error)
that happens while the client is still connected is forwarded through
the queue and re-raised here, so the original provider exception (and
its status) reaches the proxy's failure handling unchanged rather than
being masked by a generic incomplete-stream event.
This method provides the common logic for both Anthropic and Bedrock implementations.
"""
collected_chunks: Final = []
saw_terminal_event = False
queue: Final[asyncio.Queue[bytes | None | BaseException]] = asyncio.Queue(
maxsize=ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE
)
client_detached: Final = asyncio.Event()
pump_task: Final = asyncio.create_task(self._pump_upstream_to_queue(completion_stream, queue, client_detached))
_UPSTREAM_PUMP_TASKS.add(pump_task)
pump_task.add_done_callback(_UPSTREAM_PUMP_TASKS.discard)
reached_end = False # rebind-ok: flipped once the relay consumes the end-of-stream sentinel
try:
while True:
item = await queue.get()
if item is None:
reached_end = True
break
if isinstance(item, BaseException):
raise item
yield item
finally:
client_detached.set()
if not reached_end:
self._dispatch_pending_deferred_logging()
def _dispatch_pending_deferred_logging(self) -> None:
"""Fire deferred billing that a torn-down response would otherwise drop.
When the pump finishes draining while the client is still connected it
stores the logging coroutine for ProxyLogging._fire_deferred_stream_logging,
which the proxy only fires on a normally completed response: a client
disconnect (GeneratorExit / CancelledError) re-raises past it. Without
this dispatch that window loses the spend row entirely.
"""
deferred_cb: Final = getattr(self.litellm_logging_obj, "_on_deferred_stream_complete", None)
deferred_args: Final = getattr(self.litellm_logging_obj, "_deferred_stream_complete_args", None)
if deferred_cb is None or deferred_args is None:
return
self.litellm_logging_obj._on_deferred_stream_complete = None
self.litellm_logging_obj._deferred_stream_complete_args = None
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=deferred_cb(*deferred_args))
async def _bill_collected_chunks(
self,
collected_chunks: list[bytes], # mutable-ok: SSE buffer forwarded to list-typed _handle_streaming_logging
*,
stream_teardown: bool,
) -> None:
from litellm._logging import verbose_proxy_logger
try:
await self._handle_streaming_logging(collected_chunks, stream_teardown=stream_teardown)
except Exception as exc: # noqa: BLE001 # billing is best-effort; never crash the pump
verbose_proxy_logger.warning(
"async_sse_wrapper billing failed after %d chunks: %s(%s)",
len(collected_chunks),
type(exc).__name__,
exc,
)
@staticmethod
async def _abort_upstream(
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
) -> None:
"""Close the upstream provider stream so it stops generating and billing."""
from litellm._logging import verbose_proxy_logger
try:
await aclose_if_supported(completion_stream)
except Exception as exc: # noqa: BLE001 # abort is best-effort; log and continue
verbose_proxy_logger.warning(
"async_sse_wrapper failed to abort upstream stream: %s(%s)",
type(exc).__name__,
exc,
)
@staticmethod
async def _enqueue_for_client(
queue: "asyncio.Queue[bytes | None | BaseException]",
client_detached: "asyncio.Event",
item: bytes | None | BaseException,
) -> bool:
"""Deliver one item to the client, applying backpressure.
Returns True if the item was queued, False if the client disconnected
before there was room (the item is then dropped, since a gone client
can't receive it). Never blocks once the client has detached.
"""
if client_detached.is_set():
return False
try:
queue.put_nowait(item)
except asyncio.QueueFull:
pass
else:
return True
put_task: Final = asyncio.ensure_future(queue.put(item))
detached_task: Final = asyncio.ensure_future(client_detached.wait())
try:
await asyncio.wait(frozenset((put_task, detached_task)), return_when=asyncio.FIRST_COMPLETED)
finally:
if not detached_task.done():
detached_task.cancel()
if put_task.done() and not put_task.cancelled():
return True
put_task.cancel()
return False
async def _pump_upstream_to_queue(
self,
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
queue: "asyncio.Queue[bytes | None | BaseException]",
client_detached: "asyncio.Event",
) -> None:
"""Drain the whole upstream into ``queue`` (backpressured) and bill once.
Runs detached so a client disconnect can't interrupt the upstream read;
see ``async_sse_wrapper`` for the full rationale. On a completed drain
the success billing (or deferred park) happens before the end-of-stream
sentinel is enqueued: the relay can only tear down after consuming the
sentinel, so its teardown can never outrun the park and get mistaken
for a client disconnect, and a sentinel the client never consumes falls
back to dispatching the parked billing here.
"""
from litellm._logging import verbose_proxy_logger
collected_chunks: Final[list[bytes]] = [] # mutable-ok: SSE billing buffer appended to across the drain
saw_terminal_event = False # rebind-ok: accumulates across the upstream loop
draining_detached = False # rebind-ok: set once this pump claims a detached-drain slot
try:
async for chunk in completion_stream:
if self.completion_start_time is None:
@ -432,17 +617,62 @@ class BaseAnthropicMessagesStreamingIterator:
saw_terminal_event = saw_terminal_event or _is_terminal_stream_chunk(chunk)
encoded_chunk = self._convert_chunk_to_sse_format(chunk)
collected_chunks.append(encoded_chunk)
yield encoded_chunk
except (GeneratorExit, asyncio.CancelledError):
# A client disconnect tears the generator down at the yield, so the
# post-loop logging below never runs and the tokens already streamed
# (and billed by the provider) would never reach spend tracking. See LIT-5839.
if collected_chunks:
await self._handle_streaming_logging(collected_chunks, stream_teardown=True)
raise
if not client_detached.is_set():
await self._enqueue_for_client(queue, client_detached, encoded_chunk)
continue
if not draining_detached:
if not _try_claim_detached_drain_slot():
verbose_proxy_logger.warning(
"async_sse_wrapper: detached-drain cap (%d) reached; billing %d partial "
"chunks and aborting the upstream stream to stop provider billing",
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS,
len(collected_chunks),
)
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
await self._abort_upstream(completion_stream)
return
draining_detached = True
except Exception as exc: # noqa: BLE001 # upstream errors are handled/forwarded by _handle_pump_upstream_error
await self._handle_pump_upstream_error(queue, client_detached, collected_chunks, exc)
return
if not saw_terminal_event:
yield _incomplete_stream_error_sse_event()
if client_detached.is_set():
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
return
if not saw_terminal_event and not await self._enqueue_for_client(
queue, client_detached, _incomplete_stream_error_sse_event()
):
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
return
await self._bill_collected_chunks(collected_chunks, stream_teardown=False)
if not await self._enqueue_for_client(queue, client_detached, None):
self._dispatch_pending_deferred_logging()
# Handle logging after all chunks are processed
await self._handle_streaming_logging(collected_chunks)
async def _handle_pump_upstream_error(
self,
queue: "asyncio.Queue[bytes | None | BaseException]",
client_detached: "asyncio.Event",
collected_chunks: list[bytes], # mutable-ok: SSE buffer forwarded to list-typed _bill_collected_chunks
exc: BaseException,
) -> None:
"""Forward a provider error to a still-connected client, else salvage partial spend.
Handing the original exception to the client-facing generator lets it
re-raise so the proxy's failure handling keeps the provider status and
owns logging (no success-bill). If the client already went away, or
disconnects before ever consuming the queued exception, no failure hook
runs, so bill the partial instead of dropping the request.
"""
from litellm._logging import verbose_proxy_logger
if not client_detached.is_set() and await self._enqueue_for_client(queue, client_detached, exc):
await client_detached.wait()
if not _exception_left_unconsumed(queue, exc):
return
verbose_proxy_logger.warning(
"async_sse_wrapper upstream pump failed after client disconnect (%d chunks): %s(%s)",
len(collected_chunks),
type(exc).__name__,
exc,
)
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)

View file

@ -1631,6 +1631,11 @@ class AmazonConverseConfig(BaseConfig):
bedrock_tool_config["toolChoice"] = tool_choice_values
self._drop_tool_choice_type_conflicting_with_tool_config(additional_request_params)
config_block_entries: Final = tuple(
(config_name, config_class, inference_params.pop(config_name, None))
for config_name, config_class in self.get_config_blocks().items()
)
data: Final[CommonRequestObject] = {
"inferenceConfig": self._transform_inference_params(inference_params=inference_params),
}
@ -1641,9 +1646,7 @@ class AmazonConverseConfig(BaseConfig):
if system_content_blocks:
data["system"] = system_content_blocks
# Handle all config blocks
for config_name, config_class in self.get_config_blocks().items():
config_value = inference_params.pop(config_name, None)
for config_name, config_class, config_value in config_block_entries:
if config_value is not None:
data[config_name] = config_class(**config_value)

View file

@ -7,13 +7,18 @@ This uses aws_sdk_bedrock_runtime for bidirectional streaming with Nova Sonic.
import asyncio
import contextlib
import json
from collections.abc import AsyncIterator, Mapping
from typing import Final, Protocol
from pydantic import JsonValue, TypeAdapter
import litellm
from litellm._logging import _redact_string, verbose_proxy_logger
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.litellm_core_utils.realtime_streaming import DefaultLoggedRealTimeEventTypes
from litellm.types.llms.openai import OpenAIRealtimeEvents
from litellm.types.realtime import RealtimeResponseTransformInput
from ..base_aws_llm import BaseAWSLLM
@ -32,6 +37,17 @@ def _json_str(value: JsonValue) -> str | None:
return value if isinstance(value, str) else None
def _should_log_event(openai_message: Mapping[str, object]) -> bool:
logged_types: Final = (
litellm.logged_real_time_event_types
if litellm.logged_real_time_event_types is not None
else DefaultLoggedRealTimeEventTypes
)
if logged_types == "*":
return True
return openai_message.get("type") in logged_types
class RealtimeClientWebSocket(Protocol):
"""The client-facing websocket surface the realtime bridge talks to."""
@ -205,16 +221,22 @@ class BedrockRealtime(BaseAWSLLM):
)
)
bedrock_to_client_task: Final = asyncio.create_task(
self._forward_bedrock_to_client(
bedrock_stream,
websocket,
transformation_config,
model,
logging_obj,
session_state,
async def forward_bedrock_and_collect_logged_events() -> tuple[OpenAIRealtimeEvents, ...]:
return tuple(
[
event
async for event in self._forward_bedrock_to_client(
bedrock_stream,
websocket,
transformation_config,
model,
logging_obj,
session_state,
)
]
)
)
bedrock_to_client_task: Final = asyncio.create_task(forward_bedrock_and_collect_logged_events())
# Wait for both tasks to complete
await asyncio.gather(
@ -223,6 +245,27 @@ class BedrockRealtime(BaseAWSLLM):
return_exceptions=True,
)
forwarded_logged_events: Final = (
bedrock_to_client_task.result()
if not bedrock_to_client_task.cancelled() and bedrock_to_client_task.exception() is None
else ()
)
logged_events: Final = (
*forwarded_logged_events,
*(
leftover_event
for leftover_event in transformation_config.leftover_usage_done_events()
if _should_log_event(leftover_event)
),
)
if logged_events:
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
logging_obj.dispatch_success_handlers(
list(logged_events), # mutable-ok: realtime spend logging requires a list result
prefer_async_handlers=True,
)
)
except Exception as e:
verbose_proxy_logger.exception("Error in BedrockRealtime.async_realtime: %s", e)
try:
@ -304,8 +347,8 @@ class BedrockRealtime(BaseAWSLLM):
model: str,
logging_obj: LiteLLMLogging,
session_state: RealtimeResponseTransformInput,
):
"""Forward messages from Bedrock stream to client WebSocket."""
) -> AsyncIterator[OpenAIRealtimeEvents]:
"""Forward messages from Bedrock to the client, yielding the ones to record for spend logging."""
try:
while True:
# Receive from Bedrock
@ -353,11 +396,14 @@ class BedrockRealtime(BaseAWSLLM):
)
# Send transformed messages to client
openai_messages = transformed_response.get("response", [])
response_value = transformed_response["response"]
openai_messages = response_value if isinstance(response_value, list) else (response_value,)
for openai_message in openai_messages:
message_json = json.dumps(openai_message)
await client_ws.send_text(message_json)
verbose_proxy_logger.debug("Bedrock Realtime: Sent to client: %s", message_json[:200])
if _should_log_event(openai_message):
yield openai_message
except Exception as e:
verbose_proxy_logger.debug("Bedrock to client forwarding ended: %s", e, exc_info=True)

View file

@ -7,7 +7,7 @@ Transforms between OpenAI Realtime API format and Bedrock Nova Sonic format.
import base64
import json
import uuid as uuid_lib
from typing import Any, Final
from typing import Any, Final, cast
from pydantic import BaseModel
@ -20,29 +20,54 @@ from litellm.types.llms.openai import (
OpenAIRealtimeContentPartDone,
OpenAIRealtimeDoneEvent,
OpenAIRealtimeEvents,
OpenAIRealtimeInputAudioBufferSpeechEvent,
OpenAIRealtimeInputAudioTranscriptionCompleted,
OpenAIRealtimeInputAudioTranscriptionDelta,
OpenAIRealtimeOutputItemDone,
OpenAIRealtimeResponseAudioDone,
OpenAIRealtimeResponseContentPartAdded,
OpenAIRealtimeResponseDelta,
OpenAIRealtimeResponseDoneObject,
OpenAIRealtimeResponseTextDone,
OpenAIRealtimeResponseUsage,
OpenAIRealtimeStreamResponseBaseObject,
OpenAIRealtimeStreamResponseOutputItemAdded,
OpenAIRealtimeStreamSession,
OpenAIRealtimeStreamSessionEvents,
OpenAIRealtimeUsageTokenDetails,
)
from litellm.types.realtime import (
ALL_DELTA_TYPES,
RealtimeResponseTransformInput,
RealtimeResponseTypedDict,
)
from litellm.utils import get_empty_usage
class BedrockContentEnd(BaseModel):
stopReason: str | None = None
class BedrockUsageTokenDetails(BaseModel):
speechTokens: int = 0
textTokens: int = 0
class BedrockUsageDetailsTotal(BaseModel):
input: BedrockUsageTokenDetails = BedrockUsageTokenDetails()
output: BedrockUsageTokenDetails = BedrockUsageTokenDetails()
class BedrockUsageDetails(BaseModel):
total: BedrockUsageDetailsTotal = BedrockUsageDetailsTotal()
class BedrockUsageEvent(BaseModel):
totalInputTokens: int = 0
totalOutputTokens: int = 0
totalTokens: int = 0
details: BedrockUsageDetails = BedrockUsageDetails()
TRIGGER_AUDIO_SAMPLE_RATE_HERTZ: Final = 16000
TRIGGER_AUDIO_BYTES_PER_SECOND: Final = TRIGGER_AUDIO_SAMPLE_RATE_HERTZ * 2
TRIGGER_LEADING_SILENCE: Final = bytes(TRIGGER_AUDIO_BYTES_PER_SECOND // 2)
@ -87,6 +112,15 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
# Text configuration
self.text_media_type = "text/plain"
# Response-stream state (Bedrock events carry no role on textOutput,
# so the USER/ASSISTANT split from contentStart is tracked here)
self._user_transcript_active = False
self._user_transcript_generation_stage: str | None = None
self._user_item_id: str | None = None
self._user_transcript_buffer = ""
self._cumulative_usage = BedrockUsageEvent()
self._reported_usage = BedrockUsageEvent()
def validate_environment(self, headers: dict, model: str, api_key: str | None = None) -> dict:
"""Validate environment - no special validation needed for Bedrock."""
return headers
@ -691,6 +725,11 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
role: Final = content_start.get("role")
if role != "ASSISTANT":
if role == "USER" and content_start.get("type") == "TEXT":
self._user_transcript_active = True
self._user_transcript_generation_stage = self._parse_generation_stage(
content_start.get("additionalModelFields")
)
return (
[],
current_response_id,
@ -700,6 +739,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
)
verbose_logger.debug("Handling ASSISTANT contentStart")
is_new_response: Final = current_response_id is None
# Initialize IDs if needed
if not current_response_id:
@ -715,7 +755,8 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
returned_messages: Final[list[OpenAIRealtimeEvents]] = []
# Send response.created
# Send response.created only once per response (a response can contain
# multiple content blocks, e.g. TEXT then AUDIO)
response_created: Final = OpenAIRealtimeStreamResponseBaseObject(
type="response.created",
event_id=f"event_{uuid.uuid4()}",
@ -727,7 +768,8 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
"conversation_id": current_conversation_id,
},
)
returned_messages.append(response_created)
if is_new_response:
returned_messages.append(response_created)
# Send response.output_item.added
output_item_added: Final = OpenAIRealtimeStreamResponseOutputItemAdded(
@ -767,6 +809,108 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
current_delta_type,
)
@staticmethod
def _parse_generation_stage(additional_model_fields: object) -> str | None:
if not isinstance(additional_model_fields, str):
return None
try:
parsed: Final = json.loads(additional_model_fields)
except json.JSONDecodeError:
return None
stage: Final = parsed.get("generationStage") if isinstance(parsed, dict) else None
return stage if isinstance(stage, str) else None
def _current_user_item_id(self, new_utterance: bool = False) -> str:
"""Item id shared by all events of one user utterance (speech boundaries and transcript)."""
if new_utterance or self._user_item_id is None:
self._user_item_id = f"item_{uuid.uuid4()}"
return self._user_item_id
def transform_user_speech_event(self, is_speech_start: bool) -> tuple[OpenAIRealtimeEvents, ...]:
"""Transform Bedrock userSpeechStart/userSpeechEnd to OpenAI speech boundary events."""
verbose_logger.debug("Handling userSpeech%s", "Start" if is_speech_start else "End")
speech_event: Final[OpenAIRealtimeInputAudioBufferSpeechEvent] = {
"type": "input_audio_buffer.speech_started" if is_speech_start else "input_audio_buffer.speech_stopped",
"event_id": f"event_{uuid.uuid4()}",
"item_id": self._current_user_item_id(new_utterance=is_speech_start),
}
return (speech_event,)
def transform_usage_event(self, usage_event: BedrockUsageEvent) -> None:
"""Record Bedrock's session-cumulative usage totals for the next response.done."""
verbose_logger.debug("Handling usageEvent")
self._cumulative_usage = usage_event
def _take_usage_delta(self) -> OpenAIRealtimeResponseUsage:
"""Usage for the response now completing: cumulative totals minus what prior response.done events reported."""
prior: Final = self._reported_usage
latest: Final = self._cumulative_usage
self._reported_usage = latest
input_details: Final[OpenAIRealtimeUsageTokenDetails] = {
"audio_tokens": latest.details.total.input.speechTokens - prior.details.total.input.speechTokens,
"text_tokens": latest.details.total.input.textTokens - prior.details.total.input.textTokens,
"cached_tokens": 0,
}
output_details: Final[OpenAIRealtimeUsageTokenDetails] = {
"audio_tokens": latest.details.total.output.speechTokens - prior.details.total.output.speechTokens,
"text_tokens": latest.details.total.output.textTokens - prior.details.total.output.textTokens,
}
usage_delta: Final[OpenAIRealtimeResponseUsage] = {
"input_tokens": latest.totalInputTokens - prior.totalInputTokens,
"output_tokens": latest.totalOutputTokens - prior.totalOutputTokens,
"total_tokens": latest.totalTokens - prior.totalTokens,
"input_token_details": input_details,
"output_token_details": output_details,
}
return usage_delta
def leftover_usage_done_events(self) -> tuple[OpenAIRealtimeEvents, ...]:
"""Logged-only response.done for usage Bedrock reports after the final turn's contentEnd."""
if self._cumulative_usage == self._reported_usage:
return ()
usage: Final = self._take_usage_delta()
leftover_done: Final = OpenAIRealtimeDoneEvent(
type="response.done",
event_id=f"event_{uuid.uuid4()}",
response=OpenAIRealtimeResponseDoneObject(
object="realtime.response",
id=f"resp_{uuid.uuid4()}",
status="completed",
conversation_id=f"conv_{uuid.uuid4()}",
usage=dict(usage), # mutable-ok: OpenAIRealtimeResponseDoneObject types usage as plain dict
),
)
return (leftover_done,)
def transform_user_transcript_event(self, transcript: str) -> tuple[OpenAIRealtimeEvents, ...]:
"""Transform a USER-role Bedrock textOutput (ASR transcript) to an OpenAI transcription delta."""
verbose_logger.debug("Handling USER textOutput (ASR transcript)")
delta_event: Final[OpenAIRealtimeInputAudioTranscriptionDelta] = {
"type": "conversation.item.input_audio_transcription.delta",
"event_id": f"event_{uuid.uuid4()}",
"item_id": self._current_user_item_id(),
"content_index": 0,
"delta": transcript,
}
if self._user_transcript_generation_stage != "SPECULATIVE":
self._user_transcript_buffer += transcript
return (delta_event,)
def user_transcript_completed_events(self) -> tuple[OpenAIRealtimeEvents, ...]:
"""One completed event with the full transcript once the FINAL user content block ends."""
transcript: Final = self._user_transcript_buffer
if not transcript:
return ()
self._user_transcript_buffer = ""
completed_event: Final[OpenAIRealtimeInputAudioTranscriptionCompleted] = {
"type": "conversation.item.input_audio_transcription.completed",
"event_id": f"event_{uuid.uuid4()}",
"item_id": self._current_user_item_id(),
"content_index": 0,
"transcript": transcript,
}
return (completed_event,)
def transform_text_output_event(
self,
event: dict,
@ -985,7 +1129,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
if not current_response_id or not current_conversation_id:
return [], None, None, None
usage_obj: Final = get_empty_usage()
usage: Final = self._take_usage_delta()
response_done: Final = OpenAIRealtimeDoneEvent(
type="response.done",
event_id=f"event_{uuid.uuid4()}",
@ -995,11 +1139,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
status="completed",
output=[],
conversation_id=current_conversation_id,
usage={
"prompt_tokens": usage_obj.prompt_tokens,
"completion_tokens": usage_obj.completion_tokens,
"total_tokens": usage_obj.total_tokens,
},
usage=dict(usage),
),
)
@ -1042,8 +1182,6 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
# Create a function call arguments done event
# This is a custom event format that matches what clients expect
from typing import cast
function_call_event: Final[dict[str, Any]] = {
"type": "response.function_call_arguments.done",
"event_id": f"event_{uuid.uuid4()}",
@ -1194,18 +1332,26 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
returned_messages.extend(events)
elif "textOutput" in event:
events, current_delta_chunks = self.transform_text_output_event(
event,
current_output_item_id,
current_response_id,
current_delta_chunks,
)
returned_messages.extend(events)
if self._user_transcript_active:
returned_messages.extend(self.transform_user_transcript_event(event["textOutput"].get("content", "")))
else:
events, current_delta_chunks = self.transform_text_output_event(
event,
current_output_item_id,
current_response_id,
current_delta_chunks,
)
returned_messages.extend(events)
elif "audioOutput" in event:
events = self.transform_audio_output_event(event, current_output_item_id, current_response_id)
returned_messages.extend(events)
elif "contentEnd" in event and self._user_transcript_active:
self._user_transcript_active = False
self._user_transcript_generation_stage = None
returned_messages.extend(self.user_transcript_completed_events())
elif "contentEnd" in event:
events, current_delta_chunks = self.transform_content_end_event(
event,
@ -1224,6 +1370,12 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
) = self._response_done_events(current_response_id, current_conversation_id)
returned_messages.extend(done_events)
elif "userSpeechStart" in event or "userSpeechEnd" in event:
returned_messages.extend(self.transform_user_speech_event("userSpeechStart" in event))
elif "usageEvent" in event:
self.transform_usage_event(BedrockUsageEvent.model_validate(event["usageEvent"]))
elif "toolUse" in event:
events, tool_call_id, tool_name = self.transform_tool_use_event(
event, current_output_item_id, current_response_id

View file

@ -15,9 +15,11 @@ API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/overview
from .chat.transformation import GigaChatConfig, GigaChatError
from .embedding.transformation import GigaChatEmbeddingConfig
from .passthrough.transformation import GigaChatPassthroughConfig
__all__ = [
__all__ = (
"GigaChatConfig",
"GigaChatEmbeddingConfig",
"GigaChatError",
]
"GigaChatPassthroughConfig",
)

View file

@ -7,6 +7,7 @@ Based on official GigaChat SDK authentication flow.
import time
import uuid
from collections.abc import Mapping
from typing import Final
import httpx
@ -16,7 +17,7 @@ from litellm.caching.caching import InMemoryCache
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
_get_httpx_client,
_get_httpx_client, # pyright: ignore[reportPrivateUsage] # house cached-client factory has no public alias
get_async_httpx_client,
)
from litellm.secret_managers.main import get_secret_str
@ -63,6 +64,7 @@ def get_access_token(
credentials: str | None = None,
scope: str | None = None,
auth_url: str | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> str:
"""
Get valid access token, using cache if available.
@ -78,71 +80,88 @@ def get_access_token(
Raises:
GigaChatAuthError: If authentication fails
"""
credentials = credentials or _get_credentials()
if not credentials:
if not litellm_params:
litellm_params = {} # mutable-ok: empty dict default; rebind-ok: provide default
access_token: Final = litellm_params.get("gigachat_access_token") or get_secret_str("GIGACHAT_ACCESS_TOKEN")
if access_token:
return access_token
effective_credentials: Final = credentials or _get_credentials()
if not effective_credentials:
raise GigaChatAuthError(
status_code=401,
message="GigaChat credentials not provided. Set GIGACHAT_CREDENTIALS or GIGACHAT_API_KEY environment variable.",
)
scope = scope or _get_scope()
auth_url = auth_url or _get_auth_url()
effective_scope: Final = scope or litellm_params.get("gigachat_scope") or _get_scope()
effective_auth_url: Final = auth_url or litellm_params.get("gigachat_auth_url") or _get_auth_url()
# Check cache
cache_key: Final = f"gigachat_token:{credentials[:16]}"
cache_key: Final = f"gigachat_token:{effective_credentials[:16]}"
cached: Final = _token_cache.get_cache(cache_key)
if cached:
token, expires_at = cached
_token, _expires_at = cached
# Check if token is still valid (with buffer)
if time.time() * 1000 < expires_at - TOKEN_EXPIRY_BUFFER_MS:
if time.time() * 1000 < _expires_at - TOKEN_EXPIRY_BUFFER_MS:
verbose_logger.debug("Using cached GigaChat access token")
return token
return _token
# Request new token
token, expires_at = _request_token_sync(credentials, scope, auth_url)
new_token, new_expires_at = _request_token_sync(effective_credentials, effective_scope, effective_auth_url) # pyright: ignore[reportArgumentType] # credential keys may be broader than str
# Cache token
ttl_seconds: Final = max(0, (expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000)
if ttl_seconds > 0:
_token_cache.set_cache(cache_key, (token, expires_at), ttl=ttl_seconds)
if new_expires_at:
# Cache token
ttl_seconds: Final = max(0, (new_expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000)
if ttl_seconds > 0:
_token_cache.set_cache(cache_key, (new_token, new_expires_at), ttl=ttl_seconds)
return token
return new_token
async def get_access_token_async(
credentials: str | None = None,
scope: str | None = None,
auth_url: str | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> str:
"""Async version of get_access_token."""
credentials = credentials or _get_credentials()
if not credentials:
if not litellm_params:
litellm_params = {} # mutable-ok: empty dict default; rebind-ok: provide default
access_token: Final = litellm_params.get("gigachat_access_token") or get_secret_str("GIGACHAT_ACCESS_TOKEN")
if access_token:
return access_token
effective_credentials: Final = credentials or _get_credentials()
if not effective_credentials:
raise GigaChatAuthError(
status_code=401,
message="GigaChat credentials not provided. Set GIGACHAT_CREDENTIALS or GIGACHAT_API_KEY environment variable.",
)
scope = scope or _get_scope()
auth_url = auth_url or _get_auth_url()
effective_scope: Final = scope or litellm_params.get("gigachat_scope") or _get_scope()
effective_auth_url: Final = auth_url or litellm_params.get("gigachat_auth_url") or _get_auth_url()
# Check cache
cache_key: Final = f"gigachat_token:{credentials[:16]}"
cache_key: Final = f"gigachat_token:{effective_credentials[:16]}"
cached: Final = _token_cache.get_cache(cache_key)
if cached:
token, expires_at = cached
if time.time() * 1000 < expires_at - TOKEN_EXPIRY_BUFFER_MS:
_token, _expires_at = cached
if time.time() * 1000 < _expires_at - TOKEN_EXPIRY_BUFFER_MS:
verbose_logger.debug("Using cached GigaChat access token")
return token
return _token
# Request new token
token, expires_at = await _request_token_async(credentials, scope, auth_url)
new_token, new_expires_at = await _request_token_async(effective_credentials, effective_scope, effective_auth_url) # pyright: ignore[reportArgumentType] # credential keys may be broader than str
# Cache token
ttl_seconds: Final = max(0, (expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000)
if ttl_seconds > 0:
_token_cache.set_cache(cache_key, (token, expires_at), ttl=ttl_seconds)
if new_expires_at:
# Cache token
ttl_seconds: Final = max(0, (new_expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000)
if ttl_seconds > 0:
_token_cache.set_cache(cache_key, (new_token, new_expires_at), ttl=ttl_seconds)
return token
return new_token
def _request_token_sync(
@ -154,7 +173,7 @@ def _request_token_sync(
Request new access token from GigaChat OAuth endpoint (sync).
Returns:
Tuple of (access_token, expires_at_ms)
tuple of (access_token, expires_at_ms)
"""
headers: Final = {
"Authorization": f"Basic {credentials}",
@ -169,7 +188,7 @@ def _request_token_sync(
client: Final = _get_http_client()
response: Final = client.post(auth_url, headers=headers, data=data, timeout=30)
response.raise_for_status()
return _parse_token_response(response)
return _parse_token_response(response) # pyright: ignore[reportArgumentType] # httpx Response may be None at type level
except httpx.HTTPStatusError as e:
raise GigaChatAuthError(
status_code=e.response.status_code,
@ -204,7 +223,7 @@ async def _request_token_async(
)
response: Final = await client.post(auth_url, headers=headers, data=data, timeout=30)
response.raise_for_status()
return _parse_token_response(response)
return _parse_token_response(response) # pyright: ignore[reportArgumentType] # httpx Response may be None at type level
except httpx.HTTPStatusError as e:
raise GigaChatAuthError(
status_code=e.response.status_code,
@ -223,7 +242,7 @@ def _parse_token_response(response: httpx.Response) -> tuple[str, int]:
# GigaChat returns either 'tok'/'exp' or 'access_token'/'expires_at'
access_token: Final = data.get("tok") or data.get("access_token")
expires_at = data.get("exp") or data.get("expires_at")
expires_at_raw: Final = data.get("exp") or data.get("expires_at")
if not access_token:
raise GigaChatAuthError(
@ -232,8 +251,11 @@ def _parse_token_response(response: httpx.Response) -> tuple[str, int]:
)
# expires_at is in milliseconds
if isinstance(expires_at, str):
expires_at = int(expires_at)
expires_at: int # rebind-ok: conditionally assigned from str or int
if isinstance(expires_at_raw, str):
expires_at = int(expires_at_raw) # rebind-ok: conditionally assigned from str or int
else:
expires_at = expires_at_raw # pyright: ignore[reportAssignmentType] # raw value is int or str; converted above; rebind-ok: conditionally assigned from str or int
verbose_logger.debug("GigaChat access token obtained successfully")
return access_token, expires_at

View file

@ -5,8 +5,8 @@ GigaChat Chat Module
from .streaming import GigaChatModelResponseIterator
from .transformation import GigaChatConfig, GigaChatError
__all__ = [
__all__ = (
"GigaChatConfig",
"GigaChatError",
"GigaChatModelResponseIterator",
]
)

View file

@ -4,13 +4,15 @@ GigaChat Streaming Response Handler
import json
import uuid
from collections.abc import Mapping, Sequence
from typing import Any, Final
from litellm.llms.gigachat.utils import convert_usage
from litellm.types.llms.openai import (
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
)
from litellm.types.utils import GenericStreamingChunk
from litellm.types.utils import ChatCompletionUsageBlock, GenericStreamingChunk
class GigaChatModelResponseIterator:
@ -26,14 +28,9 @@ class GigaChatModelResponseIterator:
self.response_iterator = self.streaming_response
self.json_mode = json_mode
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk:
def chunk_parser(self, chunk: Mapping[str, object]) -> GenericStreamingChunk:
"""Parse a single streaming chunk from GigaChat."""
text = ""
tool_use: ChatCompletionToolCallChunk | None = None
is_finished = False
finish_reason: str | None = None
choices: Final = chunk.get("choices", [])
choices: Sequence = chunk.get("choices") or () # mutable-ok: tuple literal as default
if not choices:
return GenericStreamingChunk(
text="",
@ -45,40 +42,63 @@ class GigaChatModelResponseIterator:
)
choice: Final = choices[0]
delta: Final = choice.get("delta", {})
finish_reason = choice.get("finish_reason")
delta: Mapping[str, object] = choice.get("delta") or {} # mutable-ok: empty dict default for get
chunk_finish_reason: Final = choice.get("finish_reason")
# Extract text content
text = delta.get("content", "") or ""
text: Final = delta.get("content", "") or ""
usage_block: ChatCompletionUsageBlock | None = None # rebind-ok: conditionally assigned after stop detection
tool_use: ChatCompletionToolCallChunk | None = None # rebind-ok: conditionally assigned on function_call
finish_reason: str | None = chunk_finish_reason
# Handle function_call in stream
if finish_reason == "function_call" and delta.get("function_call"):
func_call: Final = delta["function_call"]
args = func_call.get("arguments", {})
if isinstance(args, dict):
args = json.dumps(args, ensure_ascii=False)
raw_function_call: Final = delta.get("function_call")
if chunk_finish_reason == "function_call" and isinstance(raw_function_call, Mapping) and raw_function_call:
func_call: Final[Mapping[str, object]] = raw_function_call
args_raw: Final[object] = func_call.get("arguments") or {}
args_str: str # rebind-ok: conditionally assigned from dict or str
if isinstance(args_raw, dict):
args_str = json.dumps(args_raw, ensure_ascii=False) # rebind-ok: build from dict
else:
args_str = str(args_raw)
name_raw: Final = func_call.get("name")
tool_use = ChatCompletionToolCallChunk(
id=f"call_{uuid.uuid4().hex[:24]}",
type="function",
function=ChatCompletionToolCallFunctionChunk(
name=func_call.get("name", ""),
arguments=args,
name=name_raw if isinstance(name_raw, str) else "",
arguments=args_str,
),
index=0,
)
finish_reason = "tool_calls"
if finish_reason is not None:
is_finished = True
usage_data: Final = chunk.get("usage") or {} # mutable-ok: empty dict default
if usage_data and isinstance(usage_data, dict):
validated_usage: Final = {k: int(v) for k, v in usage_data.items()}
usage = convert_usage(validated_usage)
_prompt_details: dict | None = (
usage.prompt_tokens_details.model_dump() if usage.prompt_tokens_details else None
) # rebind-ok: conditional
_completion_details: dict | None = (
usage.completion_tokens_details.model_dump() if usage.completion_tokens_details else None
) # rebind-ok: conditional
usage_block = ChatCompletionUsageBlock( # pyright: ignore[reportCallIssue] # TypedDict kwarg constructor
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
total_tokens=usage.total_tokens,
prompt_tokens_details=_prompt_details,
completion_tokens_details=_completion_details,
)
return GenericStreamingChunk(
text=text,
text=str(text),
tool_use=tool_use,
is_finished=is_finished,
is_finished=chunk_finish_reason is not None,
finish_reason=finish_reason or "",
usage=None,
usage=usage_block,
index=choice.get("index", 0),
)

View file

@ -4,19 +4,22 @@ GigaChat Chat Transformation
Transforms OpenAI-format requests to GigaChat format and back.
"""
from __future__ import annotations
import json
import time
import uuid
from collections.abc import AsyncIterator, Iterator
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
import httpx
from litellm._logging import verbose_logger
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.gigachat.utils import convert_usage, get_api_base
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Choices, Message, ModelResponse, Usage
from litellm.types.utils import Choices, Message, ModelResponse
from ..authenticator import get_access_token
from ..file_handler import upload_file_sync
@ -30,9 +33,6 @@ if TYPE_CHECKING:
else:
LiteLLMLoggingObj = Any
# GigaChat API endpoint
GIGACHAT_BASE_URL: Final = "https://gigachat.devices.sberbank.ru/api/v1"
def is_valid_json(value: str) -> bool:
"""Checks whether the value passed is a valid serialized JSON string"""
@ -90,30 +90,30 @@ class GigaChatConfig(BaseConfig):
api_base: str | None,
api_key: str | None,
model: str,
optional_params: dict,
litellm_params: dict,
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
stream: bool | None = None,
) -> str:
"""Get complete API URL for chat completions."""
base: Final = api_base or get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL
base: Final = get_api_base(api_base)
return f"{base}/chat/completions"
def validate_environment(
self,
headers: dict,
headers: dict, # mutable-ok: mutates in place per GigaChat OAuth setup
model: str,
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
) -> dict: # mutable-ok: base class contract returns dict for httpx
"""
Set up headers with OAuth token.
"""
# Get access token
credentials: Final = api_key or get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY")
access_token: Final = get_access_token(credentials=credentials)
access_token: Final = get_access_token(credentials=credentials, litellm_params=litellm_params)
# Store credentials for image uploads
self._current_credentials = credentials
@ -125,9 +125,9 @@ class GigaChatConfig(BaseConfig):
return headers
def get_supported_openai_params(self, model: str) -> list[str]:
def get_supported_openai_params(self, model: str) -> list[str]: # mutable-ok: base class contract returns list
"""Return list of supported OpenAI parameters."""
return [
return [ # mutable-ok: base class contract returns list
"stream",
"temperature",
"top_p",
@ -143,11 +143,11 @@ class GigaChatConfig(BaseConfig):
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
non_default_params: Mapping[str, object],
optional_params: dict, # mutable-ok: mutated in place per GigaChat mapping
model: str,
drop_params: bool,
) -> dict:
) -> dict: # mutable-ok: base class contract returns dict
"""Map OpenAI parameters to GigaChat parameters."""
for param, value in non_default_params.items():
if param == "stream":
@ -167,42 +167,50 @@ class GigaChatConfig(BaseConfig):
pass
elif param == "tools":
# Convert tools to functions format
optional_params["functions"] = self._convert_tools_to_functions(value)
if isinstance(value, Sequence):
optional_params["functions"] = self._convert_tools_to_functions(value)
elif param == "tool_choice":
# Map OpenAI tool_choice to GigaChat function_call
mapped_choice = self._map_tool_choice(value)
if mapped_choice is not None:
optional_params["function_call"] = mapped_choice
if isinstance(value, (str, Mapping)):
mapped_choice = self._map_tool_choice(value)
if mapped_choice is not None:
optional_params["function_call"] = mapped_choice
elif param == "functions":
optional_params["functions"] = value
elif param == "function_call":
optional_params["function_call"] = value
elif param == "response_format":
# Handle structured output via function calling
if value.get("type") == "json_schema":
if isinstance(value, Mapping) and value.get("type") == "json_schema":
json_schema = value.get("json_schema", {})
schema_name = json_schema.get("name", "structured_output")
schema = json_schema.get("schema", {})
function_def = {
function_def = { # mutable-ok: request payload for httpx
"name": schema_name,
"description": f"Output structured response: {schema_name}",
"parameters": schema,
}
if "functions" not in optional_params:
optional_params["functions"] = []
optional_params["functions"].append(function_def)
optional_params["function_call"] = {"name": schema_name}
existing_functions = optional_params.get("functions")
optional_params["functions"] = [
*(
existing_functions
if isinstance(existing_functions, Sequence) and not isinstance(existing_functions, str)
else ()
),
function_def,
]
optional_params["function_call"] = {"name": schema_name} # mutable-ok: request payload
optional_params["_structured_output"] = True
return optional_params
def _convert_tools_to_functions(self, tools: list[dict]) -> list[dict]:
def _convert_tools_to_functions(self, tools: Sequence) -> Sequence[dict]:
"""Convert OpenAI tools format to GigaChat functions format."""
functions: Final = []
functions: Final[list[dict]] = [] # mutable-ok: accumulator for building functions list
for tool in tools:
if tool.get("type") == "function":
if isinstance(tool, dict) and tool.get("type") == "function":
func = tool.get("function", {})
functions.append(
{
@ -213,7 +221,7 @@ class GigaChatConfig(BaseConfig):
)
return functions
def _map_tool_choice(self, tool_choice: str | dict) -> str | dict | None:
def _map_tool_choice(self, tool_choice: str | Mapping[str, object]) -> str | Mapping[str, object] | None:
"""
Map OpenAI tool_choice to GigaChat function_call format.
@ -246,8 +254,9 @@ class GigaChatConfig(BaseConfig):
# OpenAI format: {"type": "function", "function": {"name": "func_name"}}
# GigaChat format: {"name": "func_name"}
if tool_choice.get("type") == "function":
func_name: Final = tool_choice.get("function", {}).get("name")
if func_name:
function_spec: Final = tool_choice.get("function")
func_name: Final = function_spec.get("name") if isinstance(function_spec, Mapping) else None
if isinstance(func_name, str) and func_name:
return {"name": func_name}
# Default to None (don't set function_call)
@ -273,20 +282,51 @@ class GigaChatConfig(BaseConfig):
verbose_logger.error("Failed to upload image: %s", e)
return None
def _transform_list_content(self, content: Sequence) -> tuple[str, Sequence[str]]:
"""
Extract text and image attachments from a multimodal message content list.
Args:
content: List of content parts (OpenAI multimodal format)
Returns:
Tuple of (combined text, list of attachment file ids)
"""
texts: Final[list[str]] = [] # mutable-ok: accumulator
attachments: Final[list[str]] = [] # mutable-ok: accumulator
for part in content:
if isinstance(part, dict):
if part.get("type") == "text":
texts.append(part.get("text", ""))
elif part.get("type") == "image_url":
# Extract image URL and upload to GigaChat
image_url: object = part.get("image_url", {})
upload_url: str
if isinstance(image_url, str):
upload_url = image_url
else:
upload_url = str(image_url.get("url", "")) if isinstance(image_url, dict) else ""
if upload_url:
file_id = self._upload_image(upload_url)
if file_id:
attachments.append(file_id)
text: Final = "\n".join(texts) if texts else ""
return text, attachments
def transform_request(
self,
model: str,
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
headers: Mapping[str, object],
) -> dict: # mutable-ok: request payload sent to httpx
"""Transform OpenAI request to GigaChat format."""
# Transform messages
giga_messages: Final = self._transform_messages(messages)
# Build request
request_data: Final = {
request_data: Final[dict[str, object]] = {
"model": model.replace("gigachat/", ""),
"messages": giga_messages,
}
@ -311,9 +351,9 @@ class GigaChatConfig(BaseConfig):
return request_data
def _transform_messages(self, messages: list[AllMessageValues]) -> list[dict]:
def _transform_messages(self, messages: Sequence[AllMessageValues]) -> Sequence[dict]:
"""Transform OpenAI messages to GigaChat format."""
transformed: Final = []
transformed: Final[list[dict]] = [] # mutable-ok: accumulator for building transformed messages
for i, msg in enumerate(messages):
message = dict(msg)
@ -341,24 +381,7 @@ class GigaChatConfig(BaseConfig):
# Handle list content (multimodal) - extract text and images
content = message.get("content")
if isinstance(content, list):
texts = []
attachments = []
for part in content:
if isinstance(part, dict):
if part.get("type") == "text":
texts.append(part.get("text", ""))
elif part.get("type") == "image_url":
# Extract image URL and upload to GigaChat
image_url = part.get("image_url", {})
if isinstance(image_url, str):
url = image_url
else:
url = image_url.get("url", "")
if url:
file_id = self._upload_image(url)
if file_id:
attachments.append(file_id)
message["content"] = "\n".join(texts) if texts else ""
message["content"], attachments = self._transform_list_content(content)
if attachments:
message["attachments"] = attachments
@ -393,7 +416,7 @@ class GigaChatConfig(BaseConfig):
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: "tiktoken.Encoding | None",
encoding: tiktoken.Encoding | None,
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:
@ -408,7 +431,7 @@ class GigaChatConfig(BaseConfig):
is_structured_output: Final = optional_params.get("_structured_output", False)
choices: Final = []
choices: Final[list[Choices]] = [] # mutable-ok: accumulator for building response choices
for choice in response_json.get("choices", []):
message_data = choice.get("message", {})
finish_reason = choice.get("finish_reason", "stop")
@ -462,11 +485,7 @@ class GigaChatConfig(BaseConfig):
# Build usage
usage_data: Final = response_json.get("usage", {})
usage: Final = Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
)
usage: Final = convert_usage(usage_data)
model_response.id = response_json.get("id", f"chatcmpl-{uuid.uuid4().hex[:12]}")
model_response.created = response_json.get("created", int(time.time()))

View file

@ -5,6 +5,8 @@ Transforms OpenAI /v1/embeddings format to GigaChat format.
API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/reference/rest/post-embeddings
"""
from __future__ import annotations
import types
from typing import Final
@ -14,14 +16,12 @@ from litellm import LlmProviders
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.llms.gigachat.utils import get_api_base
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
from litellm.types.utils import EmbeddingResponse
from ..authenticator import get_access_token
# GigaChat API endpoint
GIGACHAT_BASE_URL: Final = "https://gigachat.devices.sberbank.ru/api/v1"
class GigaChatEmbeddingError(BaseLLMException):
"""GigaChat Embedding API error."""
@ -78,9 +78,9 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
Returns provider info for GigaChat.
Returns:
Tuple of (custom_llm_provider, api_base, dynamic_api_key)
tuple of (custom_llm_provider, api_base, dynamic_api_key)
"""
api_base = api_base or GIGACHAT_BASE_URL
api_base = get_api_base(api_base)
return LlmProviders.GIGACHAT.value, api_base, api_key
def get_complete_url(
@ -93,7 +93,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
stream: bool | None = None,
) -> str:
"""Get the complete URL for embeddings endpoint."""
base: Final = api_base or GIGACHAT_BASE_URL
base: Final = get_api_base(api_base)
return f"{base}/embeddings"
def transform_embedding_request(
@ -114,14 +114,12 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
"""
# Normalize input to list
if isinstance(input, str):
input_list: list = [input]
elif isinstance(input, list):
input_list = input
input_list: list = [input] # rebind-ok: locally scoped conversion
else:
input_list = [input]
input_list = input
# Remove gigachat/ prefix from model if present
model = model.removeprefix("gigachat/")
model = model.removeprefix("gigachat/") # rebind-ok: parameter reassignment for normalization
return {
"model": model,
@ -191,7 +189,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
Set up headers with OAuth token for GigaChat.
"""
# Get access token via OAuth
access_token: Final = get_access_token(api_key)
access_token: Final = get_access_token(credentials=api_key, litellm_params=litellm_params)
default_headers: Final = {
"Content-Type": "application/json",

View file

@ -9,6 +9,7 @@ import base64
import hashlib
import re
import uuid
from collections.abc import Mapping
from typing import Final
from litellm._logging import verbose_logger
@ -16,13 +17,11 @@ from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
)
from litellm.llms.gigachat.utils import get_api_base
from litellm.types.utils import LlmProviders
from .authenticator import get_access_token, get_access_token_async
# GigaChat API endpoint
GIGACHAT_BASE_URL: Final = "https://gigachat.devices.sberbank.ru/api/v1"
# Simple in-memory cache for file IDs
_file_cache: Final[dict[str, str]] = {}
@ -82,6 +81,7 @@ def upload_file_sync(
image_url: str,
credentials: str | None = None,
api_base: str | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> str | None:
"""
Upload file to GigaChat and return file_id (sync).
@ -114,10 +114,10 @@ def upload_file_sync(
filename: Final = f"{uuid.uuid4()}.{ext}"
# Get access token
access_token: Final = get_access_token(credentials)
access_token: Final = get_access_token(credentials=credentials, litellm_params=litellm_params)
# Upload to GigaChat
base_url: Final = api_base or GIGACHAT_BASE_URL
base_url: Final = get_api_base(api_base)
upload_url: Final = f"{base_url}/files"
client: Final = _get_httpx_client(params={"ssl_verify": False})
@ -147,6 +147,7 @@ async def upload_file_async(
image_url: str,
credentials: str | None = None,
api_base: str | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> str | None:
"""
Upload file to GigaChat and return file_id (async).
@ -179,10 +180,10 @@ async def upload_file_async(
filename: Final = f"{uuid.uuid4()}.{ext}"
# Get access token
access_token: Final = await get_access_token_async(credentials)
access_token: Final = await get_access_token_async(credentials=credentials, litellm_params=litellm_params)
# Upload to GigaChat
base_url: Final = api_base or GIGACHAT_BASE_URL
base_url: Final = get_api_base(api_base)
upload_url: Final = f"{base_url}/files"
client: Final = get_async_httpx_client(

View file

@ -0,0 +1,7 @@
"""
GigaChat passthrough Module
"""
from .transformation import GigaChatPassthroughConfig
__all__ = ("GigaChatPassthroughConfig",)

View file

@ -0,0 +1,213 @@
from __future__ import annotations
import json
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Final
import httpx
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.gigachat.authenticator import get_access_token
from litellm.llms.gigachat.chat.streaming import GigaChatModelResponseIterator
from litellm.llms.gigachat.utils import GIGACHAT_BASE_URL
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import EmbeddingResponse
if TYPE_CHECKING:
from httpx import URL, Response
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.utils import CostResponseTypes
class GigaChatPassthroughConfig(BasePassthroughConfig):
def is_streaming_request(self, endpoint: str, request_data: Mapping[str, object]) -> bool:
return request_data.get("stream", False)
def get_complete_url(
self,
api_base: str | None,
api_key: str | None,
model: str,
endpoint: str,
request_query_params: Mapping[str, object] | None,
litellm_params: Mapping[str, object],
) -> tuple[URL, str]:
"""Get complete API URL for chat completions."""
base_target_url: Final = self.get_api_base(api_base)
if base_target_url is None:
raise Exception("GigaChat api base not found")
complete_url: Final = f"{base_target_url}/{endpoint.lstrip('/')}"
return (
httpx.URL(complete_url),
base_target_url,
)
def validate_environment(
self,
headers: dict, # mutable-ok: mutates in place to set OAuth headers
model: str,
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
) -> dict: # mutable-ok: base class contract returns dict for httpx
"""
Set up headers with OAuth token.
"""
# Get access token
access_token: Final = get_access_token(credentials=api_key, litellm_params=litellm_params)
headers["Authorization"] = f"Bearer {access_token}" # rebind-ok: mutating for OAuth setup
headers["Content-Type"] = "application/json" # rebind-ok: mutating for OAuth setup
headers["Accept"] = "application/json" # rebind-ok: mutating for OAuth setup
return headers
def logging_non_streaming_response(
self,
model: str,
custom_llm_provider: str,
httpx_response: Response,
request_data: Mapping[str, object],
logging_obj: LiteLLMLoggingObj,
endpoint: str,
) -> CostResponseTypes | None:
from litellm import encoding
from litellm.types.utils import LlmProviders, ModelResponse
from litellm.utils import ProviderConfigManager
# cost tracking only for completions and embeddings
if "completions" in endpoint:
provider_chat_config: Final = ProviderConfigManager.get_provider_chat_config(
provider=LlmProviders(custom_llm_provider),
model=model,
)
if provider_chat_config is None:
raise ValueError(f"No provider config found for model: {model}")
raw_messages: Final = request_data.get("messages")
litellm_model_response: Final = provider_chat_config.transform_response(
model=model,
messages=list(raw_messages)
if isinstance(raw_messages, list)
else [], # mutable-ok: transform_response wants a list
raw_response=httpx_response,
model_response=ModelResponse(),
logging_obj=logging_obj,
optional_params={}, # mutable-ok: empty dict kwarg for transform_response
litellm_params={}, # mutable-ok: empty dict kwarg for transform_response
api_key="",
request_data=dict(request_data), # mutable-ok: transform_response wants a dict
encoding=encoding,
)
return litellm_model_response
if "embeddings" in endpoint:
provider_embedding_config: Final = ProviderConfigManager.get_provider_embedding_config(
provider=LlmProviders(custom_llm_provider),
model=model,
)
if provider_embedding_config is None:
raise ValueError(f"No provider config found for model: {model}")
litellm_embedding_response: Final[EmbeddingResponse] = (
provider_embedding_config.transform_embedding_response(
model=model,
raw_response=httpx_response,
model_response=EmbeddingResponse(),
logging_obj=logging_obj,
optional_params={}, # mutable-ok: empty dict kwarg for transform_embedding_response
api_key="",
request_data=dict(request_data), # mutable-ok: transform_embedding_response wants a dict
litellm_params={}, # mutable-ok: empty dict kwarg for transform_embedding_response
)
)
return litellm_embedding_response
return None
def handle_logging_collected_chunks(
self,
all_chunks: Sequence[str],
litellm_logging_obj: LiteLLMLoggingObj,
model: str,
custom_llm_provider: str,
endpoint: str,
) -> CostResponseTypes | None:
"""
1. Convert all_chunks to a ModelResponseStream
2. combine model_response_stream to model_response
3. Return the model_response
"""
from litellm.litellm_core_utils.streaming_handler import (
convert_generic_chunk_to_model_response_stream,
generic_chunk_has_all_required_fields,
)
from litellm.main import stream_chunk_builder
from litellm.types.utils import ModelResponseStream
all_translated_chunks: Final[list[object]] = [] # mutable-ok: accumulator
for chunk in all_chunks:
chunk = chunk.strip()
if not chunk or chunk == "[DONE]":
continue
chunk = chunk.removeprefix("data: ")
try:
message = json.loads(chunk)
except json.JSONDecodeError:
continue
gigachat_iterator = GigaChatModelResponseIterator(
streaming_response=None,
sync_stream=False,
)
translated_chunk = gigachat_iterator.chunk_parser(chunk=message)
if isinstance(translated_chunk, dict) and generic_chunk_has_all_required_fields( # pyright: ignore[reportUnnecessaryIsInstance] # runtime guard for patched chunk_parser
dict(translated_chunk)
):
chunk_obj = convert_generic_chunk_to_model_response_stream(
translated_chunk # pyright: ignore[reportArgumentType] # validated TypedDict
)
elif isinstance(translated_chunk, ModelResponseStream):
chunk_obj = translated_chunk
else:
continue
all_translated_chunks.append(chunk_obj)
if len(all_translated_chunks) > 0:
return stream_chunk_builder(
chunks=all_translated_chunks,
logging_obj=litellm_logging_obj,
)
return None
@staticmethod
def get_api_base(api_base: str | None = None) -> str | None:
return api_base or get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL
@staticmethod
def get_api_key(
api_key: str | None = None,
) -> str | None:
return api_key or get_secret_str("GIGACHAT_API_KEY")
@staticmethod
def get_base_model(model: str) -> str | None:
return model
def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]:
return list(super().get_models(api_key, api_base))

View file

@ -0,0 +1,26 @@
from collections.abc import Mapping
from typing import Final
from litellm.secret_managers.main import get_secret_str
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
# GigaChat API endpoint
GIGACHAT_BASE_URL: Final = "https://gigachat.devices.sberbank.ru/api/v1"
def convert_usage(usage_data: Mapping[str, int]) -> Usage:
precached_prompt_tokens: Final = usage_data.get("precached_prompt_tokens", 0)
prompt_tokens_details: Final = (
PromptTokensDetailsWrapper(cached_tokens=precached_prompt_tokens) if precached_prompt_tokens > 0 else None
)
return Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0) + precached_prompt_tokens,
completion_tokens=usage_data.get("completion_tokens", 0),
prompt_tokens_details=prompt_tokens_details,
total_tokens=usage_data.get("total_tokens", 0) + precached_prompt_tokens,
)
def get_api_base(api_base: str | None = None) -> str | None:
return api_base or get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL

View file

@ -9,7 +9,7 @@ response parsing, and streaming chunk parsing for models served with
import datetime
import json
from collections.abc import Iterable, Mapping, Sequence
from typing import Any, Final
from typing import Final
import httpx
from pydantic import JsonValue, TypeAdapter, ValidationError
@ -76,7 +76,7 @@ def _content_text(content: str | Iterable[Mapping[str, object]] | None) -> str:
return str(content)
def _extract_text_content(content: Any) -> str:
def _extract_text_content(content: str | Iterable[Mapping[str, object]] | None) -> str:
"""Return the plain-text representation of a message content value."""
return _content_text(content)

View file

@ -268,6 +268,7 @@ class BaseOpenAILLM:
"max_retries",
"organization",
"api_base",
"workload_identity_config",
)
openai_client_fields: Final = (
BaseOpenAILLM.get_openai_client_initialization_param_fields(client_type=client_type)

View file

@ -51,6 +51,7 @@ from .common_utils import (
drop_params_from_unprocessable_entity_error,
is_output_token_limit_error,
)
from .workload_identity import resolve_openai_workload_identity_config
openaiOSeriesConfig: Final = OpenAIOSeriesConfig()
openAIGPT5Config: Final = OpenAIGPT5Config()
@ -349,6 +350,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client: OpenAI | AsyncOpenAI | None = None,
shared_session: Optional["ClientSession"] = None,
) -> OpenAI | AsyncOpenAI | None:
workload_identity_config: Final = resolve_openai_workload_identity_config(api_key=api_key, api_base=api_base)
client_initialization_params: Final[dict] = locals()
if client is None:
if not isinstance(max_retries, int):
@ -364,28 +366,49 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
if cached_client:
if isinstance(cached_client, OpenAI) or isinstance(cached_client, AsyncOpenAI):
return cached_client
http_client: Final[httpx.Client | httpx.AsyncClient | None] = (
OpenAIChatCompletion._get_async_http_client(shared_session=shared_session)
if is_async
else OpenAIChatCompletion._get_sync_http_client()
)
if is_async:
_new_client: OpenAI | AsyncOpenAI = AsyncOpenAI(
api_key=api_key,
base_url=api_base,
http_client=http_client,
timeout=timeout,
max_retries=max_retries,
organization=organization,
async_http_client: Final = OpenAIChatCompletion._get_async_http_client(shared_session=shared_session)
http_client: httpx.Client | httpx.AsyncClient | None = async_http_client
_new_client: OpenAI | AsyncOpenAI = (
AsyncOpenAI(
workload_identity=workload_identity_config.to_sdk_workload_identity(),
base_url=api_base,
http_client=async_http_client,
timeout=timeout,
max_retries=max_retries,
organization=organization,
)
if workload_identity_config is not None
else AsyncOpenAI(
api_key=api_key,
base_url=api_base,
http_client=async_http_client,
timeout=timeout,
max_retries=max_retries,
organization=organization,
)
)
else:
_new_client = OpenAI(
api_key=api_key,
base_url=api_base,
http_client=http_client,
timeout=timeout,
max_retries=max_retries,
organization=organization,
sync_http_client: Final = OpenAIChatCompletion._get_sync_http_client()
http_client = sync_http_client
_new_client = (
OpenAI(
workload_identity=workload_identity_config.to_sdk_workload_identity(),
base_url=api_base,
http_client=sync_http_client,
timeout=timeout,
max_retries=max_retries,
organization=organization,
)
if workload_identity_config is not None
else OpenAI(
api_key=api_key,
base_url=api_base,
http_client=sync_http_client,
timeout=timeout,
max_retries=max_retries,
organization=organization,
)
)
## SAVE CACHE KEY

View file

@ -4,7 +4,100 @@ OpenAI Responses API token counting transformation logic.
This module handles the transformation of requests to OpenAI's /v1/responses/input_tokens endpoint.
"""
from typing import Any, Final
from collections.abc import Mapping, Sequence
from typing import Any, Final, Literal
from typing_extensions import ReadOnly, TypedDict
class ResponsesInputTextPart(TypedDict):
type: ReadOnly[Literal["input_text"]]
text: ReadOnly[str]
class ResponsesInputImagePart(TypedDict):
type: ReadOnly[Literal["input_image"]]
image_url: ReadOnly[str]
detail: ReadOnly[str]
class ResponsesInputFilePart(TypedDict):
type: ReadOnly[Literal["input_file"]]
filename: ReadOnly[str]
file_data: ReadOnly[str]
ResponsesInputPart = ResponsesInputTextPart | ResponsesInputImagePart | ResponsesInputFilePart
ResponsesContentRole = Literal["user", "assistant"]
def _chat_image_block_to_responses_part(image_url: object) -> ResponsesInputImagePart | None:
url: Final = image_url.get("url") if isinstance(image_url, Mapping) else image_url
if not isinstance(url, str) or not url:
return None
detail: Final = image_url.get("detail") if isinstance(image_url, Mapping) else None
part: Final[ResponsesInputImagePart] = {
"type": "input_image",
"image_url": url,
"detail": detail if isinstance(detail, str) and detail else "auto",
}
return part
def _chat_file_block_to_responses_part(file_value: object) -> ResponsesInputFilePart | None:
"""Only an inline file round trips: OpenAI rejects `file_data` without the `filename` beside it."""
if not isinstance(file_value, Mapping):
return None
filename: Final = file_value.get("filename")
file_data: Final = file_value.get("file_data")
if not isinstance(filename, str) or not filename or not isinstance(file_data, str) or not file_data:
return None
part: Final[ResponsesInputFilePart] = {
"type": "input_file",
"filename": filename,
"file_data": file_data,
}
return part
def _chat_block_to_responses_part(block: object, role: ResponsesContentRole) -> ResponsesInputPart | None:
if isinstance(block, str):
bare: Final[ResponsesInputTextPart] = {"type": "input_text", "text": block}
return bare
if not isinstance(block, Mapping):
return None
match block.get("type"):
case "text":
text_value: Final = block.get("text")
text: Final[ResponsesInputTextPart] = {
"type": "input_text",
"text": text_value if isinstance(text_value, str) else "",
}
return text
case "image_url" if role == "user":
return _chat_image_block_to_responses_part(block.get("image_url"))
case "file" if role == "user":
return _chat_file_block_to_responses_part(block.get("file"))
case _:
return None
def chat_content_blocks_to_responses_content(
content: Sequence[object],
role: ResponsesContentRole,
) -> str | tuple[ResponsesInputPart, ...]:
"""Text-only content collapses to a joined string, which every role accepts and counts identically.
Only a user turn may carry an image or file part: the Responses API rejects any part but
output_text and refusal inside an assistant turn.
"""
parts: Final = tuple(
part for part in (_chat_block_to_responses_part(block, role) for block in content) if part is not None
)
if any(part["type"] != "input_text" for part in parts):
return parts
return "\n".join(part["text"] for part in parts if part["type"] == "input_text")
class OpenAICountTokensConfig:
@ -120,18 +213,13 @@ class OpenAICountTokensConfig:
instructions_parts.append("\n".join(text_parts))
elif role == "user":
if isinstance(content, list):
# Extract text from content blocks for Responses API
text_parts = []
for block in content:
if isinstance(block, dict) and block.get("type") == "text":
text_parts.append(block.get("text", ""))
elif isinstance(block, str):
text_parts.append(block)
content = "\n".join(text_parts)
content = chat_content_blocks_to_responses_content(content, "user")
input_items.append({"role": "user", "content": content})
elif role == "assistant":
# Map tool_calls to Responses API function_call items
tool_calls = msg.get("tool_calls")
if isinstance(content, list):
content = chat_content_blocks_to_responses_content(content, "assistant")
if content:
input_items.append({"role": "assistant", "content": content})
if tool_calls:

View file

@ -21,6 +21,7 @@ from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
from ..common_utils import OpenAIError
from ..workload_identity import get_workload_identity_bearer_token, resolve_openai_workload_identity_config
OPENAI_RESPONSES_API_MIN_MAX_OUTPUT_TOKENS: Final = 16
@ -392,6 +393,14 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
litellm_params = litellm_params or GenericLiteLLMParams()
api_key = litellm_params.api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY")
headers.setdefault("Content-Type", "application/json")
workload_identity_config: Final = (
resolve_openai_workload_identity_config(api_key=api_key, api_base=litellm_params.api_base)
if self.custom_llm_provider is LlmProviders.OPENAI
else None
)
if workload_identity_config is not None:
headers["Authorization"] = f"Bearer {get_workload_identity_bearer_token(workload_identity_config)}"
return headers
headers["Authorization"] = f"Bearer {api_key}"
return headers

View file

@ -0,0 +1,100 @@
from __future__ import annotations
from dataclasses import dataclass
from functools import lru_cache
from typing import TYPE_CHECKING, Final
from urllib.parse import urlparse
import litellm
from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str
from .common_utils import OpenAIError
if TYPE_CHECKING:
from collections.abc import Callable
from openai.auth import SubjectTokenProvider, WorkloadIdentity, WorkloadIdentityAuth
OPENAI_WIF_CLIENT_ID: Final = "litellm"
_OPENAI_API_HOST: Final = "api.openai.com"
_SDK_UPGRADE_MESSAGE: Final = (
"OpenAI workload identity federation requires openai>=2.32.0. "
"Upgrade the installed openai package to use OPENAI_IDENTITY_PROVIDER_ID / "
"OPENAI_SERVICE_ACCOUNT_ID / OPENAI_IDENTITY_TOKEN_FILE."
)
@dataclass(frozen=True, slots=True)
class OpenAIWorkloadIdentityConfig:
identity_provider_id: str
service_account_id: str
token_file: str
def to_sdk_workload_identity(self) -> WorkloadIdentity:
k8s_token_provider: Final = _load_sdk_k8s_token_provider()
workload_identity: Final[WorkloadIdentity] = {
"client_id": OPENAI_WIF_CLIENT_ID,
"identity_provider_id": self.identity_provider_id,
"service_account_id": self.service_account_id,
"provider": k8s_token_provider(self.token_file),
}
return workload_identity
def resolve_openai_workload_identity_config(
api_key: str | None,
api_base: str | None,
) -> OpenAIWorkloadIdentityConfig | None:
static_api_key: Final = normalize_nonempty_secret_str(api_key) or normalize_nonempty_secret_str(
get_secret_str("OPENAI_API_KEY")
)
if static_api_key is not None:
return None
effective_api_base: Final = (
api_base or litellm.api_base or get_secret_str("OPENAI_BASE_URL") or get_secret_str("OPENAI_API_BASE")
)
if not _targets_openai_api(effective_api_base):
return None
identity_provider_id: Final = get_secret_str("OPENAI_IDENTITY_PROVIDER_ID")
service_account_id: Final = get_secret_str("OPENAI_SERVICE_ACCOUNT_ID")
token_file: Final = get_secret_str("OPENAI_IDENTITY_TOKEN_FILE")
if not identity_provider_id or not service_account_id or not token_file:
return None
return OpenAIWorkloadIdentityConfig(
identity_provider_id=identity_provider_id,
service_account_id=service_account_id,
token_file=token_file,
)
def get_workload_identity_bearer_token(config: OpenAIWorkloadIdentityConfig) -> str:
return _workload_identity_auth(config).get_token()
def _targets_openai_api(api_base: str | None) -> bool:
if api_base is None:
return True
parsed: Final = urlparse(api_base)
return parsed.scheme == "https" and parsed.hostname == _OPENAI_API_HOST
@lru_cache(maxsize=16)
def _workload_identity_auth(config: OpenAIWorkloadIdentityConfig) -> WorkloadIdentityAuth:
sdk_workload_identity_auth: Final = _load_sdk_workload_identity_auth()
return sdk_workload_identity_auth(workload_identity=config.to_sdk_workload_identity())
def _load_sdk_workload_identity_auth() -> type[WorkloadIdentityAuth]:
try:
from openai.auth import WorkloadIdentityAuth as sdk_workload_identity_auth
except ImportError as e:
raise OpenAIError(status_code=500, message=_SDK_UPGRADE_MESSAGE) from e
return sdk_workload_identity_auth
def _load_sdk_k8s_token_provider() -> Callable[[str], SubjectTokenProvider]:
try:
from openai.auth import k8s_service_account_token_provider
except ImportError as e:
raise OpenAIError(status_code=500, message=_SDK_UPGRADE_MESSAGE) from e
return k8s_service_account_token_provider

View file

@ -160,14 +160,12 @@ class RunwayMLVideoConfig(BaseVideoConfig):
**self._prompt_image_param(video_create_optional_params),
**self._ratio_param(video_create_optional_params),
**self._duration_param(video_create_optional_params),
# Pass through other parameters that aren't OpenAI-specific
**{key: value for key, value in video_create_optional_params.items() if key not in supported_openai_params},
}
@staticmethod
def _prompt_image_param(video_create_optional_params: VideoCreateOptionalRequestParams) -> Mapping[str, object]:
# Handle input_reference parameter - map to promptImage
# RunwayML supports URLs and data URIs directly
if "input_reference" in video_create_optional_params:
return {"promptImage": video_create_optional_params["input_reference"]}
return {}

View file

@ -182,7 +182,6 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM):
else None
)
# Generation config with proper structure for image editing
generation_config: Final[dict[str, object]] = {
key: value for key, value in (("response_modalities", ["IMAGE"]), ("image_config", image_config)) if value
}

View file

@ -203,7 +203,6 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
if value is not None
}
# Build the request body for Vertex AI RAG API
query_body: Final[Mapping[str, object]] = {
key: value
for key, value in (("text", query), ("rag_retrieval_config", rag_retrieval_config or None))
@ -294,7 +293,6 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
# Add metadata if provided
metadata: Final = vector_store_create_optional_params.get("metadata")
# Build the request body for Vertex AI RAG Corpus creation
request_body: Final[dict[str, object]] = {
key: value
for key, value in (

View file

@ -27,6 +27,15 @@ from .common_utils import (
get_vertex_base_url,
)
def _graft_default_vertex_path(api_base: str, default_url: str) -> str:
parsed_api_base: Final = urlparse(api_base)
default_segments: Final = urlparse(default_url).path.lstrip("/").split("/")
graft_segments: Final = default_segments[1:] if default_segments[0] in ("v1", "v1beta1") else default_segments
grafted_path: Final = parsed_api_base.path.rstrip("/") + "/" + "/".join(graft_segments)
return parsed_api_base._replace(path=grafted_path).geturl()
GOOGLE_IMPORT_ERROR_MESSAGE: Final = (
"Google Cloud SDK not found. Install it with: pip install 'litellm[google]' or pip install google-cloud-aiplatform"
)
@ -621,8 +630,9 @@ class VertexBase:
Handles custom api_base for:
1. Gemini (Google AI Studio) - constructs /models/{model}:{endpoint}
2. Vertex AI with standard proxies - constructs {api_base}:{endpoint};
if api_base has no path (bare host), grafts the default vertex URL path onto it
2. Vertex AI with standard proxies - grafts the default vertex URL path onto the
api_base when its path is empty or only an API version (/v1, /v1beta1);
otherwise constructs {api_base}:{endpoint}
3. Vertex AI with PSC endpoints - constructs full path structure
{api_base}/v1/projects/{project}/locations/{location}/endpoints/{model}:{endpoint}
(only when use_psc_endpoint_format=True)
@ -669,10 +679,14 @@ class VertexBase:
)
elif urlparse(api_base).path in ("", "/"):
url = api_base.rstrip("/") + urlparse(url).path
elif urlparse(api_base).path.rstrip("/") in ("/v1", "/v1beta1") and "/projects/" in urlparse(url).path:
url = _graft_default_vertex_path(api_base=api_base, default_url=url)
else:
url = f"{api_base}:{endpoint}"
if stream is True:
url = url + "?alt=sse"
parsed_stream_url: Final = urlparse(url)
stream_query: Final = f"{parsed_stream_url.query}&alt=sse" if parsed_stream_url.query else "alt=sse"
url = parsed_stream_url._replace(query=stream_query).geturl()
return auth_header, url
def _get_token_and_url(

View file

@ -5507,6 +5507,9 @@ def completion(
tpm=kwargs.get("tpm"),
rpm=kwargs.get("rpm"),
use_xai_oauth=kwargs.get("use_xai_oauth", False),
gigachat_scope=kwargs.get("gigachat_scope"),
gigachat_auth_url=kwargs.get("gigachat_auth_url"),
gigachat_access_token=kwargs.get("gigachat_access_token"),
**{key: kwargs[key] for key in FORWARDED_KWARGS_KEYS if key in kwargs},
)
cast(LiteLLMLoggingObj, logging).update_environment_variables(

File diff suppressed because it is too large Load diff

View file

@ -29,6 +29,8 @@ class LiteLLM_ProxyModelTable(LiteLLMPydanticObjectBase):
@model_validator(mode="before")
@classmethod
def check_potential_json_str(cls, values):
if not isinstance(values, dict):
return values
if isinstance(values.get("litellm_params"), str):
try:
values["litellm_params"] = json.loads(values["litellm_params"])

View file

@ -2,17 +2,22 @@
This module is used to pass through requests to the LLM APIs.
"""
from __future__ import annotations
import asyncio
import contextvars
from collections.abc import AsyncGenerator, Coroutine, Generator
from collections.abc import AsyncGenerator, AsyncIterator, Coroutine, Generator, Iterator
from functools import partial
from typing import TYPE_CHECKING, Any, Final, Optional, cast
from types import TracebackType
from typing import Any, Final, cast
import httpx
from httpx._types import CookieTypes, QueryParamTypes, RequestFiles
from httpx._types import CookieTypes, QueryParamTypes, RequestContent, RequestFiles
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.passthrough.utils import CommonUtils
@ -21,9 +26,222 @@ from litellm.utils import client
base_llm_http_handler = BaseLLMHTTPHandler()
from .utils import BasePassthroughUtils
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
async def _as_async_generator(iterable: AsyncIterator[bytes]) -> AsyncGenerator[bytes, Any]:
async for chunk in iterable:
yield chunk
def _as_generator(iterable: Iterator[bytes]) -> Generator[bytes, Any, Any]:
yield from iterable
class AsyncPassthroughStreamingResponse(AsyncGenerator[Any, Any]):
def __init__(
self,
response: Coroutine[Any, Any, httpx.Response],
litellm_logging_obj: LiteLLMLoggingObj,
provider_config: BasePassthroughConfig,
) -> None:
self._initialized = False
self._status_code: int = 0
self._headers = httpx.Headers()
self._response_coro = response
self._response: httpx.Response
self._iterator: AsyncGenerator[bytes, Any]
self._litellm_logging_obj = litellm_logging_obj
self._provider_config = provider_config
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
self._flush_scheduled = False
self._background_tasks: set[asyncio.Task] = set() # mutable-ok: instance set for background task tracking
self._hidden_params: dict[str, object] = {} # mutable-ok: router attaches response headers here in place
@property
def status_code(self) -> int:
if not self._initialized:
raise RuntimeError("AsyncPassthroughStreamingResponse must be awaited before accessing status_code")
return self._status_code
@status_code.setter
def status_code(self, value: int) -> None:
self._status_code = value
@property
def headers(self) -> httpx.Headers:
if not self._initialized:
raise RuntimeError("AsyncPassthroughStreamingResponse must be awaited before accessing headers")
return self._headers
@headers.setter
def headers(self, value: httpx.Headers) -> None:
self._headers = value
def __await__(self) -> Iterator[Any]:
async def _init():
if not self._initialized:
self._response = await self._response_coro
self.headers = self._response.headers
self.status_code = self._response.status_code
self._initialized = True
try:
self._response.raise_for_status()
self._iterator = _as_async_generator(self._response.aiter_bytes())
except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic
try:
await self._response.aread()
except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic
pass
try:
await self._response.aclose()
except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic
pass
raise
return self
return _init().__await__()
def _start_flush(self) -> None:
if self._flush_scheduled or not self._raw_bytes:
return
self._flush_scheduled = True
try:
task: Final = asyncio.create_task(
self._litellm_logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=self._raw_bytes,
provider_config=self._provider_config,
)
)
# Compliant: Save a strong reference to prevent GC
self._background_tasks.add(task)
# Remove the task from the set when it finishes to avoid memory leaks
task.add_done_callback(self._background_tasks.discard)
except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush; %d buffered chunks dropped: %s",
len(self._raw_bytes),
e,
)
def __aiter__(self) -> AsyncPassthroughStreamingResponse:
return self
def aiter_bytes(self) -> AsyncPassthroughStreamingResponse:
return self
async def __anext__(self) -> bytes:
if not self._initialized:
await self # pyright: ignore[reportGeneralTypeIssues] # structural type check misses __await__
try:
chunk: Final = await anext(self._iterator)
self._raw_bytes.append(chunk)
except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic
self._start_flush()
try:
await self._response.aclose()
except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic
pass
raise
else:
return chunk
async def asend(self, value: bytes) -> bytes:
if not self._initialized:
await self # pyright: ignore[reportGeneralTypeIssues] # structural type check misses __await__
return await self._iterator.asend(value)
async def athrow(
self,
typ: BaseException | type[BaseException],
val: BaseException | object = None,
tb: TracebackType | None = None,
) -> bytes:
if not self._initialized:
await self # pyright: ignore[reportGeneralTypeIssues] # structural type check misses __await__
return await self._iterator.athrow(typ, val, tb) # pyright: ignore[reportCallIssue, reportArgumentType] # matches one of the athrow overloads
async def aclose(self) -> None:
self._start_flush()
try:
if self._initialized:
await self._iterator.aclose()
await self._response.aclose()
except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic
pass
class PassthroughStreamingResponse(Generator[Any, Any, Any]):
def __init__(
self,
response: httpx.Response,
litellm_logging_obj: LiteLLMLoggingObj,
provider_config: BasePassthroughConfig,
) -> None:
self._response = response
self.headers = response.headers
self.status_code = response.status_code
self._litellm_logging_obj = litellm_logging_obj
self._provider_config = provider_config
self._iterator: Generator[bytes, Any, Any] = _as_generator(response.iter_bytes())
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
self._flush_scheduled = False
def _start_flush(self) -> None:
if self._flush_scheduled or not self._raw_bytes:
return
self._flush_scheduled = True
from litellm.utils import executor
try:
executor.submit(
self._litellm_logging_obj.flush_passthrough_collected_chunks,
raw_bytes=self._raw_bytes,
provider_config=self._provider_config,
)
except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush; %d buffered chunks dropped: %s",
len(self._raw_bytes),
e,
)
def __iter__(self) -> PassthroughStreamingResponse:
return self
def __next__(self) -> bytes:
try:
chunk: Final = next(self._iterator)
self._raw_bytes.append(chunk)
except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic
self._start_flush()
try:
self._response.close()
except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic
pass
raise
else:
return chunk
def send(self, value: bytes) -> bytes:
return self._iterator.send(value)
def throw(
self,
typ: BaseException | type[BaseException],
val: BaseException | object = None,
tb: TracebackType | None = None,
) -> bytes:
return self._iterator.throw(typ, val, tb) # pyright: ignore[reportCallIssue, reportArgumentType] # matches one of the throw overloads
def close(self) -> None:
self._start_flush()
try:
self._response.close()
except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic
pass
@client
@ -37,10 +255,10 @@ async def allm_passthrough_route(
api_key: str | None = None,
request_query_params: dict | None = None,
request_headers: dict | None = None,
content: Any | None = None,
content: RequestContent | None = None,
data: dict | None = None,
files: RequestFiles | None = None,
json: Any | None = None,
json: object | None = None,
params: QueryParamTypes | None = None,
cookies: CookieTypes | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
@ -64,7 +282,7 @@ async def allm_passthrough_route(
from litellm.utils import ProviderConfigManager
provider_config = cast(
Optional["BasePassthroughConfig"], kwargs.get("provider_config")
BasePassthroughConfig | None, kwargs.get("provider_config")
) or ProviderConfigManager.get_provider_passthrough_config(
provider=LlmProviders(custom_llm_provider),
model=model,
@ -132,12 +350,12 @@ async def allm_passthrough_route(
if resolved_custom_llm_provider:
try:
provider_config = cast(
Optional["BasePassthroughConfig"], kwargs.get("provider_config")
BasePassthroughConfig | None, kwargs.get("provider_config")
) or ProviderConfigManager.get_provider_passthrough_config(
provider=LlmProviders(resolved_custom_llm_provider),
model=model,
)
except Exception:
except Exception: # noqa: BLE001 S110
# If we can't get provider config, pass None
pass
@ -162,10 +380,10 @@ def llm_passthrough_route(
api_key: str | None = None,
request_query_params: dict | None = None,
request_headers: dict | None = None,
content: Any | None = None,
content: RequestContent | None = None,
data: dict | None = None,
files: RequestFiles | None = None,
json: Any | None = None,
json: object | None = None,
params: QueryParamTypes | None = None,
cookies: CookieTypes | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
@ -190,7 +408,9 @@ def llm_passthrough_route(
_is_async: Final = bool(kwargs.get("allm_passthrough_route", False))
litellm_logging_obj: Final = cast("LiteLLMLoggingObj", kwargs.get("litellm_logging_obj"))
litellm_logging_obj: Final = cast(
LiteLLMLoggingObj, kwargs.get("litellm_logging_obj")
) # cast-ok: logging obj is constructed upstream; tests inject mocks
model, custom_llm_provider, api_key, api_base = get_llm_provider(
model=model,
@ -235,7 +455,7 @@ def llm_passthrough_route(
)
provider_config: Final = cast(
Optional["BasePassthroughConfig"], kwargs.get("provider_config")
BasePassthroughConfig | None, kwargs.get("provider_config")
) or ProviderConfigManager.get_provider_passthrough_config(
provider=LlmProviders(custom_llm_provider),
model=model,
@ -276,10 +496,13 @@ def llm_passthrough_route(
forward_headers=False,
)
_request_data: dict | None = (
data if isinstance(data, dict) else (json if isinstance(json, dict) else None)
) # rebind-ok: conditional
headers, signed_json_body = provider_config.sign_request(
headers=headers,
litellm_params=litellm_params_dict,
request_data=data if data else json,
request_data=_request_data,
api_base=str(updated_url),
model=model,
)
@ -301,9 +524,12 @@ def llm_passthrough_route(
)
## IS STREAMING REQUEST
_streaming_request_data: dict = (
data if isinstance(data, dict) else (json if isinstance(json, dict) else {})
) # rebind-ok: conditional
is_streaming_request: Final = provider_config.is_streaming_request(
endpoint=endpoint,
request_data=data or json or {},
request_data=_streaming_request_data,
)
# Update logging object with streaming status
@ -334,18 +560,26 @@ def llm_passthrough_route(
else:
# Sync path - client.client.send returns Response directly
response: httpx.Response = client.client.send(request=request, stream=is_streaming_request)
response.raise_for_status()
try:
response.raise_for_status()
except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic
try:
response.read()
except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic
pass
try:
response.close()
except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic
pass
raise
if (
hasattr(response, "iter_bytes") and is_streaming_request
): # yield the chunk, so we can store it in the logging object
return _sync_streaming(response, litellm_logging_obj, provider_config)
if hasattr(response, "iter_bytes") and is_streaming_request:
return PassthroughStreamingResponse(response, litellm_logging_obj, provider_config)
else:
# For non-streaming responses, yield the entire response
return response
except Exception as e:
if provider_config is None:
raise e
# provider_config is guaranteed non-None here due to the earlier guard
assert provider_config is not None
raise base_llm_http_handler._handle_error(
e=e,
provider_config=provider_config,
@ -356,8 +590,8 @@ async def _async_passthrough_request(
client: HTTPHandler | AsyncHTTPHandler,
request: httpx.Request,
is_streaming_request: bool,
litellm_logging_obj: "LiteLLMLoggingObj",
provider_config: "BasePassthroughConfig",
litellm_logging_obj: LiteLLMLoggingObj,
provider_config: BasePassthroughConfig,
) -> httpx.Response | AsyncGenerator[Any, Any]:
"""
Handle async passthrough requests.
@ -369,8 +603,7 @@ async def _async_passthrough_request(
# Check if it's a coroutine and await it
if asyncio.iscoroutine(response_result):
if is_streaming_request:
# Pass the coroutine to _async_streaming which will await it
return _async_streaming(
return await AsyncPassthroughStreamingResponse( # pyright: ignore[reportGeneralTypeIssues] # structural type check misses __await__
response=response_result,
litellm_logging_obj=litellm_logging_obj,
provider_config=provider_config,
@ -383,84 +616,3 @@ async def _async_passthrough_request(
else:
# Fallback for sync-like behavior (shouldn't happen in async path)
raise Exception("Expected coroutine from async client")
def _sync_streaming(
response: httpx.Response,
litellm_logging_obj: "LiteLLMLoggingObj",
provider_config: "BasePassthroughConfig",
):
from litellm.utils import executor
raw_bytes: Final[list[bytes]] = []
flush_scheduled = False
try:
for chunk in response.iter_bytes():
raw_bytes.append(chunk)
yield chunk
finally:
if not flush_scheduled and raw_bytes:
flush_scheduled = True
try:
executor.submit(
litellm_logging_obj.flush_passthrough_collected_chunks,
raw_bytes=raw_bytes,
provider_config=provider_config,
)
except Exception as e:
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush "
"in _sync_streaming; %d buffered chunks dropped: %s",
len(raw_bytes),
e,
)
async def _async_streaming(
response: Coroutine[Any, Any, httpx.Response],
litellm_logging_obj: "LiteLLMLoggingObj",
provider_config: "BasePassthroughConfig",
):
iter_response: Final = await response
try:
iter_response.raise_for_status()
except Exception:
try:
await iter_response.aclose()
except Exception:
pass
raise
raw_bytes: Final[list[bytes]] = []
flush_scheduled = False
try:
async for chunk in iter_response.aiter_bytes():
raw_bytes.append(chunk)
yield chunk
except Exception:
try:
await iter_response.aclose()
except Exception:
pass
raise
finally:
# GeneratorExit (raised on client disconnect) is not caught by
# `except Exception`; the finally block ensures partial usage
# still gets flushed for spend tracking. See LIT-2642.
if not flush_scheduled and raw_bytes:
flush_scheduled = True
try:
asyncio.create_task(
litellm_logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=raw_bytes,
provider_config=provider_config,
)
)
except Exception as e:
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush "
"in _async_streaming; %d buffered chunks dropped: %s",
len(raw_bytes),
e,
)

View file

@ -422,6 +422,9 @@ class LiteLLMRoutes(enum.Enum):
"/responses/{response_id}/cancel",
"/v1/responses/{response_id}/cancel",
"/openai/v1/responses/{response_id}/cancel",
"/responses/input_tokens",
"/v1/responses/input_tokens",
"/openai/v1/responses/input_tokens",
# vector stores
"/vector_stores",
"/v1/vector_stores",
@ -471,6 +474,7 @@ class LiteLLMRoutes(enum.Enum):
"/vllm",
"/mistral",
"/milvus",
"/gigachat",
"/watsonx",
]
@ -3549,6 +3553,19 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
)
class SpendLogsRouterMetadata(TypedDict):
"""
Router provenance stamped on spend logs for deployments flagged with
model_info.internal_router_model, correlating the requested model group
with the provider deployment that served the call
"""
requested_model: ReadOnly[str | None]
selected_model: ReadOnly[str | None]
selected_provider: ReadOnly[str | None]
router_correlation_id: ReadOnly[str | None]
class SpendLogsMetadata(TypedDict):
"""
Specific metadata k,v pairs logged to spendlogs for easier cost tracking
@ -3591,6 +3608,7 @@ class SpendLogsMetadata(TypedDict):
compression_savings: CompressionSavingsMetadata | None
autorouter_savings: ReadOnly[float | None] # stamped by the logging payload; None = not auto-routed
litellm_gateway_injected_cache: ReadOnly[str | None]
router_metadata: ReadOnly[SpendLogsRouterMetadata | None] # None = deployment not flagged internal_router_model
class SpendLogsPayload(TypedDict):

View file

@ -2,13 +2,14 @@
Handles Authentication Errors
"""
import logging
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
from fastapi import HTTPException, Request, status
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._logging import verbose_proxy_logger, verbose_proxy_stdout_logger
from litellm.constants import EMPTY_MAPPING
from litellm.integrations.otel.runtime import seed_request_identity
from litellm.litellm_core_utils.core_helpers import is_expected_client_error
@ -18,7 +19,11 @@ from litellm.proxy._types import (
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_utils import _get_request_ip_address
from litellm.proxy.auth.auth_utils import (
_get_request_ip_address,
is_invalid_virtual_key_error,
mark_invalid_virtual_key_error,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.types.services import ServiceTypes
@ -36,6 +41,41 @@ else:
Span = Any
def _as_proxy_exception(e: Exception) -> ProxyException:
"""Convert an authentication failure into the ProxyException the client receives."""
if isinstance(e, litellm.BudgetExceededError):
return ProxyException(
message=e.message,
type=ProxyErrorTypes.budget_exceeded,
param=None,
code=getattr(e, "status_code", status.HTTP_429_TOO_MANY_REQUESTS),
)
if isinstance(e, HTTPException):
return ProxyException(
message=getattr(e, "detail", f"Authentication Error({e})"),
type=ProxyErrorTypes.auth_error,
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_401_UNAUTHORIZED),
)
if isinstance(e, ProxyException):
return e
if PrismaDBExceptionHandler.is_database_service_unavailable_error(e):
return ProxyException(
message=(
"Service Unavailable, the authentication database is temporarily unreachable. Please retry shortly."
),
type=ProxyErrorTypes.no_db_connection,
param="None",
code=status.HTTP_503_SERVICE_UNAVAILABLE,
)
return ProxyException(
message="Authentication Error, " + str(e),
type=ProxyErrorTypes.auth_error,
param=getattr(e, "param", "None"),
code=status.HTTP_401_UNAUTHORIZED,
)
def _with_requester_ip_address(request_data: dict[str, object], requester_ip: str | None) -> dict[str, object]:
"""Auth gate rejections are raised before `add_litellm_data_to_request` records the
caller IP, so their failure logs would otherwise carry no IP nor key/user identity."""
@ -110,16 +150,21 @@ class UserAPIKeyAuthExceptionHandler:
request=request,
use_x_forwarded_for=general_settings.get("use_x_forwarded_for") is True,
)
log_fn: Final = (
verbose_proxy_logger.error
if is_expected_client_error(e) and not litellm.log_client_error_tracebacks
else verbose_proxy_logger.exception
)
log_fn(
# Log authentication failures before identity seeding and callbacks, so the log
# survives a raising callback pipeline. Classify and route malformed virtual-key
# rejections to WARNING on stdout (suppressible via LITELLM_LOG=ERROR).
log_extra: Final = {"requester_ip": requester_ip}
is_invalid_virtual_key: Final = is_invalid_virtual_key_error(e)
is_quiet_log: Final = is_invalid_virtual_key and not litellm.log_client_error_tracebacks
logger: Final = verbose_proxy_stdout_logger if is_quiet_log else verbose_proxy_logger
logger.log(
logging.WARNING if is_quiet_log else logging.ERROR,
"litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - %s\nRequester IP Address:%s",
e,
requester_ip,
extra={"requester_ip": requester_ip},
exc_info=True if litellm.log_client_error_tracebacks or not is_expected_client_error(e) else None,
extra=log_extra,
)
# Log this exception to OTEL, Datadog etc. Reuse the identity resolved
@ -167,35 +212,13 @@ class UserAPIKeyAuthExceptionHandler:
if transformed_exception is not None:
e = transformed_exception
if isinstance(e, litellm.BudgetExceededError):
raise ProxyException(
message=e.message,
type=ProxyErrorTypes.budget_exceeded,
param=None,
code=getattr(e, "status_code", status.HTTP_429_TOO_MANY_REQUESTS),
final_exception: Final = mark_invalid_virtual_key_error(_as_proxy_exception(e), is_invalid_virtual_key)
# If a quiet-logged malformed-key transform yields non-401, escalate to ERROR
if is_quiet_log and str(final_exception.code) != str(status.HTTP_401_UNAUTHORIZED):
verbose_proxy_logger.error(
"litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - %s\nRequester IP Address:%s",
final_exception,
requester_ip,
extra=log_extra,
)
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "detail", f"Authentication Error({e})"),
type=ProxyErrorTypes.auth_error,
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_401_UNAUTHORIZED),
)
elif isinstance(e, ProxyException):
raise e
if PrismaDBExceptionHandler.is_database_service_unavailable_error(e):
raise ProxyException(
message=(
"Service Unavailable, the authentication database is "
"temporarily unreachable. Please retry shortly."
),
type=ProxyErrorTypes.no_db_connection,
param="None",
code=status.HTTP_503_SERVICE_UNAVAILABLE,
)
raise ProxyException(
message="Authentication Error, " + str(e),
type=ProxyErrorTypes.auth_error,
param=getattr(e, "param", "None"),
code=status.HTTP_401_UNAUTHORIZED,
)
raise final_exception

View file

@ -15,6 +15,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import (
BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY,
EMPTY_MAPPING,
INVALID_VIRTUAL_KEY_ERROR_MARKER,
MINIMUM_CUSTOM_KEY_LENGTH,
STANDARD_CUSTOMER_ID_HEADERS,
)
@ -34,6 +35,43 @@ from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS
from litellm.types.utils import CustomPricingLiteLLMParams
def is_invalid_virtual_key_error(exception: BaseException | None) -> bool:
"""True when an authentication error rejects a malformed virtual key.
Classifies only by the marker stamped where that 401 is raised. Message
content is never inspected: other 401s interpolate caller-supplied values
(vector store ids, organization ids) into their messages, so a phrase
match would let a request body demote an authorization failure to the
quiet log path.
"""
if not isinstance(exception, (HTTPException, ProxyException)):
return False
code: Final[object] = getattr(exception, "code", None)
status_code: Final[object] = code if code is not None else getattr(exception, "status_code", None)
if str(status_code) != str(status.HTTP_401_UNAUTHORIZED):
return False
return getattr(exception, INVALID_VIRTUAL_KEY_ERROR_MARKER, False) is True
def mark_invalid_virtual_key_error(exception: ProxyException, is_invalid_virtual_key: bool) -> ProxyException:
"""Return an independently marked malformed-key exception after callback transformations."""
if not is_invalid_virtual_key or str(exception.code) != str(status.HTTP_401_UNAUTHORIZED):
return exception
marked_exception: Final = ProxyException(
message=exception.message,
type=exception.type,
param=exception.param,
code=exception.code,
headers=exception.headers.copy(),
openai_code=None if exception.openai_code is None else str(exception.openai_code),
provider_specific_fields=exception.provider_specific_fields,
)
setattr(marked_exception, INVALID_VIRTUAL_KEY_ERROR_MARKER, True)
return marked_exception
def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None:
client_ip = None
if use_x_forwarded_for is True and "x-forwarded-for" in request.headers:

View file

@ -19,12 +19,15 @@ import fastapi
import orjson
from fastapi import HTTPException, Request, WebSocket, status
from fastapi.security.api_key import APIKeyHeader
from starlette.exceptions import WebSocketException
import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._service_logger import ServiceLogging
from litellm.constants import (
GLOBAL_PROXY_SPEND_CACHE_KEY,
INVALID_VIRTUAL_KEY_ERROR_MARKER,
INVALID_VIRTUAL_KEY_ERROR_MESSAGE,
LITELLM_PROXY_BUDGET_NAME,
LITELLM_PROXY_MASTER_KEY_ALIAS,
)
@ -65,6 +68,7 @@ from litellm.proxy.auth.auth_utils import (
get_model_from_request,
get_request_route,
get_request_route_template,
is_invalid_virtual_key_error,
iter_request_fallback_targets,
normalize_request_route,
pre_db_read_auth_checks,
@ -539,6 +543,8 @@ async def user_api_key_auth_websocket(websocket: WebSocket):
try:
return await user_api_key_auth(request=request, api_key=f"Bearer {api_key}")
except Exception as e:
if is_invalid_virtual_key_error(e):
raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION)
verbose_proxy_logger.exception(e)
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
raise HTTPException(status_code=403, detail=str(e))
@ -1867,13 +1873,17 @@ async def _user_api_key_auth_builder(
_masked_key: Final = f"{api_key[:4]}****{api_key[-4:]}" if len(api_key) > 8 else "****"
if not api_key.startswith("sk-"):
_hint = _JWT_AUTH_DISABLED_HINT if not enable_jwt_auth and JWTHandler.is_jwt(token=api_key) else ""
raise HTTPException(
_malformed_key_error = HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=(
f"LiteLLM Virtual Key expected. Received={_masked_key}, "
f"{INVALID_VIRTUAL_KEY_ERROR_MESSAGE}. Received={_masked_key}, "
f"expected to start with 'sk-'.{_hint}"
),
) # prevent token hashes from being used
# Stamp provenance here so log routing classifies this 401 by
# where it was raised, never by its message text.
setattr(_malformed_key_error, INVALID_VIRTUAL_KEY_ERROR_MARKER, True)
raise _malformed_key_error
else:
verbose_logger.warning(
"litellm.proxy.proxy_server.user_api_key_auth(): Warning - Key is not a string. Got type={}".format(

View file

@ -1495,6 +1495,35 @@ class ProxyBaseLLMRequestProcessing:
def __init__(self, data: dict):
self.data = data
@staticmethod
def _merge_passthrough_streaming_headers(
response_headers: httpx.Headers | dict | None,
custom_headers: dict,
) -> dict:
"""
Merge upstream passthrough headers with proxy/custom headers.
Proxy/custom headers win on key collisions.
"""
excluded_headers: Final = { # mutable-ok: set of header names to exclude from forwarding
"transfer-encoding",
"content-encoding",
"set-cookie",
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailer",
"upgrade",
}
merged_headers: Final = { # mutable-ok: dict comprehension for merged headers forwarded to httpx
key: value for key, value in dict(response_headers or {}).items() if key.lower() not in excluded_headers
}
merged_headers.update(custom_headers)
return merged_headers
@staticmethod
def get_custom_headers(
*,
@ -2389,6 +2418,16 @@ class ProxyBaseLLMRequestProcessing:
)
if route_type == "allm_passthrough_route":
upstream_response_headers: Final = getattr(response, "headers", None)
streaming_headers: Final = (
ProxyBaseLLMRequestProcessing._merge_passthrough_streaming_headers(
response_headers=upstream_response_headers,
custom_headers=custom_headers,
)
if upstream_response_headers is not None
else custom_headers
)
# Check if response is an async generator
if self._is_streaming_response(response):
if asyncio.iscoroutine(response):
@ -2418,11 +2457,11 @@ class ProxyBaseLLMRequestProcessing:
# For passthrough routes, stream directly without error parsing
# since we're dealing with raw binary data (e.g., AWS event streams)
return StreamingResponse(
content=generator,
status_code=status.HTTP_200_OK,
return _UpstreamClosingStreamingResponse(
content=generator, # pyright: ignore[reportArgumentType] # generator-configured StreamingResponse
status_code=getattr(response, "status_code", status.HTTP_200_OK),
media_type=self._passthrough_event_stream_media_type(),
headers=custom_headers,
headers=streaming_headers,
)
else:
_early = await self._handle_non_streaming_allm_passthrough_route(
@ -2437,7 +2476,7 @@ class ProxyBaseLLMRequestProcessing:
return StreamingResponse(
content=response.aiter_bytes(),
status_code=response.status_code,
headers=custom_headers,
headers=streaming_headers,
)
elif route_type == "anthropic_messages":
# Check if response is actually a streaming response (async generator)

View file

@ -7,20 +7,15 @@ instead of aggregating LiteLLM_SpendLogs every time a window counter goes cold
(issue #35766). Raw SQL rather than the Prisma upsert helper because the
conditional roll cannot be expressed through the query builder.
Seeding a row that does not exist yet reads LiteLLM_SpendLogs once, excluding
the requests whose increments are in the same batch so neither source counts
them twice. One gap survives that exclusion: without the Redis transaction
buffer every pod flushes its own increments, so a row seeded by one pod can
include spend logs whose increments are still queued on another pod, and those
increments are added again when that pod flushes. That is bounded by a single
flush interval, happens at most once per window row, and only ever over-counts:
the seed never omits spend, because every increment not yet in the row still
reaches it on its own pod's next flush. A row therefore lags real spend by at
most one flush interval of queued increments, the same lag the SpendLogs
aggregate it replaces (and every other spend column) already has.
Seeding a row that does not exist yet reads LiteLLM_SpendLogs once and takes
off what the increments being flushed will add, so neither source counts the
same request twice. A row therefore lags real spend by at most one flush
interval of increments queued elsewhere: the same lag the SpendLogs aggregate
it replaces (and every other spend column) already has.
"""
from collections.abc import Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Final, Protocol
@ -67,33 +62,46 @@ _ROLL_WINDOW_SPEND_SQL: Final = (
)
_SEED_FROM_SPEND_LOGS_KEY_SQL: Final = (
'SELECT COALESCE(SUM(spend), 0.0) AS total FROM "LiteLLM_SpendLogs" '
"WHERE api_key = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC') "
"AND NOT (request_id = ANY($3::text[]) AND \"startTime\" >= ($4::timestamptz AT TIME ZONE 'UTC'))"
"SELECT COALESCE(SUM(spend), 0.0) AS total, "
"COALESCE(SUM(spend) FILTER (WHERE \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC')), 0.0) AS before_batch "
'FROM "LiteLLM_SpendLogs" '
"WHERE api_key = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
)
_SEED_FROM_SPEND_LOGS_TEAM_SQL: Final = (
'SELECT COALESCE(SUM(spend), 0.0) AS total FROM "LiteLLM_SpendLogs" '
"WHERE team_id = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC') "
"AND NOT (request_id = ANY($3::text[]) AND \"startTime\" >= ($4::timestamptz AT TIME ZONE 'UTC'))"
"SELECT COALESCE(SUM(spend), 0.0) AS total, "
"COALESCE(SUM(spend) FILTER (WHERE \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC')), 0.0) AS before_batch "
'FROM "LiteLLM_SpendLogs" '
"WHERE team_id = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
)
_SEED_FROM_SPEND_LOGS_KEY_UNBOUNDED_SQL: Final = (
'SELECT COALESCE(SUM(spend), 0.0) AS total FROM "LiteLLM_SpendLogs" '
"SELECT COALESCE(SUM(spend), 0.0) AS total, COALESCE(SUM(spend), 0.0) AS before_batch "
'FROM "LiteLLM_SpendLogs" '
"WHERE api_key = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
)
_SEED_FROM_SPEND_LOGS_TEAM_UNBOUNDED_SQL: Final = (
'SELECT COALESCE(SUM(spend), 0.0) AS total FROM "LiteLLM_SpendLogs" '
"SELECT COALESCE(SUM(spend), 0.0) AS total, COALESCE(SUM(spend), 0.0) AS before_batch "
'FROM "LiteLLM_SpendLogs" '
"WHERE team_id = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
)
_UPSERT_TRANSACTION_TIMEOUT: Final = timedelta(seconds=60)
@dataclass(frozen=True, slots=True)
class WindowSeedTotals:
"""The two sums a seed needs: everything persisted for the window, and the
part of it that predates the batch being flushed."""
total: float
before_batch: float
class WindowSpendLogsAggregate(Protocol):
"""Sums LiteLLM_SpendLogs for one entity since window_start, ignoring the
requests whose ids are handed in.
"""Sums LiteLLM_SpendLogs for one entity since window_start, split at the
batch's earliest request.
Injected so the flush can be exercised without a database and so the
expensive aggregate stays swappable.
@ -105,21 +113,19 @@ class WindowSpendLogsAggregate(Protocol):
entity_type: str,
entity_id: str,
window_start: datetime,
exclude_request_ids: Sequence[str],
exclude_started_at: datetime | None,
) -> float | None: ...
batch_started_at: datetime | None,
) -> WindowSeedTotals | None: ...
async def spend_logs_total_excluding(
async def spend_logs_seed_totals(
prisma_client: "PrismaClient",
entity_type: str,
entity_id: str,
window_start: datetime,
exclude_request_ids: Sequence[str],
exclude_started_at: datetime | None,
) -> float | None:
"""LiteLLM_SpendLogs spend for one entity since window_start, minus the
requests already accounted for by the increments being flushed.
batch_started_at: datetime | None,
) -> WindowSeedTotals | None:
"""LiteLLM_SpendLogs spend for one entity since window_start, both in full
and up to the start of the batch being flushed, in one scan.
The spend log writer drains its own queue on a ~2s poll whenever anything
is queued, while window increments flush on the much slower batch tick, so
@ -127,13 +133,12 @@ async def spend_logs_total_excluding(
already in the table. Counting them in the seed and again in the increment
is what made a fresh row land at twice the true spend.
The exclusion is bounded to rows that started at or after the batch's
earliest request. request_id can be chosen by the client
(x-litellm-call-id), so an unbounded exclusion would let a replayed old id
erase a historical row from the seed while its increment still lands.
Without a known start the batch's ids are not excluded at all: that can
only over-count once, which enforcement tolerates, whereas under-counting
is a budget bypass.
Both halves are needed because neither is safe alone: the full sum
double-counts this batch, and the sum before the batch drops spend another
pod has already persisted but not yet incremented. _seed_base picks between
them. Without a known batch start the two are the same sum, so the seed
counts everything: that can only over-count once, which enforcement
tolerates, whereas under-counting is a budget bypass.
"""
if entity_type == Litellm_EntityType.KEY.value:
bounded_sql, unbounded_sql = _SEED_FROM_SPEND_LOGS_KEY_SQL, _SEED_FROM_SPEND_LOGS_KEY_UNBOUNDED_SQL
@ -143,21 +148,23 @@ async def spend_logs_total_excluding(
return None
rows: Final = (
await prisma_client.db.query_raw(unbounded_sql, entity_id, window_start)
if exclude_started_at is None or not exclude_request_ids
if batch_started_at is None
else await prisma_client.db.query_raw(
bounded_sql,
entity_id,
window_start,
tuple(exclude_request_ids),
_exclusion_lower_bound(exclude_started_at),
_exclusion_upper_bound(batch_started_at),
)
)
if not rows:
return 0.0
return float(rows[0].get("total") or 0.0)
return WindowSeedTotals(total=0.0, before_batch=0.0)
return WindowSeedTotals(
total=float(rows[0].get("total") or 0.0),
before_batch=float(rows[0].get("before_batch") or 0.0),
)
def _exclusion_lower_bound(started_at: datetime) -> datetime:
def _exclusion_upper_bound(started_at: datetime) -> datetime:
"""LiteLLM_SpendLogs.startTime is TIMESTAMP(3); floor to the second so a
millisecond rounding of the batch's own earliest row cannot slip under it."""
return to_naive_utc(started_at).replace(microsecond=0)
@ -194,20 +201,33 @@ async def _seed_base_for_missing_row(
This is the LiteLLM_SpendLogs aggregate the window counter reseed runs on
every cold counter today, but here it runs once per window lifetime and off
the request path, and it excludes this batch's own requests so they are
counted by their increments alone.
the request path, and it discounts the queued increments so they are
counted once.
"""
if _primary_key(transaction) in existing_primary_keys:
return 0.0
base: Final = await spend_logs_aggregate(
totals: Final = await spend_logs_aggregate(
prisma_client=prisma_client,
entity_type=transaction["entity_type"],
entity_id=transaction["entity_id"],
window_start=datetime.fromisoformat(transaction["window_start"]).replace(tzinfo=timezone.utc),
exclude_request_ids=transaction["request_ids"],
exclude_started_at=_transaction_started_at(transaction),
batch_started_at=_transaction_started_at(transaction),
)
return float(base or 0.0)
if totals is None:
return 0.0
return _seed_base(totals=totals, batch_spend=transaction["spend"])
def _seed_base(totals: WindowSeedTotals, batch_spend: float) -> float:
"""What the window already held before the increments about to be applied.
Subtracting the batch's own spend from the full sum keeps every other
request in the seed, including the ones another pod persisted and has not
incremented yet, which a plain cutoff would drop for good if that pod died.
When this batch's own log rows have not landed yet the subtraction takes
spend that was never counted, so the sum before the batch is the floor.
"""
return max(totals.total - batch_spend, totals.before_batch)
def _transaction_started_at(transaction: WindowSpendTransaction) -> datetime | None:
@ -241,7 +261,7 @@ def _upsert_params(
async def commit_window_spend_updates(
prisma_client: "PrismaClient",
transactions: Sequence[WindowSpendTransaction],
spend_logs_aggregate: WindowSpendLogsAggregate = spend_logs_total_excluding,
spend_logs_aggregate: WindowSpendLogsAggregate = spend_logs_seed_totals,
) -> None:
"""Apply aggregated window increments to LiteLLM_BudgetWindowSpend.

View file

@ -215,11 +215,7 @@ class DBSpendUpdateWriter:
start_time: datetime | None,
end_time: datetime | None,
response_cost: float | None,
) -> str | None:
"""Returns the LiteLLM_SpendLogs request_id this call was recorded
under, so the caller can tell the budget-window writer which log rows
its increments already cover. None when the payload could not be built.
"""
) -> None:
from litellm.proxy.proxy_server import (
disable_spend_logs,
litellm_proxy_budget_name,
@ -236,7 +232,7 @@ class DBSpendUpdateWriter:
team_id,
)
if ProxyUpdateSpend.disable_spend_updates() is True:
return None
return
if token is not None and isinstance(token, str) and token.startswith("sk-"):
hashed_token = hash_token(token=token)
else:
@ -310,7 +306,6 @@ class DBSpendUpdateWriter:
)
verbose_proxy_logger.debug("Runs spend update on all tables")
return payload.get("request_id")
except Exception:
spend_log_error(
"Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue "
@ -323,7 +318,7 @@ class DBSpendUpdateWriter:
org_id,
end_user_id,
)
return None
return
async def _enqueue_tool_usage_transaction(
self,

View file

@ -46,6 +46,7 @@ from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdate
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
WindowSpendTransaction,
WindowSpendUpdateQueue,
to_wire_payload,
)
from litellm.secret_managers.main import str_to_bool
from litellm.types.caching import (
@ -298,7 +299,7 @@ class RedisUpdateBuffer:
ServiceTypes.REDIS_DAILY_AGENT_SPEND_UPDATE_QUEUE,
),
(
window_spend_update_transactions,
tuple(map(to_wire_payload, window_spend_update_transactions)),
REDIS_WINDOW_SPEND_UPDATE_BUFFER_KEY,
ServiceTypes.REDIS_WINDOW_SPEND_UPDATE_QUEUE,
),
@ -484,7 +485,12 @@ class RedisUpdateBuffer:
(daily_end_user_spend_update_transactions, REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY),
(daily_agent_spend_update_transactions, REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY),
(daily_tag_spend_update_transactions, REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY),
(window_spend_update_transactions, REDIS_WINDOW_SPEND_UPDATE_BUFFER_KEY),
(
None
if window_spend_update_transactions is None
else tuple(map(to_wire_payload, window_spend_update_transactions)),
REDIS_WINDOW_SPEND_UPDATE_BUFFER_KEY,
),
)
rpush_list: Final = tuple(

View file

@ -26,17 +26,12 @@ class WindowSpendTransaction(TypedDict):
window_start is an ISO-8601 string rather than a datetime so the
transaction survives the JSON round trip through the Redis buffer.
request_ids carries the LiteLLM_SpendLogs ids this spend came from. The
one-time seed for a window that has no row yet subtracts them from its
LiteLLM_SpendLogs aggregate, because the spend log writer flushes on its
own ~2s poll and will usually have persisted these rows before the window
queue flushes; without the exclusion the seed and the increment would each
count them.
started_at is the earliest request start in the batch. The seed only
subtracts a request_id whose LiteLLM_SpendLogs.startTime is at or after it,
so a client that replays an old id through x-litellm-call-id cannot make the
seed drop the historical row that id already paid for.
started_at is the earliest request start in the batch. The one-time seed for
a window that has no row yet uses it to tell this batch's own
LiteLLM_SpendLogs rows from everything else, because the spend log writer
flushes on its own ~2s poll and will usually have persisted this batch's
rows before the window queue flushes; without that split the seed and the
increment would each count them.
"""
entity_type: ReadOnly[str]
@ -44,10 +39,37 @@ class WindowSpendTransaction(TypedDict):
window_duration: ReadOnly[str]
window_start: ReadOnly[str]
spend: ReadOnly[float]
request_ids: ReadOnly[Sequence[str]]
started_at: ReadOnly[str | None]
class WindowSpendWirePayload(WindowSpendTransaction):
"""How an increment is encoded in the shared Redis buffer.
request_ids is dead weight here: workers built before this field was
dropped index it while merging whatever they pop, and the pop is
destructive, so a leader still running one of those during a rolling deploy
would raise on a payload without the key and lose those increments. It is
always empty, which only makes such a leader seed without exclusions.
TODO: remove once no supported version reads it, i.e. one release after the
field stopped being written.
"""
request_ids: ReadOnly[Sequence[str]]
def to_wire_payload(transaction: WindowSpendTransaction) -> WindowSpendWirePayload:
return WindowSpendWirePayload(
entity_type=transaction["entity_type"],
entity_id=transaction["entity_id"],
window_duration=transaction["window_duration"],
window_start=transaction["window_start"],
spend=transaction["spend"],
started_at=transaction.get("started_at"),
request_ids=(),
)
def to_naive_utc(value: datetime) -> datetime:
"""LiteLLM_BudgetWindowSpend.window_start is TIMESTAMP(3), which holds naive UTC."""
if value.tzinfo is None:
@ -72,7 +94,6 @@ def build_window_spend_transaction(
window_duration: str,
window_start: datetime,
spend: float,
request_id: str | None = None,
started_at: datetime | None = None,
) -> WindowSpendTransaction:
return WindowSpendTransaction(
@ -81,7 +102,6 @@ def build_window_spend_transaction(
window_duration=window_duration,
window_start=to_naive_utc(window_start).isoformat(timespec="microseconds"),
spend=spend,
request_ids=() if request_id is None else (request_id,),
started_at=None
if started_at is None
else to_naive_utc(started_at.astimezone(timezone.utc)).isoformat(timespec="microseconds"),
@ -101,7 +121,6 @@ def _merge_window_spend_transactions(
window_duration=first["window_duration"],
window_start=first["window_start"],
spend=math.fsum(payload["spend"] for payload in payloads),
request_ids=tuple(sorted(frozenset(chain.from_iterable(payload["request_ids"] for payload in payloads)))),
started_at=min(started_ats) if started_ats else None,
)

View file

@ -433,7 +433,7 @@ class HeadroomGuardrail(CustomGuardrail):
payload["model"] = model
try:
raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped
url=f"{self.headroom_api_base}/v1/compress",
json=payload,
headers=self._request_headers(),
@ -570,7 +570,7 @@ class HeadroomGuardrail(CustomGuardrail):
params["query"] = query
try:
raw_response: HttpxResponse = await self.async_handler.get( # pyright: ignore[reportUnknownMemberType]
raw_response: HttpxResponse = await self.async_handler.get( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.get is untyped
url=f"{self.headroom_api_base}/v1/retrieve/{hash_value}",
params=params,
headers=self._request_headers(),

View file

@ -156,6 +156,31 @@ def _header_value(headers: Mapping[str, str], key: str, default: str) -> str:
return headers.get(key, default)
def _is_image_part(item: object) -> bool:
"""Whether a structured-message content part carries an image rather than text."""
if not isinstance(item, Mapping):
return False
part: Final[Mapping[object, object]] = item
return part.get("type") == "image_url"
def _scannable_text(content: object) -> str:
"""Flatten a structured message's content into the single string the v1 detection endpoint takes.
Image parts are dropped: the endpoint accepts one string, so an image would only reach it as
its stringified source (a base64 blob or a URL), which is not text the scanner can evaluate.
"""
if not isinstance(content, list):
return str(content or "")
parts: Final[Sequence[object]] = content
text_parts: Final = [item for item in parts if not _is_image_part(item)] # mutable-ok: sent as a list repr
return str(text_parts or "")
def is_saas(host: str) -> bool:
"""Checks whether the connection is to the SaaS platform"""
@ -270,7 +295,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
"messages": [
{
"role": last_msg.get("role", "user"),
"content": str(last_msg.get("content", "")),
"content": _scannable_text(last_msg.get("content")),
}
]
},

View file

@ -20,6 +20,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
file_sanitization_fail_open=getattr(litellm_params, "file_sanitization_fail_open", None),
)
litellm.logging_callback_manager.add_litellm_callback(_prompt_security_callback)

View file

@ -4,10 +4,12 @@ import os
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Final, Literal, Optional
import httpx
from fastapi import HTTPException
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.exceptions import Timeout as LiteLLMTimeout
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
@ -24,6 +26,9 @@ if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
_SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS: Final = 30.0
class PromptSecurityGuardrailMissingSecrets(Exception):
pass
@ -63,6 +68,13 @@ class _SanitizeStatusResponse(TypedDict, total=False):
metadata: ReadOnly[_SanitizeMetadata]
class _SanitizeResult(TypedDict):
action: ReadOnly[str]
content: ReadOnly[str | None]
metadata: ReadOnly[_SanitizeMetadata]
violations: ReadOnly[Sequence[str]]
class PromptSecurityGuardrail(CustomGuardrail):
@classmethod
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
@ -79,6 +91,8 @@ class PromptSecurityGuardrail(CustomGuardrail):
user: str | None = None,
system_prompt: str | None = None,
check_tool_results: bool | None = None,
file_sanitization_timeout: float = _SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS,
file_sanitization_fail_open: bool | None = None,
**kwargs,
):
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
@ -108,6 +122,8 @@ class PromptSecurityGuardrail(CustomGuardrail):
# Configuration for file sanitization
self.max_poll_attempts = 30 # Maximum number of polling attempts
self.poll_interval = 2 # Seconds between polling attempts
self.file_sanitization_timeout = file_sanitization_timeout
self.file_sanitization_fail_open = file_sanitization_fail_open is not False
super().__init__(**kwargs)
@ -397,6 +413,39 @@ class PromptSecurityGuardrail(CustomGuardrail):
Sanitize file content using Prompt Security API.
Returns: dict with keys 'action', 'content', 'metadata'
"""
try:
return await asyncio.wait_for(
self._sanitize_file_content(file_data, filename, user_api_key_alias),
timeout=self.file_sanitization_timeout,
)
except (asyncio.TimeoutError, httpx.TimeoutException, LiteLLMTimeout) as exc:
if not self.file_sanitization_fail_open:
verbose_proxy_logger.error(
"Prompt Security Guardrail: file sanitization for %s timed out with %s; failing closed",
filename,
type(exc).__name__,
)
raise HTTPException(status_code=408, detail="File sanitization timeout") from exc
verbose_proxy_logger.error(
"Prompt Security Guardrail: file sanitization for %s timed out with %s; failing open",
filename,
type(exc).__name__,
)
fail_open_result: Final[_SanitizeResult] = {
"action": "allow",
"content": None,
"metadata": {},
"violations": (),
}
return fail_open_result
async def _sanitize_file_content(
self,
file_data: bytes,
filename: str,
user_api_key_alias: str | None,
) -> _SanitizeResult:
headers: Final = {"APP-ID": self.api_key}
if user_api_key_alias:
headers["X-LiteLLM-Key-Alias"] = user_api_key_alias

View file

@ -197,7 +197,7 @@ class RepelloAIGuardrail(CustomGuardrail):
repelloai_response: RepelloAIAnalyzeResponse | None = None
try:
verbose_proxy_logger.debug("RepelloAI Argus request: %s", request)
response: Final[HttpxResponse] = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
response: Final[HttpxResponse] = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped
url=endpoint,
headers={"X-API-Key": self.repelloai_api_key},
json=request,

View file

@ -587,7 +587,7 @@ async def _update_database_and_spend_counters(
model_access_groups: Sequence[str] | None = None,
) -> None:
try:
spend_log_request_id = await proxy_logging_obj.db_spend_update_writer.update_database(
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
@ -623,7 +623,6 @@ async def _update_database_and_spend_counters(
budget_reservation=budget_reservation,
end_user_id=end_user_id,
tags=request_tags,
request_id=spend_log_request_id,
request_started_at=start_time,
model_access_groups=model_access_groups,
)

View file

@ -7,7 +7,7 @@ POST /auto_router/validate_complexity_router_config - Dry-run the complexity-rou
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from itertools import groupby
from itertools import chain, groupby
from operator import attrgetter
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Protocol
@ -58,10 +58,11 @@ from litellm.types.management_endpoints.auto_router_endpoints import (
ComplexityRouterConfigValidationResponse,
RequestComplexityRouterConfig,
ShadowEvalDirection,
ShadowEvalJobKeyResponse,
ShadowEvalJobResponse,
ShadowEvalJobTargetResponse,
ShadowEvalResult,
ShadowEvalSlice,
ShadowEvalTargetType,
StartShadowEvalRequest,
)
@ -104,10 +105,43 @@ class _VerificationTokenTable(Protocol):
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_VerificationTokenRow]: ...
class _TeamRow(Protocol):
@property
def team_id(self) -> str: ...
@property
def team_alias(self) -> str | None: ...
class _TeamRowsTable(Protocol):
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_TeamRow]: ...
class _UserRow(Protocol):
@property
def user_id(self) -> str: ...
@property
def user_email(self) -> str | None: ...
class _UserRowsTable(Protocol):
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_UserRow]: ...
class _ShadowEvalJobRow(Protocol):
@property
def id(self) -> str: ...
@property
def group_id(self) -> str: ...
@property
def target_type(self) -> str: ...
@property
def target_id(self) -> str: ...
class _ShadowEvalJobTable(Protocol):
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_ShadowEvalJobRow]: ...
@ -138,6 +172,14 @@ def _verification_tokens(prisma_client: "PrismaClient") -> _VerificationTokenTab
return prisma_client.db.litellm_verificationtoken
def _team_rows(prisma_client: "PrismaClient") -> _TeamRowsTable:
return prisma_client.db.litellm_teamtable
def _user_rows(prisma_client: "PrismaClient") -> _UserRowsTable:
return prisma_client.db.litellm_usertable
def _shadow_eval_jobs(prisma_client: "PrismaClient") -> _ShadowEvalJobTable:
return prisma_client.db.litellm_shadowevaljob
@ -836,7 +878,7 @@ def _validate_judge_is_not_a_candidate(
def _is_unique_violation(error: Exception) -> bool:
"""Whether a Prisma create failed on a unique index. One active job per key and
"""Whether a Prisma create failed on a unique index. One active job per target and
direction lives in a partial unique index (raw SQL in the migration; schema.prisma
cannot express partial indexes), so the read-then-create check above it is advisory:
two concurrent starts pass the read, and the loser must surface as the same 409
@ -885,7 +927,7 @@ _ATTEMPT_AGG_BY_LEG_SQL: Final = "SELECT job_id AS grp," + _ATTEMPT_AGG_SELECT
# direction, and mid-deploy rows from old pods price as judge-only until the deploy ends).
_SWEEP_FINISHED_JOBS_SQL: Final = """
UPDATE "LiteLLM_ShadowEvalJob" j SET stopped_at = (NOW() AT TIME ZONE 'utc')
WHERE j.api_key_id = ANY($1::text[]) AND j.stopped_at IS NULL
WHERE j.target_type = $2 AND j.target_id = ANY($1::text[]) AND j.stopped_at IS NULL
AND (
j.ends_at <= (NOW() AT TIME ZONE 'utc')
OR (SELECT COUNT(*) FROM "LiteLLM_ShadowEvalAttempt" a WHERE a.job_id = j.id) >= j.max_turns
@ -966,10 +1008,10 @@ WHERE group_id IN (
)
"""
_LIST_LEGS_BY_KEY_SQL: Final = """
_LIST_LEGS_BY_TARGET_SQL: Final = """
SELECT * FROM "LiteLLM_ShadowEvalJob"
WHERE group_id IN (
SELECT group_id FROM "LiteLLM_ShadowEvalJob" WHERE api_key_id = $2
SELECT group_id FROM "LiteLLM_ShadowEvalJob" WHERE target_type = $2 AND target_id = $3
GROUP BY group_id ORDER BY MAX(created_at) DESC LIMIT $1::int
)
"""
@ -1007,15 +1049,16 @@ def _slices(rows: Sequence[_AttemptAggRow]) -> tuple[ShadowEvalSlice, ...]:
class _LegRow(BaseModel):
"""One LiteLLM_ShadowEvalJob row, validated off the untyped prisma record. A row is
one key's leg of a job; the legs of a job share group_id and identical config, written
together by one create_many. The API's job id is the group id, so leg ids never leave
the server (attempts reference them internally)."""
one target's leg of a job; the legs of a job share group_id and identical config,
written together by one create_many. The API's job id is the group id, so leg ids
never leave the server (attempts reference them internally)."""
model_config = ConfigDict(from_attributes=True)
id: str
group_id: str
api_key_id: str
target_type: ShadowEvalTargetType
target_id: str
router_name: str
direction: ShadowEvalDirection
baseline_model: str | None = None
@ -1068,16 +1111,17 @@ def _group_response(
first: Final = legs[0]
return ShadowEvalJobResponse(
job_id=group_id,
keys=tuple(
ShadowEvalJobKeyResponse(
api_key_id=leg.api_key_id,
targets=tuple(
ShadowEvalJobTargetResponse(
target_type=leg.target_type,
target_id=leg.target_id,
max_turns=leg.max_turns,
max_budget=leg.max_budget,
stopped_at=leg.stopped_at,
attempt_count=stats.attempt_count if (stats := attempt_counts.get(leg.id)) else 0,
spend=round(stats.spend, 6) if stats else 0.0,
)
for leg in sorted(legs, key=lambda leg: leg.api_key_id)
for leg in sorted(legs, key=lambda leg: (leg.target_type, leg.target_id))
),
router_name=first.router_name,
direction=first.direction,
@ -1090,34 +1134,85 @@ def _group_response(
)
_NO_KEY_LABELS: Final[tuple[str | None, str | None]] = (None, None)
_NO_TARGET_LABELS: Final[tuple[str | None, str | None]] = (None, None)
async def _with_key_labels(
def _target_labels(
key_rows: Sequence[_VerificationTokenRow],
team_rows: Sequence[_TeamRow],
user_rows: Sequence[_UserRow],
) -> Mapping[tuple[str, str], tuple[str | None, str | None]]:
"""Display labels by (target_type, target_id): a key's (alias, masked name), a
team's (alias, None), a user's (email, None)."""
return MappingProxyType(
{ # mutable-ok: MappingProxyType needs a dict to wrap
key: value
for key, value in chain(
((("key", row.token), (row.key_alias, row.key_name)) for row in key_rows),
((("team", row.team_id), (row.team_alias, None)) for row in team_rows),
((("user", row.user_id), (row.user_email, None)) for row in user_rows),
)
}
)
def _target_ids_of(responses: Sequence[ShadowEvalJobResponse], target_type: ShadowEvalTargetType) -> tuple[str, ...]:
return tuple(
sorted(
frozenset(
target.target_id
for response in responses
for target in response.targets
if target.target_type == target_type
)
)
)
async def _with_target_labels(
prisma_client: "PrismaClient", responses: Sequence[ShadowEvalJobResponse]
) -> tuple[ShadowEvalJobResponse, ...]:
"""Resolve every scoped key's hash to its alias and masked name in one batched read,
so the UI can say whose traffic a job shadows. Deleted keys resolve to None."""
"""Resolve every scoped target's id to a display label in one batched read per kind,
so the UI can say whose traffic a job shadows: a key's alias and masked name, a
team's alias, a user's email. Deleted targets resolve to None."""
if not responses:
return ()
tokens: Final = sorted(frozenset(key.api_key_id for response in responses for key in response.keys))
key_rows: Final = await _verification_tokens(prisma_client).find_many(
where={"token": {"in": tokens}} # mutable-ok: Prisma filter
tokens: Final = _target_ids_of(responses, "key")
team_ids: Final = _target_ids_of(responses, "team")
user_ids: Final = _target_ids_of(responses, "user")
key_rows: Final = (
await _verification_tokens(prisma_client).find_many(
where={"token": {"in": list(tokens)}} # mutable-ok: Prisma filter
)
if tokens
else ()
)
labels: Final[Mapping[str, tuple[str | None, str | None]]] = {
row.token: (row.key_alias, row.key_name) for row in key_rows or ()
}
team_rows: Final = (
await _team_rows(prisma_client).find_many(
where={"team_id": {"in": list(team_ids)}} # mutable-ok: Prisma filter
)
if team_ids
else ()
)
user_rows: Final = (
await _user_rows(prisma_client).find_many(
where={"user_id": {"in": list(user_ids)}} # mutable-ok: Prisma filter
)
if user_ids
else ()
)
labels: Final = _target_labels(key_rows or (), team_rows or (), user_rows or ())
return tuple(
response.model_copy(
update={ # mutable-ok: pydantic update payload
"keys": tuple(
key.model_copy(
"targets": tuple(
target.model_copy(
update={ # mutable-ok: pydantic update payload
"key_alias": labels.get(key.api_key_id, _NO_KEY_LABELS)[0],
"key_name": labels.get(key.api_key_id, _NO_KEY_LABELS)[1],
"target_alias": labels.get((target.target_type, target.target_id), _NO_TARGET_LABELS)[0],
"key_name": labels.get((target.target_type, target.target_id), _NO_TARGET_LABELS)[1],
}
)
for key in response.keys
for target in response.targets
)
}
)
@ -1125,29 +1220,37 @@ async def _with_key_labels(
)
async def _shadow_eval_results(prisma_client: "PrismaClient", legs: Sequence[_LegRow]) -> ShadowEvalResult | None:
"""All three stratifications of one job's verdicts. Tier answers "where does the router
do well"; the model stratification groups by whichever model served the real arm, so it
answers "which of the models these keys use today would the router beat" forward, and
"for the turns the router sent to X, did X beat the baseline" in reverse; key answers
"which key's traffic does the router suit". Reads are bounded by the job's own attempts
(<= the sum of its keys' max_turns) via the job_id index."""
async def _shadow_eval_results(
prisma_client: "PrismaClient", legs: Sequence[_LegRow]
) -> tuple[ShadowEvalResult | None, Mapping[tuple[str, str], ShadowEvalSlice]]:
"""One job's stratified verdicts, plus each target's own slice keyed by the
(target_type, target_id) pair so a key, team, and user sharing an id can never
collapse into one entry. Tier answers "where does the router do well"; the model
stratification groups by whichever model served the real arm, so it answers "which
of the models these targets use today would the router beat" forward, and "for the
turns the router sent to X, did X beat the baseline" in reverse; the per-target
slices answer "which target's traffic does the router suit". Reads are bounded by
the job's own attempts (<= the sum of its targets' max_turns) via the job_id index."""
leg_ids: Final = [leg.id for leg in legs] # mutable-ok: query param
by_tier: Final = _ATTEMPT_AGG_ROWS.validate_python(
await _query_raw(prisma_client, _ATTEMPT_AGG_BY_TIER_SQL, leg_ids) or ()
)
if not by_tier:
return None
return None, MappingProxyType({})
by_model: Final = _ATTEMPT_AGG_ROWS.validate_python(
await _query_raw(prisma_client, _ATTEMPT_AGG_BY_MODEL_SQL, leg_ids) or ()
)
key_by_leg: Final = MappingProxyType({leg.id: leg.api_key_id for leg in legs})
target_by_leg: Final = MappingProxyType({leg.id: (leg.target_type, leg.target_id) for leg in legs})
by_leg: Final = _ATTEMPT_AGG_ROWS.validate_python(
await _query_raw(prisma_client, _ATTEMPT_AGG_BY_LEG_SQL, leg_ids) or ()
)
by_key: Final = tuple(
row.model_copy(update={"grp": key_by_leg[row.grp]}) # mutable-ok: pydantic update payload
for row in by_leg
verdicts_by_target: Final[Mapping[tuple[str, str], ShadowEvalSlice]] = MappingProxyType(
{ # mutable-ok: MappingProxyType needs a dict to wrap
target_by_leg[slice.group]: slice.model_copy(
update={"group": target_by_leg[slice.group][1]} # mutable-ok: pydantic update payload
)
for slice in _slices(by_leg)
}
)
total_turns: Final = sum(r.turn_count for r in by_tier)
funnel_rows: Final = await _query_raw(prisma_client, _FUNNEL_TOTALS_SQL, leg_ids)
@ -1155,10 +1258,9 @@ async def _shadow_eval_results(prisma_client: "PrismaClient", legs: Sequence[_Le
# Coverage only when EVERY leg has a funnel row: a partial seed (one leg's insert
# failed) must read as unknown, not as job-level counts missing a leg's traffic.
funnel: Final = counted if counted is not None and counted.legs_with_rows == len(leg_ids) else None
return ShadowEvalResult(
result: Final = ShadowEvalResult(
by_tier=_slices(by_tier),
by_current_model=_slices(by_model),
by_key=_slices(by_key),
overall_shadow_win_rate_pct=_pct_of(sum(r.shadow_wins for r in by_tier), total_turns),
overall_tie_rate_pct=_pct_of(sum(r.ties for r in by_tier), total_turns),
sampled_real_spend=sum(r.real_spend for r in by_tier),
@ -1168,6 +1270,7 @@ async def _shadow_eval_results(prisma_client: "PrismaClient", legs: Sequence[_Le
shed_count=funnel.shed if funnel is not None else None,
withheld_count=funnel.withheld if funnel is not None else None,
)
return result, verdicts_by_target
@router.post(
@ -1182,22 +1285,29 @@ async def start_shadow_eval(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> ShadowEvalJobResponse:
"""
Start a shadow eval: duplicate a sampled slice of one or more keys' live traffic against
a second arm, judge the two responses blind, and stratify win rates by tier, by the model
that served the real arm, and by key.
Start a shadow eval: duplicate a sampled slice of one or more targets' live traffic
against a second arm, judge the two responses blind, and stratify win rates by tier,
by the model that served the real arm, and by target.
A forward job answers whether the keys should adopt router_name: it samples the requests
the router did not serve and duplicates them through it. A reverse job answers whether a
key already on the router still gains from it: it samples the requests the router did
serve and duplicates them against baseline_model. A key can hold one active job per
direction, so both questions can run at once.
A target is a virtual key, a team, or a user. Team and user targets match on the
identity every request resolves to at auth time, so they cover JWT-authenticated
traffic, which presents no virtual key; a user target samples that user's traffic
across all their teams, whether it arrives on a JWT or a key they own.
Shadow responses are never served to users. Each key samples until its recorded eval
spend, the shadow and judge calls' own cost, reaches max_budget dollars, the job's
window ends, or the job is stopped, so one key running out of budget does not end
sampling for the others; sampling changes propagate to pods within about 10 seconds.
Shadow and judge calls bill to the shadowed key but are excluded from request counts
and auto-router adoption metrics.
A forward job answers whether the targets should adopt router_name: it samples the
requests the router did not serve and duplicates them through it. A reverse job
answers whether a target already on the router still gains from it: it samples the
requests the router did serve and duplicates them against baseline_model. A target
can hold one active job per direction, so both questions can run at once, and a
request matching several jobs' targets (say its key and its team) is sampled by
each, separately budgeted.
Shadow responses are never served to users. Each target samples until its recorded
eval spend, the shadow and judge calls' own cost, reaches max_budget dollars, the
job's window ends, or the job is stopped, so one target running out of budget does
not end sampling for the others; sampling changes propagate to pods within about 10
seconds. Shadow and judge calls bill to the sampled request's own identity but are
excluded from request counts and auto-router adoption metrics.
"""
from litellm.proxy.proxy_server import llm_router, prisma_client
@ -1206,35 +1316,88 @@ async def start_shadow_eval(
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
if llm_router is None or not _is_configured_pre_routing_strategy(llm_router, data.router_name):
raise HTTPException(status_code=400, detail=f"'{data.router_name}' is not a configured auto-router")
token_rows: Final = await _verification_tokens(prisma_client).find_many(
where={"token": {"in": list(data.api_key_ids)}} # mutable-ok: Prisma filter
)
unknown: Final = tuple(sorted(frozenset(data.api_key_ids) - frozenset(row.token for row in token_rows or ())))
if unknown:
raise HTTPException(
status_code=400,
detail=(
f"api_key_ids not on this proxy: {', '.join(unknown)}; pass each key's token hash, "
"the value the key list and key info endpoints report"
),
token_rows: Final = (
await _verification_tokens(prisma_client).find_many(
where={"token": {"in": list(data.api_key_ids)}} # mutable-ok: Prisma filter
)
if data.api_key_ids
else ()
)
team_rows: Final = (
await _team_rows(prisma_client).find_many(
where={"team_id": {"in": list(data.team_ids)}} # mutable-ok: Prisma filter
)
if data.team_ids
else ()
)
user_rows: Final = (
await _user_rows(prisma_client).find_many(
where={"user_id": {"in": list(data.user_ids)}} # mutable-ok: Prisma filter
)
if data.user_ids
else ()
)
unknown_keys: Final = sorted(frozenset(data.api_key_ids) - frozenset(row.token for row in token_rows or ()))
unknown_teams: Final = sorted(frozenset(data.team_ids) - frozenset(row.team_id for row in team_rows or ()))
unknown_users: Final = sorted(frozenset(data.user_ids) - frozenset(row.user_id for row in user_rows or ()))
unknown_parts: Final = tuple(
part
for part in (
(
f"api_key_ids not on this proxy: {', '.join(unknown_keys)}; pass each key's token hash, "
"the value the key list and key info endpoints report"
)
if unknown_keys
else None,
f"team_ids not on this proxy: {', '.join(unknown_teams)}" if unknown_teams else None,
f"user_ids not on this proxy: {', '.join(unknown_users)}" if unknown_users else None,
)
if part is not None
)
if unknown_parts:
raise HTTPException(status_code=400, detail=". ".join(unknown_parts))
# Every model check below runs once per team the job samples for, since that is the
# identity the shadow and judge calls carry and therefore what the router selects on.
team_ids: Final = tuple(dict.fromkeys(row.team_id for row in token_rows or ()))
# A user target's traffic can span teams, so it validates unscoped (None); each
# sampled attempt still resolves the judge under its own request's team at eval time.
team_ids: Final = tuple(
dict.fromkeys(
(
*(row.team_id for row in token_rows or ()),
*data.team_ids,
*((None,) if data.user_ids else ()),
)
)
)
_validate_plain_model(llm_router, data.judge_model, "judge_model", team_ids)
if data.baseline_model is not None:
_validate_plain_model(llm_router, data.baseline_model, "baseline_model", team_ids)
_validate_judge_is_not_a_candidate(llm_router, data, team_ids)
requested_targets: Final[tuple[tuple[ShadowEvalTargetType, str], ...]] = (
*(("key", key) for key in data.api_key_ids),
*(("team", team) for team in data.team_ids),
*(("user", user) for user in data.user_ids),
)
requested_by_type: Final[tuple[tuple[ShadowEvalTargetType, tuple[str, ...]], ...]] = tuple(
(target_type, ids)
for target_type, ids in (("key", data.api_key_ids), ("team", data.team_ids), ("user", data.user_ids))
if ids
)
# A job whose window passed or whose budget ran out stopped sampling on its own,
# but its legs still hold their slots in the per-key, per-direction partial unique index
# until stamped; free them so a new eval can start. Sweeping both directions is deliberate.
requested: Final = list(data.api_key_ids) # mutable-ok: query param
await prisma_client.db.execute_raw(_SWEEP_FINISHED_JOBS_SQL, requested)
# but its legs still hold their slots in the per-target, per-direction partial unique
# index until stamped; free them so a new eval can start. Sweeping both directions is
# deliberate. Sweep and claim filter on exact (target_type, id) pairs so a team id
# that happens to equal a key hash never matches the other kind's slot.
for target_type, ids in requested_by_type:
await prisma_client.db.execute_raw(_SWEEP_FINISHED_JOBS_SQL, list(ids), target_type) # mutable-ok: query param
claimed: Final = await _shadow_eval_jobs(prisma_client).find_many(
where={ # mutable-ok: Prisma filter
"api_key_id": {"in": requested}, # mutable-ok: Prisma filter
"OR": [ # mutable-ok: Prisma filter
{"target_type": target_type, "target_id": {"in": list(ids)}} # mutable-ok: Prisma filter
for target_type, ids in requested_by_type
],
"direction": data.direction,
"stopped_at": None,
},
@ -1244,7 +1407,7 @@ async def start_shadow_eval(
status_code=409,
detail=(
f"Already in an active {data.direction} shadow eval job: "
+ ", ".join(sorted(f"{row.api_key_id} (job {row.group_id})" for row in claimed))
+ ", ".join(sorted(f"{row.target_type} {row.target_id} (job {row.group_id})" for row in claimed))
+ ". Stop it first."
),
)
@ -1268,10 +1431,16 @@ async def start_shadow_eval(
# Leg ids are minted here rather than by the DB default so the funnel seed below
# writes from the same values with no read-back, which a lagging read replica
# (DATABASE_URL_READ_REPLICA) could otherwise return empty.
leg_ids: Final = tuple(str(uuid4()) for _ in data.api_key_ids)
leg_ids: Final = tuple(str(uuid4()) for _ in requested_targets)
await _shadow_eval_jobs(prisma_client).create_many(
data=[ # mutable-ok: Prisma payload
{**shared_config, "id": leg_id, "api_key_id": key} for leg_id, key in zip(leg_ids, data.api_key_ids)
{ # mutable-ok: Prisma payload
**shared_config,
"id": leg_id,
"target_type": target_type,
"target_id": target_id,
} # mutable-ok: Prisma payload
for leg_id, (target_type, target_id) in zip(leg_ids, requested_targets)
]
)
except Exception as e:
@ -1280,7 +1449,8 @@ async def start_shadow_eval(
raise HTTPException(
status_code=409,
detail=(
f"A requested key was claimed by another {data.direction} shadow eval job concurrently. Stop it first."
f"A requested target was claimed by another {data.direction} shadow eval job concurrently. "
"Stop it first."
),
) from e
# Seed a zero funnel row per leg NOW: a fully covered job never skips a request, so
@ -1293,18 +1463,19 @@ async def start_shadow_eval(
)
except Exception as seed_err: # noqa: BLE001 # coverage is advisory; the job must still start
verbose_proxy_logger.error("shadow_eval: funnel seed failed for job %s: %s", group_id, seed_err)
labels: Final = MappingProxyType({row.token: row for row in token_rows})
labels: Final = _target_labels(token_rows or (), team_rows or (), user_rows or ())
return ShadowEvalJobResponse(
job_id=group_id,
keys=tuple(
ShadowEvalJobKeyResponse(
api_key_id=api_key_id,
targets=tuple(
ShadowEvalJobTargetResponse(
target_type=target_type,
target_id=target_id,
max_turns=SHADOW_EVAL_TURN_VALVE,
max_budget=data.max_budget,
key_alias=labels[api_key_id].key_alias,
key_name=labels[api_key_id].key_name,
target_alias=labels.get((target_type, target_id), _NO_TARGET_LABELS)[0],
key_name=labels.get((target_type, target_id), _NO_TARGET_LABELS)[1],
)
for api_key_id in sorted(data.api_key_ids)
for target_type, target_id in sorted(requested_targets)
),
router_name=data.router_name,
direction=data.direction,
@ -1324,22 +1495,29 @@ async def start_shadow_eval(
)
async def list_shadow_eval_jobs(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
api_key_id: Annotated[
str | None, Query(description="Filter to jobs that shadow this key, alone or alongside others")
target_type: Annotated[
ShadowEvalTargetType | None, Query(description="Kind of target to filter on; requires target_id")
] = None,
target_id: Annotated[
str | None, Query(description="Filter to jobs that shadow this target, alone or alongside others")
] = None,
limit: Annotated[int, Query(ge=1, le=200, description="Newest jobs to return")] = 50,
) -> tuple[ShadowEvalJobResponse, ...]:
"""List shadow eval jobs, newest first, each key with its attempt count so status is
accurate. Judged counts, spend, and results ride the detail endpoint only."""
"""List shadow eval jobs, newest first, each target with its attempt count so status
is accurate. Judged counts, spend, and results ride the detail endpoint only."""
from litellm.proxy.proxy_server import prisma_client
_require_admin_viewer(user_api_key_dict, "view shadow evals")
if prisma_client is None:
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
filter_type: Final = target_type if isinstance(target_type, str) else None
filter_id: Final = target_id if isinstance(target_id, str) else None
if (filter_type is None) != (filter_id is None):
raise HTTPException(status_code=400, detail="target_type and target_id filter together; pass both or neither")
legs: Final = _LEG_ROWS.validate_python(
(
await _query_raw(prisma_client, _LIST_LEGS_BY_KEY_SQL, limit, api_key_id)
if api_key_id
await _query_raw(prisma_client, _LIST_LEGS_BY_TARGET_SQL, limit, filter_type, filter_id)
if filter_type and filter_id
else await _query_raw(prisma_client, _LIST_LEGS_SQL, limit)
)
or ()
@ -1354,7 +1532,7 @@ async def list_shadow_eval_jobs(
by_group, key=lambda group_id: max(leg.created_at for leg in by_group[group_id]), reverse=True
)
counts: Final = await _leg_attempt_counts(prisma_client, legs)
return await _with_key_labels(
return await _with_target_labels(
prisma_client, tuple(_group_response(group_id, by_group[group_id], counts) for group_id in newest_first)
)
@ -1391,16 +1569,25 @@ async def get_shadow_eval_job(
where={"job_id": {"in": leg_ids}, "outcome": "error"}, # mutable-ok: Prisma filter
order={"created_at": "desc"}, # mutable-ok: Prisma order
)
labeled: Final = await _with_key_labels(
labeled: Final = await _with_target_labels(
prisma_client, (_group_response(job_id, legs, await _leg_attempt_counts(prisma_client, legs)),)
)
results, verdicts_by_target = await _shadow_eval_results(prisma_client, legs)
return labeled[0].model_copy(
update={ # mutable-ok: pydantic update payload
"judged_count": totals[0].judged_count if totals else 0,
"error_count": totals[0].error_count if totals else 0,
"judge_spend": round(totals[0].judge_spend, 6) if totals else 0.0,
"last_error": latest_error.error if latest_error else None,
"results": await _shadow_eval_results(prisma_client, legs),
"results": results,
"targets": tuple(
target.model_copy(
update={ # mutable-ok: pydantic update payload
"verdicts": verdicts_by_target.get((target.target_type, target.target_id))
}
)
for target in labeled[0].targets
),
}
)
@ -1415,8 +1602,8 @@ async def stop_shadow_eval_job(
job_id: str,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> ShadowEvalJobResponse:
"""Stop an active shadow eval job, every key it scopes at once. Attempts are kept;
sampling halts within ~10s. Keys that already stopped on their own budget keep the
"""Stop an active shadow eval job, every target it scopes at once. Attempts are kept;
sampling halts within ~10s. Targets that already stopped on their own budget keep the
stopped_at they earned. The statement is the whole state machine: it claims the job
only while a leg still samples inside the window with no stop recorded, so a racing
operator, a same-instant budget spend, and a repeat stop all read the same 400 with
@ -1443,5 +1630,5 @@ async def stop_shadow_eval_job(
current: Final = _group_response(job_id, legs, counts)
if claimed == 0:
raise HTTPException(status_code=400, detail=f"Job {job_id} is already {current.status}")
labeled: Final = await _with_key_labels(prisma_client, (current,))
labeled: Final = await _with_target_labels(prisma_client, (current,))
return labeled[0]

View file

@ -626,11 +626,7 @@ async def update_end_user(
# get non default values for key
non_default_values: Final = dict[str, object]()
for k, v in data_json.items():
if v is not None and v not in (
[],
{},
0,
): # models default to [], spend defaults to 0, we should not reset these values
if v is not None and ((isinstance(v, bool) and k in data.fields_set()) or v not in ([], {}, 0)):
non_default_values[k] = v
## Get end user table data ##

View file

@ -1027,6 +1027,14 @@ async def user_info_v2(
This is the v2 replacement for /user/info, designed to avoid the "god endpoint" problem
where the old endpoint loaded all keys and teams into memory.
Note on `spend`: this is the user's running budget counter, which the budget reset job
resets whenever `budget_reset_at` elapses (see `budget_duration`): to zero by default,
or to the overage above `max_budget` when `budget_rollover` is enabled. It is NOT
lifetime or per-period historical spend. For historical spend over a date range, use
`/user/daily/activity` or `/user/daily/activity/aggregated`, which read daily spend
records that only ever accumulate and are never reset. The two values are expected to
diverge once a budget reset has occurred within the queried period.
Access control:
- Proxy admins can query any user
- Team admins can query users within their teams
@ -2726,6 +2734,11 @@ async def get_user_daily_activity(
Meant to optimize querying spend data for analytics for a user.
Reads daily spend records that only ever accumulate and are never affected by budget
resets. Their total can legitimately exceed the `spend` field returned by
`/v2/user/info`, which is a running budget counter that every budget reset sets back
to zero (or to the overage above `max_budget` when `budget_rollover` is enabled).
Returns:
(by date)
- spend
@ -2839,6 +2852,11 @@ async def get_user_daily_activity_aggregated(
"""
Aggregated analytics for a user's daily activity without pagination.
Returns the same response shape as the paginated endpoint with page metadata set to single-page.
Reads daily spend records that only ever accumulate and are never affected by budget
resets. Their total can legitimately exceed the `spend` field returned by
`/v2/user/info`, which is a running budget counter that every budget reset sets back
to zero (or to the overage above `max_budget` when `budget_rollover` is enabled).
"""
from litellm.proxy.proxy_server import prisma_client

View file

@ -19,13 +19,13 @@ import re
import secrets
import traceback
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from contextlib import AbstractAsyncContextManager
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeVar, cast
import fastapi
import yaml
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
@ -155,6 +155,7 @@ from litellm.types.utils import (
)
if TYPE_CHECKING:
import prisma
from prisma import Prisma
from prisma import models as prisma_models
@ -182,6 +183,14 @@ class _TxTables(Protocol):
litellm_proxymodeltable: TableActions[object]
class _ModelParamsUpdate(TypedDict):
litellm_params: ReadOnly["prisma.Json"]
class _ModelRowWhere(TypedDict):
model_id: ReadOnly[str]
class _ConfigTableActions(Protocol):
"""Config table surface this module needs; the shared repository seam exposes no ``update``."""
@ -273,12 +282,6 @@ def _env_vars_param_value(param: _EnvVarsParam) -> Mapping[str, str] | None:
return param.param_value
def _tx_tables_context(
open_tx: Callable[[], AbstractAsyncContextManager[_TxTables]],
) -> AbstractAsyncContextManager[_TxTables]:
return open_tx()
async def _check_custom_key_allowed(custom_key_value: str | None) -> None:
"""Raise 403 if custom API keys are disabled and a custom key was provided."""
if custom_key_value is None:
@ -4484,27 +4487,29 @@ async def _rotate_master_key(
if models:
decrypted_models: Final = proxy_config.decrypt_model_list_from_db(new_models=models)
verbose_proxy_logger.debug("ABLE TO DECRYPT MODELS - len(decrypted_models): %s", len(decrypted_models))
new_models: Final[list[dict[str, object]]] = []
for model in decrypted_models:
new_model = await _add_model_to_db(
model_params=Deployment(**model),
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
new_encryption_key=new_master_key,
should_create_model_in_db=False,
)
if new_model:
_dumped = dict[str, object](_as_object_dict(new_model.model_dump(exclude_none=True)))
_dumped["litellm_params"] = prisma.Json(_dumped["litellm_params"])
_dumped["model_info"] = prisma.Json(_dumped["model_info"])
new_models.append(_dumped)
verbose_proxy_logger.debug("Resetting proxy model table")
async with _tx_tables_context(prisma_client.db.tx) as tx:
await tx.litellm_proxymodeltable.delete_many()
verbose_proxy_logger.debug("Creating %s models", len(new_models))
await tx.litellm_proxymodeltable.create_many(
data=new_models,
)
reencrypted_models: Final = tuple(
[
reencrypted
for model in decrypted_models
if (
reencrypted := await _add_model_to_db(
model_params=Deployment(**model),
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
new_encryption_key=new_master_key,
should_create_model_in_db=False,
)
)
]
)
verbose_proxy_logger.debug("Re-encrypting litellm_params on %s model rows", len(reencrypted_models))
async with prisma_client.db.tx(timeout=timedelta(minutes=2)) as tx_ctx:
tx: Final[_TxTables] = tx_ctx
for reencrypted_model in reencrypted_models:
await tx.litellm_proxymodeltable.update_many(
data=_ModelParamsUpdate(litellm_params=prisma.Json(reencrypted_model.litellm_params)),
where=_ModelRowWhere(model_id=reencrypted_model.model_id),
)
await publish_config_change(redis_cache=coordination_redis_cache(), object_type="litellm_proxymodeltable")
# 3. process config table
try:

View file

@ -252,7 +252,7 @@ async def add_team_callbacks(
Use this if if you want different teams to have different success/failure callbacks
Parameters:
- callback_name (Literal["langfuse", "langsmith", "gcs"], required): The name of the callback to add
- callback_name (str, required): The name of the callback to add, e.g. "langfuse", "langsmith", "gcs", "newrelic". The value is validated against the callbacks that support team-scoped credentials
- callback_type (Literal["success", "failure", "success_and_failure"], required): The type of callback to add. One of:
- "success": Callback for successful LLM calls
- "failure": Callback for failed LLM calls
@ -268,6 +268,8 @@ async def add_team_callbacks(
- langsmith_api_key: The API key for the Langsmith callback
- langsmith_project: The project for the Langsmith callback
- langsmith_base_url: The base URL for the Langsmith callback
- newrelic_api_key: The ingest license key for the team's New Relic account; routes both LLM/agent traces and cost metrics to that account. Requires the proxy to run with LITELLM_OTEL_V2=true, otherwise this callback is rejected with a 400
- newrelic_region: The New Relic region for the team's account ("us" or "eu"), riding the team's own key
Example curl:
```

View file

@ -6,6 +6,8 @@ Provider-specific Pass-Through Endpoints
Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
"""
from __future__ import annotations
import hmac
import json
import os
@ -28,6 +30,7 @@ from litellm.constants import (
)
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
from litellm.proxy.auth.handle_jwt import JWTHandler
@ -51,6 +54,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
create_websocket_passthrough_route,
websocket_passthrough_request,
)
from litellm.proxy.utils import ProxyLogging as ProxyLoggingType
from litellm.proxy.utils import is_known_model
from litellm.proxy.vector_store_endpoints.utils import (
assert_proxy_admin_for_vector_store_index_management,
@ -70,13 +74,17 @@ from litellm.utils import ProviderConfigManager
from .passthrough_endpoint_router import PassthroughEndpointRouter
if TYPE_CHECKING:
from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig
from litellm.router import Router
ProxyConfig = _ProxyConfig # rebind-ok: conditional type alias
else:
ProxyConfig = Any # rebind-ok: runtime fallback
vertex_llm_base: Final = VertexBase()
router: Final = APIRouter()
openai_passthrough_router: Final = APIRouter()
default_vertex_config: Final = None
passthrough_endpoint_router: Final = PassthroughEndpointRouter()
@ -495,8 +503,14 @@ async def milvus_proxy_route(
request_body: Final = await get_request_body(request)
# check collectionName
collection_name: Final = cast(str | None, request_body.get("collectionName"))
extra_headers = {}
_raw_collection_name: Final = request_body.get("collectionName")
if _raw_collection_name is not None and not isinstance(_raw_collection_name, str):
raise HTTPException(
status_code=400,
detail=f"collectionName must be a string. Got {type(_raw_collection_name).__name__}",
)
collection_name: str | None = _raw_collection_name # rebind-ok: locally scoped conversion
extra_headers = {} # mutable-ok: dict for extra headers; rebind-ok: reassigned later from credentials
base_target_url: str | None = None
if not collection_name:
raise HTTPException(
@ -1273,7 +1287,7 @@ def _resolve_vertex_model_from_router(
vertex_location: Current vertex location (may be from URL)
Returns:
Tuple of (encoded_endpoint, endpoint, vertex_project, vertex_location)
tuple of (encoded_endpoint, endpoint, vertex_project, vertex_location)
with resolved values from router config
"""
if not llm_router:
@ -1702,7 +1716,7 @@ def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict:
def get_vertex_pass_through_handler(
call_type: Literal["discovery", "aiplatform"],
call_type: Literal["discovery", "aiplatform"], # noqa: UP037
) -> BaseVertexAIPassThroughHandler:
if call_type == "discovery":
return VertexAIDiscoveryPassThroughHandler()
@ -1726,7 +1740,7 @@ def _override_vertex_params_from_router_credentials(
vertex_location: Current vertex location (from URL)
Returns:
Tuple of (vertex_project, vertex_location) with overridden values if applicable
tuple of (vertex_project, vertex_location) with overridden values if applicable
"""
if router_credentials is None:
return vertex_project, vertex_location
@ -1893,12 +1907,12 @@ async def _prepare_vertex_auth_headers(
authenticated them is stripped on the credential-less branch
Returns:
Tuple containing:
tuple containing:
- headers: dict - Authentication headers to use
- base_target_url: Optional[str] - Updated base target URL
- base_target_url: str | None - Updated base target URL
- headers_passed_through: bool - Whether headers were passed through from request
- vertex_project: Optional[str] - Updated vertex project ID
- vertex_location: Optional[str] - Updated vertex location
- vertex_project: str | None - Updated vertex project ID
- vertex_location: str | None - Updated vertex location
"""
vertex_llm_base: Final = VertexBase()
headers_passed_through = False
@ -2546,7 +2560,7 @@ def _vertex_publisher_model_suffix(model: str) -> str:
return f"{VERTEX_PUBLISHER_MODEL_PREFIX}{model.rsplit('/', 1)[-1]}"
def _get_llm_router() -> "Router | None":
def _get_llm_router() -> Router | None:
from litellm.proxy.proxy_server import llm_router
return llm_router
@ -2586,7 +2600,7 @@ def _resolve_vertex_live_credentials(
def _build_vertex_live_setup_model_rewriter(
vertex_project: str | None,
vertex_location: str | None,
llm_router: "Router | None",
llm_router: Router | None,
) -> Callable[[str], str] | None:
"""
Rewrite the ``setup`` frame's model into the full Vertex resource path the Live API requires.
@ -2606,7 +2620,7 @@ def _build_vertex_live_setup_model_rewriter(
return rewrite
def _resolve_alias_to_upstream_model(setup_model: str, llm_router: "Router | None") -> str:
def _resolve_alias_to_upstream_model(setup_model: str, llm_router: Router | None) -> str:
"""
The Live SDK wraps whatever the caller typed as ``models/<name>``, so a gateway alias arrives prefixed
"""
@ -2796,6 +2810,238 @@ def create_generic_websocket_passthrough_endpoint(
)
@router.api_route(
"/gigachat/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"], # mutable-ok: FastAPI route methods
tags=["Gigachat Pass-through", "pass-through"], # mutable-ok: FastAPI route tags
)
async def gigachat_proxy_route(
endpoint: str,
request: Request,
fastapi_response: Response,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> Response:
"""
[Docs](https://docs.litellm.ai/docs/pass_through/gigachat)
"""
from litellm.proxy.proxy_server import (
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
select_data_generator,
user_api_base,
user_max_tokens,
user_model,
user_request_timeout,
user_temperature,
version,
)
## check for streaming
request_body: Final[dict[str, object]] = await get_request_body(request)
is_router_model = False # rebind-ok: conditionally set to True when model uses router
raw_model: Final = request_body.get("model")
model: Final = raw_model if isinstance(raw_model, str) else None
if model:
is_router_model = is_passthrough_request_using_router_model(
request_body, llm_router
) # rebind-ok: conditionally set to True
elif any(word in endpoint for word in ("completions", "embeddings")):
raise HTTPException(
status_code=400, detail={"error": "Model is required in request body"}
) # mutable-ok: HTTPException detail dict
# If router model, use dedicated router passthrough handler
# This uses the same common processing path as non-router models
if model and is_router_model and llm_router:
return await handle_gigachat_passthrough_router_model(
model=model,
endpoint=endpoint,
request=request,
request_body=request_body,
fastapi_response=fastapi_response,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
verbose_proxy_logger.debug(
"Gigachat passthrough: Using direct Gigachat model '%s' for endpoint '%s'", model, endpoint
)
from litellm.llms.gigachat.authenticator import get_access_token
from litellm.llms.gigachat.utils import GIGACHAT_BASE_URL
base_target_url: Final = get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL
request_path: Final = httpx.URL(endpoint).path
encoded_endpoint: Final = request_path if request_path.startswith("/") else f"/{request_path}"
base_url: Final = httpx.URL(base_target_url)
updated_url: Final = base_url.copy_with(
path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, encoded_endpoint)
)
is_streaming_request: Final = await is_streaming_request_fn(request)
endpoint_func: Final = create_pass_through_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers={"Authorization": f"Bearer {get_access_token()}"},
is_streaming_request=is_streaming_request,
)
return await endpoint_func(
request,
fastapi_response,
user_api_key_dict,
)
async def handle_gigachat_passthrough_router_model(
model: str,
endpoint: str,
request: Request,
request_body: dict,
fastapi_response: Response,
llm_router: litellm.Router,
user_api_key_dict: UserAPIKeyAuth,
proxy_logging_obj: ProxyLoggingType,
general_settings: dict,
proxy_config: ProxyConfig,
select_data_generator: Callable,
user_model: str | None,
user_temperature: float | None,
user_request_timeout: float | None,
user_max_tokens: int | None,
user_api_base: str | None,
version: str | None,
) -> Response | StreamingResponse:
"""
Handle Gigachat passthrough for router models (models defined in config.yaml).
Uses the same common processing path as non-router models to ensure
metadata and hooks are properly initialized.
Args:
model: The router model name (e.g., "gigachat/gigachat-2")
endpoint: The Gigachat endpoint path (e.g., "/chat/completions")
request: The FastAPI request object
request_body: The parsed request body
llm_router: The LiteLLM router instance
user_api_key_dict: The user API key authentication dictionary
proxy_logging_obj: Proxy logging
general_settings: Proxy general settings
proxy_config: Proxy config
select_data_generator: Select data generator function
(additional args for common processing)
Returns:
Response or StreamingResponse depending on endpoint type
"""
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
# Detect streaming based on request body
is_streaming: Final = request_body.get("stream", False) # pyright: ignore[reportUnknownVariableType] # request_body is dict[Unknown, Unknown]
data: dict[str, Any] = await _read_request_body(
request=request
) # mutable-ok: mutated in place by proxy pipeline; pyright: ignore[reportExplicitAny] # Any needed for proxy pipeline
if user_api_key_dict is not None:
auth_metadata: Final = {
metadata_key: value
for metadata_key, value in (
("user_api_key_user_id", getattr(user_api_key_dict, "user_id", None)),
("user_api_key_team_id", getattr(user_api_key_dict, "team_id", None)),
("user_api_key_org_id", getattr(user_api_key_dict, "org_id", None)),
("agent_id", getattr(user_api_key_dict, "agent_id", None)),
)
if value is not None
}
existing_metadata: Final = data.get("metadata")
data["metadata"] = {
**(existing_metadata if isinstance(existing_metadata, dict) else {}),
**auth_metadata,
}
verbose_proxy_logger.debug(
"Gigachat router passthrough: model='%s', endpoint='%s', streaming=%s", model, endpoint, is_streaming
)
# Use the common processing path (same as non-router models)
# This ensures all metadata, hooks, and logging are properly initialized
data["model"] = model
data["method"] = request.method
data["endpoint"] = endpoint
data["json"] = request_body
data["custom_llm_provider"] = "gigachat"
# Remove sensitive keys from data
keys: Final = [ # mutable-ok: list of keys to remove from data
"gigachat_auth_url",
"gigachat_access_token",
"gigachat_scope",
"api_base",
"api_key",
]
for key in keys:
data.pop(key, None)
client: Final = get_async_httpx_client(
llm_provider=LlmProviders.GIGACHAT,
params={ # mutable-ok: httpx client params
"timeout": httpx.Timeout(timeout=600.0, connect=5.0),
},
)
data["client"] = client
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
# Use the common passthrough processing to handle metadata and hooks
# This also handles all response formatting (streaming/non-streaming) and exceptions
try:
result = await base_llm_response_processor.base_passthrough_process_llm_request( # rebind-ok: assigned once in try block
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=model,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except Exception as e: # noqa: BLE001 # Safe catch-all for handle exception
# Use common exception handling
raise await base_llm_response_processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
)
else:
if isinstance(result, StreamingResponse):
if result.headers.get("Content-Type") is None:
result.headers["Content-Type"] = "text/event-stream; charset=utf-8"
return result
@router.api_route(
"/watsonx/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],

View file

@ -6,9 +6,10 @@ import posixpath
import traceback
from base64 import b64encode
from collections.abc import AsyncGenerator, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from itertools import groupby
from typing import Any, Final, TypedDict, cast
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
from urllib.parse import urlencode, urlparse
import httpx
@ -47,6 +48,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
)
from litellm.litellm_core_utils.initialize_dynamic_callback_params import validate_no_callback_env_reference
from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
@ -78,7 +80,10 @@ from litellm.proxy.common_utils.http_parsing_utils import (
from litellm.proxy.common_utils.sse_keepalive import (
wrap_passthrough_sse_bytes_with_keepalive_pings,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
_get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above
)
from litellm.proxy.utils import normalize_route_for_root_path
from litellm.repositories.team_repository import TeamRepository
from litellm.secret_managers.main import get_secret_str
@ -90,7 +95,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
EndpointType,
PassthroughStandardLoggingPayload,
)
from litellm.types.utils import Usage
from litellm.types.utils import TRUSTED_CALLBACK_VARS_FIELD, Usage
from .streaming_handler import PassThroughStreamingHandler
from .success_handler import PassThroughEndpointLogging
@ -99,6 +104,9 @@ from .upstream_usage_headers import (
apply_upstream_reported_usage,
)
if TYPE_CHECKING:
from litellm.proxy.proxy_server import ProxyConfig
router: Final = APIRouter()
pass_through_endpoint_logging: Final = PassThroughEndpointLogging()
@ -752,6 +760,67 @@ def _build_passthrough_failure_request_payload(
return request_payload
@dataclass(frozen=True, slots=True)
class _TeamCallbackWiring:
success_callbacks: "list[str | Callable | CustomLogger] | None" = None # mutable-ok: Logging.__init__ arg
failure_callbacks: "list[str | Callable | CustomLogger] | None" = None # mutable-ok: Logging.__init__ arg
logging_kwargs: dict[str, str | dict[str, str]] | None = None # mutable-ok: Logging.__init__ arg
def _resolve_team_callback_wiring(
user_api_key_dict: UserAPIKeyAuth,
proxy_config: "ProxyConfig",
route_description: str,
) -> _TeamCallbackWiring:
"""Resolve key/team dynamic logging callbacks for a passthrough request.
Mirrors add_litellm_data_to_request: callback_vars are unpacked top-level
(read by initialize_standard_callback_dynamic_params) and also stamped on
the proxy-owned trusted-vars field (read by get_trusted_callback_params).
Fails open: a callback resolution or validation error is logged at error
level and the request proceeds without dynamic callbacks, since a broken
logging config must not fail the customer's upstream call (and the
websocket is already accepted by the time this runs on that path). The
env-reference check runs here because the deprecated callback_settings
branch skips AddTeamCallback validation, and Logging.__init__ would
otherwise reject the vars mid-request.
"""
try:
callback_settings_obj: Final = _get_dynamic_logging_metadata(
user_api_key_dict=user_api_key_dict, proxy_config=proxy_config
)
if callback_settings_obj and callback_settings_obj.callback_vars:
for (
item
) in callback_settings_obj.callback_vars.items(): # rebind-ok: dict.items iteration for env-ref validation
validate_no_callback_env_reference(item[0], item[1], source="key/team callback metadata")
except Exception: # noqa: BLE001 - a broken logging config must never fail the passthrough request
verbose_proxy_logger.exception(
"%s: failed to resolve team logging callbacks, continuing without them",
route_description,
)
return _TeamCallbackWiring()
if callback_settings_obj is None:
return _TeamCallbackWiring()
callback_vars: Final = callback_settings_obj.callback_vars
success_callbacks: Final = callback_settings_obj.success_callback
failure_callbacks: Final = callback_settings_obj.failure_callback
logging_kwargs: Final = (
None
if not callback_vars
else { # mutable-ok: Logging arg
**callback_vars,
TRUSTED_CALLBACK_VARS_FIELD: callback_vars,
}
)
return _TeamCallbackWiring(
success_callbacks=None if success_callbacks is None else [*success_callbacks], # mutable-ok: Logging arg
failure_callbacks=None if failure_callbacks is None else [*failure_callbacks], # mutable-ok: Logging arg
logging_kwargs=logging_kwargs,
)
async def _log_passthrough_upstream_failure(
response: httpx.Response,
user_api_key_dict: UserAPIKeyAuth,
@ -845,7 +914,7 @@ async def pass_through_request(
from litellm.proxy.pass_through_endpoints.passthrough_guardrails import (
PassthroughGuardrailHandler,
)
from litellm.proxy.proxy_server import proxy_logging_obj
from litellm.proxy.proxy_server import proxy_config, proxy_logging_obj
#########################################################
# Initialize variables
@ -930,6 +999,11 @@ async def pass_through_request(
# read e.g. ``chat gpt-4o`` instead of ``chat unknown``.
passthrough_model: Final = (_parsed_body.get("model") if isinstance(_parsed_body, dict) else None) or "unknown"
start_time: Final = datetime.now()
team_callbacks: Final = _resolve_team_callback_wiring(
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
route_description="pass_through_endpoint",
)
logging_obj = Logging(
model=passthrough_model,
messages=[{"role": "user", "content": safe_dumps(_parsed_body)}],
@ -938,6 +1012,9 @@ async def pass_through_request(
start_time=start_time,
litellm_call_id=litellm_call_id,
function_id="1245",
dynamic_success_callbacks=team_callbacks.success_callbacks,
dynamic_failure_callbacks=team_callbacks.failure_callbacks,
kwargs=team_callbacks.logging_kwargs,
)
# Store passthrough guardrails config on logging_obj for field targeting
@ -2022,7 +2099,7 @@ async def websocket_passthrough_request(
setup_model_rewriter: Optional rewrite of the setup frame's model before it reaches the upstream
"""
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.proxy_server import proxy_logging_obj
from litellm.proxy.proxy_server import proxy_config, proxy_logging_obj
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
PassthroughStandardLoggingPayload,
)
@ -2055,6 +2132,11 @@ async def websocket_passthrough_request(
upstream_headers[header_name] = header_value
# Initialize logging object similar to HTTP passthrough
team_callbacks: Final = _resolve_team_callback_wiring(
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
route_description="websocket_passthrough",
)
logging_obj: Final = Logging(
model="unknown",
messages=[{"role": "user", "content": "WebSocket connection"}],
@ -2063,6 +2145,9 @@ async def websocket_passthrough_request(
start_time=start_time,
litellm_call_id=litellm_call_id,
function_id="websocket_passthrough",
dynamic_success_callbacks=team_callbacks.success_callbacks,
dynamic_failure_callbacks=team_callbacks.failure_callbacks,
kwargs=team_callbacks.logging_kwargs,
)
# Create passthrough logging payload

View file

@ -2658,7 +2658,6 @@ async def increment_spend_counters(
budget_reservation: dict | None = None,
end_user_id: str | None = None,
tags: list[str] | None = None,
request_id: str | None = None,
request_started_at: datetime | None = None,
model_access_groups: Sequence[str] | None = None,
):
@ -2733,7 +2732,6 @@ async def increment_spend_counters(
window_duration=duration,
window_start=key_window_start,
increment=cost,
request_id=request_id,
request_started_at=request_started_at,
)
@ -2777,7 +2775,6 @@ async def increment_spend_counters(
window_duration=duration,
window_start=team_window_start,
increment=cost,
request_id=request_id,
request_started_at=request_started_at,
)
@ -3005,16 +3002,15 @@ async def _enqueue_window_spend_row_update(
window_duration: str,
window_start: datetime | None,
increment: float,
request_id: str | None,
request_started_at: datetime | None,
) -> None:
"""Queue this request's cost against the LiteLLM_BudgetWindowSpend row for
the window, so enforcement can read a maintained total instead of
aggregating LiteLLM_SpendLogs.
request_id is the LiteLLM_SpendLogs id this cost was recorded under and
request_started_at its startTime; the flush uses them to keep the one-time
seed from counting a request that its increment already covers.
request_started_at is this request's LiteLLM_SpendLogs startTime; the flush
stops the one-time seed there so a request its increment already covers is
not counted twice.
Enqueued even when the cache increment was skipped for a reserved counter:
the reservation only pre-charged the counter, and the row still owes the
@ -3035,7 +3031,6 @@ async def _enqueue_window_spend_row_update(
window_duration=window_duration,
window_start=window_start,
spend=increment,
request_id=request_id,
started_at=request_started_at,
)
)

View file

@ -1318,6 +1318,68 @@
],
"default_model_placeholder": "gpt-3.5-turbo"
},
{
"provider": "GIGACHAT",
"provider_display_name": "GigaChat",
"litellm_provider": "gigachat",
"credential_fields": [
{
"key": "api_base",
"label": "API Base",
"placeholder": null,
"tooltip": null,
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "api_key",
"label": "API Key",
"placeholder": null,
"tooltip": null,
"required": false,
"field_type": "password",
"options": null,
"default_value": null
},
{
"key": "gigachat_scope",
"label": "Scope",
"placeholder": null,
"tooltip": null,
"required": false,
"field_type": "select",
"options": [
"GIGACHAT_API_PERS",
"GIGACHAT_API_B2B",
"GIGACHAT_API_CORP"
],
"default_value": "GIGACHAT_API_PERS"
},
{
"key": "gigachat_auth_url",
"label": "Auth URL",
"placeholder": null,
"tooltip": null,
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "gigachat_access_token",
"label": "Access token",
"placeholder": null,
"tooltip": "Disable OAuth, provide value to authorization.",
"required": false,
"field_type": "password",
"options": null,
"default_value": null
}
],
"default_model_placeholder": "GigaChat-2"
},
{
"provider": "GITHUB",
"provider_display_name": "Github",

View file

@ -1,14 +1,18 @@
import asyncio
import json
import time
from collections.abc import AsyncIterator, Mapping
from collections.abc import AsyncIterator, Awaitable, Mapping
from enum import Enum
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NamedTuple, cast, get_args
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, cast, get_args
from uuid import uuid4
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi.responses import JSONResponse
from openai.types.responses.response_create_params import ResponseInputParam
from starlette.websockets import WebSocket, WebSocketDisconnect
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import ModifyResponseException
@ -26,8 +30,13 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_set_request_parsed_body,
)
from litellm.types.llms.openai import REASONING_EFFORT, ResponsesAPIResponse
from litellm.types.llms.openai import (
REASONING_EFFORT,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
)
from litellm.types.responses.main import DeleteResponseResult
from litellm.types.utils import TokenCountResponse
if TYPE_CHECKING:
from litellm.router import Router
@ -35,7 +44,7 @@ if TYPE_CHECKING:
router: Final = APIRouter()
_user_api_key_auth_dep: Final = Depends(user_api_key_auth)
_RESPONSES_TAGS: Final = ["responses"] # mutable-ok: fastapi's route signature requires List[str] tags
_RESPONSES_TAGS: Final[list[str | Enum]] = ["responses"] # mutable-ok: fastapi's route signature requires list tags
_TOOL_PAYLOAD_KEYS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
{
@ -1017,6 +1026,152 @@ async def compact_response(
)
class _ResponsesApiErrorDetail(TypedDict):
message: ReadOnly[str]
type: ReadOnly[str]
param: ReadOnly[str | None]
code: ReadOnly[str | None]
class _ResponsesApiErrorBody(TypedDict):
error: ReadOnly[_ResponsesApiErrorDetail]
class _ResponsesInputTokensResult(TypedDict):
object: ReadOnly[str]
input_tokens: ReadOnly[int]
class _TokenCountPayload(TypedDict):
model: ReadOnly[str]
messages: ReadOnly[tuple[Mapping[str, object], ...]]
tools: ReadOnly[object]
class _TokenCounter(Protocol):
def __call__(self, request: TokenCountRequest, call_endpoint: bool) -> Awaitable[TokenCountResponse]: ...
def _proxy_token_counter() -> _TokenCounter:
from litellm.proxy.proxy_server import token_counter
return token_counter
_token_counter_dep: Final = Depends(_proxy_token_counter)
def _responses_invalid_request_response(message: str, param: str | None, code: str | None) -> JSONResponse:
body: Final[_ResponsesApiErrorBody] = {
"error": {
"message": message,
"type": "invalid_request_error",
"param": param,
"code": code,
}
}
return JSONResponse(status_code=400, content=body)
def _missing_responses_param_response(param: str) -> JSONResponse:
return _responses_invalid_request_response(
message=f"Missing required parameter: '{param}'.",
param=param,
code="missing_required_parameter",
)
def _responses_input_as_token_count_messages(
input_value: str | ResponseInputParam,
instructions: str | None,
) -> tuple[Mapping[str, object], ...]:
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
request_params: Final[ResponsesAPIOptionalRequestParams] = {"instructions": instructions}
transformed: Final = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=input_value,
responses_api_request=request_params,
)
return tuple(
message if isinstance(message, dict) else message.model_dump(exclude_none=True) for message in transformed
)
@router.post(
"/v1/responses/input_tokens",
dependencies=(_user_api_key_auth_dep,),
tags=_RESPONSES_TAGS,
)
@router.post(
"/responses/input_tokens",
dependencies=(_user_api_key_auth_dep,),
tags=_RESPONSES_TAGS,
)
@router.post(
"/openai/v1/responses/input_tokens",
dependencies=(_user_api_key_auth_dep,),
tags=_RESPONSES_TAGS,
)
async def responses_input_tokens(
request: Request,
token_counter: _TokenCounter = _token_counter_dep,
):
"""
Count the input tokens of a Responses API request without calling the model.
Follows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/input-tokens
```bash
curl -X POST http://localhost:4000/v1/responses/input_tokens \
-H "Content-Type: application/json" \
-H "Authorization: Bearer sk-1234" \
-d '{
"model": "gpt-4o",
"input": "Hello, how are you?"
}'
```
Returns: `{"object": "response.input_tokens", "input_tokens": <count>}`
"""
data: Final = await _read_request_body(request=request)
model_name: Final = data.get("model")
input_value: Final = data.get("input")
if not isinstance(model_name, str) or not model_name:
return _missing_responses_param_response("model")
if input_value is None:
return _missing_responses_param_response("input")
if isinstance(input_value, (str, list)) and not input_value:
return _responses_invalid_request_response(
message="""One of "input" or "previous_response_id" or 'prompt' or 'conversation' must be provided.""",
param=None,
code="missing_required_parameter",
)
try:
payload: Final[_TokenCountPayload] = {
"model": model_name,
"messages": _responses_input_as_token_count_messages(
input_value=input_value,
instructions=data.get("instructions"),
),
"tools": data.get("tools"),
}
token_request: Final = TokenCountRequest.model_validate(payload)
except Exception as e:
return _responses_invalid_request_response(
message=f"Invalid request for token counting: {e}", param=None, code=None
)
token_response: Final = await token_counter(request=token_request, call_endpoint=True)
result: Final[_ResponsesInputTokensResult] = {
"object": "response.input_tokens",
"input_tokens": token_response.total_tokens,
}
return result
@router.post(
"/v1/responses/{response_id}/cancel",
dependencies=[Depends(user_api_key_auth)],

View file

@ -144,10 +144,8 @@ async def background_streaming_task(
# Process streaming response following OpenAI events format
# https://platform.openai.com/docs/api-reference/responses-streaming
output_items: Final = dict[str, _OutputItem]() # Track output items by ID
accumulated_text: Final = dict[
tuple[str, int], str
]() # Track accumulated text deltas by (item_id, content_index)
output_items: Final = dict[str, _OutputItem]()
accumulated_text: Final = dict[tuple[str, int], str]()
# ResponsesAPIResponse fields to extract from response.completed
usage_data = None
@ -262,7 +260,6 @@ async def background_streaming_task(
if "content" in delta_item:
content_list = delta_item["content"]
if content_index < len(content_list):
# Update existing content part with accumulated text
content_entry = content_list[content_index]
if isinstance(content_entry, dict):
content_entry["text"] = accumulated_text[key]

View file

@ -1529,14 +1529,15 @@ model LiteLLM_AutoRouterSession {
model LiteLLM_ShadowEvalJob {
id String @id @default(cuid())
group_id String // legs of one job share this; the API's job id
api_key_id String // hashed virtual key whose traffic this leg shadows
target_type String @default("key") // key | team | user
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
router_name String // the auto-router under evaluation, in either direction
direction String @default("forward") // forward | reverse
baseline_model String? // reverse only: the fixed model the router is judged against
judge_model String
shadow_percentage Float
max_turns Int // sample-count ceiling: the whole budget on pre-max_budget jobs, the error-loop valve otherwise
max_budget Float? // per-key USD cap on the eval's own shadow + judge spend; null on jobs from before spend budgets
max_budget Float? // per-target USD cap on the eval's own shadow + judge spend; null on jobs from before spend budgets
created_at DateTime @default(now())
created_by String?
ends_at DateTime
@ -1544,7 +1545,7 @@ model LiteLLM_ShadowEvalJob {
stopped_by String? // operator who stopped it early; null when it ended on its own
@@index([group_id])
@@index([api_key_id])
@@index([target_type, target_id])
@@index([created_at])
}

View file

@ -172,7 +172,14 @@ async def reserve_budget_for_request(
) -> dict | None:
if valid_token is None or not RouteChecks.is_llm_api_route(route=route):
return None
if route in {"/models", "/v1/models", "/utils/token_counter"}:
if route in {
"/models",
"/v1/models",
"/utils/token_counter",
"/responses/input_tokens",
"/v1/responses/input_tokens",
"/openai/v1/responses/input_tokens",
}:
return None
if get_model_from_request(request_body, route, llm_router=llm_router) is None:
return None

View file

@ -30,7 +30,7 @@ from litellm.litellm_core_utils.litellm_logging import (
request_model_access_groups_from_litellm_params,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsRouterMetadata
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
from litellm.proxy.utils import PrismaClient, hash_token
from litellm.types.utils import (
@ -93,6 +93,24 @@ def _redact_logged_api_key(value: str | None, *, already_redacted: bool = False)
return hash_token(stripped)
def _get_router_metadata_for_spend_log(
metadata: Mapping[str, object] | None,
requested_model: str | None,
selected_model: str | None,
selected_provider: str | None,
router_correlation_id: str | None,
) -> SpendLogsRouterMetadata | None:
model_info: Final = metadata.get("model_info") if metadata is not None else None
if not isinstance(model_info, Mapping) or model_info.get("internal_router_model") is not True:
return None
return SpendLogsRouterMetadata(
requested_model=requested_model or None,
selected_model=selected_model or None,
selected_provider=selected_provider or None,
router_correlation_id=router_correlation_id,
)
def _get_spend_logs_metadata(
metadata: dict | None,
applied_guardrails: list[str] | None = None,
@ -109,6 +127,7 @@ def _get_spend_logs_metadata(
cost_breakdown: CostBreakdown | None = None,
litellm_call_id: str | None = None,
autorouter_savings: float | None = None,
router_metadata: SpendLogsRouterMetadata | None = None,
) -> SpendLogsMetadata:
if metadata is None:
return SpendLogsMetadata(
@ -148,13 +167,17 @@ def _get_spend_logs_metadata(
autorouter_savings=autorouter_savings,
litellm_gateway_injected_cache=None,
litellm_call_id=litellm_call_id,
router_metadata=router_metadata,
)
verbose_proxy_logger.debug(
"getting payload for SpendLogs, available keys in metadata: " + str(list(metadata.keys()))
)
# Filter the metadata dictionary to include only the specified keys
clean_metadata: Final = SpendLogsMetadata(**{key: metadata.get(key) for key in SpendLogsMetadata.__annotations__})
clean_metadata: Final = SpendLogsMetadata(
**{key: metadata.get(key) for key in SpendLogsMetadata.__annotations__ if key != "router_metadata"},
router_metadata=router_metadata,
)
_raw_key: Final = clean_metadata.get("user_api_key")
_trusted_hash: Final = metadata.get("user_api_key_hash")
_already_redacted: Final = (
@ -375,6 +398,20 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
hidden_params: Final = standard_logging_payload.get("hidden_params", {})
litellm_overhead_time_ms = hidden_params.get("litellm_overhead_time_ms")
custom_llm_provider: Final = (
kwargs.get("custom_llm_provider")
or _sl_attribution_fallback(standard_logging_payload, "custom_llm_provider")
or None
)
raw_model: Final = cast(str, kwargs.get("model") or "")
model_name: Final = (
standard_logging_payload.get("model") if standard_logging_payload is not None else None
) or reconstruct_model_name(raw_model, custom_llm_provider, metadata or {})
litellm_call_id: Final = cast(
str | None,
kwargs.get("litellm_call_id") or litellm_params.get("litellm_call_id"),
)
# clean up litellm metadata
clean_metadata = _get_spend_logs_metadata(
metadata,
@ -433,9 +470,13 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
autorouter_savings=(
standard_logging_payload.get("autorouter_savings", None) if standard_logging_payload is not None else None
),
litellm_call_id=cast(
str | None,
kwargs.get("litellm_call_id") or litellm_params.get("litellm_call_id"),
litellm_call_id=litellm_call_id,
router_metadata=_get_router_metadata_for_spend_log(
metadata=metadata,
requested_model=_model_group,
selected_model=model_name,
selected_provider=custom_llm_provider,
router_correlation_id=litellm_call_id,
),
)
@ -480,15 +521,6 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
# Extract agent_id for A2A requests (set directly on model_call_details)
agent_id: Final[str | None] = kwargs.get("agent_id") or metadata.get("agent_id")
custom_llm_provider: Final = (
kwargs.get("custom_llm_provider")
or _sl_attribution_fallback(standard_logging_payload, "custom_llm_provider")
or None
)
raw_model: Final = cast(str, kwargs.get("model") or "")
model_name: Final = (
standard_logging_payload.get("model") if standard_logging_payload is not None else None
) or reconstruct_model_name(raw_model, custom_llm_provider, metadata or {})
try:
payload: Final[SpendLogsPayload] = SpendLogsPayload(

View file

@ -186,7 +186,6 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase):
base_url: Final = get_vertex_base_url(self.location)
url: Final = f"{base_url}/v1beta1/projects/{self.project_id}/locations/{self.location}/ragCorpora"
# Build request body with camelCase keys (Vertex AI API format)
vector_db_config: Final = self.vector_store_config.get("vector_db_config")
embedding_model: Final = self.vector_store_config.get("embedding_model")
embedding_model_config: Final = (
@ -447,7 +446,6 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase):
# Add max embedding requests per minute if specified
max_embedding_qpm: Final = self.vector_store_config.get("max_embedding_requests_per_min")
# Build request body with camelCase keys (Vertex AI API format)
chunking_config: Final = (
{"chunkSize": chunk_size or 1024, "chunkOverlap": chunk_overlap or 200}
if chunk_size or chunk_overlap

View file

@ -1629,6 +1629,8 @@ class LiteLLMCompletionResponsesConfig:
file_dict["file_id"] = file_id
if item.get("file_data"):
file_dict["file_data"] = item["file_data"]
if item.get("filename"):
file_dict["filename"] = item["filename"]
new_item: Final[dict[str, object]] = {"type": "file", "file": file_dict}
if "cache_control" in item:

View file

@ -6486,6 +6486,8 @@ class Router:
**kwargs,
)
elif call_type == "allm_passthrough_route":
if client:
kwargs["client"] = client
return await self._ageneric_api_call_with_fallbacks(
original_function=original_function,
passthrough_on_no_deployment=True,
@ -9152,7 +9154,8 @@ class Router:
if _deployment_on_router is not None:
# deployment with this model_id exists on the router
if (
deployment.litellm_params == _deployment_on_router.litellm_params
deployment.model_name == _deployment_on_router.model_name
and deployment.litellm_params == _deployment_on_router.litellm_params
and deployment.model_info == _deployment_on_router.model_info
):
# No need to update

View file

@ -154,6 +154,9 @@ model_list:
# Fallback model if tier cannot be determined
default_model: gpt-4o
# Replace a routed model that cannot take image input (default: false)
modality_routing: true
```
## Usage
@ -178,6 +181,25 @@ response = litellm.completion(
## Special Behaviors
### Modality-based capability routing
The classifier reads text alone, so a request carrying an image can classify cheap and land on a
text-only model, which rejects it with a provider 400 no fallback catches. With
`modality_routing: true`, one gate inspects every decided placement: when the routed model is
explicitly declared `supports_vision: false` (deployment `model_info` first, the model cost map
otherwise; unmapped names stay routable, and a multi-deployment group must accept on every
deployment), the request is re-placed on the nearest HIGHER tier holding a capable model, with
routing plugins still applied to the re-pick, then on `default_model` (never on plugin routers
and never for a plan-floored decision), and otherwise rejected with a clear 400 naming the
router. The walk only ever goes up, so a plan-mode floor cannot be undercut; a router whose only
vision model sits below the decided tier gets the 400 and an actionable message instead.
A same-tier re-pick keeps the decision's cause and adds `modality:image` to `signals`; a tier
change or default takeover records `cause: modality_escalation` with the displaced placement
(`modality_escalated_from:<TIER>` or `modality_displaced_default_model`). Escalations are never
pinned by session affinity, and a KEPT session pin bypasses the gate entirely: a session pinned
to a text-only model keeps it even when an image arrives.
### Heuristic-first chaining
`classifier_type: heuristic_first` runs the local scorer on every request and only calls the LLM

View file

@ -30,6 +30,7 @@ from litellm.constants import EMPTY_MAPPING, RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata
from litellm.litellm_core_utils.prompt_templates.common_utils import request_contains_image_content
from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.types.utils import (
@ -479,6 +480,33 @@ def _last_human_ask_index(
)
def _newest_turn_is_human_ask(
messages: Sequence[Mapping[str, object]] | None,
marker_pairs: tuple[tuple[str, str], ...] = _DEFAULT_REMINDER_MARKERS,
) -> bool:
"""Whether the request's newest turn carries a real human ask, i.e. this is a new ask rather
than an agent loop's continuation traffic.
Anchored on `_last_human_ask_index` so every surface's plumbing reads as a continuation:
chat-completions tool turns are role=tool, Messages-surface tool_result turns flatten to empty
human text, and a hybrid turn carrying an ask alongside a tool_result still counts as an ask.
Compared against the newest non-system message rather than the raw tail, because Claude Code
appends a system-role reminder after the human turn; that trailing plumbing is neither an ask
nor loop traffic and must not turn a fresh ask into a continuation. An unreadable request (no
messages) is treated as a continuation: there is no ask to classify, which is the same reading
`_extract_current_ask_and_system_prompt` gives it downstream.
"""
if not messages:
return False
newest_non_system: Final = next(
(index for index in range(len(messages) - 1, -1, -1) if messages[index].get("role") != "system"),
None,
)
if newest_non_system is None:
return False
return _last_human_ask_index(messages, marker_pairs) == newest_non_system
def _iter_system_scope_texts(
body_system: object,
messages: Sequence[Mapping[str, object]],
@ -706,11 +734,25 @@ def _decision_is_pinnable(decision: StandardLoggingRoutingDecision | None) -> bo
of the three: an agent names the conversation on its first turn, so the cheapest tier would be
the pin every session starts with, and the real work that follows would run there for the whole
TTL. It describes what that one call is, never what the session's traffic looks like.
A context-window escalation describes the prompt's size, not the session's complexity, and
size shrinks again the moment the client compacts: pinning the escalated tier would hold the
session on the big-window model long after the oversized context that forced it is gone. The
gate re-fires per request, so leaving these unpinned costs nothing but the classifier call.
A modality escalation is transient the same way: it describes what this one call carries (an
image), not what the session's traffic looks like, and pinning it would hold every following
text turn on the vision-capable model the image forced.
"""
return decision is None or decision.get("cause") not in (
"default_model_fallback",
"plan_mode",
"housekeeping",
return decision is None or (
decision.get("cause")
not in (
"default_model_fallback",
"plan_mode",
"housekeeping",
"modality_escalation",
)
and not decision.get("context_escalated")
)
@ -759,6 +801,39 @@ class ClassificationOutcome(NamedTuple):
classifier_cost: float | None = None
def _allowed(models: tuple[str, ...], fit_filter: frozenset[str] | None) -> tuple[str, ...]:
return models if fit_filter is None else tuple(model for model in models if model in fit_filter)
def _apply_context_placement(
tier: ComplexityTier | str, signals: tuple[str, ...], placement: _ContextWindowPlacement | None
) -> tuple[ComplexityTier | str, tuple[str, ...], ComplexityTier | str | None]:
"""(final tier, signals, original tier when the gate escalated, else None)."""
if placement is None:
return tier, signals, None
if _tier_name(placement.tier) == _tier_name(tier):
return placement.tier, signals, None
return placement.tier, (*signals, "context_escalation"), tier
def _window_can_hold(window: int | None, needed: int, buffer: float) -> bool:
return window is None or needed <= int(window * buffer)
def _group_provably_fits(facts: tuple[int | None, bool], needed: int, buffer: float) -> bool:
window, has_unknown = facts
return window is not None and not has_unknown and needed <= int(window * buffer)
class _ContextWindowPlacement(NamedTuple):
"""Where the context-window gate placed the request: the placement tier, the subset of its
pool the pick may use, and every configured group not provably misfit (the adaptive filter)."""
tier: ComplexityTier | str
allowed_models: tuple[str, ...]
holdable_models: frozenset[str]
class _SessionAffinityPin(NamedTuple):
model: str
tier: ComplexityTier | None
@ -1195,6 +1270,7 @@ class ComplexityRouter(CustomLogger):
classifier_cost: float | None = None,
conversation_continuing: bool = True,
tier_litellm_params: Mapping[str, object] | None = None,
context_escalation_original_tier: ComplexityTier | str | None = None,
) -> StandardLoggingRoutingDecision:
"""Assemble the per-request provenance record for this router's decision.
@ -1244,6 +1320,12 @@ class ComplexityRouter(CustomLogger):
decision["classifier_model"] = classifier_model
if classifier_cost is not None:
decision["classifier_cost"] = classifier_cost
if context_escalation_original_tier is not None:
# The pair travels together: the flag says the gate moved the request off its
# decided tier on prompt size, and the original tier names where the decision
# (classifier, keyword rule, or session pin) had placed it before physics did.
decision["context_escalated"] = True
decision["context_escalation_original_tier"] = _tier_name(context_escalation_original_tier)
if tier_litellm_params:
masked_tier_litellm_params: Final = mask_credentials_in_payload(tier_litellm_params)
if isinstance(masked_tier_litellm_params, Mapping):
@ -1644,7 +1726,7 @@ class ComplexityRouter(CustomLogger):
return entry.litellm_params if entry is not None else MappingProxyType({})
@staticmethod
def _pick_from_tier_value(model: str | list[str], tier_key: str) -> str:
def _pick_from_tier_value(model: str | Sequence[str], tier_key: str) -> str:
if isinstance(model, str):
return model
if not model:
@ -1660,15 +1742,21 @@ class ComplexityRouter(CustomLogger):
raw_messages: list[dict[str, Any]] | None,
resolved_messages: list[dict[str, Any]] | None,
request_kwargs: dict,
allowed_models: tuple[str, ...] | None = None,
) -> str:
if not self.config.plugins:
if allowed_models is not None:
return self._pick_from_tier_value(allowed_models, _tier_name(tier))
return self.get_model_for_tier(tier)
from litellm.types.router import RoutingContext
tier_key: Final = _tier_name(tier)
metadata_key: Final = get_metadata_variable_name_from_kwargs(request_kwargs)
pool: Final = tuple(self._tier_pools().get(tier_key, ()))
full_pool: Final = tuple(self._tier_pools().get(tier_key, ()))
pool: Final = (
tuple(model for model in full_pool if model in allowed_models) if allowed_models is not None else full_pool
)
if not pool:
# Nothing for the plugins to filter. Falling through would raise the
# plugin-filtering error below and send the operator hunting for a policy
@ -1762,6 +1850,7 @@ class ComplexityRouter(CustomLogger):
request_kwargs: dict[str, Any] | None = None,
hard_floor: ComplexityTier | str | None = None,
hard_ceiling: ComplexityTier | str | None = None,
fit_filter: frozenset[str] | None = None,
) -> str:
"""hard_floor excludes every candidate whose tiers all sit below it, turning this pick's
soft floors (a distance penalty a high-scoring cheap model can outweigh) into a hard
@ -1774,7 +1863,10 @@ class ComplexityRouter(CustomLogger):
tier because that is all it is worth, so a bandit trading cost for quality has nothing to
win and must not reach above it. Without it the distance penalty is the only thing holding
the tier, and a deployment that lowers tier_distance_penalty silently gets the expensive
model back while the routing decision still reads as the cheapest tier."""
model back while the routing decision still reads as the cheapest tier.
fit_filter excludes candidates the context-window gate proved cannot hold the prompt,
in every phase including cold start and the tier fallbacks."""
from litellm.router_strategy.adaptive_router.bandit import (
normalized_cost,
thompson_sample,
@ -1785,12 +1877,12 @@ class ComplexityRouter(CustomLogger):
if adaptive is None or not isinstance(classified_tier, ComplexityTier):
# Custom tier names have no severity index; adaptive is rejected alongside
# tier_definitions, so this guard is the contract for any future caller.
return self.get_model_for_tier(classified_tier)
return self._fitting_tier_fallback(classified_tier, fit_filter)
request_type: Final = classify_prompt(user_message)
classified_idx: Final = TIER_SEVERITY_ORDER.index(classified_tier)
pools: Final = self._tier_pools()
classified_candidates: Final = tuple(pools.get(_tier_name(classified_tier), ()))
classified_candidates: Final = _allowed(tuple(pools.get(_tier_name(classified_tier), ())), fit_filter)
cold_start_candidates: Final = tuple(
model for model in classified_candidates if adaptive._cells[(request_type, model)].total_samples == 0
)
@ -1820,9 +1912,9 @@ class ComplexityRouter(CustomLogger):
if self.config.adaptive_eligible == "classified_tier":
candidates = list(classified_candidates)
if not candidates:
return self.get_model_for_tier(classified_tier)
return self._fitting_tier_fallback(classified_tier, fit_filter)
else:
candidates = list(adaptive.config.available_models)
candidates = list(_allowed(tuple(adaptive.config.available_models), fit_filter))
all_costs: Final = [adaptive.model_to_cost.get(m, 0.0) for m in candidates]
quality_weight: Final = self.config.adaptive_weights.quality
@ -1869,7 +1961,7 @@ class ComplexityRouter(CustomLogger):
best_score = score
best_model = model
if best_model is None:
return self.get_model_for_tier(classified_tier)
return self._fitting_tier_fallback(classified_tier, fit_filter)
if request_kwargs is not None:
metadata = request_kwargs.setdefault("metadata", {})
if isinstance(metadata, dict):
@ -1886,6 +1978,12 @@ class ComplexityRouter(CustomLogger):
}
return best_model
def _fitting_tier_fallback(self, classified_tier: ComplexityTier | str, fit_filter: frozenset[str] | None) -> str:
fitting: Final = _allowed(tuple(self._tier_pools().get(_tier_name(classified_tier), ())), fit_filter)
if fit_filter is not None and fitting:
return self._pick_from_tier_value(fitting, _tier_name(classified_tier))
return self.get_model_for_tier(classified_tier)
def _resolve_plan_mode_floor(self) -> ComplexityTier | str | None:
"""The configured floor as an active tier: the built-in enum member, or the defined
name itself for a custom tier set; None when the feature is off."""
@ -1956,6 +2054,163 @@ class ComplexityRouter(CustomLogger):
return None
return name if self.config.has_custom_tiers else ComplexityTier(name)
def _deployment_window(self, group: str, deployment: Mapping[str, object]) -> int | None:
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
deployment_model_info: Final = deployment.get("model_info")
declared: Final = (
deployment_model_info.get("max_input_tokens") if isinstance(deployment_model_info, Mapping) else None
)
if isinstance(declared, int):
return declared
litellm_params: Final = deployment.get("litellm_params")
params: Final = litellm_params if isinstance(litellm_params, Mapping) else EMPTY_MAPPING
provider_override: Final = params.get("custom_llm_provider")
# get_router_model_info resolves the provider, and get_llm_provider runs the OAuth device
# flow for github_copilot/chatgpt, so a metadata question must never reach it for those.
if declared_authenticating_provider(
str(params.get("model") or ""), provider_override if isinstance(provider_override, str) else None
):
return None
try:
model_info: Final = self.litellm_router_instance.get_router_model_info(
deployment=cast(dict, deployment), # cast-ok: router deployments are plain dicts
received_model_name=group,
)
window: Final = model_info.get("max_input_tokens")
except Exception: # noqa: BLE001 # best-effort: an unmappable deployment must not hide the others
return None
return window if isinstance(window, int) else None
def _group_window_facts(self, group: str) -> tuple[int | None, bool]:
"""(smallest declared context window across the group's deployments, whether any deployment
declares none). The core router picks a deployment within the group without a fit check, so
the group is only as safe as its smallest member."""
list_models: Final = getattr(self.litellm_router_instance, "get_model_list", None)
deployments: Final = list_models(model_name=group) if callable(list_models) else None
if not isinstance(deployments, list) or not deployments:
return (None, True)
windows: Final = tuple(
window for deployment in deployments if (window := self._deployment_window(group, deployment)) is not None
)
return (min(windows) if windows else None, len(windows) < len(deployments))
@staticmethod
def _out_of_band_request_text(request_kwargs: Mapping[str, object]) -> str:
"""Prompt content the resolved message list never carries: the Responses API's
`instructions`, the /v1/messages top-level `system` block, and tool definitions.
A coding agent's context is dominated by these."""
import json
instructions: Final = request_kwargs.get("instructions")
proxy_request: Final = request_kwargs.get("proxy_server_request")
body: Final = proxy_request.get("body") if isinstance(proxy_request, Mapping) else None
system: Final = body.get("system") if isinstance(body, Mapping) else None
tools: Final = (
body.get("tools") if isinstance(body, Mapping) and body.get("tools") else request_kwargs.get("tools")
)
tools_text = ""
if tools:
try:
tools_text = json.dumps(tools, default=str)
except (TypeError, ValueError):
tools_text = str(tools)
return (
(instructions if isinstance(instructions, str) else "")
+ (str(system) if system is not None else "")
+ tools_text
)
def _request_byte_upper_bound(
self, resolved_messages: Sequence[Mapping[str, object]] | None, request_kwargs: Mapping[str, object]
) -> int:
"""UTF-8 byte length of all prompt content. BPE emits at least one byte per token in every
script, so the token count never exceeds this and 'bytes fit' soundly skips counting."""
content_bytes: Final = sum(len(str(m.get("content") or "").encode()) for m in resolved_messages or ())
return content_bytes + len(self._out_of_band_request_text(request_kwargs).encode())
async def _counted_request_tokens(
self, resolved_messages: Sequence[Mapping[str, object]], request_kwargs: Mapping[str, object]
) -> int | None:
"""Real-tokenizer count of the resolved messages plus the out-of-band carriers, off the
event loop; None when counting fails, and the gate then leaves the placement alone."""
import litellm
from litellm.litellm_core_utils.asyncify import asyncify
out_of_band: Final = self._out_of_band_request_text(request_kwargs)
try:
counted: Final = await asyncify(litellm.token_counter)(
messages=cast(list, resolved_messages) # cast-ok: token_counter only iterates the sequence
)
return counted + (await asyncify(litellm.token_counter)(text=out_of_band) if out_of_band else 0)
except Exception as e: # noqa: BLE001 # best-effort: an uncountable prompt must not fail the request
verbose_router_logger.debug("ComplexityRouter: context-window token count failed. Got - %s", e)
return None
async def _context_window_placement(
self,
tier: ComplexityTier | str,
resolved_messages: Sequence[Mapping[str, object]] | None,
request_kwargs: Mapping[str, object],
pool_override: tuple[str, ...] | None = None,
) -> _ContextWindowPlacement | None:
"""Correct a decided placement whose models provably cannot hold the prompt, or None
(the placement stands). Only a real tokenizer count ever moves a request, escalation
lands only on groups whose every deployment declares a fitting window, and a group
with no resolvable window is never moved on faith in either direction."""
if not self.config.enable_context_window_escalation or not resolved_messages:
return None
pools: Final = self._tier_pools()
pool: Final = pool_override if pool_override is not None else tuple(pools.get(_tier_name(tier), ()))
if not pool:
return None
facts: Final = MappingProxyType({group: self._group_window_facts(group) for group in pool})
known_windows: Final = tuple(window for window, _ in facts.values() if window is not None)
if not known_windows:
return None
buffer: Final = self.config.context_window_escalation_buffer
if self._request_byte_upper_bound(resolved_messages, request_kwargs) <= int(min(known_windows) * buffer):
return None
needed: Final = await self._counted_request_tokens(resolved_messages, request_kwargs)
if needed is None:
return None
return self._placement_for_tokens(tier=tier, pool=pool, pools=pools, facts=facts, needed=needed)
def _placement_for_tokens(
self,
*,
tier: ComplexityTier | str,
pool: tuple[str, ...],
pools: Mapping[str, list[str]],
facts: Mapping[str, tuple[int | None, bool]],
needed: int,
) -> _ContextWindowPlacement | None:
buffer: Final = self.config.context_window_escalation_buffer
in_tier: Final = tuple(group for group in pool if _window_can_hold(facts[group][0], needed, buffer))
if in_tier and len(in_tier) == len(pool):
return None
holdable: Final = frozenset(
group
for tier_pool in pools.values()
for group in tier_pool
if _window_can_hold(self._group_window_facts(group)[0], needed, buffer)
)
if in_tier:
return _ContextWindowPlacement(tier=tier, allowed_models=in_tier, holdable_models=holdable)
for name in self.config.tier_names()[self._active_tier_severity(tier) + 1 :]:
proven = tuple(
group
for group in pools.get(name, ())
if _group_provably_fits(self._group_window_facts(group), needed, buffer)
)
if proven:
return _ContextWindowPlacement(
tier=name if self.config.has_custom_tiers else ComplexityTier(name),
allowed_models=proven,
holdable_models=holdable,
)
return None
def _apply_plan_mode_floor(self, tier: ComplexityTier | str) -> ComplexityTier | str:
"""The higher of the decided tier and the plan-mode floor; identity when the floor is unset."""
floor: Final = self._resolve_plan_mode_floor()
@ -2025,6 +2280,175 @@ class ComplexityRouter(CustomLogger):
return pinned_model
return self.get_model_for_tier(escalated_tier)
def _model_accepts_image_input(self, model_name: str) -> bool:
"""Whether a routed model or pool entry can serve an image request.
Resolved through the deployments that would actually serve the name; a name with no
deployment on the router is served by the SDK directly and is checked against the model
cost map itself. Only an explicit supports_vision false excludes, a deployment-level
model_info override first and the map otherwise, so unmapped custom names stay routable.
A multi-deployment group must accept on EVERY deployment: the router picks a deployment
inside the group after this gate runs, so a mixed group marked eligible could still hand
the image to its text-only member and fail with the exact 400 the gate exists to prevent.
"""
from litellm.utils import is_vision_explicitly_disabled
def deployment_accepts(deployment: Mapping[str, Any]) -> bool:
declared: Final = (deployment.get("model_info") or EMPTY_MAPPING).get("supports_vision")
if declared is not None:
return declared is True
litellm_model: Final = (deployment.get("litellm_params") or EMPTY_MAPPING).get("model") or model_name
return not is_vision_explicitly_disabled(litellm_model)
deployments: Final = self.litellm_router_instance.get_model_list(model_name=model_name)
if not deployments:
return not is_vision_explicitly_disabled(model_name)
return all(deployment_accepts(deployment) for deployment in deployments)
def _modality_eligible_models(self) -> frozenset[str]:
"""Every configured pool entry, plus default_model, that can serve an image request."""
names: Final = frozenset(entry for pool in self._tier_pools().values() for entry in pool) | frozenset(
name for name in (self.config.default_model,) if name
)
return frozenset(name for name in names if self._model_accepts_image_input(name))
async def _gate_response_modality(
self,
response: PreRoutingHookResponse,
messages: list[dict[str, Any]] | None, # mutable-ok: forwarded verbatim to the list-typed re-pick
resolved_messages: Sequence[Mapping[str, object]] | None,
request_kwargs: dict, # mutable-ok: same shape the hook receives
) -> PreRoutingHookResponse:
"""Replace a routed model that cannot accept this request's image input.
The single modality owner, applied to the decided response at the hook's exits so every
routing path is covered uniformly. A KEPT session pin is exempt by design (its cause);
replacement picks and every other path are just responses. The re-placement walks
UPWARD-ONLY from the decision's tier (so a plan-mode floor can never be undercut), picks
through `_pick_model_for_tier` so routing plugins still apply, then falls to
default_model (never on plugin routers, and never on a plan-floored decision, since
default_model carries no tier guarantee), else raises the clear 400. The rewritten
decision keeps its cause on a same-tier repick and becomes modality_escalation when the
tier moved or default_model took over, with the displaced placement in signals.
"""
decision: Final = response.routing_decision
if (
not self.config.modality_routing
or not resolved_messages
or response.model is None
or (decision is not None and decision.get("cause") == "session_affinity_pin")
or not request_contains_image_content(resolved_messages)
or self._model_accepts_image_input(response.model)
):
return response
eligible: Final = self._modality_eligible_models()
names: Final = self.config.tier_names()
pools: Final = self._tier_pools()
decided: Final = decision.get("tier") if decision is not None else None
start: Final = names.index(decided) if isinstance(decided, str) and decided in names else 0
capable: Final = next(
(name for name in names[start:] if any(entry in eligible for entry in pools.get(name, ()))), None
)
if capable is not None:
new_tier: ComplexityTier | str | None = capable if self.config.has_custom_tiers else ComplexityTier(capable)
repick_messages: Final = list(resolved_messages) # mutable-ok: the pick's param is list-typed
new_model = await self._pick_model_for_tier(
new_tier,
messages,
repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them
request_kwargs,
allowed_models=tuple(entry for entry in pools.get(capable, ()) if entry in eligible),
)
elif self._modality_default_model_usable(request_kwargs, resolved_messages, eligible):
new_tier = None
new_model = self._placed_default_model()
else:
import litellm
raise litellm.BadRequestError(
message=(
f"Auto-router {self.model_name} received a request with image input, but no model "
f"at or above the decided tier accepts images and modality_routing is enabled. "
f"Tiers checked: {', '.join(names[start:])}. Add a vision-capable model to a tier, "
f"or set a vision-capable default_model, or remove the image content."
),
model=self.model_name,
llm_provider="",
)
self._restamp_adaptive_choice(request_kwargs, response.model, new_model)
same_tier: Final = capable is not None and decided == capable
base_cause: Final = (decision.get("cause") if decision is not None else None) or "default_fallback"
displaced_default: Final = decided is None and response.model == self.config.default_model
markers: Final = (
"modality:image",
*((f"modality_escalated_from:{decided}",) if not same_tier and isinstance(decided, str) else ()),
*(("modality_displaced_default_model",) if not same_tier and displaced_default else ()),
)
old_signals: Final = tuple(decision.get("signals") or ()) if decision is not None else ()
new_decision: Final = self._build_routing_decision(
routed_model=new_model,
cause=base_cause if same_tier else "modality_escalation",
tier=new_tier,
score=decision.get("score") if decision is not None else None,
signals=(*old_signals, *markers),
matched_keyword=decision.get("matched_keyword") if decision is not None else None,
escalation_keyword=decision.get("escalation_keyword") if decision is not None else None,
escalated=bool(decision.get("escalated", False)) if decision is not None else False,
classifier_model=decision.get("classifier_model") if decision is not None else None,
classifier_cost=decision.get("classifier_cost") if decision is not None else None,
conversation_continuing=bool(decision.get("conversation_continuing", True))
if decision is not None
else True,
tier_litellm_params=self._litellm_params_for_model(new_tier, new_model),
context_escalation_original_tier=(
decision.get("context_escalation_original_tier") if decision is not None else None
),
)
from litellm.types.router import PreRoutingHookResponse as HookResponse
return HookResponse(
model=new_model,
messages=response.messages,
litellm_params=self._litellm_params_for_model(new_tier, new_model),
routing_decision=new_decision,
)
def _modality_default_model_usable(
self,
request_kwargs: Mapping[str, object],
resolved_messages: Sequence[Mapping[str, object]] | None,
eligible: frozenset[str],
) -> bool:
"""default_model may serve a gated request only when it is configured, plugin-free
(it is never checked against the plugin pipeline), capability-eligible, and the turn
carries no plan-mode sentinel. The sentinel is re-detected here rather than read off
the decision record, because the record only marks turns the floor RAISED; a sentinel
turn already at or above the floor keeps its ordinary cause, and default_model carries
no tier the floor could vouch for on any sentinel turn."""
return (
bool(self.config.default_model)
and not self.config.plugins
and self.config.default_model in eligible
and self._matched_plan_mode_signal(request_kwargs, resolved_messages) is None
)
def _placed_default_model(self) -> str:
"""The default_model behind a usable-default verdict; the raise is the type-level
proof, not a reachable path."""
model: Final = self.config.default_model
if model is None:
raise ValueError(f"Auto-router {self.model_name}: modality gate routed to an unset default_model")
return model
@staticmethod
def _restamp_adaptive_choice(request_kwargs: Mapping[str, object], old_model: str, new_model: str) -> None:
"""The adaptive feedback loop reads its chosen-model marker from request metadata; a
gate rewrite must move the marker with the model or rewards land on the displaced one."""
metadata: Final = request_kwargs.get("metadata")
if isinstance(metadata, dict) and metadata.get("adaptive_router_chosen_model") == old_model:
metadata["adaptive_router_chosen_model"] = new_model
def _lexical_tier_override(self, user_message: str) -> KeywordOverride | None:
"""When keyword_tier_rules match literally, the most-severe matched tier wins.
@ -2247,14 +2671,18 @@ class ComplexityRouter(CustomLogger):
@property
def _uses_tier_pin(self) -> bool:
return bool(self.config.session_affinity and not self.config.plugins)
"""classification_mode 'user_turn' implies the tier pin machinery: the pin write after each
pinnable classification is what gives a continuation a held decision to replay."""
return bool(
(self.config.session_affinity or self.config.classification_mode == "user_turn") and not self.config.plugins
)
@property
def _uses_deployment_pin(self) -> bool:
"""session_affinity implies the deployment pin: a session frozen onto one model
"""The tier pin implies the deployment pin: a session frozen onto one model
group but load-balanced across its deployments would still go cache-cold, which
is the exact failure both flags exist to prevent."""
return bool((self.config.deployment_affinity or self.config.session_affinity) and not self.config.plugins)
return bool(self.config.deployment_affinity and not self.config.plugins) or self._uses_tier_pin
def _with_session_deployment_affinity(
self, response: PreRoutingHookResponse | None
@ -2282,6 +2710,11 @@ class ComplexityRouter(CustomLogger):
pins the model chosen on the session's first turn and reuses it for every later
turn, skipping classification entirely. Otherwise delegates to `_classify_and_route`.
When `classification_mode` is 'user_turn', the same pin is replayed only on
continuation turns (an agent loop's tool traffic); a new human ask always falls
through to classification, so the session can still move tiers between asks.
With both knobs on, session_affinity's pin-first behavior wins.
Skipped entirely when `plugins` are configured: reusing a stale pin would bypass
the plugin pipeline on every turn after the first, since a pinned model was never
re-checked against a policy plugin whose decision can change between turns (e.g. a
@ -2305,7 +2738,13 @@ class ComplexityRouter(CustomLogger):
session_id: Final = self._get_session_id_from_request_kwargs(request_kwargs) if use_session_affinity else None
cache_key = self._get_session_affinity_cache_key(session_id, request_kwargs) if session_id is not None else None
if cache_key is not None:
# In 'user_turn' mode a held pin is replayed only on continuation turns; a new human
# ask falls through and re-classifies. session_affinity restores pin-first for asks too.
pin_replay_allowed: Final = bool(self.config.session_affinity) or not _newest_turn_is_human_ask(
resolved_messages, self._reminder_markers
)
if cache_key is not None and pin_replay_allowed:
pinned_value: Final = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
pinned_pin: Final = _parse_session_affinity_pin(pinned_value)
if pinned_pin is not None:
@ -2339,6 +2778,26 @@ class ComplexityRouter(CustomLogger):
session_model: Final = routed_model
if plan_floored and pinned_tier is not None:
routed_model = self.get_model_for_tier(self._apply_plan_mode_floor(pinned_tier))
pin_source_tier: Final = self._tier_for_model(routed_model)
pin_placement: Final = (
await self._context_window_placement(
pin_source_tier, resolved_messages, request_kwargs, pool_override=(routed_model,)
)
if pin_source_tier is not None
else None
)
pin_context_original_tier: Final = (
pin_source_tier
if pin_placement is not None
and pin_source_tier is not None
and _tier_name(pin_placement.tier) != _tier_name(pin_source_tier)
else None
)
if pin_placement is not None and pin_context_original_tier is not None:
# The stored pin below keeps the session's own model on purpose.
routed_model = self._pick_from_tier_value(
pin_placement.allowed_models, _tier_name(pin_placement.tier)
)
# Refresh the TTL on every hit so an active session doesn't lose its
# pin mid-conversation just because it outlives the original write.
await self.litellm_router_instance.cache.async_set_cache(
@ -2354,36 +2813,47 @@ class ComplexityRouter(CustomLogger):
kwargs_metadata: Final = request_kwargs.setdefault("metadata", {})
if isinstance(kwargs_metadata, dict):
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = routed_model
replay_cause: Final[RoutingDecisionCause] = (
"session_affinity_pin" if self.config.session_affinity else "user_turn_continuation"
)
cause: RoutingDecisionCause = (
"plan_mode"
if plan_floored
else ("session_affinity_escalation" if escalated else "session_affinity_pin")
"plan_mode" if plan_floored else ("session_affinity_escalation" if escalated else replay_cause)
)
verbose_router_logger.info(
"ComplexityRouter: routing decision cause=%s, routed_model=%s", cause, routed_model
)
routed_pin_tier: Final = self._tier_for_model(routed_model) if plan_floored else resolved_pin_tier
routed_pin_tier: Final = (
pin_placement.tier
if pin_placement is not None and pin_context_original_tier is not None
else (self._tier_for_model(routed_model) if plan_floored else resolved_pin_tier)
)
session_tier_litellm_params: Final = self._litellm_params_for_model(routed_pin_tier, routed_model)
has_original_messages: Final = messages is not None and len(messages) > 0
return self._with_session_deployment_affinity(
PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=session_tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
cause=cause,
tier=routed_pin_tier,
matched_keyword=pin_plan_sentinel if plan_floored else None,
escalation_keyword=pin_escalation_keyword,
escalated=escalated,
conversation_continuing=conversation_continuing,
tier_litellm_params=session_tier_litellm_params,
await self._gate_response_modality(
PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=session_tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
cause=cause,
tier=routed_pin_tier,
matched_keyword=pin_plan_sentinel if plan_floored else None,
escalation_keyword=pin_escalation_keyword,
escalated=escalated,
conversation_continuing=conversation_continuing,
tier_litellm_params=session_tier_litellm_params,
context_escalation_original_tier=pin_context_original_tier,
),
),
messages,
resolved_messages,
request_kwargs,
)
)
response: Final = await self._classify_and_route(
routed_response: Final = await self._classify_and_route(
model=model,
request_kwargs=request_kwargs,
messages=messages,
@ -2392,6 +2862,11 @@ class ComplexityRouter(CustomLogger):
conversation_continuing=conversation_continuing,
resolved_messages=resolved_messages,
)
response: Final = (
await self._gate_response_modality(routed_response, messages, resolved_messages, request_kwargs)
if routed_response is not None
else None
)
# Sentinel presence, not the plan_mode cause, gates the pin write: a plan-mode turn
# classified at or above the floor keeps its ordinary cause, yet on an adaptive router
# the hard floor constrained its pick, so pinning it would carry a plan-mode-shaped
@ -2573,6 +3048,8 @@ class ComplexityRouter(CustomLogger):
plan_floored: Final = tier != pre_floor_tier
if plan_floored:
signals = (*signals, "plan_mode_floor")
context_placement: Final = await self._context_window_placement(tier, resolved_messages, request_kwargs)
tier, signals, context_original_tier = _apply_context_placement(tier, signals, context_placement)
score_repr: Final = f"{score:.3f}" if score is not None else "n/a"
fallback_model: Final = self.config.default_model if not self.config.plugins else None
# A sentinel-carrying request skips the failure exit below, whether or not the floor
@ -2619,8 +3096,15 @@ class ComplexityRouter(CustomLogger):
# the cheapest tier would then contradict the floor and bound the pick below the tier
# the decision reports.
housekeeping_ceiling: Final = tier if outcome.cause == "housekeeping" else None
# A context-escalated tier becomes the hard floor: a floor the bandit can slide
# under is not a floor.
routed_model = self._soft_floor_pick(
tier, user_message, request_kwargs, hard_floor=plan_floor, hard_ceiling=housekeeping_ceiling
tier,
user_message,
request_kwargs,
hard_floor=tier if context_original_tier is not None else plan_floor,
hard_ceiling=housekeeping_ceiling,
fit_filter=context_placement.holdable_models if context_placement is not None else None,
)
adaptive: Final = self._ensure_adaptive_router()
if adaptive is not None:
@ -2637,7 +3121,13 @@ class ComplexityRouter(CustomLogger):
routed_model,
)
else:
routed_model = await self._pick_model_for_tier(tier, messages, resolved_messages, request_kwargs)
routed_model = await self._pick_model_for_tier(
tier,
messages,
resolved_messages,
request_kwargs,
allowed_models=context_placement.allowed_models if context_placement is not None else None,
)
verbose_router_logger.info(
"ComplexityRouter: routing decision cause=%s, tier=%s, score=%s, signals=%s, routed_model=%s",
outcome.cause,
@ -2690,5 +3180,6 @@ class ComplexityRouter(CustomLogger):
classifier_model=classifier_model,
classifier_cost=outcome.classifier_cost,
tier_litellm_params=tier_litellm_params,
context_escalation_original_tier=context_original_tier,
),
)

View file

@ -823,6 +823,44 @@ class ComplexityRouterConfig(BaseModel):
),
)
enable_context_window_escalation: bool = Field(
default=True,
description=(
"Escalate a request off a tier whose models provably cannot hold its prompt, before "
"dispatch. The classifier scores complexity and never prompt size, so a long agentic "
"session whose newest ask is trivial lands on a small-window tier and the provider "
"rejects it with a context-window 400 that nothing retries. When every model of the "
"decided tier has a declared window smaller than the estimated prompt, the request "
"moves to the lowest configured tier with a model whose declared window fits; when "
"only some of the tier's models fit, the pick is restricted to those and the tier "
"keeps the request. Models with no resolvable window are never escalated away from "
"and never escalated onto. Set false to dispatch on complexity alone, as before."
),
)
context_window_escalation_buffer: float = Field(
default=0.95,
gt=0,
le=1,
description=(
"Fraction of a model's declared context window the estimated prompt must fit within. "
"The token count is an estimate, so fitting against the full window would dispatch "
"prompts that the provider's own tokenizer then rejects; 0.95 leaves room for that "
"drift plus the response tokens."
),
)
modality_routing: bool = Field(
default=False,
description=(
"Route image-bearing requests only to models that can accept image input. The "
"classifier reads text alone, so an image request whose text classifies cheap "
"otherwise lands on a text-only model and fails with a provider 400. When enabled, "
"a routed model explicitly declared supports_vision false (deployment model_info "
"or the model cost map; unmapped names stay routable) is replaced by the nearest "
"HIGHER tier holding a capable model, then default_model, else a clear 400. A kept "
"session-affinity pin still wins even when an image arrives."
),
)
# Semantic (embedding) matching for keyword_tier_rules instead of literal text matching
semantic_keyword_matching: bool = Field(
default=False,
@ -839,6 +877,21 @@ class ComplexityRouterConfig(BaseModel):
description="Minimum cosine similarity for a semantic keyword match",
)
classification_mode: Literal["every_request", "user_turn"] = Field(
default="every_request",
description=(
"When to run the complexity classifier. 'every_request' (the default) classifies every "
"inference request, including the tool-result continuation turns of an agentic loop. "
"'user_turn' classifies only requests whose newest turn is a new human ask and replays "
"the session's held routing decision on continuation turns, which cuts classifier "
"spend and eliminates mid-loop model switches. Continuations with no held decision to "
"replay (no resolvable session_id, expired pin, fresh restart) still classify. Unlike "
"session_affinity, a new human ask always re-classifies, so a session can still move "
"tiers between asks. Suppressed when plugins are configured, for the same reason "
"session_affinity is: a replayed decision would bypass the plugin pipeline."
),
)
# Session affinity: pin the first turn's routed model for the rest of the session
session_affinity: bool = Field(
default=False,

View file

@ -2162,6 +2162,42 @@ class OpenAIRealtimeDoneEvent(TypedDict):
type: Literal["response.done"]
class OpenAIRealtimeInputAudioBufferSpeechEvent(TypedDict):
type: ReadOnly[Literal["input_audio_buffer.speech_started", "input_audio_buffer.speech_stopped"]]
event_id: ReadOnly[str]
item_id: ReadOnly[str]
class OpenAIRealtimeInputAudioTranscriptionDelta(TypedDict):
type: ReadOnly[Literal["conversation.item.input_audio_transcription.delta"]]
event_id: ReadOnly[str]
item_id: ReadOnly[str]
content_index: ReadOnly[int]
delta: ReadOnly[str]
class OpenAIRealtimeInputAudioTranscriptionCompleted(TypedDict):
type: ReadOnly[Literal["conversation.item.input_audio_transcription.completed"]]
event_id: ReadOnly[str]
item_id: ReadOnly[str]
content_index: ReadOnly[int]
transcript: ReadOnly[str]
class OpenAIRealtimeUsageTokenDetails(TypedDict):
audio_tokens: ReadOnly[int]
text_tokens: ReadOnly[int]
cached_tokens: NotRequired[ReadOnly[int]]
class OpenAIRealtimeResponseUsage(TypedDict):
input_tokens: ReadOnly[int]
output_tokens: ReadOnly[int]
total_tokens: ReadOnly[int]
input_token_details: NotRequired[ReadOnly[OpenAIRealtimeUsageTokenDetails]]
output_token_details: NotRequired[ReadOnly[OpenAIRealtimeUsageTokenDetails]]
class OpenAIRealtimeEventTypes(Enum):
SESSION_CREATED = "session.created"
# Beta delta event names
@ -2199,6 +2235,9 @@ OpenAIRealtimeEvents = (
| OpenAIRealtimeOutputItemDone
| OpenAIRealtimeFunctionCallArgumentsDone
| OpenAIRealtimeDoneEvent
| OpenAIRealtimeInputAudioBufferSpeechEvent
| OpenAIRealtimeInputAudioTranscriptionDelta
| OpenAIRealtimeInputAudioTranscriptionCompleted
)
OpenAIRealtimeStreamList = list[OpenAIRealtimeEvents]

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