Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_azure_mai_image_generation

This commit is contained in:
mateo-berri 2026-06-08 20:58:43 +00:00
commit 2b6cb6d5d1
No known key found for this signature in database
1260 changed files with 120392 additions and 25039 deletions

View file

@ -111,6 +111,28 @@ commands:
- wait_for_service:
url: tcp://localhost:6379
timeout: "60"
start_openai_record_replay_proxy:
description: "Start the record/replay proxy (tests/_openai_record_replay_proxy.py) on host port 8090 and wait until healthy. Models whose api_base points here replay recorded provider responses, so the E2E run neither pays for nor depends on the live provider. The default upstream is OpenAI; a non-OpenAI model must point its api_base at /__recorder_upstream/<host>/ so the recorder forwards there instead of defaulting to OpenAI. Run after uv deps are synced."
steps:
- run:
name: Start record/replay proxy
background: true
command: |
CASSETTE_REDIS_URL="$CASSETTE_REDIS_URL" \
RECORDER_UPSTREAM_BASE_URL="https://api.openai.com" \
uv run --no-sync python tests/_openai_record_replay_proxy.py --host 0.0.0.0 --port 8090
- run:
name: Wait for record/replay proxy
command: |
for i in $(seq 1 30); do
if curl -sf http://localhost:8090/__recorder_health >/dev/null 2>&1; then
echo "record/replay proxy is up"
exit 0
fi
sleep 1
done
echo "record/replay proxy did not become ready" >&2
exit 1
setup_litellm_enterprise_pip:
steps:
- run:
@ -452,6 +474,120 @@ jobs:
- auth_ui_unit_tests_coverage.xml
- auth_ui_unit_tests_coverage
proxy_behavior_tests:
docker:
- *python312_image
- image: cimg/postgres:16.0@sha256:b125148bc76e8e8eee5eb3ad6020a3a14110a14e8192f1c645128afebe2e2f84
environment:
POSTGRES_USER: postgres
POSTGRES_PASSWORD: postgres
POSTGRES_DB: litellm_test
working_directory: ~/project
environment:
DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/litellm_test"
steps:
- checkout
- setup_google_dns
- install_uv
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
- run:
name: Seed DB schema via prisma db push
command: |
uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
- run:
name: Generate Prisma Client
command: uv run --no-sync python -m prisma generate
- run:
name: Run proxy management behavior tests
command: |
mkdir -p test-results
uv run --no-sync python -m pytest tests/proxy_behavior \
-v --junitxml=test-results/junit.xml --durations=10
no_output_timeout: 15m
- store_test_results:
path: test-results
proxy_security_tests:
docker:
- *python312_image
- image: cimg/postgres:16.0@sha256:b125148bc76e8e8eee5eb3ad6020a3a14110a14e8192f1c645128afebe2e2f84
environment:
POSTGRES_USER: postgres
POSTGRES_PASSWORD: postgres
POSTGRES_DB: litellm_test
working_directory: ~/project
environment:
DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/litellm_test"
steps:
- checkout
- setup_google_dns
- install_uv
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
- run:
name: Seed DB schema via prisma db push
command: |
uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
- run:
name: Generate Prisma Client
command: uv run --no-sync python -m prisma generate
- run:
name: Run proxy security tests
command: |
mkdir -p test-results
uv run --no-sync python -m pytest tests/proxy_security_tests \
-v --junitxml=test-results/junit.xml --durations=10
no_output_timeout: 15m
- store_test_results:
path: test-results
schema_migration_check:
docker:
- *python312_image
- image: cimg/postgres:16.0@sha256:b125148bc76e8e8eee5eb3ad6020a3a14110a14e8192f1c645128afebe2e2f84
environment:
POSTGRES_USER: postgres
POSTGRES_PASSWORD: postgres
POSTGRES_DB: litellm_test
working_directory: ~/project
environment:
# An empty database; the test applies every committed migration itself.
DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/litellm_test"
steps:
- checkout
- setup_google_dns
- install_uv
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
- run:
name: Generate Prisma Client
command: uv run --no-sync python -m prisma generate
- run:
name: Check schema.prisma is in sync with committed migrations
command: |
mkdir -p test-results
uv run --no-sync python -m pytest tests/proxy_migration_tests \
-v --junitxml=test-results/junit.xml --durations=10
no_output_timeout: 15m
- store_test_results:
path: test-results
litellm_router_testing: # Runs all tests with the "router" keyword
docker:
- *python312_image
@ -1511,6 +1647,7 @@ jobs:
command: |
zstd -d litellm-docker-database.tar.zst --stdout | docker load
docker tag litellm-docker-database:ci my-app:latest
- start_openai_record_replay_proxy
- run:
name: Run Docker container
command: |
@ -1541,6 +1678,7 @@ jobs:
-e LANGFUSE_PROJECT2_PUBLIC=$LANGFUSE_PROJECT2_PUBLIC \
-e LANGFUSE_PROJECT1_SECRET=$LANGFUSE_PROJECT1_SECRET \
-e LANGFUSE_PROJECT2_SECRET=$LANGFUSE_PROJECT2_SECRET \
-e RECORDER_OPENAI_BASE_URL=http://host.docker.internal:8090/v1 \
--add-host host.docker.internal:host-gateway \
--name my-app \
-v $(pwd)/proxy_server_config.yaml:/app/config.yaml \
@ -1678,6 +1816,7 @@ jobs:
command: |
zstd -d litellm-docker-database.tar.zst --stdout | docker load
docker images | grep litellm-docker-database
- start_openai_record_replay_proxy
- run:
name: Run Docker container
# intentionally give bad redis credentials here
@ -1701,6 +1840,7 @@ jobs:
-e DD_SITE=$DD_SITE \
-e AWS_REGION_NAME=$AWS_REGION_NAME \
-e COHERE_API_KEY=$COHERE_API_KEY \
-e RECORDER_COHERE_BASE_URL=http://host.docker.internal:8090/__recorder_upstream/api.cohere.com \
-e GCS_FLUSH_INTERVAL="1" \
--add-host host.docker.internal:host-gateway \
--name my-app \
@ -2266,6 +2406,7 @@ jobs:
command: |
zstd -d litellm-docker-database.tar.zst --stdout | docker load
docker images | grep litellm-docker-database
- start_openai_record_replay_proxy
- run:
name: Run Docker container with test config
command: |
@ -2274,6 +2415,7 @@ jobs:
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
-e LITELLM_MASTER_KEY="sk-1234" \
-e ANTHROPIC_API_KEY=$ANTHROPIC_API_KEY \
-e RECORDER_ANTHROPIC_BASE_URL=http://host.docker.internal:8090/__recorder_upstream/api.anthropic.com \
-e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \
-e AWS_SECRET_ACCESS_KEY=$AWS_SECRET_ACCESS_KEY \
-e AWS_REGION_NAME="us-east-1" \
@ -2643,6 +2785,12 @@ workflows:
filters: *main_branches
- auth_ui_unit_tests:
filters: *main_branches
- proxy_behavior_tests:
filters: *main_branches
- proxy_security_tests:
filters: *main_branches
- schema_migration_check:
filters: *main_branches
- build_docker_database_image:
filters: *main_branches
- e2e_ui_testing:

View file

@ -8,3 +8,6 @@
# Update pydantic code to fix warnings (GH-3600)
876840e9957bc7e9f7d6a2b58c4d7c53dad16481
# style(ui): run prettier --write across the dashboard (#29622)
7edf3a9cb55548b143df1692f4ed7c4681d7fcf7

3
.gitattributes vendored
View file

@ -1 +1,2 @@
*.ipynb linguist-vendored
*.ipynb linguist-vendored
ui/litellm-dashboard/src/lib/http/schema.d.ts linguist-generated

View file

@ -27,6 +27,11 @@ on:
required: false
type: number
default: 10
dist:
description: "pytest-xdist distribution mode (loadscope|load|worksteal|loadfile|no)"
required: false
type: string
default: "loadscope"
artifact-name:
description: "Unique name for the coverage artifact (must be unique per run)"
required: true
@ -82,18 +87,31 @@ jobs:
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
DIST: ${{ inputs.dist }}
run: |
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
-n "${WORKERS}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--dist=loadscope \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
if [ "${WORKERS}" = "0" ]; then
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
else
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
-n "${WORKERS}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--dist="${DIST}" \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
fi
- name: Save coverage report
if: always()

View file

@ -1,190 +0,0 @@
name: _Unit Test Services Base (Reusable)
on:
workflow_call:
inputs:
test-path:
description: "Pytest path(s) to run"
required: true
type: string
workers:
description: "Number of pytest-xdist workers (0 = no parallelism)"
required: false
type: number
default: 2
reruns:
description: "Number of reruns for flaky tests"
required: false
type: number
default: 2
timeout-minutes:
description: "Job timeout in minutes"
required: false
type: number
default: 20
max-failures:
description: "Stop after this many failures"
required: false
type: number
default: 10
enable-postgres:
description: "Start a local Postgres service container and run Prisma migrations"
required: false
type: boolean
default: false
dist:
description: "pytest-xdist distribution mode (loadscope|load|worksteal|loadfile|no)"
required: false
type: string
default: "loadscope"
artifact-name:
description: "Unique name for the coverage artifact (must be unique per run)"
required: false
type: string
default: "run"
permissions:
contents: read
# The postgres service container below is spawned per-job on localhost and
# destroyed with the job. Nothing outside the runner can reach it. The
# user/password/database here are not secrets — they're bootstrap values
# for a throwaway container — so we hardcode them instead of attaching
# every matrix shard to a GHA environment just to read three "secrets"
# (which also produces a "temporarily deployed to …" notification on the
# PR timeline per shard per push).
jobs:
run:
name: Run tests
runs-on: ubuntu-latest
timeout-minutes: ${{ inputs.timeout-minutes }}
services:
postgres:
image: postgres@sha256:705a5d5b5836f3fcba0d02c4d281e6a7dd9ed2dd4078640f08a1e1e9896e097d # postgres:14
env:
POSTGRES_USER: litellm
POSTGRES_PASSWORD: litellm
POSTGRES_DB: litellm_test
ports:
- 5432:5432
options: >-
--health-cmd "pg_isready"
--health-interval 10s
--health-timeout 5s
--health-retries 5
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
with:
version: "0.10.9"
- name: Cache uv dependencies
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-services-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-services-
- name: Install dependencies
run: |
uv sync --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Run Prisma migrations
if: ${{ inputs.enable-postgres }}
env:
DATABASE_URL: "postgresql://litellm:litellm@localhost:5432/litellm_test"
run: |
uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
- name: Run tests
env:
TEST_PATH: ${{ inputs.test-path }}
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
DIST: ${{ inputs.dist }}
DATABASE_URL: ${{ inputs.enable-postgres && 'postgresql://litellm:litellm@localhost:5432/litellm_test' || '' }}
run: |
if [ "${WORKERS}" = "0" ]; then
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
else
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
-n "${WORKERS}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--dist="${DIST}" \
--durations=20 \
--cov=./litellm \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
fi
- name: Save coverage report
if: always()
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
path: coverage.xml
retention-days: 1
upload-coverage:
name: Upload coverage to Codecov
needs: run
if: always()
runs-on: ubuntu-latest
permissions:
contents: read
id-token: write
pull-requests: write
steps:
- name: Checkout code
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Download coverage report
uses: actions/download-artifact@95815c38cf2ff2164869cbab79da8d1f422bc89e # v4.2.1
with:
pattern: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
path: coverage-reports
merge-multiple: true
- name: Upload to Codecov
uses: codecov/codecov-action@75cd11691c0faa626561e295848008c8a7dddffe # v5.5.4
with:
use_oidc: true
directory: coverage-reports
root_dir: ${{ github.workspace }}
flags: ${{ inputs.artifact-name }}
fail_ci_if_error: false

View file

@ -0,0 +1,84 @@
name: Check UI API Types Sync
on:
pull_request:
paths:
- "litellm/proxy/**"
- "litellm/types/**"
- "ui/litellm-dashboard/src/lib/http/schema.d.ts"
- "ui/litellm-dashboard/scripts/gen-api-types.mjs"
- "ui/litellm-dashboard/package.json"
- "ui/litellm-dashboard/package-lock.json"
- ".github/workflows/check-ui-api-types.yml"
permissions:
contents: read
jobs:
check-sync:
name: Verify schema.d.ts matches the proxy OpenAPI spec
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
with:
version: "0.10.9"
- name: Cache uv dependencies
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Install backend dependencies
run: uv sync --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Set up Node.js
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0
with:
node-version: "20"
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dashboard dependencies
working-directory: ui/litellm-dashboard
run: npm ci
- name: Regenerate types from the live spec
working-directory: ui/litellm-dashboard
env:
LITELLM_PYTHON: "uv run --no-sync python"
run: npm run gen:api
- name: Fail if types are stale
run: |
if ! git diff --exit-code -- ui/litellm-dashboard/src/lib/http/schema.d.ts; then
echo "::error file=ui/litellm-dashboard/src/lib/http/schema.d.ts::Generated API types are out of sync with the proxy OpenAPI spec."
echo ""
echo "A backend route or model changed without regenerating the dashboard types."
echo "To fix, run from ui/litellm-dashboard:"
echo " npm run gen:api"
echo "then commit the updated src/lib/http/schema.d.ts."
exit 1
fi
echo "schema.d.ts is in sync with the proxy OpenAPI spec."

View file

@ -28,6 +28,8 @@ jobs:
tests/test_litellm/completion_extras
tests/test_litellm/containers
tests/test_litellm/experimental_mcp_client
tests/test_litellm/models
tests/test_litellm/repositories
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/passthrough

View file

@ -1,9 +1,10 @@
name: "Unit Tests: Proxy DB Operations"
# Uses DATABASE_URL secret — only runs on trusted branches, not PRs.
on:
push:
branches: [main, "litellm_**"]
pull_request:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
@ -30,9 +31,6 @@ concurrency:
# xdist balances its 188 parametrized cases across workers instead of
# pinning the whole file to one worker (the default --dist=loadscope
# behavior for single-file targets).
# * test_db_schema_migration.py is isolated because one test in it
# (test_aaaasschema_migration_check) takes ~170s — by itself it
# determines the shard's wall-clock floor.
jobs:
# Fast guard — fails the workflow if a test_*.py file under
# tests/proxy_unit_tests/ is not referenced by any matrix entry below.
@ -166,18 +164,6 @@ jobs:
dist: loadscope
timeout: 15
# ---- db-and-spend: isolate the 170s schema-migration test ----
# test_db_schema_migration.py has exactly one test, and that test
# is mostly waiting on `prisma migrate deploy` / `prisma migrate
# diff` subprocesses (~170s). It does no CPU-bound Python work
# inside the test. Running with workers=0 (serial, no xdist)
# skips the 4-worker cold-start cost we'd otherwise pay for a
# single test, saving ~4 minutes of wall-clock.
- test-group: schema-migration
test-path: "tests/proxy_unit_tests/test_db_schema_migration.py"
workers: 0
dist: loadscope
timeout: 15
- test-group: db-and-spend
test-path: >-
tests/proxy_unit_tests/test_prisma_client_backoff_retry.py
@ -232,12 +218,11 @@ jobs:
workers: 4
dist: loadscope
timeout: 15
uses: ./.github/workflows/_test-unit-services-base.yml
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
workers: ${{ matrix.workers }}
reruns: 2
timeout-minutes: ${{ matrix.timeout }}
enable-postgres: true
dist: ${{ matrix.dist }}
artifact-name: proxy-db-${{ matrix.test-group }}

View file

@ -1,34 +0,0 @@
name: "Unit Tests: Proxy Management-Endpoint Behavior Pinning"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
proxy-mgmt-behavior:
uses: ./.github/workflows/_test-unit-services-base.yml
with:
test-path: tests/proxy_behavior
# workers=0 (no xdist): the world seed is a single shared Postgres
# state — two xdist workers both call seed_world() and race on the
# ``behavior-pin-budget`` row, producing UniqueViolation + cascading
# missing-membership FK failures. The whole suite is ~7s sequentially,
# so the cost of disabling parallelism here is negligible.
workers: 0
reruns: 0
enable-postgres: true
artifact-name: proxy-mgmt-behavior
timeout-minutes: 15

View file

@ -1,28 +0,0 @@
name: "Unit Tests: Security"
# Kept push-only (was previously required by DATABASE_URL secret scoping;
# now the postgres credentials are ephemeral localhost values but the
# push-trigger stays to match the proxy-db workflow cadence).
on:
push:
branches: [main, "litellm_**"]
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
jobs:
security:
uses: ./.github/workflows/_test-unit-services-base.yml
with:
test-path: "tests/proxy_security_tests/"
workers: 1
reruns: 2
timeout-minutes: 20
enable-postgres: true
artifact-name: security

View file

@ -240,6 +240,24 @@ graph LR
7. `DBSpendUpdateWriter.update_database()` queues spend increments to Redis
8. Background job `update_spend` flushes queued spend to PostgreSQL every 60s
### Data Access Layer (Models & Repositories)
Database entities and the operations on them live in two packages at the root of `litellm/` so both the gateway (`proxy/`) and the SDK can use them without importing proxy internals:
- `litellm/models/` holds the canonical Pydantic definitions for every persisted entity (`LiteLLM_VerificationToken`, `LiteLLM_TeamTable`, `LiteLLM_UserTable`, etc.). `proxy/_types.py` re-exports these for backwards compatibility, so existing imports keep working.
- `litellm/repositories/` holds the data-access layer. `BaseRepository[T]` provides the generic CRUD (`find_by_id`, `find_many`, `create`, `update`, `delete`, `count`, `exists`); entity repositories such as `VerificationTokenRepository`, `TeamRepository`, and `UserRepository` add domain-specific queries and writes on top of it.
Conventions to follow when touching this layer:
| Concern | How it's handled |
|---------|------------------|
| JSON columns | Prisma `Json` columns are stored as JSON strings. Repositories `json.dumps()` on write and `json.loads()` on read (see `_to_model` and the `_build_*_data` helpers). |
| Archive-then-delete | `delete_team` / `delete_token` copy the row into the `LiteLLM_Deleted*` table and delete the original inside a single `prisma_client.db.tx()` transaction. Archive payloads are built explicitly so only columns that exist on the archive table are written. |
| Column vs. field names | Where a model field differs from its DB column (for example `org_id` maps to the `organization_id` column), the repository translates in both directions rather than relying on Pydantic to guess. |
| Array mutations | Adds use Prisma's atomic `push` (`add_member`, `add_admin`, `add_models`) to avoid read-modify-write races. Removals fall back to read-modify-write because Prisma has no atomic array remove. |
To add a new entity, define the model under `litellm/models/`, re-export it from `proxy/_types.py` if existing code imports it from there, and add a repository under `litellm/repositories/` (subclass `BaseRepository` for plain CRUD, or add bespoke methods when the entity needs encryption, archiving, or atomic array updates). Mirror the tests in `tests/test_litellm/repositories/`.
---
## 2. SDK Request Flow

View file

@ -42,7 +42,7 @@ When you must use real LLM models to, for example, write e2e tests, write a QA r
If you're an internal contributor, when creating a new PR, the typical flow is to branch off litellm_internal_staging and create a branch prefixed with litellm_. Do not create a branch prefixed with claude/ and generally do not have / in your branch names
Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages. Never use a `claude/` prefix or put a `/` in a branch name. Do not add "Generated with Claude Code" (or any similar attribution) to PR descriptions. Do not create a new PR/branch off the existing PR to fix/add something that is related and could've just been committed directly to the existing PR's branch
Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages. Never use a `claude/` prefix or put a `/` in a branch name. Do not add "Generated with Claude Code" (or any similar attribution) to PR descriptions or comments. Do not create a new PR/branch off the existing PR to fix/add something that is related and could've just been committed directly to the existing PR's branch
When working on a PR, keep the PR description in sync with new commits being made

View file

@ -407,7 +407,7 @@ Support for more providers. Missing a provider or LLM Platform, raise a [feature
### Run in Developer Mode
#### Services
1. Setup .env file in root
2. Run dependant services `docker-compose up db prometheus`
2. Run dependent services `docker-compose up db prometheus`
#### Backend
1. (In root) create virtual environment `python -m venv .venv`

View file

@ -12,9 +12,14 @@ spec:
{{- include "litellm.backend.selectorLabels" . | nindent 6 }}
template:
metadata:
{{- with .Values.backend.podAnnotations }}
{{- if or .Values.gateway.config.create .Values.backend.podAnnotations }}
annotations:
{{- if .Values.gateway.config.create }}
checksum/config: {{ include (print $.Template.BasePath "/gateway/configmap.yaml") . | sha256sum }}
{{- end }}
{{- with .Values.backend.podAnnotations }}
{{- toYaml . | nindent 8 }}
{{- end }}
{{- end }}
labels:
{{- include "litellm.backend.selectorLabels" . | nindent 8 }}
@ -35,7 +40,17 @@ spec:
protocol: TCP
env:
{{- include "litellm.serverEnv" (dict "root" $ "component" .Values.backend) | nindent 12 }}
{{- if .Values.gateway.config.create }}
- name: CONFIG_FILE_PATH
value: /app/config/config.yaml
{{- end }}
{{- include "litellm.envFrom" .Values.backend | nindent 10 }}
{{- if .Values.gateway.config.create }}
volumeMounts:
- name: gateway-config
mountPath: /app/config/config.yaml
subPath: config.yaml
{{- end }}
{{- with .Values.backend.livenessProbe }}
livenessProbe:
{{- toYaml . | nindent 12 }}
@ -46,6 +61,12 @@ spec:
{{- end }}
resources:
{{- toYaml .Values.backend.resources | nindent 12 }}
{{- if .Values.gateway.config.create }}
volumes:
- name: gateway-config
configMap:
name: {{ include "litellm.gateway.fullname" . }}-config
{{- end }}
{{- with .Values.backend.nodeSelector }}
nodeSelector:
{{- toYaml . | nindent 8 }}

View file

@ -0,0 +1,23 @@
-- AlterTable: add admin-configured env_vars to MCP server table
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "env_vars" JSONB DEFAULT '[]';
-- CreateTable: per-user env var values for MCP servers
CREATE TABLE IF NOT EXISTS "LiteLLM_MCPUserEnvVars" (
"id" TEXT NOT NULL,
"user_id" TEXT NOT NULL,
"server_id" TEXT NOT NULL,
"values_b64" TEXT NOT NULL,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_MCPUserEnvVars_pkey" PRIMARY KEY ("id")
);
-- CreateIndex
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_MCPUserEnvVars_user_id_server_id_key" ON "LiteLLM_MCPUserEnvVars"("user_id", "server_id");
-- CreateIndex
CREATE INDEX IF NOT EXISTS "LiteLLM_MCPUserEnvVars_user_id_idx" ON "LiteLLM_MCPUserEnvVars"("user_id");
-- CreateIndex
CREATE INDEX IF NOT EXISTS "LiteLLM_MCPUserEnvVars_server_id_idx" ON "LiteLLM_MCPUserEnvVars"("server_id");

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "oauth2_flow" TEXT;

View file

@ -0,0 +1,3 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "timeout" DOUBLE PRECISION;

View file

@ -311,6 +311,11 @@ model LiteLLM_MCPServerTable {
tool_name_to_description Json? @default("{}")
extra_headers String[] @default([])
static_headers Json? @default("{}")
// Admin-configured environment variables interpolated into static_headers
// via ${NAME} syntax. Stored as an array of
// {name, value, scope, description}. scope is "global" (value used as-is)
// or "user" (value supplied per-user via LiteLLM_MCPUserEnvVars).
env_vars Json? @default("[]")
// Health check status
status String? @default("unknown")
last_health_check DateTime?
@ -322,6 +327,7 @@ model LiteLLM_MCPServerTable {
authorization_url String?
token_url String?
registration_url String?
oauth2_flow String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)
@ -330,6 +336,7 @@ model LiteLLM_MCPServerTable {
byok_description String[] @default([])
byok_api_key_help_url String?
source_url String?
timeout Float?
// BYOM submission lifecycle
approval_status String? @default("active")
submitted_by String?
@ -364,6 +371,21 @@ model LiteLLM_MCPUserCredentials {
@@unique([user_id, server_id])
}
// Per-user environment variable values for MCP servers.
// values_b64 is an encrypted JSON object: {VAR_NAME: "value", ...}.
model LiteLLM_MCPUserEnvVars {
id String @id @default(uuid())
user_id String
server_id String
values_b64 String
created_at DateTime @default(now())
updated_at DateTime @default(now()) @updatedAt
@@unique([user_id, server_id])
@@index([user_id])
@@index([server_id])
}
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.73"
version = "0.4.74"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.73"
version = "0.4.74"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -16,8 +16,17 @@ import os
# Load .env before any other litellm imports so env vars (e.g. LITELLM_UI_SESSION_DURATION) are available
import dotenv as _dotenv
def _dev_env_hot_reload_enabled() -> bool:
"""The proxy exports this flag when started with ``--reload``. A reloaded
worker is a fresh process that inherits the reloader's environment, so an
edited ``.env`` value stays masked by the stale inherited one unless we
let the file win; overriding makes the edit take effect on reload."""
return os.getenv("LITELLM_DEV_ENV_HOT_RELOAD") == "True"
if os.getenv("LITELLM_MODE", "DEV") == "DEV":
_dotenv.load_dotenv()
_dotenv.load_dotenv(override=_dev_env_hot_reload_enabled())
from typing import (
Callable,
@ -433,6 +442,13 @@ custom_prometheus_metadata_labels: List[str] = []
custom_prometheus_tags: List[str] = []
prometheus_metrics_config: Optional[List] = None
prometheus_emit_stream_label: bool = False
# Opt-in: emit `rate_limit_category` and `rate_limit_type` labels on
# `litellm_proxy_failed_requests_metric`. Off by default to preserve the
# pre-unification label set so existing dashboards / recording rules keyed on
# that metric keep matching after upgrade. Enable when downstream consumers
# are ready to split 429s by source (vendor vs. litellm) and dimension
# (RPM/TPM/concurrent/budget).
prometheus_emit_rate_limit_labels: bool = False
prometheus_user_budget_label_include_email_alias: bool = False
prometheus_end_user_metrics_max_series_per_metric: Optional[int] = 10000
prometheus_end_user_metrics_ttl_seconds: Optional[float] = 3600.0
@ -612,6 +628,7 @@ cerebras_models: Set = set()
galadriel_models: Set = set()
nvidia_nim_models: Set = set()
nvidia_riva_models: Set = set()
soniox_models: Set = set()
sambanova_models: Set = set()
sambanova_embedding_models: Set = set()
novita_models: Set = set()
@ -844,6 +861,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
nvidia_nim_models.add(key)
elif value.get("litellm_provider") == "nvidia_riva":
nvidia_riva_models.add(key)
elif value.get("litellm_provider") == "soniox":
soniox_models.add(key)
elif value.get("litellm_provider") == "sambanova":
sambanova_models.add(key)
elif value.get("litellm_provider") == "sambanova-embedding-models":
@ -1009,6 +1028,7 @@ model_list = list(
| galadriel_models
| nvidia_nim_models
| nvidia_riva_models
| soniox_models
| sambanova_models
| azure_text_models
| novita_models
@ -1109,6 +1129,7 @@ models_by_provider: dict = {
"galadriel": galadriel_models,
"nvidia_nim": nvidia_nim_models,
"nvidia_riva": nvidia_riva_models,
"soniox": soniox_models,
"sambanova": sambanova_models | sambanova_embedding_models,
"novita": novita_models,
"nebius": nebius_models | nebius_embedding_models,
@ -1289,6 +1310,8 @@ from .exceptions import (
NotFoundError,
PermissionDeniedError,
RateLimitError,
RateLimitErrorCategory,
RateLimitType,
ServiceUnavailableError,
BadGatewayError,
OpenAIError,
@ -1740,6 +1763,9 @@ if TYPE_CHECKING:
from .llms.openrouter.responses.transformation import (
OpenRouterResponsesAPIConfig as OpenRouterResponsesAPIConfig,
)
from .llms.bedrock_mantle.responses.transformation import (
BedrockMantleResponsesAPIConfig as BedrockMantleResponsesAPIConfig,
)
from .llms.gemini.interactions.transformation import (
GoogleAIStudioInteractionsConfig as GoogleAIStudioInteractionsConfig,
)

View file

@ -237,6 +237,7 @@ LLM_CONFIG_NAMES = (
"PerplexityResponsesConfig",
"DatabricksResponsesAPIConfig",
"OpenRouterResponsesAPIConfig",
"BedrockMantleResponsesAPIConfig",
"GoogleAIStudioInteractionsConfig",
"OpenAIOSeriesConfig",
"AnthropicSkillsConfig",
@ -320,6 +321,7 @@ LLM_CONFIG_NAMES = (
"LemonadeChatConfig",
"SnowflakeEmbeddingConfig",
"AmazonNovaChatConfig",
"SonioxAudioTranscriptionConfig",
)
# Types that support lazy loading via _lazy_import_types
@ -958,6 +960,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
".llms.openrouter.responses.transformation",
"OpenRouterResponsesAPIConfig",
),
"BedrockMantleResponsesAPIConfig": (
".llms.bedrock_mantle.responses.transformation",
"BedrockMantleResponsesAPIConfig",
),
"GoogleAIStudioInteractionsConfig": (
".llms.gemini.interactions.transformation",
"GoogleAIStudioInteractionsConfig",
@ -1190,6 +1196,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
".llms.amazon_nova.chat.transformation",
"AmazonNovaChatConfig",
),
"SonioxAudioTranscriptionConfig": (
".llms.soniox.audio_transcription.transformation",
"SonioxAudioTranscriptionConfig",
),
}
# Import map for utils module lazy imports

View file

@ -309,9 +309,13 @@ class Cache:
param_value = kwargs[param]
cache_key += f"{str(param)}: {str(param_value)}"
verbose_logger.debug("\nCreated cache key: %s", cache_key)
hashed_cache_key = Cache._get_hashed_cache_key(cache_key)
hashed_cache_key = self._add_namespace_to_cache_key(hashed_cache_key, **kwargs)
verbose_logger.debug(
"\nCreated cache key: %s (source material length: %d)",
hashed_cache_key,
len(cache_key),
)
# Remove preset_cache_key from kwargs to avoid "got multiple values" TypeError
# when kwargs already contains preset_cache_key from upstream callers
kwargs_for_preset = {k: v for k, v in kwargs.items() if k != "preset_cache_key"}
@ -497,6 +501,34 @@ class Cache:
return cached_response
return cached_result
@staticmethod
def _get_safe_cache_lookup_kwargs(kwargs: Dict[str, Any]) -> Dict[str, Any]:
cache_lookup_kwargs: Dict[str, Any] = {}
for prompt_kwarg in ("messages", "input"):
if prompt_kwarg in kwargs:
cache_lookup_kwargs[prompt_kwarg] = kwargs[prompt_kwarg]
if isinstance(kwargs.get("metadata"), dict):
cache_lookup_kwargs["metadata"] = {}
return cache_lookup_kwargs
@staticmethod
def _update_metadata_from_cache_lookup_kwargs(
original_kwargs: Dict[str, Any], cache_lookup_kwargs: Dict[str, Any]
) -> None:
original_metadata = original_kwargs.get("metadata")
cache_lookup_metadata = cache_lookup_kwargs.get("metadata")
if not isinstance(original_metadata, dict) or not isinstance(
cache_lookup_metadata, dict
):
return
if "semantic-similarity" in cache_lookup_metadata:
original_metadata["semantic-similarity"] = cache_lookup_metadata[
"semantic-similarity"
]
def get_cache(self, dynamic_cache_object: Optional[BaseCache] = None, **kwargs):
"""
Retrieves the cached result for the given arguments.
@ -511,7 +543,6 @@ class Cache:
try: # never block execution
if self.should_use_cache(**kwargs) is not True:
return
messages = kwargs.get("messages", [])
if "cache_key" in kwargs:
cache_key = kwargs["cache_key"]
else:
@ -523,12 +554,19 @@ class Cache:
or cache_control_args.get("s-max-age")
or float("inf")
)
cache_lookup_kwargs = self._get_safe_cache_lookup_kwargs(kwargs)
if dynamic_cache_object is not None:
cached_result = dynamic_cache_object.get_cache(
cache_key, messages=messages
cache_key, **cache_lookup_kwargs
)
else:
cached_result = self.cache.get_cache(cache_key, messages=messages)
cached_result = self.cache.get_cache(
cache_key, **cache_lookup_kwargs
)
self._update_metadata_from_cache_lookup_kwargs(
original_kwargs=kwargs,
cache_lookup_kwargs=cache_lookup_kwargs,
)
return self._get_cache_logic(
cached_result=cached_result, max_age=max_age
)
@ -549,7 +587,6 @@ class Cache:
if self.should_use_cache(**kwargs) is not True:
return
kwargs.get("messages", [])
if "cache_key" in kwargs:
cache_key = kwargs["cache_key"]
else:

View file

@ -213,6 +213,78 @@ class RedisSemanticCache(BaseCache):
ttl = int(ttl)
return ttl
@classmethod
def _get_prompt_from_kwargs(cls, **kwargs) -> Optional[str]:
"""
Extract a semantic-cache prompt from chat or Responses API request kwargs.
"""
messages = kwargs.get("messages")
if messages:
return get_str_from_messages(messages)
if "input" not in kwargs:
return None
prompt_parts: List[str] = []
cls._collect_responses_input_text(kwargs.get("input"), prompt_parts)
prompt = "\n".join(prompt_parts).strip()
return prompt or None
@classmethod
def _collect_responses_input_text(cls, value: Any, prompt_parts: List[str]) -> None:
value = cls._coerce_response_input_value(value)
if value is None:
return
if isinstance(value, str):
stripped_value = value.strip()
if stripped_value:
prompt_parts.append(stripped_value)
return
if isinstance(value, (list, tuple)):
for item in value:
cls._collect_responses_input_text(item, prompt_parts)
return
if isinstance(value, dict):
content = value.get("content")
if content is not None:
cls._collect_responses_input_text(content, prompt_parts)
return
for text_key in ("text", "output", "input_text", "output_text"):
text_value = value.get(text_key)
if isinstance(text_value, str):
stripped_text = text_value.strip()
if stripped_text:
prompt_parts.append(stripped_text)
return
return
content = getattr(value, "content", None)
if content is not None:
cls._collect_responses_input_text(content, prompt_parts)
return
for text_key in ("text", "output", "input_text", "output_text"):
text_value = getattr(value, text_key, None)
if isinstance(text_value, str):
stripped_text = text_value.strip()
if stripped_text:
prompt_parts.append(stripped_text)
return
@staticmethod
def _coerce_response_input_value(value: Any) -> Any:
model_dump = getattr(value, "model_dump", None)
if callable(model_dump):
return model_dump()
dict_method = getattr(value, "dict", None)
if callable(dict_method):
return dict_method()
return value
def _get_embedding(self, prompt: str) -> List[float]:
"""
Generate an embedding vector for the given prompt using the configured embedding model.
@ -278,13 +350,11 @@ class RedisSemanticCache(BaseCache):
value_str: Optional[str] = None
try:
# Extract the prompt from messages
messages = kwargs.get("messages", [])
if not messages:
print_verbose("No messages provided for semantic caching")
prompt = self._get_prompt_from_kwargs(**kwargs)
if prompt is None:
print_verbose("No prompt provided for semantic caching")
return
prompt = get_str_from_messages(messages)
value_str = str(value)
store_kwargs: Dict[str, Any] = {
@ -315,14 +385,12 @@ class RedisSemanticCache(BaseCache):
print_verbose(f"Redis semantic-cache get_cache, kwargs: {kwargs}")
try:
# Extract the prompt from messages
messages = kwargs.get("messages", [])
if not messages:
print_verbose("No messages provided for semantic cache lookup")
prompt = self._get_prompt_from_kwargs(**kwargs)
if prompt is None:
print_verbose("No prompt provided for semantic cache lookup")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
prompt = get_str_from_messages(messages)
# Check the cache for semantically similar prompts in this exact
# LiteLLM cache-key scope.
check_kwargs: Dict[str, Any] = {
@ -428,13 +496,11 @@ class RedisSemanticCache(BaseCache):
print_verbose(f"Async Redis semantic-cache set_cache, kwargs: {kwargs}")
try:
# Extract the prompt from messages
messages = kwargs.get("messages", [])
if not messages:
print_verbose("No messages provided for semantic caching")
prompt = self._get_prompt_from_kwargs(**kwargs)
if prompt is None:
print_verbose("No prompt provided for semantic caching")
return
prompt = get_str_from_messages(messages)
value_str = str(value)
# Generate embedding for the value (response) to cache
@ -471,15 +537,12 @@ class RedisSemanticCache(BaseCache):
print_verbose(f"Async Redis semantic-cache get_cache, kwargs: {kwargs}")
try:
# Extract the prompt from messages
messages = kwargs.get("messages", [])
if not messages:
print_verbose("No messages provided for semantic cache lookup")
prompt = self._get_prompt_from_kwargs(**kwargs)
if prompt is None:
print_verbose("No prompt provided for semantic cache lookup")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
prompt = get_str_from_messages(messages)
# Generate embedding for the prompt
prompt_embedding = await self._get_async_embedding(prompt, **kwargs)

View file

@ -402,6 +402,20 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
instructions,
) = self.convert_chat_completion_messages_to_responses_api(messages)
# OpenAI's Responses API rejects an empty input. For a system-only
# request, carry the system message as a system-role input item instead
# of instructions, mirroring how non-string system content is already
# handled in convert_chat_completion_messages_to_responses_api.
if not input_items and instructions is not None:
input_items = [
{
"type": "message",
"role": "system",
"content": [{"type": "input_text", "text": instructions}],
}
]
instructions = None
optional_params = self._extract_extra_body_params(optional_params)
# Build responses API request using the reverse transformation logic

View file

@ -831,6 +831,7 @@ openai_compatible_providers: List = [
"nano-gpt", # Nano-GPT - JSON-configured provider
"poe", # Poe - JSON-configured provider
"chutes", # Chutes - JSON-configured provider
"parasail", # Parasail - JSON-configured provider
"featherless_ai",
"nscale",
"nebius",

View file

@ -2425,12 +2425,11 @@ class BaseTokenUsageProcessor:
if not attr.startswith("_") and not callable(
getattr(usage.completion_tokens_details, attr)
):
current_val = getattr(
combined.completion_tokens_details, attr, 0
current_val = (
getattr(combined.completion_tokens_details, attr, 0) or 0
)
new_val = getattr(usage.completion_tokens_details, attr, 0)
if new_val is not None and current_val is not None:
new_val = getattr(usage.completion_tokens_details, attr, 0) or 0
if isinstance(new_val, (int, float)):
setattr(
combined.completion_tokens_details,
attr,

View file

@ -9,13 +9,109 @@
## LiteLLM versions of the OpenAI Exception Types
from typing import Any, Dict, Optional
import enum
from typing import Any, Dict, Optional, Union
import httpx
import openai
from litellm.types.utils import LiteLLMCommonStrings
class RateLimitErrorCategory(str, enum.Enum):
"""
Category of a rate limit error, allowing callers to distinguish where the rate
limit originated. Exposed on every :class:`RateLimitError` instance via the
``category`` attribute.
Use these values to switch on the rate limit source, e.g.::
try:
...
except litellm.RateLimitError as e:
if e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT:
... # litellm's own limiter (key/team/user/model RPM/TPM/budget)
elif e.category == RateLimitErrorCategory.VENDOR_RATE_LIMIT:
... # the upstream LLM provider returned 429
"""
VENDOR_RATE_LIMIT = "vendor_rate_limit"
"""The upstream LLM provider returned a rate-limit response (e.g. OpenAI 429)."""
VENDOR_BATCH_RATE_LIMIT = "vendor_batch_rate_limit"
"""The upstream LLM provider returned a rate-limit response on a batch endpoint."""
LITELLM_RATE_LIMIT = "litellm_rate_limit"
"""LiteLLM's own rate limiter (key/team/user/model RPM/TPM, budget, parallel-requests, etc.) blocked the request."""
LITELLM_BATCH_RATE_LIMIT = "litellm_batch_rate_limit"
"""LiteLLM's own batch rate limiter (token/request budget across a batch input file) blocked the request."""
class RateLimitType(str, enum.Enum):
"""
The dimension that was exceeded when a rate-limit error fired.
This is orthogonal to :class:`RateLimitErrorCategory` — *category* tells
callers **who** rate-limited the request (the upstream vendor vs. one of
litellm's own limiters), while *type* tells them **which limit dimension**
was exceeded (an RPM ceiling, a TPM ceiling, a max-parallel-requests
ceiling, a budget cap, or a max-iterations cap).
Surfaced both on every :class:`RateLimitError` instance via the
``rate_limit_type`` attribute and on the structured
``StandardLoggingPayload.error_information.error_rate_limit_type`` field
so custom callbacks / metrics consumers can split rate-limit failures by
cause without parsing free-text error messages.
"""
REQUESTS = "requests"
"""Requests-per-minute (RPM) or requests-per-window ceiling exceeded."""
TOKENS = "tokens"
"""Tokens-per-minute (TPM) or tokens-per-window ceiling exceeded."""
CONCURRENT_REQUESTS = "concurrent_requests"
"""``max_parallel_requests`` — too many in-flight requests at once."""
BUDGET = "budget"
"""Spend budget cap reached (key, team, user, or per-session)."""
MAX_ITERATIONS = "max_iterations"
"""Per-session max-iterations cap reached (agent-style flows)."""
_RATE_LIMIT_CATEGORY_VALUES = frozenset(c.value for c in RateLimitErrorCategory)
_RATE_LIMIT_TYPE_VALUES = frozenset(t.value for t in RateLimitType)
def validate_rate_limit_category(value: Any) -> Optional[str]:
"""Return ``value`` only if it matches a known :class:`RateLimitErrorCategory`.
Used at duck-typed read sites (StandardLoggingPayload extraction, Prometheus
labels) to reject `.category` strings set by unrelated third-party exceptions
— otherwise those would leak into custom-callback payloads and Prometheus
label cardinality.
"""
if isinstance(value, RateLimitErrorCategory):
return value.value
if isinstance(value, str) and value in _RATE_LIMIT_CATEGORY_VALUES:
return value
return None
def validate_rate_limit_type(value: Any) -> Optional[str]:
"""Return ``value`` only if it matches a known :class:`RateLimitType`.
See :func:`validate_rate_limit_category` for the rationale.
"""
if isinstance(value, RateLimitType):
return value.value
if isinstance(value, str) and value in _RATE_LIMIT_TYPE_VALUES:
return value
return None
_MINIMAL_ERROR_RESPONSE: Optional[httpx.Response] = None
@ -321,6 +417,18 @@ class PermissionDeniedError(openai.PermissionDeniedError): # type: ignore
class RateLimitError(openai.RateLimitError): # type: ignore
"""
Unified rate-limit error.
Every rate-limit condition surfaced by litellm — whether it originated from
an upstream LLM provider, a vendor batch endpoint, or one of litellm's own
proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,
max-iterations, etc.) — is raised as an instance of this class.
The :attr:`category` attribute lets callers distinguish the source. See
:class:`RateLimitErrorCategory` for the available values.
"""
def __init__(
self,
message,
@ -330,6 +438,12 @@ class RateLimitError(openai.RateLimitError): # type: ignore
litellm_debug_info: Optional[str] = None,
max_retries: Optional[int] = None,
num_retries: Optional[int] = None,
category: Union[str, RateLimitErrorCategory] = (
RateLimitErrorCategory.VENDOR_RATE_LIMIT
),
rate_limit_type: Optional[Union[str, RateLimitType]] = None,
headers: Optional[Dict[str, str]] = None,
detail: Any = None,
):
self.status_code = 429
self.message = "litellm.RateLimitError: {}".format(message)
@ -338,9 +452,39 @@ class RateLimitError(openai.RateLimitError): # type: ignore
self.litellm_debug_info = litellm_debug_info
self.max_retries = max_retries
self.num_retries = num_retries
self.category = (
category.value if isinstance(category, RateLimitErrorCategory) else category
)
# Which dimension was exceeded — request count, token count, parallel
# requests, budget, max iterations. None when the source didn't
# classify the failure (e.g. legacy vendor 429 with no header hints).
self.rate_limit_type: Optional[str] = (
rate_limit_type.value
if isinstance(rate_limit_type, RateLimitType)
else rate_limit_type
)
# Headers explicitly attached to the error (e.g. retry-after,
# rate_limit_type, reset_at). Preserved across the proxy boundary so
# clients can react appropriately.
#
# IMPORTANT: we deliberately do NOT auto-populate self.headers from
# response.headers when only `response` is provided. A vendor 429 can
# set arbitrary response headers (Set-Cookie, CORS overrides, …); if
# those leaked into e.headers and a downstream proxy serializer
# forwarded them to the client, a malicious upstream could inject
# browser-interpreted headers for the proxy origin. Vendor response
# headers stay reachable on `e.response.headers` for callers that
# explicitly want them; only the proxy-supplied `headers=` kwarg
# makes it onto `self.headers`.
_response_headers = (
getattr(response, "headers", None) if response is not None else None
)
self.headers: Optional[Dict[str, str]] = (
{k: str(v) for k, v in headers.items()} if headers else None
)
# Mirrors FastAPI HTTPException.detail so the same instance can be
# serialized through both the ProxyException and HTTPException paths.
self.detail = detail if detail is not None else self.message
self.response = httpx.Response(
status_code=429,
headers=_response_headers,
@ -843,11 +987,24 @@ LITELLM_EXCEPTION_TYPES = [
class BudgetExceededError(Exception):
def __init__(
self, current_cost: float, max_budget: float, message: Optional[str] = None
self,
current_cost: float,
max_budget: float,
message: Optional[str] = None,
llm_provider: Optional[str] = None,
):
self.current_cost = current_cost
self.max_budget = max_budget
self.status_code = 429
self.llm_provider = llm_provider or ""
# Surface unified rate-limit fields without joining the RateLimitError
# hierarchy so existing `except BudgetExceededError:` handlers keep
# working; custom callbacks reading StandardLoggingPayload pick these
# up via the same `category` / `rate_limit_type` attributes the rest
# of the unified rate-limit error path uses. Stored as plain strings
# to match the normalization RateLimitError.__init__ performs.
self.category: str = RateLimitErrorCategory.LITELLM_RATE_LIMIT.value
self.rate_limit_type: str = RateLimitType.BUDGET.value
message = (
message
or f"Budget has been exceeded! Current cost: {current_cost}, Max budget: {max_budget}"
@ -1062,3 +1219,37 @@ class GuardrailInterventionNormalStringError(
def __repr__(self):
return self.__str__()
class SensitiveDataRouteException(Exception):
"""
Exception raised when a guardrail detects sensitive data and wants to reroute the request.
Instead of blocking the request, this exception signals that the request should be
routed to a different model (typically an on-premise model for data privacy).
The proxy catches this exception and:
1. Reroutes the current request to the specified model
2. When sticky_session_routing is True, stores the routing decision in session
cache so all subsequent requests in the same session are routed to the same model
"""
def __init__(
self,
route_to_model: str,
session_id: str,
guardrail_name: Optional[str] = None,
detection_info: Optional[Dict[str, Any]] = None,
message: Optional[str] = None,
sticky_session_routing: bool = True,
):
self.route_to_model = route_to_model
self.session_id = session_id
self.guardrail_name = guardrail_name
self.detection_info = detection_info or {}
self.sticky_session_routing = sticky_session_routing
self.message = (
message
or f"Sensitive data detected by {guardrail_name}. Routing to model: {route_to_model}"
)
super().__init__(self.message)

View file

@ -4,6 +4,7 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers.
import asyncio
import base64
import os
from typing import (
Any,
Awaitable,
@ -16,7 +17,6 @@ from typing import (
TypeVar,
Union,
)
import httpx
from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParameters
from mcp.client.sse import sse_client
@ -42,9 +42,8 @@ from mcp.types import (
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
from litellm._logging import verbose_logger
from litellm.constants import MCP_CLIENT_TIMEOUT
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
from litellm.types.llms.custom_http import VerifyTypes
from litellm.types.mcp import (
@ -61,13 +60,33 @@ def to_basic_auth(auth_value: str) -> str:
return base64.b64encode(auth_value.encode("utf-8")).decode()
def _strip_header_whitespace(headers: Dict[str, str]) -> Dict[str, str]:
return {
(key.strip() if isinstance(key, str) else key): (
value.strip() if isinstance(value, str) else value
)
for key, value in headers.items()
}
def _first_non_cancelled_cause(exc: BaseException) -> Optional[BaseException]:
queue: List[BaseException] = [exc]
while queue:
current = queue.pop(0)
nested = getattr(current, "exceptions", None)
if nested:
queue.extend(nested)
elif not isinstance(current, asyncio.CancelledError):
return current
return None
TSessionResult = TypeVar("TSessionResult")
class MCPSigV4Auth(httpx.Auth):
"""
httpx Auth class that signs each request with AWS SigV4.
This is used for MCP servers that require AWS SigV4 authentication,
such as AWS Bedrock AgentCore MCP servers. httpx calls auth_flow()
for every outgoing request, enabling per-request signature computation.
@ -92,10 +111,8 @@ class MCPSigV4Auth(httpx.Auth):
"Missing botocore to use AWS SigV4 authentication. "
"Run 'pip install boto3'."
)
self.service_name = aws_service_name or "bedrock-agentcore"
self.region_name = aws_region_name or "us-east-1"
# Note: os.environ/ prefixed values are already resolved by
# ProxyConfig._check_for_os_environ_vars() at config load time.
# Values arrive here as plain strings.
@ -143,20 +160,17 @@ class MCPSigV4Auth(httpx.Auth):
session_name = (
aws_session_name or f"litellm-mcp-{int(__import__('time').time())}"
)
sts_kwargs: dict = {"region_name": aws_region_name}
if aws_access_key_id and aws_secret_access_key:
sts_kwargs["aws_access_key_id"] = aws_access_key_id
sts_kwargs["aws_secret_access_key"] = aws_secret_access_key
if aws_session_token:
sts_kwargs["aws_session_token"] = aws_session_token
sts_client = boto3.client("sts", **sts_kwargs)
sts_response = sts_client.assume_role(
RoleArn=aws_role_name,
RoleSessionName=session_name,
)
sts_creds = sts_response["Credentials"]
return Credentials(
access_key=sts_creds["AccessKeyId"],
@ -178,17 +192,14 @@ class MCPSigV4Auth(httpx.Auth):
data=request.content,
headers=dict(request.headers),
)
# Sign the request — SigV4Auth.add_auth() adds Authorization,
# X-Amz-Date, and X-Amz-Security-Token (if session token present).
# Host header is derived automatically from the URL.
sigv4 = SigV4Auth(self.credentials, self.service_name, self.region_name)
sigv4.add_auth(aws_request)
# Copy SigV4 headers back to the httpx request
for header_name, header_value in aws_request.headers.items():
request.headers[header_name] = header_value
yield request
@ -198,6 +209,8 @@ class MCPClient:
SSE and HTTP transports
Authentication via Bearer token, Basic Auth, or API Key
Tool calling with error handling and result parsing
Sampling callbacks for upstream server LLM requests
Elicitation callbacks for upstream server user-input requests
"""
def __init__(
@ -211,6 +224,9 @@ class MCPClient:
extra_headers: Optional[Dict[str, str]] = None,
ssl_verify: Optional[VerifyTypes] = None,
aws_auth: Optional[httpx.Auth] = None,
sampling_callback: Optional[Callable] = None,
elicitation_callback: Optional[Callable] = None,
logging_callback: Optional[Callable] = None,
):
self.server_url: str = server_url
self.transport_type: MCPTransport = transport_type
@ -222,6 +238,9 @@ class MCPClient:
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
self._aws_auth: Optional[httpx.Auth] = aws_auth
self._last_initialize_instructions: Optional[str] = None
self._sampling_callback: Optional[Callable] = sampling_callback
self._elicitation_callback: Optional[Callable] = elicitation_callback
self._logging_callback: Optional[Callable] = logging_callback
# handle the basic auth value if provided
if auth_value:
self.update_auth_value(auth_value)
@ -231,23 +250,20 @@ class MCPClient:
) -> Tuple[Any, Optional[httpx.AsyncClient]]:
"""
Create the appropriate transport context based on transport type.
Returns:
Tuple of (transport_context, http_client).
http_client is only set for HTTP transport and needs cleanup.
"""
http_client: Optional[httpx.AsyncClient] = None
if self.transport_type == MCPTransport.stdio:
if not self.stdio_config:
raise ValueError("stdio_config is required for stdio transport")
server_params = StdioServerParameters(
command=self.stdio_config.get("command", ""),
args=self.stdio_config.get("args", []),
env=self.stdio_config.get("env", {}),
env=self._get_safe_stdio_env(self.stdio_config.get("env")),
)
return stdio_client(server_params), None
if self.transport_type == MCPTransport.sse:
headers = self._get_auth_headers()
httpx_client_factory = self._create_httpx_client_factory()
@ -260,14 +276,12 @@ class MCPClient:
),
None,
)
# HTTP transport (default)
if streamable_http_client is None:
raise ImportError(
"streamable_http_client is not available. "
"Please install mcp with HTTP support."
)
headers = self._get_auth_headers()
httpx_client_factory = self._create_httpx_client_factory()
verbose_logger.debug("litellm headers for streamable_http_client: %s", headers)
@ -281,6 +295,54 @@ class MCPClient:
)
return transport_ctx, http_client
def _get_safe_stdio_env(
self, provided_env: Optional[Dict[str, str]]
) -> Optional[Dict[str, str]]:
"""
Return a safe environment for the stdio subprocess.
If provided_env is set, we use it as-is.
If provided_env is None, we return a minimal allowlist from the parent environment
to avoid leaking sensitive LiteLLM keys (OPENAI_API_KEY, etc.) to sub-processes.
"""
if provided_env is not None:
return provided_env
# Minimal allowlist of safe/standard environment variables
safe_keys = {
"PATH",
"HOME",
"USER",
"LOGNAME",
"TMPDIR",
"TMP",
"TEMP",
"SHELL",
"LANG",
"LC_ALL",
# Node/Package manager caches
"NPM_CONFIG_CACHE",
"PNPM_HOME",
"XDG_CACHE_HOME",
"XDG_CONFIG_HOME",
"XDG_DATA_HOME",
# System info
"SYSTEMROOT",
"COMSPEC",
"PATHEXT",
"WINDIR",
}
safe_env = {}
for key in safe_keys:
if key in os.environ:
safe_env[key] = os.environ[key]
if "NPM_CONFIG_CACHE" not in safe_env:
safe_env["NPM_CONFIG_CACHE"] = MCP_NPM_CACHE_DIR
return safe_env
async def _execute_session_operation(
self,
transport_ctx: Any,
@ -288,13 +350,24 @@ class MCPClient:
) -> TSessionResult:
"""
Execute an operation within a transport and session context.
Handles entering/exiting contexts and running the operation.
Passes sampling/elicitation/logging callbacks to the ClientSession
so that upstream MCP servers can request LLM inference (sampling),
user input (elicitation), or send log messages.
"""
transport = await transport_ctx.__aenter__()
in_flight_error: Optional[BaseException] = None
try:
read_stream, write_stream = transport[0], transport[1]
session_ctx = ClientSession(read_stream, write_stream)
# Build session kwargs with optional callbacks
session_kwargs: Dict[str, Any] = {}
if self._sampling_callback is not None:
session_kwargs["sampling_callback"] = self._sampling_callback
if self._elicitation_callback is not None:
session_kwargs["elicitation_callback"] = self._elicitation_callback
if self._logging_callback is not None:
session_kwargs["logging_callback"] = self._logging_callback
session_ctx = ClientSession(read_stream, write_stream, **session_kwargs)
session = await session_ctx.__aenter__()
try:
init_result = await session.initialize()
@ -309,11 +382,21 @@ class MCPClient:
await session_ctx.__aexit__(None, None, None)
except BaseException as e:
verbose_logger.debug(f"Error during session context exit: {e}")
except BaseException as e:
in_flight_error = e
raise
finally:
try:
await transport_ctx.__aexit__(None, None, None)
except BaseException as e:
verbose_logger.debug(f"Error during transport context exit: {e}")
except BaseException as exit_error:
verbose_logger.debug(
f"Error during transport context exit: {exit_error}"
)
root_cause = _first_non_cancelled_cause(exit_error)
if root_cause is not None and isinstance(
in_flight_error, asyncio.CancelledError
):
raise root_cause from in_flight_error
async def run_with_session(
self, operation: Callable[[ClientSession], Awaitable[TSessionResult]]
@ -351,7 +434,6 @@ class MCPClient:
def _get_auth_headers(self) -> dict:
"""Generate authentication headers based on auth type."""
headers = {}
if self._mcp_auth_value:
if isinstance(self._mcp_auth_value, str):
if self.auth_type == MCPAuth.bearer_token:
@ -373,17 +455,14 @@ class MCPClient:
# Note: aws_sigv4 auth is not handled here — SigV4 requires per-request
# signing (including the body hash), so it uses httpx.Auth flow instead
# of static headers. See MCPSigV4Auth and _create_httpx_client_factory().
# update the headers with the extra headers
if self.extra_headers:
headers.update(self.extra_headers)
return headers
return _strip_header_whitespace(headers)
def _create_httpx_client_factory(self) -> Callable[..., httpx.AsyncClient]:
"""
Create a custom httpx client factory that uses LiteLLM's SSL configuration.
This factory follows the same CA bundle path logic as http_handler.py:
1. Check ssl_verify parameter (can be SSLContext, bool, or path to CA bundle)
2. Check SSL_VERIFY environment variable
@ -400,17 +479,14 @@ class MCPClient:
"""Create an httpx.AsyncClient with LiteLLM's SSL configuration."""
# Get unified SSL configuration using the same logic as http_handler.py
ssl_config = get_ssl_configuration(self.ssl_verify)
verbose_logger.debug(
f"MCP client using SSL configuration: {type(ssl_config).__name__}"
)
# Use SigV4 auth if configured and no explicit auth provided.
# The MCP SDK's sse_client and streamable_http_client call this
# factory without passing auth=, so self._aws_auth is used.
# For non-SigV4 clients, self._aws_auth is None — no behavior change.
effective_auth = auth if auth is not None else self._aws_auth
return httpx.AsyncClient(
headers=headers,
timeout=timeout,
@ -458,7 +534,6 @@ class MCPClient:
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
@ -491,7 +566,6 @@ class MCPClient:
f"MCP Tool '{call_tool_request_params.name}' progress: "
f"{progress}/{total} ({percentage:.0f}%) - {message or ''}"
)
# Forward to Host if callback provided
if host_progress_callback:
try:
@ -514,14 +588,15 @@ class MCPClient:
)
return tool_result
except asyncio.CancelledError:
verbose_logger.warning("MCP client tool call was cancelled")
verbose_logger.warning(
f"MCP client tool call timed out after {self.timeout}s for {self.server_url}"
)
raise
except Exception as e:
import traceback
error_trace = traceback.format_exc()
verbose_logger.debug(f"MCP client tool call traceback:\n{error_trace}")
# Log detailed error information
error_type = type(e).__name__
verbose_logger.error(
@ -532,14 +607,12 @@ class MCPClient:
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream - "
"the MCP server may have crashed, disconnected, or timed out."
)
# Return a default error result instead of raising
return MCPCallToolResult(
content=[
@ -577,14 +650,12 @@ class MCPClient:
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during list_tools - "
"the MCP server may have crashed, disconnected, or timed out"
)
# Return empty list instead of raising to allow graceful degradation
return []
@ -617,7 +688,6 @@ class MCPClient:
error_trace = traceback.format_exc()
verbose_logger.debug(f"MCP client get_prompt traceback:\n{error_trace}")
# Log detailed error information
error_type = type(e).__name__
verbose_logger.error(
@ -628,14 +698,12 @@ class MCPClient:
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during get_prompt - "
"the MCP server may have crashed, disconnected, or timed out."
)
raise
async def list_resources(self) -> list[Resource]:
@ -667,14 +735,12 @@ class MCPClient:
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during list_resources - "
"the MCP server may have crashed, disconnected, or timed out"
)
# Return empty list instead of raising to allow graceful degradation
return []
@ -709,14 +775,12 @@ class MCPClient:
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during list_resource_templates - "
"the MCP server may have crashed, disconnected, or timed out"
)
# Return empty list instead of raising to allow graceful degradation
return []
@ -742,7 +806,6 @@ class MCPClient:
error_trace = traceback.format_exc()
verbose_logger.debug(f"MCP client read_resource traceback:\n{error_trace}")
# Log detailed error information
error_type = type(e).__name__
verbose_logger.error(
@ -753,12 +816,10 @@ class MCPClient:
f"Server: {self.server_url or 'stdio'}, "
f"Transport: {self.transport_type}"
)
# Check if it's a stream/connection error
if "BrokenResourceError" in error_type or "Broken" in error_type:
verbose_logger.error(
"MCP client detected broken connection/stream during read_resource - "
"the MCP server may have crashed, disconnected, or timed out."
)
raise

View file

@ -37,6 +37,8 @@ from litellm.proxy._types import (
VirtualKeyEvent,
WebhookEvent,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
from litellm.types.integrations.slack_alerting import *
from ..email_templates.templates import *
@ -1231,7 +1233,7 @@ Model Info:
and recipient_user_id is not None
and prisma_client is not None
):
user_row = await prisma_client.db.litellm_usertable.find_unique(
user_row = await UserRepository(prisma_client).table.find_unique(
where={"user_id": recipient_user_id}
)
@ -1263,7 +1265,7 @@ Model Info:
team_id = webhook_event.team_id
team_name = "Default Team"
if team_id is not None and prisma_client is not None:
team_row = await prisma_client.db.litellm_teamtable.find_unique(
team_row = await TeamRepository(prisma_client).table.find_unique(
where={"team_id": team_id}
)
if team_row is not None:

View file

@ -104,6 +104,51 @@
},
"description": "Datadog Custom Metrics Integration"
},
{
"id": "galileo",
"displayName": "Galileo",
"logo": "galileo.ico",
"supports_key_team_logging": false,
"dynamic_params": {
"GALILEO_API_KEY": {
"type": "password",
"ui_name": "API Key",
"description": "Galileo Cloud API key (app.galileo.ai). Omit for enterprise username/password auth.",
"required": false
},
"GALILEO_PROJECT_ID": {
"type": "text",
"ui_name": "Project ID",
"description": "Galileo project ID to log traces to",
"required": true
},
"GALILEO_LOG_STREAM_ID": {
"type": "text",
"ui_name": "Log Stream ID",
"description": "Galileo log stream ID for v2 spans logging (optional)",
"required": false
},
"GALILEO_BASE_URL": {
"type": "text",
"ui_name": "Base URL",
"description": "Galileo API base URL (e.g. https://api.galileo.ai for Cloud, or your enterprise API URL)",
"required": false
},
"GALILEO_USERNAME": {
"type": "text",
"ui_name": "Username",
"description": "Galileo enterprise username (legacy Observe auth; use instead of API key)",
"required": false
},
"GALILEO_PASSWORD": {
"type": "password",
"ui_name": "Password",
"description": "Galileo enterprise password (legacy Observe auth)",
"required": false
}
},
"description": "Galileo AI Observability Integration"
},
{
"id": "datadog_cost_management",
"displayName": "Datadog Cost Management",

View file

@ -47,9 +47,29 @@ from litellm.exceptions import (
BlockedPiiEntityError,
GuardrailRaisedException,
ModifyResponseException,
SensitiveDataRouteException,
)
def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]:
"""Extract session_id from request data (litellm_session_id or metadata)."""
session_id = request_data.get("litellm_session_id")
if session_id:
return str(session_id)
metadata = request_data.get("metadata") or {}
session_id = metadata.get("session_id")
if session_id:
return str(session_id)
litellm_metadata = request_data.get("litellm_metadata") or {}
session_id = litellm_metadata.get("session_id")
if session_id:
return str(session_id)
return None
class CustomGuardrail(CustomLogger):
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
use_native_during_call_hook: ClassVar[bool] = False
@ -68,6 +88,9 @@ class CustomGuardrail(CustomLogger):
end_session_after_n_fails: Optional[int] = None,
on_violation: Optional[str] = None,
realtime_violation_message: Optional[str] = None,
on_sensitive_data: Optional[str] = None,
sensitive_data_route_to_model: Optional[str] = None,
sticky_session_routing: bool = True,
**kwargs,
):
"""
@ -83,6 +106,9 @@ class CustomGuardrail(CustomLogger):
end_session_after_n_fails: For /v1/realtime sessions, end the session after this many violations
on_violation: For /v1/realtime sessions, 'warn' or 'end_session'
realtime_violation_message: Message the bot speaks aloud when a /v1/realtime guardrail fires
on_sensitive_data: Action when sensitive data is detected. 'block' (default) or 'route'
sensitive_data_route_to_model: Model to route to when on_sensitive_data='route'
sticky_session_routing: When True, all subsequent requests in the session use the same model
"""
self.guardrail_name = guardrail_name
self.supported_event_hooks = supported_event_hooks
@ -96,6 +122,11 @@ class CustomGuardrail(CustomLogger):
self.end_session_after_n_fails: Optional[int] = end_session_after_n_fails
self.on_violation: Optional[str] = on_violation
self.realtime_violation_message: Optional[str] = realtime_violation_message
self.on_sensitive_data: Optional[str] = on_sensitive_data
self.sensitive_data_route_to_model: Optional[str] = (
sensitive_data_route_to_model
)
self.sticky_session_routing: bool = sticky_session_routing
if supported_event_hooks:
## validate event_hook is in supported_event_hooks
@ -167,6 +198,108 @@ class CustomGuardrail(CustomLogger):
detection_info=detection_info,
)
def raise_sensitive_data_route_exception(
self,
route_to_model: str,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
) -> None:
"""
Raise an exception to reroute the request to a different model.
Use this when sensitive data is detected and the guardrail is configured
to route to an on-premise model instead of blocking.
The exception will reroute this request to the specified model. When
sticky_session_routing is enabled (the default), it also stores the
routing decision so subsequent requests in this session reuse the model.
Args:
route_to_model: The model to route this request (and session) to
request_data: The original request data dictionary
detection_info: Optional non-sensitive detection metadata (e.g. matched
entity types, rule ids, scores). This is surfaced in request metadata
and logs, so it must not contain the raw detected sensitive values.
Raises:
SensitiveDataRouteException: Always raises to trigger rerouting
"""
session_id = self._get_session_id_from_request_data(request_data)
if not session_id:
raise ValueError(
"Cannot route sensitive data without a session_id. "
"Ensure the request includes a session_id in metadata or headers."
)
raise SensitiveDataRouteException(
route_to_model=route_to_model,
session_id=session_id,
guardrail_name=self.guardrail_name,
detection_info=detection_info,
sticky_session_routing=self.sticky_session_routing,
)
def _get_session_id_from_request_data(
self, request_data: Dict[str, Any]
) -> Optional[str]:
"""Extract session_id from request data."""
return get_session_id_from_request_data(request_data)
def should_route_on_sensitive_data(self) -> bool:
"""
Returns True if this guardrail is configured to route requests
to a different model when sensitive data is detected.
"""
return (
self.on_sensitive_data == "route"
and self.sensitive_data_route_to_model is not None
)
def handle_sensitive_data_detection(
self,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
) -> None:
"""
Handle sensitive data detection based on guardrail configuration.
If on_sensitive_data='route', raises SensitiveDataRouteException to reroute.
Otherwise, raises GuardrailRaisedException to block. When routing is
configured but the request carries no session_id, routing is not possible
so the request falls back to a graceful block.
Args:
request_data: The request data dictionary
detection_info: Optional non-sensitive detection metadata. When routing,
this is surfaced in request metadata and logs, so it must not contain
the raw detected sensitive values.
Raises:
SensitiveDataRouteException: When configured to route and a session_id is present
GuardrailRaisedException: When configured to block, or when routing is
configured but no session_id is available
"""
if self.should_route_on_sensitive_data():
try:
self.raise_sensitive_data_route_exception(
route_to_model=self.sensitive_data_route_to_model, # type: ignore
request_data=request_data,
detection_info=detection_info,
)
except ValueError:
raise GuardrailRaisedException(
message=(
f"Sensitive data detected by {self.guardrail_name} "
"(routing skipped: request has no session_id)"
),
guardrail_name=self.guardrail_name,
)
else:
raise GuardrailRaisedException(
message=f"Sensitive data detected by {self.guardrail_name}",
guardrail_name=self.guardrail_name,
)
@staticmethod
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
"""
@ -753,12 +886,20 @@ class CustomGuardrail(CustomLogger):
Guardrails signal intentional blocks by raising:
- GuardrailRaisedException (generic guardrail API, tool permission)
- BlockedPiiEntityError (Presidio PII detection)
- SensitiveDataRouteException (sensitive-data reroute to on-premise model)
- HTTPException with status 400 (content policy violation)
- ModifyResponseException (passthrough mode violation)
"""
if isinstance(e, ModifyResponseException):
return True
if isinstance(e, (GuardrailRaisedException, BlockedPiiEntityError)):
if isinstance(
e,
(
GuardrailRaisedException,
BlockedPiiEntityError,
SensitiveDataRouteException,
),
):
return True
if (
HTTPException is not None

View file

@ -7,6 +7,7 @@ from typing import List, Optional
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm.proxy._types import WebhookEvent
from litellm.repositories.team_repository import TeamRepository
# we use this for the email header, please send a test email if you change this. verify it looks good on email
LITELLM_LOGO_URL = "https://litellm-listing.s3.amazonaws.com/litellm_logo.png"
@ -24,7 +25,7 @@ async def get_all_team_member_emails(team_id: Optional[str] = None) -> list:
if prisma_client is None:
raise Exception("Not connected to DB!")
team_row = await prisma_client.db.litellm_teamtable.find_unique(
team_row = await TeamRepository(prisma_client).table.find_unique(
where={
"team_id": team_id,
}

View file

@ -2,12 +2,14 @@
from .base import FocusDestination, FocusTimeWindow
from .factory import FocusDestinationFactory
from .gcs_destination import FocusGCSDestination
from .s3_destination import FocusS3Destination
from .vantage_destination import FocusVantageDestination
__all__ = [
"FocusDestination",
"FocusDestinationFactory",
"FocusGCSDestination",
"FocusTimeWindow",
"FocusS3Destination",
"FocusVantageDestination",

View file

@ -6,6 +6,7 @@ import os
from typing import Any, Dict, Optional
from .base import FocusDestination
from .gcs_destination import FocusGCSDestination
from .s3_destination import FocusS3Destination
from .vantage_destination import FocusVantageDestination
@ -29,6 +30,8 @@ class FocusDestinationFactory:
return FocusS3Destination(prefix=prefix, config=normalized_config)
if provider_lower == "vantage":
return FocusVantageDestination(prefix=prefix, config=normalized_config)
if provider_lower == "gcs":
return FocusGCSDestination(prefix=prefix, config=normalized_config)
raise NotImplementedError(
f"Provider '{provider}' not supported for Focus export"
)
@ -72,6 +75,18 @@ class FocusDestinationFactory:
"VANTAGE_INTEGRATION_TOKEN must be provided for Vantage exports"
)
return {k: v for k, v in resolved.items() if v is not None}
if provider == "gcs":
resolved = {
"bucket_name": overrides.get("bucket_name")
or os.getenv("FOCUS_GCS_BUCKET_NAME"),
"service_account_json": overrides.get("service_account_json")
or os.getenv("FOCUS_GCS_PATH_SERVICE_ACCOUNT"),
}
if not resolved.get("bucket_name"):
raise ValueError(
"FOCUS_GCS_BUCKET_NAME must be provided for GCS exports"
)
return {k: v for k, v in resolved.items() if v is not None}
raise NotImplementedError(
f"Provider '{provider}' not supported for Focus export configuration"
)

View file

@ -0,0 +1,74 @@
"""GCS destination for Focus export — reuses GCSBucketBase auth and httpx client."""
from __future__ import annotations
from datetime import timezone
from typing import Any, Optional
from litellm._logging import verbose_logger
from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase
from litellm.litellm_core_utils.cloud_storage_security import (
encode_gcs_object_name_for_url,
)
from .base import FocusDestination, FocusTimeWindow
class FocusGCSDestination(GCSBucketBase, FocusDestination):
"""Upload serialized Focus exports to GCS using the GCS JSON API."""
def __init__(
self,
*,
prefix: str,
config: Optional[dict[str, Any]] = None,
) -> None:
config = config or {}
bucket_name = config.get("bucket_name")
if not bucket_name:
raise ValueError("bucket_name must be provided for GCS destination")
super().__init__(bucket_name=bucket_name)
service_account_json = config.get("service_account_json")
if service_account_json is not None:
self.path_service_account_json = service_account_json
self.prefix = prefix.rstrip("/")
async def deliver(
self,
*,
content: bytes,
time_window: FocusTimeWindow,
filename: str,
) -> None:
object_name = self._build_object_key(time_window=time_window, filename=filename)
headers = await self.construct_request_headers(
service_account_json=self.path_service_account_json
)
headers["Content-Type"] = "application/octet-stream"
encoded_name = encode_gcs_object_name_for_url(object_name)
url = (
f"https://storage.googleapis.com/upload/storage/v1/b/"
f"{self.BUCKET_NAME}/o?uploadType=media&name={encoded_name}"
)
response = await self.async_httpx_client.post(
url=url, headers=headers, data=content
)
if response.status_code != 200:
raise RuntimeError(
f"GCS upload failed: status={response.status_code} body={response.text}"
)
verbose_logger.debug(
"Focus GCS: uploaded %d bytes to gs://%s/%s",
len(content),
self.BUCKET_NAME,
object_name,
)
def _build_object_key(self, *, time_window: FocusTimeWindow, filename: str) -> str:
start_utc = time_window.start_time.astimezone(timezone.utc)
date_component = f"date={start_utc.strftime('%Y-%m-%d')}"
parts = [self.prefix, date_component]
if time_window.frequency == "hourly":
parts.append(f"hour={start_utc.strftime('%H')}")
key_prefix = "/".join(filter(None, parts))
return f"{key_prefix}/{filename}" if key_prefix else filename

View file

@ -1,8 +1,13 @@
from __future__ import annotations
import json
import os
import re
from typing import Any, Dict, List, Optional, Tuple, cast
import uuid
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional, Tuple, Union, cast
import httpx
from pydantic import BaseModel, Field
import litellm
@ -12,11 +17,16 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
get_content_from_model_response,
)
from litellm.types.llms.openai import (
AllMessageValues,
HttpxBinaryResponseContent,
ResponsesAPIResponse,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus
GALILEO_CLOUD_API_BASE_URL = "https://api.galileo.ai"
# Cap the in-memory buffer so persistent flush failures (e.g. Galileo
@ -33,6 +43,11 @@ class LLMResponse(BaseModel):
model: str
num_input_tokens: int
num_output_tokens: int
num_total_tokens: int
cost: Optional[float] = Field(
default=None,
description="Total cost of the LLM call in USD as computed by LiteLLM.",
)
output_logprobs: Optional[Dict[str, Any]] = Field(
default=None,
description="Optional. When available, logprobs are used to compute Uncertainty.",
@ -75,6 +90,52 @@ class GalileoObserve(CustomLogger):
return bool(self.api_key)
return bool(self.username and self.password)
async def async_health_check(self) -> IntegrationHealthCheckStatus:
try:
if not self.project_id:
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message="GALILEO_PROJECT_ID environment variable not set",
)
if not self.base_url:
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message="GALILEO_BASE_URL environment variable not set",
)
if not self.use_v2_api and (not self.username or not self.password):
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message=(
"GALILEO_API_KEY or GALILEO_USERNAME and GALILEO_PASSWORD "
"environment variables must be set"
),
)
if not await self._ensure_headers():
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message="Galileo authentication failed",
)
response = await self.async_httpx_handler.get(
url=f"{self.base_url}/current_user",
headers=self.headers,
)
if response.status_code >= 400:
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message=(f"Galileo API returned HTTP {response.status_code}"),
)
return IntegrationHealthCheckStatus(status="healthy", error_message=None)
except Exception as e:
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message=f"Galileo health check failed: {str(e)}",
)
async def async_set_galileo_headers(self) -> None:
galileo_login_response = await self.async_httpx_handler.post(
url=f"{self.base_url}/login",
@ -121,10 +182,14 @@ class GalileoObserve(CustomLogger):
@staticmethod
def _galileo_input_messages(
messages: Optional[List[Any]], input_text: str
messages: Optional[Any], input_text: str
) -> List[Dict[str, str]]:
if isinstance(messages, dict):
messages = messages.get("messages")
if not messages:
return [{"role": "user", "content": input_text}]
if not isinstance(messages, list):
return [{"role": "user", "content": input_text}]
galileo_messages: List[Dict[str, str]] = []
for message in messages:
@ -147,13 +212,59 @@ class GalileoObserve(CustomLogger):
return [{"role": "user", "content": input_text}]
@staticmethod
def _record_to_v2_span(record: Dict[str, Any]) -> Dict[str, Any]:
created_at = record.get("created_at", "")
def _local_timezone():
return datetime.now().astimezone().tzinfo or timezone.utc
@staticmethod
def _format_created_at(dt: Union[datetime, Any]) -> str:
"""Serialize timestamps as UTC ISO-8601 for Galileo."""
if not isinstance(dt, datetime):
return str(dt)
if dt.tzinfo is None:
# LiteLLM often passes naive datetimes in local time; convert to UTC
# instead of appending Z to local time (which shifts Traces tab sorting).
dt = dt.replace(tzinfo=GalileoObserve._local_timezone())
return dt.astimezone(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
@staticmethod
def _normalize_created_at(created_at: str) -> str:
if created_at and not re.search(r"(Z|[+-]\d{2}:?\d{2})$", created_at):
created_at = f"{created_at}Z"
return f"{created_at}Z"
return created_at
@staticmethod
def _token_metrics_from_record(record: Dict[str, Any]) -> Dict[str, Any]:
num_input_tokens = int(record.get("num_input_tokens") or 0)
num_output_tokens = int(record.get("num_output_tokens") or 0)
num_total_tokens = int(record.get("num_total_tokens") or 0)
if num_total_tokens == 0 and (num_input_tokens or num_output_tokens):
num_total_tokens = num_input_tokens + num_output_tokens
metrics: Dict[str, Any] = {
"num_input_tokens": num_input_tokens,
"num_output_tokens": num_output_tokens,
"num_total_tokens": num_total_tokens,
}
cost = record.get("cost")
if cost is not None:
metrics["cost"] = float(cost)
return metrics
@staticmethod
def _record_to_v2_span(
record: Dict[str, Any],
*,
trace_id: str,
span_id: str,
) -> Dict[str, Any]:
created_at = GalileoObserve._normalize_created_at(record.get("created_at", ""))
span: Dict[str, Any] = {
"type": "llm",
"id": span_id,
"trace_id": trace_id,
"parent_id": trace_id,
"name": record.get("node_type", "litellm"),
"created_at": created_at,
"input": GalileoObserve._galileo_input_messages(
@ -167,14 +278,49 @@ class GalileoObserve(CustomLogger):
"model": record.get("model"),
"metrics": {
"duration_ns": int(record.get("latency_ms", 0)) * 1_000_000,
"num_input_tokens": record.get("num_input_tokens"),
"num_output_tokens": record.get("num_output_tokens"),
**GalileoObserve._token_metrics_from_record(record),
},
}
if record.get("tags"):
span["tags"] = record["tags"]
return span
@staticmethod
def _record_to_v2_trace(record: Dict[str, Any]) -> Dict[str, Any]:
trace_id = str(uuid.uuid4())
span_id = str(uuid.uuid4())
created_at = GalileoObserve._normalize_created_at(record.get("created_at", ""))
return {
"type": "trace",
"id": trace_id,
"name": record.get("node_type", "litellm"),
"created_at": created_at,
"input": record.get("input_text", ""),
"output": record.get("output_text", ""),
"status_code": record.get("status_code", 200),
"metrics": {
"duration_ns": int(record.get("latency_ms", 0)) * 1_000_000,
**GalileoObserve._token_metrics_from_record(record),
},
"spans": [
GalileoObserve._record_to_v2_span(
record, trace_id=trace_id, span_id=span_id
)
],
}
def _build_traces_payload(self, records: List[dict]) -> Dict[str, Any]:
payload: Dict[str, Any] = {
"traces": [self._record_to_v2_trace(record) for record in records],
"logging_method": "api_direct",
"reliable": False,
"is_complete": True,
}
if self.log_stream_id:
payload["log_stream_id"] = self.log_stream_id
return payload
def _get_ingest_request(self) -> Optional[Tuple[str, Dict[str, Any]]]:
if not self.base_url or not self.project_id:
return None
@ -184,105 +330,453 @@ class GalileoObserve(CustomLogger):
# flush_in_memory_records) aren't silently dropped when we later clear
# the in-memory buffer.
records = list(self.in_memory_records)
payload = self._build_traces_payload(records)
if self.use_v2_api:
payload: Dict[str, Any] = {
"spans": [self._record_to_v2_span(record) for record in records],
"reliable": False,
}
if self.log_stream_id:
payload["log_stream_id"] = self.log_stream_id
return (
f"{self.base_url}/v2/projects/{self.project_id}/spans",
f"{self.base_url}/ingest/traces/{self.project_id}",
payload,
)
# Username/password auth logs in for a JWT and uses the standard v2 traces API.
return (
f"{self.base_url}/projects/{self.project_id}/observe/ingest",
{"records": records},
f"{self.base_url}/v2/projects/{self.project_id}/traces",
payload,
)
@staticmethod
def _redact_headers(headers: Optional[Dict[str, str]]) -> Dict[str, str]:
if not headers:
return {}
redacted: Dict[str, str] = {}
for key, value in headers.items():
if key.lower() in {"authorization", "galileo-api-key"} and value:
redacted[key] = (
f"{value[:8]}...{value[-4:]}" if len(value) > 12 else "***"
)
else:
redacted[key] = value
return redacted
def _log_flush_config(self) -> None:
verbose_logger.debug(
"Galileo Logger flush config: use_v2_api=%s base_url=%s project_id=%s "
"log_stream_id=%s api_key_set=%s username_set=%s record_count=%s",
self.use_v2_api,
self.base_url,
self.project_id,
self.log_stream_id,
bool(self.api_key),
bool(self.username),
len(self.in_memory_records),
)
@staticmethod
def _log_v2_payload_validation(payload: Dict[str, Any]) -> None:
missing_fields: List[str] = []
traces = payload.get("traces", [])
if not traces:
missing_fields.append("traces")
for trace_index, trace in enumerate(traces):
if not isinstance(trace, dict):
continue
for field in ("id", "type", "spans"):
if field not in trace:
missing_fields.append(f"traces[{trace_index}].{field}")
trace_id = trace.get("id")
for span_index, span in enumerate(trace.get("spans", [])):
if not isinstance(span, dict):
continue
for field in ("id", "trace_id", "parent_id"):
if field not in span:
missing_fields.append(
f"traces[{trace_index}].spans[{span_index}].{field}"
)
if trace_id and span.get("trace_id") != trace_id:
missing_fields.append(
f"traces[{trace_index}].spans[{span_index}].trace_id mismatch"
)
if missing_fields:
verbose_logger.debug(
"Galileo Logger: ingest /traces payload validation issues: %s",
missing_fields,
)
def _log_flush_payload(self, url: str, payload: Dict[str, Any]) -> None:
traces = payload.get("traces", [])
verbose_logger.debug(
"Galileo Logger flush URL: %s trace_count=%s",
url,
len(traces) if isinstance(traces, list) else 0,
)
if self.use_v2_api and "/ingest/traces/" in url:
self._log_v2_payload_validation(payload)
@staticmethod
def _log_http_status_error(error: httpx.HTTPStatusError, url: str) -> None:
response = error.response
verbose_logger.debug(
"Galileo Logger HTTP error: status=%s url=%s",
response.status_code,
url,
)
verbose_logger.debug(
"Galileo Logger HTTP error response body: %s",
response.text,
)
try:
verbose_logger.debug(
"Galileo Logger HTTP error response json: %s",
response.json(),
)
except Exception:
pass
@staticmethod
def _build_prompt(kwargs: Dict[str, Any]) -> Dict[str, Any]:
optional_params = kwargs.get("optional_params", {}) or {}
prompt: Dict[str, Any] = {"messages": kwargs.get("messages")}
if optional_params.get("functions") is not None:
prompt["functions"] = optional_params["functions"]
if optional_params.get("tools") is not None:
prompt["tools"] = optional_params["tools"]
return prompt
@staticmethod
def _serialize_galileo_output(value: Any) -> str:
if value is None:
return ""
if isinstance(value, str):
return value
def _json_default(obj: Any) -> Any:
if hasattr(obj, "model_dump"):
return obj.model_dump()
return str(obj)
return json.dumps(value, default=_json_default)
@staticmethod
def _prompt_to_input_text(prompt: Dict[str, Any]) -> str:
messages = prompt.get("messages")
if messages is not None:
text = GalileoObserve._input_text_from_messages(messages)
if text:
return text
return json.dumps(prompt, default=str)
@staticmethod
def _get_chat_content_for_galileo(response_obj: litellm.ModelResponse) -> Any:
if response_obj.choices and len(response_obj.choices) > 0:
message = response_obj["choices"][0]["message"]
if hasattr(message, "json"):
message_json = message.json()
if isinstance(message_json, str):
return json.loads(message_json)
return message_json
return message
return None
@staticmethod
def _get_text_completion_content_for_galileo(
response_obj: litellm.TextCompletionResponse,
) -> Optional[str]:
if response_obj.choices and len(response_obj.choices) > 0:
return response_obj.choices[0].text
return None
@staticmethod
def _get_responses_api_content_for_galileo(
response_obj: ResponsesAPIResponse,
) -> Any:
if hasattr(response_obj, "output") and response_obj.output:
return response_obj.output
return None
@staticmethod
def _langfuse_style_rerank_prompt(kwargs: Dict[str, Any]) -> Dict[str, Any]:
"""Match Langfuse rerank input: prompt = {"messages": kwargs.get("messages")}."""
return {"messages": kwargs.get("messages")}
def _get_galileo_input_output_content(
self,
kwargs: Dict[str, Any],
response_obj: Any,
level: str = "DEFAULT",
status_message: Optional[str] = None,
) -> Tuple[str, str, Any]:
"""
Mirror Langfuse _get_langfuse_input_output_content for Galileo ingest.
Returns (input_text, output_text, messages_for_span).
"""
call_type = kwargs.get("call_type")
prompt = self._build_prompt(kwargs)
if (
level == "ERROR"
and status_message is not None
and isinstance(status_message, str)
):
return self._prompt_to_input_text(prompt), status_message, prompt
if response_obj is not None and (
call_type in ("embedding", "aembedding")
or isinstance(response_obj, litellm.EmbeddingResponse)
):
# Match Langfuse OTEL: log embeddings without serializing vectors.
return self._prompt_to_input_text(prompt), "embedding-output", prompt
if response_obj is not None and isinstance(response_obj, litellm.ModelResponse):
output = self._get_chat_content_for_galileo(response_obj)
return (
self._prompt_to_input_text(prompt),
self._serialize_galileo_output(output),
kwargs.get("messages") or [],
)
if response_obj is not None and isinstance(
response_obj, HttpxBinaryResponseContent
):
return self._prompt_to_input_text(prompt), "speech-output", prompt
if response_obj is not None and isinstance(
response_obj, litellm.TextCompletionResponse
):
output = self._get_text_completion_content_for_galileo(response_obj)
return (
self._prompt_to_input_text(prompt),
self._serialize_galileo_output(output),
kwargs.get("messages") or [],
)
if response_obj is not None and isinstance(response_obj, litellm.ImageResponse):
output = response_obj.get("data", None)
return (
self._prompt_to_input_text(prompt),
self._serialize_galileo_output(output),
prompt,
)
if response_obj is not None and isinstance(
response_obj, litellm.TranscriptionResponse
):
output = response_obj.get("text", None)
return (
self._prompt_to_input_text(prompt),
self._serialize_galileo_output(output),
prompt,
)
if response_obj is not None and isinstance(
response_obj, litellm.RerankResponse
):
output = response_obj.results
rerank_prompt = self._langfuse_style_rerank_prompt(kwargs)
return (
json.dumps(rerank_prompt, default=str),
self._serialize_galileo_output(output),
rerank_prompt,
)
if response_obj is not None and isinstance(response_obj, ResponsesAPIResponse):
output = self._get_responses_api_content_for_galileo(response_obj)
return (
self._prompt_to_input_text(prompt),
self._serialize_galileo_output(output),
kwargs.get("messages") or [],
)
if (
call_type == "_arealtime"
and response_obj is not None
and isinstance(response_obj, list)
):
input_val = kwargs.get("input")
return (
self._serialize_galileo_output(input_val),
self._serialize_galileo_output(response_obj),
input_val,
)
if (
call_type == "pass_through_endpoint"
and response_obj is not None
and isinstance(response_obj, dict)
):
output = response_obj.get("response", "")
return (
self._prompt_to_input_text(prompt),
self._serialize_galileo_output(output),
prompt,
)
if response_obj is not None and isinstance(response_obj, dict):
output = get_content_from_model_response(response_obj)
return (
self._prompt_to_input_text(prompt),
self._serialize_galileo_output(output),
kwargs.get("messages") or [],
)
return self._prompt_to_input_text(prompt), "", kwargs.get("messages") or []
def get_output_str_from_response(
self, response_obj: Any, kwargs: Dict[str, Any]
) -> Optional[str]:
if response_obj is None:
return None
if kwargs.get("call_type", None) == "embedding" or isinstance(
response_obj, litellm.EmbeddingResponse
):
return None
if isinstance(response_obj, litellm.TextCompletionResponse):
return response_obj.choices[0].text
if isinstance(response_obj, litellm.ImageResponse):
return json.dumps(response_obj["data"], default=str)
if isinstance(response_obj, (litellm.ModelResponse, dict)):
return get_content_from_model_response(response_obj)
return None
) -> str:
_, output_text, _ = self._get_galileo_input_output_content(
kwargs=kwargs, response_obj=response_obj
)
return output_text
@staticmethod
def _input_text_from_messages(messages: Any) -> str:
"""Return a plain-string summary of the input suitable for the trace-level input field."""
if isinstance(messages, str):
return messages
if not isinstance(messages, list):
return ""
# Use the last user/human message so the trace table shows the actual prompt
for msg in reversed(messages):
if not isinstance(msg, dict):
continue
if str(msg.get("role", "")).lower() in ("user", "human"):
content = msg.get("content") or ""
if isinstance(content, list):
content = " ".join(
b.get("text", "") if isinstance(b, dict) else str(b)
for b in content
)
if content:
return str(content)
# Fallback: first non-empty content of any role
for msg in messages:
if isinstance(msg, dict):
content = msg.get("content") or ""
if isinstance(content, list):
content = " ".join(
b.get("text", "") if isinstance(b, dict) else str(b)
for b in content
)
if content:
return str(content)
return ""
async def async_log_success_event(
self, kwargs: Any, response_obj: Any, start_time: Any, end_time: Any
):
verbose_logger.debug("On Async Success")
try:
await self._async_log_success_event_impl(
kwargs=kwargs,
response_obj=response_obj,
start_time=start_time,
end_time=end_time,
)
except Exception:
verbose_logger.exception(
"Galileo Logger: unexpected error in async_log_success_event"
)
async def _async_log_success_event_impl(
self, kwargs: Any, response_obj: Any, start_time: Any, end_time: Any
):
if not self._is_configured():
verbose_logger.debug(
"Galileo Logger: skipping flush — set GALILEO_PROJECT_ID and "
"either GALILEO_API_KEY (hosted) or GALILEO_USERNAME/GALILEO_PASSWORD "
"(enterprise Observe)."
"Galileo Logger: skipping — GALILEO_PROJECT_ID=%s GALILEO_API_KEY=%s GALILEO_BASE_URL=%s",
bool(self.project_id),
bool(self.api_key),
bool(self.base_url),
)
return
_latency_ms = int((end_time - start_time).total_seconds() * 1000)
_call_type = kwargs.get("call_type", "litellm")
input_text = litellm.utils.get_formatted_prompt(
data=kwargs, call_type=_call_type
slo: Optional[Dict[str, Any]] = kwargs.get("standard_logging_object")
if slo is None:
verbose_logger.debug(
"Galileo Logger: no standard_logging_object in kwargs, skipping"
)
return
_call_type: str = str(
slo.get("call_type") or kwargs.get("call_type") or "litellm"
)
_usage = response_obj.get("usage", {}) or {}
num_input_tokens = _usage.get("prompt_tokens", 0)
num_output_tokens = _usage.get("completion_tokens", 0)
output_text = self.get_output_str_from_response(
response_obj=response_obj, kwargs=kwargs
input_text, output_text, messages = self._get_galileo_input_output_content(
kwargs=kwargs, response_obj=response_obj
)
if output_text is not None:
request_record = LLMResponse(
latency_ms=_latency_ms,
status_code=200,
input_text=input_text,
output_text=output_text,
node_type=_call_type,
model=kwargs.get("model", "-"),
num_input_tokens=num_input_tokens,
num_output_tokens=num_output_tokens,
created_at=start_time.strftime(
"%Y-%m-%dT%H:%M:%S"
), # timestamp str constructed in "%Y-%m-%dT%H:%M:%S" format
raw_start = slo.get("startTime")
raw_end = slo.get("endTime")
if raw_start is None or raw_end is None:
verbose_logger.debug(
"Galileo Logger: standard_logging_object missing startTime/endTime, "
"falling back to start_time/end_time params"
)
if not isinstance(start_time, datetime) or not isinstance(
end_time, datetime
):
return
start_ts = start_time
end_ts = end_time
if start_ts.tzinfo is None:
start_ts = start_ts.replace(tzinfo=GalileoObserve._local_timezone())
if end_ts.tzinfo is None:
end_ts = end_ts.replace(tzinfo=GalileoObserve._local_timezone())
start_ts = start_ts.astimezone(timezone.utc)
end_ts = end_ts.astimezone(timezone.utc)
else:
start_ts = datetime.fromtimestamp(float(raw_start), tz=timezone.utc)
end_ts = datetime.fromtimestamp(float(raw_end), tz=timezone.utc)
_latency_ms = max(0, int((end_ts - start_ts).total_seconds() * 1000))
num_input_tokens = int(slo.get("prompt_tokens") or 0)
num_output_tokens = int(slo.get("completion_tokens") or 0)
num_total_tokens = int(slo.get("total_tokens") or 0)
if num_total_tokens == 0 and (num_input_tokens or num_output_tokens):
num_total_tokens = num_input_tokens + num_output_tokens
request_record = LLMResponse(
latency_ms=_latency_ms,
status_code=200,
input_text=input_text,
output_text=output_text,
node_type=_call_type,
model=str(slo.get("model") or kwargs.get("model") or "-"),
num_input_tokens=num_input_tokens,
num_output_tokens=num_output_tokens,
num_total_tokens=num_total_tokens,
cost=slo.get("response_cost"),
created_at=GalileoObserve._format_created_at(start_ts),
)
request_dict = request_record.model_dump()
if isinstance(messages, dict):
messages = messages.get("messages")
if isinstance(messages, list) and messages:
request_dict["messages"] = messages
self.in_memory_records.append(request_dict)
verbose_logger.debug(
"Galileo Logger: queued record, in_memory=%d", len(self.in_memory_records)
)
# Bound the buffer so persistent flush failures cannot grow it
# without limit. Drop the oldest records once we exceed the cap.
if len(self.in_memory_records) > GALILEO_MAX_IN_MEMORY_RECORDS:
dropped = len(self.in_memory_records) - GALILEO_MAX_IN_MEMORY_RECORDS
self.in_memory_records = self.in_memory_records[
-GALILEO_MAX_IN_MEMORY_RECORDS:
]
verbose_logger.warning(
"Galileo Logger: in-memory buffer exceeded %s records; "
"dropped %s oldest record(s). Check Galileo connectivity/credentials.",
GALILEO_MAX_IN_MEMORY_RECORDS,
dropped,
)
request_dict = request_record.model_dump()
messages = kwargs.get("messages")
if messages:
request_dict["messages"] = messages
self.in_memory_records.append(request_dict)
# Bound the buffer so persistent flush failures cannot grow it
# without limit. Drop the oldest records once we exceed the cap.
if len(self.in_memory_records) > GALILEO_MAX_IN_MEMORY_RECORDS:
dropped = len(self.in_memory_records) - GALILEO_MAX_IN_MEMORY_RECORDS
self.in_memory_records = self.in_memory_records[
-GALILEO_MAX_IN_MEMORY_RECORDS:
]
verbose_logger.warning(
"Galileo Logger: in-memory buffer exceeded %s records; "
"dropped %s oldest record(s). Check Galileo connectivity/credentials.",
GALILEO_MAX_IN_MEMORY_RECORDS,
dropped,
)
if len(self.in_memory_records) >= self.batch_size:
await self.flush_in_memory_records()
if len(self.in_memory_records) >= self.batch_size:
await self.flush_in_memory_records()
async def flush_in_memory_records(self):
if not self.in_memory_records:
@ -296,15 +790,23 @@ class GalileoObserve(CustomLogger):
ingest_request = self._get_ingest_request()
if ingest_request is None:
verbose_logger.debug(
"Galileo Logger: missing GALILEO_BASE_URL or GALILEO_PROJECT_ID"
"Galileo Logger: missing GALILEO_BASE_URL or GALILEO_PROJECT_ID — skipping flush"
)
return
if not await self._ensure_headers():
verbose_logger.debug("Galileo Logger: could not set request headers")
verbose_logger.debug(
"Galileo Logger: could not set request headers — skipping flush"
)
return
url, payload = ingest_request
self._log_flush_config()
self._log_flush_payload(url=url, payload=payload)
verbose_logger.debug(
"Galileo Logger flush headers: %s",
self._redact_headers(self.headers),
)
verbose_logger.debug("flushing in memory records to %s", url)
try:
@ -313,6 +815,12 @@ class GalileoObserve(CustomLogger):
headers=self.headers,
json=payload,
)
except httpx.HTTPStatusError as e:
self._log_http_status_error(error=e, url=url)
verbose_logger.debug(
"Galileo Logger: failed to flush in memory records: %s", e
)
return
except Exception as e:
verbose_logger.debug(
"Galileo Logger: failed to flush in memory records: %s", e
@ -323,6 +831,11 @@ class GalileoObserve(CustomLogger):
verbose_logger.debug(
"Galileo Logger: successfully flushed in memory records"
)
verbose_logger.debug(
"Galileo Logger flush response: status=%s body=%s",
response.status_code,
response.text,
)
del self.in_memory_records[:records_in_payload]
else:
verbose_logger.debug("Galileo Logger: failed to flush in memory records")

View file

@ -102,6 +102,18 @@ def langfuse_client_init(
if Version(langfuse.version.__version__) >= Version("2.6.0"):
parameters["sdk_integration"] = "litellm"
if Version(langfuse.version.__version__) >= Version("2.7.3"):
import httpx
import litellm
from ...llms.custom_httpx.http_handler import get_ssl_configuration
parameters["httpx_client"] = httpx.Client(
verify=get_ssl_configuration(),
cert=os.getenv("SSL_CERTIFICATE", litellm.ssl_certificate),
)
client = Langfuse(**parameters)
return client

View file

@ -65,7 +65,15 @@ class OpenMeterLogger(CustomLogger):
"total_tokens": response_obj["usage"].get("total_tokens"),
}
user_param = kwargs.get("user", None) # end-user passed in via 'user' param
# OPENMETER_TRUST_REQUEST_USER (default "true"): when set to "false",
# the request-supplied `user` field is ignored and the subject is
# resolved solely from the key-bound user_api_key_user_id. Proxies
# serving multi-tenant traffic enable this to prevent clients from
# forging attribution by setting `user` in the request body.
trust_request_user = (
os.getenv("OPENMETER_TRUST_REQUEST_USER", "true").lower() != "false"
)
user_param = kwargs.get("user", None) if trust_request_user else None
# If no user provided directly, try to get it from token user_id
if user_param is None:

View file

@ -15,6 +15,7 @@ from litellm.integrations.otel.model.baggage import promoted_baggage
from litellm.integrations.otel.model.config import OpenTelemetryV2Config
from litellm.integrations.otel.plumbing.context import (
is_recordable_span,
request_root_span,
resolve_parent_context,
resolve_request_span_context,
set_request_baggage,
@ -435,8 +436,12 @@ class OpenTelemetryV2(CustomLogger):
attach(set_request_baggage(bag, context=get_current()))
# The server span was started by the instrumentor before this ran,
# so the Baggage processor (which only fires at span start) won't
# backfill it — stamp identity on it directly.
server_span = get_current_span()
# backfill it — stamp identity on it directly. Prefer the anchored
# root span over the ambient one so identity still lands on the
# server span when seeding from inside the live ``auth`` phase span
# (the auth-failure path), where ``get_current_span`` is the phase
# span, not the request's root.
server_span = request_root_span() or get_current_span()
if is_recordable_span(server_span):
# Re-capture the anchor here too: this runs post-auth with the
# server span active and covers entrypoints that bypass

View file

@ -24,14 +24,18 @@ from typing import (
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import (
BoundedPrometheusSeriesTracker,
from litellm.exceptions import (
validate_rate_limit_category,
validate_rate_limit_type,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prometheus_helpers import (
PrometheusLabelFactoryContext,
_get_cached_end_user_id_for_cost_tracking,
)
from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import (
BoundedPrometheusSeriesTracker,
)
from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
get_metadata_variable_name_from_kwargs,
@ -42,6 +46,9 @@ from litellm.proxy._types import (
LiteLLM_UserTable,
UserAPIKeyAuth,
)
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
from litellm.types.integrations.prometheus import *
from litellm.types.integrations.prometheus import (
_sanitize_prometheus_label_name,
@ -78,6 +85,20 @@ class PrometheusLogger(CustomLogger):
# Always initialize label_filters, even for non-premium users
self.label_filters = self._parse_prometheus_config()
# Cache resolved label sets per metric. Several entries in
# ``PrometheusMetricLabels.get_labels`` read module-level toggles
# (e.g. ``litellm.prometheus_emit_stream_label``,
# ``litellm.prometheus_emit_rate_limit_labels``) that can be
# changed at runtime. Prometheus counters/gauges/histograms are
# created with a *fixed* ``labelnames`` set; if a runtime call
# to ``get_labels_for_metric`` returned a different set, the
# subsequent ``counter.labels(**_labels)`` would raise a
# ``ValueError`` from the prometheus client. Snapshotting at
# logger init time pins the label set for the lifetime of the
# logger so toggling these flags only takes effect after a
# restart, keeping init-time and runtime label sets in sync.
self._cached_metric_labels: Dict[str, List[str]] = {}
_custom_buckets = litellm.prometheus_latency_buckets
self.latency_buckets = (
tuple(_custom_buckets)
@ -1033,13 +1054,27 @@ class PrometheusLogger(CustomLogger):
self, metric_name: DEFINED_PROMETHEUS_METRICS
) -> List[str]:
"""
Get the labels for a metric, filtered if configured
Get the labels for a metric, filtered if configured.
The result is cached on the instance so the label set used to
construct each Prometheus metric at ``__init__`` time stays in lock
step with the label set passed to ``counter.labels(...)`` at
runtime, even if the underlying module-level toggles consulted by
:meth:`PrometheusMetricLabels.get_labels` (e.g.
``litellm.prometheus_emit_rate_limit_labels``,
``litellm.prometheus_emit_stream_label``) are flipped after the
logger has been created.
"""
cached = self._cached_metric_labels.get(metric_name)
if cached is not None:
return cached
# Get default labels for this metric from PrometheusMetricLabels
default_labels = PrometheusMetricLabels.get_labels(metric_name)
# If no label filtering is configured for this metric, use default labels
if metric_name not in self.label_filters:
self._cached_metric_labels[metric_name] = default_labels
return default_labels
# Get configured labels for this metric
@ -1050,6 +1085,7 @@ class PrometheusLogger(CustomLogger):
label for label in default_labels if label in configured_labels
]
self._cached_metric_labels[metric_name] = filtered_labels
return filtered_labels
def _track_end_user_metric_series(
@ -2029,14 +2065,8 @@ class PrometheusLogger(CustomLogger):
Proxy level tracking - failed client side requests
labelnames=[
"end_user",
"hashed_api_key",
"api_key_alias",
REQUESTED_MODEL,
"team",
"team_alias",
] + EXCEPTION_LABELS,
See :attr:`PrometheusMetricLabels.litellm_proxy_failed_requests_metric`
for the authoritative list of labels emitted on this metric.
"""
from litellm.litellm_core_utils.litellm_logging import (
StandardLoggingPayloadSetup,
@ -2059,6 +2089,9 @@ class PrometheusLogger(CustomLogger):
model_id = _metadata.get("model_info", {}).get("id") or request_data.get(
"model_info", {}
).get("id")
rate_limit_category, rate_limit_type = self._extract_rate_limit_labels(
original_exception
)
enum_values = UserAPIKeyLabelValues(
end_user=user_api_key_dict.end_user_id,
user=user_api_key_dict.user_id,
@ -2073,6 +2106,8 @@ class PrometheusLogger(CustomLogger):
status_code=str(status_code),
exception_status=str(status_code),
exception_class=self._get_exception_class_name(original_exception),
rate_limit_category=rate_limit_category,
rate_limit_type=rate_limit_type,
tags=_tags,
route=user_api_key_dict.request_route,
client_ip=_metadata.get("requester_ip_address"),
@ -2690,7 +2725,7 @@ class PrometheusLogger(CustomLogger):
Args:
guardrail_name: Name of the guardrail
latency_seconds: Execution latency in seconds
status: "success" or "error"
status: "success", "error", or "intervened"
error_type: Type of error if any, None otherwise
hook_type: "pre_call", "during_call", or "post_call"
"""
@ -2843,6 +2878,33 @@ class PrometheusLogger(CustomLogger):
@staticmethod
def _get_exception_class_name(exception: Exception) -> str:
# Some exception types pin the ``exception_class`` label to a legacy
# value for back-compat with existing dashboards (e.g. proxy-side 429s
# keep reporting as "HTTPException"). Honor that opt-in marker before
# deriving the label from the runtime class name. Reading it via
# ``getattr`` keeps this core integrations module free of a transitive
# ``fastapi`` dependency.
legacy_class_name = getattr(exception, "prometheus_exception_class_name", None)
if isinstance(legacy_class_name, str) and legacy_class_name:
return legacy_class_name
# Same back-compat reasoning for ``BudgetExceededError``: the unified
# rate-limit error work attached ``.llm_provider`` to budget errors
# too (so callbacks reading ``StandardLoggingPayload`` get provider
# attribution). Without this short-circuit, the provider prefix below
# would silently flip the label from "BudgetExceededError" to e.g.
# "Openai.BudgetExceededError" and break dashboards keyed on the
# original value.
try:
from litellm.exceptions import BudgetExceededError
except ImportError:
BudgetExceededError = None # type: ignore[assignment,misc]
if BudgetExceededError is not None and isinstance(
exception, BudgetExceededError
):
return "BudgetExceededError"
exception_class_name = ""
if hasattr(exception, "llm_provider"):
exception_class_name = getattr(exception, "llm_provider") or ""
@ -2857,6 +2919,27 @@ class PrometheusLogger(CustomLogger):
exception_class_name += exception.__class__.__name__
return exception_class_name
@staticmethod
def _extract_rate_limit_labels(
exception: Optional[Exception],
) -> Tuple[Optional[str], Optional[str]]:
"""
Pull the unified ``category`` / ``rate_limit_type`` fields off any
exception that declares them (``litellm.RateLimitError`` and bare-
Exception subclasses like ``BudgetExceededError``).
Values are validated against the :class:`RateLimitErrorCategory` /
:class:`RateLimitType` enums so unrelated third-party exceptions that
happen to declare ``.category`` / ``.rate_limit_type`` string attributes
can't leak garbage into Prometheus label cardinality.
"""
if exception is None:
return None, None
return (
validate_rate_limit_category(getattr(exception, "category", None)),
validate_rate_limit_type(getattr(exception, "rate_limit_type", None)),
)
async def log_success_fallback_event(
self, original_model_group: str, kwargs: dict, original_exception: Exception
):
@ -3198,12 +3281,12 @@ class PrometheusLogger(CustomLogger):
page_size: int, page: int
) -> Tuple[List[LiteLLM_UserTable], Optional[int]]:
skip = (page - 1) * page_size
users = await prisma_client.db.litellm_usertable.find_many(
users = await UserRepository(prisma_client).table.find_many(
skip=skip,
take=page_size,
order={"created_at": "desc"},
)
total_count = await prisma_client.db.litellm_usertable.count()
total_count = await UserRepository(prisma_client).table.count()
return users, total_count
await self._initialize_budget_metrics(
@ -3226,13 +3309,13 @@ class PrometheusLogger(CustomLogger):
async def fetch_orgs(page_size: int, page: int) -> Tuple[list, Optional[int]]:
skip = (page - 1) * page_size
orgs = await prisma_client.db.litellm_organizationtable.find_many(
orgs = await OrganizationRepository(prisma_client).table.find_many(
skip=skip,
take=page_size,
order={"created_at": "desc"},
include={"litellm_budget_table": True},
)
total_count = await prisma_client.db.litellm_organizationtable.count()
total_count = await OrganizationRepository(prisma_client).table.count()
return orgs, total_count
await self._initialize_budget_metrics(
@ -3300,14 +3383,14 @@ class PrometheusLogger(CustomLogger):
try:
# Get total user count
total_users = await prisma_client.db.litellm_usertable.count()
total_users = await UserRepository(prisma_client).table.count()
self.litellm_total_users_metric.set(total_users)
verbose_logger.debug(
f"Prometheus: set litellm_total_users to {total_users}"
)
# Get total team count
total_teams = await prisma_client.db.litellm_teamtable.count()
total_teams = await TeamRepository(prisma_client).table.count()
self.litellm_teams_count_metric.set(total_teams)
verbose_logger.debug(
f"Prometheus: set litellm_teams_count to {total_teams}"

View file

@ -244,6 +244,9 @@ search_tools:
- search_tool_name: "my-tavily-tool"
litellm_params:
search_provider: "tavily"
- search_tool_name: "my-you-com-tool"
litellm_params:
search_provider: "you_com"
```
---

View file

@ -655,7 +655,11 @@ def exception_type( # type: ignore # noqa: PLR0915
custom_llm_provider == "anthropic"
or custom_llm_provider == "anthropic_text"
): # one of the anthropics
if "prompt is too long" in error_str or "prompt: length" in error_str:
if (
"prompt is too long" in error_str
or "prompt: length" in error_str
or ExceptionCheckers.is_error_str_context_window_exceeded(error_str)
):
exception_mapping_worked = True
raise ContextWindowExceededError(
message="AnthropicError - {}".format(error_str),

View file

@ -659,6 +659,11 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
or get_secret_str("NVIDIA_RIVA_API_KEY")
or get_secret_str("NVIDIA_NIM_API_KEY")
)
elif custom_llm_provider == "soniox":
api_base = (
api_base or get_secret_str("SONIOX_API_BASE") or "https://api.soniox.com"
)
dynamic_api_key = api_key or get_secret_str("SONIOX_API_KEY")
elif custom_llm_provider == "cerebras":
api_base = (
api_base or get_secret("CEREBRAS_API_BASE") or "https://api.cerebras.ai/v1"

View file

@ -22,9 +22,11 @@ def get_supported_openai_params( # noqa: PLR0915
```
Args:
base_model: For Azure, the true underlying model (e.g. ``"azure/gpt-5.2"``)
when the deployment name differs. Used for model-type detection so that
non-standard deployment names route to the correct config.
base_model: An optional capability hint for deployments whose ``model``
label isn't recognized on its own (e.g. an Azure deployment name, or a
friendly Bedrock alias). It is additive: the result is the union of the
params supported by ``model`` and by ``base_model``, so a hint can only
add capabilities, never strip ones the real model already supports.
Returns:
- List if custom_llm_provider is mapped
@ -52,7 +54,15 @@ def get_supported_openai_params( # noqa: PLR0915
provider_config = None
if provider_config and request_type == "chat_completion":
return provider_config.get_supported_openai_params(model=base_model or model)
supported_params = provider_config.get_supported_openai_params(model=model)
if base_model and base_model != model:
base_model_params = provider_config.get_supported_openai_params(
model=base_model
)
supported_params = list(
dict.fromkeys([*supported_params, *base_model_params])
)
return supported_params
if custom_llm_provider == "bedrock":
return litellm.AmazonConverseConfig().get_supported_openai_params(model=model)
@ -331,6 +341,11 @@ def get_supported_openai_params( # noqa: PLR0915
return ElevenLabsAudioTranscriptionConfig().get_supported_openai_params(
model=model
)
elif custom_llm_provider == "soniox":
if request_type == "transcription":
return litellm.SonioxAudioTranscriptionConfig().get_supported_openai_params(
model=model
)
elif custom_llm_provider in litellm._custom_providers:
if request_type == "chat_completion":
provider_config = litellm.ProviderConfigManager.get_provider_chat_config(

View file

@ -37,6 +37,10 @@ from litellm import (
turn_off_message_logging,
)
from litellm._logging import _is_debugging_on, _redact_string, verbose_logger
from litellm.exceptions import (
validate_rate_limit_category,
validate_rate_limit_type,
)
from litellm._uuid import uuid
from litellm.batches.batch_utils import _handle_completed_batch
from litellm.caching.caching import DualCache, InMemoryCache
@ -3503,7 +3507,9 @@ class Logging(LiteLLMLoggingBaseClass):
else:
return None
def _handle_anthropic_messages_response_logging(self, result: Any) -> ModelResponse:
def _handle_anthropic_messages_response_logging(
self, result: Any
) -> Union[ModelResponse, ResponsesAPIResponse]:
"""
Handles logging for Anthropic messages responses.
@ -3522,6 +3528,15 @@ class Logging(LiteLLMLoggingBaseClass):
return result
elif isinstance(result, ModelResponse):
return result
elif isinstance(
result,
(ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent),
):
# anthropic_messages() can route to OpenAI Responses API; in that path
# the assembled streaming result is one of these terminal events rather than
# a ModelResponse. Return the inner response so downstream handlers
# (_transform_usage_objects, normalize_logging_result) can process it.
return result.response
httpx_response = self.model_call_details.get("httpx_response", None)
if httpx_response and isinstance(httpx_response, httpx.Response):
@ -5307,12 +5322,27 @@ class StandardLoggingPayloadSetup:
else str(original_exception)
)
# Duck-typed read so bare-Exception subclasses like
# `litellm.BudgetExceededError` can participate without joining the
# RateLimitError hierarchy (which would break `except BudgetExceededError`).
# Validated against the enum value sets so a third-party exception that
# happens to declare a `.category` or `.rate_limit_type` string attribute
# can't leak garbage into the payload or Prometheus label cardinality.
rate_limit_category = validate_rate_limit_category(
getattr(original_exception, "category", None)
)
rate_limit_type = validate_rate_limit_type(
getattr(original_exception, "rate_limit_type", None)
)
return StandardLoggingPayloadErrorInformation(
error_code=error_status,
error_class=error_class,
llm_provider=_llm_provider_in_exception,
traceback=traceback_info,
error_message=error_message if original_exception else "",
error_rate_limit_category=rate_limit_category,
error_rate_limit_type=rate_limit_type,
)
@staticmethod

View file

@ -3653,17 +3653,13 @@ from litellm.types.llms.bedrock import ContentBlock as BedrockContentBlock
from litellm.types.llms.bedrock import DocumentBlock as BedrockDocumentBlock
from litellm.types.llms.bedrock import ImageBlock as BedrockImageBlock
from litellm.types.llms.bedrock import SourceBlock as BedrockSourceBlock
from litellm.types.llms.bedrock import BedrockToolSpec
from litellm.types.llms.bedrock import ToolBlock as BedrockToolBlock
from litellm.types.llms.bedrock import (
ToolInputSchemaBlock as BedrockToolInputSchemaBlock,
)
from litellm.types.llms.bedrock import ToolJsonSchemaBlock as BedrockToolJsonSchemaBlock
from litellm.types.llms.bedrock import SearchResultBlock
from litellm.types.llms.bedrock import ToolResultBlock as BedrockToolResultBlock
from litellm.types.llms.bedrock import (
ToolResultContentBlock as BedrockToolResultContentBlock,
)
from litellm.types.llms.bedrock import ToolSpecBlock as BedrockToolSpecBlock
from litellm.types.llms.bedrock import ToolUseBlock as BedrockToolUseBlock
from litellm.types.llms.bedrock import VideoBlock as BedrockVideoBlock
@ -5496,6 +5492,7 @@ def _bedrock_tools_pt(
]
"""
from litellm.llms.bedrock.common_utils import (
get_bedrock_base_model,
normalize_json_schema_custom_types_to_object,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_defs
@ -5503,6 +5500,11 @@ def _bedrock_tools_pt(
_valid_json_schema_root_types = frozenset(
("array", "boolean", "integer", "null", "number", "object", "string")
)
# Only Claude on Bedrock honours strict tool schemas; other families
# (Nova, Llama, GPT-OSS) reject the strict field outright.
supports_strict_tools = bool(
model and get_bedrock_base_model(model).startswith("anthropic")
)
tool_block_list: List[BedrockToolBlock] = []
for tool_idx, tool in enumerate(tools):
# Check if tool is already a BedrockToolBlock (e.g., systemTool for Nova grounding)
@ -5548,17 +5550,16 @@ def _bedrock_tools_pt(
normalize_json_schema_custom_types_to_object(parameters)
if parameters.get("type") not in _valid_json_schema_root_types:
parameters["type"] = "object"
tool_input_schema = BedrockToolInputSchemaBlock(
json=BedrockToolJsonSchemaBlock(
type=parameters["type"],
properties=parameters.get("properties", {}),
required=parameters.get("required", []),
)
tool_block = cast(
BedrockToolBlock,
BedrockToolSpec(
name=name,
description=description,
parameters=parameters,
strict=tool.get("function", {}).get("strict", None),
supports_strict_tools=supports_strict_tools,
),
)
tool_spec = BedrockToolSpecBlock(
inputSchema=tool_input_schema, name=name, description=description
)
tool_block = BedrockToolBlock(toolSpec=tool_spec)
tool_block_list.append(tool_block)
## ADD CACHE POINT TOOL BLOCK ##

View file

@ -1149,6 +1149,32 @@ class CustomStreamWrapper:
completion_obj: Dict[str, Any] = {"content": ""}
from litellm.types.utils import GenericStreamingChunk as GChunk
if (
isinstance(chunk, ModelResponseStream)
and self.custom_llm_provider is not None
and self.custom_llm_provider in litellm._custom_providers
):
_has_content = bool(
chunk.choices
and chunk.choices[0].delta is not None
and (
chunk.choices[0].delta.content
or chunk.choices[0].delta.tool_calls
)
)
if self.received_finish_reason is not None:
if not _has_content:
raise StopIteration
if chunk.choices and chunk.choices[0].finish_reason:
self.received_finish_reason = chunk.choices[0].finish_reason
if not _has_content:
return None
# Strip finish_reason from the content chunk so it appears
# only on the trailing empty-delta chunk (OpenAI spec).
# finish_reason_handler() will emit the proper terminal chunk.
chunk.choices[0].finish_reason = None # type: ignore[assignment]
return chunk
if (
isinstance(chunk, dict)
and generic_chunk_has_all_required_fields(

View file

@ -81,7 +81,6 @@ from litellm.types.utils import (
from litellm.utils import (
ModelResponse,
Usage,
_supports_factory,
add_dummy_tool,
any_assistant_message_has_thinking_blocks,
get_max_tokens,
@ -337,50 +336,6 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
v in model_lower for v in ("opus-4-7", "opus_4_7", "opus-4.7", "opus_4.7")
)
@staticmethod
def _supports_model_capability(model: str, key: str) -> bool:
"""Check a boolean capability ``key`` in the model map.
Strips bedrock/vertex prefixes so a provider-routed Claude still
resolves to the Anthropic model-map entry.
"""
try:
if _supports_factory(
model=model,
custom_llm_provider="anthropic",
key=key,
):
return True
except Exception:
pass
candidates = [model]
for prefix in (
"bedrock/converse/",
"bedrock/invoke/",
"bedrock/",
"vertex_ai/",
):
if model.startswith(prefix):
candidates.append(model[len(prefix) :])
try:
from litellm.llms.bedrock.common_utils import BedrockModelInfo
base = BedrockModelInfo.get_base_model(model)
if base:
candidates.append(base)
candidates.append(f"bedrock/{base}")
except Exception:
pass
try:
for cand in candidates:
if cand in litellm.model_cost and (
litellm.model_cost[cand].get(key) is True
):
return True
except Exception:
pass
return False
@staticmethod
def _supports_effort_level(model: str, level: str) -> bool:
"""Check ``supports_{level}_reasoning_effort`` in the model map."""
@ -918,7 +873,39 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
anthropic_tools = []
mcp_servers = []
for tool in tools:
if "input_schema" in tool: # assume in anthropic format
if tool.get("type") == "namespace":
# Namespace is a grouping container (e.g. codex's multi_agent_v1).
# Extract its nested tools and map them individually.
for nested in tool.get("tools") or []:
if "input_schema" in nested:
# Already in Anthropic format.
anthropic_tools.append(nested)
elif "function" not in nested and "name" in nested:
# Flat format: {type, name, description, parameters, ...}.
# Normalize to OpenAI-wrapped format before mapping.
wrapped = cast(
ChatCompletionToolParam,
{
"type": nested.get("type", "function"),
"function": {
k: v for k, v in nested.items() if k != "type"
},
},
)
nested_tool, nested_mcp = self._map_tool_helper(wrapped)
if nested_tool is not None:
anthropic_tools.append(nested_tool)
if nested_mcp is not None:
mcp_servers.append(nested_mcp)
elif "function" in nested:
nested_tool, nested_mcp = self._map_tool_helper(
cast(ChatCompletionToolParam, nested)
)
if nested_tool is not None:
anthropic_tools.append(nested_tool)
if nested_mcp is not None:
mcp_servers.append(nested_mcp)
elif "input_schema" in tool: # assume in anthropic format
anthropic_tools.append(tool)
else: # assume openai tool call
new_tool, mcp_server_tool = self._map_tool_helper(tool)
@ -1620,6 +1607,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
return _tool
def should_strip_billing_metadata(self) -> bool:
"""
Whether to drop x-anthropic-billing-header system blocks before sending upstream.
The first-party Anthropic API uses these blocks for Claude Code attribution, so the
base config keeps them. Providers that reject them (e.g. Bedrock) override this to True.
"""
return False
def translate_system_message(
self, messages: List[AllMessageValues]
) -> List[AnthropicSystemMessageContent]:
@ -1627,7 +1623,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
Translate system message to anthropic format.
Removes system message from the original list and returns a new list of anthropic system message content.
Filters out system messages containing x-anthropic-billing-header metadata.
When should_strip_billing_metadata() is True, x-anthropic-billing-header system blocks are dropped.
"""
system_prompt_indices = []
anthropic_system_message_list: List[AnthropicSystemMessageContent] = []
@ -1639,10 +1635,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
# Skip empty text blocks - Anthropic API raises errors for empty text
if not system_message_block["content"]:
continue
# Skip system messages containing x-anthropic-billing-header metadata
if system_message_block["content"].startswith(
"x-anthropic-billing-header:"
):
if self.should_strip_billing_metadata() and system_message_block[
"content"
].startswith("x-anthropic-billing-header:"):
continue
anthropic_system_message_content = AnthropicSystemMessageContent(
type="text",
@ -1661,9 +1656,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
text_value = _content.get("text")
if _content.get("type") == "text" and not text_value:
continue
# Skip system messages containing x-anthropic-billing-header metadata
if (
_content.get("type") == "text"
self.should_strip_billing_metadata()
and _content.get("type") == "text"
and text_value
and text_value.startswith("x-anthropic-billing-header:")
):
@ -1978,6 +1973,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
# Remove internal LiteLLM parameters that should not be sent to Anthropic API
optional_params.pop("is_vertex_request", None)
optional_params.pop("client_metadata", None)
data = {
"model": model,

View file

@ -272,19 +272,63 @@ class AnthropicModelInfo(BaseLLMModelInfo):
)
@staticmethod
def _is_adaptive_thinking_model(model: str) -> bool:
"""Claude 4.6+ models use adaptive thinking with ``output_config.effort``."""
def _supports_model_capability(model: str, key: str) -> bool:
"""Check a boolean capability ``key`` in the model map.
Strips bedrock/vertex prefixes so a provider-routed Claude still
resolves to the Anthropic model-map entry.
"""
from litellm.utils import _supports_factory
try:
if _supports_factory(
model=model,
custom_llm_provider=None,
key="supports_adaptive_thinking",
custom_llm_provider="anthropic",
key=key,
):
return True
except Exception:
pass
candidates = [model]
for prefix in (
"bedrock/converse/",
"bedrock/invoke/",
"bedrock/",
"vertex_ai/",
):
if model.startswith(prefix):
candidates.append(model[len(prefix) :])
try:
from litellm.llms.bedrock.common_utils import BedrockModelInfo
base = BedrockModelInfo.get_base_model(model)
if base:
candidates.append(base)
candidates.append(f"bedrock/{base}")
except Exception:
pass
try:
for cand in candidates:
if cand in litellm.model_cost and (
litellm.model_cost[cand].get(key) is True
):
return True
except Exception:
pass
return False
@staticmethod
def _is_adaptive_thinking_model(model: str) -> bool:
"""Claude 4.6+ models use adaptive thinking with ``output_config.effort``.
Driven by the ``supports_adaptive_thinking`` flag in the model map; the
4.6/4.7 name checks remain only as a fallback for provider-routed ids
whose map entries predate the flag.
"""
if AnthropicModelInfo._supports_model_capability(
model, "supports_adaptive_thinking"
):
return True
return AnthropicModelInfo._is_claude_4_6_model(
model
) or AnthropicModelInfo._is_claude_4_7_model(model)

View file

@ -84,6 +84,15 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
if isinstance(content, list):
_process_content_list(content)
def should_strip_billing_metadata(self) -> bool:
"""
Whether to drop x-anthropic-billing-header system blocks before sending upstream.
The first-party Anthropic API uses these blocks for Claude Code attribution, so the
base config keeps them. Providers that reject them override this to True.
"""
return False
@staticmethod
def _filter_billing_headers_from_system(system_param):
"""
@ -286,14 +295,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
optional_params=anthropic_messages_optional_request_params,
)
# Filter out x-anthropic-billing-header from system messages
system_param = anthropic_messages_optional_request_params.get("system")
if system_param is not None:
if self.should_strip_billing_metadata() and system_param is not None:
filtered_system = self._filter_billing_headers_from_system(system_param)
if filtered_system is not None and len(filtered_system) > 0:
anthropic_messages_optional_request_params["system"] = filtered_system
else:
# Remove system parameter if all content was filtered out
anthropic_messages_optional_request_params.pop("system", None)
# Transform context_management from OpenAI format to Anthropic format if needed

View file

@ -97,8 +97,15 @@ class AzureImageEditConfig(OpenAIImageEditConfig):
)
original_url = httpx.URL(api_base)
# Extract api_version or use default
api_version = cast(Optional[str], litellm_params.get("api_version"))
# Resolve api_version: litellm_params > litellm.api_version > AZURE_API_VERSION env > default.
# Mirrors the fallback chain used by the Azure chat path in common_utils.py,
# so callers that set a global / env api_version don't get an unversioned URL.
api_version = (
cast(Optional[str], litellm_params.get("api_version"))
or litellm.api_version
or get_secret_str("AZURE_API_VERSION")
or litellm.AZURE_DEFAULT_API_VERSION
)
# Create a new dictionary with existing params
query_params = dict(original_url.params)

View file

@ -21,6 +21,9 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig):
and Azure endpoint format.
"""
def should_strip_billing_metadata(self) -> bool:
return True
def validate_anthropic_messages_environment(
self,
headers: dict,

View file

@ -40,6 +40,9 @@ class AzureAnthropicConfig(AnthropicConfig):
def custom_llm_provider(self) -> Optional[str]:
return "azure_ai"
def should_strip_billing_metadata(self) -> bool:
return True
def validate_environment(
self,
headers: dict,

View file

@ -62,6 +62,26 @@ class BaseResponsesAPIConfig(ABC):
"""
return False
def sign_request(
self,
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
api_key: Optional[str] = None,
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
) -> Tuple[dict, Optional[bytes]]:
"""Sign the request after the body is finalized.
Default is a no-op (returns headers unchanged, no signed body). Providers
whose endpoint requires request signing (e.g. Bedrock Mantle SigV4)
override this and return the signed body bytes so the handler sends those
exact bytes.
"""
return headers, None
@abstractmethod
def get_supported_openai_params(self, model: str) -> list:
pass

View file

@ -1649,12 +1649,14 @@ class AmazonConverseConfig(BaseConfig):
bedrock_tool_config["toolChoice"] = tool_choice_values
data: CommonRequestObject = {
"additionalModelRequestFields": additional_request_params,
"system": system_content_blocks,
"inferenceConfig": self._transform_inference_params(
inference_params=inference_params
),
}
if additional_request_params:
data["additionalModelRequestFields"] = additional_request_params
if system_content_blocks:
data["system"] = system_content_blocks
# Handle all config blocks
for config_name, config_class in self.get_config_blocks().items():

View file

@ -60,6 +60,9 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
def custom_llm_provider(self) -> Optional[str]:
return "bedrock"
def should_strip_billing_metadata(self) -> bool:
return True
def get_supported_openai_params(self, model: str) -> List[str]:
return AnthropicConfig.get_supported_openai_params(self, model)

View file

@ -17,6 +17,9 @@ class BedrockClaudePlatformConfig(BedrockClaudePlatformMixin, AnthropicConfig):
def custom_llm_provider(self) -> Optional[str]:
return "bedrock"
def should_strip_billing_metadata(self) -> bool:
return True
def validate_environment(
self,
headers: dict,

View file

@ -0,0 +1,171 @@
"""
Amazon Bedrock Mantle - Responses API backend.
gpt-5.5 / gpt-5.4 on Mantle are exposed ONLY on the `/openai/v1/responses`
path (not the standard `/v1/responses`). Payloads and SSE follow the OpenAI
Responses spec, so this config inherits OpenAIResponsesAPIConfig and overrides
only the endpoint URL and authentication.
Auth: Bearer token (BEDROCK_MANTLE_API_KEY or the standard
AWS_BEARER_TOKEN_BEDROCK, or litellm_params.api_key) when present; otherwise
AWS SigV4 (service name "bedrock") using the standard credential chain (IAM
role / access key / profile / web identity), signed via the shared
BaseAWSLLM._sign_request after the request body is finalized.
"""
import re
from typing import Optional, Tuple
from botocore.exceptions import (
CredentialRetrievalError,
NoCredentialsError,
PartialCredentialsError,
ProfileNotFound,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
BEDROCK_MANTLE_DEFAULT_REGION = "us-east-1"
# Checked longest/most-specific first so a full endpoint URL collapses to host
# in one pass and the appended path never doubles.
_BASE_SUFFIXES_TO_STRIP = (
"/openai/v1/responses",
"/v1/responses",
"/responses",
"/openai/v1",
"/v1",
)
# Standard Mantle host: https://bedrock-mantle.<region>.api.aws (group 1 = region).
_MANTLE_HOST_RE = re.compile(
r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE
)
class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
def __init__(self, aws_signer: Optional[BaseAWSLLM] = None):
super().__init__()
self._aws_signer = aws_signer or BaseAWSLLM()
@property
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.BEDROCK_MANTLE
@staticmethod
def _resolve_region(params: dict) -> str:
region = params.get("aws_region_name")
if region:
return region
base = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE")
if base:
match = _MANTLE_HOST_RE.match(base.rstrip("/"))
if match:
return match.group(1)
return (
get_secret_str("BEDROCK_MANTLE_REGION")
or get_secret_str("AWS_REGION_NAME")
or get_secret_str("AWS_REGION")
or BEDROCK_MANTLE_DEFAULT_REGION
)
def get_complete_url(
self,
api_base: Optional[str],
litellm_params: dict,
) -> str:
region = self._resolve_region({**litellm_params, "api_base": api_base})
base = (
api_base
or get_secret_str("BEDROCK_MANTLE_API_BASE")
or f"https://bedrock-mantle.{region}.api.aws"
)
base = base.rstrip("/")
for suffix in _BASE_SUFFIXES_TO_STRIP:
if base.endswith(suffix):
base = base[: -len(suffix)]
break
# For the standard Mantle host (including the default-region base that
# responses/main.py auto-injects into litellm_params.api_base), pin to the
# single resolved region so aws_region_name wins; preserve custom proxy hosts.
if _MANTLE_HOST_RE.match(base):
base = f"https://bedrock-mantle.{region}.api.aws"
return f"{base}/openai/v1/responses"
def validate_environment(
self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]
) -> dict:
litellm_params = litellm_params or GenericLiteLLMParams()
api_key = (
litellm_params.api_key
or get_secret_str("BEDROCK_MANTLE_API_KEY")
or get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
)
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
return headers
def supports_native_file_search(self) -> bool:
return False
def supports_native_websocket(self) -> bool:
return False
def sign_request(
self,
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
api_key: Optional[str] = None,
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
) -> Tuple[dict, Optional[bytes]]:
bearer = (
api_key
or get_secret_str("BEDROCK_MANTLE_API_KEY")
or get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
)
if not bearer:
# SigV4 path. Pin the credential-scope region to the region of the actual
# signing URL (api_base, already region-resolved by get_complete_url) so the
# SigV4 scope and the URL host can never disagree. Resolve from api_base first,
# then fall back to the regular precedence. Also drop any caller Authorization
# so _sign_request's restore-original-Authorization step cannot override the
# SigV4 header.
optional_params = {
**optional_params,
"aws_region_name": self._resolve_region(
{**optional_params, "api_base": api_base}
),
}
headers = {k: v for k, v in headers.items() if k.lower() != "authorization"}
try:
return self._aws_signer._sign_request(
service_name="bedrock",
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base,
api_key=bearer,
model=model,
stream=stream,
fake_stream=fake_stream,
)
except (
NoCredentialsError,
PartialCredentialsError,
ProfileNotFound,
CredentialRetrievalError,
) as e:
raise ValueError(
"Bedrock Mantle auth failed: no Bearer token and no usable AWS "
"credentials. Set BEDROCK_MANTLE_API_KEY (or AWS_BEARER_TOKEN_BEDROCK) "
"or pass api_key for Bearer auth, or provide AWS credentials "
"(IAM role / access key / profile / web identity) for SigV4."
) from e

View file

@ -120,6 +120,7 @@ class CohereV2ChatConfig(OpenAIGPTConfig):
"stream",
"temperature",
"max_tokens",
"max_completion_tokens",
"top_p",
"frequency_penalty",
"presence_penalty",
@ -143,7 +144,12 @@ class CohereV2ChatConfig(OpenAIGPTConfig):
optional_params["stream"] = value
if param == "temperature":
optional_params["temperature"] = value
if param == "max_tokens":
if (
param == "max_tokens"
and "max_completion_tokens" not in non_default_params
):
optional_params["max_tokens"] = value
if param == "max_completion_tokens":
optional_params["max_tokens"] = value
if param == "n":
optional_params["num_generations"] = value

View file

@ -589,6 +589,7 @@ class AsyncHTTPHandler:
params: Optional[dict] = None,
headers: Optional[dict] = None,
follow_redirects: Optional[bool] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
):
# Set follow_redirects to UseClientDefault if None
_follow_redirects = (
@ -599,7 +600,11 @@ class AsyncHTTPHandler:
params.update(HTTPHandler.extract_query_params(url))
response = await self.client.get(
url, params=params, headers=headers, follow_redirects=_follow_redirects # type: ignore
url,
params=params,
headers=headers, # type: ignore
follow_redirects=_follow_redirects, # type: ignore
timeout=timeout if timeout is not None else USE_CLIENT_DEFAULT,
)
return response
@ -1115,6 +1120,7 @@ class HTTPHandler:
params: Optional[dict] = None,
headers: Optional[dict] = None,
follow_redirects: Optional[bool] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
):
# Set follow_redirects to UseClientDefault if None
_follow_redirects = (
@ -1128,6 +1134,7 @@ class HTTPHandler:
params=params,
headers=headers,
follow_redirects=_follow_redirects,
timeout=timeout if timeout is not None else USE_CLIENT_DEFAULT,
)
return response

View file

@ -1751,6 +1751,7 @@ class BaseLLMHTTPHandler:
api_base=api_base,
optional_params=optional_params,
data=data,
api_key=api_key,
)
## LOGGING
@ -1833,6 +1834,7 @@ class BaseLLMHTTPHandler:
api_base=api_base,
optional_params=optional_params,
data=data,
api_key=api_key,
)
## LOGGING
@ -2316,6 +2318,31 @@ class BaseLLMHTTPHandler:
# but never included in the outbound provider payload.
request_context["litellm_params"] = dict(litellm_params)
is_stream_request = bool(stream)
if is_stream_request and fake_stream is True:
stream, data = self._prepare_fake_stream_request(
stream=stream,
data=data,
fake_stream=fake_stream,
)
# Sign after the body is final (post-transform/normalize/extra_body and post
# fake-stream prep) so signed bytes match what we send. No-op for providers
# that inherit the default sign_request.
headers, signed_body = responses_api_provider_config.sign_request(
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
api_base=api_base,
api_key=litellm_params.api_key,
model=model,
stream=stream,
fake_stream=fake_stream,
)
body_kwargs: Dict[str, Any] = (
{"data": signed_body} if signed_body is not None else {"json": data}
)
## LOGGING
logging_obj.pre_call(
input=input,
@ -2328,22 +2355,14 @@ class BaseLLMHTTPHandler:
)
try:
if stream:
# For streaming, use stream=True in the request
if fake_stream is True:
stream, data = self._prepare_fake_stream_request(
stream=stream,
data=data,
fake_stream=fake_stream,
)
if is_stream_request:
response = sync_httpx_client.post(
url=api_base,
headers=headers,
json=data,
timeout=timeout
or float(response_api_optional_request_params.get("timeout", 0)),
stream=stream,
**body_kwargs,
)
if fake_stream is True:
return MockResponsesAPIStreamingIterator(
@ -2368,13 +2387,12 @@ class BaseLLMHTTPHandler:
call_type=CallTypes.responses.value,
)
else:
# For non-streaming requests
response = sync_httpx_client.post(
url=api_base,
headers=headers,
json=data,
timeout=timeout
or float(response_api_optional_request_params.get("timeout", 0)),
**body_kwargs,
)
except Exception as e:
raise self._handle_error(
@ -2462,6 +2480,28 @@ class BaseLLMHTTPHandler:
# but never included in the outbound provider payload.
request_context["litellm_params"] = dict(litellm_params)
is_stream_request = bool(stream)
if is_stream_request and fake_stream is True:
stream, data = self._prepare_fake_stream_request(
stream=stream,
data=data,
fake_stream=fake_stream,
)
headers, signed_body = responses_api_provider_config.sign_request(
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
api_base=api_base,
api_key=litellm_params.api_key,
model=model,
stream=stream,
fake_stream=fake_stream,
)
body_kwargs: Dict[str, Any] = (
{"data": signed_body} if signed_body is not None else {"json": data}
)
## LOGGING
logging_obj.pre_call(
input=input,
@ -2474,22 +2514,14 @@ class BaseLLMHTTPHandler:
)
try:
if stream:
# For streaming, we need to use stream=True in the request
if fake_stream is True:
stream, data = self._prepare_fake_stream_request(
stream=stream,
data=data,
fake_stream=fake_stream,
)
if is_stream_request:
response = await async_httpx_client.post(
url=api_base,
headers=headers,
json=data,
timeout=timeout
or float(response_api_optional_request_params.get("timeout", 0)),
stream=stream,
**body_kwargs,
)
if fake_stream is True:
@ -2516,13 +2548,12 @@ class BaseLLMHTTPHandler:
call_type=CallTypes.responses.value,
)
else:
# For non-streaming, proceed as before
response = await async_httpx_client.post(
url=api_base,
headers=headers,
json=data,
timeout=timeout
or float(response_api_optional_request_params.get("timeout", 0)),
**body_kwargs,
)
except Exception as e:
@ -2586,6 +2617,8 @@ class BaseLLMHTTPHandler:
headers=headers,
)
headers.setdefault("Content-Type", "application/json")
## LOGGING
logging_obj.pre_call(
input=input,
@ -2676,6 +2709,8 @@ class BaseLLMHTTPHandler:
headers=headers,
)
headers.setdefault("Content-Type", "application/json")
## LOGGING
logging_obj.pre_call(
input=input,
@ -3999,6 +4034,18 @@ class BaseLLMHTTPHandler:
)
data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data)
headers, signed_body = responses_api_provider_config.sign_request(
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
api_base=url,
api_key=litellm_params.api_key,
model=model,
)
body_kwargs: Dict[str, Any] = (
{"data": signed_body} if signed_body is not None else {"json": data}
)
## LOGGING
logging_obj.pre_call(
input=input,
@ -4012,7 +4059,7 @@ class BaseLLMHTTPHandler:
try:
response = sync_httpx_client.post(
url=url, headers=headers, json=data, timeout=timeout
url=url, headers=headers, timeout=timeout, **body_kwargs
)
except Exception as e:
@ -4082,6 +4129,18 @@ class BaseLLMHTTPHandler:
)
data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data)
headers, signed_body = responses_api_provider_config.sign_request(
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
api_base=url,
api_key=litellm_params.api_key,
model=model,
)
body_kwargs: Dict[str, Any] = (
{"data": signed_body} if signed_body is not None else {"json": data}
)
## LOGGING
logging_obj.pre_call(
input=input,
@ -4095,7 +4154,7 @@ class BaseLLMHTTPHandler:
try:
response = await async_httpx_client.post(
url=url, headers=headers, json=data, timeout=timeout
url=url, headers=headers, timeout=timeout, **body_kwargs
)
except Exception as e:

View file

@ -26,6 +26,9 @@ class DeepSeekAnthropicMessagesConfig(AnthropicMessagesConfig):
def custom_llm_provider(self) -> Optional[str]:
return "deepseek"
def should_strip_billing_metadata(self) -> bool:
return True
@staticmethod
def get_api_key(api_key: Optional[str] = None) -> Optional[str]:
return api_key or get_secret_str("DEEPSEEK_API_KEY") or litellm.api_key

View file

@ -7,6 +7,7 @@ from .flux_pro_v11_transformation import FalAIFluxProV11Config
from .flux_pro_v11_ultra_transformation import FalAIFluxProV11UltraConfig
from .flux_schnell_transformation import FalAIFluxSchnellConfig
from .imagen4_transformation import FalAIImagen4Config
from .nano_banana_transformation import FalAINanoBananaConfig
from .recraft_v3_transformation import FalAIRecraftV3Config
from .ideogram_v3_transformation import FalAIIdeogramV3Config
from .stable_diffusion_transformation import FalAIStableDiffusionConfig
@ -20,6 +21,7 @@ __all__ = [
"FalAIBaseConfig",
"FalAIImageGenerationConfig",
"FalAIImagen4Config",
"FalAINanoBananaConfig",
"FalAIRecraftV3Config",
"FalAIBriaConfig",
"FalAIFluxProV11Config",
@ -45,7 +47,9 @@ def get_fal_ai_image_generation_config(model: str) -> BaseImageGenerationConfig:
model_lower = model.lower()
# Map model names to their corresponding configuration classes
if "imagen4" in model_lower or "imagen-4" in model_lower:
if "nano-banana" in model_lower or "gemini-25-flash-image" in model_lower:
return FalAINanoBananaConfig()
elif "imagen4" in model_lower or "imagen-4" in model_lower:
return FalAIImagen4Config()
elif "recraft" in model_lower:
return FalAIRecraftV3Config()

View file

@ -0,0 +1,105 @@
from typing import List, Optional
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
from .transformation import FalAIBaseConfig
class FalAINanoBananaConfig(FalAIBaseConfig):
"""
Configuration for Fal AI's Nano Banana / Gemini 2.5 Flash Image models.
Serves the imagen4 deprecation migration path. The same underlying model is
exposed under two endpoints that share an identical schema:
- fal-ai/nano-banana
- fal-ai/gemini-25-flash-image
Documentation: https://fal.ai/models/fal-ai/nano-banana
"""
SUPPORTED_ASPECT_RATIOS: List[str] = [
"21:9",
"16:9",
"3:2",
"4:3",
"5:4",
"1:1",
"4:5",
"3:4",
"2:3",
"9:16",
]
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
base_url: str = (
api_base or get_secret_str("FAL_AI_API_BASE") or self.DEFAULT_BASE_URL
).rstrip("/")
endpoint = model if model.startswith("fal-ai/") else f"fal-ai/{model}"
return f"{base_url}/{endpoint}"
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
return ["n", "response_format", "size"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
supported_params = self.get_supported_openai_params(model)
for key, value in non_default_params.items():
if key == "response_format":
continue
elif key == "n":
if "num_images" not in optional_params:
optional_params["num_images"] = value
elif key == "size":
if "aspect_ratio" not in optional_params:
optional_params["aspect_ratio"] = self._map_aspect_ratio(value)
elif key not in optional_params and not drop_params:
raise ValueError(
f"Parameter {key} is not supported for model {model}. "
f"Supported parameters are {supported_params}. "
"Set drop_params=True to drop unsupported parameters."
)
return optional_params
def _map_aspect_ratio(self, size: str) -> str:
if not isinstance(size, str) or "x" not in size:
return "1:1"
try:
width, height = (int(part) for part in size.split("x"))
target = width / height
except (ValueError, ZeroDivisionError):
return "1:1"
def ratio_of(aspect_ratio: str) -> float:
w, h = (int(part) for part in aspect_ratio.split(":"))
return w / h
return min(
self.SUPPORTED_ASPECT_RATIOS,
key=lambda aspect_ratio: abs(ratio_of(aspect_ratio) - target),
)
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
return {"prompt": prompt, **optional_params}

View file

@ -170,11 +170,6 @@ class FireworksAIConfig(OpenAIGPTConfig):
is_response_format_supported=False,
enforce_tool_choice=False, # tools and response_format are both set, don't enforce tool_choice
)
elif "json_schema" in value:
optional_params["response_format"] = {
"type": "json_object",
"schema": value["json_schema"]["schema"],
}
else:
optional_params["response_format"] = value
elif param == "max_completion_tokens":

View file

@ -265,7 +265,11 @@ class GeminiVideoConfig(BaseVideoConfig):
{
"instances": [
{
"prompt": "A cat playing with a ball of yarn"
"prompt": "A cat playing with a ball of yarn",
"image": {
"bytesBase64Encoded": "...",
"mimeType": "image/jpeg"
}
}
],
"parameters": {
@ -275,13 +279,18 @@ class GeminiVideoConfig(BaseVideoConfig):
}
}
"""
instance = GeminiVideoGenerationInstance(prompt=prompt)
instance: GeminiVideoGenerationInstance = {"prompt": prompt}
params_copy = video_create_optional_request_params.copy()
if "image" in params_copy and params_copy["image"] is not None:
image_data = _convert_image_to_gemini_format(params_copy["image"])
params_copy["image"] = image_data
if "image" in params_copy:
image = params_copy.pop("image")
if image is not None:
if isinstance(image, dict):
image_data = image
else:
image_data = _convert_image_to_gemini_format(image)
instance["image"] = image_data
parameters = GeminiVideoGenerationParameters(**params_copy)

View file

@ -2,7 +2,7 @@
GitHub Copilot Responses API Configuration.
This module provides the configuration for GitHub Copilot's Responses API,
which is required for models like gpt-5.1-codex that only support the /responses endpoint.
which is required for models like gpt-5.3-codex that only support the /responses endpoint.
Implementation based on analysis of the copilot-api project by caozhiyuan:
https://github.com/caozhiyuan/copilot-api
@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional, Union
import os
import litellm
from litellm._logging import verbose_logger
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.exceptions import AuthenticationError
@ -22,6 +23,7 @@ from litellm.types.llms.openai import (
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
from litellm.utils import _cached_get_model_info_helper
from ..authenticator import Authenticator
from ..common_utils import (
@ -38,6 +40,47 @@ else:
LiteLLMLoggingObj = Any
def github_copilot_supports_responses_api(model: str) -> bool:
"""
Gate native /v1/responses dispatch per github_copilot model.
Resolution (first match wins): mode "responses" -> True; mode "chat" ->
False (opt-out wins for dual-endpoint models); "/v1/responses" in
supported_endpoints -> True; else False. Unknown model -> False (the bridge
always works since every Copilot model supports /chat/completions).
Reads merged model info (per-deployment model_info applied via the router's
register_model, which also clears the cache used here).
"""
try:
info = _cached_get_model_info_helper(
model=model, custom_llm_provider="github_copilot"
)
except Exception as e:
verbose_logger.debug(
"github_copilot_supports_responses_api: get_model_info failed "
"for %s: %s",
model,
e,
)
return False
mode = info.get("mode")
if mode == "responses":
return True
if mode == "chat":
return False
# supported_endpoints is dropped by ModelInfoBase; read it from the raw
# model_cost entry via the resolved key.
key = info.get("key")
raw_info = litellm.model_cost.get(key) if isinstance(key, str) else None
endpoints = (
raw_info.get("supported_endpoints") if isinstance(raw_info, dict) else None
)
return isinstance(endpoints, list) and "/v1/responses" in endpoints
class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
"""
Configuration for GitHub Copilot's Responses API.

View file

@ -239,7 +239,7 @@ class HuggingFaceEmbedding(BaseLLM):
model_response.model = model
input_tokens = 0
for text in input:
input_tokens += len(encoding.encode(text))
input_tokens += len(encoding.encode(text, disallowed_special=()))
setattr(
model_response,

View file

@ -17,6 +17,7 @@ from litellm.proxy.common_utils.resource_ownership import (
is_proxy_admin,
user_can_access_resource_owner,
)
from litellm.repositories.table_repositories import SkillsRepository
# Skills are looked up on every chat completion that has skills enabled
# (`SkillsInjectionHook` calls ``fetch_skill_from_db``). 60s LRU/TTL cache
@ -107,7 +108,7 @@ class LiteLLMSkillsHandler:
f"LiteLLMSkillsHandler: Creating skill {skill_id} with title={data.display_title}"
)
new_skill = await prisma_client.db.litellm_skillstable.create(data=skill_data)
new_skill = await SkillsRepository(prisma_client).table.create(data=skill_data)
return _prisma_skill_to_litellm(new_skill)
@staticmethod
@ -133,7 +134,7 @@ class LiteLLMSkillsHandler:
return []
find_many_kwargs["where"] = {"created_by": {"in": owner_scopes}}
skills = await prisma_client.db.litellm_skillstable.find_many(
skills = await SkillsRepository(prisma_client).table.find_many(
**find_many_kwargs
)
return [_prisma_skill_to_litellm(s) for s in skills]
@ -150,7 +151,7 @@ class LiteLLMSkillsHandler:
return cached
prisma_client = await LiteLLMSkillsHandler._get_prisma_client()
skill = await prisma_client.db.litellm_skillstable.find_unique(
skill = await SkillsRepository(prisma_client).table.find_unique(
where={"skill_id": skill_id}
)
_SKILL_CACHE.set_cache(
@ -189,7 +190,7 @@ class LiteLLMSkillsHandler:
):
raise ValueError(f"Skill not found: {skill_id}")
await prisma_client.db.litellm_skillstable.delete(where={"skill_id": skill_id})
await SkillsRepository(prisma_client).table.delete(where={"skill_id": skill_id})
_SKILL_CACHE.set_cache(skill_id, _NEGATIVE_SKILL_SENTINEL)
return {"id": skill_id, "type": "skill_deleted"}

View file

@ -28,6 +28,9 @@ class MinimaxMessagesConfig(AnthropicMessagesConfig):
def custom_llm_provider(self) -> Optional[str]:
return "minimax"
def should_strip_billing_metadata(self) -> bool:
return True
@staticmethod
def get_api_key(api_key: Optional[str] = None) -> Optional[str]:
"""

View file

@ -134,11 +134,15 @@ class MoonshotChatConfig(OpenAIGPTConfig):
##########################################
# temperature limitations
# 1. `temperature` on KIMI API is [0, 1] but OpenAI is [0, 2]
# 2. If temperature < 0.3 and n > 1, KIMI will raise an exception.
# 1. reasoning models (kimi-k2.5, kimi-k2.6, ...) reject every temperature
# except 1, so the param is dropped and the model's default is used
# 2. `temperature` on KIMI API is [0, 1] but OpenAI is [0, 2]
# 3. If temperature < 0.3 and n > 1, KIMI will raise an exception.
# If we enter this condition, we set the temperature to 0.3 as suggested by Moonshot AI
##########################################
if "temperature" in optional_params:
if supports_reasoning(model=model, custom_llm_provider="moonshot"):
optional_params.pop("temperature", None)
elif "temperature" in optional_params:
if optional_params["temperature"] > 1:
optional_params["temperature"] = 1
if optional_params["temperature"] < 0.3 and optional_params.get("n", 1) > 1:

View file

@ -187,6 +187,7 @@ def create_responses_config_class(provider: SimpleProviderConfig):
from litellm.llms.openai_like.responses.transformation import (
OpenAILikeResponsesConfig,
)
from litellm.types.llms.openai import ResponseInputParam
from litellm.types.router import GenericLiteLLMParams
class JSONProviderResponsesConfig(OpenAILikeResponsesConfig):
@ -223,5 +224,23 @@ def create_responses_config_class(provider: SimpleProviderConfig):
api_base = api_base.rstrip("/")
return f"{api_base}/responses"
def transform_responses_api_request(
self,
model: str,
input: Union[str, ResponseInputParam],
response_api_optional_request_params: dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> dict:
if provider.special_handling.get("force_store_false"):
response_api_optional_request_params["store"] = False
return super().transform_responses_api_request(
model=model,
input=input,
response_api_optional_request_params=response_api_optional_request_params,
litellm_params=litellm_params,
headers=headers,
)
_responses_config_cache[provider.slug] = JSONProviderResponsesConfig
return JSONProviderResponsesConfig

View file

@ -115,6 +115,15 @@
"max_completion_tokens": "max_tokens"
}
},
"neosantara": {
"base_url": "https://api.neosantara.xyz/v1",
"api_key_env": "NEOSANTARA_API_KEY",
"api_base_env": "NEOSANTARA_API_BASE",
"param_mappings": {
"max_completion_tokens": "max_tokens"
},
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"]
},
"tensormesh": {
"base_url": "https://serverless.tensormesh.ai/v1",
"api_key_env": "TENSORMESH_INFERENCE_API_KEY",
@ -123,5 +132,14 @@
"param_mappings": {
"max_completion_tokens": "max_tokens"
}
},
"parasail": {
"base_url": "https://api.parasail.io/v1",
"api_key_env": "PARASAIL_API_KEY",
"api_base_env": "PARASAIL_API_BASE",
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
"special_handling": {
"force_store_false": true
}
}
}

View file

@ -25,6 +25,7 @@ class SnowflakeBaseConfig:
"temperature",
"max_tokens",
"top_p",
"stream",
"response_format",
"tools",
"tool_choice",

View file

@ -0,0 +1 @@
"""Soniox LLM provider implementation."""

View file

@ -0,0 +1 @@
"""Soniox audio transcription implementation."""

View file

@ -0,0 +1,802 @@
"""
Handler for Soniox async speech-to-text transcription.
Soniox's async transcription API requires multiple HTTP calls:
1. (optional) POST /v1/files — upload a local audio file
2. POST /v1/transcriptions — create a transcription job
3. GET /v1/transcriptions/{id} — poll until status == "completed"
4. GET /v1/transcriptions/{id}/transcript — fetch the transcript
5. (optional) DELETE /v1/transcriptions/{id} — cleanup
6. (optional) DELETE /v1/files/{id} — cleanup
Because this does not fit the single-request shape of
`base_llm_http_handler.audio_transcriptions`, the dispatch in
`litellm.main.transcription()` routes Soniox requests directly to this
handler (analogous to the OpenAI / Azure transcription handlers).
"""
import asyncio
import math
import time
from typing import (
TYPE_CHECKING,
Any,
Coroutine,
Dict,
List,
Optional,
Tuple,
Union,
)
import httpx
from litellm.litellm_core_utils.audio_utils.utils import (
get_audio_file_name,
process_audio_file,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
_get_httpx_client,
get_async_httpx_client,
)
from litellm.llms.soniox.audio_transcription.transformation import (
SonioxAudioTranscriptionConfig,
)
from litellm.llms.soniox.common_utils import (
SONIOX_DEFAULT_CLEANUP,
SONIOX_DEFAULT_MAX_POLL_ATTEMPTS,
SONIOX_DEFAULT_POLL_INTERVAL,
SONIOX_MAX_POLL_ATTEMPTS,
SONIOX_MAX_POLL_INTERVAL,
SONIOX_MIN_POLL_INTERVAL,
SONIOX_SECRET_FIELDS,
SonioxException,
get_soniox_api_base,
)
from litellm.types.utils import FileTypes, TranscriptionResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import (
Logging as LiteLLMLoggingObj,
)
else:
LiteLLMLoggingObj = Any
class SonioxAudioTranscriptionHandler:
"""Orchestrates the Soniox async transcription flow."""
# ------------------------------------------------------------------
# Public entry points
# ------------------------------------------------------------------
def audio_transcriptions(
self,
model: str,
audio_file: Optional[FileTypes],
optional_params: dict,
litellm_params: dict,
model_response: TranscriptionResponse,
timeout: float,
max_retries: int,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_base: Optional[str],
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
atranscription: bool = False,
headers: Optional[Dict[str, Any]] = None,
provider_config: Optional[SonioxAudioTranscriptionConfig] = None,
) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]:
"""Sync/async dispatch for Soniox transcription requests.
Note: ``max_retries`` is accepted for signature compatibility with
``litellm.transcription`` but is **not yet implemented** for the Soniox
async pipeline. Transient HTTP failures during upload, create, poll,
or fetch will surface immediately. Wrap calls with the standard
``litellm.Router`` / ``num_retries`` mechanism for retry behaviour.
"""
config = provider_config or SonioxAudioTranscriptionConfig()
if atranscription is True:
return self._async_audio_transcriptions(
model=model,
audio_file=audio_file,
optional_params=optional_params,
litellm_params=litellm_params,
model_response=model_response,
timeout=timeout,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
client=client if isinstance(client, AsyncHTTPHandler) else None,
headers=headers or {},
provider_config=config,
)
return self._sync_audio_transcriptions(
model=model,
audio_file=audio_file,
optional_params=optional_params,
litellm_params=litellm_params,
model_response=model_response,
timeout=timeout,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
client=client if isinstance(client, HTTPHandler) else None,
headers=headers or {},
provider_config=config,
)
# ------------------------------------------------------------------
# Helpers shared between sync and async paths
# ------------------------------------------------------------------
def _prepare(
self,
audio_file: Optional[FileTypes],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str],
api_base: Optional[str],
provider_config: SonioxAudioTranscriptionConfig,
headers: Dict[str, Any],
) -> Tuple[
Dict[str, str], # auth headers
str, # api_base (no trailing slash)
Dict[str, Any], # body for POST /v1/transcriptions (without file_id/audio_url)
Dict[str, Any], # handler-only options (poll interval, cleanup, ...)
]:
# Validate env -> auth headers.
auth_headers = provider_config.validate_environment(
headers=headers,
model="", # unused
messages=[],
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
)
base_url = get_soniox_api_base(api_base)
# Operate on a local copy so we don't mutate the caller's dict
# (the caller may reuse `optional_params` for retries or logging).
params = dict(optional_params)
# Pull handler-only kwargs out of params so they aren't sent
# to Soniox.
poll_interval = float(
params.pop("soniox_polling_interval", SONIOX_DEFAULT_POLL_INTERVAL)
)
try:
max_attempts = int(
params.pop(
"soniox_max_polling_attempts", SONIOX_DEFAULT_MAX_POLL_ATTEMPTS
)
)
except (ValueError, OverflowError):
max_attempts = SONIOX_DEFAULT_MAX_POLL_ATTEMPTS
cleanup_raw = params.pop("soniox_cleanup", SONIOX_DEFAULT_CLEANUP)
if cleanup_raw is None:
cleanup: List[str] = []
elif isinstance(cleanup_raw, str):
cleanup = [cleanup_raw]
else:
cleanup = list(cleanup_raw)
filename_override = params.pop("filename", None)
# Server-side clamps. Caller-supplied poll settings (from request kwargs)
# are bounded so an authenticated caller cannot force a worker into a
# tight poll loop (zero interval) or pin it indefinitely (huge attempt
# count). Total polling time is bounded by
# SONIOX_MAX_POLL_ATTEMPTS * SONIOX_MAX_POLL_INTERVAL.
if not math.isfinite(poll_interval):
poll_interval = SONIOX_DEFAULT_POLL_INTERVAL
clamped_poll_interval = max(
SONIOX_MIN_POLL_INTERVAL, min(poll_interval, SONIOX_MAX_POLL_INTERVAL)
)
clamped_max_attempts = max(1, min(max_attempts, SONIOX_MAX_POLL_ATTEMPTS))
handler_opts: Dict[str, Any] = {
"poll_interval": clamped_poll_interval,
"max_attempts": clamped_max_attempts,
"cleanup": cleanup,
"filename_override": filename_override,
"audio_url": params.pop("audio_url", None),
"file_id": params.pop("file_id", None),
}
# Soniox does not accept `language` directly; map_openai_params should
# already have translated it, but drop any leftover to be safe.
params.pop("language", None)
# response_format is handled by LiteLLM post-processing, not Soniox.
handler_opts["response_format"] = params.pop("response_format", None)
return auth_headers, base_url, params, handler_opts
def _build_create_body(
self,
model: str,
optional_params: dict,
handler_opts: Dict[str, Any],
file_id: Optional[str],
) -> Dict[str, Any]:
body: Dict[str, Any] = {"model": model}
# Soniox-native passthrough fields
for key, value in optional_params.items():
if value is None:
continue
body[key] = value
if handler_opts.get("audio_url"):
body["audio_url"] = handler_opts["audio_url"]
if file_id:
body["file_id"] = file_id
return body
@staticmethod
def _redact_body_for_logging(body: Dict[str, Any]) -> Dict[str, Any]:
"""Return a shallow copy of ``body`` with secret fields redacted.
Soniox's create-transcription body can include
``webhook_auth_header_value`` (a shared secret used to authenticate
webhook callbacks). Forwarding that value to logging callbacks would
let anyone with read access to those sinks forge webhook requests, so
we replace any value of a known secret-bearing field with the literal
``"[REDACTED]"`` before logging. Non-secret fields are passed through
unchanged.
"""
if not body:
return body
redacted = dict(body)
for field in SONIOX_SECRET_FIELDS:
if field in redacted and redacted[field] is not None:
redacted[field] = "[REDACTED]"
return redacted
@staticmethod
def _safe_log_pre_call(
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_base: str,
body: Dict[str, Any],
) -> None:
try:
logging_obj.pre_call(
input=None,
api_key=api_key,
additional_args={
"api_base": f"{api_base}/v1/transcriptions",
"atranscription": True,
"complete_input_dict": SonioxAudioTranscriptionHandler._redact_body_for_logging(
body
),
},
)
except Exception:
# Logging hooks are best-effort: a misbehaving callback or third-party
# observability integration must never break a real Soniox call.
pass
@staticmethod
def _safe_log_post_call(
logging_obj: LiteLLMLoggingObj,
audio_file: Optional[FileTypes],
api_key: Optional[str],
body: Dict[str, Any],
original_response: Any,
) -> None:
try:
logging_obj.post_call(
input=get_audio_file_name(audio_file) if audio_file else None,
api_key=api_key,
additional_args={
"complete_input_dict": SonioxAudioTranscriptionHandler._redact_body_for_logging(
body
)
},
original_response=original_response,
)
except Exception:
# Logging hooks are best-effort: a misbehaving callback or third-party
# observability integration must never break a real Soniox call.
pass
@staticmethod
def _raise_for_response(
response: httpx.Response,
provider_config: SonioxAudioTranscriptionConfig,
action: str,
) -> None:
if response.status_code >= 400:
try:
payload = response.json()
message = (
payload.get("error_message")
or payload.get("error")
or response.text
)
except Exception:
message = response.text
raise provider_config.get_error_class(
error_message=f"Soniox {action} failed (HTTP {response.status_code}): {message}",
status_code=response.status_code,
headers=response.headers,
)
# ------------------------------------------------------------------
# Sync flow
# ------------------------------------------------------------------
def _sync_audio_transcriptions(
self,
model: str,
audio_file: Optional[FileTypes],
optional_params: dict,
litellm_params: dict,
model_response: TranscriptionResponse,
timeout: float,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_base: Optional[str],
client: Optional[HTTPHandler],
headers: Dict[str, Any],
provider_config: SonioxAudioTranscriptionConfig,
) -> TranscriptionResponse:
auth_headers, base_url, opt_params, handler_opts = self._prepare(
audio_file=audio_file,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
provider_config=provider_config,
headers=headers,
)
http_client = (
client
if isinstance(client, HTTPHandler)
else (
_get_httpx_client(
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
)
)
)
file_id = handler_opts.get("file_id")
uploaded_file_id: Optional[str] = None
transcription_id: Optional[str] = None
try:
if not file_id and not handler_opts.get("audio_url"):
if audio_file is None:
raise SonioxException(
message=(
"Soniox transcription requires one of: a file argument, "
"an `audio_url` kwarg, or a `file_id` kwarg."
),
status_code=400,
headers=None,
)
uploaded_file_id = self._sync_upload_file(
http_client=http_client,
base_url=base_url,
auth_headers=auth_headers,
audio_file=audio_file,
filename_override=handler_opts.get("filename_override"),
timeout=timeout,
provider_config=provider_config,
)
file_id = uploaded_file_id
body = self._build_create_body(model, opt_params, handler_opts, file_id)
self._safe_log_pre_call(logging_obj, api_key, base_url, body)
create_resp = http_client.post(
url=f"{base_url}/v1/transcriptions",
headers=auth_headers,
json=body,
timeout=timeout,
)
self._raise_for_response(
create_resp, provider_config, "create transcription"
)
transcription_id = create_resp.json()["id"]
transcription_meta = self._sync_poll_until_completed(
http_client=http_client,
base_url=base_url,
auth_headers=auth_headers,
transcription_id=transcription_id,
poll_interval=handler_opts["poll_interval"],
max_attempts=handler_opts["max_attempts"],
timeout=timeout,
provider_config=provider_config,
)
transcript_resp = http_client.get(
url=f"{base_url}/v1/transcriptions/{transcription_id}/transcript",
headers=auth_headers,
timeout=timeout,
)
self._raise_for_response(
transcript_resp, provider_config, "fetch transcript"
)
transcript = transcript_resp.json()
payload = {"transcription": transcription_meta, "transcript": transcript}
response = provider_config._build_response_from_payload(
payload,
model_response=model_response,
response_format=handler_opts.get("response_format"),
)
self._safe_log_post_call(logging_obj, audio_file, api_key, body, payload)
audio_duration_ms = transcription_meta.get("audio_duration_ms")
response._hidden_params.update(
{
"model": model,
"custom_llm_provider": "soniox",
"audio_transcription_duration": (
float(audio_duration_ms) / 1000.0
if audio_duration_ms is not None
else None
),
}
)
return response
finally:
self._sync_cleanup(
http_client=http_client,
base_url=base_url,
auth_headers=auth_headers,
cleanup=handler_opts["cleanup"],
file_id_to_cleanup=uploaded_file_id,
transcription_id=transcription_id,
timeout=timeout,
)
def _sync_upload_file(
self,
http_client: HTTPHandler,
base_url: str,
auth_headers: Dict[str, str],
audio_file: FileTypes,
filename_override: Optional[str],
timeout: float,
provider_config: SonioxAudioTranscriptionConfig,
) -> str:
processed = process_audio_file(audio_file)
filename = filename_override or processed.filename
files = {
"file": (filename, processed.file_content, processed.content_type),
}
# `Authorization` header is fine; httpx sets multipart Content-Type.
upload_headers = {"Authorization": auth_headers["Authorization"]}
resp = http_client.post(
url=f"{base_url}/v1/files",
headers=upload_headers,
files=files,
timeout=timeout,
)
self._raise_for_response(resp, provider_config, "upload file")
return resp.json()["id"]
def _sync_poll_until_completed(
self,
http_client: HTTPHandler,
base_url: str,
auth_headers: Dict[str, str],
transcription_id: str,
poll_interval: float,
max_attempts: int,
timeout: float,
provider_config: SonioxAudioTranscriptionConfig,
) -> Dict[str, Any]:
for _ in range(max_attempts):
resp = http_client.get(
url=f"{base_url}/v1/transcriptions/{transcription_id}",
headers=auth_headers,
timeout=timeout,
)
self._raise_for_response(resp, provider_config, "poll transcription")
data = resp.json()
status = data.get("status")
if status == "completed":
return data
if status == "error":
raise provider_config.get_error_class(
error_message=(
f"Soniox transcription {transcription_id} failed: "
f"{data.get('error_message') or data.get('error_type') or 'unknown error'}"
),
status_code=500,
headers=resp.headers,
)
time.sleep(poll_interval)
raise provider_config.get_error_class(
error_message=(
f"Soniox transcription {transcription_id} did not complete after "
f"{max_attempts} polling attempts (interval={poll_interval}s)."
),
status_code=504,
headers={},
)
def _sync_cleanup(
self,
http_client: HTTPHandler,
base_url: str,
auth_headers: Dict[str, str],
cleanup: List[str],
file_id_to_cleanup: Optional[str],
transcription_id: Optional[str],
timeout: float,
) -> None:
if not cleanup:
return
if "transcription" in cleanup and transcription_id:
try:
http_client.delete(
url=f"{base_url}/v1/transcriptions/{transcription_id}",
headers=auth_headers,
timeout=timeout,
)
except Exception:
# Cleanup is best-effort: a failed delete leaves stale data on
# Soniox but must not mask the original transcription result
# (or, on the error path, the original error).
pass
if "file" in cleanup and file_id_to_cleanup:
try:
http_client.delete(
url=f"{base_url}/v1/files/{file_id_to_cleanup}",
headers=auth_headers,
timeout=timeout,
)
except Exception:
# Cleanup is best-effort; see comment above.
pass
# ------------------------------------------------------------------
# Async flow
# ------------------------------------------------------------------
async def _async_audio_transcriptions(
self,
model: str,
audio_file: Optional[FileTypes],
optional_params: dict,
litellm_params: dict,
model_response: TranscriptionResponse,
timeout: float,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_base: Optional[str],
client: Optional[AsyncHTTPHandler],
headers: Dict[str, Any],
provider_config: SonioxAudioTranscriptionConfig,
) -> TranscriptionResponse:
import litellm
auth_headers, base_url, opt_params, handler_opts = self._prepare(
audio_file=audio_file,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
provider_config=provider_config,
headers=headers,
)
http_client = (
client
if isinstance(client, AsyncHTTPHandler)
else (
get_async_httpx_client(
llm_provider=litellm.LlmProviders.SONIOX,
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
)
)
)
file_id = handler_opts.get("file_id")
uploaded_file_id: Optional[str] = None
transcription_id: Optional[str] = None
try:
if not file_id and not handler_opts.get("audio_url"):
if audio_file is None:
raise SonioxException(
message=(
"Soniox transcription requires one of: a file argument, "
"an `audio_url` kwarg, or a `file_id` kwarg."
),
status_code=400,
headers=None,
)
uploaded_file_id = await self._async_upload_file(
http_client=http_client,
base_url=base_url,
auth_headers=auth_headers,
audio_file=audio_file,
filename_override=handler_opts.get("filename_override"),
timeout=timeout,
provider_config=provider_config,
)
file_id = uploaded_file_id
body = self._build_create_body(model, opt_params, handler_opts, file_id)
self._safe_log_pre_call(logging_obj, api_key, base_url, body)
create_resp = await http_client.post(
url=f"{base_url}/v1/transcriptions",
headers=auth_headers,
json=body,
timeout=timeout,
)
self._raise_for_response(
create_resp, provider_config, "create transcription"
)
transcription_id = create_resp.json()["id"]
transcription_meta = await self._async_poll_until_completed(
http_client=http_client,
base_url=base_url,
auth_headers=auth_headers,
transcription_id=transcription_id,
poll_interval=handler_opts["poll_interval"],
max_attempts=handler_opts["max_attempts"],
timeout=timeout,
provider_config=provider_config,
)
transcript_resp = await http_client.get(
url=f"{base_url}/v1/transcriptions/{transcription_id}/transcript",
headers=auth_headers,
timeout=timeout,
)
self._raise_for_response(
transcript_resp, provider_config, "fetch transcript"
)
transcript = transcript_resp.json()
payload = {"transcription": transcription_meta, "transcript": transcript}
response = provider_config._build_response_from_payload(
payload,
model_response=model_response,
response_format=handler_opts.get("response_format"),
)
self._safe_log_post_call(logging_obj, audio_file, api_key, body, payload)
audio_duration_ms = transcription_meta.get("audio_duration_ms")
response._hidden_params.update(
{
"model": model,
"custom_llm_provider": "soniox",
"audio_transcription_duration": (
float(audio_duration_ms) / 1000.0
if audio_duration_ms is not None
else None
),
}
)
return response
finally:
await self._async_cleanup(
http_client=http_client,
base_url=base_url,
auth_headers=auth_headers,
cleanup=handler_opts["cleanup"],
file_id_to_cleanup=uploaded_file_id,
transcription_id=transcription_id,
timeout=timeout,
)
async def _async_upload_file(
self,
http_client: AsyncHTTPHandler,
base_url: str,
auth_headers: Dict[str, str],
audio_file: FileTypes,
filename_override: Optional[str],
timeout: float,
provider_config: SonioxAudioTranscriptionConfig,
) -> str:
processed = process_audio_file(audio_file)
filename = filename_override or processed.filename
files = {
"file": (filename, processed.file_content, processed.content_type),
}
upload_headers = {"Authorization": auth_headers["Authorization"]}
resp = await http_client.post(
url=f"{base_url}/v1/files",
headers=upload_headers,
files=files,
timeout=timeout,
)
self._raise_for_response(resp, provider_config, "upload file")
return resp.json()["id"]
async def _async_poll_until_completed(
self,
http_client: AsyncHTTPHandler,
base_url: str,
auth_headers: Dict[str, str],
transcription_id: str,
poll_interval: float,
max_attempts: int,
timeout: float,
provider_config: SonioxAudioTranscriptionConfig,
) -> Dict[str, Any]:
for _ in range(max_attempts):
resp = await http_client.get(
url=f"{base_url}/v1/transcriptions/{transcription_id}",
headers=auth_headers,
timeout=timeout,
)
self._raise_for_response(resp, provider_config, "poll transcription")
data = resp.json()
status = data.get("status")
if status == "completed":
return data
if status == "error":
raise provider_config.get_error_class(
error_message=(
f"Soniox transcription {transcription_id} failed: "
f"{data.get('error_message') or data.get('error_type') or 'unknown error'}"
),
status_code=500,
headers=resp.headers,
)
await asyncio.sleep(poll_interval)
raise provider_config.get_error_class(
error_message=(
f"Soniox transcription {transcription_id} did not complete after "
f"{max_attempts} polling attempts (interval={poll_interval}s)."
),
status_code=504,
headers={},
)
async def _async_cleanup(
self,
http_client: AsyncHTTPHandler,
base_url: str,
auth_headers: Dict[str, str],
cleanup: List[str],
file_id_to_cleanup: Optional[str],
transcription_id: Optional[str],
timeout: float,
) -> None:
if not cleanup:
return
if "transcription" in cleanup and transcription_id:
try:
await http_client.delete(
url=f"{base_url}/v1/transcriptions/{transcription_id}",
headers=auth_headers,
timeout=timeout,
)
except Exception:
# Cleanup is best-effort: a failed delete leaves stale data on
# Soniox but must not mask the original transcription result
# (or, on the error path, the original error).
pass
if "file" in cleanup and file_id_to_cleanup:
try:
await http_client.delete(
url=f"{base_url}/v1/files/{file_id_to_cleanup}",
headers=auth_headers,
timeout=timeout,
)
except Exception:
# Cleanup is best-effort; see comment above.
pass

View file

@ -0,0 +1,281 @@
"""
Translates between OpenAI's `/v1/audio/transcriptions` shape and Soniox's
async transcription API (https://soniox.com/docs/stt/async/async-transcription).
This config covers parameter mapping, env validation and response shaping.
The actual orchestration (file upload -> create -> poll -> fetch -> cleanup)
lives in `litellm.llms.soniox.audio_transcription.handler`, because Soniox's
async API requires multiple HTTP calls and does not fit the single-request
contract of `base_llm_http_handler.audio_transcriptions`.
"""
from typing import Any, Dict, List, Optional, Union
from httpx import Headers, Response
from litellm.llms.base_llm.audio_transcription.transformation import (
AudioTranscriptionRequestData,
BaseAudioTranscriptionConfig,
)
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.soniox.common_utils import (
SonioxException,
get_soniox_api_base,
get_soniox_api_key,
render_soniox_tokens,
render_soniox_tokens_as_srt,
render_soniox_tokens_as_vtt,
)
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIAudioTranscriptionOptionalParams,
)
from litellm.types.utils import FileTypes, TranscriptionResponse
# Soniox-native kwargs the user can pass through `litellm.transcription(..., **kwargs)`
# in addition to the standard OpenAI params.
SONIOX_PASSTHROUGH_PARAMS: List[str] = [
"language_hints",
"language_hints_strict",
"enable_language_identification",
"enable_speaker_diarization",
"context",
"translation",
"client_reference_id",
"webhook_url",
"webhook_auth_header_name",
"webhook_auth_header_value",
"audio_url",
"file_id",
]
# Handler-only kwargs (consumed by the handler, not sent to Soniox).
SONIOX_HANDLER_ONLY_PARAMS: List[str] = [
"soniox_polling_interval",
"soniox_max_polling_attempts",
"soniox_cleanup",
"filename",
]
class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
"""Configuration for Soniox async speech-to-text transcription."""
def get_supported_openai_params(
self, model: str
) -> List[OpenAIAudioTranscriptionOptionalParams]:
# `language` is mapped onto Soniox's `language_hints`.
# `response_format` is handled by LiteLLM (Soniox doesn't support
# SRT/VTT natively but we synthesize them from token timestamps).
return ["language", "response_format"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
# Translate the OpenAI `language` param into Soniox `language_hints`.
if "language" in non_default_params and non_default_params["language"]:
language = non_default_params["language"]
existing_hints = optional_params.get("language_hints")
if not existing_hints:
optional_params["language_hints"] = [language]
elif language not in existing_hints:
optional_params["language_hints"] = [language] + list(existing_hints)
# Capture response_format for post-processing (not sent to Soniox API).
if "response_format" in non_default_params:
optional_params["response_format"] = non_default_params["response_format"]
# Pass through Soniox-native kwargs unchanged.
for key in SONIOX_PASSTHROUGH_PARAMS + SONIOX_HANDLER_ONLY_PARAMS:
if key in non_default_params and non_default_params[key] is not None:
optional_params[key] = non_default_params[key]
return optional_params
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, Headers]
) -> BaseLLMException:
return SonioxException(
message=error_message, status_code=status_code, headers=headers
)
def validate_environment(
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
resolved_key = get_soniox_api_key(api_key)
if not resolved_key:
raise SonioxException(
message=(
"Missing Soniox API key. Set the SONIOX_API_KEY environment "
"variable or pass api_key=... to litellm.transcription()."
),
status_code=401,
headers=None,
)
merged_headers: Dict[str, str] = {
"Authorization": f"Bearer {resolved_key}",
}
if headers:
merged_headers.update(headers)
return merged_headers
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
# The handler builds per-call URLs (uploads, create, poll, fetch, delete);
# we just return the resolved base.
return get_soniox_api_base(api_base)
def transform_audio_transcription_request(
self,
model: str,
audio_file: FileTypes,
optional_params: dict,
litellm_params: dict,
) -> AudioTranscriptionRequestData:
"""
Build the JSON body for `POST /v1/transcriptions`.
The handler is responsible for the file upload (if `audio_file` is bytes)
and for filling in `file_id`/`audio_url`. This method exists so the
config can be exercised in isolation by unit tests.
"""
body: Dict[str, Any] = {"model": model}
for key in SONIOX_PASSTHROUGH_PARAMS:
value = optional_params.get(key)
if value is not None:
body[key] = value
return AudioTranscriptionRequestData(
data=body, files=None, content_type="application/json"
)
def transform_audio_transcription_response(
self,
raw_response: Response,
model_response: Optional[TranscriptionResponse] = None,
) -> TranscriptionResponse:
"""
Build a TranscriptionResponse from a Soniox transcript payload.
`raw_response.json()` may be either:
- a Soniox transcript object: `{"id": "...", "text": "...", "tokens": [...]}`
- or a merged envelope: `{"transcription": {...}, "transcript": {...}}`
produced by the handler so transcription metadata is also available.
"""
try:
payload = raw_response.json()
except Exception as exc:
raise SonioxException(
message=f"Failed to parse Soniox response: {exc}",
status_code=getattr(raw_response, "status_code", 500),
headers=getattr(raw_response, "headers", None),
)
return self._build_response_from_payload(payload, model_response=model_response)
def _build_response_from_payload(
self,
payload: Dict[str, Any],
model_response: Optional[TranscriptionResponse] = None,
response_format: Optional[str] = None,
) -> TranscriptionResponse:
"""Shared response-building logic (also used by the handler)."""
transcription_meta: Dict[str, Any] = {}
transcript: Dict[str, Any]
if isinstance(payload, dict) and "transcript" in payload:
transcription_meta = payload.get("transcription") or {}
transcript = payload.get("transcript") or {}
else:
transcript = payload if isinstance(payload, dict) else {}
tokens: List[Dict[str, Any]] = transcript.get("tokens") or []
# Decide what to put in `text` based on response_format:
# - "srt": render tokens as SRT subtitles (synthesized from timestamps)
# - "vtt": render tokens as WebVTT subtitles (synthesized from timestamps)
# - "verbose_json": return JSON with word-level timing (handled below)
# - "text" / "json" / None: default plain text rendering
if response_format == "srt" and tokens:
text = render_soniox_tokens_as_srt(tokens)
elif response_format == "vtt" and tokens:
text = render_soniox_tokens_as_vtt(tokens)
else:
# Default text rendering (also used for "json", "text",
# "verbose_json")
has_speaker = any(t.get("speaker") is not None for t in tokens)
has_language = any(t.get("language") is not None for t in tokens)
if (has_speaker or has_language) and tokens:
text = render_soniox_tokens(tokens)
elif transcript.get("text"):
text = transcript["text"]
elif tokens:
text = "".join(t.get("text", "") for t in tokens)
else:
text = ""
response = model_response or TranscriptionResponse(text=text)
response.text = text
response["task"] = "transcribe"
# Best-effort metadata fields matching OpenAI's verbose_json shape.
if transcription_meta.get("audio_duration_ms") is not None:
try:
response["duration"] = (
float(transcription_meta["audio_duration_ms"]) / 1000.0
)
except (TypeError, ValueError):
pass
# Surface a representative language if all tokens agree.
has_language = any(t.get("language") is not None for t in tokens)
if has_language:
languages = {t.get("language") for t in tokens if t.get("language")}
if len(languages) == 1:
response["language"] = next(iter(languages))
# For verbose_json, include word-level timing from tokens.
if response_format == "verbose_json" and tokens:
words: List[Dict[str, Any]] = []
for token in tokens:
word_entry: Dict[str, Any] = {"word": token.get("text", "")}
if token.get("start_ms") is not None:
word_entry["start"] = float(token["start_ms"]) / 1000.0
if token.get("end_ms") is not None:
word_entry["end"] = float(token["end_ms"]) / 1000.0
words.append(word_entry)
if words:
response["words"] = words
# Stash the raw Soniox payload so power-users can read tokens, segments,
# speaker/language data, etc.
response._hidden_params.update(
{
"soniox_raw": {
"transcription": transcription_meta,
"transcript": transcript,
}
}
)
return response

View file

@ -0,0 +1,274 @@
"""
Shared utilities for the Soniox provider (https://soniox.com).
"""
from typing import Any, Dict, List, Optional
from litellm.llms.base_llm.chat.transformation import BaseLLMException
# Soniox API base URL.
SONIOX_API_BASE: str = "https://api.soniox.com"
# Default polling interval in seconds when waiting for an async transcription
# to finish. Mirrors the Soniox SDK default.
SONIOX_DEFAULT_POLL_INTERVAL: float = 1.0
# Minimum polling interval (in seconds) the server will accept from caller-
# supplied `soniox_polling_interval` kwargs. Prevents an authenticated caller
# from forcing a worker into a tight poll loop with a zero/near-zero interval.
SONIOX_MIN_POLL_INTERVAL: float = 0.5
# Maximum polling interval (in seconds). Prevents a caller from setting an
# excessively large or non-finite interval that would keep a worker sleeping
# far longer than necessary between status checks.
SONIOX_MAX_POLL_INTERVAL: float = 60.0
# Default maximum number of polling attempts (1800 attempts * 1s ~= 30 minutes).
SONIOX_DEFAULT_MAX_POLL_ATTEMPTS: int = 1800
# Hard upper bound on polling attempts. Combined with `SONIOX_MIN_POLL_INTERVAL`
# this caps total polling time per request at ~3000s (50 minutes), preventing a
# caller from pinning a worker indefinitely via a huge attempt count.
SONIOX_MAX_POLL_ATTEMPTS: int = 6000
# Default cleanup behaviour: delete both the uploaded file (if any) and the
# transcription record after the transcript has been fetched.
SONIOX_DEFAULT_CLEANUP: List[str] = ["file", "transcription"]
# Body fields that may carry secrets and must be redacted before being
# forwarded to logging callbacks. Soniox accepts a webhook auth header value
# alongside the create-transcription request; that value lets the recipient
# authenticate webhook callbacks and must not leak into observability sinks.
SONIOX_SECRET_FIELDS: List[str] = ["webhook_auth_header_value"]
class SonioxException(BaseLLMException):
"""Provider-specific exception class for Soniox."""
pass
def get_soniox_api_key(api_key: Optional[str] = None) -> Optional[str]:
"""Resolve the Soniox API key from arg or env var."""
# Local import to avoid a circular import: litellm.secret_managers.main
# imports from litellm at top-level.
from litellm.secret_managers.main import get_secret_str
return api_key or get_secret_str("SONIOX_API_KEY")
def get_soniox_api_base(api_base: Optional[str] = None) -> str:
"""Resolve the Soniox API base URL from arg or env var (defaults to public API)."""
from litellm.secret_managers.main import get_secret_str
base = api_base or get_secret_str("SONIOX_API_BASE") or SONIOX_API_BASE
return base.rstrip("/")
def render_soniox_tokens(tokens: List[Dict[str, Any]]) -> str:
"""
Render a list of Soniox tokens to a readable transcript string.
Mirrors the behaviour of the official Soniox SDK's `renderTokens` helper:
- When the speaker changes, a `Speaker N:` tag is inserted.
- When the language changes, a `[lang]` (or `[Translation][lang]`) tag is
inserted.
If neither speaker nor language information is present on any token (i.e.
diarization and language identification are disabled), the function simply
concatenates the token texts.
"""
if not tokens:
return ""
text_parts: List[str] = []
current_speaker: Optional[Any] = None
current_language: Optional[Any] = None
for token in tokens:
text = token.get("text", "")
speaker = token.get("speaker")
language = token.get("language")
is_translation = token.get("translation_status") == "translation"
# Speaker changed -> emit a speaker tag.
if speaker is not None and speaker != current_speaker:
if current_speaker is not None:
text_parts.append("\n\n")
current_speaker = speaker
current_language = None # reset language whenever speaker changes
text_parts.append(f"Speaker {current_speaker}:")
# Language changed -> emit a language (or translation) tag.
if language is not None and language != current_language:
current_language = language
prefix = "[Translation] " if is_translation else ""
text_parts.append(f"\n{prefix}[{current_language}] ")
text = text.lstrip() if isinstance(text, str) else text
text_parts.append(text)
return "".join(text_parts)
# ---------------------------------------------------------------------------
# SRT / VTT subtitle rendering
# ---------------------------------------------------------------------------
# Maximum number of tokens to group into a single subtitle cue.
_CUE_MAX_TOKENS: int = 15
# Maximum duration (in ms) for a single cue before forcing a break.
_CUE_MAX_DURATION_MS: int = 5000
def _format_timestamp_srt(ms: int) -> str:
"""Format milliseconds as SRT timestamp: HH:MM:SS,mmm"""
if ms < 0:
ms = 0
hours = ms // 3_600_000
ms %= 3_600_000
minutes = ms // 60_000
ms %= 60_000
seconds = ms // 1_000
millis = ms % 1_000
return f"{hours:02d}:{minutes:02d}:{seconds:02d},{millis:03d}"
def _format_timestamp_vtt(ms: int) -> str:
"""Format milliseconds as VTT timestamp: HH:MM:SS.mmm"""
if ms < 0:
ms = 0
hours = ms // 3_600_000
ms %= 3_600_000
minutes = ms // 60_000
ms %= 60_000
seconds = ms // 1_000
millis = ms % 1_000
return f"{hours:02d}:{minutes:02d}:{seconds:02d}.{millis:03d}"
def _group_tokens_into_cues(
tokens: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
"""
Group Soniox tokens into subtitle cues.
Each cue has:
- start_ms: int
- end_ms: int
- text: str
Grouping heuristics:
- A new cue starts when token count exceeds _CUE_MAX_TOKENS.
- A new cue starts when duration exceeds _CUE_MAX_DURATION_MS.
- A new cue starts when the speaker changes (if diarization is on).
- Tokens without timestamps are appended to the current cue.
"""
cues: List[Dict[str, Any]] = []
current_tokens: List[str] = []
current_start: Optional[int] = None
current_end: Optional[int] = None
current_speaker: Optional[Any] = None
def _flush() -> None:
if current_tokens and current_start is not None:
text = "".join(current_tokens).strip()
if text:
cues.append(
{
"start_ms": current_start,
"end_ms": (
current_end if current_end is not None else current_start
),
"text": text,
}
)
for token in tokens:
start_ms = token.get("start_ms")
end_ms = token.get("end_ms")
text = token.get("text", "")
speaker = token.get("speaker")
# Skip tokens with no timestamp data entirely if we have no cue started
if start_ms is None and current_start is None:
continue
# Speaker change forces a new cue
if speaker is not None and speaker != current_speaker:
_flush()
current_tokens = []
current_start = start_ms
current_end = end_ms
current_speaker = speaker
current_tokens.append(text)
continue
# Duration or token count exceeded -> flush
should_break = False
if len(current_tokens) >= _CUE_MAX_TOKENS:
should_break = True
elif (
current_start is not None
and start_ms is not None
and (start_ms - current_start) >= _CUE_MAX_DURATION_MS
):
should_break = True
if should_break:
_flush()
current_tokens = []
current_start = start_ms
current_end = end_ms
current_tokens.append(text)
else:
if current_start is None:
current_start = start_ms
if end_ms is not None:
current_end = end_ms
current_tokens.append(text)
_flush()
return cues
def render_soniox_tokens_as_srt(tokens: List[Dict[str, Any]]) -> str:
"""
Render Soniox tokens as SRT (SubRip) subtitle format.
Returns an empty string if no tokens have timestamp data.
"""
cues = _group_tokens_into_cues(tokens)
if not cues:
return ""
lines: List[str] = []
for idx, cue in enumerate(cues, start=1):
start = _format_timestamp_srt(cue["start_ms"])
end = _format_timestamp_srt(cue["end_ms"])
lines.append(str(idx))
lines.append(f"{start} --> {end}")
lines.append(cue["text"])
lines.append("") # blank line between cues
return "\n".join(lines)
def render_soniox_tokens_as_vtt(tokens: List[Dict[str, Any]]) -> str:
"""
Render Soniox tokens as WebVTT subtitle format.
Returns the VTT header even if no cues are present.
"""
cues = _group_tokens_into_cues(tokens)
lines: List[str] = ["WEBVTT", ""]
for cue in cues:
start = _format_timestamp_vtt(cue["start_ms"])
end = _format_timestamp_vtt(cue["end_ms"])
lines.append(f"{start} --> {end}")
lines.append(cue["text"])
lines.append("") # blank line between cues
return "\n".join(lines)

View file

@ -337,6 +337,7 @@ class ContextCachingEndpoints(VertexBase):
return messages, optional_params, None
tools = optional_params.pop("tools", None)
tool_choice = optional_params.pop("tool_choice", None)
## AUTHORIZATION ##
token, url = self._get_token_and_url_context_caching(
@ -371,7 +372,7 @@ class ContextCachingEndpoints(VertexBase):
## CHECK IF CACHED ALREADY
generated_cache_key = local_cache_obj.get_cache_key(
messages=cached_messages, tools=tools, model=model
messages=cached_messages, tools=tools, tool_choice=tool_choice, model=model
)
google_cache_name = self.check_cache(
cache_key=generated_cache_key,
@ -402,6 +403,8 @@ class ContextCachingEndpoints(VertexBase):
)
cached_content_request_body["tools"] = tools
if tool_choice is not None:
cached_content_request_body["toolConfig"] = tool_choice
## LOGGING
logging_obj.pre_call(
@ -487,6 +490,7 @@ class ContextCachingEndpoints(VertexBase):
return messages, optional_params, None
tools = optional_params.pop("tools", None)
tool_choice = optional_params.pop("tool_choice", None)
## AUTHORIZATION ##
token, url = self._get_token_and_url_context_caching(
@ -518,7 +522,7 @@ class ContextCachingEndpoints(VertexBase):
## CHECK IF CACHED ALREADY
generated_cache_key = local_cache_obj.get_cache_key(
messages=cached_messages, tools=tools, model=model
messages=cached_messages, tools=tools, tool_choice=tool_choice, model=model
)
google_cache_name = await self.async_check_cache(
cache_key=generated_cache_key,
@ -550,6 +554,8 @@ class ContextCachingEndpoints(VertexBase):
)
cached_content_request_body["tools"] = tools
if tool_choice is not None:
cached_content_request_body["toolConfig"] = tool_choice
## LOGGING
logging_obj.pre_call(

View file

@ -1111,6 +1111,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
{
"voice": "alloy",
"format": "mp3",
"language_code": "en-US",
}
Expected output:
@ -1119,7 +1120,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
prebuiltVoiceConfig: {
voiceName: "alloy",
}
}
},
languageCode: "en-US",
}
"""
from litellm.types.llms.vertex_ai import (
@ -1145,6 +1147,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
voice_config: VoiceConfig = {"prebuiltVoiceConfig": prebuilt_voice_config}
speech_config["voiceConfig"] = voice_config
if "language_code" in value:
speech_config["languageCode"] = value["language_code"]
return cast(dict, speech_config)
@staticmethod

View file

@ -17,6 +17,9 @@ from ..output_params_utils import sanitize_vertex_anthropic_output_params
class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, VertexBase):
def should_strip_billing_metadata(self) -> bool:
return True
def validate_anthropic_messages_environment(
self,
headers: dict,

View file

@ -52,6 +52,9 @@ class VertexAIAnthropicConfig(AnthropicConfig):
def custom_llm_provider(self) -> Optional[str]:
return "vertex_ai"
def should_strip_billing_metadata(self) -> bool:
return True
def _add_context_management_beta_headers(
self, beta_set: set, context_management: dict
) -> None:

View file

View file

@ -0,0 +1,7 @@
"""
You.com Search API module.
"""
from litellm.llms.you_com.search.transformation import YouComSearchConfig
__all__ = ["YouComSearchConfig"]

View file

@ -0,0 +1,193 @@
"""
Calls You.com's /v1/search endpoint to search the web.
You.com API Reference: https://you.com/docs/api-reference/search/v1-search
OpenAPI spec: https://you.com/specs/openapi_search_v1.yaml
"""
from typing import Dict, List, Optional, TypedDict, Union
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.search.transformation import (
BaseSearchConfig,
SearchResponse,
SearchResult,
)
from litellm.secret_managers.main import get_secret_str
class _YouComSearchRequestRequired(TypedDict):
"""Required fields for You.com Search API request."""
query: str
class YouComSearchRequest(_YouComSearchRequestRequired, total=False):
"""
You.com Search API request format.
Based on: https://you.com/specs/openapi_search_v1.yaml
"""
count: int
country: str
language: str
freshness: str
include_domains: List[str]
exclude_domains: List[str]
safesearch: str
class YouComSearchConfig(BaseSearchConfig):
# Keyed tier (higher rate limits): authenticate with X-API-Key.
YOU_COM_API_BASE = "https://ydc-index.io"
# Keyless free tier: IP-throttled (100 queries/day) and requires no auth.
# Used automatically when YOUCOM_API_KEY is not set.
YOU_COM_FREE_API_BASE = "https://api.you.com/v1/agents/search"
@staticmethod
def ui_friendly_name() -> str:
return "You.com"
def validate_environment(
self,
headers: Dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
**kwargs,
) -> Dict:
"""
Set headers for the You.com Search API.
If YOUCOM_API_KEY (or an explicit api_key) is present, use the keyed
endpoint with the `X-API-Key` header. Otherwise fall through to the
keyless free tier; no auth header is required.
"""
api_key = api_key or get_secret_str("YOUCOM_API_KEY")
headers["Content-Type"] = "application/json"
# Pin Accept-Encoding to identity: the keyless `api.you.com/v1/agents/search`
# endpoint advertises gzip content-encoding but returns body bytes the
# decoder rejects, which surfaces as httpx.DecodingError through litellm's
# http handler. Identity is harmless on the keyed endpoint.
headers.setdefault("Accept-Encoding", "identity")
if api_key:
headers["X-API-Key"] = api_key
return headers
def get_complete_url(
self,
api_base: Optional[str],
optional_params: dict,
data: Optional[Union[Dict, List[Dict]]] = None,
**kwargs,
) -> str:
"""
Pick the endpoint based on whether an API key is configured.
- api_base explicit override -> use it as-is (normalized)
- YOUCOM_API_KEY set -> keyed endpoint (ydc-index.io/v1/search)
- no key -> keyless free tier (api.you.com/v1/agents/search)
"""
if api_base is None:
api_base = get_secret_str("YOUCOM_API_BASE")
if api_base is None:
api_key = kwargs.get("api_key") or get_secret_str("YOUCOM_API_KEY")
if api_key:
api_base = self.YOU_COM_API_BASE
else:
# Keyless free tier already includes the full path.
return self.YOU_COM_FREE_API_BASE
api_base = api_base.rstrip("/")
if not api_base.endswith("/v1/search") and not api_base.endswith(
"/v1/agents/search"
):
api_base = f"{api_base}/v1/search"
return api_base
def transform_search_request(
self,
query: Union[str, List[str]],
optional_params: dict,
**kwargs,
) -> Dict:
"""
Transform Search request to You.com API format.
Perplexity unified spec → You.com mappings:
- query → query
- max_results → count
- search_domain_filter → include_domains
- country → country
- max_tokens_per_page → (not applicable, ignored)
"""
if isinstance(query, list):
query = " ".join(query)
request_data: YouComSearchRequest = {
"query": query,
}
if "max_results" in optional_params:
request_data["count"] = optional_params["max_results"]
if "search_domain_filter" in optional_params:
request_data["include_domains"] = optional_params["search_domain_filter"]
if "country" in optional_params:
request_data["country"] = optional_params["country"].lower()
result_data = dict(request_data)
for param, value in optional_params.items():
if (
param not in self.get_supported_perplexity_optional_params()
and param not in result_data
):
result_data[param] = value
return result_data
def transform_search_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs,
) -> SearchResponse:
"""
Transform You.com API response to LiteLLM unified SearchResponse format.
You.com → LiteLLM mappings (for both `results.web[]` and `results.news[]`):
- title → SearchResult.title
- url → SearchResult.url
- snippets[0] → SearchResult.snippet (falls back to `description`)
- page_age → SearchResult.date
"""
response_json = raw_response.json()
raw_results = response_json.get("results") or {}
web_results = raw_results.get("web") or []
news_results = raw_results.get("news") or []
results: List[SearchResult] = []
for item in list(web_results) + list(news_results):
snippets = item.get("snippets") or []
snippet = snippets[0] if snippets else item.get("description", "")
results.append(
SearchResult(
title=item.get("title", ""),
url=item.get("url", ""),
snippet=snippet,
date=item.get("page_age"),
last_updated=None,
)
)
return SearchResponse(
results=results,
object="search",
)

View file

@ -1322,7 +1322,9 @@ def completion( # type: ignore # noqa: PLR0915
preset_cache_key = kwargs.get("preset_cache_key", None)
hf_model_name = kwargs.get("hf_model_name", None)
supports_system_message = kwargs.get("supports_system_message", None)
base_model = kwargs.get("base_model", None)
base_model = kwargs.get("base_model", None) or (
model_info.get("base_model") if isinstance(model_info, dict) else None
)
### DISABLE FLAGS ###
disable_add_transform_inline_image_block = kwargs.get(
"disable_add_transform_inline_image_block", None
@ -1534,11 +1536,7 @@ def completion( # type: ignore # noqa: PLR0915
"logit_bias": logit_bias,
"user": user,
# params to identify the model
"model": (
model_info.get("base_model")
if isinstance(model_info, dict) and model_info.get("base_model")
else model
),
"model": model,
"custom_llm_provider": custom_llm_provider,
"response_format": response_format,
"seed": seed,
@ -6657,7 +6655,7 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse:
@client
def transcription(
def transcription( # noqa: PLR0915
model: str,
file: FileTypes,
## OPTIONAL OPENAI PARAMS ##
@ -6849,6 +6847,35 @@ def transcription(
else None
),
)
elif custom_llm_provider == "soniox":
from litellm.llms.soniox.audio_transcription.handler import (
SonioxAudioTranscriptionHandler,
)
response = SonioxAudioTranscriptionHandler().audio_transcriptions(
model=model,
audio_file=file,
optional_params=optional_params,
litellm_params=litellm_params_dict,
model_response=model_response,
atranscription=atranscription,
client=(
client
if client is not None
and (
isinstance(client, HTTPHandler)
or isinstance(client, AsyncHTTPHandler)
)
else None
),
timeout=timeout,
max_retries=max_retries,
logging_obj=litellm_logging_obj,
api_base=api_base,
api_key=api_key,
headers=extra_headers,
provider_config=provider_config, # type: ignore[arg-type]
)
elif provider_config is not None:
response = base_llm_http_handler.audio_transcriptions(
model=model,

View file

@ -1319,6 +1319,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1350,6 +1351,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1381,6 +1383,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1412,6 +1415,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1443,6 +1447,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1458,6 +1463,36 @@
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"jp.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@ -2194,6 +2229,7 @@
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -14015,6 +14051,22 @@
"/v1/images/generations"
]
},
"fal_ai/fal-ai/nano-banana": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.039,
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/gemini-25-flash-image": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.039,
"supported_endpoints": [
"/v1/images/generations"
]
},
"featherless_ai/featherless-ai/Qwerky-72B": {
"litellm_provider": "featherless_ai",
"max_input_tokens": 32768,
@ -24127,6 +24179,21 @@
"max_input_tokens": 200000,
"max_output_tokens": 8192
},
"minimax/MiniMax-M3": {
"input_cost_per_token": 6e-07,
"output_cost_per_token": 2.4e-06,
"cache_read_input_token_cost": 1.2e-07,
"litellm_provider": "minimax",
"mode": "chat",
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_vision": true,
"max_input_tokens": 512000,
"max_output_tokens": 128000
},
"mistral.devstral-2-123b": {
"input_cost_per_token": 4e-07,
"litellm_provider": "bedrock_converse",
@ -25009,6 +25076,7 @@
},
"moonshot/kimi-k2-0711-preview": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2026-05-25",
"input_cost_per_token": 6e-07,
"litellm_provider": "moonshot",
"max_input_tokens": 131072,
@ -25023,6 +25091,7 @@
},
"moonshot/kimi-k2-0905-preview": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2026-05-25",
"input_cost_per_token": 6e-07,
"litellm_provider": "moonshot",
"max_input_tokens": 262144,
@ -25037,6 +25106,7 @@
},
"moonshot/kimi-k2-turbo-preview": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2026-05-25",
"input_cost_per_token": 1.15e-06,
"litellm_provider": "moonshot",
"max_input_tokens": 262144,
@ -25061,6 +25131,7 @@
"source": "https://platform.moonshot.ai/docs/guide/kimi-k2-5-quickstart",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true
@ -25077,12 +25148,14 @@
"source": "https://platform.kimi.ai/docs/pricing/chat-k26",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true
},
"moonshot/kimi-latest": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2026-01-28",
"input_cost_per_token": 2e-06,
"litellm_provider": "moonshot",
"max_input_tokens": 131072,
@ -25097,6 +25170,7 @@
},
"moonshot/kimi-latest-128k": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2026-01-28",
"input_cost_per_token": 2e-06,
"litellm_provider": "moonshot",
"max_input_tokens": 131072,
@ -25111,6 +25185,7 @@
},
"moonshot/kimi-latest-32k": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2026-01-28",
"input_cost_per_token": 1e-06,
"litellm_provider": "moonshot",
"max_input_tokens": 32768,
@ -25125,6 +25200,7 @@
},
"moonshot/kimi-latest-8k": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2026-01-28",
"input_cost_per_token": 2e-07,
"litellm_provider": "moonshot",
"max_input_tokens": 8192,
@ -25139,6 +25215,7 @@
},
"moonshot/kimi-thinking-preview": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2025-11-11",
"input_cost_per_token": 6e-07,
"litellm_provider": "moonshot",
"max_input_tokens": 131072,
@ -25151,6 +25228,7 @@
},
"moonshot/kimi-k2-thinking": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2026-05-25",
"input_cost_per_token": 6e-07,
"litellm_provider": "moonshot",
"max_input_tokens": 262144,
@ -25166,6 +25244,7 @@
},
"moonshot/kimi-k2-thinking-turbo": {
"cache_read_input_token_cost": 1.5e-07,
"deprecation_date": "2026-05-25",
"input_cost_per_token": 1.15e-06,
"litellm_provider": "moonshot",
"max_input_tokens": 262144,
@ -25189,9 +25268,11 @@
"output_cost_per_token": 5e-06,
"source": "https://platform.moonshot.ai/docs/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"moonshot/moonshot-v1-128k-0430": {
"deprecation_date": "2024-04-30",
"input_cost_per_token": 2e-06,
"litellm_provider": "moonshot",
"max_input_tokens": 131072,
@ -25213,6 +25294,7 @@
"output_cost_per_token": 5e-06,
"source": "https://platform.moonshot.ai/docs/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
@ -25226,9 +25308,11 @@
"output_cost_per_token": 3e-06,
"source": "https://platform.moonshot.ai/docs/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"moonshot/moonshot-v1-32k-0430": {
"deprecation_date": "2024-04-30",
"input_cost_per_token": 1e-06,
"litellm_provider": "moonshot",
"max_input_tokens": 32768,
@ -25250,6 +25334,7 @@
"output_cost_per_token": 3e-06,
"source": "https://platform.moonshot.ai/docs/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
@ -25263,9 +25348,11 @@
"output_cost_per_token": 2e-06,
"source": "https://platform.moonshot.ai/docs/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"moonshot/moonshot-v1-8k-0430": {
"deprecation_date": "2024-04-30",
"input_cost_per_token": 2e-07,
"litellm_provider": "moonshot",
"max_input_tokens": 8192,
@ -25287,6 +25374,7 @@
"output_cost_per_token": 2e-06,
"source": "https://platform.moonshot.ai/docs/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
@ -25300,6 +25388,7 @@
"output_cost_per_token": 5e-06,
"source": "https://platform.moonshot.ai/docs/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"morph/morph-v3-fast": {
@ -33964,6 +34053,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -33992,6 +34082,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -36041,7 +36132,8 @@
"supports_prompt_caching": true,
"supports_response_schema": false,
"supports_tool_choice": true,
"supports_web_search": true
"supports_web_search": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-3-beta": {
"cache_read_input_token_cost": 7.5e-07,
@ -36240,7 +36332,8 @@
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true,
"supports_web_search": true
"supports_web_search": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-4-fast-non-reasoning": {
"cache_read_input_token_cost": 5e-08,
@ -36257,7 +36350,8 @@
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true,
"supports_web_search": true
"supports_web_search": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-4-0709": {
"input_cost_per_token": 3e-06,
@ -36273,7 +36367,8 @@
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true,
"supports_web_search": true
"supports_web_search": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-4-latest": {
"input_cost_per_token": 3e-06,
@ -36331,7 +36426,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
"supports_web_search": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-4-1-fast-reasoning-latest": {
"cache_read_input_token_cost": 5e-08,
@ -36352,7 +36448,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
"supports_web_search": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-4-1-fast-non-reasoning": {
"cache_read_input_token_cost": 5e-08,
@ -36372,7 +36469,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
"supports_web_search": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-4-1-fast-non-reasoning-latest": {
"cache_read_input_token_cost": 5e-08,
@ -36392,7 +36490,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
"supports_web_search": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-4.20-multi-agent-beta-0309": {
"cache_read_input_token_cost": 2e-07,
@ -36543,7 +36642,8 @@
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-code-fast-1-0825": {
"cache_read_input_token_cost": 2e-08,
@ -36558,7 +36658,8 @@
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"deprecation_date": "2026-05-15"
},
"xai/grok-vision-beta": {
"input_cost_per_image": 5e-06,
@ -41260,6 +41361,44 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"bedrock_mantle/openai.gpt-5.5": {
"input_cost_per_token": 5.5e-06,
"cache_read_input_token_cost": 5.5e-07,
"output_cost_per_token": 3.3e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"supported_endpoints": ["/v1/responses"],
"supported_modalities": ["text", "image"],
"supported_output_modalities": ["text"],
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"bedrock_mantle/openai.gpt-5.4": {
"input_cost_per_token": 2.75e-06,
"cache_read_input_token_cost": 2.75e-07,
"output_cost_per_token": 1.65e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"supported_endpoints": ["/v1/responses"],
"supported_modalities": ["text", "image"],
"supported_output_modalities": ["text"],
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"volcengine/doubao-seed-2-0-pro-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,
@ -41538,5 +41677,18 @@
"supports_vision": true,
"supports_native_structured_output": true,
"supports_pdf_input": true
},
"soniox/stt-async-v4": {
"litellm_provider": "soniox",
"max_output_tokens": 8000,
"max_tokens": 8000,
"input_cost_per_second": 0.0,
"output_cost_per_second": 0.0000277778,
"mode": "audio_transcription",
"source": "https://soniox.com/pricing",
"supported_endpoints": [
"/v1/audio/transcriptions"
],
"supports_audio_input": true
}
}
}

View file

@ -0,0 +1,66 @@
"""
Domain models for LiteLLM backend.
"""
from litellm.models.access_group import LiteLLM_AccessGroupTable
from litellm.models.budget import (
LiteLLM_BudgetTable,
LiteLLM_BudgetTableFull,
LiteLLM_TeamMemberTable,
)
from litellm.models.config import LiteLLM_Config
from litellm.models.credentials import (
CreateCredentialItem,
CredentialBase,
CredentialItem,
)
from litellm.models.end_user import LiteLLM_EndUserTable
from litellm.models.managed_files import (
LiteLLM_ManagedFileTable,
LiteLLM_ManagedObjectTable,
LiteLLM_ManagedVectorStoresTable,
LiteLLM_ManagedVectorStoreTable,
)
from litellm.models.mcp_server import LiteLLM_MCPServerTable
from litellm.models.model import LiteLLM_ProxyModelTable
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.models.organization import LiteLLM_OrganizationTable
from litellm.models.organization_membership import LiteLLM_OrganizationMembershipTable
from litellm.models.project import LiteLLM_ProjectTable
from litellm.models.skills import LiteLLM_SkillsTable
from litellm.models.spend_logs import LiteLLM_ErrorLogs, LiteLLM_SpendLogs
from litellm.models.tag import LiteLLM_TagTable
from litellm.models.team import LiteLLM_TeamTable
from litellm.models.team_membership import LiteLLM_TeamMembership
from litellm.models.user import LiteLLM_UserTable
from litellm.models.verification_token import LiteLLM_VerificationToken
__all__ = [
"LiteLLM_AccessGroupTable",
"LiteLLM_BudgetTable",
"LiteLLM_BudgetTableFull",
"LiteLLM_TeamMemberTable",
"LiteLLM_Config",
"CredentialBase",
"CredentialItem",
"CreateCredentialItem",
"LiteLLM_EndUserTable",
"LiteLLM_ManagedFileTable",
"LiteLLM_ManagedObjectTable",
"LiteLLM_ManagedVectorStoreTable",
"LiteLLM_ManagedVectorStoresTable",
"LiteLLM_MCPServerTable",
"LiteLLM_ProxyModelTable",
"LiteLLM_ObjectPermissionTable",
"LiteLLM_OrganizationTable",
"LiteLLM_OrganizationMembershipTable",
"LiteLLM_ProjectTable",
"LiteLLM_SkillsTable",
"LiteLLM_ErrorLogs",
"LiteLLM_SpendLogs",
"LiteLLM_TagTable",
"LiteLLM_TeamTable",
"LiteLLM_TeamMembership",
"LiteLLM_UserTable",
"LiteLLM_VerificationToken",
]

View file

@ -0,0 +1,26 @@
"""
Access group table model.
Canonical definition for ``litellm_accessgrouptable``. Re-exported from
``litellm.proxy._types`` for backwards compatibility.
"""
from datetime import datetime
from typing import List, Optional
from litellm.types.llms.base import LiteLLMPydanticObjectBase
class LiteLLM_AccessGroupTable(LiteLLMPydanticObjectBase):
access_group_id: str
access_group_name: str
description: Optional[str] = None
access_model_names: List[str] = []
access_mcp_server_ids: List[str] = []
access_agent_ids: List[str] = []
assigned_team_ids: List[str] = []
assigned_key_ids: List[str] = []
created_at: Optional[datetime] = None
created_by: Optional[str] = None
updated_at: Optional[datetime] = None
updated_by: Optional[str] = None

38
litellm/models/base.py Normal file
View file

@ -0,0 +1,38 @@
"""
Base model class for domain models.
"""
from datetime import datetime
from typing import Any, Dict, Optional
from pydantic import BaseModel, ConfigDict
class DomainModel(BaseModel):
"""Base class for all domain models."""
model_config = ConfigDict(
from_attributes=True,
protected_namespaces=(),
extra="ignore",
)
created_at: Optional[datetime] = None
updated_at: Optional[datetime] = None
@classmethod
def from_db_record(cls, record: Any) -> "DomainModel":
"""Create a domain model from a database record."""
if record is None:
raise ValueError("Cannot create domain model from None record")
if isinstance(record, dict):
return cls(**record)
if hasattr(record, "model_dump") and callable(record.model_dump):
return cls(**record.model_dump())
if hasattr(record, "dict") and callable(record.dict):
return cls(**record.dict())
return cls(**dict(record))
def to_db_dict(self, exclude_unset: bool = False) -> Dict[str, Any]:
"""Convert domain model to a dictionary for database operations."""
return self.model_dump(exclude_none=True, exclude_unset=exclude_unset)

56
litellm/models/budget.py Normal file
View file

@ -0,0 +1,56 @@
"""
Budget table model.
Canonical definition for ``litellm_budgettable``. Re-exported from
``litellm.proxy._types`` for backwards compatibility.
"""
from datetime import datetime
from typing import List, Optional
from pydantic import ConfigDict
from litellm.types.llms.base import LiteLLMPydanticObjectBase
class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase):
"""Represents user-controllable params for a LiteLLM_BudgetTable record.
Budget-write paths use `model_fields.keys()` on this class as an allowlist
for user input. Keep server-managed fields (e.g. `budget_reset_at`) on
`LiteLLM_BudgetTableFull` so they aren't user-settable.
"""
budget_id: Optional[str] = None
soft_budget: Optional[float] = None
max_budget: Optional[float] = None
max_parallel_requests: Optional[int] = None
tpm_limit: Optional[int] = None
rpm_limit: Optional[int] = None
model_max_budget: Optional[dict] = None
budget_duration: Optional[str] = None
allowed_models: Optional[List[str]] = (
None # per-member model scope; empty = inherit team models
)
model_config = ConfigDict(protected_namespaces=())
class LiteLLM_BudgetTableFull(LiteLLM_BudgetTable):
"""LiteLLM_BudgetTable + server-managed fields returned on API responses."""
budget_reset_at: Optional[datetime] = None
created_at: datetime
class LiteLLM_TeamMemberTable(LiteLLM_BudgetTable):
"""
Used to track spend of a user_id within a team_id
"""
spend: Optional[float] = None
user_id: Optional[str] = None
team_id: Optional[str] = None
budget_id: Optional[str] = None
model_config = ConfigDict(protected_namespaces=())

15
litellm/models/config.py Normal file
View file

@ -0,0 +1,15 @@
"""
Config table model.
Canonical definition for ``litellm_config``. Re-exported from
``litellm.proxy._types`` for backwards compatibility.
"""
from typing import Dict
from litellm.types.llms.base import LiteLLMPydanticObjectBase
class LiteLLM_Config(LiteLLMPydanticObjectBase):
param_name: str
param_value: Dict

View file

@ -0,0 +1,31 @@
"""
Credential table models.
These are the canonical credential types for the proxy. They live in the model
layer; ``litellm.types.utils`` re-exports them for backwards compatibility.
"""
from typing import Optional
from pydantic import BaseModel, model_validator
class CredentialBase(BaseModel):
credential_name: str
credential_info: dict
class CredentialItem(CredentialBase):
credential_values: dict
class CreateCredentialItem(CredentialBase):
credential_values: Optional[dict] = None
model_id: Optional[str] = None
@model_validator(mode="before")
@classmethod
def check_credential_params(cls, values):
if not values.get("credential_values") and not values.get("model_id"):
raise ValueError("Either credential_values or model_id must be set")
return values

View file

@ -0,0 +1,35 @@
"""
End-user table model.
Canonical definition for ``litellm_endusertable``. Re-exported from
``litellm.proxy._types`` for backwards compatibility.
"""
from typing import Literal, Optional
from pydantic import ConfigDict, model_validator
from litellm.models.budget import LiteLLM_BudgetTable
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.types.llms.base import LiteLLMPydanticObjectBase
class LiteLLM_EndUserTable(LiteLLMPydanticObjectBase):
user_id: str
blocked: bool
alias: Optional[str] = None
spend: float = 0.0
allowed_model_region: Optional[Literal["eu", "us"]] = None
default_model: Optional[str] = None
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
object_permission_id: Optional[str] = None
object_permission: Optional[LiteLLM_ObjectPermissionTable] = None
@model_validator(mode="before")
@classmethod
def set_model_info(cls, values):
if values.get("spend") is None:
values.update({"spend": 0.0})
return values
model_config = ConfigDict(protected_namespaces=())

View file

@ -0,0 +1,62 @@
"""
Managed file, object, and vector store table models.
Canonical definitions for the ``litellm_managed*`` tables. Re-exported from
``litellm.proxy._types`` for backwards compatibility.
"""
from datetime import datetime
from typing import Any, Dict, List, Literal, Optional, Union
from litellm.types.llms.base import LiteLLMPydanticObjectBase
from litellm.types.llms.openai import OpenAIFileObject, ResponsesAPIResponse
from litellm.types.utils import LiteLLMBatch, LiteLLMFineTuningJob
class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase):
unified_file_id: str
file_object: Optional[OpenAIFileObject] = None
model_mappings: Dict[str, str]
flat_model_file_ids: List[str]
created_by: Optional[str] = None
team_id: Optional[str] = None
updated_by: Optional[str] = None
storage_backend: Optional[str] = None
storage_url: Optional[str] = None
class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase):
unified_object_id: str
model_object_id: str
file_purpose: Literal["batch", "fine-tune", "response", "container"]
file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob, ResponsesAPIResponse]
created_by: Optional[str] = None
team_id: Optional[str] = None
class LiteLLM_ManagedVectorStoreTable(LiteLLMPydanticObjectBase):
"""Table for managing vector stores with target_model_names support."""
unified_resource_id: str
resource_object: Optional[Any] = None
model_mappings: Dict[str, str]
flat_model_resource_ids: List[str]
created_by: Optional[str] = None
team_id: Optional[str] = None
updated_by: Optional[str] = None
storage_backend: Optional[str] = None
storage_url: Optional[str] = None
class LiteLLM_ManagedVectorStoresTable(LiteLLMPydanticObjectBase):
vector_store_id: str
custom_llm_provider: str
vector_store_name: Optional[str]
vector_store_description: Optional[str]
vector_store_metadata: Optional[Dict[str, Any]]
created_at: Optional[datetime]
updated_at: Optional[datetime]
litellm_credential_name: Optional[str]
litellm_params: Optional[Dict[str, Any]]
team_id: Optional[str]
user_id: Optional[str]

View file

@ -0,0 +1,103 @@
"""
MCP server table model.
Canonical definition for ``litellm_mcpservertable``. Re-exported from
``litellm.proxy._types`` for backwards compatibility.
"""
import enum
from datetime import datetime
from typing import Dict, List, Literal, Optional
from pydantic import Field
from litellm.types.llms.base import LiteLLMPydanticObjectBase
from litellm.types.mcp import MCPAuthType, MCPCredentials, MCPTransportType
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
class MCPEnvVarScope(str, enum.Enum):
"""Scope for an MCP server environment variable.
- ``global``: value is provided by the admin and used for all users.
- ``user``: each user must provide their own value via the per-user
env-var endpoint. The admin-supplied ``value`` is treated as a
placeholder/hint and is not used at request time.
"""
global_ = "global"
user = "user"
class MCPEnvVar(LiteLLMPydanticObjectBase):
"""One environment variable for an MCP server.
Variables can be interpolated into ``static_headers`` using ``${NAME}``
syntax. ``scope=global`` values are stored on the server. ``scope=user``
values are stored per-user in ``LiteLLM_MCPUserEnvVars`` and supplied by
each user.
"""
name: str
value: str = ""
scope: MCPEnvVarScope = MCPEnvVarScope.global_
description: Optional[str] = None
class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
"""Represents a LiteLLM_MCPServerTable record"""
server_id: str
server_name: Optional[str] = None
alias: Optional[str] = None
description: Optional[str] = None
url: Optional[str] = None
spec_path: Optional[str] = None
transport: MCPTransportType
auth_type: Optional[MCPAuthType] = None
credentials: Optional[MCPCredentials] = None
instructions: Optional[str] = None
created_at: Optional[datetime] = None
created_by: Optional[str] = None
updated_at: Optional[datetime] = None
updated_by: Optional[str] = None
teams: List[Dict[str, Optional[str]]] = Field(default_factory=list)
mcp_access_groups: List[str] = Field(default_factory=list)
allowed_tools: List[str] = Field(default_factory=list)
tool_name_to_display_name: Optional[Dict[str, str]] = None
tool_name_to_description: Optional[Dict[str, str]] = None
extra_headers: List[str] = Field(default_factory=list)
mcp_info: Optional[MCPInfo] = None
static_headers: Optional[Dict[str, str]] = None
env_vars: Optional[List[MCPEnvVar]] = None
status: Optional[Literal["healthy", "unhealthy", "unknown"]] = Field(
default="unknown",
description="Health status: 'healthy', 'unhealthy', 'unknown'",
)
last_health_check: Optional[datetime] = None
health_check_error: Optional[str] = None
command: Optional[str] = None
args: List[str] = Field(default_factory=list)
env: Dict[str, str] = Field(default_factory=dict)
authorization_url: Optional[str] = None
token_url: Optional[str] = None
registration_url: Optional[str] = None
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
oauth_passthrough: bool = False
is_byok: bool = False
byok_description: List[str] = Field(default_factory=list)
byok_api_key_help_url: Optional[str] = None
has_user_credential: Optional[bool] = None
source_url: Optional[str] = None
timeout: Optional[float] = None
approval_status: Optional[str] = Field(
default="active",
description="Approval status: 'pending_review', 'active', 'rejected'",
)
submitted_by: Optional[str] = None
submitted_at: Optional[datetime] = None
reviewed_at: Optional[datetime] = None
review_notes: Optional[str] = None

59
litellm/models/model.py Normal file
View file

@ -0,0 +1,59 @@
"""
Proxy model table model.
Canonical definition for ``litellm_proxymodeltable``. Re-exported from
``litellm.proxy._types`` for backwards compatibility.
"""
import json
from datetime import datetime
from typing import Optional
from pydantic import ConfigDict, model_validator
from litellm.types.llms.base import LiteLLMPydanticObjectBase
class LiteLLM_ProxyModelTable(LiteLLMPydanticObjectBase):
model_id: str
model_name: str
litellm_params: dict
model_info: Optional[dict] = None
blocked: bool = False
created_at: Optional[datetime] = None
created_by: Optional[str] = None
updated_at: Optional[datetime] = None
updated_by: Optional[str] = None
model_config = ConfigDict(protected_namespaces=())
@model_validator(mode="before")
@classmethod
def check_potential_json_str(cls, values):
if isinstance(values.get("litellm_params"), str):
try:
values["litellm_params"] = json.loads(values["litellm_params"])
except json.JSONDecodeError:
pass
if isinstance(values.get("model_info"), str):
try:
values["model_info"] = json.loads(values["model_info"])
except json.JSONDecodeError:
pass
return values
@property
def is_blocked(self) -> bool:
return self.blocked
@property
def team_id(self) -> Optional[str]:
if self.model_info:
return self.model_info.get("team_id")
return None
@property
def team_public_model_name(self) -> Optional[str]:
if self.model_info:
return self.model_info.get("team_public_model_name")
return None

View file

@ -0,0 +1,26 @@
"""
Object permission table model.
Canonical definition for ``litellm_objectpermissiontable``. Re-exported from
``litellm.proxy._types`` for backwards compatibility.
"""
from typing import Dict, List, Optional
from litellm.types.llms.base import LiteLLMPydanticObjectBase
class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase):
"""Represents a LiteLLM_ObjectPermissionTable record"""
object_permission_id: str
mcp_servers: Optional[List[str]] = []
mcp_access_groups: Optional[List[str]] = []
mcp_tool_permissions: Optional[Dict[str, List[str]]] = None
vector_stores: Optional[List[str]] = []
agents: Optional[List[str]] = []
agent_access_groups: Optional[List[str]] = []
models: Optional[List[str]] = []
mcp_toolsets: Optional[List[str]] = None
blocked_tools: Optional[List[str]] = []
search_tools: Optional[List[str]] = []

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