Merge remote-tracking branch 'origin/main' into litellm_ptu_shares_per_team
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

This commit is contained in:
mateo-berri 2026-09-24 18:20:36 -07:00
commit e9c8e62d73
565 changed files with 12777 additions and 68994 deletions

View file

@ -31,7 +31,7 @@ while IFS= read -r file || [ -n "$file" ]; do
case "$file" in
model_prices_and_context_window.json | litellm/model_prices_and_context_window_backup.json | model_prices_and_context_window.schema.json)
has_cost_map=true ;;
tests/test_litellm/* | tests/proxy_unit_tests/*) : ;;
tests/test_litellm/* | tests/proxy_unit_tests/* | tests/unit/proxy/*) : ;;
*) outside_cost_map_set=true ;;
esac
done

View file

@ -0,0 +1,140 @@
#!/usr/bin/env bash
set -euo pipefail
flag="${1:?usage: unit_selection.sh <codecov flag>}"
legacy_flags=(
caching-local
enterprise-package
enterprise-routing
mcp-integration
proxy-db-auth-checks
proxy-db-budgets
proxy-db-custom-logging
proxy-db-db-and-spend
proxy-db-endpoints-and-responses
proxy-db-guardrails-hooks
proxy-db-jwt-and-keys
proxy-db-key-generation
proxy-db-logging-misc
proxy-db-proxy-runtime
proxy-db-proxy-server-core
proxy-db-proxy-utils
proxy-extras
proxy-infra
)
legacy_paths() {
case "$1" in
caching-local) echo tests/unit/caching ;;
enterprise-package)
echo tests/unit/enterprise/integrations
echo tests/unit/enterprise/proxy/auth
echo tests/unit/enterprise/proxy/guardrails
echo tests/unit/enterprise/proxy/hooks
echo tests/unit/enterprise/proxy/management_endpoints
echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
enterprise-routing)
echo tests/unit/enterprise/enterprise_callbacks/send_emails
echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py
echo tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py
echo tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py
echo tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py
echo tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py
echo tests/unit/enterprise/proxy/test_enterprise_routes.py
echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py
echo tests/unit/enterprise/proxy/test_managed_files_access_check.py
echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;;
mcp-integration)
echo tests/unit/proxy/_experimental/mcp_server
echo tests/unit/responses/mcp
echo tests/mcp_tests/test_proxy_mcp_e2e.py ;;
proxy-db-auth-checks)
echo tests/unit/proxy/auth/test_auth_checks.py
echo tests/unit/proxy/auth/test_user_api_key_auth.py
echo tests/unit/proxy/test_deprecated_key_grace_period.py ;;
proxy-db-budgets)
echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py
echo tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py
echo tests/unit/proxy/test_zero_cost_model_budget_bypass.py ;;
proxy-db-custom-logging)
echo tests/unit/proxy/test_custom_callback_input.py
echo tests/unit/proxy/test_custom_logger_s3_gcs.py ;;
proxy-db-db-and-spend)
echo tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py
echo tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py
echo tests/unit/proxy/db/test_update_daily_tag_spend.py
echo tests/unit/proxy/test_db_schema_changes.py
echo tests/unit/proxy/test_prisma_client_backoff_retry.py
echo tests/unit/proxy/test_update_spend.py
echo tests/unit/skills/test_skills_db.py ;;
proxy-db-endpoints-and-responses)
echo tests/unit/proxy/auth/test_models_fallback_endpoint.py
echo tests/unit/proxy/common_utils/test_check_batch_cost.py
echo tests/unit/proxy/common_utils/test_check_responses_cost.py
echo tests/unit/proxy/common_utils/test_realtime_cache.py
echo tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py
echo tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py
echo tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py
echo tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py
echo tests/unit/proxy/response_polling/test_response_polling_handler.py
echo tests/unit/proxy/test_custom_tokenizer_bug.py
echo tests/unit/proxy/test_get_favicon.py
echo tests/unit/proxy/test_get_image.py
echo tests/unit/proxy/test_prompt_test_endpoint.py
echo tests/unit/proxy/test_reducto_ocr_route.py
echo tests/unit/proxy/test_response_polling_pre_call_checks.py
echo tests/unit/proxy/test_ui_path_detection.py ;;
proxy-db-guardrails-hooks)
echo tests/unit/proxy/hooks/test_banned_keyword_list.py
echo tests/unit/proxy/test_proxy_setting_guardrails.py
echo tests/unit/proxy/test_unit_test_proxy_hooks.py ;;
proxy-db-jwt-and-keys)
echo tests/unit/proxy/auth/test_jwt.py
echo tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py
echo tests/unit/proxy/test_proxy_custom_auth.py ;;
proxy-db-key-generation) echo tests/unit/proxy/management_endpoints/test_key_generate_prisma.py ;;
proxy-db-logging-misc)
echo tests/unit/proxy/management_helpers/test_audit_logs_proxy.py
echo tests/unit/proxy/spend_tracking/test_search_api_logging.py
echo tests/unit/proxy/test_proxy_reject_logging.py ;;
proxy-db-proxy-runtime)
echo tests/unit/proxy/auth/test_multipart_bypass_repro.py
echo tests/unit/proxy/auth/test_proxy_routes.py
echo tests/unit/proxy/middleware/test_request_size_limit_middleware.py
echo tests/unit/proxy/test_proxy_config_unit_test.py
echo tests/unit/proxy/test_proxy_token_counter.py
echo tests/unit/proxy/test_server_root_path.py ;;
proxy-db-proxy-server-core)
echo tests/unit/proxy/test_aproxy_startup.py
echo tests/unit/proxy/test_proxy_server.py ;;
proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;;
proxy-extras) echo tests/unit/litellm_proxy_extras ;;
proxy-infra) echo tests/unit/gateway ;;
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
esac
}
expand() {
while read -r path; do
if [ -d "$path" ]; then
find "$path" -name 'test_*.py'
elif [ -f "$path" ]; then
echo "$path"
else
echo "unit_selection.sh: $path does not exist" >&2
exit 1
fi
done
}
if [ "$flag" = unit ]; then
comm -23 \
<(find tests/unit -name 'test_*.py' | sort) \
<(for legacy in "${legacy_flags[@]}"; do legacy_paths "$legacy"; done | expand | sort)
exit 0
fi
legacy_paths "$flag" | expand | sort

View file

@ -74,6 +74,7 @@ commands:
steps:
- run:
name: Install Codecov CLI (pinned v11.3.1)
when: always
command: |
curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov
curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM
@ -90,7 +91,6 @@ commands:
uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)"
setup_test_deps:
steps:
- checkout
- install_uv
- install_rust
- restore_cache:
@ -165,42 +165,72 @@ commands:
jobs:
unit:
parameters:
tests_path:
type: string
default: tests/unit
flag:
type: string
default: unit
shards:
type: integer
default: 6
workers:
type: integer
default: 4
dist:
type: string
default: loadscope
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
legacy_mcp_peer:
type: boolean
default: false
reruns:
type: integer
default: 0
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
parallelism: << parameters.shards >>
environment:
COVERAGE_CORE: sysmon
LITELLM_LOCAL_MODEL_COST_MAP: "True"
steps:
- setup_test_deps
- checkout
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- setup_test_deps
- when:
condition: << parameters.legacy_mcp_peer >>
steps:
- run:
name: Install MCP SDK1 peer
command: |
uv venv --python 3.12 .venv-mcp-peer
uv pip install --python .venv-mcp-peer 'mcp==1.28.1' 'langchain-mcp-adapters==0.2.1'
echo "export MCP_TEST_PEER_PYTHON=$PWD/.venv-mcp-peer/bin/python" >> "$BASH_ENV"
- run:
name: "Run << parameters.tests_path >> shard"
name: "Run << parameters.flag >> shard"
no_output_timeout: 20m
command: |
mkdir -p test-results/<< parameters.flag >>
mapfile -t files < <(find << parameters.tests_path >> -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)
if [ "${#files[@]}" -eq 0 ]; then echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.tests_path >> files; nothing to run"; exit 0; fi
selection="$(bash .circleci/scripts/unit_selection.sh << parameters.flag >>)" || { echo "unit_selection.sh failed for << parameters.flag >>"; exit 1; }
[ -n "${selection}" ] || { echo "unit_selection.sh produced no files for << parameters.flag >>"; exit 1; }
shard="$(printf '%s\n' "${selection}" | circleci tests split --split-by=timings --timings-type=filename)" || { echo "circleci tests split failed for << parameters.flag >>"; exit 1; }
[ -n "${shard}" ] || { echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.flag >> files; nothing to run"; exit 0; }
mapfile -t files < <(printf '%s\n' "${shard}")
xdist_args=()
if [ "<< parameters.workers >>" -gt 0 ]; then xdist_args=(-n << parameters.workers >> --dist=<< parameters.dist >>); fi
rerun_args=(-p no:rerunfailures)
if [ "<< parameters.reruns >>" -gt 0 ]; then rerun_args=(--reruns << parameters.reruns >> --reruns-delay 1 --rerun-except "from pytest-timeout"); fi
test_env=(PATH="$PATH" HOME="$HOME" CI=true COVERAGE_CORE="$COVERAGE_CORE" LITELLM_LOCAL_MODEL_COST_MAP="$LITELLM_LOCAL_MODEL_COST_MAP")
if [ -n "${MCP_TEST_PEER_PYTHON:-}" ]; then test_env+=(MCP_TEST_PEER_PYTHON="$MCP_TEST_PEER_PYTHON"); fi
set +e
uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
env -i "${test_env[@]}" \
uv run --no-sync pytest "${files[@]}" "${rerun_args[@]}" -p no:pytest-retry --timeout=90 "${xdist_args[@]}" --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
status=$?
set -e
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from the shard; passing"; exit 0; fi
@ -224,6 +254,7 @@ jobs:
resource_class: large
working_directory: ~/project
steps:
- checkout
- setup_test_deps
- run:
name: Checkout litellm-docs
@ -250,16 +281,17 @@ jobs:
resource_class: large
working_directory: ~/project
steps:
- setup_test_deps
- checkout
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- setup_test_deps
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
- start_redis
- run:
name: Run owned integration contracts
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
command: env -i PATH="$PATH" HOME="$HOME" CIRCLE_SHA1="$CIRCLE_SHA1" CIRCLE_WORKFLOW_ID="$CIRCLE_WORKFLOW_ID" bash .circleci/scripts/run_integration.sh << parameters.suite >>
no_output_timeout: 15m
- run:
name: Stop owned database and Redis
@ -282,6 +314,61 @@ workflows:
- unit:
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-<< matrix.flag >>
shards: 1
workers: 2
reruns: 2
matrix:
parameters:
flag: [caching-local, proxy-extras, enterprise-routing]
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-mcp-integration
flag: mcp-integration
shards: 1
workers: 2
legacy_mcp_peer: true
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-<< matrix.flag >>
shards: 1
reruns: 2
matrix:
parameters:
flag:
- enterprise-package
- proxy-infra
- proxy-db-auth-checks
- proxy-db-jwt-and-keys
- proxy-db-proxy-server-core
- proxy-db-proxy-runtime
- proxy-db-custom-logging
- proxy-db-logging-misc
- proxy-db-db-and-spend
- proxy-db-guardrails-hooks
- proxy-db-budgets
- proxy-db-endpoints-and-responses
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-proxy-db-proxy-utils
flag: proxy-db-proxy-utils
shards: 1
reruns: 2
dist: worksteal
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-proxy-db-key-generation
flag: proxy-db-key-generation
shards: 1
workers: 0
reruns: 2
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- documentation
- integration:
name: integration-<< matrix.suite >>

View file

@ -6,7 +6,8 @@
## TLDR
<!-- Fill in the bullets below and keep each one short and concrete: one line per bullet, roughly 10 words max -->
<!-- Fill in the bullets below and keep each one short and concrete: one line per bullet, roughly 10 words max
If the PR intentionally changes what existing users see or how a screen behaves, add a line under the bullets that starts "Intentional product change:" describing what changes, why, and what users lose. Reviewers must never have to infer a deliberate UX change from the diff -->
Problem this solves:
@ -28,7 +29,8 @@ How it solves it:
No LiteLLM internals: never name functions, files, DB tables, config classes, hooks, callbacks, or code paths. "The upload hands back an ID that looks like OpenAI's own `file-abc123` instead of the scrambled one the gateway returned" is right, "no managed-file row was registered" is wrong
Keep the two lists step-for-step identical until they diverge, so the changed step is obvious
If the bug had a security or authorization consequence, end each list with what another user could or could no longer do
Regenerate this section whenever new commits change the PR's behavior, so it never describes an older revision
Regenerate this section, screenshots included, whenever new commits change the PR's behavior, so it never describes an older revision
If the PR changes what an Admin UI page shows, embed a before and an after screenshot of that page right after its list, taken at the same URL on the same data, with the rows, fields, or controls that changed boxed in red so a reader spots the difference without reading the steps. These are the UI screenshots for Screenshots / Proof of Fix too: embed them once here and have that section's Before and After steps point back to them instead of repeating the images
Example:

View file

@ -34,7 +34,6 @@ GLOB_CHARS = frozenset("*?")
# tests has to be named by some shard or it runs nowhere. A child listed here is
# itself decomposed one level deeper and is checked through its own entry.
SHARDED_ROOTS: tuple[str, ...] = (
"tests/proxy_unit_tests",
"tests/test_litellm",
"tests/test_litellm/proxy",
)
@ -120,6 +119,13 @@ def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
)
def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]:
script: Final = repo_root / ".circleci/scripts/unit_selection.sh"
if not script.is_file():
return frozenset()
return frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(_uncommented(script.read_text())))
def _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
return frozenset(
match.group(0)
@ -611,7 +617,10 @@ def main() -> int:
scalars = _all_scalars()
integration_paths, ownership_findings = _integration_ownership()
test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars) | integration_paths) + ownership_findings
test_findings = (
_uncovered_tests(allowlist, _invoked_test_tokens(scalars) | _unit_selection_tokens() | integration_paths)
+ ownership_findings
)
dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars))
stale_findings = _stale_allowlist_paths(allowlist, test_files=_test_files(), dockerfiles=_dockerfiles())

View file

@ -13,6 +13,15 @@ on:
have its path existence-checked like any other token.
required: true
type: string
fork-flag:
description: >-
Codecov flag of the `.circleci/tests.yml` job that now owns part of
this shard. CircleCI does not run on pull requests from forks, so on
those events this shard also runs the files
`.circleci/scripts/unit_selection.sh` lists for the flag.
required: false
type: string
default: ""
workers:
description: "Number of pytest-xdist workers"
required: false
@ -92,6 +101,7 @@ jobs:
pull-requests: read
outputs:
decision: ${{ steps.changes.outputs.decision }}
has-coverage: ${{ steps.tests.outputs.has-coverage }}
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
@ -160,10 +170,13 @@ jobs:
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Run tests
id: tests
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: ${{ inputs.timeout-minutes }}
env:
TEST_PATH: ${{ inputs.test-path }}
FORK_FLAG: ${{ inputs.fork-flag }}
IS_FORK: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository }}
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
@ -171,9 +184,18 @@ jobs:
DIST: ${{ inputs.dist }}
COVERAGE_CORE: sysmon
run: |
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
selection="${TEST_PATH}"
if [ "${IS_FORK}" = "true" ] && [ -n "${FORK_FLAG}" ]; then
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${FORK_FLAG}" | tr '\n' ' ')"
fi
if [ -z "${selection// /}" ]; then
echo "shard selection is empty on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run"
exit 0
fi
pytest_args=()
existing_paths=0
for token in ${TEST_PATH:?}; do
for token in ${selection}; do
case "${token}" in
-*) pytest_args+=("${token}") ;;
*)
@ -187,7 +209,7 @@ jobs:
esac
done
if [ "${existing_paths}" -eq 0 ]; then
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
echo "No path in the selection exists (${selection}); nothing to run"
exit 0
fi
xdist_args=()
@ -209,8 +231,11 @@ jobs:
--cov-config=pyproject.toml
status=$?
set -e
if [ -f coverage.xml ]; then
echo "has-coverage=true" >> "$GITHUB_OUTPUT"
fi
if [ "$status" -eq 5 ]; then
echo "pytest collected no tests from ${TEST_PATH}; passing"
echo "pytest collected no tests from ${selection}; passing"
exit 0
fi
exit "$status"
@ -226,7 +251,7 @@ jobs:
upload-coverage:
name: Upload coverage to Codecov
needs: run
if: always() && needs.run.outputs.decision != 'skip'
if: always() && needs.run.outputs.decision != 'skip' && needs.run.outputs.has-coverage == 'true'
runs-on: ubuntu-latest
permissions:
contents: read

View file

@ -4,6 +4,7 @@ on:
pull_request:
paths:
- tests/e2e/claude_code/cron_vm/**
- tests/e2e/claude_code/pr_gate_version_resolver.py
- .github/workflows/compat-matrix-image.yml
workflow_dispatch:
@ -28,6 +29,14 @@ jobs:
- name: Build the Render cron image
run: docker build -f tests/e2e/claude_code/cron_vm/Dockerfile -t compat-matrix:${{ github.sha }} tests/e2e
- name: Run the pinned binaries as the cron user
- name: Resolve and install the Claude Code CLI as the cron user
run: |
docker run --rm compat-matrix:${{ github.sha }} bash -c 'set -e; whoami; claude --version; gh --version; uv --version'
docker run --rm compat-matrix:${{ github.sha }} bash -c '
set -euo pipefail
whoami
gh --version
uv --version
version="$(uv run --no-project --python 3.12 python /opt/litellm/tests/e2e/claude_code/pr_gate_version_resolver.py)"
/opt/litellm/tests/e2e/claude_code/cron_vm/install_claude_code.sh "${version}" /tmp/claude-cli
/tmp/claude-cli/claude --version
'

View file

@ -180,6 +180,18 @@ jobs:
echo "No changed tests/e2e Python files; skipping."
fi
- name: Run the claude_code harness unit tests
if: steps.changes.outputs.decision != 'skip'
run: |
if ! git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- ':(glob)tests/e2e/claude_code/**/*.py' ':(glob)tests/e2e/*.py' tests/e2e/claude_code/cron_vm/install_claude_code.sh pyproject.toml uv.lock .github/workflows/test-linting.yml | grep -q .; then
echo "No changed claude_code harness files; skipping."
exit 0
fi
retry() { "$@" || { sleep 15; "$@"; } || { sleep 30; "$@"; }; }
CLAUDE_VERSION="$(retry uv run --no-sync python tests/e2e/claude_code/pr_gate_version_resolver.py)"
tests/e2e/claude_code/cron_vm/install_claude_code.sh "$CLAUDE_VERSION" "$RUNNER_TEMP/claude-cli"
PATH="$RUNNER_TEMP/claude-cli:$PATH" uv run --no-sync pytest -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures tests/e2e/claude_code/_*_unit_tests
- name: Check for circular imports
if: steps.changes.outputs.decision != 'skip'
run: |

View file

@ -20,6 +20,12 @@ concurrency:
# rather than alphabetical letter ranges. Adding a new test file means adding it
# to whichever group it belongs to, not reshuffling slices.
#
# `.circleci/tests.yml` runs each group's files on same-repo events under the
# `proxy-db-<group>` Codecov flag; `.circleci/scripts/unit_selection.sh` holds
# the file lists. CircleCI does not build pull requests from forks, so `fork-flag`
# makes the shard run that list there. `test-path` keeps the files that still
# reach real providers and never left tests/proxy_unit_tests.
#
# Design targets:
# * Every shard runs in <= 7 minutes of wall-clock on the default runner.
# Most of a shard's time is pytest plugin load + xdist worker imports +
@ -58,7 +64,7 @@ jobs:
proxy-db:
needs: assert-shard-coverage
# Display only the semantic shard name in the checks UI instead of GHA's
# default "proxy-db (key-generation, tests/proxy_unit_tests/…, 0, loadscope, 20)"
# default "proxy-db (key-generation, tests/unit/proxy/…, 0, loadscope, 20)"
# which includes every matrix field and gets truncated past the test-path.
name: ${{ matrix.test-group }}
permissions:
@ -71,132 +77,93 @@ jobs:
include:
# Must run serially — event-loop conflict with the logging worker.
- test-group: key-generation
test-path: "tests/proxy_unit_tests/test_key_generate_prisma.py"
test-path: ""
fork-flag: proxy-db-key-generation
workers: 0
dist: loadscope
timeout: 20
# ---- auth: split into 2 shards ----
- test-group: auth-checks
test-path: >-
tests/proxy_unit_tests/test_auth_checks.py
tests/proxy_unit_tests/test_user_api_key_auth.py
tests/proxy_unit_tests/test_deprecated_key_grace_period.py
test-path: ""
fork-flag: proxy-db-auth-checks
workers: 4
dist: loadscope
timeout: 15
- test-group: jwt-and-keys
test-path: >-
tests/proxy_unit_tests/test_jwt.py
tests/proxy_unit_tests/test_jwt_key_mapping.py
tests/proxy_unit_tests/test_proxy_custom_auth.py
tests/proxy_unit_tests/test_key_generate_dynamodb.py
test-path: ""
fork-flag: proxy-db-jwt-and-keys
workers: 4
dist: loadscope
timeout: 15
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
- test-group: proxy-utils
test-path: "tests/proxy_unit_tests/test_proxy_utils.py"
test-path: ""
fork-flag: proxy-db-proxy-utils
workers: 4
dist: worksteal
timeout: 15
# ---- proxy server: split into 2 shards ----
- test-group: proxy-server-core
test-path: >-
tests/proxy_unit_tests/test_proxy_server.py
tests/proxy_unit_tests/test_aproxy_startup.py
test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py"
fork-flag: proxy-db-proxy-server-core
workers: 4
dist: loadscope
timeout: 15
- test-group: proxy-runtime
test-path: >-
tests/proxy_unit_tests/test_proxy_config_unit_test.py
tests/proxy_unit_tests/test_proxy_routes.py
tests/proxy_unit_tests/test_server_root_path.py
tests/proxy_unit_tests/test_proxy_pass_user_config.py
tests/proxy_unit_tests/test_proxy_token_counter.py
tests/proxy_unit_tests/test_request_size_limit_middleware.py
tests/proxy_unit_tests/test_multipart_bypass_repro.py
test-path: ""
fork-flag: proxy-db-proxy-runtime
workers: 4
dist: loadscope
timeout: 15
# ---- logging: split into 2 shards ----
- test-group: custom-logging
test-path: >-
tests/proxy_unit_tests/test_custom_callback_input.py
tests/proxy_unit_tests/test_custom_logger_s3_gcs.py
tests/proxy_unit_tests/test_proxy_custom_logger.py
test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py"
fork-flag: proxy-db-custom-logging
workers: 4
dist: loadscope
timeout: 15
- test-group: logging-misc
test-path: >-
tests/proxy_unit_tests/test_proxy_reject_logging.py
tests/proxy_unit_tests/test_audit_logs_proxy.py
tests/proxy_unit_tests/test_search_api_logging.py
test-path: ""
fork-flag: proxy-db-logging-misc
workers: 4
dist: loadscope
timeout: 15
- test-group: db-and-spend
test-path: >-
tests/proxy_unit_tests/test_prisma_client_backoff_retry.py
tests/proxy_unit_tests/test_db_schema_changes.py
tests/proxy_unit_tests/test_e2e_pod_lock_manager.py
tests/proxy_unit_tests/test_skills_db.py
tests/proxy_unit_tests/test_update_daily_tag_spend.py
tests/proxy_unit_tests/test_update_spend.py
tests/proxy_unit_tests/test_proxy_encrypt_decrypt.py
test-path: ""
fork-flag: proxy-db-db-and-spend
workers: 4
dist: loadscope
timeout: 15
# ---- guardrails + budget + hooks: split into 2 ----
- test-group: guardrails-hooks
test-path: >-
tests/proxy_unit_tests/test_proxy_setting_guardrails.py
tests/proxy_unit_tests/test_banned_keyword_list.py
tests/proxy_unit_tests/test_unit_test_proxy_hooks.py
test-path: ""
fork-flag: proxy-db-guardrails-hooks
workers: 4
dist: loadscope
timeout: 15
- test-group: budgets
test-path: >-
tests/proxy_unit_tests/test_default_end_user_budget_simple.py
tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py
tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py
test-path: ""
fork-flag: proxy-db-budgets
workers: 4
dist: loadscope
timeout: 15
- test-group: endpoints-and-responses
test-path: >-
tests/proxy_unit_tests/test_blog_posts_endpoint.py
tests/proxy_unit_tests/test_models_fallback_endpoint.py
tests/proxy_unit_tests/test_google_endpoint_routing.py
tests/proxy_unit_tests/test_google_gemini_proxy_request.py
tests/proxy_unit_tests/test_gemini_agents_endpoints.py
tests/proxy_unit_tests/test_get_favicon.py
tests/proxy_unit_tests/test_get_image.py
tests/proxy_unit_tests/test_reducto_ocr_route.py
tests/proxy_unit_tests/test_ui_path_detection.py
tests/proxy_unit_tests/test_prompt_test_endpoint.py
tests/proxy_unit_tests/test_check_batch_cost.py
tests/proxy_unit_tests/test_check_responses_cost.py
tests/proxy_unit_tests/test_response_polling_handler.py
tests/proxy_unit_tests/test_response_polling_pre_call_checks.py
tests/proxy_unit_tests/test_realtime_cache.py
tests/proxy_unit_tests/test_proxy_exception_mapping.py
tests/proxy_unit_tests/test_custom_tokenizer_bug.py
test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py"
fork-flag: proxy-db-endpoints-and-responses
workers: 4
dist: loadscope
timeout: 15
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
fork-flag: ${{ matrix.fork-flag }}
workers: ${{ matrix.workers }}
reruns: 2
timeout-minutes: ${{ matrix.timeout }}

View file

@ -31,10 +31,14 @@ concurrency:
# number, so a partially-specified entry would fail the call rather than fall
# back to the default.
#
# tests/proxy_unit_tests keeps its own caller (test-unit-proxy-db.yml): it is
# already a matrix and carries a shard-coverage guard that reads that file by
# name. Folding it in here is a follow-up, together with generalising that guard
# into assert_ci_coverage.py.
# tests/unit/proxy keeps its own caller (test-unit-proxy-db.yml): it is already
# a matrix and carries a shard-coverage guard that reads that file by name.
# Folding it in here is a follow-up, together with generalising that guard into
# assert_ci_coverage.py.
#
# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the
# shard under the same Codecov flag. CircleCI does not build pull requests from
# forks, so the shard still runs those files there and skips them elsewhere.
jobs:
unit:
name: ${{ matrix.shard }}
@ -49,6 +53,7 @@ jobs:
- shard: mcp-integration
artifact-name: mcp-integration
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
fork-flag: mcp-integration
workers: 2
reruns: 0
timeout-minutes: 20
@ -65,10 +70,10 @@ jobs:
- shard: enterprise-routing
artifact-name: enterprise-routing
test-path: >-
tests/test_litellm/enterprise
tests/test_litellm/google_genai
tests/test_litellm/router_utils
tests/test_litellm/router_strategy
fork-flag: enterprise-routing
workers: 2
reruns: 2
timeout-minutes: 20
@ -200,7 +205,7 @@ jobs:
tests/test_litellm/proxy/types_utils
tests/test_litellm/proxy/logging_endpoints
tests/test_litellm/proxy/test_*.py
tests/test_gateway
fork-flag: proxy-infra
workers: 4
reruns: 2
timeout-minutes: 20
@ -208,11 +213,8 @@ jobs:
- shard: caching-local
artifact-name: caching-local
test-path: >-
tests/local_testing/test_cache_preset_key.py
tests/local_testing/test_caching_handler.py
tests/local_testing/test_responses_stream_cache_keys.py
tests/local_testing/test_unit_test_caching.py
test-path: ""
fork-flag: caching-local
workers: 2
reruns: 2
timeout-minutes: 20
@ -220,7 +222,8 @@ jobs:
- shard: proxy-extras
artifact-name: proxy-extras
test-path: "tests/litellm-proxy-extras"
test-path: ""
fork-flag: proxy-extras
workers: 2
reruns: 2
timeout-minutes: 20
@ -228,7 +231,8 @@ jobs:
- shard: enterprise-package
artifact-name: enterprise-package
test-path: "tests/enterprise"
test-path: ""
fork-flag: enterprise-package
workers: 4
reruns: 2
timeout-minutes: 20
@ -247,6 +251,7 @@ jobs:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
fork-flag: ${{ matrix.fork-flag || '' }}
workers: ${{ matrix.workers }}
reruns: ${{ matrix.reruns }}
timeout-minutes: ${{ matrix.timeout-minutes }}

View file

@ -51,8 +51,8 @@ help:
@echo " make test-unit-core-utils - Run core utils tests (~32 files)"
@echo " make test-unit-other - Run other tests (caching, responses, etc., ~69 files)"
@echo " make test-unit-root - Run root-level tests (~34 files)"
@echo " make test-proxy-unit-a - Run proxy_unit_tests (a-o, ~20 files)"
@echo " make test-proxy-unit-b - Run proxy_unit_tests (p-z, ~28 files)"
@echo " make test-proxy-unit-a - Run tests/unit/proxy (a-o)"
@echo " make test-proxy-unit-b - Run tests/unit/proxy (p-z)"
@echo " make test-integration - Run integration tests"
@echo " make test-unit-helm - Run helm unit tests"
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
@ -332,17 +332,17 @@ test-unit-core-utils: install-test-deps
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
test-unit-other: install-test-deps
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/test_litellm/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/unit/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
test-unit-root: install-test-deps
$(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
# Proxy unit tests (tests/proxy_unit_tests split alphabetically)
# Proxy unit tests (tests/unit/proxy split alphabetically)
test-proxy-unit-a: install-test-deps
$(UV_RUN) pytest tests/proxy_unit_tests/test_[a-o]*.py --tb=short -vv -n 2 --durations=20
$(UV_RUN) pytest tests/unit/proxy --ignore-glob='tests/unit/proxy/test_[p-z]*.py' --tb=short -vv -n 2 --durations=20
test-proxy-unit-b: install-test-deps
$(UV_RUN) pytest tests/proxy_unit_tests/test_[p-z]*.py --tb=short -vv -n 2 --durations=20
$(UV_RUN) pytest tests/unit/proxy/test_[p-z]*.py tests/unit/skills --tb=short -vv -n 2 --durations=20
test-integration: install-test-deps
$(UV_RUN) pytest tests/ -k "not test_litellm"

View file

@ -1697,6 +1697,63 @@
"title": "litellm_video_duration_seconds_metric rate",
"type": "timeseries"
},
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"description": "Share of the provider's bill LiteLLM captured as spend over the scheduled capture-rate check's window (needs general_settings.spend_capture_rate_check); NaN while no rate is available",
"fieldConfig": {
"defaults": {
"color": {
"mode": "palette-classic"
},
"custom": {
"drawStyle": "line",
"fillOpacity": 10,
"lineWidth": 1,
"showPoints": "never",
"spanNulls": false
},
"unit": "percentunit"
},
"overrides": []
},
"gridPos": {
"h": 8,
"w": 12,
"x": 12,
"y": 107
},
"id": 111,
"options": {
"legend": {
"calcs": [],
"displayMode": "list",
"placement": "bottom",
"showLegend": true
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "max by (api_provider) (litellm_spend_capture_rate)",
"legendFormat": "{{api_provider}}",
"range": true,
"refId": "A"
}
],
"title": "litellm_spend_capture_rate",
"type": "timeseries"
},
{
"collapsed": false,
"gridPos": {

View file

@ -1,6 +1,6 @@
# LiteLLM All Prometheus Metrics dashboard
Every `litellm_*` metric family the proxy can expose on `/metrics` (134 families across 95 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard

View file

@ -0,0 +1,43 @@
-- One-shot backfill of LiteLLM_VerificationToken.total_spend (lifetime spend)
-- for keys created before the column was introduced in LiteLLM v1.103.0.
--
-- The column was added with DEFAULT 0 and no backfill, so keys that predate
-- the upgrade report lifetime spend below their current period spend. New
-- deployments do not need this script: total_spend is updated at request
-- time from the moment the release is deployed. Run it only if you want
-- pre-upgrade keys to show their historical lifetime spend. It sets lifetime
-- spend to at least the current spend on every key, active and archived,
-- because current period spend is a valid lower bound on lifetime spend.
-- For keys with no budget reset that is already the exact lifetime value;
-- for resetting keys it only recovers the current period. It is idempotent:
-- it only touches rows where total_spend is below spend, so re-running is a
-- no-op. It touches no spend logs and runs in seconds.
--
-- IMPORTANT caveats before running:
--
-- 1. Take a backup of the affected tables first:
-- pg_dump "$DATABASE_URL" -t '"LiteLLM_VerificationToken"' -t '"LiteLLM_DeletedVerificationToken"' > key_total_spend_backup.sql
--
-- 2. A key "resets" when its own budget_duration IS NOT NULL, or when its
-- budget_id links to a LiteLLM_BudgetTable row whose budget_duration IS
-- NOT NULL (a linked budget resets the key's spend each period too). For
-- those keys this script only recovers the current period;
-- db_scripts/backfill_key_total_spend_from_spend_logs.sql is an optional
-- follow-up that rebuilds the earlier periods from LiteLLM_SpendLogs.
--
-- 3. No proxy restart is needed. The proxy picks up the corrected values on
-- its next read of each key.
--
-- Usage:
-- psql "$DATABASE_URL" -f db_scripts/backfill_key_total_spend.sql
UPDATE "LiteLLM_VerificationToken"
SET total_spend = spend
WHERE total_spend < spend;
UPDATE "LiteLLM_DeletedVerificationToken"
SET total_spend = spend
WHERE total_spend < spend;
-- Verify: this should return 0.
-- SELECT count(*) FROM "LiteLLM_VerificationToken" WHERE total_spend < spend;

View file

@ -0,0 +1,89 @@
-- Optional follow-up to db_scripts/backfill_key_total_spend.sql. Run that
-- script first; this one rebuilds earlier budget periods for the keys it
-- can only partially fix: keys whose spend resets each period, because their own
-- budget_duration IS NOT NULL or because their budget_id links to a
-- LiteLLM_BudgetTable row whose budget_duration IS NOT NULL.
--
-- For those keys the "spend" column only covers the current period, so
-- lifetime spend is reconstructed from LiteLLM_SpendLogs. The join matches
-- l.api_key against both the stored token and its second sha256
-- (encode(sha256(convert_to(token, 'UTF8')), 'hex')), because spend logs
-- written by older paths recorded the re-hashed digest instead of the
-- token. It is idempotent and never lowers a value: every statement only
-- touches rows where total_spend is below the rebuilt sum, so re-running is
-- a no-op, and a key whose log history is shorter than its current period
-- keeps the value backfill_key_total_spend.sql already gave it.
--
-- IMPORTANT caveats before running:
--
-- 1. Take a backup of the affected tables first:
-- pg_dump "$DATABASE_URL" -t '"LiteLLM_VerificationToken"' -t '"LiteLLM_DeletedVerificationToken"' > key_total_spend_backup.sql
--
-- 2. It requires spend logs to have been enabled, and coverage is bounded
-- by maximum_spend_logs_retention_period: spend older than the retention
-- window is already gone and cannot be recovered.
--
-- 3. On a large SpendLogs table the join scan is slow, so run it off peak.
--
-- 4. Run it while the proxy is idle (or with traffic paused). The proxy
-- flushes spend logs in batches, so a request that already raised
-- total_spend but whose log is still queued is missing from the sum, and
-- the rebuilt value would be short by that in-flight amount.
--
-- 5. A custom token can be deleted and recreated, so the archived table can
-- hold several lifetimes of one token. The update only rewrites archived
-- rows that reset, and the log sum covers every lifetime of that token.
--
-- 6. No proxy restart is needed. The proxy picks up the corrected values on
-- its next read of each key.
--
-- Usage:
-- psql "$DATABASE_URL" -f db_scripts/backfill_key_total_spend_from_spend_logs.sql
-- Active keys whose spend resets (own budget_duration, or a linked
-- LiteLLM_BudgetTable row with one). Rebuild from LiteLLM_SpendLogs,
-- matching api_key against the stored token and its second sha256 digest.
UPDATE "LiteLLM_VerificationToken" k
SET total_spend = s.sum_spend
FROM (
SELECT k2.token, SUM(l.spend) AS sum_spend
FROM "LiteLLM_VerificationToken" k2
JOIN "LiteLLM_SpendLogs" l
ON l.api_key IN (k2.token, encode(sha256(convert_to(k2.token, 'UTF8')), 'hex'))
WHERE k2.budget_duration IS NOT NULL
OR k2.budget_id IN (
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
)
GROUP BY k2.token
) s
WHERE k.token = s.token
AND k.total_spend < s.sum_spend;
-- Archived tokens are not unique, so collapse them to one row per token
-- before joining spend logs; the update then hits every resetting archived
-- row.
UPDATE "LiteLLM_DeletedVerificationToken" k
SET total_spend = s.sum_spend
FROM (
SELECT k2.token, SUM(l.spend) AS sum_spend
FROM (
SELECT DISTINCT token
FROM "LiteLLM_DeletedVerificationToken"
WHERE budget_duration IS NOT NULL
OR budget_id IN (
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
)
) k2
JOIN "LiteLLM_SpendLogs" l
ON l.api_key IN (k2.token, encode(sha256(convert_to(k2.token, 'UTF8')), 'hex'))
GROUP BY k2.token
) s
WHERE k.token = s.token
AND k.total_spend < s.sum_spend
AND (k.budget_duration IS NOT NULL
OR k.budget_id IN (
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
));
-- Verify: this should return 0.
-- SELECT count(*) FROM "LiteLLM_VerificationToken" WHERE total_spend < spend;

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "kill_switch" JSONB;

View file

@ -72,6 +72,7 @@ model LiteLLM_AgentsTable {
agent_card_params Json
static_headers Json? @default("{}")
extra_headers String[] @default([])
kill_switch Json?
agent_access_groups String[] @default([])
access_group_ids String[] @default([])
object_permission_id String?

View file

@ -5,15 +5,15 @@ use litellm_llms::{
transformation::TextractDetectTextConfig,
},
azure_ai::ocr::{
cohere_parse_transformation::AzureAICohereParseConfig,
cohere_parse_transformation::{AZURE_COHERE_PARSE_PATH, AzureAICohereParseConfig},
document_intelligence::transformation::AzureDocumentIntelligenceOcrConfig,
transformation::AzureAiOcrConfig,
transformation::{AZURE_AI_OCR_PATH, AzureAiOcrConfig},
},
base_llm::ocr::{
error::Error,
handler::{self, CallHooks, OcrClient},
transformation::{
BaseOcrConfig, LiteLLMOcrResponse, OcrCredentialInputs, OcrDocument,
BaseOcrConfig, LiteLLMOcrResponse, OcrCredentialInputs, OcrDocument, OcrResponseFormat,
PreparedOcrRequest, ResolvedOcrCredentials,
},
},
@ -157,6 +157,36 @@ pub fn get_health_check_document(
.get_health_check_document())
}
/// Normalize a relayed Azure AI response into the LiteLLM OCR shape when
/// `endpoint` is the OCR route of the model's resolved config.
pub fn passthrough_response(
model: &str,
endpoint: &str,
body: &[u8],
) -> Result<Option<LiteLLMOcrResponse>, Error> {
let (model, config) = resolve_provider_config(model, Some("azure_ai"))?;
let segments: Vec<&str> = endpoint
.split('/')
.filter(|segment| !segment.is_empty())
.collect();
let is_ocr_endpoint = match config {
OcrConfigKind::AzureAi => segments == AZURE_AI_OCR_PATH,
OcrConfigKind::AzureCohere => segments == AZURE_COHERE_PARSE_PATH,
OcrConfigKind::AzureDocumentIntelligence => {
segments == AzureDocumentIntelligenceOcrConfig::analyze_path(&model)?
}
other => {
let provider: &'static str = other.provider().into();
return Err(Error::InvalidProvider(provider.to_owned()));
}
};
if !is_ocr_endpoint {
return Ok(None);
}
with_config!(config, config => config.transform_ocr_response(&model, body, OcrResponseFormat::Litellm))
.map(Some)
}
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
#[strum(serialize_all = "snake_case")]
pub(crate) enum OcrProvider {
@ -520,4 +550,68 @@ mod tests {
assert!(matches!(&error, Error::InvalidProvider(provider) if provider == "not_a_provider"));
assert_eq!(error.http_status_code(), Some(400));
}
#[rstest]
#[case("azure_ai/mistral-document-ai-2512", "providers/mistral/azure/ocr")]
#[case("azure_ai/mistral-document-ai-2512", "/providers/mistral/azure/ocr/")]
#[case("azure_ai/Cohere-parse-v5", "providers/cohere/v2/parse")]
#[case(
"azure_ai/doc-intelligence/prebuilt-layout",
"documentintelligence/documentModels/prebuilt-layout:analyze"
)]
fn passthrough_response_recognizes_the_resolved_config_ocr_route(
#[case] model: &str,
#[case] endpoint: &str,
) {
assert!(passthrough_response(model, endpoint, b"not json").is_err());
}
#[rstest]
#[case("azure_ai/mistral-document-ai-2512", "models/info")]
#[case("azure_ai/mistral-document-ai-2512", "providers/cohere/v2/parse")]
#[case("azure_ai/Cohere-parse-v5", "providers/mistral/azure/ocr")]
#[case(
"azure_ai/doc-intelligence/prebuilt-layout",
"documentintelligence/documentModels/prebuilt-read:analyze"
)]
fn passthrough_response_skips_other_routes(#[case] model: &str, #[case] endpoint: &str) {
assert!(
passthrough_response(model, endpoint, b"not json")
.unwrap()
.is_none()
);
}
#[test]
fn passthrough_response_normalizes_the_mistral_body() {
let body = br#"{
"pages": [{"index": 0, "markdown": "page one"}, {"index": 1, "markdown": "page two"}],
"model": "mistral-document-ai-2512",
"usage_info": {"pages_processed": 2}
}"#;
let json = passthrough_response(
"azure_ai/mistral-document-ai-2512",
"providers/mistral/azure/ocr",
body,
)
.unwrap()
.unwrap()
.into_json();
assert_eq!(json["usage_info"]["pages_processed"], 2);
assert_eq!(json["pages"][0]["markdown"], "page one");
}
#[test]
fn passthrough_response_counts_cohere_billed_pages() {
let body = br#"{"id": "parse-1", "pages": [], "meta": {"billed_units": {"pages": 3}}}"#;
let json = passthrough_response(
"azure_ai/Cohere-parse-v5",
"providers/cohere/v2/parse",
body,
)
.unwrap()
.unwrap()
.into_json();
assert_eq!(json["usage_info"]["pages_processed"], 3);
}
}

View file

@ -12,7 +12,7 @@ pyo3.workspace = true
pyo3-async-runtimes.workspace = true
pythonize.workspace = true
serde.workspace = true
tokio = { workspace = true, features = ["sync"] }
tokio = { workspace = true, features = ["rt", "sync"] }
[dev-dependencies]
rstest.workspace = true

View file

@ -225,29 +225,12 @@ mod tests {
use pyo3::exceptions::PyLookupError;
use pyo3::panic::PanicException;
use pyo3::types::{PyDict, PyModule};
use rstest::{fixture, rstest};
use rstest::rstest;
use serde::Serializer;
use tokio::runtime::Builder;
use super::*;
struct InitializedPython;
impl InitializedPython {
fn attach<F, R>(&self, f: F) -> R
where
F: for<'py> FnOnce(Python<'py>) -> R,
{
Python::attach(f)
}
}
#[fixture]
#[once]
fn initialized_python() -> InitializedPython {
crate::initialize_python();
InitializedPython
}
use crate::{InitializedPython, initialized_python};
#[derive(Debug)]
struct Error(String);

View file

@ -1,6 +1,57 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{
Arc, Mutex,
atomic::{AtomicU64, Ordering},
};
use pyo3::prelude::*;
use pyo3::{exceptions::PyRuntimeError, prelude::*};
/// The caller's `contextvars` context, captured at the Python entry point so blocking Python
/// work started from Rust sees the same request-local values as the Python caller.
#[derive(Clone)]
pub struct PythonContext(Arc<Py<PyAny>>);
impl PythonContext {
pub fn capture(py: Python<'_>) -> PyResult<Self> {
Ok(Self(Arc::new(
py.import("contextvars")?
.call_method0("copy_context")?
.unbind(),
)))
}
/// Runs `f` inside a fresh copy of the captured context. The copy is what lets two blocking
/// calls run concurrently: a `contextvars.Context` cannot be entered twice at once.
pub fn enter<T, F>(&self, py: Python<'_>, f: F) -> PyResult<T>
where
F: FnOnce(Python<'_>) -> T + Send + Sync + 'static,
T: Send + 'static,
{
let copy = self.0.bind(py).call_method0("copy")?;
let body = Arc::new(Mutex::new(Some(f)));
let value = Arc::new(Mutex::new(None::<T>));
let callback = {
let body = Arc::clone(&body);
let value = Arc::clone(&value);
pyo3::types::PyCFunction::new_closure(py, None, None, move |_, _| {
Python::attach(|py| -> PyResult<()> {
let body = body
.lock()
.expect("context body slot poisoned")
.take()
.expect("the context body ran more than once");
*value.lock().expect("context value slot poisoned") = Some(body(py));
Ok(())
})
})?
};
copy.call_method1("run", (callback,))?;
value
.lock()
.expect("context value slot poisoned")
.take()
.ok_or_else(|| PyRuntimeError::new_err("the context body produced no value"))
}
}
static GIL_RELEASES: AtomicU64 = AtomicU64::new(0);
@ -19,3 +70,270 @@ where
pub fn release_count() -> u64 {
GIL_RELEASES.load(Ordering::Relaxed)
}
/// Runs Python work that may block, such as a secret manager read or a callback that does
/// I/O, on the runtime's blocking pool so the async workers stay free to poll other calls.
/// The work runs inside a copy of `context` so request-local `contextvars` survive the hop.
pub async fn attach_blocking<T, F>(context: PythonContext, f: F) -> PyResult<T>
where
F: for<'py> FnOnce(Python<'py>) -> T + Send + Sync + 'static,
T: Send + 'static,
{
match tokio::task::spawn_blocking(move || Python::attach(|py| context.enter(py, f))).await {
Ok(value) => value,
Err(error) if error.is_panic() => std::panic::resume_unwind(error.into_panic()),
Err(error) => panic!("the blocking pool dropped a Python call: {error}"),
}
}
#[cfg(test)]
mod tests {
use std::time::{Duration, Instant};
use pyo3::{exceptions::PyRuntimeError, prelude::*, types::PyDict};
use rstest::{fixture, rstest};
use super::{PythonContext, attach_blocking};
use crate::{InitializedPython, initialized_python, run_sync_value};
#[fixture]
fn namespace(#[from(initialized_python)] python: &InitializedPython) -> Py<PyDict> {
python.attach(|py| {
let namespace = PyDict::new(py);
py.run(
c"
import threading, time
finished = False
def work(seconds):
global finished
time.sleep(seconds)
finished = True
def observe(expression):
return eval(expression)
",
Some(&namespace),
None,
)
.unwrap();
namespace.unbind()
})
}
fn observe<T: for<'a, 'py> FromPyObject<'a, 'py, Error: std::fmt::Debug>>(
namespace: &Py<PyDict>,
py: Python<'_>,
expression: &str,
) -> T {
namespace
.bind(py)
.get_item("observe")
.unwrap()
.unwrap()
.call1((expression,))
.unwrap()
.extract()
.unwrap()
}
fn work(namespace: &Py<PyDict>, py: Python<'_>, seconds: f64) {
namespace
.bind(py)
.get_item("work")
.unwrap()
.unwrap()
.call1((seconds,))
.unwrap();
}
fn shared(namespace: &Py<PyDict>) -> Py<PyDict> {
Python::attach(|py| namespace.clone_ref(py))
}
fn context() -> PythonContext {
Python::attach(|py| PythonContext::capture(py).unwrap())
}
#[fixture]
fn request_context() -> (PythonContext, Py<PyAny>) {
Python::initialize();
Python::attach(|py| {
let namespace = PyDict::new(py);
py.run(
c"import contextvars\nrequest_var = contextvars.ContextVar('request_var', default='unset')",
Some(&namespace),
None,
)
.unwrap();
let var = namespace.get_item("request_var").unwrap().unwrap();
var.call_method1("set", ("request-value",)).unwrap();
(PythonContext::capture(py).unwrap(), var.unbind())
})
}
#[rstest]
#[tokio::test]
async fn blocking_python_work_leaves_the_runtime_free_to_run_other_tasks(
namespace: Py<PyDict>,
) {
let (python_done, timer_done) = tokio::join!(
attach_blocking(context(), move |py| {
work(&namespace, py, 0.3);
Instant::now()
}),
async {
tokio::time::sleep(Duration::from_millis(30)).await;
Instant::now()
}
);
assert!(
timer_done < python_done.unwrap(),
"the timer only finished after the Python call: the call ran inline on the worker"
);
}
#[rstest]
#[tokio::test]
async fn python_work_runs_off_the_thread_polling_the_future(namespace: Py<PyDict>) {
let polling: u64 = Python::attach(|py| observe(&namespace, py, "threading.get_ident()"));
let worker: u64 = attach_blocking(context(), move |py| {
observe(&namespace, py, "threading.get_ident()")
})
.await
.unwrap();
assert_ne!(worker, polling);
}
#[rstest]
#[tokio::test]
async fn a_dropped_await_never_interrupts_the_python_call(namespace: Py<PyDict>) {
let handle = shared(&namespace);
let started = tokio::time::timeout(
Duration::from_millis(10),
attach_blocking(context(), move |py| work(&handle, py, 0.1)),
)
.await;
assert!(
started.is_err(),
"the await was dropped before the call returned"
);
tokio::time::sleep(Duration::from_millis(300)).await;
let finished: bool = Python::attach(|py| observe(&namespace, py, "finished"));
assert!(finished);
}
#[rstest]
#[tokio::test]
async fn a_panic_in_python_work_reaches_the_awaiting_task(
#[from(initialized_python)] _python: &InitializedPython,
) {
let joined = tokio::spawn(attach_blocking(context(), |_| -> () {
panic!("python work failed")
}))
.await;
let error = joined.expect_err("the panic propagates instead of being swallowed");
assert!(error.is_panic());
}
#[rstest]
fn a_sync_route_can_await_python_work_without_deadlocking_on_the_gil(
#[from(initialized_python)] python: &InitializedPython,
) {
let value = python
.attach(|py| {
let context = PythonContext::capture(py).unwrap();
run_sync_value(py, async move {
tokio::time::timeout(
Duration::from_secs(5),
attach_blocking(context, |_| Python::version_str().len()),
)
.await
.map_err(|_| {
PyRuntimeError::new_err("the blocking call never re-acquired the GIL")
})
})
})
.unwrap()
.unwrap();
assert!(value > 0);
}
#[rstest]
#[tokio::test]
async fn blocking_work_sees_the_callers_contextvars(
request_context: (PythonContext, Py<PyAny>),
) {
let (context, var) = request_context;
let seen: String = attach_blocking(context, move |py| {
var.bind(py)
.call_method0("get")
.unwrap()
.extract::<String>()
.unwrap()
})
.await
.unwrap();
assert_eq!(seen, "request-value");
}
#[rstest]
#[tokio::test]
async fn concurrent_blocking_calls_each_enter_a_context_copy(
request_context: (PythonContext, Py<PyAny>),
) {
let (context, var) = request_context;
let first = Python::attach(|py| var.clone_ref(py));
let second = var;
let (first_seen, second_seen) = tokio::join!(
attach_blocking(context.clone(), move |py| {
first
.bind(py)
.call_method0("get")
.unwrap()
.extract::<String>()
.unwrap()
}),
attach_blocking(context, move |py| {
second
.bind(py)
.call_method0("get")
.unwrap()
.extract::<String>()
.unwrap()
}),
);
assert_eq!(first_seen.unwrap(), "request-value");
assert_eq!(second_seen.unwrap(), "request-value");
}
#[rstest]
#[tokio::test]
async fn writes_inside_blocking_work_do_not_leak_back_to_the_caller(
request_context: (PythonContext, Py<PyAny>),
) {
let (context, var) = request_context;
let leaked = Python::attach(|py| var.clone_ref(py));
attach_blocking(context, move |py| {
var.bind(py).call_method1("set", ("worker-value",)).unwrap();
})
.await
.unwrap();
let caller_value: String = Python::attach(|py| {
leaked
.bind(py)
.call_method0("get")
.unwrap()
.extract()
.unwrap()
});
assert_ne!(caller_value, "worker-value");
}
}

View file

@ -25,7 +25,7 @@ pub use execution::{
runtime_started,
};
pub use fork_gate::RuntimeAlreadyStarted;
pub use gil::{release_count, release_gil};
pub use gil::{PythonContext, attach_blocking, release_count, release_gil};
pub use handle::{Execution, ExecutionBody, ExecutionStep};
pub use marshal::{
Pythonized, from_py, from_py_argument, json_loads, json_object_field, panic_to_pyerr, to_py,
@ -43,3 +43,24 @@ pub(crate) fn initialize_python() {
});
});
}
#[cfg(test)]
pub(crate) struct InitializedPython;
#[cfg(test)]
impl InitializedPython {
pub(crate) fn attach<F, R>(&self, f: F) -> R
where
F: for<'py> FnOnce(pyo3::Python<'py>) -> R,
{
pyo3::Python::attach(f)
}
}
#[cfg(test)]
#[rstest::fixture]
#[once]
pub(crate) fn initialized_python() -> InitializedPython {
initialize_python();
InitializedPython
}

View file

@ -16,6 +16,8 @@ use crate::{
},
};
pub const AZURE_COHERE_PARSE_PATH: [&str; 4] = ["providers", "cohere", "v2", "parse"];
#[derive(Default)]
pub struct AzureAICohereParseConfig;
@ -131,7 +133,7 @@ impl AzureAICohereParseConfig {
}
url.set_path(path.strip_suffix("/models").unwrap_or(&path));
ApiUrl::parse(url.as_str())
.and_then(|url| url.complete_path(&["providers", "cohere", "v2", "parse"]))
.and_then(|url| url.complete_path(&AZURE_COHERE_PARSE_PATH))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base())
}

View file

@ -561,6 +561,14 @@ async fn poll_operation(
}
impl AzureDocumentIntelligenceOcrConfig {
pub fn analyze_path(model: &str) -> Result<[String; 3], Error> {
Ok([
"documentintelligence".into(),
"documentModels".into(),
format!("{}:analyze", model_id(model)?),
])
}
fn build_ocr_url(
&self,
endpoint: &str,
@ -568,9 +576,9 @@ impl AzureDocumentIntelligenceOcrConfig {
params: &DocumentIntelligenceParams,
api_version: &str,
) -> Result<String, Error> {
let model = format!("{}:analyze", model_id(model)?);
let path = Self::analyze_path(model)?;
ApiUrl::parse(endpoint)
.and_then(|url| url.complete_path(&["documentintelligence", "documentModels", &model]))
.and_then(|url| url.complete_path(&path.each_ref().map(String::as_str)))
.map(|url| {
url.append_query_pairs(
[("api-version", api_version)]

View file

@ -17,7 +17,7 @@ use crate::{
mistral::ocr::transformation::{MistralOcrConfig, MistralOcrRequest},
};
const AZURE_AI_OCR_PATH: &str = "/providers/mistral/azure/ocr";
pub const AZURE_AI_OCR_PATH: [&str; 4] = ["providers", "mistral", "azure", "ocr"];
const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY";
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
@ -179,9 +179,8 @@ impl AzureAiOcrConfig {
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
let base = Self::resolve_api_base(api_base, env_lookup)?;
let path: Vec<&str> = AZURE_AI_OCR_PATH.trim_matches('/').split('/').collect();
ApiUrl::parse(&base)
.and_then(|url| url.complete_path(&path))
.and_then(|url| url.complete_path(&AZURE_AI_OCR_PATH))
.map(|url| url.into_string())
.map_err(|_| Error::RequestField {
path: "api_base".into(),

View file

@ -7,6 +7,26 @@
- Python, Rust SDK and gateway use one lifecycle-bearing core route entrypoint; provider helpers stay private, never bridge-accessible transport drivers
- Built-in provider/config/secret/auth/document preparation stays in Rust; caller-authored callbacks and focused Python-file reads run only at core-selected points
- Target GIL-enabled CPython explicitly with `#[pymodule(gil_used = true)]`; detach Rust-only work
- GIL and tokio invariants, each pinned by a test in `host-python` (`execution.rs`,
`gil.rs`) so a regression fails there before it deadlocks a proxy:
- Never hold the GIL while waiting on the runtime. A sync entrypoint releases it with
`release_gil` around `block_on`, because every task that attaches would otherwise wait
on the thread that is waiting on them ([pyo3 parallelism](https://pyo3.rs/v0.29.2/parallelism.html))
- Never `block_on` from a tokio worker; the sync entrypoints refuse with "cannot run from
a Tokio context" instead of panicking inside the runtime ([tokio `Runtime::block_on`](https://docs.rs/tokio/latest/tokio/runtime/struct.Runtime.html#method.block_on))
- Inside a future, `Python::attach` only for GIL-cheap work: cloning a `Py<T>`, building
a small value, reading a settings snapshot. Anything that can block (a secret manager
read, a callback that does I/O, an import, a network call) goes through
`litellm_host_python::attach_blocking`, which runs it on the blocking pool so the async
workers keep polling other calls ([tokio `spawn_blocking`](https://docs.rs/tokio/latest/tokio/task/fn.spawn_blocking.html)).
`block_in_place` is not an alternative: it needs a multi-thread worker and still steals it
- `attach_blocking` work runs on a thread the interpreter did not create (pinned by the
`threading.get_ident()` test). Like any foreign-thread attach it therefore has no running
asyncio loop and a fresh `contextvars` context: do not hand it a coroutine or anything
bound to the caller's loop
- Dropping the await (an asyncio cancel) does not interrupt the Python call; it runs to
completion and its result is discarded. A panic in it reaches the awaiting task as a panic
- Add a case to `gil.rs` when a new seam changes any of these; the tests are the spec
- Free-threading requires separate runtime/concurrency validation; omitting the attribute does not opt out on PyO3 0.28+
- Preserve public argument binding and Python object provenance
- Project only consumed fields at reference read points; no eager whole-graph serialization or equality-based alias reconstruction
@ -56,6 +76,14 @@ GIL handling to `litellm-host-python`.
decides whether to raise or fall back. For a rust-only provider/route (no
Python reference), the Python side is a thin dispatch that calls Rust and
raises when the bridge is unavailable, with no fallback.
- Declare it by passing `python=NO_PYTHON` (`litellm.rust_bridge.runtime`)
to `PublicDispatch.run`/`arun` or `runtime.run`/`arun`, never a stand-in
callable that raises, and give every context of it a `RUST_REQUIRED`
catalog rule
- Any other decision, an unprojectable call, or a bypass raises
`NoPythonImplementationError` before native runs, so a misdeclared route
fails in tests instead of reaching deleted code. When deleting a route's
Python implementation, switch its dispatch to `NO_PYTHON` in the same change
- Keep the Python interface minimal (well under 100 lines per route): it only
marshals inputs and calls Rust. Do not add per-route feature flags, and do
not put provider dispatch in `litellm/main.py`; it lives in a thin dispatch

View file

@ -1,10 +0,0 @@
Native OCR uses `litellm_secrets::source::SecretSource`. Built-in secret managers resolve to retained Rust backends. Custom Python managers and overrides keep the callback path. Readable managers still require the Rust secret-manager binding to be enabled
The shared proxy initializer captures native configuration without loading the extension or doing native I/O. `_SecretManagerRuntime.from_client` constructs a backend on first use and keeps its handle on the Python client. The secret-manager dispatcher selects Python or Rust through `catalog.py`. Native reads call that handle; Rust routes extract the backend directly. Configuration changes replace the handle, while calls already bound to the previous backend keep using it. Handles cannot be reused after fork. Directly constructed LiteLLM managers are adapted on first native use. Manually supplied SDK clients keep their Python behavior because their credentials cannot be inferred safely. Provider implementations contain no bridge registration
Retention describes ownership and lifetime. `callbacks-legacy-python::PublicCall` owns Python references for one call to preserve identity. A native cache or secret-manager handle owns shared Rust state across calls to preserve connection pools and caches. Both use existing `Py<T>` and shared Rust ownership, with execution and GIL transitions handled by `litellm-host-python`
Cache and secret-manager catalog entries remain Python-only, including when `LITELLM_RUST=1`. This wiring does not change rollout policy
OCR provider requests use the shared `litellm-http` pool. AWS and Google secret-manager SDK clients keep their SDK transports, which do not yet inherit the pool's proxy, TLS, certificate, timeout, or observability configuration. Preserve those SDK transports and configure them equivalently instead of forcing them through reqwest

View file

@ -34,7 +34,7 @@ mod _native {
#[pymodule_export]
use crate::routes::messages::{amessages, messages};
#[pymodule_export]
use crate::routes::ocr::{aocr, ocr};
use crate::routes::ocr::{aocr, ocr, ocr_health_check_document, ocr_passthrough_response};
#[pymodule_export]
use crate::routes::responses::{ResponsesWebSocketConnection, aresponses, responses};
#[pymodule_export]
@ -85,6 +85,8 @@ mod tests {
"ProcessReservedForForking",
"ocr",
"aocr",
"ocr_health_check_document",
"ocr_passthrough_response",
"embedding",
"aembedding",
"transcription",

View file

@ -8,8 +8,9 @@ use std::sync::LazyLock;
use host::OcrRouteHost;
use litellm_auth_gcp::VertexAuth;
use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call};
use litellm_core::ocr::route::ocr_machine;
use litellm_core::ocr::{provider_config, route::ocr_machine};
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_host_python::to_py;
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
use pyo3::{
prelude::*,
@ -106,6 +107,30 @@ pub(crate) fn aocr(
run_ocr(py, request, args, kwargs, true)
}
#[pyfunction]
pub(crate) fn ocr_health_check_document(
py: Python<'_>,
model: &str,
custom_llm_provider: Option<&str>,
) -> PyResult<Py<PyAny>> {
let document = provider_config::get_health_check_document(model, custom_llm_provider)
.map_err(errors::to_pyerr)?;
to_py(py, &document)
}
#[pyfunction]
pub(crate) fn ocr_passthrough_response(
py: Python<'_>,
model: &str,
endpoint: &str,
body: &[u8],
) -> PyResult<Option<Py<PyAny>>> {
provider_config::passthrough_response(model, endpoint, body)
.map_err(errors::to_pyerr)?
.map(|response| to_py(py, &response.into_json()))
.transpose()
}
#[cfg(test)]
mod tests {
use pyo3::prelude::*;

View file

@ -1,6 +1,7 @@
use std::{future::Future, pin::Pin};
use std::{future::Future, pin::Pin, sync::Arc};
use litellm_core_utils::settings::Lookup;
use litellm_host_python::{PythonContext, attach_blocking};
use litellm_secrets::{
Error, ExternalSecretManager, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue,
};
@ -19,6 +20,11 @@ const ENVIRONMENT_FALLBACK_LOG: &str =
/// A secret manager whose reads execute in Python: a custom manager, a legacy compatible
/// client, or a manually assigned SDK client.
pub(crate) struct PythonSecretManager {
client: Arc<PythonClient>,
context: PythonContext,
}
struct PythonClient {
client: Py<PyAny>,
system: Option<KeyManagementSystem>,
settings: Option<Py<PyAny>>,
@ -29,14 +35,20 @@ impl PythonSecretManager {
client: Py<PyAny>,
system: Option<KeyManagementSystem>,
settings: Option<Py<PyAny>>,
context: PythonContext,
) -> Self {
Self {
client,
system,
settings,
client: Arc::new(PythonClient {
client,
system,
settings,
}),
context,
}
}
}
impl PythonClient {
fn read(&self, py: Python<'_>, name: &str) -> PyResult<Option<String>> {
let client = self.client.bind(py);
let kwargs = PyDict::new(py);
@ -76,7 +88,7 @@ fn python_name(system: KeyManagementSystem) -> &'static str {
impl ExternalSecretManager for PythonSecretManager {
fn system(&self) -> KeyManagementSystem {
self.system.unwrap_or(KeyManagementSystem::Custom)
self.client.system.unwrap_or(KeyManagementSystem::Custom)
}
fn read_secret<'a>(
@ -85,18 +97,26 @@ impl ExternalSecretManager for PythonSecretManager {
_settings: &'a KeyManagementSettings,
_environment: &'a (dyn Lookup + Send + Sync),
) -> Pin<Box<dyn Future<Output = Result<Option<Secret>, Error>> + Send + 'a>> {
let client = Arc::clone(&self.client);
let context = self.context.clone();
let name = name.to_owned();
Box::pin(async move {
Python::attach(|py| match self.read(py, name) {
match attach_blocking(context, move |py| match client.read(py, &name) {
Ok(value) => Ok(value.map(SecretValue::new).map(Secret::String)),
// `get_secret` answers a failed manager read from the process environment, but
// only for `Exception`: cancellation and other `BaseException`s propagate.
Err(error) if error.is_instance_of::<PyException>(py) => {
log_environment_fallback(py, name, &error)
log_environment_fallback(py, &name, &error)
.map_err(|error| external_error(py, error))?;
Err(read_error(py, error))
}
Err(error) => Err(external_error(py, error)),
})
.await
{
Ok(result) => result,
Err(error) => Python::attach(|py| Err(external_error(py, error))),
}
})
}
}
@ -116,8 +136,9 @@ fn log_environment_fallback(py: Python<'_>, name: &str, error: &PyErr) -> PyResu
}
#[cfg(test)]
#[allow(clippy::await_holding_lock)]
mod tests {
use std::sync::Arc;
use std::sync::{Arc, Mutex, MutexGuard};
use litellm_secrets::{
FailurePolicy, KeyManagementSettings, KeyManagementSystem, OidcResolver, SecretManager,
@ -126,15 +147,27 @@ mod tests {
use pyo3::{prelude::*, types::PyDict};
use rstest::rstest;
use litellm_host_python::PythonContext;
use super::{HANDLER_MODULE, PythonSecretManager, python_name};
use crate::secrets::python_error;
/// `sys.modules` is interpreter-global, so tests that install or rely on the handler module
/// cannot overlap with any other test on this list.
static HANDLER_LOCK: Mutex<()> = Mutex::new(());
fn handler_guard() -> MutexGuard<'static, ()> {
HANDLER_LOCK.lock().expect("handler lock poisoned")
}
/// A resolver over a Python manager whose reads raise `failure_type`, with the chained
/// exceptions Python attaches, and `fallback` as the process environment.
/// exceptions Python attaches, and `fallback` as the process environment. The returned guard
/// keeps other module-mutating tests out for the lifetime of the returned resolver.
fn failing_resolver(
failure_type: &str,
fallback: Option<&'static str>,
) -> (SecretResolver, Py<PyDict>) {
) -> (SecretResolver, Py<PyDict>, MutexGuard<'static, ()>) {
let handler = handler_guard();
Python::initialize();
let (reader, locals) = Python::attach(|py| {
let locals = PyDict::new(py);
@ -167,6 +200,7 @@ handler.get_secret_from_manager = get_secret_from_manager
locals.get_item("manager").unwrap().unwrap().unbind(),
None,
None,
PythonContext::capture(py).unwrap(),
);
(reader, locals.unbind())
});
@ -179,7 +213,7 @@ handler.get_secret_from_manager = get_secret_from_manager
OidcResolver::default(),
)
.with_failure_policy(FailurePolicy::EnvironmentFallback);
(resolver, locals)
(resolver, locals, handler)
}
#[rstest]
@ -191,7 +225,7 @@ handler.get_secret_from_manager = get_secret_from_manager
#[case] failure_type: &str,
#[case] fallback: Option<&'static str>,
) {
let (resolver, locals) = failing_resolver(failure_type, fallback);
let (resolver, locals, _handler) = failing_resolver(failure_type, fallback);
let error = resolver.get_secret("API_KEY", None).await.unwrap_err();
Python::attach(|py| {
let original = python_error(py, &error).unwrap();
@ -264,7 +298,7 @@ sys.modules.setdefault('litellm._logging', logging)
#[case] fallback: Option<&'static str>,
#[case] name: &str,
) {
let (resolver, _locals) = failing_resolver(failure_type, fallback);
let (resolver, _locals, _handler) = failing_resolver(failure_type, fallback);
Python::attach(|py| assert!(logged_errors(py, name).is_empty()));
let secret = resolver.get_secret(name, None).await.unwrap();
assert_eq!(
@ -290,6 +324,7 @@ sys.modules.setdefault('litellm._logging', logging)
/// Installs a fake `get_secret_from_manager` that records its kwargs, runs `body`, and
/// removes the fake handler again; parent package stubs persist for concurrent tests.
/// Callers hold `handler_guard` before attaching so the GIL is never held while waiting on it.
fn with_fake_handler<'py>(py: Python<'py>, body: impl FnOnce(&Bound<'py, PyDict>)) {
let locals = PyDict::new(py);
py.run(
@ -330,6 +365,7 @@ else:
#[case("123")]
#[case("{'key': 'value'}")]
fn nonstring_results_are_absent_without_a_read_failure(#[case] expression: &str) {
let _handler = handler_guard();
Python::initialize();
Python::attach(|py| {
with_fake_handler(py, |locals| {
@ -340,8 +376,13 @@ else:
Some(locals),
)
.unwrap();
let reader = PythonSecretManager::new(py.None(), None, None);
assert_eq!(reader.read(py, "KEY").unwrap(), None);
let reader = PythonSecretManager::new(
py.None(),
None,
None,
PythonContext::capture(py).unwrap(),
);
assert_eq!(reader.client.read(py, "KEY").unwrap(), None);
});
});
}
@ -365,6 +406,7 @@ else:
#[test]
fn configured_systems_dispatch_through_the_python_handler_with_the_original_settings() {
let _handler = handler_guard();
Python::initialize();
Python::attach(|py| {
with_fake_handler(py, |locals| {
@ -374,9 +416,10 @@ else:
client.clone().unbind(),
Some(KeyManagementSystem::AzureKeyVault),
Some(settings.clone().unbind()),
PythonContext::capture(py).unwrap(),
);
assert_eq!(
reader.read(py, "API_KEY").unwrap().as_deref(),
reader.client.read(py, "API_KEY").unwrap().as_deref(),
Some("handled-API_KEY")
);
assert!(py.import(HANDLER_MODULE).is_ok());
@ -416,6 +459,7 @@ else:
#[case] system: Option<KeyManagementSystem>,
#[case] key_manager: &str,
) {
let _handler = handler_guard();
Python::initialize();
Python::attach(|py| {
with_fake_handler(py, |locals| {
@ -434,9 +478,14 @@ manager = Manager()
)
.unwrap();
let manager = locals.get_item("manager").unwrap().unwrap();
let reader = PythonSecretManager::new(manager.clone().unbind(), system, None);
let reader = PythonSecretManager::new(
manager.clone().unbind(),
system,
None,
PythonContext::capture(py).unwrap(),
);
assert_eq!(
reader.read(py, "API_KEY").unwrap().as_deref(),
reader.client.read(py, "API_KEY").unwrap().as_deref(),
Some("handled-API_KEY")
);
assert_eq!(

View file

@ -5,6 +5,8 @@ use litellm_secrets_types::{AccessMode, KeyManagementSettings, KeyManagementSyst
use pyo3::prelude::*;
use serde_json::Value;
use litellm_host_python::PythonContext;
use super::callback::PythonSecretManager;
use crate::{
coercion::{Field, FieldSpec, ProjectionError},
@ -87,7 +89,7 @@ pub(crate) struct SecretManagerSnapshot {
}
impl SecretManagerSnapshot {
pub(crate) fn into_state(self) -> Arc<SecretManagerState> {
pub(crate) fn into_state(self, context: PythonContext) -> Arc<SecretManagerState> {
match self.client {
SecretManagerClient::Native(backend) => {
Arc::new(SecretManagerState::new(*backend, self.settings))
@ -98,6 +100,7 @@ impl SecretManagerSnapshot {
client,
self.system,
self.settings_object,
context,
))),
self.settings,
)),

View file

@ -4,102 +4,30 @@ mod error;
mod mutation;
mod operations;
mod provider;
mod python;
pub(crate) mod resolved;
pub(crate) mod runtime;
mod vault;
use std::sync::Arc;
use litellm_secrets::source::{EnvironmentSecrets, SecretSource};
use pyo3::prelude::*;
pub(crate) use error::python_error;
use litellm_secrets::source::SecretSource;
use pyo3::prelude::*;
use python::PythonSecrets;
use resolved::ResolvedSecrets;
use crate::{
coercion::FieldSpec,
errors::RustBridgeDeclined,
python_settings::{PythonSettings, Snapshot},
};
use crate::{coercion::FieldSpec, python_settings::PythonSettings};
const READABLE: FieldSpec<bool> = FieldSpec::new("readable", |field| field.schema_bool());
const NATIVE: FieldSpec<bool> = FieldSpec::new("native", |field| field.schema_bool());
/// Where a Rust route reads provider secrets from, as `litellm.get_secret` would.
/// Where a Rust route reads provider secrets from. Python's `get_secret_str` until a
/// `SecretManagerRule` in `catalog.py` moves the configured system off `PYTHON_ONLY`, then the
/// native secret manager.
pub(crate) fn source(py: Python<'_>) -> PyResult<Arc<dyn SecretSource>> {
select(&PythonSettings::SecretManager.read(py)?, || {
Ok(Arc::new(ResolvedSecrets::new(config::read(py)?)))
})
}
fn select(
manager: &Snapshot<'_>,
resolved: impl FnOnce() -> PyResult<Arc<dyn SecretSource>>,
) -> PyResult<Arc<dyn SecretSource>> {
if !manager.read(&READABLE)? {
return Ok(Arc::new(EnvironmentSecrets::python_compatible()));
}
if !manager.read(&NATIVE)? {
return Err(RustBridgeDeclined::new_err(
"the configured secret manager is not enabled for the Rust bridge",
));
}
resolved()
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use litellm_secrets::source::{EnvironmentSecrets, SecretSource};
use pyo3::{prelude::*, types::PyDict};
use rstest::rstest;
use super::select;
use crate::{errors::RustBridgeDeclined, python_settings::PythonSettings};
enum Selected {
Environment,
Declined,
Resolved,
}
#[rstest]
#[case::unreadable(false, false, Selected::Environment)]
#[case::unreadable_even_if_native(false, true, Selected::Environment)]
#[case::readable_python_only(true, false, Selected::Declined)]
#[case::readable_native(true, true, Selected::Resolved)]
fn readable_and_native_select_the_secret_source(
#[case] readable: bool,
#[case] native: bool,
#[case] expected: Selected,
) {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
locals.set_item("readable", readable).unwrap();
locals.set_item("native", native).unwrap();
let manager = py
.eval(
c"__import__('types').SimpleNamespace(readable=readable, native=native)",
None,
Some(&locals),
)
.unwrap();
let mut resolved_called = false;
let selected = select(&PythonSettings::SecretManager.snapshot(manager), || {
resolved_called = true;
Ok(Arc::new(EnvironmentSecrets::python_compatible()) as Arc<dyn SecretSource>)
});
match expected {
Selected::Environment => assert!(selected.is_ok() && !resolved_called),
Selected::Resolved => assert!(selected.is_ok() && resolved_called),
Selected::Declined => {
let error = selected.err().expect("the Rust route declines");
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
assert!(!resolved_called);
}
}
});
}
if PythonSettings::SecretManager.read(py)?.read(&NATIVE)? {
let context = litellm_host_python::PythonContext::capture(py)?;
return Ok(Arc::new(ResolvedSecrets::new(config::read(py)?, context)));
}
Ok(Arc::new(PythonSecrets::new(py)?))
}

View file

@ -0,0 +1,190 @@
use std::sync::Arc;
use futures_util::future::BoxFuture;
use litellm_host_python::{PythonContext, attach_blocking};
use litellm_secrets::{Error, SecretValue, source::SecretSource};
use pyo3::prelude::*;
use super::error::external_error;
/// Reads each secret through Python's `get_secret_str`, so the configured manager, the key
/// management settings and the environment fallback behave exactly as they do in Python.
pub(super) struct PythonSecrets {
get_secret_str: Arc<Py<PyAny>>,
context: PythonContext,
}
impl PythonSecrets {
pub(super) fn new(py: Python<'_>) -> PyResult<Self> {
Ok(Self::reading_with(
py.import("litellm.secret_managers.main")?
.getattr("get_secret_str")?
.unbind(),
PythonContext::capture(py)?,
))
}
fn reading_with(get_secret_str: Py<PyAny>, context: PythonContext) -> Self {
Self {
get_secret_str: Arc::new(get_secret_str),
context,
}
}
}
impl SecretSource for PythonSecrets {
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> BoxFuture<'a, Result<Option<SecretValue>, Error>> {
let get_secret_str = Arc::clone(&self.get_secret_str);
let context = self.context.clone();
let name = name.to_owned();
Box::pin(async move {
match attach_blocking(context, move |py| {
get_secret_str
.bind(py)
.call1((name,))
.and_then(|value| value.extract::<Option<String>>())
.map(|value| value.map(SecretValue::new))
.map_err(|error| external_error(py, error))
})
.await
{
Ok(result) => result,
Err(error) => Python::attach(|py| Err(external_error(py, error))),
}
})
}
}
#[cfg(test)]
mod tests {
use litellm_secrets::source::SecretSource;
use pyo3::{prelude::*, types::PyDict};
use rstest::{fixture, rstest};
use super::PythonSecrets;
use crate::secrets::python_error;
use litellm_host_python::PythonContext;
#[fixture]
fn namespace() -> Py<PyDict> {
Python::initialize();
Python::attach(|py| {
let namespace = PyDict::new(py);
py.run(
c"
import contextvars
import threading
read_on = None
request_var = contextvars.ContextVar('request_var', default=None)
seen_context_values = []
raised = KeyboardInterrupt('secret manager stopped')
def get_secret_str(name):
global read_on
read_on = threading.get_ident()
seen_context_values.append(request_var.get())
if name == 'RAISING':
raise raised
return {'MISTRAL_API_KEY': 'vault-key'}.get(name)
",
Some(&namespace),
None,
)
.unwrap();
namespace.unbind()
})
}
#[fixture]
fn secrets(namespace: Py<PyDict>) -> (PythonSecrets, Py<PyDict>) {
let (reader, context) = Python::attach(|py| {
let namespace = namespace.bind(py);
namespace
.get_item("request_var")
.unwrap()
.unwrap()
.call_method1("set", ("request-value",))
.unwrap();
(
namespace
.get_item("get_secret_str")
.unwrap()
.unwrap()
.unbind(),
PythonContext::capture(py).unwrap(),
)
});
(PythonSecrets::reading_with(reader, context), namespace)
}
fn global<T: for<'a, 'py> FromPyObject<'a, 'py, Error: std::fmt::Debug>>(
namespace: &Py<PyDict>,
py: Python<'_>,
name: &str,
) -> T {
namespace
.bind(py)
.get_item(name)
.unwrap()
.unwrap()
.extract()
.unwrap()
}
#[rstest]
#[case::found("MISTRAL_API_KEY", Some("vault-key"))]
#[case::missing("OTHER", None)]
#[tokio::test]
async fn returns_what_get_secret_str_returns(
secrets: (PythonSecrets, Py<PyDict>),
#[case] name: &str,
#[case] expected: Option<&str>,
) {
let value = secrets.0.get_secret_str(name).await.unwrap();
assert_eq!(value.as_ref().map(|value| value.expose()), expected);
}
#[rstest]
#[tokio::test]
async fn exceptions_surface_as_the_original_python_object(
secrets: (PythonSecrets, Py<PyDict>),
) {
let error = secrets.0.get_secret_str("RAISING").await.unwrap_err();
Python::attach(|py| {
let surfaced = python_error(py, &error).expect("the Python exception is preserved");
let raised: Py<PyAny> = global(&secrets.1, py, "raised");
assert!(surfaced.value(py).is(raised.bind(py)));
});
}
#[rstest]
#[tokio::test]
async fn reads_run_off_the_thread_polling_the_route(secrets: (PythonSecrets, Py<PyDict>)) {
let polling: u64 = Python::attach(|py| {
py.import("threading")
.unwrap()
.call_method0("get_ident")
.unwrap()
.extract()
.unwrap()
});
secrets.0.get_secret_str("MISTRAL_API_KEY").await.unwrap();
let read_on: u64 = Python::attach(|py| global(&secrets.1, py, "read_on"));
assert_ne!(read_on, polling);
}
#[rstest]
#[tokio::test]
async fn reads_see_the_callers_contextvars(secrets: (PythonSecrets, Py<PyDict>)) {
secrets.0.get_secret_str("MISTRAL_API_KEY").await.unwrap();
let seen: Vec<String> = Python::attach(|py| global(&secrets.1, py, "seen_context_values"));
assert_eq!(seen, vec!["request-value".to_owned()]);
}
}

View file

@ -2,6 +2,7 @@ use std::sync::Arc;
use futures_util::future::BoxFuture;
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_host_python::PythonContext;
use litellm_secrets::source::SecretSource;
use litellm_secrets::{
Error, FailurePolicy, OidcResolver, SecretManagerState, SecretResolver, SecretValue,
@ -14,8 +15,8 @@ pub(crate) struct ResolvedSecrets {
}
impl ResolvedSecrets {
pub(crate) fn new(snapshot: SecretManagerSnapshot) -> Self {
Self::from_state(snapshot.into_state())
pub(crate) fn new(snapshot: SecretManagerSnapshot, context: PythonContext) -> Self {
Self::from_state(snapshot.into_state(context))
}
fn from_state(state: Arc<SecretManagerState>) -> Self {

View file

@ -394,12 +394,10 @@ UTILS_MODULE_NAMES: Final = (
"redact_message_input_output_from_logging",
"CustomStreamWrapper",
"BaseGoogleGenAIGenerateContentConfig",
"BaseOCRConfig",
"BaseSearchConfig",
"BaseTextToSpeechConfig",
"BedrockModelInfo",
"CohereModelInfo",
"MistralOCRConfig",
"Rules",
"AsyncHTTPHandler",
"HTTPHandler",
@ -1367,7 +1365,6 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
"litellm.llms.base_llm.google_genai.transformation",
"BaseGoogleGenAIGenerateContentConfig",
),
"BaseOCRConfig": ("litellm.llms.base_llm.ocr.transformation", "BaseOCRConfig"),
"BaseSearchConfig": (
"litellm.llms.base_llm.search.transformation",
"BaseSearchConfig",
@ -1378,7 +1375,6 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
),
"BedrockModelInfo": ("litellm.llms.bedrock.common_utils", "BedrockModelInfo"),
"CohereModelInfo": ("litellm.llms.cohere.common_utils", "CohereModelInfo"),
"MistralOCRConfig": ("litellm.llms.mistral.ocr.transformation", "MistralOCRConfig"),
"Rules": ("litellm.litellm_core_utils.rules", "Rules"),
"AsyncHTTPHandler": ("litellm.llms.custom_httpx.http_handler", "AsyncHTTPHandler"),
"HTTPHandler": ("litellm.llms.custom_httpx.http_handler", "HTTPHandler"),

View file

@ -557,6 +557,8 @@ SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS: Final[float] = float(
request_timeout: float = float(os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS))))
request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ
DEFAULT_A2A_AGENT_TIMEOUT: Final[float] = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes
AGENT_KILL_SWITCH_TIMEOUT_SECONDS: Final = 10.0
AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS: Final = 2000
# Patterns that indicate a localhost/internal URL in A2A agent cards that should be
# replaced with the original base_url. This is a common misconfiguration where
# developers deploy agents with development URLs in their agent cards.
@ -2120,6 +2122,14 @@ PTU_LAPSED_ALERT_LIMIT: Final[int] = 10
DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID: Final[str] = "daily_global_spend_reconcile_job"
DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS: Final[int] = 3600
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM: Final[str] = "daily_global_spend_reconciled_through"
SPEND_CAPTURE_RATE_CHECK_JOB_ID: Final[str] = "spend_capture_rate_check_job"
SPEND_CAPTURE_RATE_CHECK_LOCK_TTL_SECONDS: Final[int] = 900
SPEND_CAPTURE_RATE_MAX_RANGE_DAYS: Final[int] = 180
SPEND_CAPTURE_RATE_DOCS_URL: Final[str] = "https://docs.litellm.ai/docs/proxy/spend_capture_rate"
OPENAI_ORGANIZATION_COSTS_URL: Final[str] = "https://api.openai.com/v1/organization/costs"
# Buckets per page the OpenAI costs endpoint allows (1 to 180, default 7), 2026-09-24
OPENAI_ORGANIZATION_COSTS_PAGE_LIMIT: Final[int] = 180
PROVIDER_BILLING_TIMEOUT_SECONDS: Final[float] = 30.0
# Slack allowed when deciding a sentinel row is stale. The row's updated_at and the
# run's cutoff are stamped by different hosts, so clock skew between them must not let
# one run delete a charge another just wrote. A stale row is hours old and a concurrent

View file

@ -421,6 +421,24 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm
pass
async def async_filter_listed_models(
self,
user_api_key_dict: UserAPIKeyAuth,
model_names: Sequence[str],
) -> Sequence[str]:
"""Runs on the model listing routes (`/v1/models`, `/v1/models/{id}`, `/model/info`,
`/model_group/info`) with the public model names the route would otherwise return, so a
lookup of one model may offer just that name: decide per name, never by position in the
sequence. Return the names to keep as a sequence of strings; a name left out disappears
from every listing, any alias of it offered in the same call goes with it, and
`/v1/models/{id}` answers 404 for it, exactly as for a model that does not exist. Names
outside `model_names` are ignored, so a callback can only narrow the listing, never widen
it. Under `use_team_public_model_name: false`, `/v1/models` and `/model_group/info` list a
team model by its internal routing name while `/model/info` keeps its public name, so hide
both names to hide it on every route.
"""
return model_names
async def async_post_call_response_headers_hook(
self,
data: dict,

View file

@ -729,6 +729,15 @@ class PrometheusLogger(CustomLogger):
labelnames=self.get_labels_for_metric("litellm_zero_cost_requests_total"),
)
self.litellm_spend_capture_rate = self._gauge_factory(
"litellm_spend_capture_rate",
(
"Share of the provider's bill LiteLLM captured as spend over the scheduled check's window "
"(captured spend / provider bill), by api_provider; NaN when the last check produced no rate"
),
labelnames=self.get_labels_for_metric("litellm_spend_capture_rate"),
)
# Cache metrics
self.litellm_cache_hits_metric = self._counter_factory(
name="litellm_cache_hits_metric",
@ -2028,6 +2037,15 @@ class PrometheusLogger(CustomLogger):
)
self.litellm_zero_cost_requests_total.labels(**labels).inc()
def set_spend_capture_rate(self, api_provider: str, capture_rate: float | None) -> None:
labels: Final = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric("litellm_spend_capture_rate"),
enum_values=UserAPIKeyLabelValues(api_provider=api_provider),
)
gauge: Final = self.litellm_spend_capture_rate
series: Final = gauge.labels(**labels) if labels else gauge
series.set(math.nan if capture_rate is None else capture_rate)
@staticmethod
def _get_remaining_from_v3_rate_limit_headers(
standard_logging_payload: StandardLoggingPayload | None,

View file

@ -139,6 +139,18 @@ def declared_authenticating_provider(model: str | None, custom_llm_provider: str
return declared if declared in PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO else None
def inferred_provider(model: str | None) -> str | None:
if not model:
return None
declared: Final = declared_authenticating_provider(model)
if declared is not None:
return declared
try:
return get_llm_provider(model=model)[1]
except Exception: # noqa: BLE001 # get_llm_provider raises for an unknown name, which then has no provider
return None
def get_llm_provider(
model: str,
custom_llm_provider: str | None = None,

View file

@ -6,8 +6,10 @@ import base64
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Final, Literal
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, DocumentType
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS, LlmProviders
from litellm.llms.base_llm.ocr.transformation import DocumentType
from litellm.rust_bridge import runtime
from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR_HEALTH_CHECK_DOCUMENT
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
@ -29,11 +31,12 @@ def get_image_file_for_health_check() -> bytes:
def _ocr_health_check_document(model: str, custom_llm_provider: str) -> DocumentType:
from litellm.utils import ProviderConfigManager
provider: Final = next((known for known in LlmProviders if known.value == custom_llm_provider), None)
config: Final = ProviderConfigManager.get_provider_ocr_config(model=model, provider=provider) if provider else None
return (config or BaseOCRConfig()).get_health_check_document()
native: Final = NATIVE_OCR_HEALTH_CHECK_DOCUMENT.load()
if native is None:
raise runtime.NoPythonImplementationError(
"ocr health check documents are resolved by the Rust extension, which is not available"
)
return native(model, custom_llm_provider)
class HealthCheckHelpers:

View file

@ -385,6 +385,10 @@ def _get_cached_prometheus_logger():
return _PrometheusLogger
class RawRequestCaptured(Exception):
pass
_DEPLOYMENT_PRICING_KEYS: Final = (
"input_cost_per_token",
"output_cost_per_token",
@ -591,6 +595,7 @@ class Logging(LiteLLMLoggingBaseClass):
kwargs: dict | None = None,
log_raw_request_response: bool = False,
supports_correlation_logging: bool = True,
raw_request_only: bool = False,
):
_input: Final[str | None] = messages # save original value of messages
if messages is not None:
@ -650,6 +655,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.streaming_chunks: list[Any] = [] # for generating complete stream response
self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response
self.log_raw_request_response = log_raw_request_response
self.raw_request_only = raw_request_only
# Initialize dynamic callbacks
self.dynamic_input_callbacks: list[str | Callable | CustomLogger] | None = dynamic_input_callbacks
@ -1476,6 +1482,9 @@ class Logging(LiteLLMLoggingBaseClass):
if capture_exception: # log this error to sentry for debugging
capture_exception(e)
if self.raw_request_only:
raise RawRequestCaptured()
def _print_llm_call_debugging_log(
self,
api_base: str,

View file

@ -1,15 +0,0 @@
"""Azure AI OCR module."""
from .cohere_parse_transformation import AzureAICohereParseConfig
from .common_utils import get_azure_ai_ocr_config
from .document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
from .transformation import AzureAIOCRConfig
__all__ = [
"AzureAICohereParseConfig",
"AzureAIOCRConfig",
"AzureDocumentIntelligenceOCRConfig",
"get_azure_ai_ocr_config",
]

View file

@ -1,91 +0,0 @@
"""Cohere Parse served from Azure AI Foundry (`/providers/cohere/v2/parse`)."""
from collections.abc import Mapping
from typing import Final
import httpx
from litellm.litellm_core_utils.prompt_templates.image_handling import (
async_convert_url_to_base64,
convert_url_to_base64,
)
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
from litellm.llms.cohere.ocr.transformation import COHERE_PARSE_PATH, CohereParseConfig
from litellm.secret_managers.main import get_secret_str
AZURE_AI_API_KEY_ENV_VAR: Final = "AZURE_AI_API_KEY"
AZURE_AI_API_BASE_ENV_VAR: Final = "AZURE_AI_API_BASE"
AZURE_AI_COHERE_PROVIDER_PATH: Final = "/providers/cohere"
AZURE_AI_MODELS_PATH_SUFFIX: Final = "/models"
class AzureAICohereParseConfig(CohereParseConfig):
"""Same request and response shape as Cohere Parse, behind Azure AI auth and URL layout.
Foundry cannot fetch external URLs, so remote images are inlined as base64 data URIs.
"""
def get_api_key_env_var(self) -> str | None:
return AZURE_AI_API_KEY_ENV_VAR
def _llm_provider(self) -> str:
return "azure_ai"
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: Mapping[str, object] | None = None,
**kwargs: object, # kwargs-ok: BaseOCRConfig.validate_environment signature
) -> dict[str, str]: # mutable-ok: BaseOCRConfig signature
resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR)
if resolved_base is None:
raise ValueError(
f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable "
"or pass api_base parameter"
)
resolved_key: Final = api_key or get_secret_str(AZURE_AI_API_KEY_ENV_VAR)
return { # mutable-ok: BaseOCRConfig signature
**get_azure_ai_auth_headers(api_key=resolved_key, litellm_params=litellm_params),
"Content-Type": "application/json",
**headers,
}
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object] | None = None,
**kwargs: object, # kwargs-ok: BaseOCRConfig.get_complete_url signature
) -> str:
resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR)
if resolved_base is None:
raise ValueError(
f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable "
"or pass api_base parameter"
)
url: Final = httpx.URL(resolved_base)
if not url.is_absolute_url:
raise ValueError(
"Azure AI API Base must be an absolute URL including scheme (e.g. "
f"'https://<resource>.services.ai.azure.com'). Got api_base={resolved_base!r}."
)
path: Final = url.path.rstrip("/")
if path.endswith(COHERE_PARSE_PATH):
return str(url.copy_with(path=path))
if path.endswith(f"{AZURE_AI_COHERE_PROVIDER_PATH}/v2"):
return str(url.copy_with(path=f"{path}/parse"))
return str(
url.copy_with(
path=f"{path.removesuffix(AZURE_AI_MODELS_PATH_SUFFIX)}{AZURE_AI_COHERE_PROVIDER_PATH}{COHERE_PARSE_PATH}"
)
)
def _resolve_image_url_sync(self, image_url: str) -> str:
return convert_url_to_base64(image_url)
async def _resolve_image_url_async(self, image_url: str) -> str:
return await async_convert_url_to_base64(image_url)

View file

@ -1,71 +0,0 @@
"""
Common utilities for Azure AI OCR providers.
This module provides routing logic to determine which OCR configuration to use
based on the model name.
"""
from typing import TYPE_CHECKING, Final, Optional
from litellm._logging import verbose_logger
if TYPE_CHECKING:
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
def is_azure_document_intelligence_model(model: str) -> bool:
"""Whether an azure_ai OCR model routes to Azure Document Intelligence.
Azure AI exposes two OCR services on the same provider; the sub-route in the
model name (`azure_ai/doc-intelligence/<model>`) selects Document Intelligence
over Mistral OCR. This is the single source of truth for that routing decision.
"""
lowered: Final = model.lower()
return "doc-intelligence" in lowered or "documentintelligence" in lowered
def is_azure_cohere_parse_model(model: str) -> bool:
lowered: Final = model.lower()
return "cohere" in lowered and "parse" in lowered
def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
"""
Determine which Azure AI OCR configuration to use based on the model name.
Azure AI supports multiple OCR services:
- Azure Document Intelligence: azure_ai/doc-intelligence/<model>
- Mistral OCR (via Azure AI): azure_ai/<model>
Args:
model: The model name (e.g., "azure_ai/doc-intelligence/prebuilt-read",
"azure_ai/pixtral-12b-2409")
Returns:
OCR configuration instance for the specified model
Examples:
>>> get_azure_ai_ocr_config("azure_ai/doc-intelligence/prebuilt-read")
<AzureDocumentIntelligenceOCRConfig object>
>>> get_azure_ai_ocr_config("azure_ai/pixtral-12b-2409")
<AzureAIOCRConfig object>
"""
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
# Check for Azure Document Intelligence models
if is_azure_document_intelligence_model(model):
verbose_logger.debug("Routing %s to Azure Document Intelligence OCR config", model)
return AzureDocumentIntelligenceOCRConfig()
if is_azure_cohere_parse_model(model):
verbose_logger.debug("Routing %s to Azure AI Cohere Parse config", model)
return AzureAICohereParseConfig()
# Default to Mistral-based OCR for other azure_ai models
verbose_logger.debug("Routing %s to Azure AI (Mistral) OCR config", model)
return AzureAIOCRConfig()

View file

@ -1,5 +0,0 @@
"""Azure Document Intelligence OCR module."""
from .transformation import AzureDocumentIntelligenceOCRConfig
__all__ = ["AzureDocumentIntelligenceOCRConfig"]

View file

@ -1,806 +0,0 @@
"""
Azure Document Intelligence OCR transformation implementation.
Azure Document Intelligence (formerly Form Recognizer) provides advanced document analysis capabilities.
This implementation transforms between Mistral OCR format and Azure Document Intelligence API v4.0.
Note: Azure Document Intelligence API is async - POST returns 202 Accepted with Operation-Location header.
The operation location must be polled until the analysis completes.
"""
import asyncio
import re
import time
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final
from urllib.parse import quote
import httpx
from pydantic import BaseModel
from litellm._logging import verbose_logger
from litellm.constants import (
AZURE_DOCUMENT_INTELLIGENCE_API_VERSION,
AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI,
AZURE_OPERATION_POLLING_TIMEOUT,
)
from litellm.exceptions import UnsupportedParamsError
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin, encode_url_path_segment
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
from litellm.llms.base_llm.ocr.transformation import (
OCR_REQUEST_FORMAT_PARAM,
BaseOCRConfig,
DocumentType,
OCRPage,
OCRPageDimensions,
OCRRequestData,
OCRRequestFormat,
OCRResponse,
OCRUsageInfo,
parse_ocr_request_format,
)
from litellm.secret_managers.main import get_secret_str
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR: Final = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"
class AzureDocumentIntelligenceLine(BaseModel):
content: str | None = None
class AzureDocumentIntelligencePage(BaseModel):
pageNumber: int | None = None
width: float | None = None
height: float | None = None
unit: str | None = None
lines: tuple[AzureDocumentIntelligenceLine, ...] = ()
class AzureDocumentIntelligenceAnalyzeResult(BaseModel):
content: str | None = None
pages: tuple[AzureDocumentIntelligencePage, ...] = ()
tables: list[dict[str, object]] | None = None
keyValuePairs: list[dict[str, object]] | None = None
class AzureDocumentIntelligenceOperation(BaseModel):
status: str | None = None
analyzeResult: AzureDocumentIntelligenceAnalyzeResult | None = None
class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
"""
Azure Document Intelligence OCR transformation configuration.
Supports Azure Document Intelligence v4.0 (2024-11-30) API.
Model route: azure_ai/doc-intelligence/<model>
Supported models:
- prebuilt-layout: Extracts text with markdown, tables, and structure (closest to Mistral OCR)
- prebuilt-read: Basic text extraction optimized for reading
- prebuilt-document: General document analysis
Reference: https://learn.microsoft.com/en-us/azure/ai-services/document-intelligence/
"""
def __init__(self) -> None:
super().__init__()
def get_api_key_env_var(self) -> str | None:
return AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR
def resolve_connection_params(
self,
*,
api_key: str | None,
api_base: str | None,
dynamic_api_key: str | None,
dynamic_api_base: str | None,
) -> tuple[str | None, str | None]:
explicit_api_key: Final = None if api_key is None else dynamic_api_key or api_key
explicit_api_base: Final = None if api_base is None else dynamic_api_base or api_base
return explicit_api_key, explicit_api_base
def get_supported_ocr_params(self, model: str) -> list:
"""
Get supported OCR parameters for Azure Document Intelligence.
Azure DI exposes a `pages` query parameter on the analyze endpoint
(1-based, e.g. "1-3,5,7-9"). To keep the public request shape
aligned with Mistral OCR, callers pass `pages` using Mistral
semantics — a list of 0-based integers — or a pre-formatted
Azure-style string. Azure DI also exposes a `features` query
parameter enabling add-on capabilities (e.g. "keyValuePairs",
"languages"), passed as a list of feature names or a
comma-separated string. Other Mistral-specific params (e.g.
`include_image_base64`) are not supported by Azure DI and are
ignored during transformation.
`req_format` selects the response shape: "litellm" (default) returns
the normalized OCR schema, "native" returns Azure DI's own analyze
operation payload as-is.
"""
return ["pages", "features", OCR_REQUEST_FORMAT_PARAM]
def map_ocr_params(
self,
non_default_params: Mapping[str, object],
optional_params: dict,
model: str,
) -> dict:
"""
Map OCR params to Azure DI format.
Translates Mistral-style `pages` (list[int], 0-based) into Azure's
`pages` query string (1-based, e.g. "1,2,3" or "1-3,5"). A raw
string that already matches Azure's format is passed through
unchanged. `features` (list[str] or comma-separated string) is
normalized into Azure's comma-joined `features` query string.
"""
pages: Final = non_default_params.get("pages")
features: Final = non_default_params.get("features")
request_format: Final = non_default_params.get(OCR_REQUEST_FORMAT_PARAM)
normalized_pages: Final = self._normalize_pages_param(pages) if pages is not None else ""
normalized_features: Final = self._normalize_features_param(features) if features is not None else ""
return {
**optional_params,
**({"pages": normalized_pages} if normalized_pages else {}),
**({"features": normalized_features} if normalized_features else {}),
**(
{OCR_REQUEST_FORMAT_PARAM: self._parse_request_format(request_format, model)}
if request_format is not None
else {}
),
}
@staticmethod
def _parse_request_format(request_format: object, model: str) -> OCRRequestFormat:
try:
return parse_ocr_request_format(request_format)
except ValueError as e:
raise UnsupportedParamsError(message=f"{e}", model=model, llm_provider="azure_ai") from e
@staticmethod
def _normalize_pages_param(pages: object) -> str:
"""
Convert a caller-provided `pages` value to Azure DI's query-string
form. Azure expects 1-based page numbers, grammar: `^(\\d+(-\\d+)?)(,\\s*(\\d+(-\\d+)?))*$`.
Accepted inputs:
- list[int]: Mistral-style 0-based indices. Converted to 1-based
and joined (e.g. [0,1,2] -> "1,2,3").
- list[str]: tokens like "1" or "3-5". Validated, joined as-is
(treated as Azure-native, i.e. 1-based).
- str: already in Azure format. Validated and whitespace-stripped.
"""
pages_pattern: Final = re.compile(r"^\s*\d+(-\d+)?(\s*,\s*\d+(-\d+)?)*\s*$")
if isinstance(pages, str):
if not pages_pattern.match(pages):
raise ValueError(
f"Invalid `pages` string for Azure Document Intelligence: "
f"{pages!r}. Expected format like '1-3,5,7-9'."
)
return pages.replace(" ", "")
if isinstance(pages, list):
if len(pages) == 0:
return ""
if any(isinstance(p, bool) for p in pages):
raise ValueError("`pages` must be integers, not booleans")
if all(isinstance(p, int) for p in pages):
if any(p < 0 for p in pages):
raise ValueError("`pages` integers must be >= 0 (Mistral 0-based indices)")
# Mistral 0-based -> Azure 1-based.
return ",".join(str(p + 1) for p in sorted(set(pages)))
if all(isinstance(p, str) for p in pages):
joined: Final = ",".join(p.strip() for p in pages)
if not pages_pattern.match(joined):
raise ValueError(
f"Invalid `pages` list for Azure Document Intelligence: "
f"{pages!r}. Expected tokens like '1' or '3-5'."
)
return joined
raise ValueError("`pages` must be a list[int] (0-based, Mistral-style) or a string like '1-3,5,7-9'.")
@staticmethod
def _normalize_features_param(features: object) -> str:
"""
Convert a caller-provided `features` value to Azure DI's query-string
form (comma-joined feature names, e.g. "keyValuePairs,languages").
Accepted inputs:
- list[str]: feature names like ["keyValuePairs", "languages"].
- str: a single feature name or comma-separated names.
"""
invalid_features_error: Final = ValueError(
f"Invalid `features` for Azure Document Intelligence: {features!r}. "
f"Expected a list of feature names or a comma-separated string like "
f"'keyValuePairs' or 'keyValuePairs,languages'."
)
if isinstance(features, str):
raw_tokens = features.split(",")
elif isinstance(features, list):
if len(features) == 0:
return ""
raw_tokens = [feature for feature in features if isinstance(feature, str)]
if len(raw_tokens) != len(features):
raise invalid_features_error
else:
raise invalid_features_error
tokens: Final = tuple(token.strip() for token in raw_tokens)
feature_pattern: Final = re.compile(r"^[A-Za-z][A-Za-z0-9]*$")
if not all(feature_pattern.match(token) for token in tokens):
raise invalid_features_error
return ",".join(tokens)
def validate_environment(
self,
headers: dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> dict:
"""
Validate environment and return headers for Azure Document Intelligence.
Authentication uses the Ocp-Apim-Subscription-Key header, or an Entra ID / OAuth bearer
token when no subscription key is set.
"""
# Get API key from environment if not provided
if api_key is None:
api_key = get_secret_str(AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR)
# Validate API base/endpoint is provided
if api_base is None:
api_base = get_secret_str("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
if api_base is None:
raise ValueError(
"Missing Azure Document Intelligence Endpoint - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT environment variable or pass api_base parameter"
)
headers = {
**get_azure_ai_auth_headers(
api_key=api_key,
litellm_params=litellm_params,
api_key_header="Ocp-Apim-Subscription-Key",
api_key_env_var=AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR,
),
"Content-Type": "application/json",
**headers,
}
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Azure Document Intelligence endpoint.
Format: {endpoint}/documentintelligence/documentModels/{modelId}:analyze?api-version=2024-11-30
Note: API version 2024-11-30 uses /documentintelligence/ path (not /formrecognizer/)
Args:
api_base: Azure Document Intelligence endpoint (e.g., https://your-resource.cognitiveservices.azure.com)
model: Model ID (e.g., "prebuilt-layout", "prebuilt-read")
optional_params: Optional parameters
Returns: Complete URL for Azure DI analyze endpoint
"""
if api_base is None:
api_base = get_secret_str("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
if api_base is None:
raise ValueError(
"Missing Azure Document Intelligence Endpoint - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT environment variable or pass api_base parameter"
)
# Ensure no trailing slash
api_base = api_base.rstrip("/")
# Extract model ID from full model path if needed
# Model can be "prebuilt-layout" or "azure_ai/doc-intelligence/prebuilt-layout"
model_id = model
if "/" in model:
# Extract the last part after the last slash
model_id = model.split("/")[-1]
encoded_model_id: Final = encode_url_path_segment(model_id, field_name="model_id")
# Azure Document Intelligence analyze endpoint
# Note: API version 2024-11-30+ uses /documentintelligence/ (not /formrecognizer/)
url: Final = (
f"{api_base}/documentintelligence/documentModels/{encoded_model_id}:analyze"
f"?api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}"
)
# Azure DI accepts `pages` (1-based, e.g. "1-3,5") and `features`
# (comma-joined names, e.g. "keyValuePairs") as query params.
# `optional_params` has already been normalized in `map_ocr_params`.
pages: Final = optional_params.get("pages") if optional_params else None
features: Final = optional_params.get("features") if optional_params else None
pages_query: Final = f"&pages={quote(str(pages), safe=',-')}" if pages else ""
features_query: Final = f"&features={quote(str(features), safe=',')}" if features else ""
return f"{url}{pages_query}{features_query}"
def _extract_base64_from_data_uri(self, data_uri: str) -> str:
"""
Extract base64 content from a data URI.
Args:
data_uri: Data URI like "data:application/pdf;base64,..."
Returns:
Base64 string without the data URI prefix
"""
# Match pattern: data:[<mediatype>][;base64],<data>
match: Final = re.match(r"data:([^;]+)(?:;base64)?,(.+)", data_uri)
if match:
return match.group(2)
return data_uri
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request to Azure Document Intelligence format.
Mistral OCR format:
{
"document": {
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}
}
Azure DI format:
{
"urlSource": "https://example.com/doc.pdf"
}
OR
{
"base64Source": "base64_encoded_content"
}
Args:
model: Model name
document: Document dict from user (Mistral format)
optional_params: Already mapped optional parameters
headers: Request headers
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug("Azure Document Intelligence transform_ocr_request - model: %s", model)
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Extract document URL from Mistral format
doc_type: Final = document.get("type")
document_url = None
if doc_type == "document_url":
document_url = document.get("document_url", "")
elif doc_type == "image_url":
document_url = document.get("image_url", "")
else:
raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url' or 'image_url'")
if not document_url:
raise ValueError("Document URL is required")
# Build Azure DI request
data: Final[dict[str, str]] = {}
# Check if it's a data URI (base64)
if document_url.startswith("data:"):
# Extract base64 content
base64_content: Final = self._extract_base64_from_data_uri(document_url)
data["base64Source"] = base64_content
verbose_logger.debug("Using base64Source for Azure Document Intelligence")
else:
# Regular URL
data["urlSource"] = document_url
verbose_logger.debug("Using urlSource for Azure Document Intelligence")
# Azure DI: `pages` is a query param (wired in get_complete_url),
# not a body field. Other Mistral-specific params (e.g.
# include_image_base64, image_limit) are unsupported and ignored.
return OCRRequestData(data=data, files=None)
def _transform_azure_page(self, azure_page: AzureDocumentIntelligencePage) -> OCRPage:
page_number: Final = azure_page.pageNumber if azure_page.pageNumber is not None else 1
markdown: Final = "\n".join(line.content or "" for line in azure_page.lines)
dimensions: Final = self._convert_dimensions(
width=azure_page.width if azure_page.width is not None else 8.5,
height=azure_page.height if azure_page.height is not None else 11,
unit=azure_page.unit if azure_page.unit is not None else "inch",
)
return OCRPage(index=page_number - 1, markdown=markdown, dimensions=dimensions)
def _convert_dimensions(self, width: float, height: float, unit: str) -> OCRPageDimensions:
"""
Convert Azure DI dimensions to pixels.
Azure DI provides dimensions in inches. We convert to pixels using configured DPI.
Args:
width: Width in specified unit
height: Height in specified unit
unit: Unit of measurement (e.g., "inch")
Returns:
OCRPageDimensions with pixel values
"""
# Convert to pixels using configured DPI
dpi: Final = AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI
if unit == "inch":
width_px = int(width * dpi)
height_px = int(height * dpi)
else:
# If unit is not inches, assume it's already in pixels
width_px = int(width)
height_px = int(height)
return OCRPageDimensions(width=width_px, height=height_px, dpi=dpi)
@staticmethod
def _check_timeout(start_time: float, timeout_secs: int) -> None:
"""
Check if operation has timed out.
Args:
start_time: Start time of the operation
timeout_secs: Timeout duration in seconds
Raises:
TimeoutError: If operation has exceeded timeout
"""
if time.time() - start_time > timeout_secs:
raise TimeoutError(f"Azure Document Intelligence operation polling timed out after {timeout_secs} seconds")
@staticmethod
def _get_retry_after(response: httpx.Response) -> int:
"""
Get retry-after duration from response headers.
Args:
response: HTTP response
Returns:
Retry-after duration in seconds (default: 2)
"""
retry_after: Final = int(response.headers.get("retry-after", "2"))
verbose_logger.debug("Retry polling after: %s seconds", retry_after)
return retry_after
@staticmethod
def _check_operation_status(response: httpx.Response) -> str:
"""
Check Azure DI operation status from response.
Args:
response: HTTP response from operation endpoint
Returns:
Operation status string
Raises:
ValueError: If operation failed or status is unknown
"""
try:
result: Final = response.json()
status: Final = result.get("status")
verbose_logger.debug("Azure DI operation status: %s", status)
if status == "succeeded":
return "succeeded"
elif status == "failed":
error_msg: Final = result.get("error", {}).get("message", "Unknown error")
raise ValueError(f"Azure Document Intelligence analysis failed: {error_msg}")
elif status in ["running", "notStarted"]:
return "running"
else:
raise ValueError(f"Unknown operation status: {status}")
except Exception as e:
if "succeeded" in str(e) or "failed" in str(e):
raise
# If we can't parse JSON, something went wrong
raise ValueError(f"Failed to parse Azure DI operation response: {e}")
def _poll_operation_sync(
self,
operation_url: str,
headers: dict[str, str],
timeout_secs: int,
) -> httpx.Response:
"""
Poll Azure Document Intelligence operation until completion (sync).
Azure DI POST returns 202 with Operation-Location header.
We need to poll that URL until status is "succeeded" or "failed".
Args:
operation_url: The Operation-Location URL to poll
headers: Request headers (including auth)
timeout_secs: Total timeout in seconds
Returns:
Final response with completed analysis
"""
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
client: Final = _get_httpx_client()
start_time: Final = time.time()
verbose_logger.debug("Polling Azure DI operation: %s", operation_url)
while True:
self._check_timeout(start_time=start_time, timeout_secs=timeout_secs)
# Poll the operation status
response = client.get(url=operation_url, headers=headers)
# Check operation status
status = self._check_operation_status(response=response)
if status == "succeeded":
return response
elif status == "running":
# Wait before polling again
retry_after = self._get_retry_after(response=response)
time.sleep(retry_after)
async def _poll_operation_async(
self,
operation_url: str,
headers: dict[str, str],
timeout_secs: int,
) -> httpx.Response:
"""
Poll Azure Document Intelligence operation until completion (async).
Args:
operation_url: The Operation-Location URL to poll
headers: Request headers (including auth)
timeout_secs: Total timeout in seconds
Returns:
Final response with completed analysis
"""
import litellm
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.AZURE_AI)
start_time: Final = time.time()
verbose_logger.debug("Polling Azure DI operation (async): %s", operation_url)
while True:
self._check_timeout(start_time=start_time, timeout_secs=timeout_secs)
# Poll the operation status
response = await client.get(url=operation_url, headers=headers)
# Check operation status
status = self._check_operation_status(response=response)
if status == "succeeded":
return response
elif status == "running":
# Wait before polling again
retry_after = self._get_retry_after(response=response)
await asyncio.sleep(retry_after)
def _get_polling_target(self, raw_response: httpx.Response) -> tuple[str, dict[str, str]]:
operation_url: Final = raw_response.headers.get("Operation-Location")
if not operation_url:
raise ValueError("Azure Document Intelligence returned 202 but no Operation-Location header found")
# Reject cross-origin polling URLs — the auth headers
# below would otherwise leak to whatever URL the upstream
# (or an attacker-controlled upstream) returns. VERIA-51.
try:
assert_same_origin(operation_url, str(raw_response.request.url))
except SSRFError as ssrf_err:
raise ValueError(f"Azure Document Intelligence: rejected polling URL ({ssrf_err})")
poll_headers: Final = {
header: raw_response.request.headers[header]
for header in ("Ocp-Apim-Subscription-Key", "Authorization")
if header in raw_response.request.headers
}
return operation_url, poll_headers
@staticmethod
def _get_request_format(optional_params: object) -> OCRRequestFormat:
if not isinstance(optional_params, dict):
return "litellm"
request_format: Final = optional_params.get(OCR_REQUEST_FORMAT_PARAM)
if request_format is None:
return "litellm"
return parse_ocr_request_format(request_format)
def _transform_completed_response(
self,
model: str,
raw_response: httpx.Response,
request_format: OCRRequestFormat,
) -> OCRResponse:
"""
Transform a completed Azure Document Intelligence analyze operation
into the Mistral OCR response shape, preserving Azure-native
`analyzeResult` fields (`content`, `tables`, `keyValuePairs`) as
top-level response fields.
When `request_format` is "native", the untouched Azure operation
payload is attached to the response's hidden params so the proxy can
return it verbatim while cost tracking still reads `usage_info`.
"""
raw_operation: Final[Mapping[str, object]] = raw_response.json()
operation: Final = AzureDocumentIntelligenceOperation.model_validate(raw_operation)
verbose_logger.debug("Azure Document Intelligence response status: %s", operation.status)
if operation.status != "succeeded":
raise ValueError(f"Azure Document Intelligence analysis failed with status: {operation.status}")
analyze_result: Final = (
operation.analyzeResult if operation.analyzeResult is not None else AzureDocumentIntelligenceAnalyzeResult()
)
mistral_pages: Final = [self._transform_azure_page(azure_page) for azure_page in analyze_result.pages]
usage_info: Final = OCRUsageInfo(pages_processed=len(mistral_pages), doc_size_bytes=None)
response: Final = OCRResponse(
pages=mistral_pages,
model=model,
usage_info=usage_info,
object="ocr",
content=analyze_result.content,
tables=analyze_result.tables,
keyValuePairs=analyze_result.keyValuePairs,
)
if request_format == "native":
response.set_provider_native_response(raw_operation)
return response
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: "LiteLLMLoggingObj",
**kwargs,
) -> OCRResponse:
"""
Transform Azure Document Intelligence response to Mistral OCR format.
Handles async operation polling: If response is 202 Accepted, polls Operation-Location
until analysis completes.
Azure DI response (after polling):
{
"status": "succeeded",
"analyzeResult": {
"content": "Full document text...",
"pages": [
{
"pageNumber": 1,
"width": 8.5,
"height": 11,
"unit": "inch",
"lines": [{"content": "text", "boundingBox": [...]}]
}
],
"tables": [...],
"keyValuePairs": [...]
}
}
Mistral OCR format (with Azure-native fields preserved):
{
"pages": [
{
"index": 0,
"markdown": "extracted text",
"dimensions": {"width": 816, "height": 1056, "dpi": 96}
}
],
"model": "azure_ai/doc-intelligence/prebuilt-layout",
"usage_info": {"pages_processed": 1},
"object": "ocr",
"content": "Full document text...",
"tables": [...],
"keyValuePairs": [...]
}
Args:
model: Model name
raw_response: Raw HTTP response from Azure DI (may be 202 Accepted)
logging_obj: Logging object
Returns:
OCRResponse in Mistral format
"""
request_format: Final = self._get_request_format(kwargs.get("optional_params"))
if raw_response.status_code != 202:
return self._transform_completed_response(
model=model, raw_response=raw_response, request_format=request_format
)
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation...")
operation_url, poll_headers = self._get_polling_target(raw_response)
completed_response: Final = self._poll_operation_sync(
operation_url=operation_url,
headers=poll_headers,
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
)
return self._transform_completed_response(
model=model, raw_response=completed_response, request_format=request_format
)
async def async_transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: "LiteLLMLoggingObj",
**kwargs,
) -> OCRResponse:
"""
Async transform Azure Document Intelligence response to Mistral OCR format.
Handles async operation polling: If response is 202 Accepted, polls Operation-Location
until analysis completes using async polling.
Args:
model: Model name
raw_response: Raw HTTP response from Azure DI (may be 202 Accepted)
logging_obj: Logging object
Returns:
OCRResponse in Mistral format
"""
request_format: Final = self._get_request_format(kwargs.get("optional_params"))
if raw_response.status_code != 202:
return self._transform_completed_response(
model=model, raw_response=raw_response, request_format=request_format
)
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation (async)...")
operation_url, poll_headers = self._get_polling_target(raw_response)
completed_response: Final = await self._poll_operation_async(
operation_url=operation_url,
headers=poll_headers,
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
)
return self._transform_completed_response(
model=model, raw_response=completed_response, request_format=request_format
)

View file

@ -1,263 +0,0 @@
"""
Azure AI OCR transformation implementation.
"""
from typing import Final
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.image_handling import (
async_convert_url_to_base64,
convert_url_to_base64,
)
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRRequestData
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.secret_managers.main import get_secret_str
AZURE_AI_OCR_API_KEY_ENV_VAR: Final = "AZURE_AI_API_KEY"
class AzureAIOCRConfig(MistralOCRConfig):
"""
Azure AI OCR transformation configuration.
Azure AI uses Mistral's OCR API but with a different endpoint format.
Inherits transformation logic from MistralOCRConfig since they use the same format.
Reference: Azure AI Foundry OCR documentation
Important: Azure AI only supports base64 data URIs (data:image/..., data:application/pdf;base64,...).
Regular URLs are not supported.
"""
def __init__(self) -> None:
super().__init__()
def get_api_key_env_var(self) -> str | None:
return AZURE_AI_OCR_API_KEY_ENV_VAR
def validate_environment(
self,
headers: dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> dict:
"""
Validate environment and return headers for Azure AI OCR.
Authenticates with AZURE_AI_API_KEY, or with an Entra ID / OAuth token when no key is set.
"""
# Get API key from environment if not provided
if api_key is None:
api_key = get_secret_str(AZURE_AI_OCR_API_KEY_ENV_VAR)
# Validate API base is provided
if api_base is None:
api_base = get_secret_str("AZURE_AI_API_BASE")
if api_base is None:
raise ValueError(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter"
)
headers = {
**get_azure_ai_auth_headers(api_key=api_key, litellm_params=litellm_params),
"Content-Type": "application/json",
**headers,
}
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Azure AI OCR endpoint.
Azure AI endpoint format: https://<api_base>/providers/mistral/azure/ocr
Args:
api_base: Azure AI API base URL
model: Model name (not used in URL construction)
optional_params: Optional parameters
Returns: Complete URL for Azure AI OCR endpoint
"""
if api_base is None:
raise ValueError(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter"
)
# Ensure no trailing slash
api_base = api_base.rstrip("/")
# Azure AI OCR endpoint format
return f"{api_base}/providers/mistral/azure/ocr"
def _convert_url_to_data_uri_sync(self, url: str) -> str:
"""
Synchronously convert a URL to a base64 data URI.
Azure AI OCR doesn't have internet access, so we need to fetch URLs
and convert them to base64 data URIs.
Args:
url: The URL to convert
Returns:
Base64 data URI string
"""
verbose_logger.debug("Azure AI OCR: Converting URL to base64 data URI (sync): %s", url)
# Fetch and convert to base64 data URI
# convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
data_uri: Final = convert_url_to_base64(url=url)
verbose_logger.debug("Azure AI OCR: Converted URL to data URI (length: %s)", len(data_uri))
return data_uri
async def _convert_url_to_data_uri_async(self, url: str) -> str:
"""
Asynchronously convert a URL to a base64 data URI.
Azure AI OCR doesn't have internet access, so we need to fetch URLs
and convert them to base64 data URIs.
Args:
url: The URL to convert
Returns:
Base64 data URI string
"""
verbose_logger.debug("Azure AI OCR: Converting URL to base64 data URI (async): %s", url)
# Fetch and convert to base64 data URI asynchronously
# async_convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
data_uri: Final = await async_convert_url_to_base64(url=url)
verbose_logger.debug("Azure AI OCR: Converted URL to data URI (length: %s)", len(data_uri))
return data_uri
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Azure AI, converting URLs to base64 data URIs (sync).
Azure AI OCR doesn't have internet access, so we automatically fetch
any URLs and convert them to base64 data URIs synchronously.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug("Azure AI OCR transform_ocr_request (sync) - model: %s", model)
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Check if we need to convert URL to base64
doc_type: Final = document.get("type")
transformed_document: Final = document.copy()
if doc_type == "document_url":
document_url: Final = document.get("document_url", "")
# If it's not already a data URI, convert it
if document_url and not document_url.startswith("data:"):
verbose_logger.debug("Azure AI OCR: Converting document URL to base64 data URI (sync)")
data_uri = self._convert_url_to_data_uri_sync(url=document_url)
transformed_document["document_url"] = data_uri
elif doc_type == "image_url":
image_url: Final = document.get("image_url", "")
# If it's not already a data URI, convert it
if image_url and not image_url.startswith("data:"):
verbose_logger.debug("Azure AI OCR: Converting image URL to base64 data URI (sync)")
data_uri = self._convert_url_to_data_uri_sync(url=image_url)
transformed_document["image_url"] = data_uri
# Call parent's transform to build the request
return super().transform_ocr_request(
model=model,
document=transformed_document,
optional_params=optional_params,
headers=headers,
**kwargs,
)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Azure AI, converting URLs to base64 data URIs (async).
Azure AI OCR doesn't have internet access, so we automatically fetch
any URLs and convert them to base64 data URIs asynchronously.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug("Azure AI OCR async_transform_ocr_request - model: %s", model)
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Check if we need to convert URL to base64
doc_type: Final = document.get("type")
transformed_document: Final = document.copy()
if doc_type == "document_url":
document_url: Final = document.get("document_url", "")
# If it's not already a data URI, convert it
if document_url and not document_url.startswith("data:"):
verbose_logger.debug("Azure AI OCR: Converting document URL to base64 data URI (async)")
data_uri = await self._convert_url_to_data_uri_async(url=document_url)
transformed_document["document_url"] = data_uri
elif doc_type == "image_url":
image_url: Final = document.get("image_url", "")
# If it's not already a data URI, convert it
if image_url and not image_url.startswith("data:"):
verbose_logger.debug("Azure AI OCR: Converting image URL to base64 data URI (async)")
data_uri = await self._convert_url_to_data_uri_async(url=image_url)
transformed_document["image_url"] = data_uri
# Call parent's transform to build the request
return super().transform_ocr_request(
model=model,
document=transformed_document,
optional_params=optional_params,
headers=headers,
**kwargs,
)

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
@ -13,7 +13,7 @@ from litellm.llms.azure_ai.common_utils import (
api_key_header_for_base,
get_azure_ai_auth_headers,
)
from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.llms.base_llm.passthrough.transformation import (
BasePassthroughConfig,
RelayShape,
@ -22,6 +22,7 @@ from litellm.llms.base_llm.passthrough.transformation import (
relayed_body,
strip_leading_model_segment,
)
from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR_PASSTHROUGH_RESPONSE, NativeOcrPassthroughResponse
from litellm.types.llms.openai import AllMessageValues
from litellm.types.rerank import RerankResponse
from litellm.types.utils import CallTypes, ImageResponse, StandardPassThroughResponseObject
@ -30,7 +31,6 @@ if TYPE_CHECKING:
from httpx import URL, Response
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
from litellm.llms.base_llm.passthrough.transformation import LoggedRelayResponse
@ -92,9 +92,9 @@ FOUNDRY_RELAY_SHAPES: Final = (
class AzureAIPassthroughConfig(AzureFoundryModelInfo, BasePassthroughConfig):
def __init__(self, ocr_config_for: Callable[[str], BaseOCRConfig | None] = get_azure_ai_ocr_config) -> None:
def __init__(self, passthrough_ocr: NativeOcrPassthroughResponse | None = None) -> None:
super().__init__()
self.ocr_config_for: Final = ocr_config_for
self._passthrough_ocr: Final = passthrough_ocr
def is_streaming_request(self, endpoint: str, request_data: Mapping[str, object]) -> bool:
return bool(request_data.get("stream"))
@ -168,28 +168,22 @@ class AzureAIPassthroughConfig(AzureFoundryModelInfo, BasePassthroughConfig):
def logged_ocr_response(
self, model: str, httpx_response: Response, logging_obj: Logging, endpoint: str
) -> OCRResponse | None:
ocr_config: Final = self.ocr_config_for(model)
if ocr_config is None or httpx_response.status_code != 200:
return None
relayed_url: Final = httpx_response.request.url
relayed_origin: Final = str(relayed_url.copy_with(path="/", query=None, fragment=None)).rstrip("/")
ocr_url: Final = httpx.URL(
ocr_config.get_complete_url(
api_base=relayed_origin,
model=model,
optional_params={}, # mutable-ok: BaseOCRConfig wants a dict
)
passthrough_ocr: Final = (
self._passthrough_ocr if self._passthrough_ocr is not None else NATIVE_OCR_PASSTHROUGH_RESPONSE.load()
)
if passthrough_ocr is None or httpx_response.status_code != 200:
return None
known_prefixes: Final = (model, model_group_from(logging_obj.litellm_params))
native_endpoint: Final = strip_leading_model_segment(endpoint, known_prefixes)
if f"/{native_endpoint.strip('/')}" != ocr_url.path:
return None
try:
ocr_response: Final = ocr_config.transform_ocr_response(
model=model, raw_response=httpx_response, logging_obj=logging_obj
result: Final = passthrough_ocr(model, native_endpoint, httpx_response.content)
ocr_response: Final = OCRResponse.model_validate(result) if result is not None else None
except (ValueError, RuntimeError) as error:
verbose_logger.warning(
"azure_ai passthrough: OCR body from %s is not costable: %s", httpx_response.request.url, error
)
except (ValueError, AttributeError) as error:
verbose_logger.warning("azure_ai passthrough: OCR body from %s is not costable: %s", ocr_url, error)
return None
if ocr_response is None:
return None
logging_obj.call_type = CallTypes.aocr.value # rebind-ok: routes cost calculation to the per-page OCR path
return ocr_response

View file

@ -0,0 +1,54 @@
from collections.abc import Iterable, Mapping
from functools import cache
from types import MappingProxyType
from typing import Final, cast, get_type_hints
from litellm.types.llms.openai import ResponseInputParam, ResponsesAPIOptionalRequestParams
def _frozen_mapping(items: Iterable[tuple[str, object]]) -> Mapping[str, object]:
return MappingProxyType(dict(items))
@cache
def _responses_request_keys() -> frozenset[str]:
return frozenset(get_type_hints(ResponsesAPIOptionalRequestParams))
def responses_batch_body_to_chat_body(
openai_request_body: Mapping[str, object],
custom_llm_provider: str | None = None,
) -> dict[str, object]: # mutable-ok: provider transforms take the bridged chat body as a plain dict
"""
Rewrite the body of an OpenAI `/v1/responses` batch record as a Chat Completions body.
Batch providers translate chat bodies into their own request shape, so a Responses
record goes through the same Responses-to-Chat bridge the real-time path uses for
providers without a native Responses API: `input`, `instructions`, `max_output_tokens`
and the tool params translate identically in batch and real time. Like real time, the
record's fields are forwarded as sent instead of validated against the SDK TypedDicts,
whose required keys (a function tool's `strict`, an image part's `detail`) clients omit.
"""
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
responses_input: Final = openai_request_body.get("input")
if responses_input is None:
raise ValueError(
"Batch record for /v1/responses is missing required `input` field: "
f"model={openai_request_body.get('model', '')}"
)
model: Final = openai_request_body.get("model")
chat_input: Final = cast(str | ResponseInputParam, responses_input) # cast-ok: forwarded as sent
responses_request: Final = cast( # cast-ok: client-supplied fields forwarded verbatim, as real time does
ResponsesAPIOptionalRequestParams,
_frozen_mapping((key, value) for key, value in openai_request_body.items() if key in _responses_request_keys()),
)
return LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # transformer declares a bare dict return
model=model if isinstance(model, str) else "",
input=chat_input,
responses_api_request=responses_request,
custom_llm_provider=custom_llm_provider,
metadata=openai_request_body.get("metadata"),
)

View file

@ -1,23 +1,19 @@
"""Base OCR transformation module."""
from .transformation import (
BaseOCRConfig,
DocumentType,
OCRPage,
OCRPageDimensions,
OCRPageImage,
OCRRequestData,
OCRResponse,
OCRUsageInfo,
)
__all__ = [
"BaseOCRConfig",
"DocumentType",
"OCRPage",
"OCRPageDimensions",
"OCRPageImage",
"OCRRequestData",
"OCRResponse",
"OCRUsageInfo",
]

View file

@ -1,27 +1,16 @@
"""
Base OCR transformation configuration.
Base OCR types shared by the Rust OCR route and Python consumers.
"""
import builtins
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import Any, Final, Literal
import httpx
from pydantic import PrivateAttr
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.base import LiteLLMPydanticObjectBase
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
# DocumentType for OCR - providers always receive a dict with
# type="document_url" or type="image_url" (str values only).
# File-type inputs are preprocessed to this format in litellm/ocr/main.py.
DocumentType = dict[str, str]
DocumentType = Mapping[str, object]
OCRRequestFormat = Literal["litellm", "native"]
@ -33,8 +22,6 @@ OCR_REQUEST_FORMAT_HEADER: Final = "x-req-format"
PROVIDER_NATIVE_RESPONSE_KEY: Final = "provider_native_response"
HEALTH_CHECK_PDF_DATA_URI: Final = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y="
def parse_ocr_request_format(value: object) -> OCRRequestFormat:
if value == "litellm":
@ -102,7 +89,6 @@ class OCRResponse(LiteLLMPydanticObjectBase):
model_config = {"extra": "allow"}
# Define private attributes using PrivateAttr
_hidden_params: dict = PrivateAttr(default_factory=dict)
def set_provider_native_response(self, native_response: Mapping[str, builtins.object]) -> None:
@ -113,203 +99,3 @@ class OCRResponse(LiteLLMPydanticObjectBase):
"""The provider's own response payload, when `req_format=native` was requested."""
native_response: Final = self._hidden_params.get(PROVIDER_NATIVE_RESPONSE_KEY)
return native_response if isinstance(native_response, dict) else None
class OCRRequestData(LiteLLMPydanticObjectBase):
"""OCR request data structure."""
data: dict | bytes | None = None
files: dict[str, Any] | None = None
class BaseOCRConfig:
"""
Base configuration for OCR transformations.
Handles provider-agnostic OCR operations.
"""
def __init__(self) -> None:
pass
def get_supported_ocr_params(self, model: str) -> list:
"""
Get supported OCR parameters for this provider.
Override this method in provider-specific implementations.
"""
return []
def get_api_key_env_var(self) -> str | None:
"""
Return the provider-specific API key environment variable name, if any.
"""
return None
def resolve_connection_params(
self,
*,
api_key: str | None,
api_base: str | None,
dynamic_api_key: str | None,
dynamic_api_base: str | None,
) -> tuple[str | None, str | None]:
return dynamic_api_key or api_key, dynamic_api_base or api_base
def get_health_check_document(self) -> DocumentType:
return { # mutable-ok: litellm.aocr rejects any document that is not a dict
"type": "document_url",
"document_url": HEALTH_CHECK_PDF_DATA_URI,
}
def map_ocr_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
) -> dict:
"""Map OCR parameters to provider-specific parameters."""
return optional_params
def validate_environment(
self,
headers: dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> dict:
"""
Validate environment and return headers.
Override in provider-specific implementations.
"""
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for OCR endpoint.
Override in provider-specific implementations.
"""
raise NotImplementedError("get_complete_url must be implemented by provider")
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request to provider-specific format.
Override in provider-specific implementations.
Note: By the time this method is called, any file-type documents have already
been converted to document_url/image_url format with base64 data URIs by
the preprocessing in litellm/ocr/main.py.
Args:
model: Model name
document: Document to process - always a dict with type="document_url" or type="image_url"
optional_params: Optional parameters for the request
headers: Request headers
Returns:
OCRRequestData with data and files fields
"""
raise NotImplementedError("transform_ocr_request must be implemented by provider")
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Async transform OCR request to provider-specific format.
Optional method - providers can override if they need async transformations
(e.g., Azure AI for URL-to-base64 conversion).
Default implementation falls back to sync transform_ocr_request.
Args:
model: Model name
document: Document to process (Mistral format dict, or file path, bytes, etc.)
optional_params: Optional parameters for the request
headers: Request headers
Returns:
OCRRequestData with data and files fields
"""
# Default implementation: call sync version
return self.transform_ocr_request(
model=model,
document=document,
optional_params=optional_params,
headers=headers,
**kwargs,
)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs,
) -> OCRResponse:
"""
Transform provider-specific OCR response to standard format.
Override in provider-specific implementations.
"""
raise NotImplementedError("transform_ocr_response must be implemented by provider")
async def async_transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs,
) -> OCRResponse:
"""
Async transform provider-specific OCR response to standard format.
Optional method - providers can override if they need async transformations
(e.g., Azure Document Intelligence for async operation polling).
Default implementation falls back to sync transform_ocr_response.
Args:
model: Model name
raw_response: Raw HTTP response
logging_obj: Logging object
Returns:
OCRResponse in standard format
"""
# Default implementation: call sync version
return self.transform_ocr_response(
model=model,
raw_response=raw_response,
logging_obj=logging_obj,
**kwargs,
)
def get_error_class(
self,
error_message: str,
status_code: int,
headers: dict,
) -> Exception:
"""Get appropriate error class for the provider."""
return BaseLLMException(
status_code=status_code,
message=error_message,
headers=headers,
)

View file

@ -1,5 +1,5 @@
import base64
from typing import Final, NoReturn
from typing import Final
import httpx
@ -16,14 +16,6 @@ from litellm.rust_bridge.transcription.native import (
from litellm.types.utils import FileTypes, TranscriptionResponse
def _no_python_implementation() -> NoReturn:
raise NotImplementedError("Bedrock audio transcription is implemented in Rust only")
async def _no_async_python_implementation() -> NoReturn:
_no_python_implementation()
class BedrockAudioTranscriptionRustDispatch:
@staticmethod
def _audio_payload(audio_file: FileTypes) -> dict[str, object]:
@ -77,7 +69,7 @@ class BedrockAudioTranscriptionRustDispatch:
RouteContext(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model),
binding=NATIVE_TRANSCRIPTION,
native=native,
python=_no_python_implementation,
python=runtime.NO_PYTHON,
)
async def async_audio_transcriptions(
@ -110,5 +102,5 @@ class BedrockAudioTranscriptionRustDispatch:
RouteContext(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model),
binding=NATIVE_ATRANSCRIPTION,
native=native,
python=_no_async_python_implementation,
python=runtime.NO_PYTHON,
)

View file

@ -8,7 +8,6 @@ from collections.abc import Iterable, Mapping, MutableMapping, Sequence
from contextlib import suppress
from dataclasses import dataclass
from datetime import datetime
from functools import cache
from itertools import chain
from types import MappingProxyType
from typing import Any, Final, Literal, TypeAlias, TypedDict
@ -17,7 +16,7 @@ from urllib.parse import quote, unquote, urlencode
import httpx
from httpx import Headers, Response
from openai.types.file_deleted import FileDeleted
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
from pydantic import BaseModel, ConfigDict, Field
from typing_extensions import ReadOnly
from litellm._logging import verbose_logger
@ -41,7 +40,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_data,
text_completion_prompt_to_messages,
)
from litellm.llms.base_llm.base_utils import map_developer_role_to_system_role
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.files.batch_records import responses_batch_body_to_chat_body
from litellm.llms.base_llm.files.transformation import (
BaseFilesConfig,
LiteLLMLoggingObj,
@ -56,11 +57,9 @@ from litellm.types.llms.openai import (
OpenAICreateFileRequestOptionalParams,
OpenAIFileObject,
PathLike,
ResponseInputParam,
ResponsesAPIOptionalRequestParams,
)
from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums
from litellm.utils import get_llm_provider
from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums, all_litellm_params
from litellm.utils import get_llm_provider, get_optional_params
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
@ -89,6 +88,14 @@ def _frozen_mapping(items: Iterable[tuple[str, object]]) -> Mapping[str, object]
return MappingProxyType(dict(items))
_LITELLM_PARAMS_THE_MAPPER_TAKES: Final = frozenset({"allowed_openai_params"})
_MAPPED_PARAMS_THE_REQUEST_HANDLER_STRIPS: Final = frozenset({"json_mode"})
def _invoke_route_model(model: str) -> str:
return f"invoke/{_strip_llm_routing_prefix(model).removeprefix('invoke/')}"
def _strip_llm_routing_prefix(model: str) -> str:
try:
stripped_model, _, _, _ = get_llm_provider(model=model, custom_llm_provider=None)
@ -130,22 +137,6 @@ class _S3UploadResponse(TypedDict, total=False):
ContentLength: ReadOnly[int]
# JSONL batch records are untyped json, so the `/v1/responses` fields are
# validated into their concrete Responses API types before being handed to the
# Responses-to-Chat bridge. Both adapters drop keys the Responses API doesn't
# define, which is what the bridge would ignore anyway. Built on first use
# rather than at import: `ResponseInputParam` is a deep union and only batch
# files carrying `/v1/responses` records need it.
@cache
def _responses_input_adapter() -> TypeAdapter[str | ResponseInputParam]:
return TypeAdapter(str | ResponseInputParam)
@cache
def _responses_request_adapter() -> TypeAdapter[ResponsesAPIOptionalRequestParams]:
return TypeAdapter(ResponsesAPIOptionalRequestParams)
class _BedrockS3RequestParams(AwsAuthParams):
"""Typed view of the credential/region params the S3 GetObject path reads."""
@ -859,33 +850,9 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
Delegates to the same Responses-to-Chat bridge the real-time path uses
for providers without a native Responses API (which is every Bedrock
model), so `input`, `instructions`, `max_output_tokens` and the tool
params translate identically in batch and real time. The bridge always
emits a `tools` key; an empty one is dropped rather than shipped as an
empty array inside `modelInput`.
params translate identically in batch and real time.
"""
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
responses_input: Final = openai_request_body.get("input")
if responses_input is None:
raise ValueError(
"Batch record for /v1/responses is missing required `input` field: "
f"model={openai_request_body.get('model', '')}"
)
chat_body: Final[Mapping[str, object]] = (
LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request(
model=openai_request_body.get("model", ""),
input=_responses_input_adapter().validate_python(responses_input),
responses_api_request=_responses_request_adapter().validate_python(
_frozen_mapping(
(key, value) for key, value in openai_request_body.items() if key not in ("model", "input")
)
),
metadata=openai_request_body.get("metadata"),
)
)
return _frozen_mapping((key, value) for key, value in chat_body.items() if key != "tools" or value)
return responses_batch_body_to_chat_body(openai_request_body)
@staticmethod
def _transform_batch_body_to_chat_body(
@ -922,7 +889,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
"""
from litellm.types.utils import LlmProviders
messages: Final = openai_request_body.get("messages", [])
messages: Final = map_developer_role_to_system_role(openai_request_body.get("messages", []))
optional_params: Final = {k: v for k, v in openai_request_body.items() if k not in ["model", "messages"]}
# --- Anthropic: use existing AmazonAnthropicClaudeConfig ---
@ -932,16 +899,24 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
)
config: Final = AmazonAnthropicClaudeConfig()
mapped_params = config.map_openai_params(
non_default_params={},
optional_params=optional_params,
model=model,
drop_params=False,
mapped_params = get_optional_params(
model=_invoke_route_model(model),
custom_llm_provider="bedrock",
messages=messages,
**MappingProxyType(
{
k: v
for k, v in optional_params.items()
if k not in all_litellm_params or k in _LITELLM_PARAMS_THE_MAPPER_TAKES
}
),
)
return config.transform_request(
model=model,
messages=messages,
optional_params=mapped_params,
optional_params={
k: v for k, v in mapped_params.items() if k not in _MAPPED_PARAMS_THE_REQUEST_HANDLER_STRIPS
},
litellm_params={},
headers={},
)

View file

@ -1,3 +0,0 @@
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
__all__ = ("CohereParseConfig",)

View file

@ -1,298 +0,0 @@
"""Cohere Parse (`POST /v2/parse`) exposed through LiteLLM's OCR interface."""
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal
import httpx
from pydantic import BaseModel, ConfigDict, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm.exceptions import BadRequestError, UnsupportedParamsError
from litellm.llms.base_llm.ocr.transformation import (
OCR_REQUEST_FORMAT_PARAM,
BaseOCRConfig,
DocumentType,
OCRPage,
OCRPageImage,
OCRRequestData,
OCRRequestFormat,
OCRResponse,
OCRUsageInfo,
parse_ocr_request_format,
)
from litellm.llms.cohere.common_utils import CohereError
from litellm.secret_managers.main import get_secret_str
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
COHERE_API_KEY_ENV_VAR: Final = "COHERE_API_KEY"
COHERE_PARSE_API_BASE: Final = "https://api.cohere.com"
COHERE_PARSE_PATH: Final = "/v2/parse"
COHERE_PARSE_OUTPUT_FORMAT_PARAM: Final = "output_format"
COHERE_PARSE_OUTPUT_FORMATS: Final = ("markdown", "blocks")
COHERE_PARSE_DEFAULT_OUTPUT_FORMAT: Final = "markdown"
COHERE_PARSE_SUPPORTED_PARAMS: Final = (COHERE_PARSE_OUTPUT_FORMAT_PARAM, OCR_REQUEST_FORMAT_PARAM)
COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI: Final = (
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
)
COHERE_PARSE_IMAGE_ONLY_MESSAGE: Final = (
"Cohere Parse only accepts `image_url` documents (an image URL or a base64 image data URI); "
"`document_url` and PDF inputs are not supported."
)
_NATIVE_RESPONSE_ADAPTER: Final = TypeAdapter(dict[str, object])
_BOUNDING_BOX_ADAPTER: Final = TypeAdapter(Mapping[str, object])
class _CohereParseDocument(TypedDict):
type: ReadOnly[Literal["image_url"]]
image_url: ReadOnly[str]
class _CohereParseRequestBody(TypedDict):
model: ReadOnly[str]
document: ReadOnly[_CohereParseDocument]
output_format: ReadOnly[str]
class _MarkdownPage(TypedDict):
index: ReadOnly[int]
markdown: ReadOnly[str]
images: ReadOnly[Sequence[OCRPageImage] | None]
class _BlocksPage(_MarkdownPage):
blocks: ReadOnly[Sequence[Mapping[str, object]]]
class _CohereParseMarkdown(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
content: str = ""
images: Sequence[Mapping[str, object]] | None = None
class _CohereParsePage(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
index: int | None = None
markdown: _CohereParseMarkdown | None = None
blocks: Sequence[Mapping[str, object]] | None = None
class _CohereParseBilledUnits(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
pages: int | None = None
class _CohereParseMeta(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
billed_units: _CohereParseBilledUnits | None = None
class _CohereParseResponse(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
pages: Sequence[_CohereParsePage] = ()
meta: _CohereParseMeta | None = None
def _requested_format(optional_params: Mapping[str, object] | None) -> OCRRequestFormat:
if optional_params is None:
return "litellm"
return "native" if optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" else "litellm"
def _page_image(image: Mapping[str, object]) -> OCRPageImage:
bounding_box: Final = image.get("bounding_box")
if not isinstance(bounding_box, Mapping):
return OCRPageImage.model_validate(image)
bbox: Final = _BOUNDING_BOX_ADAPTER.validate_python(bounding_box)
return OCRPageImage.model_validate(MappingProxyType({**image, "bbox": bbox}))
def _normalize_page(page: _CohereParsePage, position: int) -> OCRPage:
markdown: Final = page.markdown
images: Final = tuple(_page_image(image) for image in markdown.images) if markdown and markdown.images else None
normalized: Final[_MarkdownPage] = {
"index": page.index if page.index is not None else position,
"markdown": markdown.content if markdown else "",
"images": images,
}
if page.blocks is None:
return OCRPage.model_validate(normalized)
with_blocks: Final[_BlocksPage] = {**normalized, "blocks": page.blocks}
return OCRPage.model_validate(with_blocks)
def _billed_pages(parsed: _CohereParseResponse) -> int | None:
if parsed.meta is None or parsed.meta.billed_units is None:
return None
return parsed.meta.billed_units.pages
class CohereParseConfig(BaseOCRConfig):
"""Cohere Parse, an image-only document understanding endpoint returning markdown or blocks."""
def get_supported_ocr_params(self, model: str) -> list[str]: # mutable-ok: BaseOCRConfig signature
return list(COHERE_PARSE_SUPPORTED_PARAMS) # mutable-ok: BaseOCRConfig signature
def get_api_key_env_var(self) -> str | None:
return COHERE_API_KEY_ENV_VAR
def get_health_check_document(self) -> DocumentType:
return { # mutable-ok: litellm.aocr rejects any document that is not a dict
"type": "image_url",
"image_url": COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI,
}
def _llm_provider(self) -> str:
return "cohere"
def map_ocr_params(
self,
non_default_params: Mapping[str, object],
optional_params: Mapping[str, object],
model: str,
) -> dict[str, object]: # mutable-ok: BaseOCRConfig signature
output_format: Final = non_default_params.get(COHERE_PARSE_OUTPUT_FORMAT_PARAM)
if output_format is not None and output_format not in COHERE_PARSE_OUTPUT_FORMATS:
raise UnsupportedParamsError(
message=(
f"Invalid `{COHERE_PARSE_OUTPUT_FORMAT_PARAM}`: {output_format!r}. "
f"Expected one of {', '.join(COHERE_PARSE_OUTPUT_FORMATS)}."
),
model=model,
llm_provider=self._llm_provider(),
)
requested_format: Final = non_default_params.get(OCR_REQUEST_FORMAT_PARAM)
request_format: Final = parse_ocr_request_format(requested_format) if requested_format is not None else None
overrides: Final = tuple(
(key, value)
for key, value in (
(COHERE_PARSE_OUTPUT_FORMAT_PARAM, output_format),
(OCR_REQUEST_FORMAT_PARAM, request_format),
)
if value is not None
)
return {**optional_params, **dict(overrides)} # mutable-ok: BaseOCRConfig signature
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: Mapping[str, object] | None = None,
**kwargs: object, # kwargs-ok: BaseOCRConfig.validate_environment signature
) -> dict[str, str]: # mutable-ok: BaseOCRConfig signature
resolved_key: Final = api_key or get_secret_str(COHERE_API_KEY_ENV_VAR)
if resolved_key is None:
raise ValueError(
f"Missing {COHERE_API_KEY_ENV_VAR} - set it in the environment or pass api_key to "
"litellm.ocr()/litellm.aocr()"
)
return { # mutable-ok: BaseOCRConfig signature
"Authorization": f"Bearer {resolved_key}",
"Content-Type": "application/json",
**headers,
}
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object] | None = None,
**kwargs: object, # kwargs-ok: BaseOCRConfig.get_complete_url signature
) -> str:
url: Final = httpx.URL(api_base or COHERE_PARSE_API_BASE)
path: Final = url.path.rstrip("/")
if path.endswith(COHERE_PARSE_PATH):
return str(url.copy_with(path=path))
if path.endswith("/v2"):
return str(url.copy_with(path=f"{path}/parse"))
return str(url.copy_with(path=f"{path}{COHERE_PARSE_PATH}"))
def _image_url(self, document: DocumentType, model: str) -> str:
image_url: Final = document.get("image_url", "")
if document.get("type") != "image_url" or not image_url or image_url.startswith("data:application/pdf"):
raise BadRequestError(
message=COHERE_PARSE_IMAGE_ONLY_MESSAGE,
model=model,
llm_provider=self._llm_provider(),
)
return image_url
def _resolve_image_url_sync(self, image_url: str) -> str:
return image_url
async def _resolve_image_url_async(self, image_url: str) -> str:
return image_url
def _build_request(self, model: str, image_url: str, optional_params: Mapping[str, object]) -> OCRRequestData:
body: Final[_CohereParseRequestBody] = {
"model": model,
"document": {"type": "image_url", "image_url": image_url},
"output_format": str(
optional_params.get(COHERE_PARSE_OUTPUT_FORMAT_PARAM, COHERE_PARSE_DEFAULT_OUTPUT_FORMAT)
),
}
return OCRRequestData(data=dict(body), files=None) # mutable-ok: OCRRequestData.data is a dict
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: Mapping[str, object],
headers: Mapping[str, str],
**kwargs: object, # kwargs-ok: BaseOCRConfig.transform_ocr_request signature
) -> OCRRequestData:
image_url: Final = self._resolve_image_url_sync(self._image_url(document, model))
return self._build_request(model=model, image_url=image_url, optional_params=optional_params)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: Mapping[str, object],
headers: Mapping[str, str],
**kwargs: object, # kwargs-ok: BaseOCRConfig.async_transform_ocr_request signature
) -> OCRRequestData:
image_url: Final = await self._resolve_image_url_async(self._image_url(document, model))
return self._build_request(model=model, image_url=image_url, optional_params=optional_params)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: "LiteLLMLoggingObj",
optional_params: Mapping[str, object] | None = None,
**kwargs: object, # kwargs-ok: BaseOCRConfig.transform_ocr_response signature
) -> OCRResponse:
native: Final = _NATIVE_RESPONSE_ADAPTER.validate_python(raw_response.json())
parsed: Final = _CohereParseResponse.model_validate(native)
pages: Final = [ # mutable-ok: OCRResponse.pages is a list
_normalize_page(page, position) for position, page in enumerate(parsed.pages)
]
billed_pages: Final = _billed_pages(parsed)
response: Final = OCRResponse(
pages=pages,
model=model,
usage_info=OCRUsageInfo(pages_processed=billed_pages if billed_pages is not None else len(pages)),
)
if _requested_format(optional_params) == "native":
response.set_provider_native_response(native)
return response
def get_error_class(
self,
error_message: str,
status_code: int,
headers: Mapping[str, str],
) -> Exception:
return CohereError(status_code=status_code, message=error_message)

View file

@ -81,7 +81,6 @@ from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse
from litellm.llms.base_llm.realtime.http_transformation import BaseRealtimeHTTPConfig
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
@ -1636,320 +1635,6 @@ class BaseLLMHTTPHandler:
api_key=api_key,
)
def _prepare_ocr_request(
self,
model: str,
document: dict[str, str],
optional_params: dict,
logging_obj: LiteLLMLoggingObj,
api_key: str | None,
api_base: str | None,
headers: dict[str, object] | None,
provider_config: BaseOCRConfig,
litellm_params: dict,
) -> tuple[dict[str, object], str, dict[str, object], None]:
"""
Shared logic for preparing OCR requests.
Returns: (headers, complete_url, data, files)
"""
from litellm.llms.base_llm.ocr.transformation import OCRRequestData
headers = provider_config.validate_environment(
api_key=api_key,
api_base=api_base,
headers=headers or {},
model=model,
litellm_params=litellm_params,
)
complete_url: Final = provider_config.get_complete_url(
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
)
# Transform the request to get data and files
transformed_result: Final = provider_config.transform_ocr_request(
model=model,
document=document,
optional_params={key: value for key, value in optional_params.items() if key != OCR_REQUEST_FORMAT_PARAM},
headers=headers,
api_key=api_key,
api_base=api_base,
)
# All providers return OCRRequestData
if not isinstance(transformed_result, OCRRequestData):
raise ValueError(f"Provider {provider_config.__class__.__name__} must return OCRRequestData")
# Data is always a dict for Mistral OCR format
if not isinstance(transformed_result.data, dict):
raise ValueError(f"Expected dict data for OCR request, got {type(transformed_result.data)}")
data: Final = transformed_result.data
## LOGGING
logging_obj.pre_call(
input="OCR document processing",
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": complete_url,
"headers": headers,
},
)
return headers, complete_url, data, None
async def _async_prepare_ocr_request(
self,
model: str,
document: dict[str, str],
optional_params: dict,
logging_obj: LiteLLMLoggingObj,
api_key: str | None,
api_base: str | None,
headers: dict[str, object] | None,
provider_config: BaseOCRConfig,
litellm_params: dict,
) -> tuple[dict[str, object], str, dict[str, object], None]:
"""
Async version of _prepare_ocr_request for providers that need async transforms.
Returns: (headers, complete_url, data, files)
"""
from litellm.llms.base_llm.ocr.transformation import OCRRequestData
headers = provider_config.validate_environment(
api_key=api_key,
api_base=api_base,
headers=headers or {},
model=model,
litellm_params=litellm_params,
)
complete_url: Final = provider_config.get_complete_url(
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
)
# Use async transform (providers can override this method if they need async operations)
transformed_result: Final = await provider_config.async_transform_ocr_request(
model=model,
document=document,
optional_params={key: value for key, value in optional_params.items() if key != OCR_REQUEST_FORMAT_PARAM},
headers=headers,
api_key=api_key,
api_base=api_base,
)
# All providers return OCRRequestData
if not isinstance(transformed_result, OCRRequestData):
raise ValueError(f"Provider {provider_config.__class__.__name__} must return OCRRequestData")
# Data is always a dict for Mistral OCR format
if not isinstance(transformed_result.data, dict):
raise ValueError(f"Expected dict data for OCR request, got {type(transformed_result.data)}")
data: Final = transformed_result.data
## LOGGING
logging_obj.pre_call(
input="OCR document processing",
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": complete_url,
"headers": headers,
},
)
return headers, complete_url, data, None
def _transform_ocr_response(
self,
provider_config: BaseOCRConfig,
model: str,
response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
optional_params: Mapping[str, object],
) -> OCRResponse:
"""Shared logic for transforming OCR responses."""
normalized: Final = provider_config.transform_ocr_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
optional_params=optional_params,
)
return self._finalize_ocr_response(normalized, response, optional_params)
@staticmethod
def _finalize_ocr_response(
normalized: OCRResponse,
response: httpx.Response,
optional_params: Mapping[str, object],
) -> OCRResponse:
if (
optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"
and normalized.get_provider_native_response() is None
):
normalized.set_provider_native_response(response.json())
return normalized
def ocr(
self,
model: str,
document: dict[str, str],
optional_params: dict,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
client: HTTPHandler | AsyncHTTPHandler | None = None,
aocr: bool = False,
headers: dict[str, object] | None = None,
provider_config: BaseOCRConfig | None = None,
litellm_params: dict | None = None,
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
"""
Sync OCR handler.
"""
if provider_config is None:
raise ValueError(f"No provider config found for model: {model} and provider: {custom_llm_provider}")
if litellm_params is None:
litellm_params = {}
if aocr is True:
return self.async_ocr(
model=model,
document=document,
optional_params=optional_params,
timeout=timeout,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
client=client,
headers=headers,
provider_config=provider_config,
litellm_params=litellm_params,
)
# Prepare the request
headers, complete_url, data, files = self._prepare_ocr_request(
model=model,
document=document,
optional_params=optional_params,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
headers=headers,
provider_config=provider_config,
litellm_params=litellm_params,
)
if client is None or not isinstance(client, HTTPHandler):
client = _get_httpx_client()
try:
# Make the POST request with JSON data (Mistral format)
response: Final = client.post(
url=complete_url,
headers=headers,
json=data,
timeout=timeout,
)
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
logging_obj.post_call(
api_key=api_key,
original_response=response.text,
additional_args={"complete_input_dict": data},
)
return self._transform_ocr_response(
provider_config=provider_config,
model=model,
response=response,
logging_obj=logging_obj,
optional_params=optional_params,
)
async def async_ocr(
self,
model: str,
document: dict[str, str],
optional_params: dict,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
client: HTTPHandler | AsyncHTTPHandler | None = None,
headers: dict[str, object] | None = None,
provider_config: BaseOCRConfig | None = None,
litellm_params: dict | None = None,
) -> OCRResponse:
"""
Async OCR handler.
"""
if provider_config is None:
raise ValueError(f"No provider config found for model: {model} and provider: {custom_llm_provider}")
if litellm_params is None:
litellm_params = {}
# Prepare the request using async prepare method
headers, complete_url, data, files = await self._async_prepare_ocr_request(
model=model,
document=document,
optional_params=optional_params,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
headers=headers,
provider_config=provider_config,
litellm_params=litellm_params,
)
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders(custom_llm_provider),
)
else:
async_httpx_client = client
try:
# Make the async POST request with JSON data (Mistral format)
response: Final = await async_httpx_client.post(
url=complete_url,
headers=headers,
json=data,
timeout=timeout,
)
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
logging_obj.post_call(
api_key=api_key,
original_response=response.text,
additional_args={"complete_input_dict": data},
)
# Use async response transform for async operations
normalized: Final = await provider_config.async_transform_ocr_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
optional_params=optional_params,
)
return self._finalize_ocr_response(normalized, response, optional_params)
def search(
self,
query: str | list[str],
@ -6211,7 +5896,6 @@ class BaseLLMHTTPHandler:
BaseGoogleGenAIGenerateContentConfig,
BaseAnthropicMessagesConfig,
BaseBatchesConfig,
BaseOCRConfig,
BaseVideoConfig,
BaseSearchConfig,
BaseTextToSpeechConfig,
@ -6254,12 +5938,6 @@ class BaseLLMHTTPHandler:
status_code=status_code,
headers=error_headers,
)
if (
isinstance(provider_config, BaseOCRConfig)
and isinstance(provider_error, BaseLLMException)
and isinstance(error_response, httpx.Response)
):
provider_error.response = error_response
if not isinstance(received_status_code, int):
provider_error.status_code_is_synthesized = True
raise provider_error

View file

@ -369,7 +369,7 @@ model LiteLLM_SkillsTable {
Run the tests:
```bash
pytest tests/proxy_unit_tests/test_skills_db.py -v
pytest tests/unit/skills/test_skills_db.py -v
```
Tests cover:

View file

@ -1,244 +0,0 @@
"""
Mistral OCR transformation implementation.
"""
from typing import TYPE_CHECKING, Final
import httpx
from litellm._logging import verbose_logger
from litellm.llms.base_llm.ocr.transformation import (
BaseOCRConfig,
DocumentType,
OCRRequestData,
OCRResponse,
)
from litellm.secret_managers.main import get_secret_str
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
MISTRAL_OCR_API_KEY_ENV_VAR: Final = "MISTRAL_API_KEY"
class MistralOCRConfig(BaseOCRConfig):
"""
Mistral OCR transformation configuration.
Reference: https://docs.mistral.ai/api/#tag/ocr
"""
def __init__(self) -> None:
super().__init__()
def get_supported_ocr_params(self, model: str) -> list:
"""
Get supported OCR parameters for Mistral OCR.
Mistral OCR supports:
- pages: List of page numbers to process
- include_image_base64: Whether to include base64 encoded images
- image_limit: Maximum number of images to return
- image_min_size: Minimum size of images to include
- bbox_annotation_format: Format for bounding box annotations
- document_annotation_format: Format for document annotations
- document_annotation_prompt: Prompt for document annotation extraction
- extract_header: Whether to extract document header
- extract_footer: Whether to extract document footer
- table_format: Table output format ("markdown" or "html")
- confidence_scores_granularity: Confidence score level ("word" or "page")
- include_blocks: Whether to return paragraph-level bounding boxes and typed content blocks (OCR 4)
- id: Request identifier
"""
return [
"pages",
"include_image_base64",
"image_limit",
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
"document_annotation_prompt",
"extract_header",
"extract_footer",
"table_format",
"confidence_scores_granularity",
"include_blocks",
"id",
]
def get_api_key_env_var(self) -> str | None:
return MISTRAL_OCR_API_KEY_ENV_VAR
def map_ocr_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
) -> dict:
"""
Map OCR parameters to Mistral-specific format.
Mistral accepts these parameters directly, so no transformation needed.
Just filter out unsupported params.
"""
supported_params: Final = self.get_supported_ocr_params(model=model)
# Only include params that are in the supported list
mapped_params: Final = {}
for param, value in non_default_params.items():
if param in supported_params:
mapped_params[param] = value
return mapped_params
def validate_environment(
self,
headers: dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> dict:
"""
Validate environment and return headers for Mistral OCR.
"""
# Get API key from environment if not provided
if api_key is None:
api_key = get_secret_str(MISTRAL_OCR_API_KEY_ENV_VAR)
if api_key is None:
raise ValueError(
"Missing Mistral API Key - A call is being made to Mistral but no key is set either in the environment variables or via params"
)
headers = {
"Authorization": f"Bearer {api_key}",
**headers,
}
# Don't set Content-Type for multipart/form-data - httpx will handle it
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Mistral OCR endpoint.
Returns: https://api.mistral.ai/v1/ocr
"""
if api_base is None:
api_base = "https://api.mistral.ai/v1"
# Ensure no trailing slash
api_base = api_base.rstrip("/")
# Remove /v1 if it's already in the base to avoid duplication
if api_base.endswith("/v1"):
return f"{api_base}/ocr"
return f"{api_base}/v1/ocr"
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request to Mistral-specific format.
Mistral OCR API accepts:
{
"model": "mistral-ocr-latest",
"document": {
"type": "document_url",
"document_url": "<https-url or data-uri>"
},
"pages": [0], # optional
"include_image_base64": false, # optional
...
}
Args:
model: Model name (e.g., "mistral-ocr-latest")
document: Document dict from user (Mistral format) - already validated in main.py
optional_params: Already mapped optional parameters
headers: Request headers
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug("Mistral OCR transform_ocr_request - model: %s", model)
# Document parameter is the Mistral-format dict from the user
# Just pass it through as-is to the Mistral API
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Build request data - use document dict directly
data: Final = {
"model": model,
"document": document, # Pass through the Mistral-format document dict
}
# Add all optional parameters from the already-mapped optional_params
data.update(optional_params)
# No multipart files - using JSON
return OCRRequestData(data=data, files=None)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: "LiteLLMLoggingObj",
**kwargs,
) -> OCRResponse:
"""
Return Mistral OCR response in native format.
Mistral OCR is the standard format for LiteLLM OCR responses.
No transformation needed - return native response.
Mistral OCR returns:
{
"pages": [
{
"index": 0,
"markdown": "extracted text content",
"images": [...],
"dimensions": {...}
},
...
],
"model": "mistral-ocr-2505-completion",
"document_annotation": null,
"usage_info": {...}
}
"""
try:
response_json: Final = raw_response.json()
verbose_logger.debug("Mistral OCR response keys: %s", response_json.keys())
# Return native Mistral format - no transformation
return OCRResponse(
pages=response_json.get("pages", []),
model=response_json.get("model", model),
document_annotation=response_json.get("document_annotation"),
usage_info=response_json.get("usage_info"),
object="ocr",
)
except Exception as e:
verbose_logger.error("Error parsing Mistral OCR response: %s", e)
raise e

View file

@ -0,0 +1,133 @@
"""OpenAI's organization costs endpoint: the USD the organization was billed per UTC day, read with an admin key."""
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import date, datetime, timedelta, timezone
from types import MappingProxyType
from typing import Final, Literal, TypeAlias
import httpx
from pydantic import BaseModel, ConfigDict, ValidationError
from litellm.constants import (
OPENAI_ORGANIZATION_COSTS_PAGE_LIMIT,
OPENAI_ORGANIZATION_COSTS_URL,
PROVIDER_BILLING_TIMEOUT_SECONDS,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
OPENAI_ADMIN_KEY_ENV_VAR: Final = "OPENAI_ADMIN_KEY"
BillingHttpGet: TypeAlias = Callable[
[str, Mapping[str, object], Mapping[str, str]], # mutable-ok: Callable parameter list is type syntax
Awaitable[httpx.Response],
]
@dataclass(frozen=True, slots=True)
class OpenAICostsRequestFailed:
detail: str
class _OpenAICostAmount(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
value: float
currency: Literal["usd"]
class _OpenAICostResult(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
amount: _OpenAICostAmount
class _OpenAICostBucket(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
start_time: int
results: tuple[_OpenAICostResult, ...] = ()
class _OpenAICostsPage(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
data: tuple[_OpenAICostBucket, ...]
has_more: bool = False
next_page: str | None = None
async def provider_billing_get(url: str, params: Mapping[str, object], headers: Mapping[str, str]) -> httpx.Response:
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.ProviderBilling)
return await client.get(
url,
params=dict(params), # mutable-ok: AsyncHTTPHandler.get takes dict params
headers=dict(headers), # mutable-ok: AsyncHTTPHandler.get takes dict headers
timeout=PROVIDER_BILLING_TIMEOUT_SECONDS,
)
def _utc_midnight(day: date) -> int:
return int(datetime(day.year, day.month, day.day, tzinfo=timezone.utc).timestamp())
def _bucket_day(bucket: _OpenAICostBucket) -> str:
return datetime.fromtimestamp(bucket.start_time, tz=timezone.utc).date().isoformat()
async def fetch_openai_daily_costs(
start_date: date,
end_date: date,
*,
admin_key: str,
project_ids: Sequence[str] = (),
http_get: BillingHttpGet = provider_billing_get,
) -> Mapping[str, float] | OpenAICostsRequestFailed:
"""USD billed by OpenAI per UTC day (ISO date) over the closed range, following pagination to the end."""
scope: Final = (("project_ids[]", tuple(project_ids)),) if project_ids else ()
window: Final[Mapping[str, object]] = MappingProxyType(
{
key: value
for key, value in (
("start_time", _utc_midnight(start_date)),
("end_time", _utc_midnight(end_date + timedelta(days=1))),
("bucket_width", "1d"),
("limit", OPENAI_ORGANIZATION_COSTS_PAGE_LIMIT),
*scope,
)
}
)
headers: Final[Mapping[str, str]] = MappingProxyType({"Authorization": f"Bearer {admin_key}"})
async def fetch_from(page: str | None) -> tuple[_OpenAICostBucket, ...] | OpenAICostsRequestFailed:
params: Final[Mapping[str, object]] = MappingProxyType(
{key: value for key, value in (*window.items(), ("page", page)) if value is not None}
)
try:
response: Final = await http_get(OPENAI_ORGANIZATION_COSTS_URL, params, headers)
except httpx.HTTPError as exc:
return OpenAICostsRequestFailed(f"request failed: {exc}")
if response.status_code != 200:
return OpenAICostsRequestFailed(f"HTTP {response.status_code}: {response.text[:300]}")
try:
parsed: Final = _OpenAICostsPage.model_validate(response.json())
except (ValueError, ValidationError) as exc:
return OpenAICostsRequestFailed(f"unexpected response shape: {exc}")
if not parsed.has_more or parsed.next_page is None:
return parsed.data
rest: Final = await fetch_from(parsed.next_page)
return rest if isinstance(rest, OpenAICostsRequestFailed) else parsed.data + rest
buckets: Final = await fetch_from(None)
if isinstance(buckets, OpenAICostsRequestFailed):
return buckets
days: Final = frozenset(_bucket_day(bucket) for bucket in buckets)
return MappingProxyType(
{
day: sum(
result.amount.value for bucket in buckets if _bucket_day(bucket) == day for result in bucket.results
)
for day in days
}
)

View file

@ -1,149 +0,0 @@
import base64
import binascii
from collections import defaultdict
from typing import TYPE_CHECKING, Any, Final, NoReturn
import httpx
from litellm.constants import request_timeout
REDUCTO_API_BASE: Final = "https://platform.reducto.ai"
REDUCTO_ID_PREFIX: Final = "reducto://"
if TYPE_CHECKING:
from litellm.llms.base_llm.ocr.transformation import OCRPage
def _normalize_api_base(api_base: str | None) -> str:
return (api_base or REDUCTO_API_BASE).rstrip("/")
def _raise_bad_request(message: str, model: str) -> NoReturn:
import litellm
raise litellm.BadRequestError(
message=message,
model=model,
llm_provider="reducto",
)
def extract_file_id_or_bytes(
source_url: str,
model: str,
) -> tuple[str | None, bytes | None, str | None]:
if source_url.startswith(REDUCTO_ID_PREFIX):
return source_url, None, None
if source_url.startswith("http://") or source_url.startswith("https://"):
_raise_bad_request(
"Reducto requires type='file' (auto-uploaded) or a reducto:// id. Plain http(s) URLs are not supported; upload the file first.",
model=model,
)
if not source_url.startswith("data:"):
_raise_bad_request(
"Reducto requires a reducto:// id or a base64 data URI after OCR preprocessing.",
model=model,
)
try:
header, encoded = source_url.split(",", 1)
except ValueError:
_raise_bad_request("Invalid Reducto data URI provided.", model=model)
if ";base64" not in header:
_raise_bad_request("Reducto only supports base64-encoded data URIs.", model=model)
mime: Final = header.removeprefix("data:").split(";")[0] or "application/octet-stream"
try:
raw_bytes: Final = base64.b64decode(encoded, validate=True)
except (binascii.Error, ValueError):
_raise_bad_request("Invalid Reducto base64 payload provided.", model=model)
return None, raw_bytes, mime
def _extract_file_id_from_upload_response(response: httpx.Response) -> str:
try:
payload: Final = response.json()
except ValueError as exc:
raise ValueError(f"Reducto /upload returned a non-JSON 200 response: {response.text}") from exc
file_id: Final = (payload or {}).get("file_id") if isinstance(payload, dict) else None
if not isinstance(file_id, str) or not file_id:
raise ValueError(f"Reducto /upload returned 200 without a file_id; got payload={payload}")
return file_id
def upload_bytes_sync(
raw_bytes: bytes,
mime: str | None,
api_key: str,
api_base: str | None,
) -> str:
import litellm
response: Final = litellm.module_level_client.post(
url="{}{}".format(_normalize_api_base(api_base), "/upload"),
headers={"Authorization": f"Bearer {api_key}"},
files={"file": ("document", raw_bytes, mime or "application/octet-stream")},
timeout=request_timeout,
)
response.raise_for_status()
return _extract_file_id_from_upload_response(response)
async def upload_bytes_async(
raw_bytes: bytes,
mime: str | None,
api_key: str,
api_base: str | None,
) -> str:
import litellm
response: Final = await litellm.module_level_aclient.post(
url="{}{}".format(_normalize_api_base(api_base), "/upload"),
headers={"Authorization": f"Bearer {api_key}"},
files={"file": ("document", raw_bytes, mime or "application/octet-stream")},
timeout=request_timeout,
)
response.raise_for_status()
return _extract_file_id_from_upload_response(response)
def build_pages_from_reducto(result: dict[str, Any]) -> list["OCRPage"]:
from litellm.llms.base_llm.ocr.transformation import OCRPage
chunks: Final = result.get("chunks", []) or []
blocks_by_page: Final[dict[int, list[dict[str, Any]]]] = defaultdict(list)
for chunk in chunks:
for block in chunk.get("blocks", []) or []:
page_no = (block.get("bbox") or {}).get("page")
if page_no is None:
continue
try:
normalized_page = int(page_no)
except (TypeError, ValueError):
continue
blocks_by_page[normalized_page].append(block)
if not blocks_by_page:
fallback_markdown: Final = "\n\n".join(chunk.get("content", "") for chunk in chunks if chunk.get("content"))
if fallback_markdown == "":
return []
return [OCRPage(index=0, markdown=fallback_markdown)]
pages: Final[list[OCRPage]] = []
for page_no, blocks in sorted(blocks_by_page.items()):
markdown = "\n\n".join(block.get("content", "") for block in blocks if block.get("content"))
page_index = max(page_no - 1, 0)
page = OCRPage(
index=page_index,
markdown=markdown,
)
# OCRPage accepts extra keys at runtime; assign blocks after construction
# so static typing does not reject provider-specific metadata.
setattr(page, "blocks", blocks)
pages.append(page)
return pages

View file

@ -1 +0,0 @@

View file

@ -1,236 +0,0 @@
from typing import TYPE_CHECKING, Any, Final
import httpx
from litellm.llms.base_llm.ocr.transformation import (
BaseOCRConfig,
DocumentType,
OCRRequestData,
OCRResponse,
OCRUsageInfo,
)
from litellm.llms.reducto.common import (
REDUCTO_API_BASE,
build_pages_from_reducto,
extract_file_id_or_bytes,
upload_bytes_async,
upload_bytes_sync,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
class _BaseReductoOCRConfig(BaseOCRConfig):
def map_ocr_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
) -> dict:
mapped_params: Final = dict(optional_params)
supported_params: Final = self.get_supported_ocr_params(model=model)
for param, value in non_default_params.items():
if param in supported_params:
mapped_params[param] = value
return mapped_params
def validate_environment(
self,
headers: dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> dict:
from litellm.secret_managers.main import get_secret_str
resolved_key: Final = api_key or get_secret_str("REDUCTO_API_KEY")
if resolved_key is None:
raise ValueError(
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
)
return {
"Authorization": f"Bearer {resolved_key}",
"Content-Type": "application/json",
**headers,
}
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
return "{}/parse".format((api_base or REDUCTO_API_BASE).rstrip("/"))
def _get_source_url(self, document: DocumentType, model: str) -> str:
source_url: Final = document.get("document_url") or document.get("image_url")
if source_url is None:
raise ValueError(
f"Reducto expected OCR preprocessing to produce document_url or image_url for model={model}"
)
return source_url
@staticmethod
def _resolve_credentials(api_key: str | None, api_base: str | None) -> tuple[str, str]:
from litellm.secret_managers.main import get_secret_str
resolved_key: Final = api_key or get_secret_str("REDUCTO_API_KEY")
if resolved_key is None:
raise ValueError(
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
)
resolved_base: Final = (api_base or REDUCTO_API_BASE).rstrip("/")
return resolved_key, resolved_base
def _ensure_file_id_sync(
self,
model: str,
document: DocumentType,
api_key: str | None,
api_base: str | None,
) -> str:
source_url: Final = self._get_source_url(document=document, model=model)
file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model)
if file_id is not None:
return file_id
resolved_key, resolved_base = self._resolve_credentials(api_key, api_base)
return upload_bytes_sync(
raw_bytes=raw_bytes or b"",
mime=mime,
api_key=resolved_key,
api_base=resolved_base,
)
async def _ensure_file_id_async(
self,
model: str,
document: DocumentType,
api_key: str | None,
api_base: str | None,
) -> str:
source_url: Final = self._get_source_url(document=document, model=model)
file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model)
if file_id is not None:
return file_id
resolved_key, resolved_base = self._resolve_credentials(api_key, api_base)
return await upload_bytes_async(
raw_bytes=raw_bytes or b"",
mime=mime,
api_key=resolved_key,
api_base=resolved_base,
)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: "LiteLLMLoggingObj",
**kwargs,
) -> OCRResponse:
response_json: Final = raw_response.json()
result: Final = response_json.get("result", response_json) or {}
usage: Final = response_json.get("usage", {}) or {}
response: Final = OCRResponse(
pages=build_pages_from_reducto(result),
model=model,
usage_info=OCRUsageInfo(
pages_processed=usage.get("num_pages"),
credits=usage.get("credits"),
),
object="ocr",
)
response._hidden_params["reducto_raw"] = response_json
return response
class ReductoParseV3Config(_BaseReductoOCRConfig):
def get_supported_ocr_params(self, model: str) -> list:
return ["formatting", "retrieval", "settings"]
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id: Final = self._ensure_file_id_sync(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(data={"input": file_id, **optional_params}, files=None)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id: Final = await self._ensure_file_id_async(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(data={"input": file_id, **optional_params}, files=None)
class ReductoParseLegacyConfig(_BaseReductoOCRConfig):
def get_supported_ocr_params(self, model: str) -> list:
return ["enhance"]
def _build_legacy_body(self, file_id: str, optional_params: dict) -> dict[str, Any]:
body: Final[dict[str, Any]] = {"document_url": file_id}
enhance: Final = optional_params.get("enhance")
if enhance is not None:
body["options"] = {"enhance": enhance}
return body
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id: Final = self._ensure_file_id_sync(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(
data=self._build_legacy_body(file_id=file_id, optional_params=optional_params),
files=None,
)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id: Final = await self._ensure_file_id_async(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(
data=self._build_legacy_body(file_id=file_id, optional_params=optional_params),
files=None,
)

View file

@ -35,7 +35,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_data,
extract_file_metadata,
)
from litellm.llms.base_llm.base_utils import map_developer_role_to_system_role
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.files.batch_records import responses_batch_body_to_chat_body
from litellm.llms.base_llm.files.transformation import (
BaseFilesConfig,
BaseFileUploadStream,
@ -529,21 +531,30 @@ def is_passthrough_batch_upload(create_file_data: Mapping[str, object], litellm_
return create_file_data.get("purpose") == "batch" and litellm_params.get("passthrough") is True
def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool:
def _batch_entry_route_path(openai_entry: Mapping[str, object]) -> str:
"""
Whether an OpenAI batch JSONL line targets the embeddings endpoint.
The route an OpenAI batch JSONL line targets, without query string or trailing slash.
OpenAI puts the target route on each line's `url` (e.g. `/v1/embeddings`); Vertex
has no equivalent per-line field, so the route decides which Vertex request shape
the line has to be translated into.
"""
url = openai_entry.get("url")
url: Final = openai_entry.get("url")
if not isinstance(url, str):
return False
path = url.split("?")[0].rstrip("/")
return ""
return url.split("?")[0].rstrip("/")
def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool:
path: Final = _batch_entry_route_path(openai_entry)
return path == "embeddings" or path.endswith("/embeddings")
def _is_responses_batch_entry(openai_entry: Mapping[str, object]) -> bool:
path: Final = _batch_entry_route_path(openai_entry)
return path == "responses" or path.endswith("/responses")
def _openai_embedding_input_elements(
embedding_input: GeminiEmbeddingInput,
) -> tuple[str | list[str], ...]:
@ -665,10 +676,15 @@ def _openai_batch_jsonl_entry_to_vertex_rows(
return _openai_batch_jsonl_entry_to_vertex_embeddings_rows(openai_entry)
openai_request_body: Final = openai_entry.get("body") or {}
chat_request_body: Final = (
responses_batch_body_to_chat_body(openai_request_body, custom_llm_provider="vertex_ai")
if _is_responses_batch_entry(openai_entry)
else openai_request_body
)
vertex_request_body: Final = _transform_request_body(
messages=openai_request_body.get("messages", []),
model=openai_request_body.get("model", ""),
optional_params=map_openai_to_vertex_params(openai_request_body),
messages=map_developer_role_to_system_role(chat_request_body.get("messages", [])),
model=chat_request_body.get("model", ""),
optional_params=map_openai_to_vertex_params(chat_request_body),
custom_llm_provider="vertex_ai",
litellm_params={},
cached_content=None,

View file

@ -1,5 +0,0 @@
"""Vertex AI OCR module."""
from .transformation import VertexAIOCRConfig
__all__ = ["VertexAIOCRConfig"]

View file

@ -1,41 +0,0 @@
"""
Common utilities for Vertex AI OCR providers.
This module provides routing logic to determine which OCR configuration to use
based on the model name.
"""
from typing import TYPE_CHECKING, Optional
if TYPE_CHECKING:
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
def get_vertex_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
"""
Determine which Vertex AI OCR configuration to use based on the model name.
Vertex AI supports multiple OCR services:
- Vertex AI OCR: vertex_ai/<model>
Args:
model: The model name (e.g., "vertex_ai/ocr/<model>")
Returns:
OCR configuration instance for the specified model
Examples:
>>> get_vertex_ai_ocr_config("vertex_ai/deepseek-ai/deepseek-ocr-maas")
<VertexAIDeepSeekOCRConfig object>
>>> get_vertex_ai_ocr_config("vertex_ai/ocr/mistral-ocr-maas")
<VertexAIOCRConfig object>
"""
from litellm.llms.vertex_ai.ocr.deepseek_transformation import (
VertexAIDeepSeekOCRConfig,
)
from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig
if "deepseek" in model:
return VertexAIDeepSeekOCRConfig()
return VertexAIOCRConfig()

View file

@ -1,378 +0,0 @@
"""
Vertex AI DeepSeek OCR transformation implementation.
"""
import json
from typing import TYPE_CHECKING, Any, Final
import httpx
from litellm._logging import verbose_logger
from litellm.llms.base_llm.ocr.transformation import (
BaseOCRConfig,
DocumentType,
OCRPage,
OCRRequestData,
OCRResponse,
OCRUsageInfo,
)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
VERTEX_AI_DEEPSEEK_OCR_API_KEY_ENV_VAR: Final = "VERTEX_AI_API_KEY"
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class VertexAIDeepSeekOCRConfig(BaseOCRConfig):
"""
Vertex AI DeepSeek OCR transformation configuration.
This transformation converts standard LiteLLM OCR requests to the
Vertex AI DeepSeek OCR OpenAPI endpoint shape and normalizes the response.
"""
def __init__(self) -> None:
super().__init__()
self.vertex_base = VertexBase()
def get_api_key_env_var(self) -> str | None:
return VERTEX_AI_DEEPSEEK_OCR_API_KEY_ENV_VAR
def validate_environment(
self,
headers: dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> dict:
"""
Validate environment and return headers for Vertex AI OCR.
Vertex AI uses Bearer token authentication with access token from credentials.
"""
if api_key is not None:
return {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
**headers,
}
# Extract Vertex AI parameters using safe helpers from VertexBase
# Use safe_get_* methods that don't mutate litellm_params dict
litellm_params = litellm_params or {}
vertex_project: Final = VertexBase.safe_get_vertex_ai_project(litellm_params=litellm_params)
vertex_credentials: Final = VertexBase.safe_get_vertex_ai_credentials(litellm_params=litellm_params)
# Get access token from Vertex credentials
access_token, project_id = self.vertex_base.get_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
)
headers = {
"Authorization": f"Bearer {access_token}",
"Content-Type": "application/json",
**headers,
}
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Vertex AI DeepSeek OCR endpoint.
Args:
api_base: Vertex AI API base URL (optional)
model: Model name (e.g., "deepseek-ai/deepseek-ocr-maas")
optional_params: Optional parameters
litellm_params: LiteLLM parameters containing vertex_project, vertex_location
Returns: Complete URL for Vertex AI OCR endpoint
"""
# Extract Vertex AI parameters using safe helpers from VertexBase
# Use safe_get_* methods that don't mutate litellm_params dict
litellm_params = litellm_params or {}
vertex_project: Final = VertexBase.safe_get_vertex_ai_project(litellm_params=litellm_params)
vertex_location = VertexBase.safe_get_vertex_ai_location(litellm_params=litellm_params)
if vertex_project is None:
raise ValueError(
"Missing vertex_project - Set VERTEXAI_PROJECT environment variable or pass vertex_project parameter"
)
if vertex_location is None:
vertex_location = "us-central1"
# Get API base URL
if api_base is None:
api_base = "https://aiplatform.googleapis.com"
# Ensure no trailing slash
api_base = api_base.rstrip("/")
return f"{api_base}/v1/projects/{vertex_project}/locations/{vertex_location}/endpoints/openapi/chat/completions"
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Vertex AI DeepSeek OCR.
Converts OCR document format to the Vertex AI DeepSeek OCR payload:
- Input: {"type": "image_url", "image_url": "gs://..."}
- Output: {"model": "deepseek-ai/deepseek-ocr-maas", "messages": [{"role": "user", "content": [{"type": "image_url", "image_url": "gs://..."}]}]}
Args:
model: Model name (e.g., "deepseek-ai/deepseek-ocr-maas")
document: Document dict from user (Mistral OCR format)
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data for the DeepSeek OCR endpoint
"""
verbose_logger.debug("Vertex AI DeepSeek OCR transform_ocr_request (sync) called")
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Extract document type and URL
doc_type: Final = document.get("type")
image_url = None
document_url = None
if doc_type == "image_url":
image_url = document.get("image_url", "")
elif doc_type == "document_url":
document_url = document.get("document_url", "")
else:
raise ValueError(f"Unsupported document type: {doc_type}. Expected 'image_url' or 'document_url'")
# Build DeepSeek OCR message content
content_item = {}
if image_url:
content_item = {"type": "image_url", "image_url": image_url}
elif document_url:
# For document URLs, we use image_url type as well (Vertex AI supports both)
content_item = {"type": "image_url", "image_url": document_url}
# Build DeepSeek OCR request
provider_model: Final = model if model.startswith("deepseek-ai/") else f"deepseek-ai/{model}"
data: Final = {
"model": provider_model,
"messages": [{"role": "user", "content": [content_item]}],
}
# Add optional parameters (stream, temperature, etc.)
deepseek_ocr_params: Final = {}
for key, value in optional_params.items():
if key in ["stream", "temperature", "max_tokens", "top_p", "n", "stop"]:
deepseek_ocr_params[key] = value
data.update(deepseek_ocr_params)
verbose_logger.debug("Vertex AI DeepSeek OCR: Transformed request")
return OCRRequestData(data=data, files=None)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Vertex AI DeepSeek OCR (async).
Same as sync version - no async-specific logic needed.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data for the DeepSeek OCR endpoint
"""
return self.transform_ocr_request(
model=model,
document=document,
optional_params=optional_params,
headers=headers,
**kwargs,
)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs,
) -> OCRResponse:
"""
Transform Vertex AI DeepSeek OCR response to OCR format.
Vertex AI DeepSeek OCR returns an OpenAPI response:
{
"id": "...",
"choices": [{
"message": {
"role": "assistant",
"content": "<OCR result as JSON string or markdown>"
}
}],
"usage": {...}
}
We need to extract the content and convert it to OCRResponse format.
Args:
model: Model name
raw_response: Raw HTTP response from Vertex AI
logging_obj: Logging object
**kwargs: Additional arguments
Returns:
OCRResponse in standard format
"""
verbose_logger.debug("Vertex AI DeepSeek OCR transform_ocr_response called")
verbose_logger.debug("Raw response: %s", raw_response.text)
try:
response_json: Final = raw_response.json()
# Extract OCR content from provider response
choices: Final = response_json.get("choices", [])
if not choices:
raise ValueError("No choices in DeepSeek OCR response")
message: Final = choices[0].get("message", {})
content: Final = message.get("content", "")
if not content:
raise ValueError("No content in DeepSeek OCR response")
# Try to parse content as JSON (OCR result might be JSON string)
ocr_data = None
try:
# If content is a JSON string, parse it
if isinstance(content, str) and content.strip().startswith("{"):
ocr_data = json.loads(content)
elif isinstance(content, dict):
ocr_data = content
else:
# If content is markdown text, create a single page with the markdown
ocr_data = {
"pages": [{"index": 0, "markdown": content}],
"model": model,
"usage_info": response_json.get("usage", {}),
}
except json.JSONDecodeError:
# If JSON parsing fails, treat content as markdown
ocr_data = {
"pages": [{"index": 0, "markdown": content}],
"model": model,
"usage_info": response_json.get("usage", {}),
}
# Ensure we have the expected structure
if "pages" not in ocr_data:
# If OCR data doesn't have pages, wrap the content in a page
ocr_data = {
"pages": [
{
"index": 0,
"markdown": (content if isinstance(content, str) else json.dumps(content)),
}
],
"model": ocr_data.get("model", model),
"usage_info": ocr_data.get("usage_info", response_json.get("usage", {})),
}
# Convert usage info if present
usage_info = None
if "usage_info" in ocr_data:
usage_dict: Final = ocr_data["usage_info"]
if isinstance(usage_dict, dict):
usage_info = OCRUsageInfo(**usage_dict)
# Build OCRResponse
pages = []
for page_data in ocr_data.get("pages", []):
# Ensure page has required fields
if isinstance(page_data, dict):
page = OCRPage(
index=page_data.get("index", 0),
markdown=page_data.get("markdown", ""),
images=page_data.get("images"),
dimensions=page_data.get("dimensions"),
)
pages.append(page)
if not pages:
# Create a default page if none exist
pages = [OCRPage(index=0, markdown=content if isinstance(content, str) else "")]
return OCRResponse(
pages=pages,
model=ocr_data.get("model", model),
document_annotation=ocr_data.get("document_annotation"),
usage_info=usage_info,
object="ocr",
)
except Exception as e:
verbose_logger.error("Error parsing Vertex AI DeepSeek OCR response: %s", e)
raise e
async def async_transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs,
) -> OCRResponse:
"""
Async transform Vertex AI DeepSeek OCR response to OCR format.
Same as sync version - no async-specific logic needed.
Args:
model: Model name
raw_response: Raw HTTP response
logging_obj: Logging object
**kwargs: Additional arguments
Returns:
OCRResponse in standard format
"""
return self.transform_ocr_response(
model=model,
raw_response=raw_response,
logging_obj=logging_obj,
**kwargs,
)

View file

@ -1,288 +0,0 @@
"""
Vertex AI Mistral OCR transformation implementation.
"""
from typing import Final
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.image_handling import (
async_convert_url_to_base64,
convert_url_to_base64,
)
from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRRequestData
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
VERTEX_AI_OCR_API_KEY_ENV_VAR: Final = "VERTEX_AI_API_KEY"
class VertexAIOCRConfig(MistralOCRConfig):
"""
Vertex AI Mistral OCR transformation configuration.
Vertex AI uses Mistral's OCR API format through the Mistral publisher endpoint.
Inherits transformation logic from MistralOCRConfig since they use the same format.
Reference: Vertex AI Mistral OCR documentation
Important: Vertex AI OCR only supports base64 data URIs (data:image/..., data:application/pdf;base64,...).
Regular URLs are not supported.
"""
def __init__(self) -> None:
super().__init__()
self.vertex_base = VertexBase()
def get_api_key_env_var(self) -> str | None:
return VERTEX_AI_OCR_API_KEY_ENV_VAR
def validate_environment(
self,
headers: dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> dict:
"""
Validate environment and return headers for Vertex AI OCR.
Vertex AI uses Bearer token authentication with access token from credentials.
"""
if api_key is not None:
return {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
**headers,
}
# Extract Vertex AI parameters using safe helpers from VertexBase
# Use safe_get_* methods that don't mutate litellm_params dict
litellm_params = litellm_params or {}
vertex_project: Final = VertexBase.safe_get_vertex_ai_project(litellm_params=litellm_params)
vertex_credentials: Final = VertexBase.safe_get_vertex_ai_credentials(litellm_params=litellm_params)
# Get access token from Vertex credentials
access_token, project_id = self.vertex_base.get_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
)
headers = {
"Authorization": f"Bearer {access_token}",
"Content-Type": "application/json",
**headers,
}
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Vertex AI OCR endpoint.
Vertex AI endpoint format:
https://{location}-aiplatform.googleapis.com/v1/projects/{project}/locations/{location}/publishers/mistralai/ocr
Args:
api_base: Vertex AI API base URL (optional)
model: Model name (not used in URL construction)
optional_params: Optional parameters
litellm_params: LiteLLM parameters containing vertex_project, vertex_location
Returns: Complete URL for Vertex AI OCR endpoint
"""
# Extract Vertex AI parameters using safe helpers from VertexBase
# Use safe_get_* methods that don't mutate litellm_params dict
litellm_params = litellm_params or {}
vertex_project: Final = VertexBase.safe_get_vertex_ai_project(litellm_params=litellm_params)
vertex_location = VertexBase.safe_get_vertex_ai_location(litellm_params=litellm_params)
if vertex_project is None:
raise ValueError(
"Missing vertex_project - Set VERTEXAI_PROJECT environment variable or pass vertex_project parameter"
)
if vertex_location is None:
vertex_location = "us-central1"
# Get API base URL
if api_base is None:
api_base = get_vertex_base_url(vertex_location)
# Ensure no trailing slash
api_base = api_base.rstrip("/")
# Vertex AI OCR endpoint format for Mistral publisher
# Format: https://{region}-aiplatform.googleapis.com/v1/projects/{project}/locations/{region}/publishers/mistralai/models/{model}:rawPredict
return f"{api_base}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/mistralai/models/{model}:rawPredict"
def _convert_url_to_data_uri_sync(self, url: str) -> str:
"""
Synchronously convert a URL to a base64 data URI.
Vertex AI OCR doesn't have internet access, so we need to fetch URLs
and convert them to base64 data URIs.
Args:
url: The URL to convert
Returns:
Base64 data URI string
"""
verbose_logger.debug("Vertex AI OCR: Converting URL to base64 data URI (sync): %s", url)
# Fetch and convert to base64 data URI
# convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
data_uri: Final = convert_url_to_base64(url=url)
verbose_logger.debug("Vertex AI OCR: Converted URL to data URI (length: %s)", len(data_uri))
return data_uri
async def _convert_url_to_data_uri_async(self, url: str) -> str:
"""
Asynchronously convert a URL to a base64 data URI.
Vertex AI OCR doesn't have internet access, so we need to fetch URLs
and convert them to base64 data URIs.
Args:
url: The URL to convert
Returns:
Base64 data URI string
"""
verbose_logger.debug("Vertex AI OCR: Converting URL to base64 data URI (async): %s", url)
# Fetch and convert to base64 data URI asynchronously
# async_convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
data_uri: Final = await async_convert_url_to_base64(url=url)
verbose_logger.debug("Vertex AI OCR: Converted URL to data URI (length: %s)", len(data_uri))
return data_uri
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Vertex AI, converting URLs to base64 data URIs (sync).
Vertex AI OCR doesn't have internet access, so we automatically fetch
any URLs and convert them to base64 data URIs synchronously.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug("Vertex AI OCR transform_ocr_request (sync) called")
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Check if we need to convert URL to base64
doc_type: Final = document.get("type")
transformed_document: Final = document.copy()
if doc_type == "document_url":
document_url: Final = document.get("document_url", "")
# If it's not already a data URI, convert it
if document_url and not document_url.startswith("data:"):
verbose_logger.debug("Vertex AI OCR: Converting document URL to base64 data URI (sync)")
data_uri = self._convert_url_to_data_uri_sync(url=document_url)
transformed_document["document_url"] = data_uri
elif doc_type == "image_url":
image_url: Final = document.get("image_url", "")
# If it's not already a data URI, convert it
if image_url and not image_url.startswith("data:"):
verbose_logger.debug("Vertex AI OCR: Converting image URL to base64 data URI (sync)")
data_uri = self._convert_url_to_data_uri_sync(url=image_url)
transformed_document["image_url"] = data_uri
# Call parent's transform to build the request
return super().transform_ocr_request(
model=model,
document=transformed_document,
optional_params=optional_params,
headers=headers,
**kwargs,
)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Vertex AI, converting URLs to base64 data URIs (async).
Vertex AI OCR doesn't have internet access, so we automatically fetch
any URLs and convert them to base64 data URIs asynchronously.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug("Vertex AI OCR async_transform_ocr_request - model: %s", model)
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Check if we need to convert URL to base64
doc_type: Final = document.get("type")
transformed_document: Final = document.copy()
if doc_type == "document_url":
document_url: Final = document.get("document_url", "")
# If it's not already a data URI, convert it
if document_url and not document_url.startswith("data:"):
verbose_logger.debug("Vertex AI OCR: Converting document URL to base64 data URI (async)")
data_uri = await self._convert_url_to_data_uri_async(url=document_url)
transformed_document["document_url"] = data_uri
elif doc_type == "image_url":
image_url: Final = document.get("image_url", "")
# If it's not already a data URI, convert it
if image_url and not image_url.startswith("data:"):
verbose_logger.debug("Vertex AI OCR: Converting image URL to base64 data URI (async)")
data_uri = await self._convert_url_to_data_uri_async(url=image_url)
transformed_document["image_url"] = data_uri
# Call parent's transform to build the request
return super().transform_ocr_request(
model=model,
document=transformed_document,
optional_params=optional_params,
headers=headers,
**kwargs,
)

View file

@ -6092,7 +6092,7 @@
"cache_creation_input_audio_token_cost": 3e-07,
"cache_read_input_audio_token_cost": 3e-07,
"cache_read_input_token_cost": 6e-08,
"deprecation_date": "2027-06-25",
"deprecation_date": "2027-07-31",
"input_cost_per_audio_token": 1e-05,
"input_cost_per_image_token": 8e-07,
"input_cost_per_token": 6e-07,
@ -26207,12 +26207,10 @@
},
"gemini-2.5-flash-image": {
"deprecation_date": "2027-03-15",
"cache_read_input_token_cost": 3e-08,
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 3e-07,
"input_cost_per_token_batches": 1.5e-07,
"input_cost_per_token_flex": 1.5e-07,
"input_cost_per_token_priority": 5.4e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
@ -26313,10 +26311,14 @@
"gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_priority": 3.6e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_priority": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -26326,7 +26328,9 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_priority": 2.16e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
@ -27849,6 +27853,7 @@
"gemini-embedding-001": {
"deprecation_date": "2028-05-20",
"input_cost_per_token": 1.5e-07,
"input_cost_per_token_batches": 1.2e-07,
"litellm_provider": "vertex_ai-embedding-models",
"max_input_tokens": 2048,
"max_tokens": 2048,
@ -28393,6 +28398,8 @@
"input_cost_per_image": 0.0011,
"input_cost_per_token": 2e-06,
"input_cost_per_token_batches": 1e-06,
"input_cost_per_token_flex": 1e-06,
"input_cost_per_token_priority": 3.6e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 32768,
@ -28404,6 +28411,8 @@
"rpm": 1000,
"tpm": 4000000,
"output_cost_per_token_batches": 6e-06,
"output_cost_per_token_flex": 6e-06,
"output_cost_per_token_priority": 2.16e-05,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@ -41390,20 +41399,20 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro": {
"input_cost_per_token": 9.1263e-07,
"input_cost_per_token": 8.44944e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 1.82526e-06,
"output_cost_per_token": 1.689888e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"cache_read_input_token_cost": 7.60525e-08,
"cache_read_input_token_cost": 7.0412e-08,
"supports_audio_input": false,
"supports_pdf_input": false,
"supports_vision": false,
@ -49411,12 +49420,10 @@
},
"vertex_ai/gemini-2.5-flash-image": {
"deprecation_date": "2027-03-15",
"cache_read_input_token_cost": 3e-08,
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 3e-07,
"input_cost_per_token_batches": 1.5e-07,
"input_cost_per_token_flex": 1.5e-07,
"input_cost_per_token_priority": 5.4e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
@ -49494,10 +49501,14 @@
"vertex_ai/gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_priority": 3.6e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_priority": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -49507,7 +49518,9 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_priority": 2.16e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"supports_reasoning": false,
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
@ -56202,6 +56215,7 @@
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_per_second": 3.3333333333333335e-05,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -56394,6 +56408,7 @@
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_per_second": 3.3333333333333335e-05,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -60135,6 +60150,23 @@
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/routers/deepseek-v4p1-flash-us": {
"cache_read_input_token_cost": 9e-09,
"input_cost_per_token": 4.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 1.8e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-vision-exp": {
"cache_read_input_token_cost": 7e-09,
"deprecation_date": "2026-09-25",
@ -60214,6 +60246,23 @@
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/deepseek-v4p1-flash-us": {
"cache_read_input_token_cost": 9e-09,
"input_cost_per_token": 4.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 1.8e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/deepseek-v4-flash-vision-exp": {
"cache_read_input_token_cost": 7e-09,
"deprecation_date": "2026-09-25",
@ -60310,13 +60359,16 @@
},
"fireworks_ai/kimi-k3-us": {
"cache_read_input_token_cost": 4.5e-07,
"cache_read_input_token_cost_priority": 5.625e-07,
"input_cost_per_token": 4.5e-06,
"input_cost_per_token_priority": 5.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.25e-05,
"output_cost_per_token_priority": 2.8125e-05,
"reasoning_effort_levels": [
"low",
"high",
@ -60518,13 +60570,16 @@
},
"fireworks_ai/accounts/fireworks/routers/kimi-k3-us": {
"cache_read_input_token_cost": 4.5e-07,
"cache_read_input_token_cost_priority": 5.625e-07,
"input_cost_per_token": 4.5e-06,
"input_cost_per_token_priority": 5.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.25e-05,
"output_cost_per_token_priority": 2.8125e-05,
"reasoning_effort_levels": [
"low",
"high",
@ -63283,6 +63338,22 @@
"supports_tool_choice": true,
"supports_vision": false
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"input_cost_per_token": 2.1e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": false
},
"fireworks_ai/glm-5p3": {
"cache_read_input_token_cost": 2.6e-07,
"cache_read_input_token_cost_priority": 3.25e-07,
@ -63302,6 +63373,22 @@
"supports_tool_choice": true,
"supports_vision": false
},
"fireworks_ai/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"input_cost_per_token": 2.1e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": false
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-fast": {
"cache_read_input_token_cost": 3.9e-07,
"input_cost_per_token": 2.1e-06,
@ -63349,6 +63436,20 @@
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"input_cost_per_token": 2.25e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/glm-5p3-flash": {
"cache_read_input_token_cost": 3e-08,
"cache_read_input_token_cost_priority": 3.75e-08,
@ -63366,6 +63467,20 @@
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"input_cost_per_token": 2.25e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/models/inkling": {
"cache_read_input_token_cost": 1.7e-07,
"input_cost_per_token": 1e-06,
@ -65897,9 +66012,9 @@
"supports_web_search": false
},
"openrouter/z-ai/glm-5.3-flash": {
"input_cost_per_token": 1.5e-07,
"output_cost_per_token": 5e-07,
"cache_read_input_token_cost": 5e-08,
"input_cost_per_token": 4.5e-08,
"output_cost_per_token": 6e-07,
"cache_read_input_token_cost": 2.85e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
"max_output_tokens": 943718,
@ -66648,9 +66763,9 @@
"supports_web_search": true
},
"openrouter/deepseek/deepseek-v4-flash": {
"input_cost_per_token": 8.4e-08,
"output_cost_per_token": 1.68e-07,
"cache_read_input_token_cost": 1.68e-08,
"input_cost_per_token": 4.9e-08,
"output_cost_per_token": 9.8e-08,
"cache_read_input_token_cost": 9.8e-09,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
@ -69475,6 +69590,7 @@
"input_cost_per_audio_token": 3e-06,
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -69495,6 +69611,7 @@
"input_cost_per_audio_token": 3e-06,
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -73039,6 +73156,26 @@
"supports_vision": true,
"supports_web_search": false
},
"openrouter/mistralai/mistral-large-2512": {
"cache_read_input_token_cost": 5e-08,
"input_cost_per_token": 5e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 209715,
"max_tokens": 209715,
"mode": "chat",
"output_cost_per_token": 1.5e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": false,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
},
"openrouter/mistralai/mistral-large-2512:batch": {
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_token": 2.5e-07,

View file

@ -1,16 +1,15 @@
from collections.abc import Awaitable, Callable, Coroutine, Mapping
from typing import Final, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
from collections.abc import Coroutine, Mapping
from typing import Final
import httpx
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import main
from litellm.ocr.main import convert_file_document_to_url_document, get_mime_type
from litellm.rust_bridge import runtime
from litellm.rust_bridge.catalog import Route, RouteContext
from litellm.rust_bridge.dispatch import PublicDispatch, call_hook
from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, LiteLLMOcrRequest
__all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr")
__all__ = ("aocr", "ocr")
def _bind_request(
@ -42,16 +41,6 @@ def _public_request(name: str, args: tuple[object, ...], kwargs: Mapping[str, ob
raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None
_PYTHON_OCR: Final = cast( # cast-ok: forward the original call shape through the Python @client decorator
Callable[..., OCRResponse | Coroutine[object, object, OCRResponse]],
main.ocr, # noqa: TID251 # dispatch boundary owns this Python fallback
)
_PYTHON_AOCR: Final = cast( # cast-ok: forward the original call shape through the Python @client decorator
Callable[..., Awaitable[OCRResponse]],
main.aocr, # noqa: TID251 # dispatch boundary owns this Python fallback
)
def _context(request: LiteLLMOcrRequest) -> RouteContext:
prefix, separator, _ = request.model.partition("/")
provider: Final = request.custom_llm_provider or (prefix if separator else None)
@ -79,7 +68,7 @@ def ocr(
return _DISPATCH.run(
args,
kwargs,
python=_PYTHON_OCR,
python=runtime.NO_PYTHON,
binding=NATIVE_OCR,
native=call_hook,
)
@ -89,7 +78,7 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr
return await _ADISPATCH.arun(
args,
kwargs,
python=_PYTHON_AOCR,
python=runtime.NO_PYTHON,
binding=NATIVE_AOCR,
native=call_hook,
)

View file

@ -1,421 +0,0 @@
"""
Main OCR function for LiteLLM.
"""
import asyncio
import base64
import mimetypes
import os
import re
from collections.abc import Coroutine, Mapping
from dataclasses import dataclass
from io import IOBase
from types import MappingProxyType
from typing import Final, Protocol, cast # noqa: TID251 # adapters preserve the legacy untyped contracts
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.constants import request_timeout
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.ocr.transformation import (
OCR_REQUEST_FORMAT_PARAM,
BaseOCRConfig,
OCRResponse,
parse_ocr_request_format,
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import CustomPricingLiteLLMParams
from litellm.utils import ProviderConfigManager, client
base_llm_http_handler: Final = BaseLLMHTTPHandler()
class FileReader(Protocol):
def read(self) -> bytes | str: ...
@dataclass(frozen=True, slots=True)
class _PreparedOCRRequest:
model: str
document: Mapping[str, object]
api_key: str | None
api_base: str | None
custom_llm_provider: str
extra_headers: dict[str, object] | None
provider_config: BaseOCRConfig
optional_params: dict[str, object]
litellm_params: dict[str, object]
effective_timeout: float | httpx.Timeout
litellm_logging_obj: LiteLLMLoggingObj
def _prepare_ocr_request(
model: str,
document: Mapping[str, object],
api_key: str | None,
api_base: str | None,
timeout: float | httpx.Timeout | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
kwargs: dict[str, object],
) -> _PreparedOCRRequest:
litellm_logging_obj: Final = cast( # cast-ok: @client supplies the logging object; preserve legacy failure behavior
LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj")
)
litellm_call_id: Final = cast( # cast-ok: @client supplies the call id without coercion
str | None, kwargs.get("litellm_call_id", None)
)
if not isinstance(document, dict):
raise litellm.BadRequestError(
message="document must be a dict with 'type' and URL/file field",
model=model,
llm_provider=_error_provider(model, custom_llm_provider) or "",
)
normalized_document: Final = (
convert_file_document_to_url_document(document) if document.get("type") == "file" else document
)
doc_type: Final = normalized_document.get("type")
if doc_type not in ("document_url", "image_url"):
raise litellm.BadRequestError(
message=f"Invalid document type: {doc_type}. Must be 'document_url', 'image_url', or 'file'",
model=model,
llm_provider=_error_provider(model, custom_llm_provider) or "",
)
if not normalized_document.get(doc_type):
raise litellm.BadRequestError(
message="Document URL is required",
model=model,
llm_provider=_error_provider(model, custom_llm_provider) or "",
)
(
model,
custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
)
ocr_provider_config: Final = ProviderConfigManager.get_provider_ocr_config(
model=model,
provider=litellm.LlmProviders(custom_llm_provider),
)
if ocr_provider_config is None:
raise ValueError(f"OCR is not supported for provider: {custom_llm_provider}")
resolved_api_key, resolved_api_base = ocr_provider_config.resolve_connection_params(
api_key=api_key,
api_base=api_base,
dynamic_api_key=dynamic_api_key,
dynamic_api_base=dynamic_api_base,
)
verbose_logger.debug("OCR call - model: %s, provider: %s", model, custom_llm_provider)
litellm_params: Final = GenericLiteLLMParams.model_validate(kwargs)
supported_params: Final = ocr_provider_config.get_supported_ocr_params(model=model)
requested_format: Final = kwargs.get(OCR_REQUEST_FORMAT_PARAM)
if requested_format is not None:
try:
parse_ocr_request_format(requested_format)
except ValueError as e:
raise litellm.exceptions.UnsupportedParamsError(
message=f"{e}", model=model, llm_provider=custom_llm_provider
) from e
non_default_params: Final = {param: kwargs.pop(param) for param in supported_params if param in kwargs}
try:
mapped_params: Final = ocr_provider_config.map_ocr_params(
non_default_params=non_default_params,
optional_params={},
model=model,
)
except ValueError as error:
raise litellm.BadRequestError(message=str(error), model=model, llm_provider=custom_llm_provider) from error
optional_params: Final = (
mapped_params if requested_format is None else {**mapped_params, OCR_REQUEST_FORMAT_PARAM: requested_format}
)
verbose_logger.debug("OCR optional_params after mapping: %s", optional_params)
effective_timeout: Final = timeout or request_timeout
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params={
"litellm_call_id": litellm_call_id,
"api_base": resolved_api_base,
**litellm_params.model_dump(include=frozenset(CustomPricingLiteLLMParams.model_fields), exclude_none=True),
},
custom_llm_provider=custom_llm_provider,
)
return _PreparedOCRRequest(
model=model,
document=normalized_document,
api_key=resolved_api_key,
api_base=resolved_api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
provider_config=ocr_provider_config,
optional_params=cast(dict[str, object], optional_params),
litellm_params=dict(litellm_params),
effective_timeout=effective_timeout,
litellm_logging_obj=litellm_logging_obj,
)
def _error_provider(model: str, custom_llm_provider: str | None) -> str | None:
if custom_llm_provider is not None:
return custom_llm_provider
prefix: Final = model.partition("/")[0]
if prefix in ("mistral", "azure_ai", "vertex_ai"):
return prefix
return "mistral" if model.startswith("mistral-ocr") else None
@client
async def aocr(
model: str,
document: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: dict[str, object] | None = None,
**kwargs: object, # kwargs-ok: public OCR accepts provider-specific options
) -> OCRResponse:
completion_kwargs: Final[dict[str, object]] = {
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"timeout": timeout,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"kwargs": kwargs,
}
try:
prepared: Final = _prepare_ocr_request(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
)
model = prepared.model
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update(model=model, custom_llm_provider=custom_llm_provider)
response = base_llm_http_handler.ocr(
model=prepared.model,
document=cast( # cast-ok: preserve legacy document fields for provider validation
dict[str, str], prepared.document
),
optional_params=prepared.optional_params,
timeout=prepared.effective_timeout,
logging_obj=prepared.litellm_logging_obj,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared.custom_llm_provider,
aocr=True,
headers=prepared.extra_headers,
provider_config=prepared.provider_config,
litellm_params=prepared.litellm_params,
)
if asyncio.iscoroutine(response):
response = await response
if response is None:
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
return response
except Exception as e:
error_provider: Final = _error_provider(model, custom_llm_provider)
error_model: Final = model.removeprefix(f"{error_provider}/") if error_provider else model
raise litellm.exception_type(
model=error_model,
custom_llm_provider=error_provider,
original_exception=e,
completion_kwargs=completion_kwargs,
extra_kwargs=kwargs,
)
_MIME_PATTERN: Final = re.compile(r"^[\w.+-]+/[\w.+-]+$")
_MIME_TYPE_MAP: Final = MappingProxyType(
{
".pdf": "application/pdf",
".png": "image/png",
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".gif": "image/gif",
".webp": "image/webp",
".tiff": "image/tiff",
".tif": "image/tiff",
".bmp": "image/bmp",
}
)
def get_mime_type(file_path: str) -> str:
ext: Final = os.path.splitext(file_path)[1].lower()
mime: Final = _MIME_TYPE_MAP.get(ext)
if mime:
return mime
guessed, _ = mimetypes.guess_type(file_path)
return guessed or "application/octet-stream"
def _read_file(file_input: object) -> tuple[bytes, str, str | None]:
if isinstance(file_input, str):
raise ValueError(
"OCR file input does not accept bare str values. Pass bytes, "
"a pathlib.Path, or a file-like object. To OCR a local file "
"from a path, call open(path, 'rb') yourself."
)
if isinstance(file_input, os.PathLike):
file_path: Final = str(cast(object, file_input)) # cast-ok: preserve staging's str(PathLike) conversion
if not os.path.isfile(file_path):
raise FileNotFoundError(f"File not found: {file_path}")
mime_type: Final = get_mime_type(file_path)
with open(file_path, "rb") as stream:
return stream.read(), mime_type, os.path.basename(file_path)
if isinstance(file_input, bytes):
return file_input, "application/octet-stream", None
if isinstance(file_input, IOBase) or hasattr(file_input, "read"):
file_name: Final = cast( # cast-ok: retain legacy validation and errors for file-like metadata
str | None, getattr(file_input, "name", None)
)
inferred_mime: Final = get_mime_type(file_name) if file_name else "application/octet-stream"
reader: Final = cast(FileReader, file_input) # cast-ok: legacy accepts duck-typed file readers
content: Final = reader.read()
return content.encode("utf-8") if isinstance(content, str) else content, inferred_mime, file_name
raise ValueError(
f"Unsupported file input type: {type(file_input)}. Expected pathlib.Path, bytes, or a file-like object."
)
def convert_file_document_to_url_document(document: Mapping[str, object]) -> dict[str, str]:
file_input: Final = document.get("file")
if file_input is None:
raise ValueError(
"document with type='file' must include a 'file' field containing "
"a pathlib.Path, file-like object, or bytes"
)
file_bytes, inferred_mime, file_name = _read_file(file_input)
if not file_bytes:
raise ValueError("File is empty or could not be read")
mime_type: Final = cast( # cast-ok: keep staging's MIME validation errors
str, document.get("mime_type", inferred_mime)
)
if not _MIME_PATTERN.match(mime_type):
raise ValueError(f"Invalid MIME type: {mime_type}")
base64_data: Final = base64.b64encode(file_bytes).decode("utf-8")
data_uri: Final = f"data:{mime_type};base64,{base64_data}"
if mime_type.startswith("image/"):
verbose_logger.debug(
"OCR file input: Converted file to image_url data URI (mime=%s, size=%s bytes, name=%s)",
mime_type,
len(file_bytes),
file_name,
)
return {"type": "image_url", "image_url": data_uri}
verbose_logger.debug(
"OCR file input: Converted file to document_url data URI (mime=%s, size=%s bytes, name=%s)",
mime_type,
len(file_bytes),
file_name,
)
return {"type": "document_url", "document_url": data_uri}
@client
def ocr(
model: str,
document: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: dict[str, object] | None = None,
**kwargs: object, # kwargs-ok: public OCR accepts provider-specific options
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
completion_kwargs: Final[dict[str, object]] = {
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"timeout": timeout,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"kwargs": kwargs,
}
try:
_is_async: Final = kwargs.pop("aocr", False) is True
completion_kwargs["aocr"] = _is_async
prepared: Final = _prepare_ocr_request(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
kwargs=kwargs,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout=timeout,
)
model = prepared.model
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update(model=model, custom_llm_provider=custom_llm_provider)
response: Final = base_llm_http_handler.ocr(
model=prepared.model,
document=cast( # cast-ok: preserve legacy document fields for provider validation
dict[str, str], prepared.document
),
optional_params=prepared.optional_params,
timeout=prepared.effective_timeout,
logging_obj=prepared.litellm_logging_obj,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared.custom_llm_provider,
aocr=_is_async,
headers=prepared.extra_headers,
provider_config=prepared.provider_config,
litellm_params=prepared.litellm_params,
)
return response
except Exception as e:
error_provider: Final = _error_provider(model, custom_llm_provider)
error_model: Final = model.removeprefix(f"{error_provider}/") if error_provider else model
raise litellm.exception_type(
model=error_model,
custom_llm_provider=error_provider,
original_exception=e,
completion_kwargs=completion_kwargs,
extra_kwargs=kwargs,
)

View file

@ -2057,7 +2057,7 @@ class MCPRequestHandler:
@staticmethod
async def _get_team_object_permission(
user_api_key_auth: UserAPIKeyAuth | None = None,
):
) -> LiteLLM_ObjectPermissionTable | None:
"""
Get team object_permission - automatically loaded by get_team_object() in main auth flow.
@ -2289,6 +2289,10 @@ class MCPRequestHandler:
)
)
allowed_tools = _as_list(
await MCPRequestHandler._apply_agent_caller_tool_ceiling(allowed_tools, server_id, user_api_key_auth)
)
return await MCPRequestHandler._apply_agent_and_org_tool_ceilings(
allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source
)
@ -3170,6 +3174,48 @@ class MCPRequestHandler:
return list(user_tools)
return list(set(allowed_tools) & set(user_tools))
@staticmethod
async def _apply_agent_caller_tool_ceiling(
allowed_tools: Sequence[str] | None,
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
) -> Sequence[str] | None:
"""Narrow an agent key's tools on ``server_id`` to those the invoking user and team (echoed back
by the agent as ``x-litellm-user-id`` / ``x-litellm-team-id``) may call: the echoed team's tool
grants when it names any on this server, then the echoed user's own tool entitlement. The tools
axis twin of ``_apply_agent_caller_ceiling``, so the headers only ever narrow. Denies every tool
on the server when the caller's team cannot be loaded, since a caller we cannot resolve must not
read as unrestricted."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
caller_auth: Final = agent_caller_auth(user_api_key_auth) if user_api_key_auth else None
if caller_auth is None:
return allowed_tools
try:
team_obj_perm: Final = await MCPRequestHandler._get_team_object_permission(caller_auth)
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
except Exception as e: # noqa: BLE001 # an unresolved caller team must deny, not widen
verbose_logger.warning(
"MCP agent caller team tool ceiling unresolvable, denying tools on %r: %s", server_id, e
)
return ()
team_direct_tools: Final = (
global_mcp_server_manager.expand_tool_permissions(team_obj_perm.mcp_tool_permissions).get(server_id)
if team_obj_perm
else None
)
team_tools: Final = MCPRequestHandler._union_tool_grants(team_direct_tools, team_toolset_tools)
team_capped: Final = (
allowed_tools
if team_tools is None
else tuple(team_tools)
if allowed_tools is None
else tuple(frozenset(allowed_tools) & frozenset(team_tools))
)
return await MCPRequestHandler._apply_user_tool_ceiling(team_capped, server_id, caller_auth)
@staticmethod
async def _apply_end_user_tool_ceiling(
allowed_tools: Sequence[str] | None,

View file

@ -2392,6 +2392,16 @@
],
"title": "Extra Headers"
},
"kill_switch": {
"anyOf": [
{
"$ref": "#/components/schemas/AgentKillSwitchConfig"
},
{
"type": "null"
}
]
},
"litellm_params": {
"additionalProperties": true,
"title": "Litellm Params",
@ -2561,6 +2571,221 @@
"title": "AgentKeySummary",
"type": "object"
},
"AgentKillSwitchApiKeyAuth": {
"additionalProperties": false,
"properties": {
"api_key": {
"title": "Api Key",
"type": "string"
},
"header_name": {
"default": "x-api-key",
"title": "Header Name",
"type": "string"
},
"type": {
"const": "api_key",
"title": "Type",
"type": "string"
}
},
"required": [
"type",
"api_key"
],
"title": "AgentKillSwitchApiKeyAuth",
"type": "object"
},
"AgentKillSwitchBasicAuth": {
"additionalProperties": false,
"properties": {
"password": {
"title": "Password",
"type": "string"
},
"type": {
"const": "basic",
"title": "Type",
"type": "string"
},
"username": {
"title": "Username",
"type": "string"
}
},
"required": [
"type",
"username",
"password"
],
"title": "AgentKillSwitchBasicAuth",
"type": "object"
},
"AgentKillSwitchBearerAuth": {
"additionalProperties": false,
"properties": {
"token": {
"title": "Token",
"type": "string"
},
"type": {
"const": "bearer",
"title": "Type",
"type": "string"
}
},
"required": [
"type",
"token"
],
"title": "AgentKillSwitchBearerAuth",
"type": "object"
},
"AgentKillSwitchConfig": {
"additionalProperties": false,
"description": "Webhook an admin fires to shut an agent down out of band. LiteLLM only\nmakes the call; whatever the endpoint does with it is the agent's business.",
"properties": {
"auth": {
"anyOf": [
{
"discriminator": {
"mapping": {
"api_key": "#/components/schemas/AgentKillSwitchApiKeyAuth",
"basic": "#/components/schemas/AgentKillSwitchBasicAuth",
"bearer": "#/components/schemas/AgentKillSwitchBearerAuth"
},
"propertyName": "type"
},
"oneOf": [
{
"$ref": "#/components/schemas/AgentKillSwitchBearerAuth"
},
{
"$ref": "#/components/schemas/AgentKillSwitchApiKeyAuth"
},
{
"$ref": "#/components/schemas/AgentKillSwitchBasicAuth"
}
]
},
{
"type": "null"
}
],
"title": "Auth"
},
"body": {
"anyOf": [
{
"additionalProperties": true,
"type": "object"
},
{
"type": "null"
}
],
"title": "Body"
},
"headers": {
"additionalProperties": {
"type": "string"
},
"title": "Headers",
"type": "object"
},
"method": {
"default": "POST",
"enum": [
"POST",
"PUT",
"PATCH",
"DELETE",
"GET"
],
"title": "Method",
"type": "string"
},
"query_params": {
"additionalProperties": {
"type": "string"
},
"title": "Query Params",
"type": "object"
},
"url": {
"title": "Url",
"type": "string"
}
},
"required": [
"url"
],
"title": "AgentKillSwitchConfig",
"type": "object"
},
"AgentKillSwitchResult": {
"properties": {
"agent_id": {
"title": "Agent Id",
"type": "string"
},
"error": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Error"
},
"method": {
"enum": [
"POST",
"PUT",
"PATCH",
"DELETE",
"GET"
],
"title": "Method",
"type": "string"
},
"response_body": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Response Body"
},
"status_code": {
"anyOf": [
{
"type": "integer"
},
{
"type": "null"
}
],
"title": "Status Code"
},
"url": {
"title": "Url",
"type": "string"
}
},
"required": [
"agent_id",
"url",
"method"
],
"title": "AgentKillSwitchResult",
"type": "object"
},
"AgentMakePublicResponse": {
"properties": {
"message": {
@ -2775,6 +3000,16 @@
],
"title": "Keys"
},
"kill_switch": {
"anyOf": [
{
"$ref": "#/components/schemas/AgentKillSwitchConfig"
},
{
"type": "null"
}
]
},
"litellm_params": {
"anyOf": [
{
@ -3574,6 +3809,16 @@
],
"title": "Extra Headers"
},
"kill_switch": {
"anyOf": [
{
"$ref": "#/components/schemas/AgentKillSwitchConfig"
},
{
"type": "null"
}
]
},
"litellm_params": {
"additionalProperties": true,
"title": "Litellm Params",
@ -4341,6 +4586,54 @@
]
}
},
"/v1/agents/{agent_id}/kill_switch": {
"post": {
"description": "Fire the agent's configured kill switch webhook. Proxy admin only.\n\nLiteLLM only makes the configured HTTP call and reports what came back; it\ndoes not change the agent's state in LiteLLM. Returns 200 when the webhook\nanswered 2xx, 502 with the same result body otherwise. Every attempt is\nwritten to the audit log as a `kill_switch_fired` row against the agent.\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/kill_switch\" \\\n -H \"Authorization: Bearer <your_api_key>\"\n```",
"operationId": "trigger_agent_kill_switch_v1_agents__agent_id__kill_switch_post",
"parameters": [
{
"in": "path",
"name": "agent_id",
"required": true,
"schema": {
"title": "Agent Id",
"type": "string"
}
}
],
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/AgentKillSwitchResult"
}
}
},
"description": "Successful Response"
},
"422": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
},
"description": "Validation Error"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Trigger Agent Kill Switch",
"tags": [
"agents"
]
}
},
"/v1/agents/{agent_id}/make_public": {
"post": {
"description": "Make an agent publicly discoverable\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/make_public\" \\\n -H \"Authorization: Bearer <your_api_key>\" \\\n -H \"Content-Type: application/json\"\n```\n\nExample Response:\n```json\n{\n \"agent_id\": \"123e4567-e89b-12d3-a456-426614174000\",\n \"agent_name\": \"my-custom-agent\",\n \"litellm_params\": {\n \"make_public\": true\n },\n \"agent_card_params\": {...},\n \"created_at\": \"2025-11-15T10:30:00Z\",\n \"updated_at\": \"2025-11-15T10:35:00Z\",\n \"created_by\": \"user123\",\n \"updated_by\": \"user123\"\n}\n```",

View file

@ -51,6 +51,7 @@ from litellm.types.proxy.carried_budget_state import (
UserBudgetSnapshot,
)
from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry
from litellm.types.proxy.spend_capture_rate import SpendCaptureRateCheckSettings
from litellm.types.router import RouterErrors, UpdateRouterConfig
from litellm.types.router_weights import validate_router_settings_dict
from litellm.types.secret_managers.main import KeyManagementSystem
@ -238,6 +239,7 @@ class LitellmTableNames(str, enum.Enum):
CONFIG_TABLE_NAME = "LiteLLM_Config"
SSO_CONFIG_TABLE_NAME = "LiteLLM_SSOConfig"
UI_SETTINGS_TABLE_NAME = "LiteLLM_UISettings"
AGENT_TABLE_NAME = "LiteLLM_AgentsTable"
class Litellm_EntityType(enum.Enum):
@ -578,6 +580,7 @@ class LiteLLMRoutes(enum.Enum):
"/v1/agents/{agent_id}",
"/v1/agents/make_public",
"/v1/agents/{agent_id}/make_public",
"/v1/agents/{agent_id}/kill_switch",
)
# Backwards-compat union — virtual keys may be configured with
@ -777,6 +780,7 @@ class LiteLLMRoutes(enum.Enum):
"/global/spend/provider",
"/global/spend/tags",
"/global/spend/all_tag_names",
"/spend/capture_rate",
]
public_routes = frozenset(
@ -2946,6 +2950,14 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"every replica. On by default; set to tune the window, pin a job, or turn it off."
),
)
spend_capture_rate_check: SpendCaptureRateCheckSettings | None = Field(
None,
description=(
"Daily check of the spend LiteLLM captured against the provider's own bill (OpenAI via OPENAI_ADMIN_KEY). "
"Publishes litellm_spend_capture_rate per provider and alerts when the ratio over the lookback window "
"falls under the threshold (default 0.9). Off unless set."
),
)
maximum_spend_logs_retention_period: str | None = Field(
None,
description="Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted.",
@ -3688,7 +3700,7 @@ from litellm.models.spend_logs import ( # noqa: E402
)
from litellm.models.tag import LiteLLM_TagTable as LiteLLM_TagTable # noqa: E402
AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "unblocked", "rotated"]
AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "unblocked", "rotated", "kill_switch_fired"]
class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase):

View file

@ -13,13 +13,14 @@ import litellm
from litellm.constants import REDACTED_BY_LITELM_STRING
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy.agent_endpoints.kill_switch import restore_kill_switch
from litellm.proxy.management_helpers.object_permission_utils import (
handle_update_object_permission_common,
)
from litellm.proxy.utils import PrismaClient
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import AgentsRepository, ObjectPermissionRepository
from litellm.types.agents import AgentConfig, AgentResponse, PatchAgentRequest
from litellm.types.agents import AgentConfig, AgentKillSwitchConfig, AgentResponse, PatchAgentRequest
if TYPE_CHECKING:
from prisma import models as prisma_models
@ -31,6 +32,10 @@ class AgentObjectPermissionRecord(Protocol):
def dict(self) -> dict[str, object]: ...
class AgentIdWhere(TypedDict):
agent_id: ReadOnly[str]
class AgentRecordDump(TypedDict):
agent_id: str
agent_name: str
@ -38,6 +43,7 @@ class AgentRecordDump(TypedDict):
agent_card_params: dict[str, object]
static_headers: dict[str, str] | None
extra_headers: list[str] | None
kill_switch: ReadOnly[AgentKillSwitchConfig | None]
access_group_ids: ReadOnly[Sequence[str] | None]
object_permission: dict[str, object] | None
spend: float
@ -70,6 +76,9 @@ class AgentRecord(Protocol):
@property
def access_group_ids(self) -> Sequence[str] | None: ...
@property
def kill_switch(self) -> Mapping[str, object] | None: ...
@property
def spend(self) -> float: ...
@ -211,6 +220,29 @@ def parse_agent_litellm_params(value: object) -> Mapping[str, object]:
return _EMPTY_LITELLM_PARAMS
_KILL_SWITCH_ADAPTER: Final[TypeAdapter[AgentKillSwitchConfig | None]] = TypeAdapter(AgentKillSwitchConfig | None)
def parse_agent_kill_switch(value: object) -> AgentKillSwitchConfig | None:
if value is None:
return None
try:
if isinstance(value, str):
return _KILL_SWITCH_ADAPTER.validate_json(value)
return _KILL_SWITCH_ADAPTER.validate_python(value)
except ValidationError:
return None
def serialize_agent_kill_switch(incoming: object, existing: object) -> str:
"""prisma-client-py drops ``None`` from update data, so a cleared kill switch is stored as the JSON literal
``null`` (read back as ``None``), the same convention ``memory_endpoints`` uses for ``Json?`` columns."""
restored: Final = restore_kill_switch(
_KILL_SWITCH_ADAPTER.validate_python(incoming), parse_agent_kill_switch(existing)
)
return safe_dumps(restored.model_dump() if restored is not None else None)
_MISSING_AGENT_PARAM: Final = object()
_RESTORE_AGENT_PARAMS_MAX_DEPTH: Final = 10
@ -293,6 +325,12 @@ def _patched_access_group_ids(agent: PatchAgentRequest) -> Mapping[str, object]:
return MappingProxyType({"access_group_ids": tuple(dict.fromkeys(agent.get("access_group_ids") or ()))})
def _patched_kill_switch(agent: PatchAgentRequest, existing: object) -> Mapping[str, object]:
if "kill_switch" not in agent:
return MappingProxyType({})
return MappingProxyType({"kill_switch": serialize_agent_kill_switch(agent.get("kill_switch"), existing)})
def _restore_redacted_litellm_params(
incoming: Mapping[str, object],
existing: Mapping[str, object],
@ -531,6 +569,7 @@ class AgentRegistry:
"agent_name": agent_name,
"litellm_params": litellm_params,
"agent_card_params": agent_card_params,
"kill_switch": serialize_agent_kill_switch(agent.get("kill_switch"), None),
"created_by": created_by,
"updated_by": created_by,
"created_at": datetime.now(timezone.utc),
@ -613,7 +652,10 @@ class AgentRegistry:
existing_agent: Final[Mapping[str, object]] = dict(existing_record)
augment_agent: Final = {**existing_agent, **agent}
update_data: Final[dict[str, object]] = {**_patched_access_group_ids(agent)}
update_data: Final[dict[str, object]] = {
**_patched_access_group_ids(agent),
**_patched_kill_switch(agent, existing_agent.get("kill_switch")),
}
if augment_agent.get("agent_name"):
update_data["agent_name"] = augment_agent.get("agent_name")
if "litellm_params" in agent:
@ -716,6 +758,9 @@ class AgentRegistry:
)
extra_headers_val_u: Final = agent.get("extra_headers") or []
access_group_ids_val_u: Final = tuple(dict.fromkeys(agent.get("access_group_ids") or ()))
kill_switch_val_u: Final = serialize_agent_kill_switch(
agent.get("kill_switch"), existing_row.kill_switch if existing_row is not None else None
)
update_data: Final[dict[str, object]] = {
"agent_name": agent_name,
@ -723,6 +768,7 @@ class AgentRegistry:
"agent_card_params": agent_card_params,
"static_headers": static_headers_val_u,
"extra_headers": extra_headers_val_u,
"kill_switch": kill_switch_val_u,
"access_group_ids": access_group_ids_val_u,
"updated_by": updated_by,
"updated_at": datetime.now(timezone.utc),

View file

@ -33,6 +33,8 @@ from litellm.proxy.a2a.agent_card import (
normalize_protocol_version,
)
from litellm.proxy.agent_endpoints.agent_registry import (
AgentIdWhere,
parse_agent_kill_switch,
parse_agent_litellm_params,
redact_sensitive_agent_litellm_params,
)
@ -45,6 +47,15 @@ from litellm.proxy.agent_endpoints.agent_search import (
search_agents,
)
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import accessible_agents
from litellm.proxy.agent_endpoints.kill_switch import (
KillSwitchAuditLogWriter,
KillSwitchHttpClient,
build_kill_switch_audit_log,
default_kill_switch_audit_log_writer,
default_kill_switch_http_client,
fire_kill_switch,
redact_kill_switch,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
@ -53,6 +64,8 @@ from litellm.types.agents import (
AgentCard,
AgentConfig,
AgentKeySummary,
AgentKillSwitchConfig,
AgentKillSwitchResult,
AgentMakePublicResponse,
AgentResponse,
MakeAgentsPublicRequest,
@ -160,9 +173,10 @@ def _redact_sensitive_agent_fields(
) -> list[AgentResponse]:
"""
Return copies of the given agents with credential-bearing litellm_params
values replaced by a fixed marker (never returned to ANY caller,
admin included) and, for non-admin callers, virtual-key and header
fields stripped entirely. The original objects are not modified.
values and kill-switch auth secrets replaced by a fixed marker (never
returned to ANY caller, admin included) and, for non-admin callers,
virtual-key, header and kill-switch fields stripped entirely. The original
objects are not modified.
"""
redacted: Final[list[AgentResponse]] = []
for agent in agents:
@ -171,8 +185,10 @@ def _redact_sensitive_agent_fields(
copy.static_headers = None
copy.extra_headers = None
copy.keys = None
copy.kill_switch = None
if copy.litellm_params:
copy.litellm_params = _redact_agent_litellm_params_dict(copy.litellm_params)
copy.kill_switch = redact_kill_switch(copy.kill_switch)
redacted.append(copy)
return redacted
@ -872,6 +888,74 @@ async def delete_agent(
raise HTTPException(status_code=500, detail=str(e))
@router.post(
"/v1/agents/{agent_id}/kill_switch",
tags=["[beta] A2A Agents"], # mutable-ok: fastapi types tags as list[str | Enum]
dependencies=(Depends(user_api_key_auth),),
response_model=AgentKillSwitchResult,
)
async def trigger_agent_kill_switch(
agent_id: str,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
http_client: Annotated[KillSwitchHttpClient, Depends(default_kill_switch_http_client)],
audit_log_writer: Annotated[KillSwitchAuditLogWriter, Depends(default_kill_switch_audit_log_writer)],
):
"""
Fire the agent's configured kill switch webhook. Proxy admin only.
LiteLLM only makes the configured HTTP call and reports what came back; it
does not change the agent's state in LiteLLM. Returns 200 when the webhook
answered 2xx, 502 with the same result body otherwise. Every attempt is
written to the audit log as a `kill_switch_fired` row against the agent.
Example Request:
```bash
curl -X POST "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/kill_switch" \\
-H "Authorization: Bearer <your_api_key>"
```
"""
from litellm.proxy.proxy_server import litellm_proxy_admin_name
await check_feature_access_for_user(user_api_key_dict, "agents")
_check_agent_management_permission(user_api_key_dict)
resolved: Final = await _resolve_agent_kill_switch(agent_id)
if resolved is None:
raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found")
resolved_agent_id, config = resolved
if config is None:
raise HTTPException(status_code=400, detail=f"Agent with ID {agent_id} has no kill_switch configured")
result: Final = await fire_kill_switch(agent_id=resolved_agent_id, config=config, http_client=http_client)
await audit_log_writer(
build_kill_switch_audit_log(
result=result,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
)
)
if not result.succeeded:
raise HTTPException(status_code=502, detail=result.model_dump())
return result
async def _resolve_agent_kill_switch(agent_id: str) -> tuple[str, AgentKillSwitchConfig | None] | None:
"""The DB row wins over this replica's in-memory registry so a trigger never fires a webhook another
replica has since changed; config.yaml agents have no row and fall back to the registry."""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is not None:
where: Final[AgentIdWhere] = {"agent_id": agent_id}
row: Final = await agents_table(prisma_client).find_unique(where=where)
if row is not None:
return row.agent_id, parse_agent_kill_switch(row.kill_switch)
agent: Final = AGENT_REGISTRY.get_agent_by_id(agent_id=agent_id)
if agent is None:
return None
return agent.agent_id, agent.kill_switch
@router.post(
"/v1/agents/{agent_id}/make_public",
tags=["[beta] A2A Agents"],

View file

@ -0,0 +1,239 @@
from base64 import b64encode
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping
from dataclasses import dataclass
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Final, Protocol, TypeAlias
import httpx
from typing_extensions import assert_never
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import (
AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS,
AGENT_KILL_SWITCH_TIMEOUT_SECONDS,
REDACTED_BY_LITELM_STRING,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # its params arg is a bare dict in http_handler
)
from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames, UserAPIKeyAuth
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update, get_audit_log_changed_by
from litellm.types.agents import (
AgentKillSwitchApiKeyAuth,
AgentKillSwitchAuth,
AgentKillSwitchBasicAuth,
AgentKillSwitchBearerAuth,
AgentKillSwitchConfig,
AgentKillSwitchResult,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
def _with_auth(config: AgentKillSwitchConfig, auth: AgentKillSwitchAuth) -> AgentKillSwitchConfig:
return AgentKillSwitchConfig(
url=config.url,
method=config.method,
headers=config.headers,
query_params=config.query_params,
body=config.body,
auth=auth,
)
def redact_kill_switch(config: AgentKillSwitchConfig | None) -> AgentKillSwitchConfig | None:
if config is None or config.auth is None:
return config
return _with_auth(config, _redact_auth(config.auth))
def _redact_auth(auth: AgentKillSwitchAuth) -> AgentKillSwitchAuth:
match auth:
case AgentKillSwitchBearerAuth():
return AgentKillSwitchBearerAuth(type="bearer", token=REDACTED_BY_LITELM_STRING)
case AgentKillSwitchApiKeyAuth():
return AgentKillSwitchApiKeyAuth(
type="api_key", header_name=auth.header_name, api_key=REDACTED_BY_LITELM_STRING
)
case AgentKillSwitchBasicAuth():
return AgentKillSwitchBasicAuth(type="basic", username=auth.username, password=REDACTED_BY_LITELM_STRING)
case _:
assert_never(auth)
def restore_kill_switch(
incoming: AgentKillSwitchConfig | None,
existing: AgentKillSwitchConfig | None,
) -> AgentKillSwitchConfig | None:
"""Put the stored secret back behind an auth field echoed as the redaction
marker; a marker with no stored secret of the same auth type becomes ""."""
if incoming is None or incoming.auth is None:
return incoming
existing_auth: Final = existing.auth if existing is not None else None
return _with_auth(incoming, _restore_auth(incoming.auth, existing_auth))
def _restore_secret(incoming_value: str, existing_value: str | None) -> str:
if incoming_value != REDACTED_BY_LITELM_STRING:
return incoming_value
return existing_value if existing_value is not None else ""
def _restore_auth(incoming: AgentKillSwitchAuth, existing: AgentKillSwitchAuth | None) -> AgentKillSwitchAuth:
match incoming:
case AgentKillSwitchBearerAuth():
stored_token: Final = existing.token if isinstance(existing, AgentKillSwitchBearerAuth) else None
return AgentKillSwitchBearerAuth(type="bearer", token=_restore_secret(incoming.token, stored_token))
case AgentKillSwitchApiKeyAuth():
stored_key: Final = existing.api_key if isinstance(existing, AgentKillSwitchApiKeyAuth) else None
return AgentKillSwitchApiKeyAuth(
type="api_key",
header_name=incoming.header_name,
api_key=_restore_secret(incoming.api_key, stored_key),
)
case AgentKillSwitchBasicAuth():
stored_password: Final = existing.password if isinstance(existing, AgentKillSwitchBasicAuth) else None
return AgentKillSwitchBasicAuth(
type="basic",
username=incoming.username,
password=_restore_secret(incoming.password, stored_password),
)
case _:
assert_never(incoming)
@dataclass(frozen=True, slots=True)
class KillSwitchRequest:
method: str
url: str
headers: Mapping[str, str]
json_body: Mapping[str, object] | None
def _auth_headers(auth: AgentKillSwitchAuth | None) -> Mapping[str, str]:
match auth:
case None:
return MappingProxyType({})
case AgentKillSwitchBearerAuth():
return MappingProxyType({"Authorization": f"Bearer {auth.token}"})
case AgentKillSwitchApiKeyAuth():
return MappingProxyType({auth.header_name: auth.api_key})
case AgentKillSwitchBasicAuth():
credentials: Final = b64encode(f"{auth.username}:{auth.password}".encode()).decode()
return MappingProxyType({"Authorization": f"Basic {credentials}"})
case _:
assert_never(auth)
def build_kill_switch_request(config: AgentKillSwitchConfig) -> KillSwitchRequest:
url: Final = httpx.URL(config.url).copy_merge_params(config.query_params)
return KillSwitchRequest(
method=config.method,
url=str(url),
headers=MappingProxyType({**config.headers, **_auth_headers(config.auth)}),
json_body=config.body,
)
class KillSwitchHttpClient(Protocol):
def build_request(
self,
method: str,
url: str,
*,
headers: Mapping[str, str],
json: Mapping[str, object] | None,
timeout: float,
) -> httpx.Request: ...
async def send(self, request: httpx.Request, *, stream: bool, follow_redirects: bool) -> httpx.Response: ...
def default_kill_switch_http_client() -> KillSwitchHttpClient:
return get_async_httpx_client(llm_provider=httpxSpecialProvider.AgentKillSwitch).client
KillSwitchAuditLogWriter: TypeAlias = Callable[[LiteLLM_AuditLogs], Awaitable[None]] # mutable-ok: Callable params
def default_kill_switch_audit_log_writer() -> KillSwitchAuditLogWriter:
return create_audit_log_for_update
def build_kill_switch_audit_log(
*,
result: AgentKillSwitchResult,
user_api_key_dict: UserAPIKeyAuth,
litellm_proxy_admin_name: str | None,
) -> LiteLLM_AuditLogs:
return LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=get_audit_log_changed_by(
litellm_changed_by=None,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.AGENT_TABLE_NAME,
object_id=result.agent_id,
action="kill_switch_fired",
updated_values=result.model_dump_json(exclude_none=True),
)
async def fire_kill_switch(
*,
agent_id: str,
config: AgentKillSwitchConfig,
http_client: KillSwitchHttpClient,
timeout: float = AGENT_KILL_SWITCH_TIMEOUT_SECONDS,
) -> AgentKillSwitchResult:
request: Final = build_kill_switch_request(config)
reported_url: Final = str(httpx.URL(request.url).copy_with(query=None))
verbose_proxy_logger.info("Firing kill switch for agent %s: %s %s", agent_id, request.method, reported_url)
try:
response: Final = await http_client.send(
http_client.build_request(
request.method,
request.url,
headers=request.headers,
json=request.json_body,
timeout=timeout,
),
stream=True,
follow_redirects=False,
)
body: Final = await _read_text_prefix(response, AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS)
except httpx.HTTPError as exc:
verbose_proxy_logger.warning("Kill switch for agent %s failed: %s", agent_id, type(exc).__name__)
return AgentKillSwitchResult(
agent_id=agent_id,
url=reported_url,
method=config.method,
error=type(exc).__name__,
)
return AgentKillSwitchResult(
agent_id=agent_id,
url=reported_url,
method=config.method,
status_code=response.status_code,
response_body=body,
)
async def _read_text_prefix(response: httpx.Response, max_chars: int) -> str:
try:
return await _take_text(response.aiter_text(), max_chars)
finally:
await response.aclose()
async def _take_text(chunks: AsyncIterator[str], max_chars: int) -> str:
taken = "" # rebind-ok: running prefix of a stream that is abandoned once the cap is hit
async for chunk in chunks:
taken += chunk # rebind-ok: see above
if len(taken) >= max_chars:
break
return taken[:max_chars]

View file

@ -495,11 +495,12 @@ class RateLimitStatus(TypedDict):
limit_remaining: int
rate_limit_type: Literal["requests", "tokens", "max_parallel_requests"]
descriptor_key: str
# Only populated by the atomic_check_and_increment_by_n path. A caller
# matching a status back to its descriptor must key on (descriptor_key,
# descriptor_value) when this is present, not descriptor_key alone --
# e.g. a batch charging several models' project ITPM/OTPM in one call
# produces multiple statuses sharing the same descriptor_key.
# Populated by the atomic_check_and_increment_by_n and windowed
# sliding-window paths. A caller matching a status back to its
# descriptor must key on (descriptor_key, descriptor_value) when this
# is present, not descriptor_key alone -- e.g. a batch charging several
# models' project ITPM/OTPM in one call, or a request carrying multiple
# rate-limited tags, produces statuses sharing the same descriptor_key.
descriptor_value: NotRequired[ReadOnly[str]]
@ -529,6 +530,7 @@ class WindowKeyMetadata(TypedDict):
tokens_limit: int | None
window_size: int
descriptor_key: str
descriptor_value: ReadOnly[str]
class AtomicCounterMeta(TypedDict):
@ -607,6 +609,47 @@ class RequestRateLimiterStash:
batch_enqueued_reservation: BatchEnqueuedTokenReservation | None = None
batch_tpd_refund_ops: tuple[ReservationAwareIncrementOperation, ...] = ()
reservation_released: bool = False
tpm_limited_tags: frozenset[str] = field(default_factory=frozenset)
@dataclass(frozen=True, slots=True)
class TagRateLimit:
rpm_limit: int | None
tpm_limit: int | None
class TagRateLimitResolver(Protocol):
def __call__(self, tag_names: Sequence[str], /) -> Awaitable[Mapping[str, TagRateLimit]]: ...
def _tag_rate_limit_descriptor(tag: str, limit: TagRateLimit, window_size: int) -> RateLimitDescriptor:
rate_limit: Final[RateLimitDescriptorRateLimitObject] = {
"requests_per_unit": limit.rpm_limit,
"tokens_per_unit": limit.tpm_limit,
"window_size": window_size,
}
return RateLimitDescriptor(key="tag", value=tag, rate_limit=rate_limit)
async def resolve_tag_rate_limits_from_db(tag_names: Sequence[str]) -> Mapping[str, TagRateLimit]:
from litellm.proxy.auth.auth_checks import get_tag_objects_batch
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
if prisma_client is None or not tag_names:
return MappingProxyType({})
tag_objects: Final = await get_tag_objects_batch(
tag_names=tag_names,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
return MappingProxyType(
{
tag_name: TagRateLimit(rpm_limit=budget.rpm_limit, tpm_limit=budget.tpm_limit)
for tag_name, tag_object in tag_objects.items()
if (budget := tag_object.litellm_budget_table) is not None
and (budget.rpm_limit is not None or budget.tpm_limit is not None)
}
)
_request_stash: Final[ContextVar[RequestRateLimiterStash | None]] = ContextVar(
@ -672,6 +715,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
self,
internal_usage_cache: InternalUsageCache,
time_provider: Callable[[], datetime] | None = None,
tag_rate_limit_resolver: TagRateLimitResolver = resolve_tag_rate_limits_from_db,
model_group_resolver: Callable[[str], str | None] = _resolve_model_group_alias_via_proxy_router,
ptu_team_ceiling_resolver: Callable[
[str, str], PTUTeamCeiling | None
@ -679,6 +723,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
):
self.internal_usage_cache = internal_usage_cache
self._time_provider = time_provider or datetime.now
self._tag_rate_limit_resolver = tag_rate_limit_resolver
self._model_group_resolver = model_group_resolver
self._ptu_team_ceiling_resolver = ptu_team_ceiling_resolver
if self.internal_usage_cache.dual_cache.redis_cache is not None:
@ -1238,6 +1283,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"limit_remaining": limit_remaining,
"rate_limit_type": rate_limit_type,
"descriptor_key": key_metadata[window_key]["descriptor_key"],
"descriptor_value": key_metadata[window_key]["descriptor_value"],
}
)
@ -1542,6 +1588,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"tokens_limit": int(tokens_limit) if tokens_limit is not None else None,
"window_size": int(window_size),
"descriptor_key": descriptor_key,
"descriptor_value": descriptor_value,
}
return keys_to_fetch, key_metadata, gauges
@ -2790,6 +2837,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return descriptors
async def _create_tag_rate_limit_descriptors(self, data: Mapping[str, object]) -> tuple[RateLimitDescriptor, ...]:
tags: Final = tuple(dict.fromkeys(get_tags_from_request_body(data)))
if not tags:
return ()
tag_limits: Final = await self._tag_rate_limit_resolver(tags)
return tuple(
_tag_rate_limit_descriptor(tag, limit, self.window_size)
for tag in tags
if (limit := tag_limits.get(tag)) is not None
)
def _create_rate_limit_descriptors(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -3196,7 +3254,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if status["code"] == "OVER_LIMIT":
descriptor_key = status["descriptor_key"]
matching_descriptor = next(
(desc for desc in descriptors if desc["key"] == descriptor_key),
(
desc
for desc in descriptors
if desc["key"] == descriptor_key
and ((status_value := status.get("descriptor_value")) is None or desc["value"] == status_value)
),
None,
)
descriptor_value = matching_descriptor["value"] if matching_descriptor is not None else "unknown"
@ -3610,6 +3673,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return [ # mutable-ok: the shared generation reservation helpers require a list
*descriptors,
*self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model),
*await self._create_tag_rate_limit_descriptors(data),
]
async def _release_request_capacity_when_admitted(
@ -3704,6 +3768,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
data=request_data,
call_type=call_type,
)
stash.tpm_limited_tags = frozenset(
d["value"]
for d in descriptors
if d["key"] == "tag" and d["rate_limit"] is not None and d["rate_limit"].get("tokens_per_unit") is not None
)
# Only check rate limits if we have descriptors with actual limits
if descriptors:
@ -4424,6 +4493,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
standard_logging_metadata: dict[str, Any],
kwargs: object,
model_group: str | None,
tpm_limited_tags: Set[str] = frozenset(),
) -> list[tuple[str, str]]:
"""
Enumerate every (scope_key, scope_value) pair that *might* carry a
@ -4479,6 +4549,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
targets.append(("agent", agent_id))
if session_id:
targets.append(("agent_session", f"{agent_id}:{session_id}"))
targets.extend(("tag", tag) for tag in sorted(tpm_limited_tags))
return targets
def _build_reservation_aware_tpm_ops(
@ -4653,6 +4724,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
targets: Final = self._collect_tpm_scope_targets(
standard_logging_metadata=standard_logging_metadata,
kwargs=kwargs,
tpm_limited_tags=stash.tpm_limited_tags if stash is not None else frozenset(),
model_group=reconcile_model.group if reconcile_model is not None else None,
)
charged_targets: Final = (

View file

@ -1,5 +1,6 @@
#### OCR Endpoints #####
import io
import json
from collections.abc import Mapping
from typing import Final, cast
@ -15,7 +16,6 @@ from litellm.llms.base_llm.ocr.transformation import (
OCRResponse,
parse_ocr_request_format,
)
from litellm.ocr.main import convert_file_document_to_url_document, get_mime_type
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
@ -24,20 +24,24 @@ router: Final = APIRouter()
_MAX_FILE_BYTES: Final = 50 * 1024 * 1024
class _NamedUpload(io.BytesIO):
name: str | None
def __init__(self, content: bytes, name: str | None) -> None:
super().__init__(content)
self.name = name
def _build_document_from_upload(
file_content: bytes,
filename: str | None,
content_type: str | None,
) -> dict[str, str]:
) -> dict[str, object]:
supplied_mime: Final = content_type.split(";")[0].strip() if content_type else None
mime_type: Final = (
get_mime_type(filename)
if filename and (not supplied_mime or supplied_mime == "application/octet-stream")
else supplied_mime
)
return convert_file_document_to_url_document(
{"type": "file", "file": file_content, "mime_type": mime_type or "application/octet-stream"}
)
upload: Final = _NamedUpload(file_content, filename)
if supplied_mime and supplied_mime != "application/octet-stream":
return {"type": "file", "file": upload, "mime_type": supplied_mime}
return {"type": "file", "file": upload}
def _with_request_format(data: Mapping[str, object], request: Request) -> Mapping[str, object]:

View file

@ -210,6 +210,25 @@ class ProxyInitializationHelpers:
response: Final = httpx.get(url=f"http://{host}:{port}/health")
print(json.dumps(response.json(), indent=4))
@staticmethod
def _run_config_validation(config: str | None) -> None:
if config is None:
raise click.UsageError("--validate_config requires --config <path>")
import asyncio
from litellm.proxy.proxy_server import ProxyConfig
async def _load() -> int:
_, model_list, _ = await ProxyConfig().load_config(router=None, config_file_path=config)
return len(model_list)
try:
model_count: Final = asyncio.run(_load())
except Exception as error:
click.echo(f"LiteLLM: config validation failed: {error}", err=True)
raise click.exceptions.Exit(1) from error
click.echo(f"LiteLLM: config OK ({model_count} models)")
@staticmethod
def _run_test_chat_completion(
host: str,
@ -887,6 +906,12 @@ class ProxyInitializationHelpers:
default=False,
help="Skip starting the server after setup (useful for migrations only)",
)
@click.option(
"--validate_config",
is_flag=True,
default=False,
help="Load and validate the config file (including mcp_servers) without starting the server, then exit. Exit code 1 on any config error.",
)
@click.option(
"--keepalive_timeout",
default=None,
@ -1027,6 +1052,7 @@ def run_server(
log_config,
use_prisma_db_push: bool,
skip_server_startup,
validate_config: bool,
keepalive_timeout,
timeout_worker_healthcheck,
max_requests_before_restart,
@ -1069,6 +1095,9 @@ def run_server(
if version is True:
ProxyInitializationHelpers._echo_litellm_version()
return
if validate_config is True:
ProxyInitializationHelpers._run_config_validation(config)
return
if model and "ollama" in model and api_base is None:
ProxyInitializationHelpers._run_ollama_serve()
if health is True:

View file

@ -152,6 +152,7 @@ from litellm.router_utils.auto_router_tuning_baseline import (
snapshot_tuning_baselines,
tuning_limit_violation,
)
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.router_utils.routing_groups import parse_routing_groups
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.utils import (
@ -292,6 +293,7 @@ from litellm.constants import (
REALTIME_SESSION_FAILURE_LOGGED_KEY,
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG,
SPEND_CAPTURE_RATE_CHECK_JOB_ID,
USER_SPEND_ALERTS_JOB_ID,
WEEKLY_SPEND_REPORT_JOB_ID,
)
@ -757,6 +759,9 @@ from litellm.proxy.spend_tracking.budget_reservation import (
from litellm.proxy.spend_tracking.daily_global_spend_rollup import (
run_scheduled_daily_global_spend_reconcile,
)
from litellm.proxy.spend_tracking.spend_capture_rate import (
run_scheduled_spend_capture_rate_check,
)
from litellm.proxy.spend_tracking.spend_counter_batch import (
PendingSpendIncrement,
active_spend_counter_batch,
@ -855,6 +860,7 @@ from litellm.types.proxy.model_deprecation import (
DEFAULT_DEPRECATION_WARN_DAYS,
ModelDeprecationResponse,
)
from litellm.types.proxy.spend_capture_rate import SpendCaptureProvider, SpendCaptureRateCheckSettings
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.router import (
ClassifierPlugin,
@ -5059,6 +5065,11 @@ def _bind_general_settings_store(settings: SettingsStore) -> None:
general_settings = settings # pyright: ignore[reportAssignmentType] # legacy global accepts mappings
def _current_general_settings() -> Mapping[str, object]:
"""The live ``general_settings``, whichever object a config reload has bound since the caller was created."""
return general_settings
@lru_cache(maxsize=4096)
def _log_ignored_cost_map_copy(model_id: str, fields: tuple[str, ...]) -> None:
verbose_proxy_logger.warning(
@ -10447,6 +10458,13 @@ class ProxyStartupEvent:
prisma_client=prisma_client,
)
cls._initialize_spend_capture_rate_check_job(
scheduler=scheduler,
proxy_logging_obj=proxy_logging_obj,
prisma_client=prisma_client,
read_general_settings=_current_general_settings,
)
### PTU DAILY ROLLUP ###
from litellm.proxy.spend_tracking.ptu_feature_flag import (
is_ptu_cost_attribution_enabled,
@ -10821,6 +10839,64 @@ class ProxyStartupEvent:
next_run_time=datetime.now(timezone.utc) + timedelta(minutes=2),
)
@classmethod
def _initialize_spend_capture_rate_check_job(
cls,
scheduler: AsyncIOScheduler,
proxy_logging_obj: ProxyLogging,
prisma_client: PrismaClient,
read_general_settings: Callable[[], Mapping[str, object]],
) -> None:
"""The job always runs and re-reads ``spend_capture_rate_check`` each run; an absent setting clears the gauge."""
cls._spend_capture_rate_check_settings(read_general_settings())
async def alert(message: str) -> None:
await proxy_logging_obj.alerting_handler(
message=message,
level="High",
alert_type=AlertType.failed_tracking_spend,
)
def publish(provider: str, capture_rate: float | None) -> None:
from litellm.integrations.prometheus import PrometheusLogger
for logger in litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=PrometheusLogger):
if isinstance(logger, PrometheusLogger):
logger.set_spend_capture_rate(api_provider=provider, capture_rate=capture_rate)
async def check() -> None:
settings: Final = cls._spend_capture_rate_check_settings(read_general_settings())
if settings is None:
for provider in get_args(SpendCaptureProvider):
publish(provider, None)
return
await run_scheduled_spend_capture_rate_check(
prisma_client,
settings,
pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager,
alert=alert,
publish=publish,
)
scheduler.add_job(
check,
"cron",
hour=1,
minute=15,
timezone="UTC",
id=SPEND_CAPTURE_RATE_CHECK_JOB_ID,
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
next_run_time=datetime.now(timezone.utc) + timedelta(minutes=2),
)
@staticmethod
def _spend_capture_rate_check_settings(
general_settings: Mapping[str, object],
) -> SpendCaptureRateCheckSettings | None:
raw_settings: Final = general_settings.get("spend_capture_rate_check")
return None if raw_settings is None else SpendCaptureRateCheckSettings.model_validate(raw_settings)
@classmethod
async def _initialize_slack_alerting_jobs(
cls,
@ -11104,6 +11180,40 @@ class ProxyStartupEvent:
#### API ENDPOINTS ####
async def _names_hidden_by_listing_callbacks(
user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
) -> frozenset[str]:
hidden: Final = await proxy_logging_obj.hidden_by_listing_callbacks(user_api_key_dict, model_names)
if not hidden or llm_router is None:
return hidden
aliases: Final = llm_router.model_group_alias
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings)
return hidden | frozenset(
alias
for alias in aliases
if (target := resolve_model_group_alias(aliases, alias)) is not None
and internal_to_public.get(target, target) in hidden
)
async def _entries_kept_by_listing_callbacks(
entries: Sequence[tuple[str, str]], user_api_key_dict: UserAPIKeyAuth
) -> tuple[tuple[str, str], ...]:
hidden: Final = await _names_hidden_by_listing_callbacks(
user_api_key_dict, tuple(response_id for response_id, _ in entries)
)
if not hidden:
return tuple(entries)
return tuple(entry for entry in entries if entry[0] not in hidden)
async def _deployment_hidden_by_listing_callbacks(deployment: Deployment, user_api_key_dict: UserAPIKeyAuth) -> bool:
listed_name: Final = _translate_model_name_for_response(deployment.model_dump(exclude_none=True)).get("model_name")
if not isinstance(listed_name, str):
return False
return listed_name in await _names_hidden_by_listing_callbacks(user_api_key_dict, (listed_name,))
@router.get("/v1/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"])
@router.get(
"/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"]
@ -11254,7 +11364,9 @@ async def model_list(
# The internal routing key drives the metadata/fallback lookup, while the
# public name is what the client sees as the model id.
model_data = []
admin_entries: Final = TeamModelNameTranslator.listing_entries(all_models, llm_router, settings)
admin_entries: Final = await _entries_kept_by_listing_callbacks(
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict
)
for response_id, lookup_id in admin_entries:
model_info = create_model_info_response(
model_id=lookup_id,
@ -11310,7 +11422,10 @@ async def model_list(
# public name is what the client sees as the model id.
model_data = []
entries: Final = alias_listing_entries(
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), caller_aliases
await _entries_kept_by_listing_callbacks(
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict
),
caller_aliases,
)
for response_id, lookup_id in entries:
model_info = create_model_info_response(
@ -11404,13 +11519,24 @@ async def model_info(
llm_router=llm_router,
)
hidden_names: Final = blocked_names | unhealthy_names
if hidden_names:
all_models = [m for m in all_models if m not in hidden_names]
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings)
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(
user_api_key_dict,
tuple(
response_id
for response_id, _ in TeamModelNameTranslator.listing_entries(
tuple(m for m in all_models if m not in hidden_names), llm_router, settings
)
),
)
if hidden_names or callback_hidden_names:
all_models = [
m for m in all_models if m not in hidden_names and internal_to_public.get(m, m) not in callback_hidden_names
]
undiscoverable_names: Final = undiscoverable_model_names(
all_models, llm_router, user_api_key_dict, team_id or user_api_key_dict.team_id
)
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings)
aliased_model_id: Final = alias_target(
model_id,
caller_alias_maps(
@ -13715,7 +13841,7 @@ async def transform_request(request: TransformRequestBody):
except ValueError as e:
raise HTTPException(status_code=400, detail={"error": str(e)})
return return_raw_request(endpoint=request.call_type, kwargs=request.request_body)
return await asyncio.to_thread(return_raw_request, request.call_type, request.request_body)
async def _check_if_model_is_user_added(
@ -15730,7 +15856,7 @@ async def model_info_v1(
if litellm_model_id is not None:
# user is trying to get specific model from litellm router
deployment_info: Final = llm_router.get_deployment(model_id=litellm_model_id)
if deployment_info is None:
if deployment_info is None or await _deployment_hidden_by_listing_callbacks(deployment_info, user_api_key_dict):
raise HTTPException(
status_code=400,
detail={"error": f"Model id = {litellm_model_id} not found on litellm proxy"},
@ -15819,10 +15945,17 @@ async def model_info_v1(
general_settings=general_settings,
llm_router=llm_router,
)
visible_models: Final = discoverable_rows(
servable_rows: Final = discoverable_rows(
(model for model in all_models if model.get("model_name") not in hidden_names),
user_api_key_dict,
)
listed_names: Final = tuple(
dict.fromkeys(name for model in servable_rows if isinstance(name := model.get("model_name"), str))
)
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(user_api_key_dict, listed_names)
visible_models: Final = tuple(
model for model in servable_rows if model.get("model_name") not in callback_hidden_names
)
verbose_proxy_logger.debug("all_models: %s", visible_models)
return _model_info_json_response(visible_models)
@ -15871,7 +16004,7 @@ async def model_deprecations(
def _get_model_group_info(
llm_router: Router, all_models_str: list[str], model_group: str | None
llm_router: Router, all_models_str: Sequence[str], model_group: str | None
) -> list[ModelGroupInfoProxy]:
model_groups: Final[list[ModelGroupInfoProxy]] = []
@ -16104,23 +16237,34 @@ async def model_group_info(
undiscoverable_group_names: Final = undiscoverable_model_names(
all_models_str, llm_router, user_api_key_dict, user_api_key_dict.team_id
)
model_groups: list[ModelGroupInfoProxy] = _get_model_group_info(
llm_router=llm_router,
all_models_str=[name for name in all_models_str if name not in undiscoverable_group_names],
model_group=model_group,
)
listed_group_names: Final = tuple(name for name in all_models_str if name not in undiscoverable_group_names)
# Append A2A agents to model groups
from litellm.proxy.agent_endpoints.model_list_helpers import (
append_agents_to_model_group,
)
model_groups = await append_agents_to_model_group(
model_groups=model_groups,
model_groups: Final = await append_agents_to_model_group(
model_groups=_get_model_group_info(
llm_router=llm_router, all_models_str=listed_group_names, model_group=model_group
),
user_api_key_dict=user_api_key_dict,
)
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings)
public_group_names: Final = tuple(
internal_to_public.get(group.model_group, group.model_group) for group in model_groups
)
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(
user_api_key_dict, tuple(dict.fromkeys(public_group_names))
)
return {"data": model_groups}
return {
"data": [
group
for group, public_name in zip(model_groups, public_group_names, strict=True)
if public_name not in callback_hidden_names
]
}
@router.get(

View file

@ -72,6 +72,7 @@ model LiteLLM_AgentsTable {
agent_card_params Json
static_headers Json? @default("{}")
extra_headers String[] @default([])
kill_switch Json?
agent_access_groups String[] @default([])
access_group_ids String[] @default([])
object_permission_id String?

View file

@ -0,0 +1,300 @@
"""Compare the spend LiteLLM captured for a provider against what that provider billed for the same UTC days.
LiteLLM's side is ``LiteLLM_DailyUserSpend``, summed over the ``custom_llm_provider`` values that land on the
provider's bill. The provider's side is its billing API, read with the customer's own billing credential
(OpenAI: the organization costs endpoint and an admin key in ``OPENAI_ADMIN_KEY``).
"""
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import date, datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, TypeAlias
from pydantic import BaseModel, ConfigDict, TypeAdapter
from typing_extensions import assert_never
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
SPEND_CAPTURE_RATE_CHECK_JOB_ID,
SPEND_CAPTURE_RATE_CHECK_LOCK_TTL_SECONDS,
SPEND_CAPTURE_RATE_DOCS_URL,
)
from litellm.llms.openai.organization_costs import (
OPENAI_ADMIN_KEY_ENV_VAR,
BillingHttpGet,
OpenAICostsRequestFailed,
fetch_openai_daily_costs,
provider_billing_get,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.proxy.spend_capture_rate import (
CaptureRateDay,
CaptureRateReport,
SpendCaptureProvider,
SpendCaptureRateCheckSettings,
)
if TYPE_CHECKING:
from litellm.caching.redis_cache import RedisCache
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.utils import PrismaClient
OPENAI_BILLED_LITELLM_PROVIDERS: Final = ("openai", "text-completion-openai")
CaptureRatePublisher: TypeAlias = Callable[[SpendCaptureProvider, float | None], None] # mutable-ok: Callable params
_CAPTURED_SPEND_BY_DAY_SQL: Final = """
SELECT date, COALESCE(SUM(spend), 0)::float AS spend
FROM "LiteLLM_DailyUserSpend"
WHERE date >= $1 AND date <= $2 AND custom_llm_provider = ANY($3::text[])
GROUP BY date
"""
@dataclass(frozen=True, slots=True)
class ProviderBillingCredentialMissing:
provider: SpendCaptureProvider
env_var: str
@dataclass(frozen=True, slots=True)
class ProviderBillingRequestFailed:
provider: SpendCaptureProvider
detail: str
ProviderBillingFailure: TypeAlias = ProviderBillingCredentialMissing | ProviderBillingRequestFailed
CheckResult: TypeAlias = CaptureRateReport | ProviderBillingFailure
class _CapturedSpendRow(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
date: str
spend: float
_CAPTURED_SPEND_ROWS: Final = TypeAdapter(tuple[_CapturedSpendRow, ...])
async def captured_spend_by_day(
prisma_client: "PrismaClient",
*,
litellm_providers: Sequence[str],
start_date: date,
end_date: date,
) -> Mapping[str, float]:
"""LiteLLM's tracked spend per UTC day (ISO date) for the given ``custom_llm_provider`` values."""
rows: Final = await prisma_client.db.query_raw(
_CAPTURED_SPEND_BY_DAY_SQL, start_date.isoformat(), end_date.isoformat(), tuple(litellm_providers)
)
return MappingProxyType({row.date: row.spend for row in _CAPTURED_SPEND_ROWS.validate_python(rows)})
def _ratio(captured: float, billed: float) -> float | None:
return None if billed <= 0 else captured / billed
def _days(start_date: date, end_date: date) -> tuple[date, ...]:
return tuple(start_date + timedelta(days=offset) for offset in range((end_date - start_date).days + 1))
def compute_capture_rate(
*,
provider: SpendCaptureProvider,
start_date: date,
end_date: date,
captured_by_day: Mapping[str, float],
billed_by_day: Mapping[str, float],
threshold: float,
) -> CaptureRateReport:
days: Final = tuple(
CaptureRateDay(
date=day.isoformat(),
captured_spend=captured_by_day.get(day.isoformat(), 0.0),
provider_spend=billed_by_day.get(day.isoformat(), 0.0),
capture_rate=_ratio(captured_by_day.get(day.isoformat(), 0.0), billed_by_day.get(day.isoformat(), 0.0)),
)
for day in _days(start_date, end_date)
)
captured: Final = sum(day.captured_spend for day in days)
billed: Final = sum(day.provider_spend for day in days)
rate: Final = _ratio(captured, billed)
return CaptureRateReport(
provider=provider,
start_date=start_date.isoformat(),
end_date=end_date.isoformat(),
captured_spend=captured,
provider_spend=billed,
capture_rate=rate,
threshold=threshold,
below_threshold=rate is not None and rate < threshold,
days=days,
)
async def capture_rate_report(
prisma_client: "PrismaClient",
*,
provider: SpendCaptureProvider,
start_date: date,
end_date: date,
threshold: float,
openai_project_ids: Sequence[str] = (),
http_get: BillingHttpGet = provider_billing_get,
) -> CheckResult:
match provider:
case "openai":
admin_key: Final = get_secret_str(OPENAI_ADMIN_KEY_ENV_VAR)
if admin_key is None:
return ProviderBillingCredentialMissing(provider, OPENAI_ADMIN_KEY_ENV_VAR)
billed: Final = await fetch_openai_daily_costs(
start_date, end_date, admin_key=admin_key, project_ids=openai_project_ids, http_get=http_get
)
if isinstance(billed, OpenAICostsRequestFailed):
return ProviderBillingRequestFailed(provider, billed.detail)
captured: Final = await captured_spend_by_day(
prisma_client,
litellm_providers=OPENAI_BILLED_LITELLM_PROVIDERS,
start_date=start_date,
end_date=end_date,
)
return compute_capture_rate(
provider=provider,
start_date=start_date,
end_date=end_date,
captured_by_day=captured,
billed_by_day=billed,
threshold=threshold,
)
case _:
assert_never(provider)
def alert_message(result: CheckResult) -> str | None:
"""The alert a check outcome warrants, or ``None`` when the capture rate is healthy."""
match result:
case ProviderBillingCredentialMissing(provider=provider, env_var=env_var):
return (
f"Spend capture-rate check: {env_var} is not set, so the {provider} bill cannot be read. "
f"Set it or remove general_settings.spend_capture_rate_check. {SPEND_CAPTURE_RATE_DOCS_URL}"
)
case ProviderBillingRequestFailed(provider=provider, detail=detail):
return f"Spend capture-rate check: could not read the {provider} bill ({detail}). {SPEND_CAPTURE_RATE_DOCS_URL}"
case CaptureRateReport():
if not result.below_threshold or result.capture_rate is None:
return None
return (
f"Spend capture rate for {result.provider} is {result.capture_rate:.1%}, under the "
f"{result.threshold:.0%} threshold: LiteLLM captured ${result.captured_spend:,.2f} of the "
f"${result.provider_spend:,.2f} {result.provider} bill for {result.start_date} to {result.end_date}. "
f"Requests reach {result.provider} outside LiteLLM or cost tracking is dropping spend. "
f"{SPEND_CAPTURE_RATE_DOCS_URL}"
)
case _:
assert_never(result)
def _published_rate(result: CheckResult) -> float | None:
"""The gauge value: the rate, or ``None`` (NaN on the gauge) when this window produced no rate."""
return result.capture_rate if isinstance(result, CaptureRateReport) else None
async def _check_every_provider(
prisma_client: "PrismaClient",
settings: SpendCaptureRateCheckSettings,
*,
publish: CaptureRatePublisher,
today: date | None,
http_get: BillingHttpGet,
) -> tuple[CheckResult, ...]:
"""Check every configured provider over the closed days before ``today`` and publish each outcome."""
end_date: Final = (today or datetime.now(timezone.utc).date()) - timedelta(days=1)
start_date: Final = end_date - timedelta(days=settings.lookback_days - 1)
results: Final = tuple(
[
await capture_rate_report(
prisma_client,
provider=provider,
start_date=start_date,
end_date=end_date,
threshold=settings.threshold,
openai_project_ids=settings.openai_project_ids,
http_get=http_get,
)
for provider in settings.providers
]
)
for result in results:
publish(result.provider, _published_rate(result))
verbose_proxy_logger.info("Spend capture-rate check: %s", result)
return results
def _alert_messages(results: Sequence[CheckResult]) -> tuple[str, ...]:
return tuple(message for message in map(alert_message, results) if message is not None)
async def run_spend_capture_rate_check(
prisma_client: "PrismaClient",
settings: SpendCaptureRateCheckSettings,
*,
alert: Callable[[str], Awaitable[None]],
publish: CaptureRatePublisher,
today: date | None = None,
http_get: BillingHttpGet = provider_billing_get,
) -> tuple[CheckResult, ...]:
"""Check every configured provider, publish each rate, and alert on every outcome that warrants one."""
results: Final = await _check_every_provider(
prisma_client, settings, publish=publish, today=today, http_get=http_get
)
for message in _alert_messages(results):
await alert(message)
return results
async def run_scheduled_spend_capture_rate_check(
prisma_client: "PrismaClient",
settings: SpendCaptureRateCheckSettings,
*,
pod_lock_manager: "PodLockManager | None",
alert: Callable[[str], Awaitable[None]],
publish: CaptureRatePublisher,
today: date | None = None,
http_get: BillingHttpGet = provider_billing_get,
) -> tuple[CheckResult, ...]:
"""Every worker publishes its own gauge; the first replica whose finished check has an alert claims the window."""
results: Final = await _check_every_provider(
prisma_client, settings, publish=publish, today=today, http_get=http_get
)
messages: Final = _alert_messages(results)
if not messages:
return results
if not await _claims_alert_window(pod_lock_manager):
verbose_proxy_logger.info("Spend capture-rate check: another pod alerted this window")
return results
for message in messages:
await alert(message)
return results
async def _claims_alert_window(pod_lock_manager: "PodLockManager | None") -> bool:
"""The lock is left to expire, so every replica firing within its TTL of the winner stays quiet."""
redis_cache: Final = None if pod_lock_manager is None else pod_lock_manager.redis_cache
if pod_lock_manager is None or redis_cache is None:
return True
acquired: Final = await pod_lock_manager.acquire_lock(
cronjob_id=SPEND_CAPTURE_RATE_CHECK_JOB_ID, ttl=SPEND_CAPTURE_RATE_CHECK_LOCK_TTL_SECONDS
)
return acquired or not await _lock_is_held(pod_lock_manager, redis_cache)
async def _lock_is_held(pod_lock_manager: "PodLockManager", redis_cache: "RedisCache") -> bool:
try:
return bool(
await redis_cache.async_get_cache(pod_lock_manager.get_redis_lock_key(SPEND_CAPTURE_RATE_CHECK_JOB_ID))
)
except Exception as exc: # noqa: BLE001 # an unreadable lock must not silence the alert
verbose_proxy_logger.warning("Spend capture-rate check: could not read the lock: %s", exc)
return False

View file

@ -24,7 +24,7 @@ from typing import (
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
from pydantic import TypeAdapter
from typing_extensions import ReadOnly
from typing_extensions import ReadOnly, assert_never
import litellm
from litellm._logging import verbose_proxy_logger
@ -32,11 +32,17 @@ from litellm.constants import (
EMPTY_MAPPING,
LITELLM_TRUNCATED_PAYLOAD_FIELD,
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
SPEND_CAPTURE_RATE_MAX_RANGE_DAYS,
)
from litellm.litellm_core_utils.classifier_logging import classifier_audit_fields, classifier_input_snapshot
from litellm.proxy._types import *
from litellm.proxy._types import ProviderBudgetResponse, ProviderBudgetResponseObject
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.spend_tracking.spend_capture_rate import (
ProviderBillingCredentialMissing,
ProviderBillingRequestFailed,
capture_rate_report,
)
# NOTE: Avoid module-level import from common_utils: proxy_server imports this
# module while common_utils may pull proxy_server during init, which can leave
@ -52,6 +58,7 @@ from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
from litellm.types.proxy.spend_capture_rate import CaptureRateReport, SpendCaptureProvider
if TYPE_CHECKING:
from prisma import models as prisma_models
@ -1183,6 +1190,84 @@ async def get_global_activity_exceptions(
)
@router.get(
"/spend/capture_rate",
tags=["Budget & Spend Tracking"], # mutable-ok: FastAPI tags kwarg is list-typed
dependencies=(Depends(user_api_key_auth),),
response_model=CaptureRateReport,
)
async def get_spend_capture_rate(
start_date: Annotated[date, fastapi.Query(description="First UTC day of the range, YYYY-MM-DD")],
end_date: Annotated[date, fastapi.Query(description="Last UTC day of the range, YYYY-MM-DD, inclusive")],
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
provider: Annotated[
SpendCaptureProvider,
fastapi.Query(description="Provider whose bill to compare against; needs OPENAI_ADMIN_KEY set on the proxy"),
] = "openai",
threshold: Annotated[
float, fastapi.Query(gt=0, le=1, description="Ratio under which the report flags below_threshold")
] = 0.9,
project_ids: Annotated[
list[str] | None,
fastapi.Query(
description=(
"Scope the OpenAI bill to these project ids; omit to compare against the whole organization. Captured "
"spend is never scoped, so pass every project LiteLLM's OpenAI keys belong to"
)
),
] = None,
) -> CaptureRateReport:
"""
Compare the spend LiteLLM captured for a provider against that provider's own bill, per UTC day.
Admin only. Reads the provider's billing API with the billing credential set on the proxy
(OpenAI: `OPENAI_ADMIN_KEY`) and sums `LiteLLM_DailyUserSpend` for the same days.
Example:
```
curl -H "Authorization: Bearer sk-1234" \
"http://localhost:4000/spend/capture_rate?provider=openai&start_date=2026-09-17&end_date=2026-09-23"
```
"""
from litellm.proxy.proxy_server import prisma_client
if not _is_admin_view_safe(user_api_key_dict):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Only proxy admins can read the capture rate")
if prisma_client is None:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=CommonProxyErrors.db_not_connected_error.value
)
if end_date < start_date:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="end_date must not be before start_date")
if (end_date - start_date).days >= SPEND_CAPTURE_RATE_MAX_RANGE_DAYS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Date range too large; maximum is {SPEND_CAPTURE_RATE_MAX_RANGE_DAYS} days",
)
result: Final = await capture_rate_report(
prisma_client,
provider=provider,
start_date=start_date,
end_date=end_date,
threshold=threshold,
openai_project_ids=tuple(project_ids or ()),
)
match result:
case ProviderBillingCredentialMissing(env_var=env_var):
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=f"{env_var} is not set on the proxy, so the {provider} bill cannot be read",
)
case ProviderBillingRequestFailed(detail=detail):
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY, detail=f"Could not read the {provider} bill: {detail}"
)
case CaptureRateReport():
return result
case _:
assert_never(result)
@router.get(
"/global/spend/provider",
tags=["Budget & Spend Tracking"],

View file

@ -36,6 +36,7 @@ from typing import (
Final,
Generic,
Literal,
NoReturn,
Optional,
Protocol,
TypeAlias,
@ -106,7 +107,7 @@ except ImportError:
raise ImportError("backoff is not installed. Please install it via 'pip install backoff'")
from fastapi import HTTPException, status
from pydantic import TypeAdapter
from pydantic import TypeAdapter, ValidationError
import litellm
import litellm.litellm_core_utils
@ -1136,11 +1137,51 @@ class _CallbackCapabilities:
# avoids the per-request ``get_custom_logger_compatible_class`` walk for
# every string entry in ``litellm.callbacks``.
resolved_callbacks: tuple[object, ...] = field(default_factory=tuple)
listed_models_filters: tuple[CustomLogger, ...] = field(default_factory=tuple)
def _overrides_hook(callback: CustomLogger, hook_name: str) -> bool:
leaf_to_base: Final = takewhile(lambda klass: klass is not CustomLogger, type(callback).__mro__)
return any(hook_name in klass.__dict__ for klass in leaf_to_base)
def _overrides_moderation_hook(callback: CustomLogger) -> bool:
leaf_to_base: Final = takewhile(lambda klass: klass is not CustomLogger, type(callback).__mro__)
return any("async_moderation_hook" in klass.__dict__ for klass in leaf_to_base)
return _overrides_hook(callback, "async_moderation_hook")
_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...])
@dataclass(frozen=True, slots=True)
class MalformedListingFilterReturn:
callback: str
tag: Literal["malformed_listing_filter_return"] = "malformed_listing_filter_return"
async def _names_kept_by_listing_callbacks(
callbacks: Sequence[CustomLogger],
user_api_key_dict: UserAPIKeyAuth,
model_names: tuple[str, ...],
) -> tuple[str, ...] | MalformedListingFilterReturn:
if not callbacks or not model_names:
return model_names
returned: Final = await callbacks[0].async_filter_listed_models(user_api_key_dict, model_names)
try:
kept: Final = frozenset(_LISTED_MODEL_NAMES.validate_python(returned))
except ValidationError:
return MalformedListingFilterReturn(callback=type(callbacks[0]).__name__)
return await _names_kept_by_listing_callbacks(
callbacks[1:], user_api_key_dict, tuple(name for name in model_names if name in kept)
)
def _raise_malformed_listing_filter_return(error: MalformedListingFilterReturn) -> NoReturn:
raise ProxyException(
message=f"{error.callback}.async_filter_listed_models must return a sequence of model names",
type=ProxyErrorTypes.internal_server_error,
param=None,
code=500,
)
class ProxyLogging:
@ -2808,6 +2849,9 @@ class ProxyLogging:
has_moderation_override=has_moderation_override,
iterator_overrides=tuple(iterator_overrides),
resolved_callbacks=tuple(resolved_callbacks),
listed_models_filters=tuple(
callback for callback in resolved_callbacks if _overrides_hook(callback, "async_filter_listed_models")
),
)
# Limit cache to handle test churn without leaking; production
# callback lists are stable so this rarely grows past 1 entry.
@ -3715,6 +3759,18 @@ class ProxyLogging:
verbose_proxy_logger.exception("Error in post_call_response_headers_hook: %s", str(e))
return merged_headers
async def hidden_by_listing_callbacks(
self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
) -> frozenset[str]:
filters: Final = ProxyLogging._callback_capabilities().listed_models_filters
if not filters:
return frozenset()
candidates: Final = tuple(model_names)
kept: Final = await _names_kept_by_listing_callbacks(filters, user_api_key_dict, candidates)
if isinstance(kept, MalformedListingFilterReturn):
_raise_malformed_listing_filter_return(kept)
return frozenset(candidates).difference(kept)
@staticmethod
def _build_litellm_call_info(data: dict, response: object) -> dict[str, object]:
"""

View file

@ -227,6 +227,9 @@ from litellm.router_utils.pre_call_checks.deployment_affinity_check import (
DeploymentAffinityCheck,
warn_on_unknown_model_group_affinity_flags,
)
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
EncryptedContentAffinityCheck,
)
from litellm.router_utils.pre_call_checks.io_token_rate_limit_check import (
build_io_token_rate_limit_headers,
deployment_has_io_token_limits,
@ -439,6 +442,7 @@ _RUNTIME_TOGGLEABLE_PRE_CALL_CHECKS: Final[Mapping[str, type[CustomLogger]]] = M
{
"prompt_caching": PromptCachingDeploymentCheck,
"enforce_model_rate_limits": ModelRateLimitingCheck,
"encrypted_content_affinity": EncryptedContentAffinityCheck,
}
)
@ -2208,10 +2212,6 @@ class Router:
)
def _add_encrypted_content_affinity_check(self, enable_global_affinity: bool) -> None:
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
EncryptedContentAffinityCheck,
)
def _move_before_deployment_affinity(
callback_list: list[Any],
callback_to_move: EncryptedContentAffinityCheck,
@ -7888,48 +7888,15 @@ class Router:
"""
return run_async_function(self.async_function_with_fallbacks, *args, **kwargs)
def _get_fallback_model_group_from_fallbacks(
self,
fallbacks: list[dict[str, list[str]]],
model_group: str | None = None,
) -> list[str] | None:
"""
Returns the list of fallback models to use for a given model group
If no fallback model group is found, returns None
Example:
fallbacks = [{"gpt-3.5-turbo": ["gpt-4"]}, {"gpt-4o": ["gpt-3.5-turbo"]}]
model_group = "gpt-3.5-turbo"
returns: ["gpt-4"]
"""
if model_group is None:
return None
fallback_model_group: list[str] | None = None
for item in fallbacks: # [{"gpt-3.5-turbo": ["gpt-4"]}]
if list(item.keys())[0] == model_group:
fallback_model_group = item[model_group]
break
return fallback_model_group
def _get_fallback_model_group_for_lookup_groups(
self,
fallbacks: list[dict[str, list[str]]], # mutable-ok: mirrors the sibling resolver's contract
fallbacks: list[dict[str, list[str]]], # mutable-ok: mirrors the shared resolver's contract
lookup_groups: tuple[str, ...],
) -> list[str] | None: # mutable-ok: mirrors the sibling resolver's contract
"""First lookup group whose exact-key chain resolves (tier first, then requested group)."""
return next(
(
resolved
for resolved in (
self._get_fallback_model_group_from_fallbacks(fallbacks=fallbacks, model_group=group)
for group in lookup_groups
)
if resolved is not None
),
None,
) -> list[str] | None: # mutable-ok: mirrors the shared resolver's contract
fallback_model_group, _ = get_fallback_model_group_for_lookup_groups(
fallbacks=fallbacks, lookup_groups=lookup_groups
)
return fallback_model_group
def _get_first_default_fallback(self) -> str | None:
"""

View file

@ -11,6 +11,7 @@ import litellm
from litellm._logging import verbose_router_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs, safe_deep_copy
from litellm.litellm_core_utils.get_llm_provider_logic import inferred_provider
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_structure
from litellm.router_utils.add_retry_fallback_headers import (
add_fallback_headers_to_response,
@ -236,6 +237,13 @@ def _check_stripped_model_group(model_group: str, fallback_key: str) -> bool:
return False
def _provider_prefixed_model_group(model_group: str, fallback_keys: Sequence[str]) -> str | None:
if "/" in model_group or not any(key.endswith(f"/{model_group}") for key in fallback_keys):
return None
provider: Final = inferred_provider(model_group)
return f"{provider}/{model_group}" if provider else None
PRE_ROUTING_SELECTED_MODEL_KEY: Final = "pre_routing_selected_model"
_ROUTER_METADATA_BUCKETS: Final = ("metadata", "litellm_metadata")
@ -439,22 +447,26 @@ def get_fallback_model_group(fallbacks: list[Any], model_group: str) -> tuple[li
Checks:
- exact match
- stripped model group match
- provider-prefixed model group match
- generic fallback
"""
generic_fallback_idx: int | None = None
stripped_model_fallback: list[str] | None = None
fallback_model_group: list[str] | None = None
fallback_keys: Final = tuple(next(iter(item)) for item in fallbacks if isinstance(item, dict) and item)
prefixed_model_group: Final = _provider_prefixed_model_group(model_group, fallback_keys)
## check for specific model group-specific fallbacks
for idx, item in enumerate(fallbacks):
if isinstance(item, dict):
if list(item.keys())[0] == model_group: # check exact match
fallback_key = next(iter(item))
if fallback_key == model_group: # check exact match
fallback_model_group = item[model_group]
break
elif _check_stripped_model_group(
model_group=model_group, fallback_key=list(item.keys())[0]
elif fallback_key == prefixed_model_group or _check_stripped_model_group(
model_group=model_group, fallback_key=fallback_key
): # check generic fallback
stripped_model_fallback = item[list(item.keys())[0]]
elif list(item.keys())[0] == "*": # check generic fallback
stripped_model_fallback = item[fallback_key]
elif fallback_key == "*": # check generic fallback
generic_fallback_idx = idx
elif isinstance(item, str):
fallback_model_group = [item]

View file

@ -8,7 +8,7 @@ from re import Match
from typing import Final
from litellm._logging import verbose_router_logger
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider, get_llm_provider
from litellm.litellm_core_utils.get_llm_provider_logic import inferred_provider
class PatternUtils:
@ -218,18 +218,9 @@ class PatternMatchRouter:
Returns:
bool: True if pattern exists, False otherwise
"""
provider: Final = (
custom_llm_provider or declared_authenticating_provider(model) or self._resolved_provider(model)
)
provider: Final = custom_llm_provider or inferred_provider(model)
return self.route(model) or self.route(f"{provider}/{model}")
@staticmethod
def _resolved_provider(model: str | None) -> str | None:
try:
return get_llm_provider(model=model)[1] if model else None
except Exception: # noqa: BLE001 # get_llm_provider raises when the provider is unknown; the name then routes as-is
return None
def get_deployments_by_pattern(self, model: str, custom_llm_provider: str | None = None) -> list[dict]:
"""
Get the deployments by pattern

View file

@ -36,17 +36,10 @@ Safe to enable globally:
- No cache required.
"""
import time
from collections.abc import Iterator, Mapping
from typing import TYPE_CHECKING, Final, Optional, Protocol, cast
import httpx
from typing import TYPE_CHECKING, Final, Optional, cast
from litellm._logging import verbose_router_logger
from litellm.exceptions import (
RateLimitError,
ServiceUnavailableError,
)
from litellm.integrations.custom_logger import CustomLogger, Span
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -55,7 +48,6 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
strip_encrypted_reasoning_from_messages,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.router_utils.cooldown_cache import CooldownCacheValue
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import Deployment
@ -63,14 +55,6 @@ if TYPE_CHECKING:
from litellm.router import Router
class _SupportsActiveCooldowns(Protocol):
"""Cooldown-cache handle: this check only reads back the currently active cooldowns."""
async def async_get_active_cooldowns(
self, model_ids: list[str], parent_otel_span: Span | None
) -> list[tuple[str, CooldownCacheValue]]: ...
class EncryptedContentAffinityCheck(CustomLogger):
"""
Routes follow-up Responses API requests to the deployment that produced
@ -194,23 +178,6 @@ class EncryptedContentAffinityCheck(CustomLogger):
return deployment
return None
@staticmethod
def _request_team_id(request_kwargs: Mapping[str, object]) -> str | None:
containers: Final = (request_kwargs.get("metadata"), request_kwargs.get("litellm_metadata"))
team_ids: Final = (c.get("user_api_key_team_id") for c in containers if isinstance(c, Mapping))
return next((tid for tid in team_ids if isinstance(tid, str)), None)
def _routed_group_candidate_model_ids(self, request_kwargs: Mapping[str, object], model: str) -> frozenset[str]:
"""
Deployment ids that could serve this turn's routed ``model``, as the router
resolves a route (model_group_alias / routing group / model_name / team /
pattern). Delegates to the router so the full precedence is not re-derived here
and no deployment ids are written into request kwargs bound for the provider.
"""
if self.router is None:
return frozenset()
return self.router.get_candidate_model_ids_for_route(model=model, team_id=self._request_team_id(request_kwargs))
@staticmethod
def _encryption_boundary_key(
litellm_params: object,
@ -262,9 +229,6 @@ class EncryptedContentAffinityCheck(CustomLogger):
Deployments in ``healthy_deployments`` sharing the originating
deployment's ``(api_base, api_key)``, alongside the originating
deployment object (or ``None`` if it was removed / router unavailable).
Returns ``([], originating_or_None)`` when no boundary match exists,
so the caller can reuse the looked-up ``originating`` rather than
re-querying the router.
"""
if self.router is None:
return [], None
@ -294,18 +258,12 @@ class EncryptedContentAffinityCheck(CustomLogger):
"""
If the request ``input`` contains litellm-encoded item IDs, or its Anthropic
``messages`` replay a bridge-tagged thinking block, decode the embedded
``model_id`` and pin the request to that deployment. Raises
``RateLimitError`` / ``ServiceUnavailableError`` when the originating
deployment is a member of the routed model group but currently unavailable
and no encryption-boundary peer exists, rather than dispatching a doomed
request to a non-peer deployment. When the origin is not a member of the
routed group (an auto-router tier change, a model switch with no peer, a
removed deployment, or an unknown/forged marker), the encrypted reasoning is
stripped and the request dispatches with its readable history instead. The
429/503 split mirrors the originating cooldown's status:
a 429-induced cooldown surfaces as 429 (with ``Retry-After`` set to the
remaining cooldown window) so OpenAI-compatible clients back off and
retry after the deployment is eligible again.
``model_id`` and pin the request to that deployment. When the origin cannot
serve this turn and no encryption-boundary peer is configured (it is
unhealthy, the request was routed to a different group by an auto-router tier
change or model switch, or the marker is removed/unknown/forged), the
encrypted reasoning is stripped and the request dispatches to the healthy
pool with its readable history instead of failing.
"""
request_kwargs = request_kwargs or {}
typed_healthy_deployments: Final = cast(list[dict], healthy_deployments)
@ -348,7 +306,7 @@ class EncryptedContentAffinityCheck(CustomLogger):
return [deployment]
# Follow-up switched model_name (LIT-2531): pin by Azure resource instead.
boundary_matches, originating = self._find_deployments_on_same_encryption_boundary(
boundary_matches, _originating = self._find_deployments_on_same_encryption_boundary(
healthy_deployments=typed_healthy_deployments,
model_id=model_id,
)
@ -362,101 +320,17 @@ class EncryptedContentAffinityCheck(CustomLogger):
request_kwargs["_encrypted_content_affinity_pinned"] = True
return boundary_matches
# The origin cannot serve this turn's routed group and no peer shares the boundary, so its
# The origin cannot serve this turn and no peer shares its encryption boundary, so its
# encrypted reasoning can never decrypt here. Strip it, keep the readable history, and dispatch
# to the routed group instead of failing. Membership is tested by deployment id against the set
# the router actually resolved for this route, not by model-group name, so an alias, a
# provider-qualified spelling, a team-public name, or a pattern route of the same group is not
# mistaken for a tier change. An unknown origin (a removed deployment, or a forged marker) is
# treated the same as a cross-group one, which also denies an authenticated caller a
# deployment-id existence oracle: a real cross-group id and a nonexistent id both strip and
# dispatch rather than returning distinguishable responses. Only a genuine same-group member
# that is currently unavailable falls through to the fail-fast, preserving the cooldown contract.
routed_group_model_ids: Final = (
self._routed_group_candidate_model_ids(request_kwargs, model) if originating is not None else frozenset()
# to the healthy pool instead of failing the request. This also denies an authenticated caller a
# deployment-id existence oracle: a same-group id, a cross-group id, a removed id and a forged
# marker all strip and dispatch rather than returning distinguishable responses.
verbose_router_logger.warning(
"EncryptedContentAffinityCheck: model_id=%s cannot serve group %s and no deployment on the same "
"encryption boundary is configured; forwarding without its encrypted reasoning",
model_id[:64],
model,
)
if str(model_id) not in routed_group_model_ids:
verbose_router_logger.debug(
"EncryptedContentAffinityCheck: model_id=%s is not a candidate for the routed group %s; "
"forwarding without its encrypted reasoning",
model_id,
model,
)
ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
strip_encrypted_reasoning_from_messages(anthropic_messages)
return typed_healthy_deployments
# The origin is a member of the routed group but currently unavailable (cooled down); fail fast
# rather than dispatching to a non-peer, which would guarantee an upstream 400.
raise await self._unavailable_origin_error(
model=model,
model_id=model_id,
parent_otel_span=parent_otel_span,
)
async def _unavailable_origin_error(
self,
model: str,
model_id: str,
parent_otel_span: Span | None,
) -> Exception:
# Public error messages intentionally omit the originating ``model_id`` so
# an authenticated caller forging encrypted-content markers cannot use the
# error surface to enumerate which deployment IDs exist on this router.
cooldown: Final = await self._get_origin_cooldown(model_id=model_id, parent_otel_span=parent_otel_span)
if cooldown is not None and str(cooldown.get("status_code")) == "429":
retry_after: Final = self._cooldown_seconds_remaining(cooldown)
return RateLimitError(
message=(
"The deployment that produced this encrypted_content is "
f"rate-limited (cooling down for ~{retry_after}s), and no "
"deployment on the same encryption boundary is configured. "
"Retry after the Retry-After window or configure a deployment "
"with the same (api_base, api_key)."
),
llm_provider="",
model=model,
response=httpx.Response(
status_code=429,
headers={"retry-after": str(retry_after)},
request=httpx.Request("POST", "https://litellm.ai/"),
),
)
return ServiceUnavailableError(
message=(
"The deployment that produced this encrypted_content is "
"currently unavailable (likely cooled down), and no deployment "
"on the same encryption boundary is configured. Retry later or "
"configure a deployment with the same (api_base, api_key)."
),
llm_provider="",
model=model,
)
async def _get_origin_cooldown(
self,
model_id: str,
parent_otel_span: Span | None,
) -> CooldownCacheValue | None:
if self.router is None:
return None
cooldown_cache: Final[_SupportsActiveCooldowns | None] = getattr(self.router, "cooldown_cache", None)
if cooldown_cache is None:
return None
try:
active: Final = await cooldown_cache.async_get_active_cooldowns(
model_ids=[model_id], parent_otel_span=parent_otel_span
)
except Exception:
return None
for cached_model_id, value in active:
if cached_model_id == model_id:
return value
return None
@staticmethod
def _cooldown_seconds_remaining(cooldown: CooldownCacheValue) -> int:
remaining = float(cooldown.get("timestamp", 0.0)) + float(cooldown.get("cooldown_time", 0.0)) - time.time()
return max(1, int(remaining))
ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
strip_encrypted_reasoning_from_messages(anthropic_messages)
return typed_healthy_deployments

View file

@ -46,6 +46,8 @@ def aocr(
args: tuple[object, ...],
kwargs: dict[str, object],
) -> Coroutine[object, object, OCRResponse]: ...
def ocr_health_check_document(model: str, custom_llm_provider: str | None) -> dict[str, object]: ...
def ocr_passthrough_response(model: str, endpoint: str, body: bytes) -> dict[str, object] | None: ...
def embedding(
request: LiteLLMEmbeddingRequest,
args: tuple[object, ...],
@ -417,6 +419,8 @@ __all__ = [
"gil_stats",
"messages",
"ocr",
"ocr_health_check_document",
"ocr_passthrough_response",
"process_state_started",
"reserve_process_for_forking",
"responses",

View file

@ -108,8 +108,7 @@ RULES: Final[Rules] = (
LoggerRule(Rollout.RUST_OPT_IN),
RouteRule(Route.CHAT_COMPLETIONS, Rollout.PYTHON_ONLY),
RouteRule(Route.EMBEDDINGS, Rollout.PYTHON_ONLY),
RouteRule(Route.OCR, Rollout.RUST_REQUIRED, providers=frozenset({"aws_textract"})),
RouteRule(Route.OCR, Rollout.RUST_OPT_OUT),
RouteRule(Route.OCR, Rollout.RUST_REQUIRED),
RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY),
RouteRule(Route.RESPONSES, Rollout.PYTHON_ONLY),
RouteRule(Route.TOKEN_COUNTER, Rollout.PYTHON_ONLY),

View file

@ -44,17 +44,39 @@ class PublicDispatch(Generic[RequestT]):
return rollout_decision(rule.rollout) is not Decision.PYTHON
return False
def _native_request(self, args: tuple[object, ...], kwargs: Mapping[str, object]) -> RequestT:
request: Final = self.request(args, kwargs)
if request is None:
raise runtime.NoPythonImplementationError(
f"{self.route.value} has no Python implementation, so every call must project to a native request"
)
if self.bypass is not None and self.bypass(request):
raise runtime.NoPythonImplementationError(
f"{self.route.value} has no Python implementation, so a call its bypass predicate matches cannot "
"be served"
)
return request
def run(
self,
args: tuple[object, ...],
kwargs: Mapping[str, object],
*,
python: Callable[..., ResultT],
python: Callable[..., ResultT] | runtime.NoPythonImplementation,
binding: NativeBinding[NativeT],
native: Callable[[NativeT, RequestT, tuple[object, ...], Mapping[str, object]], ResultT],
rules: Rules | None = None,
) -> ResultT:
selected_rules: Final = catalog.RULES if rules is None else rules
if isinstance(python, runtime.NoPythonImplementation):
native_request: Final = self._native_request(args, kwargs)
return runtime.run(
self.context(native_request),
binding=binding,
native=lambda hook: native(hook, native_request, args, kwargs),
python=python,
rules=selected_rules,
)
if not self._requires_projection(selected_rules):
return python(*args, **kwargs)
request: Final = self.request(args, kwargs)
@ -73,12 +95,21 @@ class PublicDispatch(Generic[RequestT]):
args: tuple[object, ...],
kwargs: Mapping[str, object],
*,
python: Callable[..., Awaitable[ResultT]],
python: Callable[..., Awaitable[ResultT]] | runtime.NoPythonImplementation,
binding: NativeBinding[NativeT],
native: Callable[[NativeT, RequestT, tuple[object, ...], Mapping[str, object]], Awaitable[ResultT]],
rules: Rules | None = None,
) -> ResultT:
selected_rules: Final = catalog.RULES if rules is None else rules
if isinstance(python, runtime.NoPythonImplementation):
native_request: Final = self._native_request(args, kwargs)
return await runtime.arun(
self.context(native_request),
binding=binding,
native=lambda hook: native(hook, native_request, args, kwargs),
python=python,
rules=selected_rules,
)
if not self._requires_projection(selected_rules):
return await python(*args, **kwargs)
request: Final = self.request(args, kwargs)

View file

@ -6,7 +6,7 @@ from typing import Final, Protocol, cast # noqa: TID251 # validates dynamicall
import httpx
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRResponse
from litellm.rust_bridge.bindings import NativeBinding
@ -53,5 +53,31 @@ def _aocr_binding(value: object) -> NativeAocr | None:
return cast("NativeAocr", value) # cast-ok: callable validated at the native binding boundary
class NativeOcrHealthCheckDocument(Protocol):
def __call__(self, model: str, custom_llm_provider: str | None) -> DocumentType: ...
class NativeOcrPassthroughResponse(Protocol):
def __call__(self, model: str, endpoint: str, body: bytes) -> Mapping[str, object] | None: ...
def _health_check_document_binding(value: object) -> NativeOcrHealthCheckDocument | None:
if not callable(value):
return None
return cast("NativeOcrHealthCheckDocument", value) # cast-ok: callable validated at the native binding boundary
def _passthrough_response_binding(value: object) -> NativeOcrPassthroughResponse | None:
if not callable(value):
return None
return cast("NativeOcrPassthroughResponse", value) # cast-ok: callable validated at the native binding boundary
NATIVE_OCR: Final = NativeBinding("ocr", validate=_ocr_binding)
NATIVE_AOCR: Final = NativeBinding("aocr", validate=_aocr_binding)
NATIVE_OCR_HEALTH_CHECK_DOCUMENT: Final = NativeBinding(
"ocr_health_check_document", validate=_health_check_document_binding
)
NATIVE_OCR_PASSTHROUGH_RESPONSE: Final = NativeBinding(
"ocr_passthrough_response", validate=_passthrough_response_binding
)

View file

@ -41,29 +41,37 @@ class BridgeErrorContext:
model: str
@dataclass(frozen=True, slots=True)
class NoPythonImplementation:
pass
NO_PYTHON: Final = NoPythonImplementation()
class NoPythonImplementationError(RuntimeError):
pass
def run(
context: RouteContext,
*,
binding: NativeBinding[NativeT],
native: Callable[[NativeT], ResultT],
python: Callable[[], ResultT],
python: Callable[[], ResultT] | NoPythonImplementation,
rules: Rules | None = None,
) -> ResultT:
selected: Final = decision(context, rules)
if isinstance(python, NoPythonImplementation):
_require_rust(context, selected)
return _required(_attempt_native(context, binding, native), context)
match selected:
case Decision.PYTHON:
return python()
case Decision.RUST_WITH_FALLBACK | Decision.RUST_REQUIRED:
loaded: Final = binding.load()
result: Final = attempt(
native_call=None if loaded is None else lambda: native(loaded),
adapt=_identity,
context=_error_context(context),
)
if isinstance(result, RustHandled):
return mark_rust_response(result.value)
if selected is Decision.RUST_REQUIRED:
_raise_required(result, _error_context(context))
result: Final = _attempt_native(context, binding, native)
if isinstance(result, RustHandled) or selected is Decision.RUST_REQUIRED:
return _required(result, context)
return python()
case _:
assert_never(selected)
@ -74,29 +82,61 @@ async def arun(
*,
binding: NativeBinding[NativeT],
native: Callable[[NativeT], Awaitable[ResultT]],
python: Callable[[], Awaitable[ResultT]],
python: Callable[[], Awaitable[ResultT]] | NoPythonImplementation,
rules: Rules | None = None,
) -> ResultT:
selected: Final = decision(context, rules)
if isinstance(python, NoPythonImplementation):
_require_rust(context, selected)
return _required(await _aattempt_native(context, binding, native), context)
match selected:
case Decision.PYTHON:
return await python()
case Decision.RUST_WITH_FALLBACK | Decision.RUST_REQUIRED:
loaded: Final = binding.load()
result: Final = await aattempt(
native_call=None if loaded is None else lambda: native(loaded),
adapt=_identity,
context=_error_context(context),
)
if isinstance(result, RustHandled):
return mark_rust_response(result.value)
if selected is Decision.RUST_REQUIRED:
_raise_required(result, _error_context(context))
result: Final = await _aattempt_native(context, binding, native)
if isinstance(result, RustHandled) or selected is Decision.RUST_REQUIRED:
return _required(result, context)
return await python()
case _:
assert_never(selected)
def _require_rust(context: RouteContext, selected: Decision) -> None:
if selected is not Decision.RUST_REQUIRED:
raise NoPythonImplementationError(
f"{context.route.value} has no Python implementation, so its catalog rules must resolve to "
f"RUST_REQUIRED, but provider={context.provider!r} model={context.model!r} resolved to {selected.name}"
)
def _attempt_native(
context: RouteContext, binding: NativeBinding[NativeT], native: Callable[[NativeT], ResultT]
) -> RustAttempt[ResultT]:
loaded: Final = binding.load()
return attempt(
native_call=None if loaded is None else lambda: native(loaded),
adapt=_identity,
context=_error_context(context),
)
async def _aattempt_native(
context: RouteContext, binding: NativeBinding[NativeT], native: Callable[[NativeT], Awaitable[ResultT]]
) -> RustAttempt[ResultT]:
loaded: Final = binding.load()
return await aattempt(
native_call=None if loaded is None else lambda: native(loaded),
adapt=_identity,
context=_error_context(context),
)
def _required(result: RustAttempt[ResultT], context: RouteContext) -> ResultT:
if isinstance(result, RustHandled):
return mark_rust_response(result.value)
_raise_required(result, _error_context(context))
def _identity(value: ResultT) -> ResultT:
return value

View file

@ -1,8 +1,9 @@
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, TypeAlias
from urllib.parse import urlsplit
from pydantic import BaseModel, ConfigDict, PrivateAttr, StrictInt
from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, StrictInt, field_validator
from typing_extensions import ReadOnly, Required, TypedDict
from litellm.types.llms.base import LiteLLMPydanticObjectBase
@ -178,6 +179,74 @@ class AgentObjectPermission(TypedDict, total=False):
agents: list[str] | None
class AgentKillSwitchBearerAuth(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
type: Literal["bearer"]
token: str
class AgentKillSwitchApiKeyAuth(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
type: Literal["api_key"]
header_name: str = "x-api-key"
api_key: str
class AgentKillSwitchBasicAuth(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
type: Literal["basic"]
username: str
password: str
AgentKillSwitchAuth: TypeAlias = Annotated[
AgentKillSwitchBearerAuth | AgentKillSwitchApiKeyAuth | AgentKillSwitchBasicAuth,
Field(discriminator="type"),
]
AgentKillSwitchMethod: TypeAlias = Literal["POST", "PUT", "PATCH", "DELETE", "GET"]
class AgentKillSwitchConfig(BaseModel):
"""Webhook an admin fires to shut an agent down out of band. LiteLLM only
makes the call; whatever the endpoint does with it is the agent's business."""
model_config = ConfigDict(frozen=True, extra="forbid")
url: str
method: AgentKillSwitchMethod = "POST"
headers: Mapping[str, str] = Field(default_factory=dict)
query_params: Mapping[str, str] = Field(default_factory=dict)
body: Mapping[str, object] | None = None
auth: AgentKillSwitchAuth | None = None
@field_validator("url")
@classmethod
def _require_absolute_http_url(cls, value: str) -> str:
parts: Final = urlsplit(value)
if parts.scheme not in ("http", "https") or not parts.netloc:
raise ValueError("kill_switch.url must be an absolute http(s) URL")
return value
class AgentKillSwitchResult(BaseModel):
model_config = ConfigDict(frozen=True)
agent_id: str
url: str
method: AgentKillSwitchMethod
status_code: int | None = None
response_body: str | None = None
error: str | None = None
@property
def succeeded(self) -> bool:
return self.status_code is not None and 200 <= self.status_code < 300
class AgentConfig(TypedDict, total=False):
agent_name: Required[str]
agent_card_params: Required[AgentCard]
@ -190,6 +259,7 @@ class AgentConfig(TypedDict, total=False):
static_headers: dict[str, str] | None
extra_headers: list[str] | None
access_group_ids: ReadOnly[Sequence[str] | None]
kill_switch: ReadOnly[AgentKillSwitchConfig | None]
class PatchAgentRequest(TypedDict, total=False):
@ -204,6 +274,7 @@ class PatchAgentRequest(TypedDict, total=False):
static_headers: dict[str, str] | None
extra_headers: list[str] | None
access_group_ids: ReadOnly[Sequence[str] | None]
kill_switch: ReadOnly[AgentKillSwitchConfig | None]
AGENT_CALLER_USER_ID_HEADER: Final = "x-litellm-user-id"
@ -243,6 +314,7 @@ class AgentResponse(BaseModel):
static_headers: dict[str, str] | None = None
extra_headers: list[str] | None = None
access_group_ids: Sequence[str] | None = None
kill_switch: AgentKillSwitchConfig | None = None
keys: list[AgentKeySummary] | None = None
search_score: float | None = None
created_at: datetime | None = None

View file

@ -281,6 +281,7 @@ DEFINED_PROMETHEUS_METRICS = Literal[
"litellm_guardrail_errors_total",
"litellm_guardrail_requests_total",
"litellm_zero_cost_requests_total",
"litellm_spend_capture_rate",
# Cache metrics
"litellm_cache_hits_metric",
"litellm_cache_misses_metric",
@ -600,6 +601,8 @@ class PrometheusMetricLabels:
ZERO_COST_REASON_LABEL,
)
litellm_spend_capture_rate = (UserAPIKeyLabelNames.API_PROVIDER.value,)
litellm_input_tokens_metric = [
UserAPIKeyLabelNames.END_USER.value,
UserAPIKeyLabelNames.API_KEY_HASH.value,

View file

@ -24,8 +24,10 @@ class httpxSpecialProvider(str, Enum):
Search = "search"
MCP = "mcp"
RAG = "rag"
ProviderBilling = "provider_billing"
A2AProvider = "a2a_provider"
AgentHealthCheck = "agent_health_check"
AgentKillSwitch = "agent_kill_switch"
A2A = "a2a"
PromptManagement = "prompt_management"
UI = "ui"

View file

@ -0,0 +1,54 @@
"""The captured-spend to provider-bill ratio: the share of a provider's bill that went through LiteLLM and was priced.
``capture_rate = captured_spend / provider_spend`` over the same UTC days. 1.0 means LiteLLM saw and priced every
dollar the provider billed, lower means traffic reaches the provider outside LiteLLM or cost tracking drops spend,
higher means LiteLLM prices above the bill. ``None`` means the provider billed nothing, so there is no ratio.
"""
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field
from litellm.constants import SPEND_CAPTURE_RATE_MAX_RANGE_DAYS
SpendCaptureProvider = Literal["openai"]
class SpendCaptureRateCheckSettings(BaseModel):
"""``general_settings.spend_capture_rate_check``: the daily check of captured spend against the provider bill."""
model_config = ConfigDict(frozen=True, extra="forbid")
providers: tuple[SpendCaptureProvider, ...] = Field(("openai",), min_length=1)
threshold: float = Field(0.9, gt=0, le=1)
lookback_days: int = Field(7, ge=1, le=SPEND_CAPTURE_RATE_MAX_RANGE_DAYS)
openai_project_ids: tuple[str, ...] = Field(
(),
description=(
"Scope the OpenAI bill to these project ids; empty compares against the whole organization. Captured "
"spend is never scoped, so list every project LiteLLM's OpenAI keys belong to"
),
)
class CaptureRateDay(BaseModel):
model_config = ConfigDict(frozen=True)
date: str
captured_spend: float
provider_spend: float
capture_rate: float | None
class CaptureRateReport(BaseModel):
model_config = ConfigDict(frozen=True)
provider: SpendCaptureProvider
start_date: str
end_date: str
captured_spend: float
provider_spend: float
capture_rate: float | None
threshold: float
below_threshold: bool
days: tuple[CaptureRateDay, ...]

View file

@ -392,7 +392,6 @@ if TYPE_CHECKING:
from litellm.llms.base_llm.image_variations.transformation import (
BaseImageVariationConfig,
)
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.base_llm.realtime.http_transformation import (
BaseRealtimeHTTPConfig,
@ -430,7 +429,6 @@ if TYPE_CHECKING:
)
from litellm.llms.cohere.common_utils import CohereModelInfo
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.proxy._types import AllowedModelRegion
from litellm.router_utils.get_retry_from_policy import (
get_num_retries_from_retry_policy,
@ -9758,51 +9756,6 @@ class ProviderConfigManager:
return get_openrouter_image_edit_config(model)
return None
@staticmethod
def get_provider_ocr_config(
model: str,
provider: LlmProviders,
) -> BaseOCRConfig | None:
"""
Get OCR configuration for a given provider.
"""
from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig
# Special handling for Azure AI - distinguish between Mistral OCR and Document Intelligence
if provider == litellm.LlmProviders.AZURE_AI:
from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config
return get_azure_ai_ocr_config(model=model)
if provider == litellm.LlmProviders.VERTEX_AI:
from litellm.llms.vertex_ai.ocr.common_utils import get_vertex_ai_ocr_config
return get_vertex_ai_ocr_config(model=model)
if provider == litellm.LlmProviders.COHERE:
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
return CohereParseConfig()
if provider == litellm.LlmProviders.REDUCTO:
from litellm.llms.reducto.ocr.transformation import (
ReductoParseLegacyConfig,
ReductoParseV3Config,
)
if model == "parse-legacy":
return ReductoParseLegacyConfig()
return ReductoParseV3Config()
MistralOCRConfig: Final = litellm_utils.MistralOCRConfig
PROVIDER_TO_CONFIG_MAP: Final = {
litellm.LlmProviders.MISTRAL: MistralOCRConfig,
}
config_class: Final = PROVIDER_TO_CONFIG_MAP.get(provider, None)
if config_class is None:
return None
return config_class()
@staticmethod
def get_provider_search_config(
provider: SearchProviders,
@ -10273,7 +10226,7 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
"""
from datetime import datetime
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.litellm_logging import Logging, RawRequestCaptured
litellm_logging_obj: Final = Logging(
model="gpt-3.5-turbo",
@ -10284,6 +10237,7 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
start_time=datetime.now(),
function_id="1234",
log_raw_request_response=True,
raw_request_only=True,
)
llm_api_endpoint: Final = getattr(litellm, endpoint.value)
@ -10294,7 +10248,11 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
llm_api_endpoint(
**kwargs,
litellm_logging_obj=litellm_logging_obj,
api_key="my-fake-api-key", # 👈 ensure the request fails
api_key="my-fake-api-key",
)
except RawRequestCaptured:
received_exception = (
"raw request was not captured before the provider call; check the proxy logs for the pre_call error"
)
except Exception as e:
received_exception = str(e)

View file

@ -6092,7 +6092,7 @@
"cache_creation_input_audio_token_cost": 3e-07,
"cache_read_input_audio_token_cost": 3e-07,
"cache_read_input_token_cost": 6e-08,
"deprecation_date": "2027-06-25",
"deprecation_date": "2027-07-31",
"input_cost_per_audio_token": 1e-05,
"input_cost_per_image_token": 8e-07,
"input_cost_per_token": 6e-07,
@ -26207,12 +26207,10 @@
},
"gemini-2.5-flash-image": {
"deprecation_date": "2027-03-15",
"cache_read_input_token_cost": 3e-08,
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 3e-07,
"input_cost_per_token_batches": 1.5e-07,
"input_cost_per_token_flex": 1.5e-07,
"input_cost_per_token_priority": 5.4e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
@ -26313,10 +26311,14 @@
"gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_priority": 3.6e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_priority": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -26326,7 +26328,9 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_priority": 2.16e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
@ -27849,6 +27853,7 @@
"gemini-embedding-001": {
"deprecation_date": "2028-05-20",
"input_cost_per_token": 1.5e-07,
"input_cost_per_token_batches": 1.2e-07,
"litellm_provider": "vertex_ai-embedding-models",
"max_input_tokens": 2048,
"max_tokens": 2048,
@ -28393,6 +28398,8 @@
"input_cost_per_image": 0.0011,
"input_cost_per_token": 2e-06,
"input_cost_per_token_batches": 1e-06,
"input_cost_per_token_flex": 1e-06,
"input_cost_per_token_priority": 3.6e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 32768,
@ -28404,6 +28411,8 @@
"rpm": 1000,
"tpm": 4000000,
"output_cost_per_token_batches": 6e-06,
"output_cost_per_token_flex": 6e-06,
"output_cost_per_token_priority": 2.16e-05,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@ -41390,20 +41399,20 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro": {
"input_cost_per_token": 9.1263e-07,
"input_cost_per_token": 8.44944e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 1.82526e-06,
"output_cost_per_token": 1.689888e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"cache_read_input_token_cost": 7.60525e-08,
"cache_read_input_token_cost": 7.0412e-08,
"supports_audio_input": false,
"supports_pdf_input": false,
"supports_vision": false,
@ -49411,12 +49420,10 @@
},
"vertex_ai/gemini-2.5-flash-image": {
"deprecation_date": "2027-03-15",
"cache_read_input_token_cost": 3e-08,
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 3e-07,
"input_cost_per_token_batches": 1.5e-07,
"input_cost_per_token_flex": 1.5e-07,
"input_cost_per_token_priority": 5.4e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
@ -49494,10 +49501,14 @@
"vertex_ai/gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_priority": 3.6e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_priority": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -49507,7 +49518,9 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_priority": 2.16e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"supports_reasoning": false,
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
@ -56202,6 +56215,7 @@
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_per_second": 3.3333333333333335e-05,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -56394,6 +56408,7 @@
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_per_second": 3.3333333333333335e-05,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -60135,6 +60150,23 @@
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/routers/deepseek-v4p1-flash-us": {
"cache_read_input_token_cost": 9e-09,
"input_cost_per_token": 4.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 1.8e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-vision-exp": {
"cache_read_input_token_cost": 7e-09,
"deprecation_date": "2026-09-25",
@ -60214,6 +60246,23 @@
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/deepseek-v4p1-flash-us": {
"cache_read_input_token_cost": 9e-09,
"input_cost_per_token": 4.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 1.8e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/deepseek-v4-flash-vision-exp": {
"cache_read_input_token_cost": 7e-09,
"deprecation_date": "2026-09-25",
@ -60310,13 +60359,16 @@
},
"fireworks_ai/kimi-k3-us": {
"cache_read_input_token_cost": 4.5e-07,
"cache_read_input_token_cost_priority": 5.625e-07,
"input_cost_per_token": 4.5e-06,
"input_cost_per_token_priority": 5.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.25e-05,
"output_cost_per_token_priority": 2.8125e-05,
"reasoning_effort_levels": [
"low",
"high",
@ -60518,13 +60570,16 @@
},
"fireworks_ai/accounts/fireworks/routers/kimi-k3-us": {
"cache_read_input_token_cost": 4.5e-07,
"cache_read_input_token_cost_priority": 5.625e-07,
"input_cost_per_token": 4.5e-06,
"input_cost_per_token_priority": 5.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.25e-05,
"output_cost_per_token_priority": 2.8125e-05,
"reasoning_effort_levels": [
"low",
"high",
@ -63283,6 +63338,22 @@
"supports_tool_choice": true,
"supports_vision": false
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"input_cost_per_token": 2.1e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": false
},
"fireworks_ai/glm-5p3": {
"cache_read_input_token_cost": 2.6e-07,
"cache_read_input_token_cost_priority": 3.25e-07,
@ -63302,6 +63373,22 @@
"supports_tool_choice": true,
"supports_vision": false
},
"fireworks_ai/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"input_cost_per_token": 2.1e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": false
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-fast": {
"cache_read_input_token_cost": 3.9e-07,
"input_cost_per_token": 2.1e-06,
@ -63349,6 +63436,20 @@
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"input_cost_per_token": 2.25e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/glm-5p3-flash": {
"cache_read_input_token_cost": 3e-08,
"cache_read_input_token_cost_priority": 3.75e-08,
@ -63366,6 +63467,20 @@
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"input_cost_per_token": 2.25e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/models/inkling": {
"cache_read_input_token_cost": 1.7e-07,
"input_cost_per_token": 1e-06,
@ -65897,9 +66012,9 @@
"supports_web_search": false
},
"openrouter/z-ai/glm-5.3-flash": {
"input_cost_per_token": 1.5e-07,
"output_cost_per_token": 5e-07,
"cache_read_input_token_cost": 5e-08,
"input_cost_per_token": 4.5e-08,
"output_cost_per_token": 6e-07,
"cache_read_input_token_cost": 2.85e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
"max_output_tokens": 943718,
@ -66648,9 +66763,9 @@
"supports_web_search": true
},
"openrouter/deepseek/deepseek-v4-flash": {
"input_cost_per_token": 8.4e-08,
"output_cost_per_token": 1.68e-07,
"cache_read_input_token_cost": 1.68e-08,
"input_cost_per_token": 4.9e-08,
"output_cost_per_token": 9.8e-08,
"cache_read_input_token_cost": 9.8e-09,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
@ -69475,6 +69590,7 @@
"input_cost_per_audio_token": 3e-06,
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -69495,6 +69611,7 @@
"input_cost_per_audio_token": 3e-06,
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -73039,6 +73156,26 @@
"supports_vision": true,
"supports_web_search": false
},
"openrouter/mistralai/mistral-large-2512": {
"cache_read_input_token_cost": 5e-08,
"input_cost_per_token": 5e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 209715,
"max_tokens": 209715,
"mode": "chat",
"output_cost_per_token": 1.5e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": false,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
},
"openrouter/mistralai/mistral-large-2512:batch": {
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_token": 2.5e-07,

View file

@ -65,7 +65,5 @@ max-args = 5
"litellm.responses.main.aresponses".msg = "Import litellm.responses.dispatch.aresponses so the call routes through dispatch."
"litellm.llms.anthropic.experimental_pass_through.messages.handler.anthropic_messages".msg = "Import litellm.messages.anthropic_messages so the call routes through dispatch."
"litellm.llms.anthropic.experimental_pass_through.messages.handler.anthropic_messages_handler".msg = "Import litellm.messages.anthropic_messages_handler so the call routes through dispatch."
"litellm.ocr.main.ocr".msg = "Import litellm.ocr.dispatch.ocr so the call routes through dispatch."
"litellm.ocr.main.aocr".msg = "Import litellm.ocr.dispatch.aocr so the call routes through dispatch."
"litellm.main.completion".msg = "Import litellm.completion so the call routes through dispatch."
"litellm.main.acompletion".msg = "Import litellm.acompletion so the call routes through dispatch."

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