mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge litellm_internal_staging into fix/bedrock-guardrail-image-input
Takes upstream's _image_sources docstring, which landed via #38940.
This commit is contained in:
commit
124811085a
214 changed files with 13649 additions and 2510 deletions
23
.github/actions/cache-cargo-build/action.yml
vendored
23
.github/actions/cache-cargo-build/action.yml
vendored
|
|
@ -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-
|
||||
|
|
|
|||
230
.github/scripts/close_duplicate_issues.py
vendored
230
.github/scripts/close_duplicate_issues.py
vendored
|
|
@ -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()
|
||||
69
.github/workflows/auto-close-duplicates.yml
vendored
Normal file
69
.github/workflows/auto-close-duplicates.yml
vendored
Normal 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 }}
|
||||
40
.github/workflows/check_duplicate_issues.yml
vendored
40
.github/workflows/check_duplicate_issues.yml
vendored
|
|
@ -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.
|
||||
|
|
|
|||
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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`
|
||||
|
|
|
|||
15
Dockerfile
15
Dockerfile
|
|
@ -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}" \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}" \
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -86,6 +86,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/comprehendmedical",
|
||||
"/cohere/",
|
||||
"/gemini/",
|
||||
"/gigachat/",
|
||||
"/google/",
|
||||
"/vertex_ai/",
|
||||
"/vertex-ai/",
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
@ -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])
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -1728,6 +1751,7 @@ SENTRY_DENYLIST: Final = [
|
|||
"jwt_token",
|
||||
"private_key",
|
||||
"SLACK_WEBHOOK_URL",
|
||||
"ALERTING_WEBHOOK_URL",
|
||||
"webhook_url",
|
||||
"LANGFUSE_SECRET_KEY",
|
||||
# Email Configuration
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -1485,9 +1485,9 @@ Model Info:
|
|||
elif self.default_webhook_url is not None:
|
||||
_digest_webhook = self.default_webhook_url
|
||||
else:
|
||||
_digest_webhook = os.getenv("SLACK_WEBHOOK_URL", None)
|
||||
_digest_webhook = os.getenv("SLACK_WEBHOOK_URL") or os.getenv("ALERTING_WEBHOOK_URL")
|
||||
if _digest_webhook is None:
|
||||
raise ValueError("Missing SLACK_WEBHOOK_URL from environment")
|
||||
raise ValueError("Missing SLACK_WEBHOOK_URL / ALERTING_WEBHOOK_URL from environment")
|
||||
|
||||
digest_key: Final = f"{alert_type_name_str}:{request_model or ''}:{api_base or ''}"
|
||||
|
||||
|
|
@ -1516,10 +1516,10 @@ Model Info:
|
|||
elif self.default_webhook_url is not None:
|
||||
slack_webhook_url = self.default_webhook_url
|
||||
else:
|
||||
slack_webhook_url = os.getenv("SLACK_WEBHOOK_URL", None)
|
||||
slack_webhook_url = os.getenv("SLACK_WEBHOOK_URL") or os.getenv("ALERTING_WEBHOOK_URL")
|
||||
|
||||
if slack_webhook_url is None:
|
||||
raise ValueError("Missing SLACK_WEBHOOK_URL from environment")
|
||||
raise ValueError("Missing SLACK_WEBHOOK_URL / ALERTING_WEBHOOK_URL from environment")
|
||||
payload: Final = {"text": formatted_message}
|
||||
headers: Final = {"Content-type": "application/json"}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")),
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -861,21 +861,9 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
def _image_sources(block: Mapping[str, object]) -> tuple[str, ...]:
|
||||
"""Normalize an Anthropic image block into strings a guardrail can read.
|
||||
|
||||
`source` is one of three shapes (`AnthropicMessagesImageParam.source`):
|
||||
|
||||
{"type": "base64", "media_type": "image/png", "data": "<b64>"}
|
||||
{"type": "url", "url": "https://..."}
|
||||
{"type": "file", "file_id": "..."}
|
||||
|
||||
base64 is returned as a data URI rather than the bare payload: consumers of
|
||||
``GenericGuardrailAPIInputs["images"]`` otherwise have no way to know the
|
||||
format, and an API like Bedrock's ApplyGuardrail requires it. url is passed
|
||||
through so the consumer can fetch it under its own SSRF policy.
|
||||
|
||||
file is not resolvable here (the bytes live behind the Files API), so it
|
||||
yields nothing. That is a silent gap for any consumer that treats a missing
|
||||
entry as "no image to scan"; scanning a file_id needs a fetch this extractor
|
||||
has no client for.
|
||||
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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ GigaChat Chat Module
|
|||
from .streaming import GigaChatModelResponseIterator
|
||||
from .transformation import GigaChatConfig, GigaChatError
|
||||
|
||||
__all__ = [
|
||||
__all__ = (
|
||||
"GigaChatConfig",
|
||||
"GigaChatError",
|
||||
"GigaChatModelResponseIterator",
|
||||
]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
7
litellm/llms/gigachat/passthrough/__init__.py
Normal file
7
litellm/llms/gigachat/passthrough/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""
|
||||
GigaChat passthrough Module
|
||||
"""
|
||||
|
||||
from .transformation import GigaChatPassthroughConfig
|
||||
|
||||
__all__ = ("GigaChatPassthroughConfig",)
|
||||
213
litellm/llms/gigachat/passthrough/transformation.py
Normal file
213
litellm/llms/gigachat/passthrough/transformation.py
Normal 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))
|
||||
26
litellm/llms/gigachat/utils.py
Normal file
26
litellm/llms/gigachat/utils.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
100
litellm/llms/openai/workload_identity.py
Normal file
100
litellm/llms/openai/workload_identity.py
Normal 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
|
||||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
||||
|
|
@ -2541,7 +2545,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
)
|
||||
alerting: list | None = Field(
|
||||
None,
|
||||
description="List of alerting integrations. Today, just slack - `alerting: ['slack']`",
|
||||
description="List of alerting integrations - e.g. `alerting: ['slack', 'webhook', 'email']`. 'slack' posts Slack-format messages to any Slack-compatible webhook (Slack, Rocket.Chat, Mattermost); 'webhook' posts structured JSON budget alerts to WEBHOOK_URL",
|
||||
)
|
||||
alert_types: list[AlertType] | None = Field(
|
||||
None,
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -489,7 +489,7 @@ lite codex exec "summarize the repo"
|
|||
|
||||
Each command resolves your LiteLLM key (logging in via SSO when none is stored and you are at a terminal; otherwise it expects `LITELLM_PROXY_API_KEY` or `--api-key`), checks the key against the proxy so bad credentials fail immediately instead of deep inside the agent, exports the environment variables the agent reads, then replaces itself with the agent process.
|
||||
|
||||
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol).
|
||||
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins, and `ENABLE_TOOL_SEARCH=true` (unless you already set it) so Claude Code keeps tool search on even though the base URL is a proxy rather than a first-party Anthropic host. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol).
|
||||
|
||||
Options (these belong to the wrapper, so put them before the agent's own flags):
|
||||
|
||||
|
|
@ -505,7 +505,7 @@ The credential is short-lived by design (default 24h, configurable via `LITELLM_
|
|||
|
||||
### Route Every Claude Code Session Through the Proxy
|
||||
|
||||
`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL and `apiKeyHelper` to a `lite auth print-token` invocation, drops any stray static `ANTHROPIC_API_KEY` so the helper-issued token wins, and leaves every other setting in the file untouched. It backs up the original file before patching it.
|
||||
`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL, `env.ENABLE_TOOL_SEARCH` to `true` when that key is missing, and `apiKeyHelper` to a `lite auth print-token` invocation, drops any stray static `ANTHROPIC_API_KEY` so the helper-issued token wins, and leaves every other setting in the file untouched. It backs up the original file before patching it.
|
||||
|
||||
Two things need to already be true: you've run `lite login` (or `lite login --pkce`, whose key the helper renews on its own), since the apiKeyHelper depends on that stored token, and the proxy is already reachable, since `lite up` does not start one for you.
|
||||
|
||||
|
|
@ -529,7 +529,7 @@ Cursor is not supported: it has no equivalent file-based config to hot-patch thi
|
|||
lite --base-url https://your-proxy.example.com login --config-claude
|
||||
```
|
||||
|
||||
It writes the same two settings `lite up` does, `env.ANTHROPIC_BASE_URL` and `apiKeyHelper`, but persistently: there is no backup, nothing to restore, and no foreground process to keep alive. Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag.
|
||||
It writes the same settings `lite up` does, `env.ANTHROPIC_BASE_URL`, `env.ENABLE_TOOL_SEARCH`, and `apiKeyHelper`, but persistently: there is no backup, nothing to restore, and no foreground process to keep alive. Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag.
|
||||
|
||||
Because the credential is reached through `apiKeyHelper` rather than copied into the file, a later `lite login` refreshes it with no further action: Claude Code re-runs the helper on every request and picks up whatever token the most recent login stored. Nothing secret is written to `settings.json`.
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,8 @@ from .auth import context_secret_vault, get_stored_api_key, login
|
|||
ANTHROPIC_BASE_URL_ENV: Final = "ANTHROPIC_BASE_URL"
|
||||
ANTHROPIC_AUTH_TOKEN_ENV: Final = "ANTHROPIC_AUTH_TOKEN"
|
||||
ANTHROPIC_API_KEY_ENV: Final = "ANTHROPIC_API_KEY"
|
||||
ENABLE_TOOL_SEARCH_ENV: Final = "ENABLE_TOOL_SEARCH"
|
||||
ENABLE_TOOL_SEARCH_VALUE: Final = "true"
|
||||
OPENAI_BASE_URL_ENV: Final = "OPENAI_BASE_URL"
|
||||
OPENAI_API_KEY_ENV: Final = "OPENAI_API_KEY"
|
||||
|
||||
|
|
@ -61,7 +63,10 @@ def build_agent_env(
|
|||
Anthropic clients (Claude Code) append /v1/messages to ANTHROPIC_BASE_URL,
|
||||
so it stays the bare proxy root; OpenAI clients (Codex, OpenCode) expect the
|
||||
/v1 suffix on OPENAI_BASE_URL. ANTHROPIC_API_KEY is dropped so a stray
|
||||
Anthropic key cannot win over the bearer token we set.
|
||||
Anthropic key cannot win over the bearer token we set. ENABLE_TOOL_SEARCH
|
||||
defaults to true because Claude Code turns tool search off when
|
||||
ANTHROPIC_BASE_URL is not a first-party Anthropic host; a value already in
|
||||
the environment is left alone.
|
||||
"""
|
||||
env: Final = dict(base_env)
|
||||
root: Final = base_url.rstrip("/")
|
||||
|
|
@ -69,6 +74,8 @@ def build_agent_env(
|
|||
env[ANTHROPIC_BASE_URL_ENV] = root
|
||||
env[ANTHROPIC_AUTH_TOKEN_ENV] = api_key
|
||||
env.pop(ANTHROPIC_API_KEY_ENV, None)
|
||||
if ENABLE_TOOL_SEARCH_ENV not in env:
|
||||
env[ENABLE_TOOL_SEARCH_ENV] = ENABLE_TOOL_SEARCH_VALUE
|
||||
if PROFILE_OPENAI in profiles:
|
||||
env[OPENAI_BASE_URL_ENV] = root + "/v1"
|
||||
env[OPENAI_API_KEY_ENV] = api_key
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ API_KEY_HELPER_KEY: Final = "apiKeyHelper"
|
|||
ANTHROPIC_API_KEY_KEY: Final = "ANTHROPIC_API_KEY"
|
||||
ANTHROPIC_AUTH_TOKEN_KEY: Final = "ANTHROPIC_AUTH_TOKEN"
|
||||
ANTHROPIC_BASE_URL_KEY: Final = "ANTHROPIC_BASE_URL"
|
||||
ENABLE_TOOL_SEARCH_KEY: Final = "ENABLE_TOOL_SEARCH"
|
||||
ENABLE_TOOL_SEARCH_VALUE: Final = "true"
|
||||
# Force every one of Claude Code's own model tiers to request the auto-router by name.
|
||||
# Router's auto-router registry is keyed by the literal requested model string
|
||||
# (litellm/router.py:10711-10717) with no wildcard/pattern resolution, so a bare "*"
|
||||
|
|
@ -34,6 +36,7 @@ def merge_claude_settings_static_token(
|
|||
raw_env: Final = settings.get(ENV_KEY, {})
|
||||
base_env: Final = raw_env if isinstance(raw_env, dict) else {}
|
||||
env: Final[dict[str, JsonValue]] = {
|
||||
ENABLE_TOOL_SEARCH_KEY: ENABLE_TOOL_SEARCH_VALUE,
|
||||
**base_env,
|
||||
ANTHROPIC_BASE_URL_KEY: base_url.rstrip("/"),
|
||||
ANTHROPIC_AUTH_TOKEN_KEY: auth_token,
|
||||
|
|
|
|||
|
|
@ -21,6 +21,8 @@ ENV_KEY: Final = "env"
|
|||
API_KEY_HELPER_KEY: Final = "apiKeyHelper"
|
||||
ANTHROPIC_BASE_URL_KEY: Final = "ANTHROPIC_BASE_URL"
|
||||
ANTHROPIC_API_KEY_KEY: Final = "ANTHROPIC_API_KEY"
|
||||
ENABLE_TOOL_SEARCH_KEY: Final = "ENABLE_TOOL_SEARCH"
|
||||
ENABLE_TOOL_SEARCH_VALUE: Final = "true"
|
||||
|
||||
CLAUDE_SETTINGS_PATH: Final = Path.home() / ".claude" / "settings.json"
|
||||
BACKUP_PATH: Final = Path.home() / ".litellm" / "claude_settings_backup.json"
|
||||
|
|
@ -70,12 +72,15 @@ def merge_claude_settings(
|
|||
|
||||
Only env.ANTHROPIC_BASE_URL and the top-level apiKeyHelper are overridden; a
|
||||
stray env.ANTHROPIC_API_KEY is dropped so it cannot outrank the helper-issued
|
||||
token (same reasoning as build_agent_env in agents.py). Every other key is
|
||||
preserved untouched.
|
||||
token (same reasoning as build_agent_env in agents.py). ENABLE_TOOL_SEARCH
|
||||
defaults to true because Claude Code turns tool search off when
|
||||
ANTHROPIC_BASE_URL is not a first-party Anthropic host; an existing value is
|
||||
left alone. Every other key is preserved untouched.
|
||||
"""
|
||||
raw_env: Final = settings.get(ENV_KEY, {})
|
||||
base_env: Final = raw_env if isinstance(raw_env, dict) else {}
|
||||
env: Final = {
|
||||
ENABLE_TOOL_SEARCH_KEY: ENABLE_TOOL_SEARCH_VALUE,
|
||||
**{key: value for key, value in base_env.items() if key != ANTHROPIC_API_KEY_KEY},
|
||||
ANTHROPIC_BASE_URL_KEY: base_url.rstrip("/"),
|
||||
}
|
||||
|
|
@ -144,6 +149,8 @@ __all__ = (
|
|||
"AUTOROUTE_BACKUP_PATH",
|
||||
"BACKUP_PATH",
|
||||
"CLAUDE_SETTINGS_PATH",
|
||||
"ENABLE_TOOL_SEARCH_KEY",
|
||||
"ENABLE_TOOL_SEARCH_VALUE",
|
||||
"ENV_KEY",
|
||||
"SETTINGS_FILE_OWNERS",
|
||||
"ClaudeSettingsError",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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")),
|
||||
}
|
||||
]
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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 ##
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
```
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -1363,7 +1363,7 @@ _OPENAPI_HTTP_METHODS: Final = {
|
|||
# the UI. Kept here at module scope to match the analogous descriptor
|
||||
# `is_secret` flags in litellm.proxy.config_resolvers and the
|
||||
# `_CACHE_SENSITIVE_FIELDS` constant in the cache endpoint file.
|
||||
_ALERTING_SENSITIVE_VARS: Final[set[str]] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"}
|
||||
_ALERTING_SENSITIVE_VARS: Final[set[str]] = {"ALERTING_WEBHOOK_URL", "SLACK_WEBHOOK_URL", "SMTP_PASSWORD"}
|
||||
|
||||
|
||||
def _strip_operation_id_method_suffix(operation_id: str) -> str:
|
||||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
@ -16566,6 +16561,7 @@ async def create_config_audit_log(
|
|||
|
||||
_EXTRA_SECRET_CALLBACK_ENV_VARS: Final = frozenset(
|
||||
{
|
||||
"ALERTING_WEBHOOK_URL",
|
||||
"GALILEO_USERNAME",
|
||||
"GENERIC_LOGGER_HEADERS",
|
||||
"OTEL_HEADERS",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)],
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -645,7 +645,7 @@ class ProxyLogging:
|
|||
self.max_parallel_request_limiter = _PROXY_MaxParallelRequestsHandler(self.internal_usage_cache)
|
||||
self.max_budget_limiter = _PROXY_MaxBudgetLimiter()
|
||||
self.cache_control_check = _PROXY_CacheControlCheck()
|
||||
self.alerting: list | None = None
|
||||
self.alerting: list[str] | None = None
|
||||
self.alerting_threshold: float = 300 # default to 5 min. threshold
|
||||
self.alert_types: list[AlertType] = DEFAULT_ALERT_TYPES
|
||||
self.alert_to_webhook_url: dict | None = None
|
||||
|
|
@ -2364,7 +2364,9 @@ class ProxyLogging:
|
|||
# do nothing if alerting is not switched on (unless it's a soft_budget alert with team-specific emails)
|
||||
return
|
||||
|
||||
if self.alerting is not None and ("slack" in self.alerting or "ms_teams" in self.alerting):
|
||||
if self.alerting is not None and (
|
||||
"slack" in self.alerting or "ms_teams" in self.alerting or "webhook" in self.alerting
|
||||
):
|
||||
if self.slack_alerting_instance is not None:
|
||||
await self.slack_alerting_instance.budget_alerts(
|
||||
type=type,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -479,6 +479,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 +733,20 @@ 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.
|
||||
"""
|
||||
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",
|
||||
)
|
||||
and not decision.get("context_escalated")
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -759,6 +795,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 +1264,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 +1314,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 +1720,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 +1736,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 +1844,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 +1857,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 +1871,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 +1906,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 +1955,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 +1972,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 +2048,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()
|
||||
|
|
@ -2247,14 +2496,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 +2535,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 +2563,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 +2603,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,15 +2638,20 @@ 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(
|
||||
|
|
@ -2379,6 +2668,7 @@ class ComplexityRouter(CustomLogger):
|
|||
escalated=escalated,
|
||||
conversation_continuing=conversation_continuing,
|
||||
tier_litellm_params=session_tier_litellm_params,
|
||||
context_escalation_original_tier=pin_context_original_tier,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
|
@ -2573,6 +2863,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 +2911,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 +2936,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 +2995,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,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -823,6 +823,32 @@ 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."
|
||||
),
|
||||
)
|
||||
|
||||
# Semantic (embedding) matching for keyword_tier_rules instead of literal text matching
|
||||
semantic_keyword_matching: bool = Field(
|
||||
default=False,
|
||||
|
|
@ -839,6 +865,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,
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue