mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
merge: resolve main into litellm_mcp_continuous_tool_defaults
Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
commit
878c646e99
588 changed files with 13839 additions and 68379 deletions
|
|
@ -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
|
||||
|
|
|
|||
140
.circleci/scripts/unit_selection.sh
Executable file
140
.circleci/scripts/unit_selection.sh
Executable 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
|
||||
|
|
@ -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 >>
|
||||
|
|
|
|||
13
.github/scripts/assert_ci_coverage.py
vendored
13
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -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())
|
||||
|
||||
|
|
|
|||
33
.github/workflows/_test-unit-base.yml
vendored
33
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
12
.github/workflows/test-linting.yml
vendored
12
.github/workflows/test-linting.yml
vendored
|
|
@ -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: |
|
||||
|
|
|
|||
97
.github/workflows/test-unit-proxy-db.yml
vendored
97
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
31
.github/workflows/test-unit.yml
vendored
31
.github/workflows/test-unit.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
12
Makefile
12
Makefile
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
43
db_scripts/backfill_key_total_spend.sql
Normal file
43
db_scripts/backfill_key_total_spend.sql
Normal 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;
|
||||
89
db_scripts/backfill_key_total_spend_from_spend_logs.sql
Normal file
89
db_scripts/backfill_key_total_spend_from_spend_logs.sql
Normal 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;
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 {}),
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "kill_switch" JSONB;
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)),
|
||||
|
|
|
|||
|
|
@ -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)?))
|
||||
}
|
||||
|
|
|
|||
190
litellm-rust/crates/python-bridge/src/secrets/python.rs
Normal file
190
litellm-rust/crates/python-bridge/src/secrets/python.rs
Normal 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()]);
|
||||
}
|
||||
}
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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)
|
||||
|
|
@ -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()
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
"""Azure Document Intelligence OCR module."""
|
||||
|
||||
from .transformation import AzureDocumentIntelligenceOCRConfig
|
||||
|
||||
__all__ = ["AzureDocumentIntelligenceOCRConfig"]
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
54
litellm/llms/base_llm/files/batch_records.py
Normal file
54
litellm/llms/base_llm/files/batch_records.py
Normal 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"),
|
||||
)
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,3 +0,0 @@
|
|||
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
|
||||
|
||||
__all__ = ("CohereParseConfig",)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
133
litellm/llms/openai/organization_costs.py
Normal file
133
litellm/llms/openai/organization_costs.py
Normal 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
|
||||
}
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -1 +0,0 @@
|
|||
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -1,5 +0,0 @@
|
|||
"""Vertex AI OCR module."""
|
||||
|
||||
from .transformation import VertexAIOCRConfig
|
||||
|
||||
__all__ = ["VertexAIOCRConfig"]
|
||||
|
|
@ -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()
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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```",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
239
litellm/proxy/agent_endpoints/kill_switch.py
Normal file
239
litellm/proxy/agent_endpoints/kill_switch.py
Normal 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]
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
300
litellm/proxy/spend_tracking/spend_capture_rate.py
Normal file
300
litellm/proxy/spend_tracking/spend_capture_rate.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 ""
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue