merge: resolve main into litellm_mcp_continuous_tool_defaults

Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
Devin AI 2026-09-25 02:51:29 +00:00
commit 878c646e99
588 changed files with 13839 additions and 68379 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

@ -171,12 +171,21 @@ jobs:
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
@ -194,22 +203,34 @@ jobs:
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.flag >> shard"
no_output_timeout: 20m
command: |
mkdir -p test-results/<< parameters.flag >>
selection="$(find tests/unit -name 'test_*.py' | sort)" || { echo "test selection failed for << parameters.flag >>"; exit 1; }
[ -n "${selection}" ] || { echo "test selection produced no files for << parameters.flag >>"; exit 1; }
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
env -i "${test_env[@]}" \
uv run --no-sync pytest "${files[@]}" "${rerun_args[@]}" -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
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
@ -293,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

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

@ -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_proxy_server_gemini_pass_through.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_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

@ -96,6 +96,7 @@ Follow these coding conventions for new/updated code (a three-line fix in a lega
- No mutation; don't reassign variables, global or local. Instead of mutable lists and dicts, prefer tuples, frozen dataclasses (with slots=True), `MappingProxyType`, etc.
- Annotate every variable with `: Final` (LIT010). Unpacking and walrus targets cannot carry the annotation, so they are implicitly final. Don't rebind them. Never rebind or mutate function parameters (LIT011); `self`/`cls` attribute stores are the exception. If rebinding or in-place mutation is truly unavoidable, suppress with `# rebind-ok: <reason>`
- Qualify every TypedDict field with `ReadOnly[...]` (LIT012), which nests freely with `Required` / `NotRequired` / `Annotated` in any order. If making the key writable is truly unavoidable, suppress with `# writable-ok: <reason>`
- Comprehensions take at most one `for` clause and one `if` clause (LIT014); split stacked clauses into a helper generator, a named intermediate, or a plain loop. Suppress with `# comprehension-ok: <reason>` only when unavoidable
- Use dependency injection
- Fully typed; no `Any` or coarse types like `dict[str, Any]` or just `dict`. Every function parameter must be strongly typed
- Use tagged unions + match

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

@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Final, List, Literal, Optional, Protocol, Tupl
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import (
CLI_SESSION_KEY_PREFIX,
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS,
MAX_OBJECTS_PER_POLL_CYCLE,
)
@ -147,10 +148,12 @@ class CheckBatchCost:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}")
return {}
async def _get_key_alias(self, batch_id: str, api_key: str | None) -> str | None:
async def _get_key_alias(self, batch_id: str, api_key: str | None, created_by: str | None) -> str | None:
"""Resolve the creating virtual key's alias from its hashed token."""
if not api_key:
return None
if created_by and api_key == f"{CLI_SESSION_KEY_PREFIX}-{created_by}":
return api_key
try:
key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table(
self.prisma_client
@ -231,7 +234,7 @@ class CheckBatchCost:
**(await self._get_user_info(batch_id, job.created_by)),
}
key_alias = await self._get_key_alias(batch_id, api_key)
key_alias = await self._get_key_alias(batch_id, api_key, job.created_by)
if key_alias is not None:
metadata["user_api_key_alias"] = key_alias
team_alias = await self._get_team_alias(team_id)

View file

@ -50,6 +50,7 @@ from litellm.proxy._types import (
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.openai_files_endpoints.common_utils import (
BATCH_CREATE_HIDDEN_PARAM,
FILE_LIST_CONTINUATION_CHUNK_SIZE,
@ -359,7 +360,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
from prisma import Json
api_key = user_api_key_dict.api_key or None
api_key = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict) or None
attribution_columns = (
{
**({"api_key": api_key} if api_key is not None else {}),

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,
@ -2605,6 +2623,7 @@ class PrometheusLogger(CustomLogger):
from litellm.litellm_core_utils.litellm_logging import (
StandardLoggingPayloadSetup,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
status_code: Final = self._extract_status_code(exception=original_exception)
@ -2623,7 +2642,9 @@ class PrometheusLogger(CustomLogger):
end_user=user_api_key_dict.end_user_id,
user=user_api_key_dict.user_id,
user_email=user_api_key_dict.user_email,
hashed_api_key=None if status_code == 401 else user_api_key_dict.api_key,
hashed_api_key=None
if status_code == 401
else LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
api_key_alias=user_api_key_dict.key_alias,
team=user_api_key_dict.team_id,
team_alias=user_api_key_dict.team_alias,

View file

@ -1687,7 +1687,7 @@ class WebSearchInterceptionLogger(CustomLogger):
**user_api_key_metadata,
**parent_correlation.as_search_metadata(),
"model_group": search_tool_name,
"user_api_key": user_api_key_auth.api_key,
"user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_auth),
"user_api_key_auth": user_api_key_auth,
}

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

@ -142,14 +142,35 @@ def _resolution_key(resolution: object) -> str | None:
return str(resolution)
def _resolution_cost_per_image(entry: Mapping[str, object] | None, resolution: object) -> float | None:
resolution_key: Final = _resolution_key(resolution)
if entry is None or resolution_key is None:
return None
cost: Final = entry.get(f"output_cost_per_image_{resolution_key}")
return float(cost) if isinstance(cost, (int, float)) else None
def _requested_image_count(request_body: Mapping[str, object]) -> int:
num_images: Final = request_body.get("num_images")
return num_images if type(num_images) is int and num_images > 0 else 1
def _passthrough_cost_per_image(entry: Mapping[str, object], request_body: Mapping[str, object]) -> float | None:
resolution_cost: Final = _resolution_cost_per_image(entry, request_body.get("resolution"))
if resolution_cost is not None:
return resolution_cost
cost: Final = entry.get("output_cost_per_image")
return float(cost) if isinstance(cost, (int, float)) else None
def fal_ai_passthrough_cost(model: str, request_body: Mapping[str, object]) -> float | None:
entry: Final = _entry(f"{litellm.LlmProviders.FAL_AI.value}/{model}")
if entry is None:
return None
resolution: Final = _resolution_key(request_body.get("resolution"))
keyed_cost: Final = entry.get(f"output_cost_per_image_{resolution}") if resolution is not None else None
cost: Final = keyed_cost if isinstance(keyed_cost, (int, float)) else entry.get("output_cost_per_image")
return float(cost) if isinstance(cost, (int, float)) else None
cost_per_image: Final = _passthrough_cost_per_image(entry, request_body)
if cost_per_image is None:
return None
return cost_per_image * _requested_image_count(request_body)
def cost_calculator(
@ -172,6 +193,11 @@ def cost_calculator(
if deployment_cost_per_image is not None:
return deployment_cost_per_image * len(images)
params: Final[Mapping[str, object]] = optional_params or MappingProxyType({})
resolution_cost_per_image: Final = _resolution_cost_per_image(
_entry(f"{litellm.LlmProviders.FAL_AI.value}/{normalized_model}"), params.get("resolution")
)
if resolution_cost_per_image is not None:
return resolution_cost_per_image * len(images)
keyed_costs: Final = tuple(
_keyed_cost_per_image(
model=normalized_model,

View file

@ -8,12 +8,17 @@ from .transformation import FalAIBaseConfig
class FalAINanoBananaConfig(FalAIBaseConfig):
"""
Configuration for Fal AI's Nano Banana / Gemini 2.5 Flash Image models.
Configuration for Fal AI's Nano Banana family (Gemini Flash / Pro Image models).
Serves the imagen4 deprecation migration path. The same underlying model is
exposed under two endpoints that share an identical schema:
Serves the imagen4 deprecation migration path. Every endpoint shares the same
request schema, so one config covers all of them:
- fal-ai/nano-banana
- fal-ai/gemini-25-flash-image
- fal-ai/nano-banana-2
- fal-ai/nano-banana-pro
Provider-specific params such as ``resolution`` ("0.5K", "1K", "2K", "4K") are
forwarded as-is and drive the per-resolution price in the cost map.
Documentation: https://fal.ai/models/fal-ai/nano-banana
"""

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

@ -1,4 +1,5 @@
from collections.abc import Mapping
from collections.abc import Iterable, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol
import httpx
@ -57,23 +58,57 @@ class VertexSearchSnippet(TypedDict, total=False):
htmlSnippet: ReadOnly[str]
class VertexSearchExtractiveContent(TypedDict, total=False):
"""One ``extractive_answers`` or ``extractive_segments`` entry (opt-in via ``extractiveContentSpec``)."""
content: ReadOnly[str]
pageNumber: ReadOnly[str]
class VertexSearchDerivedStructData(TypedDict, total=False):
"""The ``derivedStructData`` blob Discovery Engine attaches to each search hit."""
"""The ``derivedStructData`` blob Discovery Engine attaches to each document hit."""
title: ReadOnly[str]
link: ReadOnly[str]
displayLink: ReadOnly[str]
formattedUrl: ReadOnly[str]
snippets: ReadOnly[list[VertexSearchSnippet]]
extractive_answers: ReadOnly[list[VertexSearchExtractiveContent]]
extractive_segments: ReadOnly[list[VertexSearchExtractiveContent]]
class VertexSearchDocument(TypedDict, total=False):
id: ReadOnly[str]
structData: ReadOnly[Mapping[str, object]]
derivedStructData: ReadOnly[VertexSearchDerivedStructData]
class VertexSearchChunkDocumentMetadata(TypedDict, total=False):
uri: ReadOnly[str]
title: ReadOnly[str]
structData: ReadOnly[Mapping[str, object]]
class VertexSearchChunkPageSpan(TypedDict, total=False):
pageStart: ReadOnly[int]
pageEnd: ReadOnly[int]
class VertexSearchChunk(TypedDict, total=False):
"""A hit when ``searchResultMode`` is ``CHUNKS``; such hits carry no ``document`` and no top-level ``id``."""
id: ReadOnly[str]
name: ReadOnly[str]
content: ReadOnly[str]
documentMetadata: ReadOnly[VertexSearchChunkDocumentMetadata]
pageSpan: ReadOnly[VertexSearchChunkPageSpan]
relevanceScore: ReadOnly[float]
class VertexSearchHit(TypedDict, total=False):
id: ReadOnly[str]
document: ReadOnly[VertexSearchDocument]
chunk: ReadOnly[VertexSearchChunk]
class VertexSearchApiResponse(TypedDict, total=False):
@ -98,6 +133,97 @@ def _vertex_search_payload(response: _VertexSearchApiSource) -> VertexSearchApiR
return response.json()
_UNKNOWN_DOCUMENT: Final = "Unknown Document"
_EMPTY_DOCUMENT: Final[VertexSearchDocument] = {}
_EMPTY_DERIVED_STRUCT_DATA: Final[VertexSearchDerivedStructData] = {}
_EMPTY_CHUNK_DOCUMENT_METADATA: Final[VertexSearchChunkDocumentMetadata] = {}
def _joined_content(entries: Sequence[VertexSearchExtractiveContent]) -> str:
return "\n\n".join(content for entry in entries if (content := entry.get("content")))
def _snippet_text(snippets: Sequence[VertexSearchSnippet]) -> str:
return " ".join(snippet.get("snippet", snippet.get("htmlSnippet", "")) for snippet in snippets)
def _document_text(derived: VertexSearchDerivedStructData) -> str:
candidates: Final = (
_joined_content(derived.get("extractive_segments", ())),
_joined_content(derived.get("extractive_answers", ())),
_snippet_text(derived.get("snippets", ())),
derived.get("title", ""),
)
return next((text for text in candidates if text), "")
def _document_id_from_chunk_name(name: str) -> str:
return name.partition("/documents/")[2].partition("/")[0]
def _non_empty_attributes(pairs: Iterable[tuple[str, object]]) -> Mapping[str, object]:
return MappingProxyType({key: value for key, value in pairs if value})
def _chunk_result(chunk: VertexSearchChunk, positional_score: float) -> VectorStoreSearchResult:
metadata: Final = chunk.get("documentMetadata", _EMPTY_CHUNK_DOCUMENT_METADATA)
uri: Final = metadata.get("uri", "")
title: Final = metadata.get("title", "")
document_id: Final = _document_id_from_chunk_name(chunk.get("name", ""))
return VectorStoreSearchResult(
score=chunk.get("relevanceScore", positional_score),
content=[VectorStoreResultContent(text=chunk.get("content", ""), type="text")],
file_id=uri or document_id,
filename=title or _UNKNOWN_DOCUMENT,
attributes={
"document_id": document_id,
**_non_empty_attributes(
(
("chunk_id", chunk.get("id", "")),
("link", uri),
("title", title),
("structData", metadata.get("structData")),
("pageSpan", chunk.get("pageSpan")),
)
),
},
)
def _document_result(hit: VertexSearchHit, score: float) -> VectorStoreSearchResult:
document: Final = hit.get("document", _EMPTY_DOCUMENT)
derived: Final = document.get("derivedStructData", _EMPTY_DERIVED_STRUCT_DATA)
link: Final = derived.get("link", "")
title: Final = derived.get("title", "")
document_id: Final = hit.get("id", "")
return VectorStoreSearchResult(
score=score,
content=[VectorStoreResultContent(text=_document_text(derived), type="text")],
file_id=link or document_id,
filename=title or _UNKNOWN_DOCUMENT,
attributes={
"document_id": document_id,
**_non_empty_attributes(
(
("link", link),
("title", title),
("displayLink", derived.get("displayLink", "")),
("formattedUrl", derived.get("formattedUrl", "")),
("structData", document.get("structData")),
)
),
},
)
def _search_result(hit: VertexSearchHit, position: int) -> VectorStoreSearchResult:
score: Final = 1.0 / (position + 1)
chunk: Final = hit.get("chunk")
if chunk is not None:
return _chunk_result(chunk, score)
return _document_result(hit, score)
class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
"""
Configuration for Vertex AI Search API Vector Store
@ -285,98 +411,19 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj
) -> VectorStoreSearchResponse:
"""
Transform Vertex AI Search API response to standard vector store search response
Transform a Discovery Engine ``:search`` response into the standard vector store search response.
Handles the format from Discovery Engine Search API which returns:
{
"results": [
{
"id": "...",
"document": {
"derivedStructData": {
"title": "...",
"link": "...",
"snippets": [...]
}
}
}
]
}
Document hits (``results[].document``) take their text from ``derivedStructData`` in a fixed order:
``extractive_segments``, then ``extractive_answers``, then ``snippets``, then ``title``; ``structData``
and the link metadata land in ``attributes``. Chunk hits (``results[].chunk``, returned when the
caller sets ``contentSearchSpec.searchResultMode`` to ``CHUNKS`` via ``extra_body``) take their text
from ``chunk.content`` and their file id and name from ``chunk.documentMetadata``.
"""
try:
response_json: Final = _vertex_search_payload(response)
# Extract results from Vertex AI Search API response
results: Final = response_json.get("results", [])
# Transform results to standard format
search_results: Final[list[VectorStoreSearchResult]] = []
for result in results:
document: VertexSearchDocument = result.get("document", {})
derived_data: VertexSearchDerivedStructData = document.get("derivedStructData", {})
# Extract text content from snippets
snippets = derived_data.get("snippets", [])
text_content = ""
if snippets:
# Combine all snippets into one text
text_parts = [snippet.get("snippet", snippet.get("htmlSnippet", "")) for snippet in snippets]
text_content = " ".join(text_parts)
# If no snippets, use title as fallback
if not text_content:
text_content = derived_data.get("title", "")
content = [
VectorStoreResultContent(
text=text_content,
type="text",
)
]
# Extract file/document information
document_link = derived_data.get("link", "")
document_title = derived_data.get("title", "")
document_id = result.get("id", "")
# Use link as file_id if available, otherwise use document ID
file_id = document_link if document_link else document_id
filename = document_title if document_title else "Unknown Document"
# Build attributes with available metadata
attributes = {
"document_id": document_id,
}
if document_link:
attributes["link"] = document_link
if document_title:
attributes["title"] = document_title
# Add display link if available
display_link = derived_data.get("displayLink", "")
if display_link:
attributes["displayLink"] = display_link
# Add formatted URL if available
formatted_url = derived_data.get("formattedUrl", "")
if formatted_url:
attributes["formattedUrl"] = formatted_url
# Note: Search API doesn't provide explicit scores in the response
# You can use the position/rank as an implicit score
score = 1.0 / (float(search_results.__len__() + 1)) # Decreasing score based on position
result_obj = VectorStoreSearchResult(
score=score,
content=content,
file_id=file_id,
filename=filename,
attributes=attributes,
)
search_results.append(result_obj)
search_results: Final = [
_search_result(hit, position) for position, hit in enumerate(response_json.get("results", ()))
]
query_view: Final[_SearchQueryView] = {"query": litellm_logging_obj.model_call_details.get("query", "")}
return VectorStoreSearchResponse(
object="vector_store.search_results.page",

View file

@ -12,6 +12,7 @@ from collections.abc import Callable
from typing import TYPE_CHECKING, Any, Final, cast
import httpx
from pydantic import ValidationError
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.custom_httpx.http_handler import (
@ -21,6 +22,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai_gemma import VertexGemmaContainerError
from litellm.types.utils import ModelResponse
if TYPE_CHECKING:
@ -29,6 +31,13 @@ if TYPE_CHECKING:
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
def parse_vertex_gemma_container_error(predictions: object) -> VertexGemmaContainerError | None:
try:
return VertexGemmaContainerError.model_validate(predictions)
except ValidationError:
return None
class VertexGemmaConfig(OpenAIGPTConfig):
"""
Configuration and transformation class for Vertex AI Gemma models
@ -123,7 +132,9 @@ class VertexGemmaConfig(OpenAIGPTConfig):
Unwrap the Vertex Gemma predictions format to OpenAI format.
Vertex Gemma wraps the OpenAI-compatible response in a 'predictions' field.
This method extracts it so the parent class can process it normally.
This method extracts it so the parent class can process it normally. A serving
container can also answer with its own OpenAI-shaped error object inside that
field, still under HTTP 200, which is raised with its own status and message.
"""
if "predictions" not in response_json:
raise BaseLLMException(
@ -131,7 +142,11 @@ class VertexGemmaConfig(OpenAIGPTConfig):
message="Invalid response format: missing 'predictions' field",
)
return response_json["predictions"]
predictions: Final = response_json["predictions"]
container_error: Final = parse_vertex_gemma_container_error(predictions)
if container_error is None:
return predictions
raise BaseLLMException(status_code=container_error.code, message=container_error.message)
@staticmethod
def _sync_post(

View file

@ -3301,12 +3301,14 @@
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"azure/ada": {
"deprecation_date": "2028-02-09",
"input_cost_per_token": 1e-07,
"litellm_provider": "azure",
"max_input_tokens": 8191,
"max_tokens": 8191,
"mode": "embedding",
"output_cost_per_token": 0.0
"output_cost_per_token": 0.0,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure/codex-mini": {
"cache_read_input_token_cost": 3.75e-07,
@ -5151,6 +5153,7 @@
"supports_tool_choice": true
},
"azure/gpt-35-turbo-16k": {
"deprecation_date": "2025-04-30",
"input_cost_per_token": 3e-06,
"litellm_provider": "azure",
"max_input_tokens": 16385,
@ -5158,9 +5161,11 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models",
"supports_tool_choice": true
},
"azure/gpt-35-turbo-16k-0613": {
"deprecation_date": "2025-04-30",
"input_cost_per_token": 3e-06,
"litellm_provider": "azure",
"max_input_tokens": 16385,
@ -5168,6 +5173,7 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models",
"supports_function_calling": true,
"supports_tool_choice": true
},
@ -5211,6 +5217,7 @@
"supports_tool_choice": true
},
"azure/gpt-4-0613": {
"deprecation_date": "2025-06-06",
"input_cost_per_token": 3e-05,
"litellm_provider": "azure",
"max_input_tokens": 8192,
@ -5218,6 +5225,7 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 6e-05,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models",
"supports_function_calling": true,
"supports_tool_choice": true
},
@ -5234,6 +5242,7 @@
"supports_tool_choice": true
},
"azure/gpt-4-32k": {
"deprecation_date": "2025-06-06",
"input_cost_per_token": 6e-05,
"litellm_provider": "azure",
"max_input_tokens": 32768,
@ -5241,9 +5250,11 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 0.00012,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models",
"supports_tool_choice": true
},
"azure/gpt-4-32k-0613": {
"deprecation_date": "2025-06-06",
"input_cost_per_token": 6e-05,
"litellm_provider": "azure",
"max_input_tokens": 32768,
@ -5251,6 +5262,7 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 0.00012,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models",
"supports_tool_choice": true
},
"azure/gpt-4-turbo": {
@ -5511,6 +5523,7 @@
},
"azure/gpt-4.5-preview": {
"cache_read_input_token_cost": 3.75e-05,
"deprecation_date": "2025-07-14",
"input_cost_per_token": 7.5e-05,
"input_cost_per_token_batches": 3.75e-05,
"litellm_provider": "azure",
@ -5520,6 +5533,7 @@
"mode": "chat",
"output_cost_per_token": 0.00015,
"output_cost_per_token_batches": 7.5e-05,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
@ -15797,6 +15811,7 @@
"supports_tool_choice": true
},
"computer-use-preview": {
"deprecation_date": "2026-07-23",
"input_cost_per_token": 3e-06,
"litellm_provider": "openai",
"max_input_tokens": 8192,
@ -22838,6 +22853,37 @@
"/v1/images/generations"
]
},
"fal_ai/fal-ai/nano-banana-2": {
"litellm_provider": "fal_ai",
"metadata": {
"comment": "priced by the request's resolution field (0.5K, 1K default, 2K, 4K); the web search and high thinking surcharges are not modeled"
},
"mode": "image_generation",
"output_cost_per_image": 0.08,
"output_cost_per_image_0.5K": 0.06,
"output_cost_per_image_1K": 0.08,
"output_cost_per_image_2K": 0.12,
"output_cost_per_image_4K": 0.16,
"source": "https://fal.ai/models/fal-ai/nano-banana-2",
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/nano-banana-pro": {
"litellm_provider": "fal_ai",
"metadata": {
"comment": "priced by the request's resolution field (1K default, 2K, 4K); the web search surcharge is not modeled"
},
"mode": "image_generation",
"output_cost_per_image": 0.15,
"output_cost_per_image_1K": 0.15,
"output_cost_per_image_2K": 0.15,
"output_cost_per_image_4K": 0.3,
"source": "https://fal.ai/models/fal-ai/nano-banana-pro",
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/openai/gpt-image-2": {
"litellm_provider": "fal_ai",
"metadata": {
@ -26207,7 +26253,6 @@
},
"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,
@ -26312,10 +26357,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,
@ -26325,7 +26374,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": [
@ -27848,6 +27899,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,
@ -28392,6 +28444,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,
@ -28403,6 +28457,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",
@ -28738,6 +28794,7 @@
"mode": "audio_speech",
"output_cost_per_audio_token": 1e-05,
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
"/v1/audio/speech"
@ -29813,6 +29870,7 @@
"mode": "chat",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2e-05,
"output_cost_per_token_batches": 1e-05,
"rpm": 10000,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_modalities": [
@ -32647,6 +32705,7 @@
},
"gpt-5-codex": {
"cache_read_input_token_cost": 1.25e-07,
"deprecation_date": "2026-07-23",
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 272000,
@ -32730,6 +32789,7 @@
},
"gpt-5.1-codex-mini": {
"cache_read_input_token_cost": 2.5e-08,
"deprecation_date": "2026-07-23",
"input_cost_per_token": 2.5e-07,
"litellm_provider": "openai",
"max_input_tokens": 400000,
@ -32760,6 +32820,7 @@
},
"gpt-5.1-codex-max": {
"cache_read_input_token_cost": 1.25e-07,
"deprecation_date": "2026-07-23",
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 400000,
@ -32791,6 +32852,7 @@
},
"gpt-5.1-codex": {
"cache_read_input_token_cost": 1.25e-07,
"deprecation_date": "2026-07-23",
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 400000,
@ -32822,6 +32884,7 @@
},
"gpt-5.1-chat-latest": {
"cache_read_input_token_cost": 1.25e-07,
"deprecation_date": "2026-07-23",
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 128000,
@ -32958,6 +33021,7 @@
},
"gpt-5.2-codex": {
"cache_read_input_token_cost": 1.75e-07,
"deprecation_date": "2026-07-23",
"input_cost_per_token": 1.75e-06,
"litellm_provider": "openai",
"max_input_tokens": 272000,
@ -32989,6 +33053,7 @@
},
"gpt-5.2-chat-latest": {
"cache_read_input_token_cost": 1.75e-07,
"deprecation_date": "2026-08-10",
"input_cost_per_token": 1.75e-06,
"litellm_provider": "openai",
"max_input_tokens": 128000,
@ -34812,6 +34877,7 @@
},
"gpt-5.3-chat-latest": {
"cache_read_input_token_cost": 1.75e-07,
"deprecation_date": "2026-08-10",
"input_cost_per_token": 1.75e-06,
"litellm_provider": "openai",
"max_input_tokens": 128000,
@ -41389,20 +41455,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,
@ -49410,7 +49476,6 @@
},
"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,
@ -49492,10 +49557,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,
@ -49505,7 +49574,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"
@ -56200,6 +56271,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,
@ -56392,6 +56464,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,
@ -56434,6 +56507,7 @@
"mode": "audio_speech",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2e-05,
"output_cost_per_token_batches": 1e-05,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
"/v1/audio/speech"
@ -56502,6 +56576,7 @@
"mode": "audio_speech",
"output_cost_per_audio_token": 1e-05,
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
"/v1/audio/speech"
@ -60133,6 +60208,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",
@ -60212,6 +60304,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",
@ -60308,13 +60417,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",
@ -60516,13 +60628,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",
@ -63281,6 +63396,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,
@ -63300,6 +63431,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,
@ -63347,6 +63494,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,
@ -63364,6 +63525,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,
@ -65895,9 +66070,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,
@ -66646,9 +66821,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,
@ -69473,6 +69648,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,
@ -69493,6 +69669,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,
@ -73037,6 +73214,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

@ -2117,7 +2117,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.
@ -2371,8 +2371,11 @@ class MCPRequestHandler:
inventory=resolved_inventory,
)
)
after_caller: Final = _as_list(
await MCPRequestHandler._apply_agent_caller_tool_ceiling(after_user, server_id, user_api_key_auth)
)
allowed_tools: Final = await MCPRequestHandler._apply_agent_and_org_tool_ceilings(
after_user, server_id, user_api_key_auth, keyless_source=keyless_source, inventory=resolved_inventory
after_caller, server_id, user_api_key_auth, keyless_source=keyless_source, inventory=resolved_inventory
)
if allowed_tools is not None:
return allowed_tools
@ -3295,6 +3298,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": [
{
@ -3569,6 +3804,16 @@
],
"title": "Extra Headers"
},
"kill_switch": {
"anyOf": [
{
"$ref": "#/components/schemas/AgentKillSwitchConfig"
},
{
"type": "null"
}
]
},
"litellm_params": {
"additionalProperties": true,
"title": "Litellm Params",
@ -4331,6 +4576,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

@ -52,6 +52,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
@ -239,6 +240,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):
@ -579,6 +581,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
@ -778,6 +781,7 @@ class LiteLLMRoutes(enum.Enum):
"/global/spend/provider",
"/global/spend/tags",
"/global/spend/all_tag_names",
"/spend/capture_rate",
]
public_routes = frozenset(
@ -2948,6 +2952,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.",
@ -3690,7 +3702,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

@ -656,7 +656,7 @@ class LiteLLMExecutedBatchRunner:
def _row_metadata(self, run: _BatchRun) -> dict[str, object]: # mutable-ok: router updates metadata in place
return { # mutable-ok: the router updates request metadata in place
**LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(run.user_api_key_dict),
"user_api_key": run.user_api_key_dict.api_key,
"user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(run.user_api_key_dict),
"user_api_end_user_max_budget": run.user_api_key_dict.end_user_max_budget,
"tags": list(run.request_tags), # mutable-ok: litellm types request tags as a list
"batch_id": run.unified_batch_id,

View file

@ -68,7 +68,7 @@ def embedding_spend_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, obj
return { # mutable-ok: the router mutates the metadata dict it is handed
**LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict),
"user_api_key": user_api_key_dict.api_key,
"user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
}

View file

@ -20,6 +20,7 @@ from litellm.proxy.auth.auth_utils import (
from litellm.proxy.auth.budget_throttle import throttled_limit
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.types.utils import Usage
if TYPE_CHECKING:
@ -250,7 +251,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
call_type: str,
):
self.print_verbose("Inside Max Parallel Request Pre-Call Hook")
api_key: Final = user_api_key_dict.api_key
api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
max_parallel_requests = user_api_key_dict.max_parallel_requests
if max_parallel_requests is None:
max_parallel_requests = sys.maxsize
@ -803,7 +804,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
"""
Retrieve the key's remaining rate limits.
"""
api_key: Final = user_api_key_dict.api_key
api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
current_date: Final = datetime.now().strftime("%Y-%m-%d")
current_hour: Final = datetime.now().strftime("%H")
current_minute: Final = datetime.now().strftime("%M")

View file

@ -472,11 +472,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]]
@ -506,6 +507,7 @@ class WindowKeyMetadata(TypedDict):
tokens_limit: int | None
window_size: int
descriptor_key: str
descriptor_value: ReadOnly[str]
class AtomicCounterMeta(TypedDict):
@ -582,6 +584,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(
@ -647,10 +690,12 @@ 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,
):
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
if self.internal_usage_cache.dual_cache.redis_cache is not None:
self.batch_rate_limiter_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
@ -1185,6 +1230,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"],
}
)
@ -1489,6 +1535,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
@ -2734,6 +2781,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,
@ -3110,7 +3168,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"
@ -3519,6 +3582,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(
@ -3613,6 +3677,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:
@ -4299,6 +4368,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
@ -4354,6 +4424,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(
@ -4528,6 +4599,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

@ -165,7 +165,7 @@ class _ProxyDBLogger(CustomLogger):
_metadata = dict(
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
)
_metadata["user_api_key"] = user_api_key_dict.api_key
_metadata["user_api_key"] = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
_metadata["status"] = "failure"
_error_information = StandardLoggingPayloadSetup.get_error_information(
original_exception=original_exception,
@ -259,7 +259,7 @@ class _ProxyDBLogger(CustomLogger):
)
await self._spend_writer().update_database(
token=user_api_key_dict.api_key,
token=LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
response_cost=recovered_response_cost,
user_id=user_api_key_dict.user_id,
end_user_id=user_api_key_dict.end_user_id,

View file

@ -1602,6 +1602,12 @@ class LiteLLMProxyRequestSetup:
data[_metadata_variable_name].update(metadata_from_headers)
return data
@staticmethod
def get_logged_api_key(user_api_key_dict: UserAPIKeyAuth) -> str | None:
if user_api_key_dict.is_session_token and user_api_key_dict.key_alias:
return user_api_key_dict.key_alias
return user_api_key_dict.api_key
@staticmethod
def get_sanitized_user_information_from_key(
user_api_key_dict: UserAPIKeyAuth,
@ -1609,7 +1615,7 @@ class LiteLLMProxyRequestSetup:
stripped_metadata: Final = strip_callback_config(user_api_key_dict.metadata)
auth_metadata: Final = cast("dict[str, str] | None", stripped_metadata) # cast-ok: metadata is free-form JSON
user_api_key_logged_metadata: Final = StandardLoggingUserAPIKeyMetadata(
user_api_key_hash=user_api_key_dict.api_key, # just the hashed token
user_api_key_hash=LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
user_api_key_alias=user_api_key_dict.key_alias,
user_api_key_spend=user_api_key_dict.spend,
user_api_key_max_budget=user_api_key_dict.max_budget,
@ -1647,7 +1653,7 @@ class LiteLLMProxyRequestSetup:
user_api_key_dict=user_api_key_dict
)
data[_metadata_variable_name].update(user_api_key_logged_metadata)
data[_metadata_variable_name]["user_api_key"] = user_api_key_dict.api_key # this is just the hashed token
data[_metadata_variable_name]["user_api_key"] = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
# Key-owned agent_id for spend attribution; keep existing (e.g. from header) if key has none
_key_agent_id: Final = getattr(user_api_key_dict, "agent_id", None)

View file

@ -15,7 +15,8 @@ from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT
from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.spend_tracking.daily_global_spend_rollup import GLOBAL_SPEND_TABLE_NAME, reconciled_through
from litellm.proxy.spend_tracking.key_metadata_recovery import (
attach_user_emails,
attach_user_details,
recover_cli_session_key_metadata,
recover_double_hashed_key_metadata,
recover_key_metadata_from_spend_logs,
)
@ -551,11 +552,12 @@ async def get_api_key_metadata(
e,
)
still_missing: Final = api_keys - frozenset(result)
from_session_keys: Final = await recover_cli_session_key_metadata(prisma_client, api_keys - frozenset(result))
still_missing: Final = api_keys - frozenset(result) - frozenset(from_session_keys)
from_reverse_hash: Final = (
await recover_double_hashed_key_metadata(prisma_client, still_missing) if still_missing else _EMPTY_KEY_METADATA
)
after_token_recovery: Final = MappingProxyType({**result, **from_reverse_hash})
after_token_recovery: Final = MappingProxyType({**result, **from_session_keys, **from_reverse_hash})
unresolved: Final = api_keys - frozenset(after_token_recovery)
from_spend_logs: Final = (
await recover_key_metadata_from_spend_logs(prisma_client, unresolved, spend_logs_window)
@ -563,7 +565,7 @@ async def get_api_key_metadata(
else _EMPTY_KEY_METADATA
)
combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs})
return await attach_user_emails(prisma_client, combined)
return await attach_user_details(prisma_client, combined)
def _adjust_dates_for_timezone(

View file

@ -26,6 +26,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeV
import fastapi
import yaml
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
from pydantic import TypeAdapter
from typing_extensions import ReadOnly, TypedDict
import litellm
@ -99,6 +100,11 @@ from litellm.proxy.management_endpoints.model_management_endpoints import (
_add_model_to_db,
)
from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights
from litellm.proxy.management_endpoints.team_admin_field_permissions import (
team_admin_key_edit_verdict,
team_admin_key_request_or_raise,
team_admin_may_edit_member_key_budgets,
)
from litellm.proxy.management_helpers.access_group_key_sync import (
sync_key_access_group_membership,
sync_key_regeneration_access_group_membership,
@ -3008,6 +3014,55 @@ async def _validate_end_user_budget_id_change(
raise HTTPException(status_code=400, detail=missing_detail)
_GENERAL_SETTINGS: Final = TypeAdapter(dict[str, object])
def _general_settings() -> Mapping[str, object]:
from litellm.proxy.proxy_server import (
general_settings, # pyright: ignore[reportUnknownVariableType] # untyped module-level dict in proxy_server
)
return _GENERAL_SETTINGS.validate_python(general_settings)
async def _acting_as_team_admin_for_key_update(
data: UpdateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
user_api_key_dict: UserAPIKeyAuth,
checked_prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
is_proxy_admin: bool,
) -> bool:
"""Whether the caller acts as a team admin on another member's team key.
Raises 403 when the caller administers the key's team but the request edits fields
outside the member_key_budgets permission (or that permission is disabled).
"""
if (
is_proxy_admin
or existing_key_row.team_id is None
or existing_key_row.user_id is None
or existing_key_row.user_id == user_api_key_dict.user_id
):
return False
team_for_grant: Final = await get_team_object(
team_id=existing_key_row.team_id,
prisma_client=checked_prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=True,
)
if not _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_for_grant):
return False
team_admin_key_request_or_raise(
team_admin_key_edit_verdict(
data=data,
existing=existing_key_row,
enabled=team_admin_may_edit_member_key_budgets(_general_settings()),
)
)
return True
async def _validate_update_key_data(
data: UpdateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
@ -3058,10 +3113,19 @@ async def _validate_update_key_data(
)
is_project_change: Final = "project_id" in data.model_fields_set and data.project_id != existing_key_row.project_id
acting_as_team_admin: Final = await _acting_as_team_admin_for_key_update(
data=data,
existing_key_row=existing_key_row,
user_api_key_dict=user_api_key_dict,
checked_prisma_client=checked_prisma_client,
user_api_key_cache=user_api_key_cache,
is_proxy_admin=_is_proxy_admin,
)
common_key_access_checks(
user_api_key_dict=user_api_key_dict,
data=data,
user_id=existing_key_row.user_id,
user_id=user_api_key_dict.user_id if acting_as_team_admin else existing_key_row.user_id,
llm_router=llm_router,
premium_user=premium_user,
)

View file

@ -1,5 +1,6 @@
"""Proxy-wide allow-list of what a team admin may do on the teams they administer: team-settings fields on
/team/update, plus the ``projects`` permission for /project/new and /project/update."""
/team/update, the ``projects`` permission for /project/new and /project/update, and the
``member_key_budgets`` permission for budget fields on other members' keys via /key/update."""
from collections.abc import Mapping
from dataclasses import dataclass
@ -12,9 +13,11 @@ from typing_extensions import assert_never
from litellm._logging import verbose_proxy_logger
from litellm.models.team import LiteLLM_TeamTable
from litellm.models.verification_token import LiteLLM_VerificationToken
from litellm.proxy._types import (
LiteLLM_ManagementEndpoint_MetadataFields,
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
UpdateKeyRequest,
UpdateTeamRequest,
)
@ -23,12 +26,20 @@ TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING: Final = "team_admin_editable_team_field
# TODO(LIT-5722): add the remaining team settings one per PR, each with its value-diff tests and dashboard field
SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS: Final[frozenset[str]] = frozenset({"tpm_limit", "rpm_limit", "max_budget"})
TEAM_ADMIN_PROJECTS_PERMISSION: Final = "projects"
TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION: Final = "member_key_budgets"
SUPPORTED_TEAM_ADMIN_PERMISSIONS: Final[frozenset[str]] = SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS | {
TEAM_ADMIN_PROJECTS_PERMISSION
TEAM_ADMIN_PROJECTS_PERMISSION,
TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION,
}
# spend is deliberately excluded: the stored row lags the live cross-pod counter, so a value-diff gate
# would let a team admin overwrite real usage.
KEY_BUDGET_FIELDS: Final[frozenset[str]] = frozenset({"max_budget", "budget_duration", "soft_budget", "budget_limits"})
_KEY_REQUEST_IDENTITY: Final[frozenset[str]] = frozenset({"key", "token", "metadata"})
_FIELD_LIST: Final = TypeAdapter(list[str])
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
_WINDOW_LIST: Final = TypeAdapter(list[dict[str, object]])
_EMPTY: Final[Mapping[str, object]] = MappingProxyType({})
_METADATA_FOLDED_FIELDS: Final[frozenset[str]] = frozenset(
(*LiteLLM_ManagementEndpoint_MetadataFields, *LiteLLM_ManagementEndpoint_MetadataFields_Premium)
@ -89,6 +100,12 @@ def team_admin_may_manage_projects(general_settings: Mapping[str, object]) -> bo
)
def team_admin_may_edit_member_key_budgets(general_settings: Mapping[str, object]) -> bool:
return TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION in resolve_team_admin_editable_fields(
general_settings, frozenset({TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION})
)
def _as_object(value: object) -> Mapping[str, object]:
try:
return _JSON_OBJECT.validate_json(value) if isinstance(value, str) else _JSON_OBJECT.validate_python(value)
@ -101,7 +118,7 @@ def _stored_metadata(existing: Mapping[str, object]) -> Mapping[str, object]:
def _submitted_metadata(
data: UpdateTeamRequest, submitted: Mapping[str, object], existing: Mapping[str, object]
data: UpdateTeamRequest | UpdateKeyRequest, submitted: Mapping[str, object], existing: Mapping[str, object]
) -> Mapping[str, object]:
"""Metadata as it would be stored: the caller's dict (or the stored one) with top-level folded fields laid over."""
base: Final = (
@ -112,7 +129,7 @@ def _submitted_metadata(
def _metadata_changes(
data: UpdateTeamRequest, submitted: Mapping[str, object], existing: Mapping[str, object]
data: UpdateTeamRequest | UpdateKeyRequest, submitted: Mapping[str, object], existing: Mapping[str, object]
) -> frozenset[str]:
merged: Final = _submitted_metadata(data, submitted, existing)
stored: Final = _stored_metadata(existing)
@ -200,3 +217,109 @@ def team_admin_request_or_raise(verdict: TeamAdminEditVerdict) -> UpdateTeamRequ
)
case _:
assert_never(verdict)
def _budget_windows(value: object) -> frozenset[tuple[object, object]] | None:
"""(budget_duration, max_budget) pairs for a stored or submitted budget_limits value.
Stored windows carry server-added keys like ``reset_at``; only the caller-owned pair matters.
``None`` means the value is not a list of windows and needs a plain comparison.
"""
if value is None:
return frozenset()
if not isinstance(value, list):
return None
try:
windows_input: Final = _WINDOW_LIST.validate_python(value)
except ValidationError:
return None
windows: Final = frozenset((window.get("budget_duration"), window.get("max_budget")) for window in windows_input)
if len(windows) != len(windows_input):
return None
return windows
def _key_column_changed(field: str, submitted: Mapping[str, object], existing: Mapping[str, object]) -> bool:
if field == "budget_limits":
sent: Final = _budget_windows(submitted.get(field))
stored: Final = _budget_windows(existing.get(field))
if sent is not None and stored is not None:
return sent != stored
if field in LiteLLM_VerificationToken.model_fields:
return submitted.get(field) != existing.get(field)
return True
def changed_key_fields(data: UpdateKeyRequest, existing_row: LiteLLM_VerificationToken) -> frozenset[str]:
"""Logical field names whose stored value the key-update request would change.
Same JSON-value comparison as :func:`changed_team_fields`: columns compare against the stored row,
fields the key endpoint folds into ``metadata`` compare against ``existing_row.metadata``, other
``metadata`` keys are attributed to ``metadata``, and fields with no stored counterpart count as
changed whenever they are sent. ``budget_limits`` compares (budget_duration, max_budget) pairs so
order and server-computed ``reset_at`` values do not read as edits.
"""
submitted: Final = _JSON_OBJECT.validate_json(data.model_dump_json(exclude_unset=True))
existing: Final = _JSON_OBJECT.validate_json(existing_row.model_dump_json())
column_fields: Final = frozenset(data.model_fields_set) - _KEY_REQUEST_IDENTITY - _METADATA_FOLDED_FIELDS
column_changes: Final = frozenset(
field for field in column_fields if _key_column_changed(field, submitted, existing)
)
return column_changes | _metadata_changes(data, submitted, existing)
@dataclass(frozen=True, slots=True)
class TeamAdminKeyEditAllowed:
changed: frozenset[str]
kind: Literal["allowed"] = "allowed"
@dataclass(frozen=True, slots=True)
class TeamAdminMemberKeyEditingDisabled:
kind: Literal["disabled"] = "disabled"
TeamAdminKeyEditVerdict: TypeAlias = (
TeamAdminKeyEditAllowed | TeamAdminMemberKeyEditingDisabled | TeamAdminFieldNotPermitted
)
def team_admin_key_edit_verdict(
data: UpdateKeyRequest,
existing: LiteLLM_VerificationToken,
enabled: bool,
) -> TeamAdminKeyEditVerdict:
if not enabled:
return TeamAdminMemberKeyEditingDisabled()
changed: Final = changed_key_fields(data, existing)
blocked: Final = sorted(
(changed | (frozenset({"spend"}) if "spend" in data.model_fields_set else frozenset())) - KEY_BUDGET_FIELDS
)
if blocked:
return TeamAdminFieldNotPermitted(field=blocked[0])
return TeamAdminKeyEditAllowed(changed=changed)
def team_admin_key_request_or_raise(verdict: TeamAdminKeyEditVerdict) -> None:
match verdict:
case TeamAdminKeyEditAllowed():
return
case TeamAdminMemberKeyEditingDisabled():
raise HTTPException(
status_code=403,
detail=(
"Team admins on this proxy cannot update budgets on other members' keys. "
f"Ask a proxy admin to enable '{TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION}' "
f"under {_SETTINGS_LOCATION}."
),
)
case TeamAdminFieldNotPermitted(field=field):
raise HTTPException(
status_code=403,
detail=(
"Team admins on this proxy may only update budget fields on other members' keys, "
f"not '{field}'. Ask a proxy admin to add it under {_SETTINGS_LOCATION}."
),
)
case _:
assert_never(verdict)

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

@ -607,7 +607,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
# body that mirrors them cannot clobber the authenticated key, the real
# parent span, or the proxy's own session-id decision.
_metadata.pop(SESSION_ID_OMITTED_METADATA_KEY, None)
_metadata["user_api_key"] = user_api_key_dict.api_key
_metadata["user_api_key"] = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
_metadata["litellm_parent_otel_span"] = user_api_key_dict.parent_otel_span
_metadata["user_api_key_budget_reservation"] = user_api_key_dict.budget_reservation
_metadata[MODEL_ACCESS_GROUP_METADATA_KEY] = user_api_key_dict.matched_model_access_groups

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,
)
@ -552,7 +554,7 @@ from litellm.proxy.list_api.common import (
problem_response,
request_validation_problem,
)
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request
from litellm.proxy.logging_endpoints.callback_logs_endpoints import (
rust_control_plane_router,
)
@ -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(
@ -10465,6 +10476,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,
@ -10839,6 +10857,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,
@ -11122,6 +11198,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"]
@ -11272,7 +11382,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,
@ -11328,7 +11440,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(
@ -11422,13 +11537,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(
@ -13733,7 +13859,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(
@ -15748,7 +15874,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"},
@ -15837,10 +15963,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)
@ -15889,7 +16022,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]] = []
@ -16122,23 +16255,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(
@ -16365,8 +16509,9 @@ async def async_queue_request(
# Covers both missing and JSON-string metadata (multipart /
# extra_body); see above for the same guard upstream.
data["metadata"] = {}
data["metadata"]["user_api_key"] = user_api_key_dict.api_key
data["metadata"]["user_api_key_hash"] = user_api_key_dict.api_key
logged_api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
data["metadata"]["user_api_key"] = logged_api_key
data["metadata"]["user_api_key_hash"] = logged_api_key
data["metadata"]["user_api_key_metadata"] = strip_callback_config(user_api_key_dict.metadata)
_headers: Final = _safe_get_request_headers(request).copy()
_headers.pop("authorization", None) # do not store the original `sk-..` api key in the db

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

@ -1,6 +1,7 @@
import asyncio
from collections.abc import Awaitable, Callable, Mapping, Sequence
from collections.abc import Set as AbstractSet
from dataclasses import dataclass
from datetime import datetime, timedelta
from types import MappingProxyType
from typing import Final, TypeVar
@ -11,6 +12,7 @@ from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
CLI_SESSION_KEY_PREFIX,
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
SPEND_LOG_KEY_METADATA_CACHE_TTL,
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL,
@ -62,6 +64,7 @@ _SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND
_SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS)
_HASHED_JWT_PREFIX: Final = "hashed-jwt-"
_CLI_SESSION_KEY_PREFIX: Final = f"{CLI_SESSION_KEY_PREFIX}-"
class KeyMetadataDict(TypedDict, total=False):
@ -109,7 +112,6 @@ _SPEND_LOG_METADATA_CACHE: Final = InMemoryCache(
)
_SPEND_LOG_QUERY_LOCK: Final = asyncio.Lock()
_EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({})
_EMPTY_EMAILS: Final[Mapping[str, str]] = MappingProxyType({})
async def _db_or_empty(
@ -149,54 +151,107 @@ async def _reverse_hash_key_metadata(
)
async def _emails_for_user_ids(
@dataclass(frozen=True, slots=True)
class _UserDetails:
email: str | None
only_team: str | None
_EMPTY_USER_DETAILS: Final[Mapping[str, _UserDetails]] = MappingProxyType({})
async def _details_for_user_ids(
prisma_client: PrismaClient,
user_ids: AbstractSet[str],
) -> Mapping[str, str]:
) -> Mapping[str, _UserDetails]:
if not user_ids:
return _EMPTY_EMAILS
return _EMPTY_USER_DETAILS
users: Final = await _db_or_empty(
lambda: UserRepository(prisma_client).table.find_many(
where={"user_id": {"in": list(user_ids)}}, # mutable-ok: Prisma find_many where= is a dict
),
"Failed user_email recovery for %d user ids: %s",
"Failed user detail recovery for %d user ids: %s",
len(user_ids),
)
if users is None:
return _EMPTY_EMAILS
return _EMPTY_USER_DETAILS
return MappingProxyType(
{
user.user_id: user.user_email
user.user_id: _UserDetails(
email=getattr(user, "user_email", None) or None,
only_team=_only_team(getattr(user, "teams", None)),
)
for user in users
if getattr(user, "user_id", None) and getattr(user, "user_email", None)
if getattr(user, "user_id", None)
}
)
def _meta_with_email(meta: KeyMetadataDict, emails: Mapping[str, str]) -> KeyMetadataDict:
if meta.get("user_email"):
return meta
def _only_team(teams: object) -> str | None:
if not isinstance(teams, list) or len(teams) != 1:
return None
team: Final = teams[0]
return team if isinstance(team, str) and team else None
def _is_cli_session_key(api_key: str) -> bool:
return api_key.startswith(_CLI_SESSION_KEY_PREFIX) and len(api_key) > len(_CLI_SESSION_KEY_PREFIX)
def _meta_with_user_details(
api_key: str, meta: KeyMetadataDict, details: Mapping[str, _UserDetails]
) -> KeyMetadataDict:
user_id: Final = meta.get("user_id")
if not isinstance(user_id, str) or user_id not in emails:
if not isinstance(user_id, str) or user_id not in details:
return meta
updated: Final[KeyMetadataDict] = {**meta, "user_email": emails[user_id]}
user: Final = details[user_id]
email: Final = meta.get("user_email") or user.email
team_id: Final = meta.get("team_id") or (user.only_team if _is_cli_session_key(api_key) else None)
updated: Final[KeyMetadataDict] = {
**meta,
**({"user_email": email} if email else {}),
**({"team_id": team_id} if team_id else {}),
}
return updated
async def attach_user_emails(
async def attach_user_details(
prisma_client: PrismaClient,
recovered: Mapping[str, KeyMetadataDict],
) -> Mapping[str, KeyMetadataDict]:
needing_email: Final = frozenset(
needing_details: Final = frozenset(
user_id
for meta in recovered.values()
for api_key, meta in recovered.items()
for user_id in (meta.get("user_id"),)
if isinstance(user_id, str) and user_id and not meta.get("user_email")
if isinstance(user_id, str)
and user_id
and (not meta.get("user_email") or (_is_cli_session_key(api_key) and not meta.get("team_id")))
)
emails: Final = await _emails_for_user_ids(prisma_client, needing_email)
if not emails:
details: Final = await _details_for_user_ids(prisma_client, needing_details)
if not details:
return recovered
return MappingProxyType({api_key: _meta_with_email(meta, emails) for api_key, meta in recovered.items()})
return MappingProxyType(
{api_key: _meta_with_user_details(api_key, meta, details) for api_key, meta in recovered.items()}
)
async def recover_cli_session_key_metadata(
prisma_client: PrismaClient,
missing_keys: AbstractSet[str],
) -> Mapping[str, KeyMetadataDict]:
candidates: Final = MappingProxyType(
{key: key.removeprefix(_CLI_SESSION_KEY_PREFIX) for key in missing_keys if _is_cli_session_key(key)}
)
if not candidates:
return _EMPTY_KEY_METADATA
known_users: Final = await _details_for_user_ids(prisma_client, frozenset(candidates.values()))
return MappingProxyType(
{
key: KeyMetadataDict(key_alias=key, user_id=user_id)
for key, user_id in candidates.items()
if user_id in known_users
}
)
async def recover_double_hashed_key_metadata(
@ -384,9 +439,15 @@ async def fill_missing_api_key_aliases(
if not missing_keys:
return tuple(rows)
recovered: Final = await attach_user_emails(
from_session_keys: Final = await recover_cli_session_key_metadata(prisma_client, missing_keys)
recovered: Final = await attach_user_details(
prisma_client,
await recover_double_hashed_key_metadata(prisma_client, missing_keys),
MappingProxyType(
{
**from_session_keys,
**await recover_double_hashed_key_metadata(prisma_client, missing_keys - frozenset(from_session_keys)),
}
),
)
if not recovered:
return tuple(rows)

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,18 @@ 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.litellm_pre_call_utils import LiteLLMProxyRequestSetup
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 +59,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 +1191,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"],
@ -1840,7 +1926,7 @@ async def get_key_spend_report(
scoped_api_key = _resolve_spend_report_scope(
user_api_key_dict=user_api_key_dict,
requested=requested,
caller_value=user_api_key_dict.api_key,
caller_value=LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
scope_name="api_key",
)
db_response: Sequence[Mapping[str, object]] | None = await _query_raw_or_none(

View file

@ -14,6 +14,7 @@ from pydantic import BaseModel, JsonValue
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
CLI_SESSION_KEY_PREFIX,
EMPTY_MAPPING,
LITELLM_PROXY_MASTER_KEY_ALIAS,
LITELLM_TRUNCATED_PAYLOAD_FIELD,
@ -102,19 +103,28 @@ _NON_SECRET_KEY_ALIASES: Final = frozenset(
)
def _is_non_secret_key_value(value: str) -> bool:
def _is_cli_session_alias(value: str, key_alias: object) -> bool:
return value.startswith(f"{CLI_SESSION_KEY_PREFIX}-") and value == key_alias
def _is_non_secret_key_value(value: str, *, key_alias: object = None) -> bool:
return (
value in _NON_SECRET_KEY_ALIASES or is_valid_sha256_hash(value) or _HASHED_JWT_RE.fullmatch(value) is not None
value in _NON_SECRET_KEY_ALIASES
or is_valid_sha256_hash(value)
or _HASHED_JWT_RE.fullmatch(value) is not None
or _is_cli_session_alias(value, key_alias)
)
def _redact_logged_api_key(value: str | None, *, already_redacted: bool = False) -> str | None:
def _redact_logged_api_key(
value: str | None, *, already_redacted: bool = False, key_alias: object = None
) -> str | None:
if not isinstance(value, str) or not value:
return None
stripped: Final = re.sub(r"(?i)^bearer ", "", value)
if not stripped:
return None
if already_redacted and _is_non_secret_key_value(stripped):
if already_redacted and _is_non_secret_key_value(stripped, key_alias=key_alias):
return stripped
return hash_token(stripped)
@ -230,10 +240,15 @@ def _get_spend_logs_metadata(
)
_raw_key: Final = clean_metadata.get("user_api_key")
_trusted_hash: Final = metadata.get("user_api_key_hash")
_key_alias: Final = metadata.get("user_api_key_alias")
_already_redacted: Final = (
isinstance(_trusted_hash, str) and _is_non_secret_key_value(_trusted_hash) and _trusted_hash == _raw_key
isinstance(_trusted_hash, str)
and _is_non_secret_key_value(_trusted_hash, key_alias=_key_alias)
and _trusted_hash == _raw_key
)
clean_metadata["user_api_key"] = _redact_logged_api_key(
_raw_key, already_redacted=_already_redacted, key_alias=_key_alias
)
clean_metadata["user_api_key"] = _redact_logged_api_key(_raw_key, already_redacted=_already_redacted)
clean_metadata["applied_guardrails"] = applied_guardrails
clean_metadata["batch_models"] = batch_models
clean_metadata["batch_successful_requests"] = batch_successful_requests
@ -537,10 +552,13 @@ def get_logging_payload(
standard_logging_completion_tokens = standard_logging_payload.get("completion_tokens", 0)
standard_logging_total_tokens = standard_logging_payload.get("total_tokens", 0)
_trusted_hash = metadata.get("user_api_key_hash")
_key_alias = metadata.get("user_api_key_alias")
_key_already_redacted = (
isinstance(_trusted_hash, str) and _is_non_secret_key_value(_trusted_hash) and _trusted_hash == api_key
isinstance(_trusted_hash, str)
and _is_non_secret_key_value(_trusted_hash, key_alias=_key_alias)
and _trusted_hash == api_key
)
api_key = _redact_logged_api_key(api_key, already_redacted=_key_already_redacted) or ""
api_key = _redact_logged_api_key(api_key, already_redacted=_key_already_redacted, key_alias=_key_alias) or ""
if (
standard_logging_payload is not None
@ -548,7 +566,9 @@ def get_logging_payload(
api_key = (
api_key
or _redact_logged_api_key(
standard_logging_payload["metadata"].get("user_api_key_hash"), already_redacted=True
standard_logging_payload["metadata"].get("user_api_key_hash"),
already_redacted=True,
key_alias=standard_logging_payload["metadata"].get("user_api_key_alias"),
)
or ""
)

View file

@ -328,6 +328,7 @@ class UISettings(BaseModel):
description=(
"Team settings fields a team admin may change on the teams they administer. "
"Include 'projects' to let team admins create and update projects for those teams. "
"Include 'member_key_budgets' to let team admins update budget fields on keys owned by other members of those teams. "
"Empty means team admins cannot edit team settings or manage projects at all. "
"Proxy admins and org admins are not affected."
),

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