Merge remote-tracking branch 'origin/main' into litellm_replica_db_opt_in

This commit is contained in:
yuneng 2026-09-24 23:11:13 +00:00
commit 2ce8d66ecd
379 changed files with 23212 additions and 8074 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -31,10 +31,14 @@ concurrency:
# number, so a partially-specified entry would fail the call rather than fall
# back to the default.
#
# tests/proxy_unit_tests keeps its own caller (test-unit-proxy-db.yml): it is
# already a matrix and carries a shard-coverage guard that reads that file by
# name. Folding it in here is a follow-up, together with generalising that guard
# into assert_ci_coverage.py.
# tests/unit/proxy keeps its own caller (test-unit-proxy-db.yml): it is already
# a matrix and carries a shard-coverage guard that reads that file by name.
# Folding it in here is a follow-up, together with generalising that guard into
# assert_ci_coverage.py.
#
# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the
# shard under the same Codecov flag. CircleCI does not build pull requests from
# forks, so the shard still runs those files there and skips them elsewhere.
jobs:
unit:
name: ${{ matrix.shard }}
@ -49,6 +53,7 @@ jobs:
- shard: mcp-integration
artifact-name: mcp-integration
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
fork-flag: mcp-integration
workers: 2
reruns: 0
timeout-minutes: 20
@ -65,10 +70,10 @@ jobs:
- shard: enterprise-routing
artifact-name: enterprise-routing
test-path: >-
tests/test_litellm/enterprise
tests/test_litellm/google_genai
tests/test_litellm/router_utils
tests/test_litellm/router_strategy
fork-flag: enterprise-routing
workers: 2
reruns: 2
timeout-minutes: 20
@ -108,6 +113,7 @@ jobs:
tests/test_litellm/completion_extras
tests/test_litellm/containers
tests/test_litellm/endpoints
tests/test_litellm/files
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/messages
@ -199,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
@ -207,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
@ -219,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
@ -227,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
@ -246,6 +251,7 @@ jobs:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
fork-flag: ${{ matrix.fork-flag || '' }}
workers: ${{ matrix.workers }}
reruns: ${{ matrix.reruns }}
timeout-minutes: ${{ matrix.timeout-minutes }}

View file

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

10
litellm-rust/AGENTS.md Normal file
View file

@ -0,0 +1,10 @@
# Rust workspace rules
## Test placement
- Never create a `tests.rs` (or `test.rs`) file under `src/`, and never `#[path = "tests.rs"] mod tests;`
- A test that reaches private items lives inline, in a `#[cfg(test)] mod tests { ... }` at the bottom of the file that owns those items
- A test that only uses the crate's public API lives in `crates/<crate>/tests/<subject>.rs`, next to `src/`
- Split a mixed test file along that line instead of widening visibility to move it
- A test for another crate's item belongs in that crate, not in a downstream one
- Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests/<name>/mod.rs` or `tests/<subject>/support.rs` so it is not picked up as a test crate of its own

View file

@ -3414,6 +3414,7 @@ dependencies = [
name = "litellm-types"
version = "0.1.0"
dependencies = [
"rstest",
"serde",
"serde_json",
]

View file

@ -4,7 +4,6 @@ version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
autotests = false
[dependencies]
litellm-host.workspace = true

File diff suppressed because it is too large Load diff

View file

@ -63,5 +63,151 @@ impl PendingLogging {
}
#[cfg(test)]
#[path = "../tests/deferred.rs"]
mod tests;
mod tests {
use std::ffi::CStr;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rstest::rstest;
use super::{PendingLogging, PendingSuccess};
use crate::PythonLogger;
use crate::test_support::{local, namespace, run};
/// A deferred success for the namespace's `logger` and `response`, bound as `pending`.
fn defer<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
let locals = namespace(py, c"response = object()");
run(py, &locals, script);
let pending = Py::new(
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: PythonLogger::new(local(&locals, "logger").unbind()),
response: Some(local(&locals, "response").unbind()),
start: py.None(),
end: Some(py.None()),
}),
},
)
.unwrap();
locals.set_item("pending", pending).unwrap();
locals
}
#[test]
fn release_enqueues_the_success_once_in_the_releasing_context() {
Python::initialize();
Python::attach(|py| {
let locals = defer(
py,
c"
from contextvars import ContextVar
marker = ContextVar('marker', default='unset')
observed = []
def on_enqueue(coroutine):
observed.append(marker.get())
pending.release(True)
logger.on_enqueue = on_enqueue
",
);
run(
py,
&locals,
c"
marker.set('release')
pending.release(True)
pending.release(True)
assert observed == ['release'], observed
assert logger.names() == ['async_success_handler', 'enqueued'], logger.calls
assert logger.calls[0][1] is response
",
);
});
}
#[test]
fn a_blocked_release_drops_the_success_for_good() {
Python::initialize();
Python::attach(|py| {
let locals = defer(py, c"");
run(
py,
&locals,
c"
pending.release(False)
pending.release(True)
assert logger.calls == [], logger.calls
",
);
});
}
#[rstest]
#[case::ordinary_error(c"RuntimeError('queue full')", false)]
#[case::cancellation(c"asyncio.CancelledError()", true)]
fn a_failed_enqueue_closes_the_coroutine_and_is_never_replayed(
#[case] failure: &CStr,
#[case] propagates: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = defer(
py,
c"
import asyncio
def on_enqueue(coroutine):
raise failure
logger.on_enqueue = on_enqueue
",
);
locals
.set_item("failure", py.eval(failure, None, Some(&locals)).unwrap())
.unwrap();
let released = local(&locals, "pending").call_method1("release", (true,));
match released {
Ok(_) => assert!(!propagates),
Err(error) => {
assert!(propagates);
assert!(error.value(py).is(local(&locals, "failure")));
}
}
locals.set_item("propagates", propagates).unwrap();
run(
py,
&locals,
c"
pending.release(True)
assert logger.names() == ['async_success_handler', 'enqueued', 'closed'], logger.calls
assert unraisable_from(logger) == ([] if propagates else [failure])
",
);
});
}
#[test]
fn an_unreleased_success_does_not_keep_its_logger_alive() {
Python::initialize();
Python::attach(|py| {
let locals = defer(py, c"");
run(
py,
&locals,
c"
import gc
import weakref
logger.pending = pending
reference = weakref.ref(logger)
del logger, pending
gc.collect()
assert reference() is None
",
);
});
}
}

View file

@ -16,13 +16,218 @@ mod deferred;
mod logger;
mod preparation;
mod python;
#[cfg(test)]
#[path = "../tests/support.rs"]
mod test_support;
pub(crate) use adapter::LegacyLogging;
pub use adapter::{LegacySurface, PassThroughStream};
pub use call::{PublicCall, run_legacy_call};
pub(crate) use callbacks::{LegacyCallbacks, is_internal_call};
pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup};
pub(crate) use preparation::prepare;
#[cfg(test)]
mod test_support {
use std::ffi::CStr;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use crate::{LegacyLogging, LegacySurface, PublicCall};
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
/// `tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
/// signatures, and [`namespace`] binds every fake call against it.
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json");
/// Stand-ins for `callbacks_legacy_python`, the only Python module the crate calls. Tests
/// share one interpreter and run concurrently, so each fake is installed idempotently and
/// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`).
/// Every fake is bound against the contract first, so a call the real module would reject
/// fails here too.
const STUBS: &CStr = c"
import contextvars
import inspect
import json
import sys
import traceback
import types
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.callbacks_legacy_python'):
sys.modules.setdefault(name, types.ModuleType(name))
legacy = sys.modules['litellm.rust_bridge.callbacks_legacy_python']
CONTRACT = json.loads(python_contract)
def contracted(name, fake):
signature = inspect.Signature(
[inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]]
)
def checked(*args, **kwargs):
signature.bind(*args, **kwargs)
return fake(*args, **kwargs)
return checked
if not hasattr(legacy, 'is_internal'):
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
FAKES = {
'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
kwargs=kwargs,
),
'check_limits': lambda arguments: arguments['logger'].check_limits(arguments),
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
custom_llm_provider=provider,
),
'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args),
'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call(
original_response, api_key, additional_args
),
'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)),
'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending),
'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls(
response, start, end
),
'failure_handler': lambda logger, error, start, end, asynchronous: (
logger.async_failure_handler if asynchronous else logger.failure_handler
)(error, ''.join(traceback.format_exception(error)), start, end),
'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)),
'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end),
'enqueue_logging': lambda coroutine: coroutine.enqueue(),
'restore_context': lambda logger: logger.record('restore', None),
'custom_pricing_fields': lambda: ('ocr_cost_per_page',),
'is_internal_call': lambda: legacy.is_internal.get(),
'credential_list': lambda: [],
'warn_unknown_credential': lambda name, loaded: None,
'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type),
'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook(
'success', response, call_type
),
'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type),
'stream_opened': lambda logger: logger.record('stream_opened', None),
'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record(
'stream_success', list(chunks)
),
'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error),
}
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
for name, fake in FAKES.items():
setattr(legacy, name, contracted(name, fake))
unraisable = sys.modules.setdefault(
'litellm_test_unraisable', types.ModuleType('litellm_test_unraisable')
)
if not hasattr(unraisable, 'events'):
unraisable.events = []
sys.unraisablehook = lambda event: unraisable.events.append((event.object, event.exc_value))
def unraisable_from(owner):
return [error for source, error in unraisable.events if source is owner]
class StubCoroutine:
def __init__(self, logger):
self.logger = logger
def enqueue(self):
self.logger.record('enqueued', None)
self.logger.on_enqueue(self)
def close(self):
self.logger.record('closed', None)
class StubLogger:
def __init__(self):
self.calls = []
self.hooks = {}
self.on_enqueue = lambda coroutine: None
def record(self, name, value):
self.calls.append((name, value))
def names(self):
return [name for name, _ in self.calls]
def hook(self, phase, value, call_type):
self.record(phase + '_hook', call_type)
return self.hooks.get(phase, lambda value: 'awaitable')(value)
def check_limits(self, arguments):
self.record('check_limits', arguments)
def failure_handler(self, error, trace, start, end):
self.record('failure_handler', error)
def async_failure_handler(self, error, trace, start, end):
self.record('async_failure_handler', error)
return 'awaitable'
def success_handler(self, response, start, end):
self.record('success_handler', response)
def async_success_handler(self, response, start, end):
self.record('async_success_handler', response)
return StubCoroutine(self)
def handle_sync_success_callbacks_for_async_calls(self, response, start, end):
self.record('sync_success_for_async_call', response)
logger = StubLogger()
";
/// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it.
pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
let locals = PyDict::new(py);
locals.set_item("python_contract", PYTHON_CONTRACT).unwrap();
py.run(STUBS, Some(&locals), Some(&locals)).unwrap();
py.run(script, Some(&locals), Some(&locals)).unwrap();
locals
}
pub(crate) fn run(py: Python<'_>, locals: &Bound<'_, PyDict>, code: &CStr) {
py.run(code, Some(locals), Some(locals)).unwrap();
}
pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyAny> {
locals.get_item(name).unwrap().unwrap()
}
/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`).
pub(crate) fn legacy_call(
py: Python<'_>,
locals: &Bound<'_, PyDict>,
asynchronous: bool,
) -> LegacyLogging {
let request = locals
.get_item("request")
.unwrap()
.unwrap_or_else(|| py.None().into_bound(py));
let kwargs = locals
.get_item("kwargs")
.unwrap()
.map(|kwargs| kwargs.cast_into::<PyDict>().unwrap())
.unwrap_or_else(|| PyDict::new(py));
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
LegacyLogging::new(
py,
LegacySurface {
call_type: "test",
input_description: "test input",
stream: None,
},
call,
asynchronous,
)
}
}

View file

@ -1,146 +0,0 @@
use std::ffi::CStr;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rstest::rstest;
use super::{PendingLogging, PendingSuccess};
use crate::PythonLogger;
use crate::test_support::{local, namespace, run};
/// A deferred success for the namespace's `logger` and `response`, bound as `pending`.
fn defer<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
let locals = namespace(py, c"response = object()");
run(py, &locals, script);
let pending = Py::new(
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: PythonLogger::new(local(&locals, "logger").unbind()),
response: Some(local(&locals, "response").unbind()),
start: py.None(),
end: Some(py.None()),
}),
},
)
.unwrap();
locals.set_item("pending", pending).unwrap();
locals
}
#[test]
fn release_enqueues_the_success_once_in_the_releasing_context() {
Python::initialize();
Python::attach(|py| {
let locals = defer(
py,
c"
from contextvars import ContextVar
marker = ContextVar('marker', default='unset')
observed = []
def on_enqueue(coroutine):
observed.append(marker.get())
pending.release(True)
logger.on_enqueue = on_enqueue
",
);
run(
py,
&locals,
c"
marker.set('release')
pending.release(True)
pending.release(True)
assert observed == ['release'], observed
assert logger.names() == ['async_success_handler', 'enqueued'], logger.calls
assert logger.calls[0][1] is response
",
);
});
}
#[test]
fn a_blocked_release_drops_the_success_for_good() {
Python::initialize();
Python::attach(|py| {
let locals = defer(py, c"");
run(
py,
&locals,
c"
pending.release(False)
pending.release(True)
assert logger.calls == [], logger.calls
",
);
});
}
#[rstest]
#[case::ordinary_error(c"RuntimeError('queue full')", false)]
#[case::cancellation(c"asyncio.CancelledError()", true)]
fn a_failed_enqueue_closes_the_coroutine_and_is_never_replayed(
#[case] failure: &CStr,
#[case] propagates: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = defer(
py,
c"
import asyncio
def on_enqueue(coroutine):
raise failure
logger.on_enqueue = on_enqueue
",
);
locals
.set_item("failure", py.eval(failure, None, Some(&locals)).unwrap())
.unwrap();
let released = local(&locals, "pending").call_method1("release", (true,));
match released {
Ok(_) => assert!(!propagates),
Err(error) => {
assert!(propagates);
assert!(error.value(py).is(local(&locals, "failure")));
}
}
locals.set_item("propagates", propagates).unwrap();
run(
py,
&locals,
c"
pending.release(True)
assert logger.names() == ['async_success_handler', 'enqueued', 'closed'], logger.calls
assert unraisable_from(logger) == ([] if propagates else [failure])
",
);
});
}
#[test]
fn an_unreleased_success_does_not_keep_its_logger_alive() {
Python::initialize();
Python::attach(|py| {
let locals = defer(py, c"");
run(
py,
&locals,
c"
import gc
import weakref
logger.pending = pending
reference = weakref.ref(logger)
del logger, pending
gc.collect()
assert reference() is None
",
);
});
}

View file

@ -1,282 +0,0 @@
use std::ffi::CStr;
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rstest::rstest;
use super::LegacyLogging;
use crate::test_support::{legacy_call, local, namespace, run};
const CALL: &CStr = c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
kwargs = {'logger': logger, 'document': document}
";
const TIMING: Timing = Timing {
start_time: 0.0,
end_time: 1.0,
};
fn begin<'py>(
py: Python<'py>,
locals: &Bound<'py, PyDict>,
asynchronous: bool,
) -> (LegacyLogging, LifecycleStep) {
let mut logging = legacy_call(py, locals, asynchronous);
let kwargs = local(locals, "kwargs")
.cast_into::<PyDict>()
.unwrap()
.unbind();
let step = logging.begin(py, kwargs, 0.0).unwrap();
(logging, step)
}
fn arguments<'py>(py: Python<'py>, step: LifecycleStep) -> Bound<'py, PyDict> {
let LifecycleStep::Arguments(arguments) = step else {
panic!("expected the prepared arguments");
};
arguments.into_bound(py)
}
fn awaits_deployment_hook(step: &LifecycleStep) -> bool {
matches!(step, LifecycleStep::Await(_))
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
fn deployment_pre_call_hook_runs_only_for_asynchronous_calls(#[case] asynchronous: bool) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, CALL);
let (_, step) = begin(py, &locals, asynchronous);
assert_eq!(awaits_deployment_hook(&step), asynchronous);
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert_eq!(names.contains(&"pre_hook".to_string()), asynchronous);
});
}
#[test]
fn kwargs_returned_by_the_pre_call_hook_are_what_the_call_prepares() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
replacement = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,ZWRpdGVk'}
kwargs = {'logger': logger, 'document': document}
replaced_kwargs = {'logger': logger, 'document': replacement, 'pages': [0]}
",
);
let (mut logging, step) = begin(py, &locals, true);
assert!(awaits_deployment_hook(&step));
let step = logging
.resume(py, Ok(local(&locals, "replaced_kwargs").unbind()))
.unwrap();
locals.set_item("prepared", arguments(py, step)).unwrap();
run(
py,
&locals,
c"
assert prepared['document'] is replacement
assert prepared['pages'] is replaced_kwargs['pages']
assert prepared['litellm_logging_obj'] is logger
assert 'litellm_logging_obj' not in replaced_kwargs
[checked] = [value for name, value in logger.calls if name == 'check_limits']
assert checked is prepared
",
);
});
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
fn a_keyword_the_bridge_never_reads_reaches_every_reader_as_the_callers_object(
#[case] asynchronous: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
opaque = object()
hooked = []
logger.hooks = {'pre': lambda kwargs: hooked.append(kwargs['vendor_extension']) or kwargs}
kwargs = {'logger': logger, 'vendor_extension': opaque}
",
);
let (mut logging, step) = begin(py, &locals, asynchronous);
let step = match step {
LifecycleStep::Await(hook_result) => logging.resume(py, Ok(hook_result)).unwrap(),
step => step,
};
locals.set_item("prepared", arguments(py, step)).unwrap();
locals.set_item("asynchronous", asynchronous).unwrap();
run(
py,
&locals,
c"
assert prepared['vendor_extension'] is opaque
[checked] = [value for name, value in logger.calls if name == 'check_limits']
assert checked['vendor_extension'] is opaque
assert hooked == ([opaque] if asynchronous else []), hooked
",
);
});
}
#[test]
fn response_returned_by_the_post_call_hook_is_finalized_and_returned() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
kwargs = {'logger': logger}
response = object()
replacement = object()
logger.hooks = {'pre': lambda kwargs: kwargs}
",
);
let (mut logging, _) = begin(py, &locals, true);
logging
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
let step = logging
.after_success(py, local(&locals, "response").unbind(), TIMING)
.unwrap();
assert!(awaits_deployment_hook(&step));
let step = logging
.resume(py, Ok(local(&locals, "replacement").unbind()))
.unwrap();
let LifecycleStep::Response(returned) = step else {
panic!("expected the finalized response");
};
assert!(returned.bind(py).is(local(&locals, "replacement")));
run(
py,
&locals,
c"
[finalized] = [value for name, value in logger.calls if name == 'finalize']
assert finalized is replacement
",
);
});
}
#[rstest]
#[case::pre_call(false)]
#[case::post_call(true)]
fn cancelling_a_deployment_hook_ends_the_call_with_that_cancellation(#[case] post_call: bool) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"kwargs = {'logger': logger}\nresponse = object()");
let (mut logging, _) = begin(py, &locals, true);
if post_call {
logging
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
logging
.after_success(py, local(&locals, "response").unbind(), TIMING)
.unwrap();
}
let cancellation = CancelledError::new_err("cancelled");
let cancelled = cancellation.value(py).clone();
let error = logging.resume(py, Err(cancellation)).err().unwrap();
assert!(error.value(py).is(&cancelled));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert!(!names.iter().any(|name| name.contains("handler")));
});
}
#[rstest]
#[case::hook_completed(false)]
#[case::hook_cancelled(true)]
fn failure_callbacks_run_after_the_failure_hook_however_it_ends(#[case] cancelled: bool) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"kwargs = {'logger': logger}\nfailure = ValueError('provider')",
);
let (mut logging, _) = begin(py, &locals, true);
logging
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
let failure = PyErr::from_value(local(&locals, "failure"));
let failed = LifecycleEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Call,
error: &failure,
};
let step = logging.emit(py, failed).unwrap();
assert!(awaits_deployment_hook(&step));
let hook_result = if cancelled {
Err(CancelledError::new_err("cancelled"))
} else {
Ok(py.None())
};
assert!(matches!(
logging.resume(py, hook_result).unwrap(),
LifecycleStep::Await(_)
));
run(
py,
&locals,
c"
assert logger.names()[-3:] == ['failure_hook', 'failure_handler', 'async_failure_handler'], logger.calls
assert all(value is failure for name, value in logger.calls if name.endswith('_handler'))
",
);
});
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
fn a_limit_rejected_before_the_call_surfaces_as_the_callers_error(#[case] asynchronous: bool) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
class BudgetExceeded(Exception):
pass
rejection = BudgetExceeded('over budget')
class LimitedLogger(StubLogger):
def check_limits(self, arguments):
raise rejection
logger = LimitedLogger()
logger.hooks = {'pre': lambda kwargs: kwargs}
kwargs = {'logger': logger}
",
);
let mut logging = legacy_call(py, &locals, asynchronous);
let kwargs = local(&locals, "kwargs")
.cast_into::<PyDict>()
.unwrap()
.unbind();
let result = logging.begin(py, kwargs, 0.0).and_then(|step| match step {
LifecycleStep::Await(_) => logging.resume(py, Ok(local(&locals, "kwargs").unbind())),
step => Ok(step),
});
let error = result.err().unwrap();
assert!(error.value(py).is(local(&locals, "rejection")));
});
}

View file

@ -1,523 +0,0 @@
use std::ffi::CStr;
use litellm_auth::SecretValue;
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py};
use proptest::prelude::*;
use pyo3::prelude::*;
use rstest::rstest;
use serde_json::{Map, Value, json};
use super::LegacyLogging;
use crate::PythonLogger;
use crate::test_support::{legacy_call, local, namespace, run};
/// The payload phases of `Logging` on top of `StubLogger`, with `pre_call` handing the
/// payload to the case's `on_pre_call`.
const PAYLOAD_LOGGER: &CStr = c"
class Request:
pass
class PayloadLogger(StubLogger):
def update_from_kwargs(self, **update):
self.update = update
def pre_call(self, input, api_key, additional_args):
self.record('pre_call', None)
self.pre = additional_args
self.pre_api_key = api_key
on_pre_call(additional_args)
def post_call(self, original_response, api_key, additional_args):
self.record('post_call', None)
self.post = (original_response, api_key, additional_args)
request = Request()
kwargs = {}
logger = PayloadLogger()
on_pre_call = lambda additional_args: None
check = lambda: None
";
const DOCUMENT: &str = "data:application/pdf;base64,YWJj";
const EDITED: &str = "data:application/pdf;base64,ZWRpdGVk";
fn document(source: &str) -> Value {
json!({"type": "document_url", "document_url": source})
}
fn before_send(script: &CStr, body: Value) -> WireRequest {
before_send_with_secrets(script, json!({}), body, &[])
}
/// Runs `before_send` over `body` for a route whose parameters are `optional_params`, with
/// the Python objects `script` binds, then delivers the provider's raw response the way the
/// driver does and runs the script's `check()`.
fn before_send_with_secrets(
script: &CStr,
optional_params: Value,
body: Value,
secret_fields: &[&str],
) -> WireRequest {
before_send_bound(&[], script, optional_params, body, secret_fields)
}
/// [`before_send_with_secrets`] with `bindings` placed in the namespace before `script` runs.
fn before_send_bound(
bindings: &[(&str, &Value)],
script: &CStr,
optional_params: Value,
body: Value,
secret_fields: &[&str],
) -> WireRequest {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, PAYLOAD_LOGGER);
for &(name, value) in bindings {
locals.set_item(name, to_py(py, value).unwrap()).unwrap();
}
run(py, &locals, script);
let mut logging = LegacyLogging {
logger: Some(PythonLogger::new(local(&locals, "logger").unbind())),
..legacy_call(py, &locals, false)
};
let context = RequestContext {
model: "model".into(),
custom_llm_provider: "provider".into(),
optional_params,
secret_fields: secret_fields.iter().map(|name| name.to_string()).collect(),
api_key: Some(SecretValue::new("route-key")),
};
let wire = WireRequest {
url: "https://provider.invalid/ocr".into(),
headers: vec![("x-route".into(), "route".into())],
body,
};
let step = logging.before_send(py, Box::new(wire), &context).unwrap();
let raw = MachineEvent::ResponseReceived {
raw: RawResponse {
body: "raw response".into(),
},
};
assert!(matches!(
logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(),
LifecycleStep::Done
));
run(py, &locals, c"check()");
let LifecycleStep::Wire(wire) = step else {
panic!("before_send did not hand back the wire request");
};
*wire
})
}
#[rstest]
#[case::caller_keyword(c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
pages = [0]
kwargs = {'document': document, 'pages': pages}
observed = []
on_pre_call = lambda args: observed.append(
(args['complete_input_dict']['document'] is document, args['complete_input_dict']['pages'] is pages)
)
def check():
assert observed == [(True, True)], observed
")]
#[case::request_attribute_behind_an_omitted_keyword(c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
pages = [0]
request.document = document
kwargs = {'pages': pages}
observed = []
on_pre_call = lambda args: observed.append(
(args['complete_input_dict']['document'] is document, args['complete_input_dict']['pages'] is pages)
)
def check():
assert observed == [(True, True)], observed
")]
fn passthrough_keys_reach_pre_call_as_the_callers_own_objects(#[case] script: &CStr) {
let body = json!({"model": "model", "document": document(DOCUMENT), "pages": [0]});
let wire = before_send(script, body.clone());
assert_eq!(wire.body, body);
}
#[test]
fn pre_call_edit_of_a_passthrough_object_reaches_the_caller_and_the_wire() {
let wire = before_send(
c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
kwargs = {'document': document}
def on_pre_call(args):
args['complete_input_dict']['document']['document_url'] = 'data:application/pdf;base64,ZWRpdGVk'
def check():
assert document['document_url'] == 'data:application/pdf;base64,ZWRpdGVk'
",
json!({"document": document(DOCUMENT)}),
);
assert_eq!(wire.body["document"], document(EDITED));
}
#[test]
fn a_body_key_the_route_rewrote_is_not_the_callers_object() {
let wire = before_send(
c"
document = {'type': 'document_url', 'document_url': 'https://example.invalid/scan.pdf'}
kwargs = {'document': document}
observed = []
def on_pre_call(args):
observed.append(args['complete_input_dict']['document'] is document)
args['complete_input_dict']['document']['document_name'] = 'edited.pdf'
def check():
assert observed == [False], observed
assert document == {'type': 'document_url', 'document_url': 'https://example.invalid/scan.pdf'}
",
json!({"document": document(DOCUMENT)}),
);
assert_eq!(
wire.body["document"],
json!({"type": "document_url", "document_url": DOCUMENT, "document_name": "edited.pdf"})
);
}
#[test]
fn a_caller_value_with_no_json_form_is_left_out_of_realiasing() {
let body = json!({"pages": [0]});
let wire = before_send(
c"
opaque = object()
kwargs = {'pages': opaque}
observed = []
on_pre_call = lambda args: observed.append(args['complete_input_dict']['pages'])
def check():
assert observed == [[0]], observed
",
body.clone(),
);
assert_eq!(wire.body, body);
}
#[rstest]
#[case::body(
c"
def on_pre_call(args):
args['complete_input_dict'] = {'replacement': True}
"
)]
#[case::headers(
c"
def on_pre_call(args):
args['headers'] = {'x-replacement': 'yes'}
"
)]
fn rebinding_the_payload_envelope_does_not_reach_the_wire(#[case] script: &CStr) {
let body = json!({"document": document(DOCUMENT)});
let wire = before_send(script, body.clone());
assert_eq!(wire.body, body);
assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]);
}
#[test]
fn pre_call_header_edit_reaches_the_wire() {
let wire = before_send(
c"
def on_pre_call(args):
args['headers']['x-callback'] = 'edited'
",
json!({}),
);
assert_eq!(
wire.headers,
[
("x-route".to_string(), "route".to_string()),
("x-callback".to_string(), "edited".to_string()),
]
);
}
#[test]
fn pre_call_receives_the_wire_request_and_the_logger_its_redacted_request() {
let body = json!({"model": "model", "document": document(DOCUMENT)});
before_send_with_secrets(
c"
logger_fn = lambda *args: None
kwargs = {
'litellm_call_id': 'call-1',
'client_secret': 'shh',
'proxy_server_request': {'body': {}},
'logger_fn': logger_fn,
'litellm_request_debug': True,
'ocr_cost_per_page': 0.05,
}
observed = []
on_pre_call = observed.append
def check():
[args] = observed
assert args['api_base'] == 'https://provider.invalid/ocr', args
assert args['complete_input_dict'] == {
'model': 'model',
'document': {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'},
}, args
update = logger.update
assert update['model'] == 'model' and update['custom_llm_provider'] == 'provider', update
assert update['litellm_params']['litellm_call_id'] == 'call-1', update
assert update['litellm_params']['api_base'] == 'https://provider.invalid/ocr', update
assert update['litellm_params']['logger_fn'] is logger_fn, update
assert update['litellm_params']['litellm_request_debug'] is True, update
assert update['litellm_params']['ocr_cost_per_page'] == 0.05, update
assert update['kwargs']['client_secret'] == '****', update
assert 'proxy_server_request' not in update['kwargs'], update
assert update['optional_params']['client_secret'] == '****', update
",
json!({"client_secret": "shh"}),
body,
&["client_secret"],
);
}
#[rstest]
#[case::added_key(
c"
def on_pre_call(args):
args['complete_input_dict']['include_image_base64'] = True
",
json!({"document": document(DOCUMENT), "include_image_base64": true})
)]
#[case::replaced_document(
c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
kwargs = {'document': document}
def on_pre_call(args):
args['complete_input_dict']['document'] = {
'type': 'document_url', 'document_url': 'data:application/pdf;base64,ZWRpdGVk'
}
def check():
assert document['document_url'] == 'data:application/pdf;base64,YWJj', document
",
json!({"document": document(EDITED)})
)]
#[case::retained_body_edited_after_rebinding(
c"
def on_pre_call(args):
retained = args['complete_input_dict']
args['complete_input_dict'] = {'rebound': True}
retained['include_image_base64'] = True
",
json!({"document": document(DOCUMENT), "include_image_base64": true})
)]
fn pre_call_body_edits_reach_the_wire(#[case] script: &CStr, #[case] expected: Value) {
let body = json!({"document": document(DOCUMENT)});
let wire = before_send(script, body);
assert_eq!(wire.body, expected);
}
#[test]
fn retained_headers_edited_after_rebinding_reach_the_wire() {
let wire = before_send(
c"
def on_pre_call(args):
retained = args['headers']
args['headers'] = {'x-rebound': 'rebound'}
retained['x-retained'] = 'sent'
",
json!({}),
);
assert_eq!(
wire.headers,
[
("x-route".to_string(), "route".to_string()),
("x-retained".to_string(), "sent".to_string()),
]
);
}
#[test]
fn post_call_receives_the_raw_response_the_route_key_and_the_body_and_headers_pre_call_saw() {
before_send(
c"
def check():
original_response, api_key, additional_args = logger.post
assert original_response == 'raw response', original_response
assert api_key == logger.pre_api_key == 'route-key', (api_key, logger.pre_api_key)
assert additional_args == {
'complete_input_dict': logger.pre['complete_input_dict'],
'headers': logger.pre['headers'],
}, additional_args
assert additional_args['complete_input_dict'] is logger.pre['complete_input_dict']
assert additional_args['headers'] is logger.pre['headers']
",
json!({"document": document(DOCUMENT)}),
);
}
#[test]
fn every_request_runs_the_full_pre_call_and_post_call() {
let wire = before_send(
c"
def on_pre_call(args):
args['complete_input_dict']['include_image_base64'] = True
def check():
assert logger.names() == ['pre_call', 'post_call'], logger.calls
",
json!({"document": document(DOCUMENT)}),
);
assert_eq!(
wire.body,
json!({"document": document(DOCUMENT), "include_image_base64": true})
);
}
/// What one pre-call callback does to the payload it is handed.
#[derive(Clone, Debug)]
enum Edit {
Nothing,
Set(String, Value),
Remove(String),
Rebind(Value),
RebindThenSetRetained(String, Value),
}
impl Edit {
fn script(&self) -> Value {
match self {
Self::Nothing => json!({"kind": "nothing"}),
Self::Set(key, value) => json!({"kind": "set", "key": key, "value": value}),
Self::Remove(key) => json!({"kind": "remove", "key": key}),
Self::Rebind(value) => json!({"kind": "rebind", "value": value}),
Self::RebindThenSetRetained(key, value) => {
json!({"kind": "rebind_then_set_retained", "key": key, "value": value})
}
}
}
/// The legacy contract: the provider is sent the body object `pre_call` received, as
/// the callback left it. Rebinding the envelope's key points the envelope elsewhere and
/// leaves that object alone.
fn sent(&self, body: &Map<String, Value>) -> Value {
let mut sent = body.clone();
match self {
Self::Nothing | Self::Rebind(_) => {}
Self::Set(key, value) | Self::RebindThenSetRetained(key, value) => {
sent.insert(key.clone(), value.clone());
}
Self::Remove(key) => {
sent.remove(key);
}
}
Value::Object(sent)
}
}
/// How the caller's keyword for a body key relates to what the route sends under it.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Caller {
PassedUnchanged,
RewrittenByTheRoute,
NotPassed,
}
const MODEL: &CStr = c"
aliased = {}
def on_pre_call(args):
body = args['complete_input_dict']
aliased.update({name: body[name] is kwargs[name] for name in unchanged})
kind = edit['kind']
if kind == 'set':
body[edit['key']] = edit['value']
elif kind == 'remove':
body.pop(edit['key'], None)
elif kind == 'rebind':
args['complete_input_dict'] = edit['value']
elif kind == 'rebind_then_set_retained':
args['complete_input_dict'] = {}
body[edit['key']] = edit['value']
def check():
assert aliased == {name: True for name in unchanged}, aliased
assert logger.names() == ['pre_call', 'post_call'], logger.calls
";
fn json_value() -> impl Strategy<Value = Value> {
let leaf = prop_oneof![
Just(Value::Null),
any::<bool>().prop_map(Value::from),
any::<i64>().prop_map(Value::from),
any::<f64>()
.prop_filter("JSON has no NaN or infinity", |number| number.is_finite())
.prop_map(Value::from),
".{0,8}".prop_map(Value::from),
];
leaf.prop_recursive(3, 24, 4, |inner| {
prop_oneof![
prop::collection::vec(inner.clone(), 0..4).prop_map(Value::from),
prop::collection::btree_map(key(), inner, 0..4)
.prop_map(|fields| Value::Object(fields.into_iter().collect())),
]
})
}
fn key() -> impl Strategy<Value = String> {
"[a-z]{1,6}"
}
fn caller() -> impl Strategy<Value = Caller> {
prop_oneof![
Just(Caller::PassedUnchanged),
Just(Caller::RewrittenByTheRoute),
Just(Caller::NotPassed),
]
}
fn edit() -> impl Strategy<Value = Edit> {
prop_oneof![
Just(Edit::Nothing),
(key(), json_value()).prop_map(|(key, value)| Edit::Set(key, value)),
key().prop_map(Edit::Remove),
json_value().prop_map(Edit::Rebind),
(key(), json_value()).prop_map(|(key, value)| Edit::RebindThenSetRetained(key, value)),
]
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(128))]
/// For any body, any caller keywords and any callback edit: every keyword the route
/// sends unchanged reaches `pre_call` as the caller's own object, and the provider is
/// sent exactly what the model says, so a callback that edits nothing changes nothing.
#[test]
fn the_wire_is_the_body_pre_call_received_as_the_callback_left_it(
fields in prop::collection::btree_map(key(), (json_value(), caller()), 0..5),
edit in edit(),
) {
let body: Map<String, Value> = fields
.iter()
.map(|(name, (value, _))| (name.clone(), value.clone()))
.collect();
let kwargs: Map<String, Value> = fields
.iter()
.filter_map(|(name, (value, caller))| match caller {
Caller::PassedUnchanged => Some((name.clone(), value.clone())),
Caller::RewrittenByTheRoute => Some((name.clone(), json!([value]))),
Caller::NotPassed => None,
})
.collect();
let unchanged: Value = fields
.iter()
.filter(|(_, (_, caller))| *caller == Caller::PassedUnchanged)
.map(|(name, _)| Value::from(name.clone()))
.collect();
let wire = before_send_bound(
&[
("kwargs", &Value::Object(kwargs)),
("unchanged", &unchanged),
("edit", &edit.script()),
],
MODEL,
json!({}),
Value::Object(body.clone()),
&[],
);
prop_assert_eq!(wire.body, edit.sent(&body));
prop_assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]);
}
}

View file

@ -1,205 +0,0 @@
use std::ffi::CStr;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use crate::{LegacyLogging, LegacySurface, PublicCall};
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
/// `tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
/// signatures, and [`namespace`] binds every fake call against it.
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json");
/// Stand-ins for `callbacks_legacy_python`, the only Python module the crate calls. Tests
/// share one interpreter and run concurrently, so each fake is installed idempotently and
/// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`).
/// Every fake is bound against the contract first, so a call the real module would reject
/// fails here too.
const STUBS: &CStr = c"
import contextvars
import inspect
import json
import sys
import traceback
import types
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.callbacks_legacy_python'):
sys.modules.setdefault(name, types.ModuleType(name))
legacy = sys.modules['litellm.rust_bridge.callbacks_legacy_python']
CONTRACT = json.loads(python_contract)
def contracted(name, fake):
signature = inspect.Signature(
[inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]]
)
def checked(*args, **kwargs):
signature.bind(*args, **kwargs)
return fake(*args, **kwargs)
return checked
if not hasattr(legacy, 'is_internal'):
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
FAKES = {
'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
kwargs=kwargs,
),
'check_limits': lambda arguments: arguments['logger'].check_limits(arguments),
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
custom_llm_provider=provider,
),
'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args),
'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call(
original_response, api_key, additional_args
),
'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)),
'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending),
'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls(
response, start, end
),
'failure_handler': lambda logger, error, start, end, asynchronous: (
logger.async_failure_handler if asynchronous else logger.failure_handler
)(error, ''.join(traceback.format_exception(error)), start, end),
'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)),
'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end),
'enqueue_logging': lambda coroutine: coroutine.enqueue(),
'restore_context': lambda logger: logger.record('restore', None),
'custom_pricing_fields': lambda: ('ocr_cost_per_page',),
'is_internal_call': lambda: legacy.is_internal.get(),
'credential_list': lambda: [],
'warn_unknown_credential': lambda name, loaded: None,
'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type),
'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook(
'success', response, call_type
),
'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type),
'stream_opened': lambda logger: logger.record('stream_opened', None),
'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record(
'stream_success', list(chunks)
),
'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error),
}
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
for name, fake in FAKES.items():
setattr(legacy, name, contracted(name, fake))
unraisable = sys.modules.setdefault(
'litellm_test_unraisable', types.ModuleType('litellm_test_unraisable')
)
if not hasattr(unraisable, 'events'):
unraisable.events = []
sys.unraisablehook = lambda event: unraisable.events.append((event.object, event.exc_value))
def unraisable_from(owner):
return [error for source, error in unraisable.events if source is owner]
class StubCoroutine:
def __init__(self, logger):
self.logger = logger
def enqueue(self):
self.logger.record('enqueued', None)
self.logger.on_enqueue(self)
def close(self):
self.logger.record('closed', None)
class StubLogger:
def __init__(self):
self.calls = []
self.hooks = {}
self.on_enqueue = lambda coroutine: None
def record(self, name, value):
self.calls.append((name, value))
def names(self):
return [name for name, _ in self.calls]
def hook(self, phase, value, call_type):
self.record(phase + '_hook', call_type)
return self.hooks.get(phase, lambda value: 'awaitable')(value)
def check_limits(self, arguments):
self.record('check_limits', arguments)
def failure_handler(self, error, trace, start, end):
self.record('failure_handler', error)
def async_failure_handler(self, error, trace, start, end):
self.record('async_failure_handler', error)
return 'awaitable'
def success_handler(self, response, start, end):
self.record('success_handler', response)
def async_success_handler(self, response, start, end):
self.record('async_success_handler', response)
return StubCoroutine(self)
def handle_sync_success_callbacks_for_async_calls(self, response, start, end):
self.record('sync_success_for_async_call', response)
logger = StubLogger()
";
/// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it.
pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
let locals = PyDict::new(py);
locals.set_item("python_contract", PYTHON_CONTRACT).unwrap();
py.run(STUBS, Some(&locals), Some(&locals)).unwrap();
py.run(script, Some(&locals), Some(&locals)).unwrap();
locals
}
pub(crate) fn run(py: Python<'_>, locals: &Bound<'_, PyDict>, code: &CStr) {
py.run(code, Some(locals), Some(locals)).unwrap();
}
pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyAny> {
locals.get_item(name).unwrap().unwrap()
}
/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`).
pub(crate) fn legacy_call(
py: Python<'_>,
locals: &Bound<'_, PyDict>,
asynchronous: bool,
) -> LegacyLogging {
let request = locals
.get_item("request")
.unwrap()
.unwrap_or_else(|| py.None().into_bound(py));
let kwargs = locals
.get_item("kwargs")
.unwrap()
.map(|kwargs| kwargs.cast_into::<PyDict>().unwrap())
.unwrap_or_else(|| PyDict::new(py));
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
LegacyLogging::new(
py,
LegacySurface {
call_type: "test",
input_description: "test input",
stream: None,
},
call,
asynchronous,
)
}

View file

@ -1,291 +0,0 @@
use std::ffi::CStr;
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
use pyo3::exceptions::PyRuntimeError;
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rstest::rstest;
use super::LegacyLogging;
use crate::PythonLogger;
use crate::test_support::{legacy_call, local, namespace, run};
const TIMING: Timing = Timing {
start_time: 0.0,
end_time: 1.0,
};
fn logged(py: Python<'_>, locals: &Bound<'_, PyDict>, asynchronous: bool) -> LegacyLogging {
LegacyLogging {
logger: Some(PythonLogger::new(local(locals, "logger").unbind())),
..legacy_call(py, locals, asynchronous)
}
}
fn succeed(
py: Python<'_>,
locals: &Bound<'_, PyDict>,
logging: &mut LegacyLogging,
) -> LifecycleStep {
let response = local(locals, "response").unbind();
logging
.emit(
py,
LifecycleEvent::Succeeded {
timing: TIMING,
response: &response,
},
)
.unwrap()
}
fn fail(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) -> LifecycleStep {
let failure = PyErr::from_value(local(locals, "failure"));
logging
.emit(
py,
LifecycleEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Host,
error: &failure,
},
)
.unwrap()
}
#[rstest]
#[case::sync_listened(false, c"", &["submit"])]
#[case::async_listened(
true,
c"",
&["async_success_handler", "enqueued", "sync_success_for_async_call"]
)]
#[case::async_deferred(true, c"logger._defer_async_logging = True", &["sync_success_for_async_call"])]
#[case::async_with_fallbacks(true, c"kwargs = {'fallbacks': ['other']}", &["sync_success_for_async_call"])]
fn success_reaches_the_logging_handlers(
#[case] asynchronous: bool,
#[case] script: &CStr,
#[case] expected: &[&str],
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"response = object()");
run(py, &locals, script);
let mut logging = logged(py, &locals, asynchronous);
assert!(matches!(
succeed(py, &locals, &mut logging),
LifecycleStep::Done
));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert_eq!(names, expected);
run(
py,
&locals,
c"
assert all(value is response for name, value in logger.calls if name.endswith('_handler'))
assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_async_logging', False)
",
);
});
}
#[rstest]
#[case::synchronous(false, &["failure_handler"])]
#[case::asynchronous(true, &[])]
fn internal_calls_skip_failure_callbacks_only_when_asynchronous(
#[case] asynchronous: bool,
#[case] expected: &[&str],
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"failure = ValueError('provider')");
let mut logging = LegacyLogging {
internal: true,
..logged(py, &locals, asynchronous)
};
assert!(matches!(
fail(py, &locals, &mut logging),
LifecycleStep::Done
));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert_eq!(names, expected);
});
}
#[test]
fn internal_async_calls_skip_the_async_success_fan_out() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"response = object()");
let mut logging = LegacyLogging {
internal: true,
..logged(py, &locals, true)
};
succeed(py, &locals, &mut logging);
run(
py,
&locals,
c"assert logger.names() == ['sync_success_for_async_call'], logger.calls",
);
});
}
#[test]
fn a_failing_success_callback_is_reported_without_replacing_the_response() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
response = object()
failure = ValueError('terminal diagnostic')
class FailingLogger(StubLogger):
def handle_sync_success_callbacks_for_async_calls(self, *args):
raise failure
logger = FailingLogger()
",
);
let mut logging = logged(py, &locals, true);
assert!(matches!(
succeed(py, &locals, &mut logging),
LifecycleStep::Done
));
assert!(
logging
.response
.as_ref()
.unwrap()
.bind(py)
.is(local(&locals, "response"))
);
run(py, &locals, c"assert unraisable_from(logger) == [failure]");
});
}
#[rstest]
#[case::sync_listened(false, c"", &["failure_handler"])]
#[case::async_listened(true, c"", &["failure_handler", "async_failure_handler"])]
fn failure_reaches_the_logging_handlers(
#[case] asynchronous: bool,
#[case] script: &CStr,
#[case] expected: &[&str],
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"failure = ValueError('provider')");
run(py, &locals, script);
let mut logging = logged(py, &locals, asynchronous);
let step = fail(py, &locals, &mut logging);
let awaits_async_handler = expected.contains(&"async_failure_handler");
assert_eq!(
matches!(step, LifecycleStep::Await(_)),
awaits_async_handler
);
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert_eq!(names, expected);
run(
py,
&locals,
c"assert all(value is failure for name, value in logger.calls if name.endswith('_handler'))",
);
});
}
#[test]
fn a_failing_sync_failure_callback_keeps_the_error_and_still_runs_the_async_family() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
failure = ValueError('selected')
class FailingLogger(StubLogger):
def failure_handler(self, error, trace, start, end):
self.record('failure_handler', error)
raise RuntimeError('handler failed')
logger = FailingLogger()
",
);
let mut logging = logged(py, &locals, true);
assert!(matches!(
fail(py, &locals, &mut logging),
LifecycleStep::Await(_)
));
assert!(
logging
.error
.as_ref()
.unwrap()
.bind(py)
.is(local(&locals, "failure"))
);
run(
py,
&locals,
c"assert logger.names() == ['failure_handler', 'async_failure_handler'], logger.calls",
);
});
}
#[rstest]
#[case::completed(None, true)]
#[case::handler_error(Some(false), true)]
#[case::cancelled(Some(true), false)]
fn the_async_failure_handler_ends_the_call_unless_it_was_cancelled(
#[case] error: Option<bool>,
#[case] done: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"failure = ValueError('provider')");
let mut logging = logged(py, &locals, true);
fail(py, &locals, &mut logging);
let result = match error {
None => Ok(py.None()),
Some(false) => Err(PyRuntimeError::new_err("handler failed")),
Some(true) => Err(CancelledError::new_err("cancelled")),
};
let expected = result.as_ref().err().map(|error| error.value(py).clone());
match logging.resume(py, result) {
Ok(step) => assert!(done && matches!(step, LifecycleStep::Done)),
Err(propagated) => {
assert!(!done);
assert!(propagated.value(py).is(expected.unwrap()));
}
}
});
}
#[test]
fn closing_restores_the_correlation_context_once() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"");
let mut logging = logged(py, &locals, true);
logging.close(py);
logging.close(py);
run(
py,
&locals,
c"assert logger.names() == ['restore'], logger.calls",
);
});
}

View file

@ -0,0 +1,274 @@
use serde_json::Value;
#[derive(Clone, Debug, PartialEq, Eq)]
enum Segment {
Field(String),
Every,
Index(usize),
}
fn parse_segments(path: &str) -> Option<Vec<Segment>> {
let mut segments = Vec::new();
let mut rest = path;
while !rest.is_empty() {
if let Some(after_open) = rest.strip_prefix('[') {
let (inside, after) = after_open.split_once(']')?;
segments.push(match inside {
"*" => Segment::Every,
index => Segment::Index(index.trim().parse().ok()?),
});
rest = after.strip_prefix('.').unwrap_or(after);
continue;
}
let end = rest.find(['.', '[']).unwrap_or(rest.len());
let (field, after) = rest.split_at(end);
if !field.is_empty() {
segments.push(Segment::Field(field.to_string()));
}
rest = after.strip_prefix('.').unwrap_or(after);
}
Some(segments)
}
fn without_path(value: Value, segments: &[Segment]) -> Value {
let Some((segment, tail)) = segments.split_first() else {
return value;
};
match (segment, value) {
(Segment::Field(name), Value::Object(object)) => Value::Object(
object
.into_iter()
.filter_map(|(key, item)| {
if key != *name {
return Some((key, item));
}
(!tail.is_empty()).then(|| (key, without_path(item, tail)))
})
.collect(),
),
(Segment::Every, Value::Array(items)) => Value::Array(
items
.into_iter()
.map(|item| without_path(item, tail))
.collect(),
),
(Segment::Index(index), Value::Array(items)) => Value::Array(
items
.into_iter()
.enumerate()
.map(|(position, item)| {
if position == *index {
without_path(item, tail)
} else {
item
}
})
.collect(),
),
(_, value) => value,
}
}
pub fn delete_nested_value(value: Value, path: &str) -> Value {
match parse_segments(path) {
Some(segments) => without_path(value, &segments),
None => value,
}
}
#[cfg(test)]
mod tests {
use rstest::{fixture, rstest};
use serde_json::json;
use super::*;
#[fixture]
fn body() -> Value {
json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
})
}
#[rstest]
#[case::top_level_field("top", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}}
}))]
#[case::whole_object("meta", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"top": 0.7
}))]
#[case::nested_field("meta.inner.drop", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"keep": 2}},
"top": 0.7
}))]
#[case::trailing_dot("meta.inner.drop.", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"keep": 2}},
"top": 0.7
}))]
#[case::leading_and_doubled_dots(".meta..inner.drop", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"keep": 2}},
"top": 0.7
}))]
#[case::field_in_every_element("tools[*].examples", json!({
"tools": [
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::whole_array_field_in_every_element("tools[*].arr", json!({
"tools": [
{"name": "t0", "examples": ["a"]},
{"name": "t1", "examples": ["b"]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::field_in_indexed_element("tools[1].examples", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::padded_index("tools[ 1 ].examples", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::field_right_after_bracket("tools[0]examples", json!({
"tools": [
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::index_then_wildcard("tools[0].arr[*].f", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::wildcard_then_index_only_where_it_exists("tools[*].arr[1].f", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::nested_wildcards("tools[*].arr[*].f", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
fn deletes_the_addressed_field(body: Value, #[case] path: &str, #[case] expected: Value) {
assert_eq!(delete_nested_value(body, path), expected);
}
#[rstest]
#[case::empty_path("")]
#[case::missing_field("missing")]
#[case::missing_parent("missing.field")]
#[case::field_through_a_scalar("top.value")]
#[case::field_on_an_array("tools.name")]
#[case::index_on_an_object("meta[0].user")]
#[case::wildcard_on_an_object("meta[*].user")]
#[case::wildcard_over_scalars("tools[*].examples[*].name")]
#[case::index_out_of_range("tools[5].name")]
#[case::every_element_itself("tools[*]")]
#[case::indexed_element_itself("tools[0]")]
#[case::nested_element_itself("tools[*].arr[0]")]
#[case::negative_index("tools[-1].name")]
#[case::non_numeric_index("tools[x].name")]
#[case::empty_index("tools[].name")]
#[case::unclosed_bracket("top[0")]
fn leaves_the_value_untouched(body: Value, #[case] path: &str) {
assert_eq!(delete_nested_value(body.clone(), path), body);
}
#[rstest]
#[case::wildcards_indices_and_nesting(
json!({"tools": [
{"name": "t0", "configs": [{"id": "c0", "remove_me": 1, "keep": 1}, {"id": "c1", "remove_me": 2, "keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
{"name": "t1", "configs": [{"id": "c0", "remove_me": 3, "keep": 3}, {"id": "c1", "remove_me": 4, "keep": 4}], "metadata": {"drop_this": 2, "preserve": 2}},
{"name": "t2", "configs": [{"id": "c0", "remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
]}),
&["tools[*].configs[1].remove_me", "tools[1].metadata.drop_this", "tools[*].configs[*].id"],
json!({"tools": [
{"name": "t0", "configs": [{"remove_me": 1, "keep": 1}, {"keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
{"name": "t1", "configs": [{"remove_me": 3, "keep": 3}, {"keep": 4}], "metadata": {"preserve": 2}},
{"name": "t2", "configs": [{"remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
]}),
)]
#[case::simple_and_wildcard_nesting(
json!({
"tools": [{"name": "t1", "simple_nested": {"remove": 1, "keep": 2}, "complex": [{"nested": {"remove": 3, "keep": 4}}]}],
"top_level_remove": "should_go",
"top_level_keep": "should_stay"
}),
&["tools[*].simple_nested.remove", "tools[*].complex[*].nested.remove"],
json!({
"tools": [{"name": "t1", "simple_nested": {"keep": 2}, "complex": [{"nested": {"keep": 4}}]}],
"top_level_remove": "should_go",
"top_level_keep": "should_stay"
}),
)]
#[case::triple_nested_wildcards(
json!({"tools": [{"name": "t1", "arr1": [
{"arr2": [{"field": 1, "keep": 1}, {"field": 2, "keep": 2}]},
{"arr2": [{"field": 3, "keep": 3}]}
]}]}),
&["tools[*].arr1[*].arr2[*].field"],
json!({"tools": [{"name": "t1", "arr1": [
{"arr2": [{"keep": 1}, {"keep": 2}]},
{"arr2": [{"keep": 3}]}
]}]}),
)]
fn applies_paths_in_sequence(
#[case] value: Value,
#[case] paths: &[&str],
#[case] expected: Value,
) {
let deleted = paths
.iter()
.fold(value, |value, path| delete_nested_value(value, path));
assert_eq!(deleted, expected);
}
}

View file

@ -0,0 +1,93 @@
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
use serde_json::{Map, Value};
pub fn get_provider_specific_headers(
provider_specific_header: Option<&ProviderSpecificHeaders>,
custom_llm_provider: &str,
) -> Map<String, Value> {
let entries: &[ProviderSpecificHeader] = match provider_specific_header {
None => &[],
Some(ProviderSpecificHeaders::One(entry)) => std::slice::from_ref(entry),
Some(ProviderSpecificHeaders::Many(entries)) => entries,
};
entries
.iter()
.filter(|entry| {
entry
.custom_llm_provider
.split(',')
.any(|scoped| scoped.trim() == custom_llm_provider)
})
.flat_map(|entry| entry.extra_headers.clone())
.collect()
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
#[rstest]
#[case::single_entry_for_the_provider(
json!({"custom_llm_provider": "anthropic", "extra_headers": {"Authorization": "Bearer t", "Custom-Header": "v"}}),
json!({"Authorization": "Bearer t", "Custom-Header": "v"}),
)]
#[case::single_entry_for_another_provider(
json!({"custom_llm_provider": "openai", "extra_headers": {"Authorization": "Bearer t"}}),
json!({}),
)]
#[case::provider_in_a_comma_separated_scope(
json!({"custom_llm_provider": "bedrock,anthropic,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}}),
json!({"anthropic-beta": "context-1m-2025-08-07"}),
)]
#[case::provider_missing_from_a_comma_separated_scope(
json!({"custom_llm_provider": "bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
json!({}),
)]
#[case::scope_with_spaces(
json!({"custom_llm_provider": "bedrock, anthropic , vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
json!({"anthropic-beta": "test"}),
)]
#[case::scope_names_must_match_exactly(
json!({"custom_llm_provider": "anthropic_text", "extra_headers": {"anthropic-beta": "test"}}),
json!({}),
)]
#[case::entries_scope_independently(
json!([
{"custom_llm_provider": "anthropic,bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}},
{"custom_llm_provider": "bedrock", "extra_headers": {"x-bedrock-only": "no"}},
{"custom_llm_provider": "anthropic", "extra_headers": {"authorization": "Bearer sk-ant-oat01-fake-token"}}
]),
json!({"anthropic-beta": "context-1m-2025-08-07", "authorization": "Bearer sk-ant-oat01-fake-token"}),
)]
#[case::later_entries_win(
json!([
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "first"}},
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "second"}}
]),
json!({"x-scoped": "second"}),
)]
#[case::empty_list(json!([]), json!({}))]
#[case::entry_without_scope(json!({"extra_headers": {"x-scoped": "yes"}}), json!({}))]
#[case::entry_without_headers(json!({"custom_llm_provider": "anthropic"}), json!({}))]
fn provider_specific_headers_match_the_scoped_provider(
#[case] configured: Value,
#[case] expected: Value,
) {
let configured: ProviderSpecificHeaders = serde_json::from_value(configured).unwrap();
assert_eq!(
Value::Object(get_provider_specific_headers(
Some(&configured),
"anthropic"
)),
expected
);
}
#[test]
fn no_configured_headers_match_nothing() {
assert_eq!(get_provider_specific_headers(None, "anthropic"), Map::new());
}
}

View file

@ -1,7 +1,9 @@
pub mod call_arguments;
pub mod core_helpers;
pub mod dot_notation_indexing;
pub mod exception_mapping_utils;
pub mod get_llm_provider_logic;
pub mod get_provider_specific_headers;
pub mod params;
pub mod prompt_templates;
pub mod secret_redaction;

View file

@ -4,7 +4,6 @@ version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
autotests = false
[dependencies]
litellm-secrets.workspace = true

View file

@ -14,6 +14,3 @@ pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Resu
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
.await
}
#[cfg(test)]
mod tests;

View file

@ -51,6 +51,3 @@ pub fn chat_completions_decline_reason(
.unsupported_reason(&messages, optional_params)
.map(|reason| reason.0)
}
#[cfg(test)]
mod tests;

View file

@ -143,3 +143,841 @@ pub(super) fn prepare_provider_request(
timeout: request.timeout,
})
}
#[cfg(test)]
mod tests {
use litellm_llms::base_llm::chat::transformation::RequestAuth;
use serde_json::{Map, Value, json};
use super::{prepare_provider_request, resolve_request};
use crate::chat_completions::{
Error,
types::{ChatCompletionsRequest, ProviderChatCompletionsRequest},
};
fn prepare_chat_completions_call(
request: ChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
prepare_provider_request(resolve_request(request)?)
}
fn request<'a>(
model: &'a str,
provider: Option<&'a str>,
messages: Value,
optional_params: Value,
) -> ChatCompletionsRequest<'a> {
ChatCompletionsRequest {
model,
messages,
optional_params: match optional_params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: provider,
extra_headers: None,
timeout: None,
}
}
/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers
/// carry resolved credentials), so unwrap the failure case by hand.
fn decline(request: ChatCompletionsRequest<'_>) -> Error {
match prepare_chat_completions_call(request) {
Err(error) => error,
Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url),
}
}
#[test]
fn resolves_the_provider_from_the_model_prefix() {
let prepared = prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.expect("prepares");
assert_eq!(prepared.model, "claude-sonnet-4-5");
assert_eq!(prepared.url, "https://api.anthropic.com/v1/messages");
assert_eq!(prepared.body["model"], json!("claude-sonnet-4-5"));
}
#[test]
fn strips_an_explicit_provider_prefix_from_the_model() {
let prepared = prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
))
.expect("prepares");
assert_eq!(prepared.model, "claude-sonnet-4-5");
}
#[test]
fn adds_the_auth_and_default_headers() {
let prepared = prepare_chat_completions_call(request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
))
.expect("prepares");
assert!(
prepared
.upstream_headers
.contains(&("x-api-key".to_string(), "sk-test".to_string()))
);
assert!(
prepared
.upstream_headers
.contains(&("anthropic-version".to_string(), "2023-06-01".to_string()))
);
assert!(matches!(
prepared.auth,
RequestAuth::Header {
name: "x-api-key",
..
}
));
}
#[test]
fn the_deployment_credential_replaces_a_caller_supplied_auth_header() {
// Python builds `{**headers, **anthropic_headers}`, so the deployment's key
// overwrites a forwarded one. Honouring the caller's would let whoever sends
// the request choose the Anthropic principal it bills to.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
"X-Api-Key".to_string(),
json!("sk-caller"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys[0].1, "sk-test");
}
#[test]
fn a_forwarded_authorization_header_suppresses_the_resolved_api_key_header() {
// Anthropic's `validate_environment` pops `x-api-key` and sets `authorization`
// for an OAuth token, so re-adding the key here would put the credential into
// a header the host removed on purpose.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
(
"Authorization".to_string(),
json!("Bearer sk-ant-oat01-token"),
),
("X-Api-Key".to_string(), json!("sk-caller")),
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert!(
!prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("x-api-key") && value == "sk-test"),
"the resolved key must not be applied over an OAuth bearer, got {:?}",
prepared.upstream_headers
);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-token")
);
}
#[test]
fn an_unrelated_forwarded_authorization_does_not_defer_the_resolved_key() {
// Only an OAuth bearer replaces the credential. Python sends the deployment's
// `x-api-key` alongside any other forwarded `authorization`, so deferring on
// the mere presence of that header would drop the deployment's auth.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
("Authorization".to_string(), json!("Bearer unrelated")),
("X-Api-Key".to_string(), json!("sk-caller")),
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys[0].1, "sk-test");
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer unrelated"),
"the unrelated authorization must survive, got {:?}",
prepared.upstream_headers
);
}
#[test]
fn declines_an_unsupported_request_before_resolving_credentials() {
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({"stream": true}),
);
call.api_key = None;
// No api_key is set and no env is consulted: the gate must run first, so the
// error is the decline rather than a missing-credential error.
assert_eq!(decline(call), Error::Unsupported("streaming"));
}
#[test]
fn rejects_an_unknown_provider() {
assert_eq!(
decline(request(
"openai/gpt-4o",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
Error::InvalidProvider("openai".to_string())
);
}
#[test]
fn rejects_a_model_with_no_resolvable_provider() {
assert!(matches!(
decline(request(
"claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
Error::InvalidProvider(_)
));
}
#[test]
fn rejects_an_empty_or_malformed_message_list() {
assert_eq!(
decline(request(
"anthropic/claude-sonnet-4-5",
None,
json!([]),
json!({}),
)),
Error::InvalidRequest("chat completions requires at least one message".to_string())
);
assert!(matches!(
decline(request(
"anthropic/claude-sonnet-4-5",
None,
json!("not a list"),
json!({}),
)),
Error::InvalidRequest(_)
));
}
#[test]
fn rejects_non_string_extra_headers() {
let mut call = request(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
assert_eq!(
decline(call),
Error::Headers(litellm_http::request::HeaderError {
context: "chat completions",
name: "x-trace".to_string(),
actual: "number",
})
);
}
#[test]
fn prepares_a_bedrock_call_without_resolving_credentials() {
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.api_key = None;
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert_eq!(
prepared.url,
"https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse"
);
assert_eq!(
prepared.auth,
RequestAuth::AwsSigV4 {
region: "us-east-1".to_string(),
service: "bedrock",
}
);
// SigV4 signs the serialized body, so prepare must not have added an
// Authorization header; the handler does it.
assert!(
!prepared
.upstream_headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("authorization"))
);
assert_eq!(prepared.body["inferenceConfig"], json!({"maxTokens": 16}));
}
#[tokio::test]
async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() {
// Python signs only the AWS header set and reattaches the rest, so a header
// the caller forwarded rides along without joining the canonical request.
// Signing it makes Converse 403 on a deployment that works on Python.
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
);
// A key would resolve to a bearer token and never reach the signer.
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(
"x-request-id".to_string(),
json!("abc-123"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let signed = crate::chat_completions::handler::outbound_request(&prepared)
.await
.expect("signs");
let authorization = signed
.header("authorization")
.expect("carries an authorization header")
.to_string();
assert!(
authorization.starts_with("AWS4-HMAC-SHA256"),
"expected a SigV4 signature, got {authorization}"
);
assert!(
!authorization.contains("x-request-id"),
"forwarded header reached SignedHeaders: {authorization}"
);
// It still goes on the wire, it is just not part of the signature.
assert!(
signed
.headers()
.iter()
.any(|(name, value)| name == "x-request-id" && value == "abc-123"),
"forwarded header was dropped instead of reattached"
);
}
#[tokio::test]
async fn a_forwarded_header_the_signer_computes_declines_to_python() {
// Reattaching the caller's copy next to the computed one puts the name on
// the wire twice and Bedrock rejects the pair, so a request carrying one
// has to go to Python instead of being signed here.
for forwarded in [
"Authorization",
"x-amz-date",
"x-amz-security-token",
"Date",
] {
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
);
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let error = crate::chat_completions::handler::outbound_request(&prepared)
.await
.expect_err("{forwarded} should decline instead of being signed");
assert!(
matches!(error, Error::Unsupported(_)),
"{forwarded} declined as {error:?}, which the host would not fall back on"
);
}
}
#[test]
fn a_bedrock_deployment_bearer_outranks_a_forwarded_authorization() {
// `get_request_headers` assigns `headers["Authorization"]` unconditionally
// once a bearer token resolves, so the deployment's identity wins on
// Python. Keeping the caller's would authorize and bill the call as a
// different principal, and only when the deployment carries `rust: true`.
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.extra_headers = Some(Map::from_iter([(
"Authorization".to_string(),
json!("Bearer caller-supplied"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let authorizations: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("authorization"))
.map(|(_, value)| value.as_str())
.collect();
assert_eq!(
authorizations,
vec!["Bearer sk-test"],
"the deployment token must be the only authorization on the wire"
);
}
#[test]
fn an_anthropic_forwarded_oauth_bearer_still_outranks_the_resolved_key() {
// The opposite precedence, and deliberate: Anthropic's own transform
// honours a forwarded OAuth bearer, so the Bedrock fix above must not be
// generalized into a rule that the configured key always wins.
//
// An OAuth bearer is the whole of that exception. This forwarded a plain
// `x-api-key` until round 17, which read as the same claim and was not:
// Python overwrites a forwarded `x-api-key` with the deployment's.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
"authorization".to_string(),
json!("Bearer sk-ant-oat01-forwarded"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.map(|(_, value)| value.as_str())
.collect();
assert!(keys.is_empty(), "got {:?}", prepared.upstream_headers);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-forwarded")
);
}
#[test]
fn a_bedrock_api_key_is_sent_as_a_bearer_token_instead_of_being_signed() {
// The configured bearer identity has its own account and quota boundary,
// so a request carrying one must not be signed as whatever principal the
// host's AWS credentials resolve to.
let prepared = prepare_chat_completions_call(request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
))
.expect("prepares");
assert_eq!(
prepared.auth,
RequestAuth::Bearer {
token: "sk-test".to_string()
}
);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-test"),
"prepare did not carry the bearer token"
);
}
fn decline_reason(
model: &str,
provider: Option<&str>,
messages: Value,
params: Value,
) -> Option<&'static str> {
let params = match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
};
crate::chat_completions::chat_completions_decline_reason(model, provider, messages, &params)
}
#[test]
fn the_gate_accepts_what_prepare_accepts() {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
),
None
);
}
#[test]
fn the_gate_declines_without_resolving_credentials_or_calling_out() {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"stream": true}),
),
Some("streaming")
);
assert_eq!(
decline_reason(
"openai/gpt-4o",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
),
Some("provider is not on the rust chat completions path")
);
assert_eq!(
decline_reason(
"claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
),
Some("provider is not on the rust chat completions path")
);
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!("nope"),
json!({})
),
Some("unreadable message list")
);
assert_eq!(
decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})),
Some("empty message list")
);
}
#[test]
fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
// A gate that accepts what prepare then declines would make the host emit
// its pre-call logging on a path that falls back, so pin the agreement.
for (messages, params) in [
(
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 8}),
),
(
json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]),
json!({"temperature": 0.1}),
),
(
json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]),
json!({}),
),
] {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
messages.clone(),
params.clone()
),
None,
"gate declined {messages}"
);
prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
None,
messages.clone(),
params,
))
.unwrap_or_else(|error| panic!("prepare declined {messages}: {error}"));
}
}
mod round_trip {
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
use super::*;
use crate::chat_completions::chat_completions;
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
let header_end = loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break request.len();
}
request.extend_from_slice(&buffer[..n]);
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n")
{
break position + 4;
}
};
let headers = String::from_utf8_lossy(&request[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
while request.len().saturating_sub(header_end) < content_length {
let n = socket.read(&mut buffer).await.expect("reads body");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
}
String::from_utf8(request).expect("request is utf8")
}
fn http_response(status: &str, body: &str) -> String {
format!(
"HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
)
}
/// Serve one request from a stub upstream and hand back what it received.
async fn serve_once(
status: &'static str,
body: &'static str,
) -> (String, tokio::task::JoinHandle<String>) {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let port = listener.local_addr().expect("addr").port();
let handle = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts");
let received = read_http_request(&mut socket).await;
socket
.write_all(http_response(status, body).as_bytes())
.await
.expect("writes response");
socket.flush().await.expect("flushes");
received
});
(format!("http://127.0.0.1:{port}/v1/messages"), handle)
}
fn call(api_base: &str, messages: Value, params: Value) -> ChatCompletionsRequest<'_> {
ChatCompletionsRequest {
model: "anthropic/claude-sonnet-4-5",
messages,
optional_params: match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: Some(api_base),
custom_llm_provider: None,
extra_headers: None,
timeout: Some(std::time::Duration::from_secs(10)),
}
}
const GOOD_BODY: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
#[tokio::test]
async fn round_trip_sends_the_translated_body_and_normalizes_the_response() {
let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await;
let response = chat_completions(call(
&api_base,
json!([
{"role": "system", "content": "be terse"},
{"role": "user", "content": "hi"}
]),
json!({"max_tokens": 16}),
))
.await
.expect("call succeeds");
let received = handle.await.expect("server task");
let sent: Value = serde_json::from_str(
received
.split_once("\r\n\r\n")
.expect("request has a body")
.1,
)
.expect("body is json");
assert_eq!(
sent["messages"],
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}])
);
assert_eq!(
sent["system"],
json!([{"type": "text", "text": "be terse"}])
);
assert_eq!(sent["max_tokens"], json!(16));
assert!(received.to_lowercase().contains("x-api-key: sk-test"));
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
assert_eq!(response.usage.total_tokens, 15);
}
#[tokio::test]
async fn a_response_it_cannot_normalize_is_reported_as_already_sent() {
// The provider was called and billed, so the host must not retry this
// on its own path. `MissingField` here would read as a pre-send
// decline and be retried; `InvalidResponse` cannot.
const NO_USAGE: &str =
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#;
let (api_base, handle) = serve_once("200 OK", NO_USAGE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
#[tokio::test]
async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() {
const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#;
let (api_base, handle) = serve_once("200 OK", TOOL_USE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
#[tokio::test]
async fn an_upstream_error_status_keeps_its_code() {
let (api_base, handle) =
serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("upstream rejects");
handle.await.expect("server task");
assert!(
matches!(
err,
Error::Transport(litellm_http::transport::Error::Http { status: 429, .. })
),
"expected a 429, got {err:?}"
);
}
#[tokio::test]
async fn a_connection_that_is_never_established_declines_instead_of_failing() {
// Nothing was sent, so nothing was billed and the host can still serve
// the request. Classing this with the post-send failures would turn a
// recoverable fallback into a user-facing error on exactly the
// deployments whose transport is configured only on the Python client.
let port = {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
listener.local_addr().expect("has an address").port()
// Dropped here, so the port is closed and the connect is refused.
};
let err = chat_completions(call(
&format!("http://127.0.0.1:{port}/v1/messages"),
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("nothing is listening");
assert!(
matches!(
err,
Error::Transport(litellm_http::transport::Error::Connect(_))
),
"expected a pre-send connect failure, got {err:?}"
);
}
#[test]
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
use crate::chat_completions::handler::as_response_error;
for original in [
Error::MissingField("usage"),
Error::Unsupported("non-text response content block"),
Error::InvalidRequest("whatever".to_string()),
Error::Auth(litellm_auth::Error::InvalidHeader),
] {
let label = format!("{original:?}");
assert!(
matches!(as_response_error(original), Error::InvalidResponse(_)),
"{label} must not stay retryable once the provider has answered"
);
}
// An upstream status is already unambiguous, so it survives intact.
assert!(matches!(
as_response_error(Error::Transport(litellm_http::transport::Error::Http {
status: 500,
body: "boom".to_string()
})),
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
));
}
}
}

View file

@ -1,833 +0,0 @@
use litellm_llms::base_llm::chat::transformation::RequestAuth;
use serde_json::{Map, Value, json};
use super::{
Error,
prepare::{prepare_provider_request, resolve_request},
};
use crate::chat_completions::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest};
fn prepare_chat_completions_call(
request: ChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
prepare_provider_request(resolve_request(request)?)
}
fn request<'a>(
model: &'a str,
provider: Option<&'a str>,
messages: Value,
optional_params: Value,
) -> ChatCompletionsRequest<'a> {
ChatCompletionsRequest {
model,
messages,
optional_params: match optional_params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: provider,
extra_headers: None,
timeout: None,
}
}
/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers
/// carry resolved credentials), so unwrap the failure case by hand.
fn decline(request: ChatCompletionsRequest<'_>) -> Error {
match prepare_chat_completions_call(request) {
Err(error) => error,
Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url),
}
}
#[test]
fn resolves_the_provider_from_the_model_prefix() {
let prepared = prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.expect("prepares");
assert_eq!(prepared.model, "claude-sonnet-4-5");
assert_eq!(prepared.url, "https://api.anthropic.com/v1/messages");
assert_eq!(prepared.body["model"], json!("claude-sonnet-4-5"));
}
#[test]
fn strips_an_explicit_provider_prefix_from_the_model() {
let prepared = prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
))
.expect("prepares");
assert_eq!(prepared.model, "claude-sonnet-4-5");
}
#[test]
fn adds_the_auth_and_default_headers() {
let prepared = prepare_chat_completions_call(request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
))
.expect("prepares");
assert!(
prepared
.upstream_headers
.contains(&("x-api-key".to_string(), "sk-test".to_string()))
);
assert!(
prepared
.upstream_headers
.contains(&("anthropic-version".to_string(), "2023-06-01".to_string()))
);
assert!(matches!(
prepared.auth,
RequestAuth::Header {
name: "x-api-key",
..
}
));
}
#[test]
fn the_deployment_credential_replaces_a_caller_supplied_auth_header() {
// Python builds `{**headers, **anthropic_headers}`, so the deployment's key
// overwrites a forwarded one. Honouring the caller's would let whoever sends
// the request choose the Anthropic principal it bills to.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
"X-Api-Key".to_string(),
json!("sk-caller"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys[0].1, "sk-test");
}
#[test]
fn a_forwarded_authorization_header_suppresses_the_resolved_api_key_header() {
// Anthropic's `validate_environment` pops `x-api-key` and sets `authorization`
// for an OAuth token, so re-adding the key here would put the credential into
// a header the host removed on purpose.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
(
"Authorization".to_string(),
json!("Bearer sk-ant-oat01-token"),
),
("X-Api-Key".to_string(), json!("sk-caller")),
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert!(
!prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("x-api-key") && value == "sk-test"),
"the resolved key must not be applied over an OAuth bearer, got {:?}",
prepared.upstream_headers
);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-token")
);
}
#[test]
fn an_unrelated_forwarded_authorization_does_not_defer_the_resolved_key() {
// Only an OAuth bearer replaces the credential. Python sends the deployment's
// `x-api-key` alongside any other forwarded `authorization`, so deferring on
// the mere presence of that header would drop the deployment's auth.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
("Authorization".to_string(), json!("Bearer unrelated")),
("X-Api-Key".to_string(), json!("sk-caller")),
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys[0].1, "sk-test");
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer unrelated"),
"the unrelated authorization must survive, got {:?}",
prepared.upstream_headers
);
}
#[test]
fn declines_an_unsupported_request_before_resolving_credentials() {
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({"stream": true}),
);
call.api_key = None;
// No api_key is set and no env is consulted: the gate must run first, so the
// error is the decline rather than a missing-credential error.
assert_eq!(decline(call), Error::Unsupported("streaming"));
}
#[test]
fn rejects_an_unknown_provider() {
assert_eq!(
decline(request(
"openai/gpt-4o",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
Error::InvalidProvider("openai".to_string())
);
}
#[test]
fn rejects_a_model_with_no_resolvable_provider() {
assert!(matches!(
decline(request(
"claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
Error::InvalidProvider(_)
));
}
#[test]
fn rejects_an_empty_or_malformed_message_list() {
assert_eq!(
decline(request(
"anthropic/claude-sonnet-4-5",
None,
json!([]),
json!({}),
)),
Error::InvalidRequest("chat completions requires at least one message".to_string())
);
assert!(matches!(
decline(request(
"anthropic/claude-sonnet-4-5",
None,
json!("not a list"),
json!({}),
)),
Error::InvalidRequest(_)
));
}
#[test]
fn rejects_non_string_extra_headers() {
let mut call = request(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
assert_eq!(
decline(call),
Error::Headers(litellm_http::request::HeaderError {
context: "chat completions",
name: "x-trace".to_string(),
actual: "number",
})
);
}
#[test]
fn prepares_a_bedrock_call_without_resolving_credentials() {
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.api_key = None;
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert_eq!(
prepared.url,
"https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse"
);
assert_eq!(
prepared.auth,
RequestAuth::AwsSigV4 {
region: "us-east-1".to_string(),
service: "bedrock",
}
);
// SigV4 signs the serialized body, so prepare must not have added an
// Authorization header; the handler does it.
assert!(
!prepared
.upstream_headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("authorization"))
);
assert_eq!(prepared.body["inferenceConfig"], json!({"maxTokens": 16}));
}
#[tokio::test]
async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() {
// Python signs only the AWS header set and reattaches the rest, so a header
// the caller forwarded rides along without joining the canonical request.
// Signing it makes Converse 403 on a deployment that works on Python.
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
);
// A key would resolve to a bearer token and never reach the signer.
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(
"x-request-id".to_string(),
json!("abc-123"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let signed = super::handler::outbound_request(&prepared)
.await
.expect("signs");
let authorization = signed
.header("authorization")
.expect("carries an authorization header")
.to_string();
assert!(
authorization.starts_with("AWS4-HMAC-SHA256"),
"expected a SigV4 signature, got {authorization}"
);
assert!(
!authorization.contains("x-request-id"),
"forwarded header reached SignedHeaders: {authorization}"
);
// It still goes on the wire, it is just not part of the signature.
assert!(
signed
.headers()
.iter()
.any(|(name, value)| name == "x-request-id" && value == "abc-123"),
"forwarded header was dropped instead of reattached"
);
}
#[tokio::test]
async fn a_forwarded_header_the_signer_computes_declines_to_python() {
// Reattaching the caller's copy next to the computed one puts the name on
// the wire twice and Bedrock rejects the pair, so a request carrying one
// has to go to Python instead of being signed here.
for forwarded in [
"Authorization",
"x-amz-date",
"x-amz-security-token",
"Date",
] {
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
);
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let error = super::handler::outbound_request(&prepared)
.await
.expect_err("{forwarded} should decline instead of being signed");
assert!(
matches!(error, Error::Unsupported(_)),
"{forwarded} declined as {error:?}, which the host would not fall back on"
);
}
}
#[test]
fn a_bedrock_deployment_bearer_outranks_a_forwarded_authorization() {
// `get_request_headers` assigns `headers["Authorization"]` unconditionally
// once a bearer token resolves, so the deployment's identity wins on
// Python. Keeping the caller's would authorize and bill the call as a
// different principal, and only when the deployment carries `rust: true`.
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.extra_headers = Some(Map::from_iter([(
"Authorization".to_string(),
json!("Bearer caller-supplied"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let authorizations: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("authorization"))
.map(|(_, value)| value.as_str())
.collect();
assert_eq!(
authorizations,
vec!["Bearer sk-test"],
"the deployment token must be the only authorization on the wire"
);
}
#[test]
fn an_anthropic_forwarded_oauth_bearer_still_outranks_the_resolved_key() {
// The opposite precedence, and deliberate: Anthropic's own transform
// honours a forwarded OAuth bearer, so the Bedrock fix above must not be
// generalized into a rule that the configured key always wins.
//
// An OAuth bearer is the whole of that exception. This forwarded a plain
// `x-api-key` until round 17, which read as the same claim and was not:
// Python overwrites a forwarded `x-api-key` with the deployment's.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
"authorization".to_string(),
json!("Bearer sk-ant-oat01-forwarded"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.map(|(_, value)| value.as_str())
.collect();
assert!(keys.is_empty(), "got {:?}", prepared.upstream_headers);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-forwarded")
);
}
#[test]
fn a_bedrock_api_key_is_sent_as_a_bearer_token_instead_of_being_signed() {
// The configured bearer identity has its own account and quota boundary,
// so a request carrying one must not be signed as whatever principal the
// host's AWS credentials resolve to.
let prepared = prepare_chat_completions_call(request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
))
.expect("prepares");
assert_eq!(
prepared.auth,
RequestAuth::Bearer {
token: "sk-test".to_string()
}
);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-test"),
"prepare did not carry the bearer token"
);
}
fn decline_reason(
model: &str,
provider: Option<&str>,
messages: Value,
params: Value,
) -> Option<&'static str> {
let params = match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
};
super::chat_completions_decline_reason(model, provider, messages, &params)
}
#[test]
fn the_gate_accepts_what_prepare_accepts() {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
),
None
);
}
#[test]
fn the_gate_declines_without_resolving_credentials_or_calling_out() {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"stream": true}),
),
Some("streaming")
);
assert_eq!(
decline_reason(
"openai/gpt-4o",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
),
Some("provider is not on the rust chat completions path")
);
assert_eq!(
decline_reason(
"claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
),
Some("provider is not on the rust chat completions path")
);
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!("nope"),
json!({})
),
Some("unreadable message list")
);
assert_eq!(
decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})),
Some("empty message list")
);
}
#[test]
fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
// A gate that accepts what prepare then declines would make the host emit
// its pre-call logging on a path that falls back, so pin the agreement.
for (messages, params) in [
(
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 8}),
),
(
json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]),
json!({"temperature": 0.1}),
),
(
json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]),
json!({}),
),
] {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
messages.clone(),
params.clone()
),
None,
"gate declined {messages}"
);
prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
None,
messages.clone(),
params,
))
.unwrap_or_else(|error| panic!("prepare declined {messages}: {error}"));
}
}
mod round_trip {
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
use super::*;
use crate::chat_completions::chat_completions;
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
let header_end = loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break request.len();
}
request.extend_from_slice(&buffer[..n]);
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") {
break position + 4;
}
};
let headers = String::from_utf8_lossy(&request[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
while request.len().saturating_sub(header_end) < content_length {
let n = socket.read(&mut buffer).await.expect("reads body");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
}
String::from_utf8(request).expect("request is utf8")
}
fn http_response(status: &str, body: &str) -> String {
format!(
"HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
)
}
/// Serve one request from a stub upstream and hand back what it received.
async fn serve_once(
status: &'static str,
body: &'static str,
) -> (String, tokio::task::JoinHandle<String>) {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let port = listener.local_addr().expect("addr").port();
let handle = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts");
let received = read_http_request(&mut socket).await;
socket
.write_all(http_response(status, body).as_bytes())
.await
.expect("writes response");
socket.flush().await.expect("flushes");
received
});
(format!("http://127.0.0.1:{port}/v1/messages"), handle)
}
fn call(api_base: &str, messages: Value, params: Value) -> ChatCompletionsRequest<'_> {
ChatCompletionsRequest {
model: "anthropic/claude-sonnet-4-5",
messages,
optional_params: match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: Some(api_base),
custom_llm_provider: None,
extra_headers: None,
timeout: Some(std::time::Duration::from_secs(10)),
}
}
const GOOD_BODY: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
#[tokio::test]
async fn round_trip_sends_the_translated_body_and_normalizes_the_response() {
let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await;
let response = chat_completions(call(
&api_base,
json!([
{"role": "system", "content": "be terse"},
{"role": "user", "content": "hi"}
]),
json!({"max_tokens": 16}),
))
.await
.expect("call succeeds");
let received = handle.await.expect("server task");
let sent: Value = serde_json::from_str(
received
.split_once("\r\n\r\n")
.expect("request has a body")
.1,
)
.expect("body is json");
assert_eq!(
sent["messages"],
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}])
);
assert_eq!(
sent["system"],
json!([{"type": "text", "text": "be terse"}])
);
assert_eq!(sent["max_tokens"], json!(16));
assert!(received.to_lowercase().contains("x-api-key: sk-test"));
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
assert_eq!(response.usage.total_tokens, 15);
}
#[tokio::test]
async fn a_response_it_cannot_normalize_is_reported_as_already_sent() {
// The provider was called and billed, so the host must not retry this
// on its own path. `MissingField` here would read as a pre-send
// decline and be retried; `InvalidResponse` cannot.
const NO_USAGE: &str =
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#;
let (api_base, handle) = serve_once("200 OK", NO_USAGE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
#[tokio::test]
async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() {
const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#;
let (api_base, handle) = serve_once("200 OK", TOOL_USE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
#[tokio::test]
async fn an_upstream_error_status_keeps_its_code() {
let (api_base, handle) =
serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("upstream rejects");
handle.await.expect("server task");
assert!(
matches!(
err,
Error::Transport(litellm_http::transport::Error::Http { status: 429, .. })
),
"expected a 429, got {err:?}"
);
}
#[tokio::test]
async fn a_connection_that_is_never_established_declines_instead_of_failing() {
// Nothing was sent, so nothing was billed and the host can still serve
// the request. Classing this with the post-send failures would turn a
// recoverable fallback into a user-facing error on exactly the
// deployments whose transport is configured only on the Python client.
let port = {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
listener.local_addr().expect("has an address").port()
// Dropped here, so the port is closed and the connect is refused.
};
let err = chat_completions(call(
&format!("http://127.0.0.1:{port}/v1/messages"),
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("nothing is listening");
assert!(
matches!(
err,
Error::Transport(litellm_http::transport::Error::Connect(_))
),
"expected a pre-send connect failure, got {err:?}"
);
}
#[test]
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
use crate::chat_completions::handler::as_response_error;
for original in [
Error::MissingField("usage"),
Error::Unsupported("non-text response content block"),
Error::InvalidRequest("whatever".to_string()),
Error::Auth(litellm_auth::Error::InvalidHeader),
] {
let label = format!("{original:?}");
assert!(
matches!(as_response_error(original), Error::InvalidResponse(_)),
"{label} must not stay retryable once the provider has answered"
);
}
// An upstream status is already unambiguous, so it survives intact.
assert!(matches!(
as_response_error(Error::Transport(litellm_http::transport::Error::Http {
status: 500,
body: "boom".to_string()
})),
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
));
}
}

View file

@ -1,5 +1,5 @@
use litellm_http::request::string_headers as shared_string_headers;
pub(super) use litellm_http::request::{has_bearer_auth, has_header, truncate_error_body};
pub(super) use litellm_http::request::truncate_error_body;
use litellm_llms::{
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
@ -26,3 +26,186 @@ pub(super) fn string_headers(
) -> Result<Vec<(String, String)>, Error> {
shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(Error::from)
}
#[cfg(test)]
mod tests {
use std::{sync::Arc, time::Duration};
use futures_util::future::BoxFuture;
use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::{Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
use super::{messages_provider_config, string_headers, truncate_error_body};
use crate::messages::{
Error,
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
types::MessagesShaping,
};
struct RecordingSecrets {
values: Vec<(&'static str, String)>,
requested: std::sync::Mutex<Vec<String>>,
}
impl SecretSource for RecordingSecrets {
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
Box::pin(async move {
self.requested.lock().unwrap().push(name.to_string());
Ok(self
.values
.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| SecretValue::new(value.clone())))
})
}
}
fn secrets_call() -> MessagesCall {
let Value::Object(body) = json!({
"model": "claude-sonnet-4-5",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hi"}]
}) else {
unreachable!("literal object")
};
MessagesCall {
model: "claude-sonnet-4-5".into(),
body,
api_key: None,
api_base: None,
custom_llm_provider: Some("anthropic".into()),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
}
}
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
let header_end = loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break request.len();
}
request.extend_from_slice(&buffer[..n]);
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") {
break position + 4;
}
};
let headers = String::from_utf8_lossy(&request[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
while request.len().saturating_sub(header_end) < content_length {
let n = socket.read(&mut buffer).await.expect("reads body");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
}
String::from_utf8(request).expect("request is utf8")
}
#[tokio::test]
async fn route_reads_the_provider_credential_and_base_from_the_secret_source() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let addr = listener.local_addr().expect("addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts request");
let request = read_http_request(&mut socket).await;
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
request
});
let secrets = Arc::new(RecordingSecrets {
values: vec![
("ANTHROPIC_API_KEY", "sk-from-manager".to_string()),
("ANTHROPIC_BASE_URL", format!("http://{addr}")),
],
requested: std::sync::Mutex::new(Vec::new()),
});
let output = litellm_host::run::run(
messages_machine(secrets.clone()),
&LocalMessagesHost::new(secrets_call()),
)
.await
.expect("messages request succeeds");
assert!(matches!(output, MessagesOutput::Message(_)));
let request = server.await.expect("server task completes");
assert!(
request
.to_ascii_lowercase()
.contains("x-api-key: sk-from-manager"),
"{request}"
);
let requested = secrets.requested.lock().unwrap().clone();
assert_eq!(
requested,
messages_provider_config("anthropic")
.unwrap()
.secret_names()
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
);
}
#[test]
fn provider_config_resolves_anthropic_and_azure_ai() {
assert!(messages_provider_config("anthropic").is_some());
assert!(messages_provider_config("azure_ai").is_some());
assert!(messages_provider_config("openai").is_none());
}
#[test]
fn truncate_error_body_caps_long_payloads() {
let body = "x".repeat(400);
let truncated = truncate_error_body(&body);
assert!(truncated.ends_with("... (truncated)"));
let prefix_chars = truncated
.strip_suffix("... (truncated)")
.expect("truncated marker present")
.chars()
.count();
assert_eq!(prefix_chars, 256);
}
#[test]
fn string_headers_rejects_non_string_values() {
let headers = json!({"x-count": 3}).as_object().unwrap().clone();
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert_eq!(
err,
Error::Headers(litellm_http::request::HeaderError {
context: "messages",
name: "x-count".to_string(),
actual: "number",
})
);
}
}

View file

@ -1,3 +1,5 @@
use std::sync::Arc;
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
@ -18,8 +20,34 @@ pub enum Error {
Transport(#[from] litellm_http::transport::Error),
#[error(transparent)]
Headers(#[from] litellm_http::request::HeaderError),
#[error(transparent)]
Secret(#[from] SecretError),
}
#[derive(Clone, Debug, thiserror::Error)]
#[error(transparent)]
pub struct SecretError(Arc<litellm_secrets::Error>);
impl SecretError {
pub fn source_error(&self) -> &litellm_secrets::Error {
&self.0
}
}
impl From<litellm_secrets::Error> for Error {
fn from(error: litellm_secrets::Error) -> Self {
Self::Secret(SecretError(Arc::new(error)))
}
}
impl PartialEq for SecretError {
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.0, &other.0)
}
}
impl Eq for SecretError {}
impl From<LlmError> for Error {
fn from(error: LlmError) -> Self {
match error {

View file

@ -12,6 +12,9 @@ mod common_utils;
mod handler;
mod prepare;
pub mod route;
use std::sync::Arc;
use litellm_secrets::source::EnvironmentSecrets;
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
use serde_json::Value;
@ -31,15 +34,15 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
api_base: request.api_base.map(Into::into),
custom_llm_provider: request.custom_llm_provider.map(Into::into),
extra_headers: request.extra_headers,
provider_specific_header: request.provider_specific_header,
timeout: request.timeout,
shaping: request.shaping,
};
match litellm_host::run::run(messages_machine(), &LocalMessagesHost::new(call)).await? {
let secrets = Arc::new(EnvironmentSecrets::python_compatible());
match litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await? {
MessagesOutput::Message(message) => Ok(*message),
MessagesOutput::Streamed => Err(Error::Unsupported(
"streamed responses need a streaming host",
)),
}
}
#[cfg(test)]
mod tests;

View file

@ -1,51 +1,102 @@
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use litellm_llms::base_llm::anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, MessagesAuthStrategy,
use litellm_core_utils::{
dot_notation_indexing::delete_nested_value,
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
get_provider_specific_headers::get_provider_specific_headers,
settings::Lookup,
};
use litellm_llms::{
anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request,
base_llm::anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, MessagesTransformContext,
},
};
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
use serde_json::{Map, Value};
use super::{
Error,
common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers},
common_utils::{messages_provider_config, string_headers},
};
use crate::messages::types::{MessagesRequest, ProviderMessagesRequest};
pub(super) fn prepare_provider_request(
request: MessagesRequest<'_>,
) -> Result<ProviderMessagesRequest, Error> {
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
pub(super) struct ResolvedProvider<'a> {
pub(super) model: &'a str,
pub(super) provider: &'a str,
pub(super) config: &'static dyn BaseAnthropicMessagesConfig,
}
pub(super) fn resolve_provider<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,
) -> Result<ResolvedProvider<'a>, Error> {
let CustomLlmProvider {
model,
custom_llm_provider: provider,
} = get_custom_llm_provider(model, custom_llm_provider)
.or_else(|| {
request
.custom_llm_provider
.map(|provider| CustomLlmProvider {
model: request.model,
custom_llm_provider: provider,
})
custom_llm_provider.map(|provider| CustomLlmProvider {
model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(
"unable to resolve custom_llm_provider for messages request".to_string(),
)
})?;
let model = provider_info.model.to_string();
let provider = provider_info.custom_llm_provider;
let config = messages_provider_config(provider)
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
let env_lookup = |key: &str| std::env::var(key).ok();
Ok(ResolvedProvider {
model,
provider,
config,
})
}
let headers =
validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?;
pub(super) fn prepare_provider_request(
request: MessagesRequest<'_>,
resolved: ResolvedProvider<'_>,
secrets: &dyn Lookup,
) -> Result<ProviderMessagesRequest, Error> {
let ResolvedProvider {
model,
provider,
config,
} = resolved;
let model = model.to_string();
let env_lookup = |key: &str| secrets.get(key);
let typed_request: AnthropicMessagesRequest =
serde_json::from_value(request.body).map_err(|err| {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
})?;
let transformed = config.transform_anthropic_messages_request(AnthropicMessagesRequest {
model: model.clone(),
..typed_request
})?;
serde_json::from_value(request.body).map_err(invalid_request)?;
let sanitized = shape_anthropic_messages_request(
AnthropicMessagesRequest {
model: model.clone(),
..typed_request
},
request.shaping.reasoning_auto_summary,
)?;
let trimmed =
without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?;
let transformed = config.transform_anthropic_messages_request(
trimmed,
&MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params),
)?;
let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider);
let forwarded = string_headers(Some(
request
.extra_headers
.into_iter()
.flatten()
.chain(scoped)
.collect(),
))?;
let authenticated = config.authenticate(forwarded, request.api_key, &env_lookup)?;
let headers = config.request_headers(
with_default_headers(authenticated, config.default_headers()),
&transformed,
);
let body = serde_json::to_value(transformed).map_err(|err| {
Error::InvalidRequest(format!(
"failed to serialize Anthropic messages request: {err}"
@ -65,33 +116,371 @@ pub(super) fn prepare_provider_request(
})
}
fn validate_environment(
config: &dyn BaseAnthropicMessagesConfig,
extra_headers: Option<Map<String, Value>>,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Vec<(String, String)>, Error> {
let mut headers = string_headers(extra_headers)?;
let auth_strategy = config.auth_strategy();
let already_authorized = has_header(&headers, auth_strategy.header_name())
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
if !already_authorized {
let api_key = config.resolve_api_key(api_key, env_lookup)?;
let auth_header = match auth_strategy {
MessagesAuthStrategy::Bearer => {
("authorization".to_string(), format!("Bearer {api_key}"))
}
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
};
headers.push(auth_header);
}
for (name, value) in config.default_headers() {
if !has_header(&headers, name) {
headers.push((name.to_string(), value.to_string()));
}
}
Ok(headers)
fn invalid_request(err: serde_json::Error) -> Error {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
}
fn without_additional_drop_params(
request: AnthropicMessagesRequest,
paths: &[String],
) -> Result<AnthropicMessagesRequest, Error> {
if paths.is_empty() {
return Ok(request);
}
let Value::Object(fields) = serde_json::to_value(request).map_err(invalid_request)? else {
return Err(Error::InvalidRequest(
"Anthropic messages request did not serialize to an object".to_string(),
));
};
let (required, optional): (Map<String, Value>, Map<String, Value>) = fields
.into_iter()
.partition(|(key, _)| matches!(key.as_str(), "model" | "messages"));
let trimmed = paths.iter().fold(Value::Object(optional), |body, path| {
delete_nested_value(body, path)
});
let merged: Map<String, Value> = required
.into_iter()
.chain(trimmed.as_object().cloned().unwrap_or_default())
.collect();
serde_json::from_value(Value::Object(merged)).map_err(invalid_request)
}
fn with_default_headers(
headers: Vec<(String, String)>,
defaults: &[(&str, &str)],
) -> Vec<(String, String)> {
let missing: Vec<(String, String)> = defaults
.iter()
.filter(|(name, _)| {
!headers
.iter()
.any(|(header, _)| header.eq_ignore_ascii_case(name))
})
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
.collect();
headers.into_iter().chain(missing).collect()
}
#[cfg(test)]
mod tests {
use litellm_types::utils::ProviderSpecificHeaders;
use rstest::{fixture, rstest};
use serde_json::json;
use super::*;
use crate::messages::types::MessagesShaping;
#[fixture]
fn shaping() -> MessagesShaping {
MessagesShaping::default()
}
fn prepare(request: MessagesRequest<'_>) -> Result<ProviderMessagesRequest, Error> {
prepare_with_secrets(request, &|_: &str| None)
}
fn prepare_with_secrets(
request: MessagesRequest<'_>,
secrets: &dyn Lookup,
) -> Result<ProviderMessagesRequest, Error> {
let resolved = resolve_provider(request.model, request.custom_llm_provider)?;
prepare_provider_request(request, resolved, secrets)
}
#[rstest]
#[case::api_key(
&[("ANTHROPIC_API_KEY", "sk-secret")],
&[("x-api-key", "sk-secret")],
"https://api.anthropic.com/v1/messages"
)]
#[case::auth_token(
&[("ANTHROPIC_AUTH_TOKEN", "token")],
&[("authorization", "Bearer token")],
"https://api.anthropic.com/v1/messages"
)]
#[case::api_base(
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_API_BASE", "https://gateway.test")],
&[("x-api-key", "sk-secret")],
"https://gateway.test/v1/messages"
)]
#[case::sdk_base_url(
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_BASE_URL", "https://sdk.test")],
&[("x-api-key", "sk-secret")],
"https://sdk.test/v1/messages"
)]
fn credentials_and_base_come_from_the_resolved_secrets(
shaping: MessagesShaping,
#[case] secrets: &[(&str, &str)],
#[case] expected_auth: &[(&str, &str)],
#[case] expected_url: &str,
) {
let lookup = |name: &str| {
secrets
.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| value.to_string())
};
let prepared = prepare_with_secrets(
MessagesRequest {
model: "claude-test",
body: json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
api_key: None,
api_base: None,
custom_llm_provider: Some("anthropic"),
extra_headers: None,
provider_specific_header: None,
timeout: None,
shaping,
},
&lookup,
)
.unwrap();
let auth: Vec<(&str, &str)> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization"))
.map(|(name, value)| (name.as_str(), value.as_str()))
.collect();
assert_eq!(
(auth.as_slice(), prepared.url.as_str()),
(expected_auth, expected_url)
);
}
fn prepared_body(body: Value, shaping: MessagesShaping) -> Result<Value, Error> {
prepare(MessagesRequest {
model: "anthropic/claude-test",
body,
api_key: Some("sk-test"),
api_base: Some("https://anthropic.test"),
custom_llm_provider: Some("anthropic"),
extra_headers: None,
provider_specific_header: None,
timeout: None,
shaping,
})
.map(|prepared| prepared.body)
}
#[rstest]
#[case::nothing_forwarded(
&[],
&[("x-version", "1"), ("content-type", "application/json")],
&[("x-version", "1"), ("content-type", "application/json")],
)]
#[case::forwarded_header_wins_in_any_case(
&[("X-Version", "custom"), ("x-api-key", "k")],
&[("x-version", "1"), ("content-type", "application/json")],
&[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")],
)]
#[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])]
fn default_headers_fill_only_missing_names(
#[case] forwarded: &[(&str, &str)],
#[case] defaults: &[(&str, &str)],
#[case] expected: &[(&str, &str)],
) {
let owned = |headers: &[(&str, &str)]| -> Vec<(String, String)> {
headers
.iter()
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
.collect()
};
assert_eq!(
with_default_headers(owned(forwarded), defaults),
owned(expected)
);
}
#[rstest]
#[case::top_level_and_nested_paths(
json!({
"max_tokens": 1024,
"thinking": {"type": "enabled", "budget_tokens": 2048},
"context_management": {"edits": [{"type": "clear_thinking_20251015"}]},
"metadata": {"user_id": "u1"},
"tools": [{"name": "lookup", "input_schema": {"type": "object"}, "input_examples": [{"q": "x"}]}]
}),
&["thinking", "context_management", "tools[*].input_examples"],
json!({
"max_tokens": 1024,
"metadata": {"user_id": "u1"},
"tools": [{"name": "lookup", "input_schema": {"type": "object"}}]
}),
)]
#[case::no_paths(
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
&[],
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
)]
#[case::model_and_messages_are_never_dropped(
json!({"max_tokens": 16}),
&["model", "messages", "messages[0].content"],
json!({"max_tokens": 16}),
)]
fn prepared_body_drops_configured_paths(
shaping: MessagesShaping,
#[case] fields: Value,
#[case] additional_drop_params: &[&str],
#[case] expected_fields: Value,
) {
let with_messages = |fields: Value| -> Value {
let Value::Object(fields) = fields else {
unreachable!()
};
Value::Object(
[
("model".to_string(), json!("claude-test")),
(
"messages".to_string(),
json!([{"role": "user", "content": "hi"}]),
),
]
.into_iter()
.chain(fields)
.collect(),
)
};
let shaping = MessagesShaping {
additional_drop_params: additional_drop_params
.iter()
.map(ToString::to_string)
.collect(),
..shaping
};
assert_eq!(
prepared_body(with_messages(fields), shaping),
Ok(with_messages(expected_fields))
);
}
#[rstest]
#[case::model_prefix_picks_the_provider(
"azure_ai/claude-test",
None,
&[("x-priority", "extra"), ("x-scoped", "azure_ai")]
)]
#[case::explicit_provider(
"claude-test",
Some("anthropic"),
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
)]
#[case::provider_prefix_on_an_anthropic_model(
"anthropic/claude-test",
None,
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
)]
fn provider_specific_headers_follow_the_resolved_provider(
shaping: MessagesShaping,
#[case] model: &str,
#[case] custom_llm_provider: Option<&str>,
#[case] expected: &[(&str, &str)],
) {
let configured: ProviderSpecificHeaders = serde_json::from_value(json!([
{"custom_llm_provider": "azure_ai", "extra_headers": {"x-scoped": "azure_ai"}},
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "anthropic", "x-priority": "scoped"}}
]))
.unwrap();
let prepared = prepare(MessagesRequest {
model,
body: json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
api_key: Some("sk-test"),
api_base: Some("https://resource.services.ai.azure.com"),
custom_llm_provider,
extra_headers: Some(serde_json::from_value(json!({"x-priority": "extra"})).unwrap()),
provider_specific_header: Some(configured),
timeout: None,
shaping,
})
.unwrap();
let caller_headers: Vec<(&str, &str)> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped"))
.map(|(name, value)| (name.as_str(), value.as_str()))
.collect();
assert_eq!(caller_headers, expected);
}
#[rstest]
fn prepared_body_carries_the_provider_stripped_model(shaping: MessagesShaping) {
assert_eq!(
prepared_body(
json!({
"model": "anthropic/claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 16
}),
shaping,
),
Ok(json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 16
}))
);
}
#[rstest]
fn dropped_thinking_display_is_not_restored_by_auto_summary(shaping: MessagesShaping) {
let shaping = MessagesShaping {
reasoning_auto_summary: true,
additional_drop_params: vec!["thinking.display".to_string()],
..shaping
};
assert_eq!(
prepared_body(
json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 4096,
"thinking": {"type": "enabled", "budget_tokens": 2048}
}),
shaping,
),
Ok(json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 4096,
"thinking": {"type": "enabled", "budget_tokens": 2048}
}))
);
}
#[rstest]
fn dropping_an_invalid_metadata_user_id_does_not_skip_its_validation(shaping: MessagesShaping) {
let shaping = MessagesShaping {
additional_drop_params: vec!["metadata.user_id".to_string()],
..shaping
};
assert!(matches!(
prepared_body(
json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 16,
"metadata": {"user_id": 123}
}),
shaping,
),
Err(Error::InvalidRequest(_))
));
}
#[rstest]
fn prepared_body_rejects_invalid_metadata_before_the_call(shaping: MessagesShaping) {
assert_eq!(
prepared_body(
json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 16,
"metadata": {"user_id": 123}
}),
shaping,
),
Err(Error::InvalidRequest(
"metadata.user_id must be a string, got 123".to_string()
))
);
}
}

View file

@ -1,4 +1,7 @@
use std::{sync::Mutex, time::Duration};
use std::{
sync::{Arc, Mutex},
time::Duration,
};
use bytes::Bytes;
use litellm_auth::SecretValue;
@ -9,15 +12,19 @@ use litellm_host::{
machine::{HostChannel, MachineFault, RouteMachine},
route::Route,
};
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use litellm_secrets::source::SecretSource;
use litellm_types::{
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
utils::ProviderSpecificHeaders,
};
use serde_json::{Map, Value};
use super::{
Error,
common_utils::messages_provider_config,
handler::{decode_response, network, provider_error, send},
prepare::prepare_provider_request,
types::MessagesRequest,
prepare::{prepare_provider_request, resolve_provider},
types::{MessagesRequest, MessagesShaping},
};
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
@ -38,7 +45,9 @@ pub struct MessagesCall {
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
pub provider_specific_header: Option<ProviderSpecificHeaders>,
pub timeout: Option<Duration>,
pub shaping: MessagesShaping,
}
impl MessagesCall {
@ -120,22 +129,33 @@ impl Host<Messages> for LocalMessagesHost {
}
}
pub fn messages_machine() -> MessagesMachine {
RouteMachine::new(|host| Box::pin(execute(host)))
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
}
async fn execute(host: MessagesHost) -> Result<MessagesOutput, Error> {
async fn execute(
host: MessagesHost,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesOutput, Error> {
let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?;
let stream = call.streams();
let request = prepare_provider_request(MessagesRequest {
model: &call.model,
body: Value::Object(call.body.clone()),
api_key: call.api_key.as_deref(),
api_base: call.api_base.as_deref(),
custom_llm_provider: call.custom_llm_provider.as_deref(),
extra_headers: call.extra_headers.clone(),
timeout: call.timeout,
})?;
let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
let secrets = secrets.resolve(resolved.config.secret_names()).await?;
let request = prepare_provider_request(
MessagesRequest {
model: &call.model,
body: Value::Object(call.body.clone()),
api_key: call.api_key.as_deref(),
api_base: call.api_base.as_deref(),
custom_llm_provider: call.custom_llm_provider.as_deref(),
extra_headers: call.extra_headers.clone(),
provider_specific_header: call.provider_specific_header.clone(),
timeout: call.timeout,
shaping: call.shaping.clone(),
},
resolved,
secrets.as_ref(),
)?;
if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER {
return Err(Error::Unsupported("streaming messages for this provider"));
}

View file

@ -1,8 +1,25 @@
use std::time::Duration;
use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig;
use litellm_llms::{
anthropic::common_utils::AnthropicModelCapabilities,
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
};
use litellm_types::utils::ProviderSpecificHeaders;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct MessagesShaping {
#[serde(default)]
pub capabilities: AnthropicModelCapabilities,
#[serde(default)]
pub drop_params: bool,
#[serde(default)]
pub reasoning_auto_summary: bool,
#[serde(default)]
pub additional_drop_params: Vec<String>,
}
pub struct MessagesRequest<'a> {
pub model: &'a str,
pub body: Value,
@ -10,7 +27,9 @@ pub struct MessagesRequest<'a> {
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub provider_specific_header: Option<ProviderSpecificHeaders>,
pub timeout: Option<Duration>,
pub shaping: MessagesShaping,
}
pub struct ProviderMessagesRequest {
@ -22,3 +41,86 @@ pub struct ProviderMessagesRequest {
pub upstream_headers: Vec<(String, String)>,
pub timeout: Option<Duration>,
}
#[cfg(test)]
mod tests {
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
use rstest::rstest;
use serde_json::json;
use super::*;
#[rstest]
#[case::nothing_projected(json!({}), MessagesShaping::default())]
#[case::only_drop_params(
json!({"drop_params": true}),
MessagesShaping { drop_params: true, ..MessagesShaping::default() },
)]
#[case::only_reasoning_auto_summary(
json!({"reasoning_auto_summary": true}),
MessagesShaping { reasoning_auto_summary: true, ..MessagesShaping::default() },
)]
#[case::only_additional_drop_params(
json!({"additional_drop_params": ["tools[*].input_examples"]}),
MessagesShaping {
additional_drop_params: vec!["tools[*].input_examples".to_string()],
..MessagesShaping::default()
},
)]
#[case::partial_capabilities(
json!({"capabilities": {"supports_reasoning": true}}),
MessagesShaping {
capabilities: AnthropicModelCapabilities {
supports_reasoning: true,
..AnthropicModelCapabilities::default()
},
..MessagesShaping::default()
},
)]
#[case::everything_the_python_host_projects(
json!({
"capabilities": {
"supports_reasoning": true,
"supports_adaptive_thinking": true,
"thinking_always_on": false,
"supports_legacy_thinking": false,
"supports_output_config": true,
"supports_sampling_params": false,
"supports_speed": true,
"effort_tiers": {"minimal": false, "low": true, "medium": true, "high": true, "xhigh": true, "max": false}
},
"drop_params": true,
"reasoning_auto_summary": true,
"additional_drop_params": ["metadata.user_id", "thinking"]
}),
MessagesShaping {
capabilities: AnthropicModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
thinking_always_on: false,
supports_legacy_thinking: false,
supports_output_config: true,
supports_sampling_params: false,
supports_speed: true,
effort_tiers: SupportedEffortTiers {
minimal: false,
low: true,
medium: true,
high: true,
xhigh: true,
max: false,
},
},
drop_params: true,
reasoning_auto_summary: true,
additional_drop_params: vec!["metadata.user_id".to_string(), "thinking".to_string()],
},
)]
fn shaping_deserializes_with_defaults_for_absent_fields(
#[case] projected: Value,
#[case] expected: MessagesShaping,
) {
let shaping: MessagesShaping = serde_json::from_value(projected).unwrap();
assert_eq!(shaping, expected);
}
}

View file

@ -248,3 +248,160 @@ mod tests {
}
}
}
#[cfg(test)]
mod document_tests {
use litellm_host::event::WireRequest;
use litellm_llms::base_llm::ocr::error::Error;
use rstest::rstest;
use serde_json::{Value, json};
use crate::ocr::route::LocalOcrHost;
use crate::ocr::test_support::{
MockResponse, SERVED_DOCUMENT, document_server, mock_server, perform_ocr_with,
request_body, wire_request_with_document,
};
#[derive(Clone, Copy, Debug)]
enum Route {
Mistral,
AzureAi,
VertexMistral,
AzureCohereParse,
Cohere,
}
impl Route {
fn model(self) -> &'static str {
match self {
Self::Mistral => "mistral/model",
Self::AzureAi => "azure_ai/model",
Self::VertexMistral => "vertex_ai/mistral-ocr-maas",
Self::AzureCohereParse => "azure_ai/cohere-parse",
Self::Cohere => "cohere/model",
}
}
fn document_type(self) -> &'static str {
match self {
Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url",
Self::AzureCohereParse | Self::Cohere => "image_url",
}
}
fn options(self) -> Value {
match self {
Self::Mistral | Self::AzureAi => json!({"pages": [0]}),
Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}),
Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}),
}
}
}
/// What the host does to the wire request in `before_send`.
#[derive(Clone, Copy, Debug)]
enum Host {
Detached,
ReplacesDocument,
}
const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ=";
impl Host {
fn before_send(self, wire: WireRequest) -> WireRequest {
let Value::Object(fields) = wire.body else {
return wire;
};
let body = fields
.into_iter()
.map(|(name, value)| match self {
Self::Detached => (name, value),
Self::ReplacesDocument if name == "document" => {
let document_type = value["type"].clone();
let key = document_type.as_str().unwrap_or_default().to_string();
(name, json!({"type": document_type, key: REPLACED_DOCUMENT}))
}
Self::ReplacesDocument => (name, value),
})
.collect();
WireRequest {
body: Value::Object(body),
..wire
}
}
}
struct Sent {
result: Result<(), Error>,
provider_body: Option<Value>,
}
async fn send(route: Route, host: Host, document_base: &str) -> Sent {
let (base, seen, provider) =
mock_server(vec![MockResponse::json(json!({"pages": []}))]).await;
let document_type = route.document_type();
let document =
json!({"type": document_type, document_type: format!("{document_base}/scan.png")});
let request = wire_request_with_document(route.model(), &base, document, route.options());
let local =
LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(host.before_send(wire)));
let result = perform_ocr_with(local).await.map(|_| ());
match result {
Ok(()) => provider.await.unwrap(),
Err(_) => provider.abort(),
}
let provider_body = seen
.lock()
.unwrap()
.first()
.map(|request| request_body(request));
Sent {
result,
provider_body,
}
}
fn served_document_uri() -> String {
use base64::Engine;
format!(
"data:image/png;base64,{}",
base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT)
)
}
#[rstest]
#[case::azure_ai(Route::AzureAi)]
#[case::vertex_mistral(Route::VertexMistral)]
#[case::azure_cohere_parse(Route::AzureCohereParse)]
#[tokio::test]
async fn inlining_routes_send_the_downloaded_document(#[case] route: Route) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::Detached, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(served_document_uri())
);
}
#[rstest]
#[tokio::test]
async fn document_replaced_by_the_host_reaches_the_provider(
#[values(
Route::Mistral,
Route::AzureAi,
Route::VertexMistral,
Route::AzureCohereParse,
Route::Cohere
)]
route: Route,
) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::ReplacesDocument, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(REPLACED_DOCUMENT)
);
}
}

View file

@ -9,36 +9,210 @@ pub mod types;
pub mod wire;
#[cfg(test)]
#[path = "../../tests/aws_textract_ocr.rs"]
mod aws_textract_tests;
pub(crate) mod test_support {
use std::sync::{Arc, Mutex};
#[cfg(test)]
#[path = "../../tests/azure_ai_ocr.rs"]
mod azure_ai_tests;
#[cfg(test)]
#[path = "../../tests/azure_document_intelligence_ocr.rs"]
mod azure_document_intelligence_tests;
#[cfg(test)]
#[path = "../../tests/cohere_ocr.rs"]
mod cohere_tests;
#[cfg(test)]
#[path = "../../tests/deepseek_ocr.rs"]
mod deepseek_tests;
#[cfg(test)]
#[path = "../../tests/ocr/document.rs"]
mod document_tests;
#[cfg(test)]
#[path = "../../tests/reducto_ocr.rs"]
mod reducto_tests;
#[cfg(test)]
#[path = "../../tests/ocr/support.rs"]
pub(crate) mod test_support;
#[cfg(test)]
#[path = "../../tests/ocr.rs"]
pub(crate) mod tests;
#[cfg(test)]
#[path = "../../tests/vertex_ai_deepseek_ocr.rs"]
mod vertex_ai_deepseek_tests;
#[cfg(test)]
#[path = "../../tests/vertex_ai_ocr.rs"]
mod vertex_ai_tests;
use futures_util::future::BoxFuture;
use litellm_host::event::WireRequest;
use litellm_llms::base_llm::ocr::{
error::Error,
handler::{CallHooks, OcrClient},
transformation::LiteLLMOcrResponse,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
};
use crate::ocr::{
route::{LocalOcrHost, ocr_machine},
types::LiteLLMOcrRequest,
wire::{OcrWireRequest, decode_request},
};
/// Stands in for a host with no hooks registered: the wire request goes out unchanged
/// and response events go nowhere.
pub(crate) struct NoHooks;
impl CallHooks<Error> for NoHooks {
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
Box::pin(async move { Ok(wire) })
}
fn response_received<'a>(&'a self, _body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
Box::pin(async { Ok(()) })
}
}
pub(crate) fn ocr_client() -> OcrClient {
let document_http = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test document client builds");
OcrClient::for_test(reqwest::Client::new(), document_http)
}
pub(crate) async fn perform_ocr(
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, Error> {
crate::ocr::client::perform(&ocr_client(), request).await
}
pub(crate) async fn perform_ocr_with(host: LocalOcrHost) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::run::run(ocr_machine(ocr_client()), &host).await
}
pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
wire_request_with_document(
model,
base,
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
options,
)
}
pub(crate) fn wire_request_with_document(
model: &str,
base: &str,
document: Value,
options: Value,
) -> LiteLLMOcrRequest {
decode_request(OcrWireRequest {
model: model.into(),
document,
api_key: Some(litellm_auth::SecretValue::new("test-key")),
api_base: Some(base.into()),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone(),
input_sources: Default::default(),
timeout_seconds: Some(2.0),
})
.unwrap()
}
pub(crate) fn resolved_request(
request: LiteLLMOcrRequest,
) -> crate::ocr::types::ResolvedOcrRequest {
request
.map_document(crate::ocr::document::prepare_document)
.unwrap()
}
pub(crate) fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest {
let request = resolved_request(request);
let document = request.document.clone().with_source(source.into());
request.with_document(document.into())
}
pub(crate) fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
pub(crate) const SERVED_DOCUMENT: &[u8] = b"\x89PNG served document";
/// Serves [`SERVED_DOCUMENT`] as `image/png` to every connection until aborted.
pub(crate) async fn document_server() -> (String, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let task = tokio::spawn(async move {
loop {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buffer = [0u8; 4096];
let _ = socket.read(&mut buffer).await.unwrap();
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Type: image/png\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
SERVED_DOCUMENT.len()
);
socket.write_all(head.as_bytes()).await.unwrap();
socket.write_all(SERVED_DOCUMENT).await.unwrap();
}
});
(base, task)
}
pub(crate) struct MockResponse {
pub status: u16,
pub headers: Vec<(&'static str, String)>,
pub body: Value,
}
impl MockResponse {
pub fn json(body: Value) -> Self {
Self {
status: 200,
headers: vec![],
body,
}
}
}
pub(crate) async fn mock_server(
responses: Vec<MockResponse>,
) -> (String, Arc<Mutex<Vec<String>>>, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let requests = Arc::new(Mutex::new(Vec::new()));
let seen = requests.clone();
let server_base = base.clone();
let task = tokio::spawn(async move {
for response in responses {
let (mut socket, _) = listener.accept().await.unwrap();
let mut bytes = Vec::new();
let mut buffer = [0u8; 4096];
let header_end = loop {
let n = socket.read(&mut buffer).await.unwrap();
assert!(n > 0);
bytes.extend_from_slice(&buffer[..n]);
if let Some(index) = bytes.windows(4).position(|s| s == b"\r\n\r\n") {
break index + 4;
}
};
let length = String::from_utf8_lossy(&bytes[..header_end])
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().unwrap())
})
.unwrap_or(0);
while bytes.len() < header_end + length {
let n = socket.read(&mut buffer).await.unwrap();
assert!(n > 0);
bytes.extend_from_slice(&buffer[..n]);
}
seen.lock()
.unwrap()
.push(String::from_utf8_lossy(&bytes).into_owned());
let body = serde_json::to_vec(&response.body).unwrap();
let headers = response
.headers
.into_iter()
.map(|(name, value)| {
format!("{name}: {}\r\n", value.replace("{base}", &server_base))
})
.collect::<String>();
let head = format!(
"HTTP/1.1 {} OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n{}\r\n",
response.status,
body.len(),
headers
);
socket.write_all(head.as_bytes()).await.unwrap();
socket.write_all(&body).await.unwrap();
}
});
(base, requests, task)
}
pub(crate) fn header<'a>(request: &'a str, name: &str) -> Option<&'a str> {
request
.lines()
.take_while(|line| !line.is_empty())
.find_map(|line| {
let (key, value) = line.split_once(':')?;
key.eq_ignore_ascii_case(name).then(|| value.trim())
})
}
}

File diff suppressed because it is too large Load diff

View file

@ -4,11 +4,9 @@ use std::{
thread,
};
use litellm_core::audio_transcription::{audio_transcription, types::AudioTranscriptionRequest};
use serde_json::{Map, json};
use super::audio_transcription;
use crate::audio_transcription::types::AudioTranscriptionRequest;
#[tokio::test]
async fn bedrock_request_is_signed_and_contains_audio() {
let listener = TcpListener::bind("127.0.0.1:0").expect("listener");

View file

@ -1,193 +0,0 @@
use std::{collections::BTreeMap, time::SystemTime};
use litellm_auth_aws::{Credentials, aws_signature_headers, sign_post};
use litellm_llms::base_llm::ocr::error::Error;
use serde_json::{Value, json};
use time::{PrimitiveDateTime, format_description};
use crate::ocr::{
route::LocalOcrHost,
test_support::{
MockResponse, header, mock_server, perform_ocr_with, request_body,
wire_request_with_document,
},
types::LiteLLMOcrRequest,
};
const ACCESS_KEY_ID: &str = "AKIDEXAMPLE";
const SECRET_ACCESS_KEY: &str = "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY";
fn textract_request(base: &str) -> LiteLLMOcrRequest {
textract_request_for("aws_textract/detect-document-text", base)
}
fn textract_request_for(model: &str, base: &str) -> LiteLLMOcrRequest {
wire_request_with_document(
model,
&format!("{base}/"),
json!({"type": "image_url", "image_url": "data:image/png;base64,b3JpZ2luYWw="}),
json!({
"aws_access_key_id": ACCESS_KEY_ID,
"aws_secret_access_key": SECRET_ACCESS_KEY,
"aws_region_name": "eu-west-1"
}),
)
}
fn textract_response() -> MockResponse {
MockResponse::json(json!({
"DocumentMetadata": {"Pages": 1},
"Blocks": [{"BlockType": "PAGE"}, {"BlockType": "LINE", "Text": "Invoice 12345"}]
}))
}
/// Recomputes SigV4 over the bytes the server received, at the time the client claimed.
fn expected_authorization(url: &str, raw_request: &str) -> String {
let format =
format_description::parse_borrowed::<2>("[year][month][day]T[hour][minute][second]Z")
.unwrap();
let signed_at: SystemTime =
PrimitiveDateTime::parse(header(raw_request, "x-amz-date").unwrap(), &format)
.unwrap()
.assume_utc()
.into();
let headers: BTreeMap<String, String> = ["content-type", "x-amz-target"]
.into_iter()
.map(|name| {
(
name.to_string(),
header(raw_request, name).unwrap().to_string(),
)
})
.collect();
let body = raw_request.split_once("\r\n\r\n").unwrap().1;
sign_post(
url,
body.as_bytes(),
&aws_signature_headers(&headers),
"eu-west-1",
"textract",
&Credentials::new(ACCESS_KEY_ID, SECRET_ACCESS_KEY, None, None, "test"),
signed_at,
)
.unwrap()["Authorization"]
.clone()
}
#[tokio::test]
async fn the_request_is_signed_for_textract_and_lines_become_the_page() {
let (base, seen, server) = mock_server(vec![textract_response()]).await;
let response = perform_ocr_with(LocalOcrHost::new(textract_request(&base)))
.await
.unwrap();
server.await.unwrap();
let raw = seen.lock().unwrap()[0].clone();
assert_eq!(
header(&raw, "x-amz-target"),
Some("Textract.DetectDocumentText")
);
assert_eq!(
header(&raw, "content-type"),
Some("application/x-amz-json-1.1")
);
assert_eq!(
request_body(&raw),
json!({"Document": {"Bytes": "b3JpZ2luYWw="}})
);
assert_eq!(
header(&raw, "authorization"),
Some(expected_authorization(&format!("{base}/"), &raw).as_str())
);
assert_eq!(response.pages[0].markdown, "Invoice 12345");
assert_eq!(response.usage_info.unwrap().pages_processed, Some(1));
}
#[tokio::test]
async fn a_body_rewritten_by_before_send_is_what_gets_signed_and_sent() {
let (base, seen, server) = mock_server(vec![textract_response()]).await;
let host = LocalOcrHost::new(textract_request(&base)).with_before_send(|mut wire, _| {
assert!(
!wire
.headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("authorization")),
"the hook ran after signing"
);
wire.body["Document"]["Bytes"] = Value::from("cmVkYWN0ZWQ=");
Ok(wire)
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
let raw = seen.lock().unwrap()[0].clone();
assert_eq!(
request_body(&raw),
json!({"Document": {"Bytes": "cmVkYWN0ZWQ="}})
);
assert_eq!(
header(&raw, "authorization"),
Some(expected_authorization(&format!("{base}/"), &raw).as_str())
);
}
#[tokio::test]
async fn a_multi_page_rejection_reaches_the_caller_with_the_single_page_limit() {
let (base, _, server) = mock_server(vec![MockResponse {
status: 400,
headers: vec![],
body: json!({
"__type": "UnsupportedDocumentException",
"Message": "Request has unsupported document format"
}),
}])
.await;
let error = perform_ocr_with(LocalOcrHost::new(textract_request(&base)))
.await
.unwrap_err();
server.await.unwrap();
let Error::Provider { status, body, .. } = error else {
panic!("expected a provider error, got {error:?}");
};
assert_eq!(status, 400);
assert!(
body.contains("multi-page documents are not supported"),
"{body}"
);
}
#[tokio::test]
async fn analyze_document_asks_for_layout_and_tables_and_returns_markdown() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"DocumentMetadata": {"Pages": 1},
"Blocks": [
{"Id": "l1", "BlockType": "LINE", "Text": "Quarterly Report"},
{"Id": "t", "BlockType": "LAYOUT_TITLE",
"Relationships": [{"Type": "CHILD", "Ids": ["l1"]}]}
]
}))])
.await;
let request = textract_request_for("aws_textract/analyze-document", &base);
let response = perform_ocr_with(LocalOcrHost::new(request)).await.unwrap();
server.await.unwrap();
let raw = seen.lock().unwrap()[0].clone();
assert_eq!(
header(&raw, "x-amz-target"),
Some("Textract.AnalyzeDocument")
);
assert_eq!(
request_body(&raw)["FeatureTypes"],
json!(["LAYOUT", "TABLES"])
);
assert_eq!(
header(&raw, "authorization"),
Some(expected_authorization(&format!("{base}/"), &raw).as_str())
);
assert_eq!(response.pages[0].markdown, "# Quarterly Report");
}

View file

@ -1,293 +0,0 @@
use litellm_llms::base_llm::ocr::error::Error;
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request};
use crate::ocr::route::LocalOcrHost;
#[tokio::test]
async fn facade_executes_azure_mistral_with_prepared_auth() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"include_image_base64":true}),
);
request.credentials.api_key = None;
request.transport.extra_headers = vec![(
"Authorization".into(),
"Bearer python-prepared-token".into(),
)];
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /providers/mistral/azure/ocr "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer python-prepared-token\r\n")
);
let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({
"model":"model",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"include_image_base64":true
})
);
}
#[tokio::test]
async fn facade_acquires_supplied_entra_token_for_final_request() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"azure_ad_token":"rust-owned-token"}),
);
request.credentials.api_key = None;
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer rust-owned-token\r\n")
);
}
#[tokio::test]
async fn rejects_non_inline_body_after_guardrails() {
let request = wire_request("azure_ai/model", "http://127.0.0.1:1", json!({}));
let host = LocalOcrHost::new(request).with_before_send(|mut wire, _| {
wire.body["document"] = json!({
"type":"document_url",
"document_url":"https://example.com/not-inline.pdf"
});
Ok(wire)
});
let error = perform_ocr_with(host).await.unwrap_err();
assert!(error.to_string().contains("data URI"));
}
mod transformation {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use litellm_auth::{
ResolvedCredential, SecretValue, TokenFuture, TokenProvider, TokenProviderHandle,
};
use rstest::rstest;
use serde_json::json;
use super::*;
use crate::ocr::{
test_support::{MockResponse, header, mock_server, perform_ocr},
types::LiteLLMOcrRequest,
wire::decode_request,
};
#[derive(Debug)]
struct CountingToken {
token: fn(usize) -> String,
calls: AtomicUsize,
}
impl CountingToken {
fn new(token: fn(usize) -> String) -> Arc<Self> {
Arc::new(Self {
token,
calls: AtomicUsize::new(0),
})
}
fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
}
impl TokenProvider for CountingToken {
fn acquire(&self) -> TokenFuture<'_> {
let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
let token = SecretValue::new((self.token)(call));
Box::pin(async move {
Ok(ResolvedCredential::AccessToken {
token,
expires_on: None,
})
})
}
}
fn numbered_token(call: usize) -> String {
format!("callback-{call}")
}
fn azure_request(
provider: &Arc<CountingToken>,
api_base: Option<&str>,
api_key: Option<&str>,
extra_headers: Value,
optional_params: Value,
) -> LiteLLMOcrRequest {
let wire = serde_json::from_value(json!({
"model": "azure_ai/mistral-ocr-latest",
"document": {"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": null,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": 2.0
}))
.unwrap();
LiteLLMOcrRequest {
azure_ad_token_provider: Some(TokenProviderHandle::new(provider.clone())),
..decode_request(wire).unwrap()
}
}
fn ocr_page() -> MockResponse {
MockResponse::json(json!({"pages":[{"index":0,"markdown":"hello"}]}))
}
#[tokio::test]
async fn token_provider_result_is_the_bearer_and_is_acquired_for_each_request() {
let provider = CountingToken::new(numbered_token);
let (base, seen, server) = mock_server(vec![ocr_page(), ocr_page()]).await;
for _ in 0..2 {
perform_ocr(azure_request(
&provider,
Some(&base),
None,
Value::Null,
json!({}),
))
.await
.unwrap();
}
server.await.unwrap();
assert_eq!(provider.calls(), 2);
let requests = seen.lock().unwrap();
assert_eq!(
requests
.iter()
.map(|request| header(request, "authorization"))
.collect::<Vec<_>>(),
[Some("Bearer callback-1"), Some("Bearer callback-2")]
);
}
#[rstest]
#[case::api_key_skips_provider(Some("resource-key"), Value::Null, json!({}), "Bearer resource-key", 0)]
#[case::provider_beats_static_token(
None,
Value::Null,
json!({"azure_ad_token":"static-token"}),
"Bearer callback-1",
1
)]
#[case::header_wins_on_the_wire_but_provider_still_runs(
None,
json!({"Authorization":"Bearer override"}),
json!({}),
"Bearer override",
1
)]
#[tokio::test]
async fn credential_precedence(
#[case] api_key: Option<&str>,
#[case] extra_headers: Value,
#[case] optional_params: Value,
#[case] expected_authorization: &str,
#[case] expected_calls: usize,
) {
let provider = CountingToken::new(numbered_token);
let (base, seen, server) = mock_server(vec![ocr_page()]).await;
perform_ocr(azure_request(
&provider,
Some(&base),
api_key,
extra_headers,
optional_params,
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(provider.calls(), expected_calls);
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert_eq!(
header(&requests[0], "authorization"),
Some(expected_authorization)
);
}
#[rstest]
#[case::missing_api_base(
false,
json!({}),
numbered_token,
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::MissingApiBase {
provider: "Azure AI",
environment_variable: "AZURE_AI_API_BASE",
})),
0
)]
#[case::unsupported_oidc_reference(
true,
json!({"azure_ad_token":"oidc/assertion","client_id":"client","tenant_id":"tenant"}),
numbered_token,
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::UnsupportedOidcReference)),
0
)]
#[case::empty_provider_token_ignores_static_token(
true,
json!({"azure_ad_token":"static-token"}),
|_| String::new(),
|error: &Error| matches!(error, Error::MissingAzureAiCredentials),
1
)]
#[tokio::test]
async fn credential_failures_send_no_provider_request(
#[case] with_api_base: bool,
#[case] optional_params: Value,
#[case] token: fn(usize) -> String,
#[case] expected: fn(&Error) -> bool,
#[case] expected_calls: usize,
) {
let provider = CountingToken::new(token);
let (base, seen, server) = mock_server(vec![ocr_page()]).await;
let error = perform_ocr(azure_request(
&provider,
with_api_base.then_some(base.as_str()),
None,
Value::Null,
optional_params,
))
.await
.unwrap_err();
server.abort();
assert!(expected(&error), "unexpected error: {error:?}");
assert_eq!(provider.calls(), expected_calls);
assert!(seen.lock().unwrap().is_empty());
}
}

View file

@ -1,712 +0,0 @@
use litellm_host::event::{CallEvent, MachineEvent};
use litellm_llms::base_llm::ocr::{error::Error, settings::OcrSettings};
use rstest::rstest;
use serde_json::{Value, json};
use super::{
test_support::{
MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request,
},
wire::{OcrWireRequest, decode_request},
};
use crate::ocr::route::LocalOcrHost;
fn query_value(url: &str, key: &str) -> Option<String> {
url::Url::parse(url)
.unwrap()
.query_pairs()
.find_map(|(name, value)| (name == key).then(|| value.into_owned()))
}
#[tokio::test]
async fn facade_maps_pages_features_and_url_document() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":{"pages":[]}
}))])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"]}),
);
request.document =
serde_json::from_value::<litellm_llms::base_llm::ocr::transformation::OcrDocument>(json!({
"type":"document_url",
"document_url":"https://example.com/document.pdf"
}))
.unwrap()
.into();
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let target = request.split_whitespace().nth(1).unwrap();
let url = format!("{base}{target}");
assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3"));
assert_eq!(
query_value(&url, "features").as_deref(),
Some("keyValuePairs,languages")
);
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({"urlSource":"https://example.com/document.pdf"})
);
}
#[rstest]
#[case(json!({"pages":[true]}), Error::Pages("expected only integers or only strings".into()))]
#[case(json!({"pages":[1,"2"]}), Error::Pages("expected only integers or only strings".into()))]
#[case(json!({"pages":[-1]}), Error::Pages("negative page index".into()))]
#[case(json!({"pages":"1&&features=bad"}), Error::Pages("invalid native page range".into()))]
#[case(json!({"features":"languages&pages=1"}), Error::Features)]
#[case(json!({"req_format":"azure"}), Error::RequestFormat)]
#[tokio::test]
async fn rejects_invalid_pages_features_and_format(
#[case] options: Value,
#[case] expected: Error,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({}))]).await;
let result = decode_request(OcrWireRequest {
model: "azure_ai/doc-intelligence/prebuilt-read".into(),
document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
api_key: Some(litellm_auth::SecretValue::new("key")),
api_base: Some(base),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone(),
input_sources: Default::default(),
timeout_seconds: Some(2.0),
});
let result = match result {
Ok(request) => perform_ocr(request).await,
Err(error) => Err(error),
};
server.abort();
let _ = server.await;
assert!(
seen.lock().unwrap().is_empty(),
"sent invalid options: {options}"
);
let error = result.unwrap_err();
assert_eq!(
std::mem::discriminant(&error),
std::mem::discriminant(&expected)
);
assert_eq!(error.http_status_code(), Some(400));
assert_eq!(error.to_string(), expected.to_string());
}
#[rstest]
#[case(json!({}))]
#[case(json!({"req_format":"litellm"}))]
#[tokio::test]
async fn missing_native_fields_keep_page_text_without_retaining_raw_response(
#[case] options: Value,
) {
let operation = json!({
"status":"succeeded",
"analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"hello"}]}]}
});
let (base, seen, server) = mock_server(vec![MockResponse::json(operation)]).await;
let response = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
options,
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(response.pages.len(), 1);
assert_eq!(response.pages[0].index, 0);
assert_eq!(response.pages[0].markdown, "hello");
assert_eq!(response.provider_native_response, None);
let serialized = response.into_json();
assert_eq!(serialized.get("content"), Some(&Value::Null));
assert_eq!(serialized.get("tables"), Some(&Value::Null));
assert_eq!(serialized.get("keyValuePairs"), Some(&Value::Null));
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
let target = requests[0].split_whitespace().nth(1).unwrap();
let url = format!("{base}{target}");
for field in ["pages", "features", "req_format"] {
assert_eq!(query_value(&url, field), None);
}
let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(body, json!({"base64Source":"YWJj"}));
}
#[tokio::test]
async fn inline_document_decodes_to_base64_source() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(body, json!({"base64Source":"YWJj"}));
}
#[tokio::test]
async fn immediate_response_normalizes_pages_and_preserves_native() {
let operation = json!({
"status":"succeeded",
"operationExtension":42,
"analyzeResult":{
"content":"A\n\nB",
"tables":[{"cells":[]}],
"keyValuePairs":[{"key":{"content":"A"}}],
"pages":[{
"pageNumber":"2",
"width":"8.5",
"height":11,
"unit":"inch",
"lines":[{"content":"A"},{"content":null},{"content":"B"}]
}]
}
});
let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await;
let result = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].index, 1);
assert_eq!(result.pages[0].markdown, "A\n\nB");
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width":816,"height":1056,"dpi":96})
);
assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1));
let serialized = result.clone().into_json();
assert_eq!(serialized["content"], "A\n\nB");
assert_eq!(serialized["tables"], json!([{"cells":[]}]));
assert_eq!(
serialized["keyValuePairs"],
json!([{"key":{"content":"A"}}])
);
assert!(serialized.get("key_value_pairs").is_none());
assert_eq!(
result.provider_native_response.map(Value::Object),
Some(operation)
);
}
#[tokio::test]
async fn client_settings_choose_the_api_version_and_the_inch_to_pixel_dpi() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":{"pages":[{"pageNumber":1,"width":8.5,"height":11,"unit":"inch"}]}
}))])
.await;
let client = ocr_client().with_settings(OcrSettings {
document_intelligence_api_version: "2099-01-01".into(),
document_intelligence_dpi: 72,
..OcrSettings::default()
});
let result = crate::ocr::client::perform(
&client,
wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})),
)
.await
.unwrap();
server.await.unwrap();
let target = seen.lock().unwrap()[0]
.split_whitespace()
.nth(1)
.unwrap()
.to_string();
assert_eq!(
query_value(&format!("{base}{target}"), "api-version").as_deref(),
Some("2099-01-01")
);
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width":612,"height":792,"dpi":72})
);
}
#[tokio::test]
async fn accepted_response_polls_to_success_with_only_credentials() {
let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}});
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "0".into())],
body: json!({"status":"running"}),
},
MockResponse::json(operation.clone()),
])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
);
request
.transport
.extra_headers
.push(("X-Trace".into(), "initial-only".into()));
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(
result.provider_native_response.map(Value::Object),
Some(operation)
);
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 3);
assert!(requests[0].to_ascii_lowercase().contains("x-trace:"));
for poll in &requests[1..] {
assert!(!poll.to_ascii_lowercase().contains("x-trace:"));
assert!(
poll.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: test-key")
);
}
}
#[tokio::test]
async fn accepted_response_emits_response_received_before_polling() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({"submitted": true}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let request_count = seen.clone();
let host = LocalOcrHost::new(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.with_observer(move |event| {
let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event else {
return;
};
match request_count.lock().unwrap().len() {
1 => assert_eq!(raw.body, r#"{"submitted":true}"#),
2 => assert!(raw.body.contains("succeeded")),
count => panic!("unexpected callback after {count} requests"),
}
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[tokio::test]
async fn polling_forwards_bearer_credentials() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.credentials.api_key = None;
request.transport.extra_headers = vec![("Authorization".into(), "Bearer token".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert!(
requests[1]
.to_ascii_lowercase()
.contains("authorization: bearer token")
);
}
#[tokio::test]
async fn polling_does_not_follow_redirects() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 302,
headers: vec![("Location", "{base}/redirected".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
assert!(error.to_string().contains("status 302"), "{error}");
assert_eq!(seen.lock().unwrap().len(), 2);
server.abort();
}
#[tokio::test]
async fn polling_rejects_terminal_failure() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"failed"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("status failed"));
}
#[tokio::test]
async fn malformed_provider_pages_report_response_paths() {
for (analysis, path) in [
(json!({"pages":null}), "pages"),
(json!({"pages":[null]}), "pages[0]"),
(json!({"pages":[{"lines":null}]}), "lines"),
(json!({"pages":[{"width":"bad"}]}), "width"),
] {
let (base, _, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":analysis
}))])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains(path), "{error}");
}
}
#[tokio::test]
async fn rejects_missing_invalid_and_cross_origin_operation_locations() {
for headers in [
Vec::new(),
vec![("Operation-Location", "/relative".into())],
vec![("Operation-Location", "http://example.com/operation".into())],
vec![(
"Operation-Location",
"http://user:password@127.0.0.1/operation".into(),
)],
] {
let (base, _, server) = mock_server(vec![MockResponse {
status: 202,
headers,
body: json!({}),
}])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("operation-location"));
}
}
#[tokio::test]
async fn polling_deadline_bounds_retry_delay() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "9999".into())],
body: json!({"status":"notStarted"}),
},
])
.await;
let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
let client = ocr_client().with_settings(OcrSettings {
poll_timeout: std::time::Duration::from_millis(100),
..OcrSettings::default()
});
let error = tokio::time::timeout(
std::time::Duration::from_secs(1),
crate::ocr::client::perform(&client, request),
)
.await
.unwrap()
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("timed out"));
}
#[tokio::test]
async fn model_id_is_encoded_and_dot_segments_are_rejected() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
perform_ocr(wire_request(
"azure_ai/doc-intelligence/a ?#é",
&base,
json!({}),
))
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("a%20%3F%23%C3%A9:analyze"));
for model in [
"azure_ai/doc-intelligence/.",
"azure_ai/doc-intelligence/..",
] {
let error = perform_ocr(wire_request(model, "http://127.0.0.1:1", json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains("dot segment"));
}
}
mod transformation {
use std::sync::{Arc, Mutex};
use litellm_host::event::{CallEvent, MachineEvent};
use litellm_llms::base_llm::ocr::transformation::OcrDocument;
use serde_json::{Value, json};
use super::*;
use crate::ocr::{
route::LocalOcrHost,
test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request},
};
#[tokio::test]
async fn facade_maps_pages_features_and_url_document() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":{"pages":[]}
}))])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"], "future_option": {"nested":null}, "extra_body":{"provider_option":false}}),
);
request.document = serde_json::from_value::<OcrDocument>(json!({
"type":"document_url",
"document_url":"https://example.com/document.pdf"
}))
.unwrap()
.into();
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let target = request.split_whitespace().nth(1).unwrap();
let url = format!("{base}{target}");
assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3"));
assert_eq!(
query_value(&url, "features").as_deref(),
Some("keyValuePairs,languages")
);
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({"urlSource":"https://example.com/document.pdf", "future_option":{"nested":null}, "provider_option":false})
);
}
#[tokio::test]
async fn rejects_invalid_pages_features_and_format() {
for options in [
json!({"pages":[true]}),
json!({"pages":[1,"2"]}),
json!({"pages":[-1]}),
json!({"pages":"1&&features=bad"}),
json!({"features":"languages&pages=1"}),
json!({"req_format":"azure"}),
] {
let request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
"http://127.0.0.1:1",
options.clone(),
);
let rejected = perform_ocr(request).await.is_err();
assert!(rejected, "accepted {options}");
}
}
#[tokio::test]
async fn immediate_response_normalizes_pages_and_preserves_native() {
let operation = json!({
"status":"succeeded",
"operationExtension":42,
"analyzeResult":{
"content":"A\n\nB",
"tables":[{"cells":[]}],
"keyValuePairs":[{"key":{"content":"A"}}],
"pages":[{
"pageNumber":"2",
"width":"8.5",
"height":11,
"unit":"inch",
"lines":[{"content":"A"},{"content":null},{"content":"B"}]
}]
}
});
let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await;
let result = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].index, 1);
assert_eq!(result.pages[0].markdown, "A\n\nB");
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width":816,"height":1056,"dpi":96})
);
assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1));
let serialized = result.clone().into_json();
assert_eq!(serialized["content"], "A\n\nB");
assert_eq!(serialized["tables"], json!([{"cells":[]}]));
assert_eq!(
serialized["keyValuePairs"],
json!([{"key":{"content":"A"}}])
);
assert!(serialized.get("key_value_pairs").is_none());
assert_eq!(
result.provider_native_response.as_ref(),
operation.as_object()
);
}
#[tokio::test]
async fn accepted_response_polls_to_success_with_only_credentials() {
let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}});
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "0".into())],
body: json!({"status":"running"}),
},
MockResponse::json(operation.clone()),
])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
);
request
.transport
.extra_headers
.push(("X-Trace".into(), "initial-only".into()));
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(
result.provider_native_response.as_ref(),
operation.as_object()
);
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 3);
assert!(requests[0].to_ascii_lowercase().contains("x-trace:"));
for poll in &requests[1..] {
assert!(!poll.to_ascii_lowercase().contains("x-trace:"));
assert!(
poll.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: test-key")
);
}
}
#[tokio::test]
async fn accepted_response_emits_response_received_for_submission_and_completed_poll() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({"submitted": true}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let responses_received = Arc::new(Mutex::new(Vec::new()));
let request_count = seen.clone();
let observed = responses_received.clone();
let host = LocalOcrHost::new(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
observed
.lock()
.unwrap()
.push((request_count.lock().unwrap().len(), raw.body.clone()));
}
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
assert_eq!(
*responses_received.lock().unwrap(),
[
(1, r#"{"submitted":true}"#.to_string()),
(2, r#"{"status":"succeeded"}"#.to_string()),
]
);
}
}

View file

@ -1,136 +0,0 @@
mod transformation {
use litellm_llms::{
base_llm::ocr::{
error::Error,
transformation::{BaseOcrConfig, OcrDocument, OcrResponseFormat},
},
cohere::ocr::transformation::*,
};
use rstest::rstest;
use serde_json::{Value, json};
#[tokio::test]
async fn composed_body_preserves_native_document_fields_and_untyped_overrides() {
let request = crate::ocr::test_support::wire_request(
"cohere/parse",
"https://example.com",
json!({
"output_format":"markdown", "timeout":30,
"extra_body":{
"output_format": {"future":true},
"document":{"type":"image_url","image_url":"https://example.com/a.png",
"provider_options":{"nested":[false,0,null]}}
}
}),
);
let request = request.with_document(
serde_json::from_value(json!({
"type":"image_url","image_url":"https://example.com/original.png"
}))
.unwrap(),
);
let request = crate::ocr::prepare::prepare_request_for_test(request);
let http = CohereParseConfig
.prepare_request(
&request,
&crate::ocr::test_support::ocr_client(),
&crate::ocr::test_support::NoHooks,
)
.await
.unwrap();
let body: Value = serde_json::from_slice(http.body()).unwrap();
assert_eq!(
body,
json!({
"model":"parse", "output_format":{"future":true},
"document":{"type":"image_url","image_url":"https://example.com/a.png",
"provider_options":{"nested":[false,0,null]}}
})
);
}
#[tokio::test]
async fn explicit_null_options_use_defaults_before_http() {
let request = crate::ocr::test_support::wire_request(
"cohere/parse",
"https://example.com",
json!({"output_format":null,"req_format":null}),
);
let request = request.with_document(
serde_json::from_value(
json!({"type":"image_url","image_url":"https://example.com/a.png"}),
)
.unwrap(),
);
assert_eq!(
request.response_format().unwrap(),
OcrResponseFormat::Litellm
);
let request = crate::ocr::prepare::prepare_request_for_test(request);
let http = CohereParseConfig
.prepare_request(
&request,
&crate::ocr::test_support::ocr_client(),
&crate::ocr::test_support::NoHooks,
)
.await
.unwrap();
let body: Value = serde_json::from_slice(http.body()).unwrap();
assert_eq!(body["output_format"], "markdown");
assert!(body.get("req_format").is_none());
}
#[rstest]
#[case::cohere("cohere/parse-v5.0", "POST /v2/parse ")]
#[case::azure_ai("azure_ai/Cohere-parse-v5.0", "POST /providers/cohere/v2/parse ")]
#[tokio::test]
async fn route_sends_image_to_its_parse_endpoint_with_the_bearer_key(
#[case] model: &str,
#[case] request_line: &str,
) {
use crate::ocr::test_support::{MockResponse, header, mock_server, perform_ocr};
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let request = crate::ocr::test_support::wire_request(model, &base, json!({}))
.with_document(
serde_json::from_value::<OcrDocument>(
json!({"type":"image_url","image_url":"data:image/png;base64,YWJj"}),
)
.unwrap()
.into(),
);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with(request_line), "{}", requests[0]);
assert_eq!(
header(&requests[0], "authorization"),
Some("Bearer test-key")
);
}
#[rstest]
#[tokio::test]
async fn route_rejects_non_image_document_without_a_request(
#[values("cohere/parse-v5.0", "azure_ai/Cohere-parse-v5.0")] model: &str,
) {
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr};
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let error = perform_ocr(crate::ocr::test_support::wire_request(
model,
&base,
json!({}),
))
.await
.unwrap_err();
server.abort();
assert!(matches!(error, Error::CohereImageOnly), "{error:?}");
assert!(seen.lock().unwrap().is_empty());
}
}

View file

@ -1,133 +0,0 @@
use litellm_llms::{
base_llm::ocr::transformation::{BaseOcrConfig, OcrDocument},
vertex_ai::ocr::deepseek_transformation::{
DeepSeekOcrParams, DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig,
normalize_response as transform_ocr_response,
},
};
use rstest::rstest;
use serde_json::{Value, json};
fn document() -> OcrDocument {
serde_json::from_value(json!({"type":"image_url","image_url":"gs://bucket/a.png"})).unwrap()
}
#[rstest]
#[case("stream", json!(true))]
#[case("temperature", json!(0.1))]
#[case("max_tokens", json!(1024))]
#[case("top_p", json!(0.9))]
#[case("n", json!(2))]
#[case("stop", json!("done"))]
#[case("stop", json!(["done", "stop"]))]
fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: DeepSeekOcrParams =
serde_json::from_value(json!({name: value.clone(), "ignored": true})).unwrap();
let result = serde_json::to_value(
VertexAIDeepSeekOCRConfig
.transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), &params, &[])
.unwrap(),
)
.unwrap();
assert_eq!(result["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/a.png"})
);
assert_eq!(result[name], value);
assert!(result.get("ignored").is_none());
}
#[rstest]
#[case(json!({"type":"image_url","image_url":"data:image/png;base64,AA=="}))]
#[case(json!({"type":"document_url","document_url":"data:application/pdf;base64,AA=="}))]
fn request_maps_both_document_types_to_image_content(#[case] document: Value) {
let source = document
.get("image_url")
.or_else(|| document.get("document_url"))
.unwrap()
.clone();
let request = VertexAIDeepSeekOCRConfig
.transform_ocr_request(
"deepseek-ai/deepseek-ocr-maas",
serde_json::from_value(document).unwrap(),
&DeepSeekOcrParams::default(),
&[],
)
.unwrap();
let result = serde_json::to_value(request).unwrap();
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":source})
);
}
#[rstest]
#[case(json!("# hello"), "# hello")]
#[case(json!("{broken"), "{broken")]
#[case(json!(" {\"pages\":[]} "), " {\"pages\":[]} ")]
#[case(json!({"pages":[]}), "")]
#[case(json!("[]"), "[]")]
#[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")]
#[case(json!({"pages":[{"markdown":"object"}]}), "object")]
fn response_codec_handles_text_json_and_objects(#[case] content: Value, #[case] expected: &str) {
let structured = content
.as_object()
.is_some_and(|object| object.contains_key("pages"))
|| content
.as_str()
.is_some_and(|text| text.contains("\"pages\""));
let response: DeepSeekOcrResponse = serde_json::from_value(
json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}),
)
.unwrap();
let result = transform_ocr_response("model", response)
.unwrap()
.into_json();
assert_eq!(result["pages"][0]["markdown"], expected);
assert_eq!(result["pages"][0]["index"], 0);
if structured {
assert!(result["usage_info"].is_null());
} else {
assert_eq!(result["usage_info"]["prompt_tokens"], 1);
}
}
#[test]
fn structured_result_maps_pages_usage_model_and_annotation() {
let response: DeepSeekOcrResponse = serde_json::from_value(json!({
"choices":[{"message":{"content":{
"pages":[{"index":2,"markdown":"page","images":[{"id":"one"}],"dimensions":{"width":10}}],
"model":"provider-model",
"usage_info":{"pages_processed":1},
"document_annotation":{"language":"en"},
"future":"kept"
}}}]
}))
.unwrap();
let result = transform_ocr_response("requested", response)
.unwrap()
.into_json();
assert_eq!(result["pages"][0]["index"], 2);
assert_eq!(result["pages"][0]["images"][0]["id"], "one");
assert_eq!(result["model"], "provider-model");
assert_eq!(result["usage_info"]["pages_processed"], 1);
assert_eq!(result["document_annotation"]["language"], "en");
assert_eq!(result["future"], "kept");
}
#[test]
fn response_codec_rejects_missing_empty_and_malformed_content() {
for value in [
json!({"choices":[{"message":{"content":{}}}]}),
json!({"choices":[]}),
json!({"choices":[{"message":{"content":""}}]}),
json!({"choices":[{"message":{"content":"{\"pages\":[{\"markdown\":42}]}"}}]}),
json!({"choices":[{"message":{"content":{"pages":[{"markdown":42}]}}}]}),
] {
let result = serde_json::from_value::<DeepSeekOcrResponse>(value)
.map_err(|_| ())
.and_then(|response| transform_ocr_response("model", response).map_err(|_| ()));
assert!(result.is_err());
}
}

View file

@ -1,19 +1,89 @@
use std::time::Duration;
use std::{sync::Arc, time::Duration};
use futures_util::future::BoxFuture;
use litellm_core::messages::{
Error, messages,
route::{LocalMessagesHost, MessagesCall, messages_machine},
types::{MessagesRequest, MessagesShaping},
};
use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::{Map, Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
use super::{
Error,
common_utils::{
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
},
messages,
};
use crate::messages::types::MessagesRequest;
struct RecordingSecrets {
values: Vec<(&'static str, String)>,
fails: bool,
requested: std::sync::Mutex<Vec<String>>,
}
impl RecordingSecrets {
fn new(values: Vec<(&'static str, String)>, fails: bool) -> Self {
Self {
values,
fails,
requested: std::sync::Mutex::new(Vec::new()),
}
}
}
impl SecretSource for RecordingSecrets {
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
Box::pin(async move {
self.requested.lock().unwrap().push(name.to_string());
if self.fails {
return Err(litellm_secrets::Error::ManagedSecretMissing);
}
Ok(self
.values
.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| SecretValue::new(value.clone())))
})
}
}
fn secrets_call() -> MessagesCall {
let Value::Object(body) = json!({
"model": "claude-sonnet-4-5",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hi"}]
}) else {
unreachable!("literal object")
};
MessagesCall {
model: "claude-sonnet-4-5".into(),
body,
api_key: None,
api_base: None,
custom_llm_provider: Some("anthropic".into()),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
}
}
#[tokio::test]
async fn route_surfaces_a_secret_manager_failure_before_the_call() {
let Err(error) = litellm_host::run::run(
messages_machine(Arc::new(RecordingSecrets::new(Vec::new(), true))),
&LocalMessagesHost::new(secrets_call()),
)
.await
else {
panic!("a secret manager failure fails the call");
};
assert!(
matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)),
"{error:?}"
);
}
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
@ -56,75 +126,6 @@ fn write_response(body: &str) -> String {
)
}
#[test]
fn provider_config_resolves_anthropic_and_azure_ai() {
assert!(messages_provider_config("anthropic").is_some());
assert!(messages_provider_config("azure_ai").is_some());
assert!(messages_provider_config("openai").is_none());
}
#[test]
fn truncate_error_body_caps_long_payloads() {
let body = "x".repeat(400);
let truncated = truncate_error_body(&body);
assert!(truncated.ends_with("... (truncated)"));
let prefix_chars = truncated
.strip_suffix("... (truncated)")
.expect("truncated marker present")
.chars()
.count();
assert_eq!(prefix_chars, 256);
}
#[test]
fn string_headers_rejects_non_string_values() {
let headers = json!({"x-count": 3}).as_object().unwrap().clone();
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert_eq!(
err,
Error::Headers(litellm_http::request::HeaderError {
context: "messages",
name: "x-count".to_string(),
actual: "number",
})
);
}
#[test]
fn has_header_is_case_insensitive() {
let headers = vec![("X-Api-Key".to_string(), "secret".to_string())];
assert!(has_header(&headers, "x-api-key"));
assert!(!has_header(&headers, "authorization"));
}
#[test]
fn has_bearer_auth_requires_a_nonempty_bearer_token() {
assert!(has_bearer_auth(&[(
"Authorization".to_string(),
"Bearer tok".to_string()
)]));
assert!(has_bearer_auth(&[(
"authorization".to_string(),
"bearer tok".to_string()
)]));
assert!(!has_bearer_auth(&[(
"authorization".to_string(),
"Bearer ".to_string()
)]));
assert!(!has_bearer_auth(&[(
"authorization".to_string(),
String::new()
)]));
assert!(!has_bearer_auth(&[(
"authorization".to_string(),
"Basic abc".to_string()
)]));
assert!(!has_bearer_auth(&[(
"x-api-key".to_string(),
"sk".to_string()
)]));
}
#[tokio::test]
async fn messages_round_trip_builds_azure_request_and_passes_response_through() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
@ -159,7 +160,9 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("messages request succeeds");
@ -215,7 +218,9 @@ async fn messages_round_trip_builds_native_anthropic_request() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("anthropic"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("messages request succeeds");
@ -268,7 +273,9 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: Some(headers),
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("messages request succeeds");
@ -322,7 +329,9 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: Some(headers),
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("entra id request succeeds without api key");
@ -346,7 +355,9 @@ async fn messages_requires_auth_when_no_key_and_no_header() {
api_base: Some("http://127.0.0.1:1"),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_millis(50)),
shaping: MessagesShaping::default(),
})
.await
.expect_err("missing auth errors");
@ -384,7 +395,9 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: Some(headers),
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("falls back to api key");
@ -425,7 +438,9 @@ async fn messages_maps_provider_error_status_to_http_error() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect_err("provider error propagates");
@ -445,7 +460,9 @@ async fn messages_rejects_unsupported_provider() {
api_base: Some("http://127.0.0.1:1"),
custom_llm_provider: Some("openai"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_millis(50)),
shaping: MessagesShaping::default(),
})
.await
.expect_err("unsupported provider errors");

File diff suppressed because it is too large Load diff

View file

@ -1,152 +0,0 @@
use litellm_host::event::WireRequest;
use litellm_llms::base_llm::ocr::error::Error;
use rstest::rstest;
use serde_json::{Value, json};
use super::test_support::{
MockResponse, SERVED_DOCUMENT, document_server, mock_server, perform_ocr_with, request_body,
wire_request_with_document,
};
use crate::ocr::route::LocalOcrHost;
#[derive(Clone, Copy, Debug)]
enum Route {
Mistral,
AzureAi,
VertexMistral,
AzureCohereParse,
Cohere,
}
impl Route {
fn model(self) -> &'static str {
match self {
Self::Mistral => "mistral/model",
Self::AzureAi => "azure_ai/model",
Self::VertexMistral => "vertex_ai/mistral-ocr-maas",
Self::AzureCohereParse => "azure_ai/cohere-parse",
Self::Cohere => "cohere/model",
}
}
fn document_type(self) -> &'static str {
match self {
Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url",
Self::AzureCohereParse | Self::Cohere => "image_url",
}
}
fn options(self) -> Value {
match self {
Self::Mistral | Self::AzureAi => json!({"pages": [0]}),
Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}),
Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}),
}
}
}
/// What the host does to the wire request in `before_send`.
#[derive(Clone, Copy, Debug)]
enum Host {
Detached,
ReplacesDocument,
}
const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ=";
impl Host {
fn before_send(self, wire: WireRequest) -> WireRequest {
let Value::Object(fields) = wire.body else {
return wire;
};
let body = fields
.into_iter()
.map(|(name, value)| match self {
Self::Detached => (name, value),
Self::ReplacesDocument if name == "document" => {
let document_type = value["type"].clone();
let key = document_type.as_str().unwrap_or_default().to_string();
(name, json!({"type": document_type, key: REPLACED_DOCUMENT}))
}
Self::ReplacesDocument => (name, value),
})
.collect();
WireRequest {
body: Value::Object(body),
..wire
}
}
}
struct Sent {
result: Result<(), Error>,
provider_body: Option<Value>,
}
async fn send(route: Route, host: Host, document_base: &str) -> Sent {
let (base, seen, provider) = mock_server(vec![MockResponse::json(json!({"pages": []}))]).await;
let document_type = route.document_type();
let document =
json!({"type": document_type, document_type: format!("{document_base}/scan.png")});
let request = wire_request_with_document(route.model(), &base, document, route.options());
let local =
LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(host.before_send(wire)));
let result = perform_ocr_with(local).await.map(|_| ());
match result {
Ok(()) => provider.await.unwrap(),
Err(_) => provider.abort(),
}
let provider_body = seen
.lock()
.unwrap()
.first()
.map(|request| request_body(request));
Sent {
result,
provider_body,
}
}
fn served_document_uri() -> String {
use base64::Engine;
format!(
"data:image/png;base64,{}",
base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT)
)
}
#[rstest]
#[case::azure_ai(Route::AzureAi)]
#[case::vertex_mistral(Route::VertexMistral)]
#[case::azure_cohere_parse(Route::AzureCohereParse)]
#[tokio::test]
async fn inlining_routes_send_the_downloaded_document(#[case] route: Route) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::Detached, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(served_document_uri())
);
}
#[rstest]
#[tokio::test]
async fn document_replaced_by_the_host_reaches_the_provider(
#[values(
Route::Mistral,
Route::AzureAi,
Route::VertexMistral,
Route::AzureCohereParse,
Route::Cohere
)]
route: Route,
) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::ReplacesDocument, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(REPLACED_DOCUMENT)
);
}

View file

@ -1,203 +0,0 @@
use std::sync::{Arc, Mutex};
use futures_util::future::BoxFuture;
use litellm_host::event::WireRequest;
use litellm_llms::base_llm::ocr::{
error::Error,
handler::{CallHooks, OcrClient},
transformation::LiteLLMOcrResponse,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
};
use crate::ocr::{
route::{LocalOcrHost, ocr_machine},
types::LiteLLMOcrRequest,
wire::{OcrWireRequest, decode_request},
};
/// Stands in for a host with no hooks registered: the wire request goes out unchanged
/// and response events go nowhere.
pub(crate) struct NoHooks;
impl CallHooks<Error> for NoHooks {
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
Box::pin(async move { Ok(wire) })
}
fn response_received<'a>(&'a self, _body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
Box::pin(async { Ok(()) })
}
}
pub(crate) fn ocr_client() -> OcrClient {
let document_http = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test document client builds");
OcrClient::for_test(reqwest::Client::new(), document_http)
}
pub(crate) async fn perform_ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
crate::ocr::client::perform(&ocr_client(), request).await
}
pub(crate) async fn perform_ocr_with(host: LocalOcrHost) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::run::run(ocr_machine(ocr_client()), &host).await
}
pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
wire_request_with_document(
model,
base,
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
options,
)
}
pub(crate) fn wire_request_with_document(
model: &str,
base: &str,
document: Value,
options: Value,
) -> LiteLLMOcrRequest {
decode_request(OcrWireRequest {
model: model.into(),
document,
api_key: Some(litellm_auth::SecretValue::new("test-key")),
api_base: Some(base.into()),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone(),
input_sources: Default::default(),
timeout_seconds: Some(2.0),
})
.unwrap()
}
pub(crate) fn resolved_request(
request: LiteLLMOcrRequest,
) -> crate::ocr::types::ResolvedOcrRequest {
request
.map_document(crate::ocr::document::prepare_document)
.unwrap()
}
pub(crate) fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest {
let request = resolved_request(request);
let document = request.document.clone().with_source(source.into());
request.with_document(document.into())
}
pub(crate) fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
pub(crate) const SERVED_DOCUMENT: &[u8] = b"\x89PNG served document";
/// Serves [`SERVED_DOCUMENT`] as `image/png` to every connection until aborted.
pub(crate) async fn document_server() -> (String, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let task = tokio::spawn(async move {
loop {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buffer = [0u8; 4096];
let _ = socket.read(&mut buffer).await.unwrap();
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Type: image/png\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
SERVED_DOCUMENT.len()
);
socket.write_all(head.as_bytes()).await.unwrap();
socket.write_all(SERVED_DOCUMENT).await.unwrap();
}
});
(base, task)
}
pub(crate) struct MockResponse {
pub status: u16,
pub headers: Vec<(&'static str, String)>,
pub body: Value,
}
impl MockResponse {
pub fn json(body: Value) -> Self {
Self {
status: 200,
headers: vec![],
body,
}
}
}
pub(crate) async fn mock_server(
responses: Vec<MockResponse>,
) -> (String, Arc<Mutex<Vec<String>>>, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let requests = Arc::new(Mutex::new(Vec::new()));
let seen = requests.clone();
let server_base = base.clone();
let task = tokio::spawn(async move {
for response in responses {
let (mut socket, _) = listener.accept().await.unwrap();
let mut bytes = Vec::new();
let mut buffer = [0u8; 4096];
let header_end = loop {
let n = socket.read(&mut buffer).await.unwrap();
assert!(n > 0);
bytes.extend_from_slice(&buffer[..n]);
if let Some(index) = bytes.windows(4).position(|s| s == b"\r\n\r\n") {
break index + 4;
}
};
let length = String::from_utf8_lossy(&bytes[..header_end])
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().unwrap())
})
.unwrap_or(0);
while bytes.len() < header_end + length {
let n = socket.read(&mut buffer).await.unwrap();
assert!(n > 0);
bytes.extend_from_slice(&buffer[..n]);
}
seen.lock()
.unwrap()
.push(String::from_utf8_lossy(&bytes).into_owned());
let body = serde_json::to_vec(&response.body).unwrap();
let headers = response
.headers
.into_iter()
.map(|(name, value)| {
format!("{name}: {}\r\n", value.replace("{base}", &server_base))
})
.collect::<String>();
let head = format!(
"HTTP/1.1 {} OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n{}\r\n",
response.status,
body.len(),
headers
);
socket.write_all(head.as_bytes()).await.unwrap();
socket.write_all(&body).await.unwrap();
}
});
(base, requests, task)
}
pub(crate) fn header<'a>(request: &'a str, name: &str) -> Option<&'a str> {
request
.lines()
.take_while(|line| !line.is_empty())
.find_map(|line| {
let (key, value) = line.split_once(':')?;
key.eq_ignore_ascii_case(name).then(|| value.trim())
})
}

View file

@ -1,584 +0,0 @@
use litellm_host::event::{CallEvent, MachineEvent, WireRequest};
use litellm_llms::base_llm::ocr::{error::Error, transformation::OcrDocument};
use rstest::rstest;
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request};
use crate::ocr::route::LocalOcrHost;
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[rstest]
#[case(
"reducto/parse-v3",
json!({
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://already.pdf",
json!({
"input":"reducto://already.pdf",
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[case(
"reducto/parse-legacy",
json!({
"enhance":{"agentic":[{"type":"table"}]},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://legacy.pdf",
json!({
"document_url":"reducto://legacy.pdf",
"options":{"enhance":{"agentic":[{"type":"table"}]}},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[tokio::test]
async fn request_mapping_matches_python(
#[case] model: &str,
#[case] options: Value,
#[case] source: &str,
#[case] expected: Value,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"result":{"chunks":[]}
}))])
.await;
let request = super::test_support::with_source(wire_request(model, &base, options), source);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert_eq!(request_body(&requests[0]), expected);
}
#[rstest]
#[case("parse-v3")]
#[case("parse-legacy")]
#[tokio::test]
async fn data_uri_upload_preserves_multipart_headers(
#[case] model: &str,
#[values("application/pdf", "image/png")] mime_type: &str,
) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[{"content":"hello"}]}})),
])
.await;
let document = if mime_type.starts_with("image/") {
json!({"type":"image_url","image_url":format!("data:{mime_type};base64,YWJj")})
} else {
json!({"type":"document_url","document_url":format!("data:{mime_type};base64,YWJj")})
};
let mut request = crate::ocr::types::LiteLLMOcrRequest {
document: serde_json::from_value::<OcrDocument>(document)
.unwrap()
.into(),
..wire_request(&format!("reducto/{model}"), &base, json!({}))
};
request.transport.extra_headers = vec![
("Content-Type".into(), "application/json".into()),
("X-Trace".into(), "upload-test".into()),
];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("content-type: multipart/form-data; boundary=")
);
assert!(requests[0].contains("x-trace: upload-test"));
let multipart = requests[0].split_once("\r\n\r\n").unwrap().1;
assert!(multipart.contains(&format!("Content-Type: {mime_type}\r\n")));
assert!(multipart.contains("\r\n\r\nabc\r\n--"));
assert!(requests[1].starts_with("POST /parse "));
let source_field = if model == "parse-legacy" {
"document_url"
} else {
"input"
};
assert_eq!(
request_body(&requests[1]),
json!({source_field:"reducto://uploaded.pdf"})
);
for request in requests.iter() {
assert!(
request
.to_ascii_lowercase()
.contains("authorization: bearer test-key\r\n")
);
}
}
#[tokio::test]
async fn response_received_stays_after_reducto_upload_and_parse() {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let request_count = seen.clone();
let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({}))).with_observer(
move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
assert_eq!(request_count.lock().unwrap().len(), 2);
assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#);
}
},
);
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[rstest]
#[case(json!({"file_id":""}))]
#[case(json!({}))]
#[case(json!({"file_id":null}))]
#[tokio::test]
async fn invalid_upload_ids_stop_before_parse(#[case] response: Value) {
let (base, seen, server) = mock_server(vec![MockResponse::json(response)]).await;
let error = perform_ocr(wire_request("reducto/parse-v3", &base, json!({})))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("file_id"));
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn upload_failure_stops_before_parse() {
let (base, seen, server) = mock_server(vec![MockResponse {
status: 503,
headers: vec![],
body: json!({"error":"unavailable"}),
}])
.await;
assert!(
perform_ocr(wire_request("reducto/parse-v3", &base, json!({})))
.await
.is_err()
);
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[rstest]
#[case("https://example.com/a.pdf", Error::ReductoSource)]
#[case("reducto://", Error::RequestField { path: "document file id".into() })]
#[case("data:application/pdf;base64", Error::InvalidDataUri)]
#[case("data:application/pdf;base64,INVALID!", Error::InvalidDataUri)]
#[tokio::test]
async fn rejects_invalid_document_sources_before_network(
#[case] source: &str,
#[case] expected: Error,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({}))]).await;
let request = super::test_support::with_source(
wire_request("reducto/parse-v3", &base, json!({})),
source,
);
let result = perform_ocr(request).await;
server.abort();
let _ = server.await;
assert!(
seen.lock().unwrap().is_empty(),
"sent invalid source: {source}"
);
let error = result.unwrap_err();
assert_eq!(
std::mem::discriminant(&error),
std::mem::discriminant(&expected)
);
assert_eq!(error.http_status_code(), Some(400));
assert_eq!(error.to_string(), expected.to_string());
}
#[test]
fn response_normalization_groups_blocks_and_distinguishes_null_result() {
use litellm_llms::reducto::ocr::transformation::{
ReductoResponse, normalize_response as transform_ocr_response,
};
let raw = json!({"usage":{"num_pages":"2","credits":"3"},"result":{"type":"full","chunks":[
{"blocks":[{
"type":"Table",
"content":"B",
"bbox":{"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4},
"confidence":"high",
"granular_confidence":{"parse_confidence":0.95,"extract_confidence":null},
"image_url":null
}]},
{"blocks":[{"content":"A","bbox":{"page":1},"type":"Text"},{"content":"C","bbox":{"page":1}}]}
]}});
let response: ReductoResponse = serde_json::from_value(raw).unwrap();
let normalized = transform_ocr_response("parse-v3", response)
.unwrap()
.into_json();
assert_eq!(normalized["pages"][0]["markdown"], "A\n\nC");
assert_eq!(normalized["pages"][1]["markdown"], "B");
assert_eq!(normalized["pages"][1]["blocks"][0]["type"], "Table");
assert_eq!(
normalized["pages"][1]["blocks"][0]["bbox"],
json!({"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4})
);
assert_eq!(normalized["pages"][1]["blocks"][0]["confidence"], "high");
assert_eq!(
normalized["pages"][1]["blocks"][0]["granular_confidence"]["parse_confidence"],
0.95
);
assert!(normalized["pages"][1]["blocks"][0]["image_url"].is_null());
assert_eq!(normalized["usage_info"]["pages_processed"], 2);
assert_eq!(normalized["usage_info"]["credits"], 3.0);
let missing: ReductoResponse =
serde_json::from_value(json!({"chunks":[{"content":"text"}]})).unwrap();
let missing = transform_ocr_response("parse-v3", missing).unwrap();
assert_eq!(missing.pages[0].markdown, "text");
let null: ReductoResponse = serde_json::from_value(
json!({"result":null,"chunks":[{"content":"ignored"}],"usage":null}),
)
.unwrap();
let null = transform_ocr_response("parse-v3", null).unwrap();
assert!(null.pages.is_empty());
}
#[tokio::test]
async fn facade_omits_native_response_by_default_and_preserves_auth_priority() {
let raw = json!({"job_id":"job-1","result":{"chunks":[]}});
let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await;
let mut request = super::test_support::with_source(
wire_request("reducto/parse-v3", &base, json!({})),
"reducto://ready.pdf",
);
request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.provider_native_response, None);
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer existing")
);
}
#[tokio::test]
async fn native_format_retains_the_provider_response() {
let raw = json!({
"result":{"chunks":[{"content":"native OCR response"}]},
"usage":{"num_pages":1}
});
let (base, _, server) = mock_server(vec![MockResponse::json(raw.clone())]).await;
let request = super::test_support::with_source(
wire_request("reducto/parse-v3", &base, json!({"req_format":"native"})),
"reducto://ready.pdf",
);
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "native OCR response");
assert_eq!(response.provider_native_response.as_ref(), raw.as_object());
}
#[tokio::test]
async fn unknown_model_reaches_parse_and_keeps_its_name() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"result":{"chunks":[{"content":"future model response"}]}
}))])
.await;
let request = super::test_support::with_source(
wire_request("reducto/future-parse-model", &base, json!({})),
"reducto://ready.pdf",
);
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.model, "future-parse-model");
assert_eq!(response.pages[0].markdown, "future model response");
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with("POST /parse "));
assert_eq!(
request_body(&requests[0]),
json!({"input":"reducto://ready.pdf"})
);
}
#[tokio::test]
async fn guardrail_rewrites_document_before_upload() {
let (base, seen, server) =
mock_server(vec![MockResponse::json(json!({"result":{"chunks":[]}}))]).await;
let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({})))
.with_before_send(|wire, _| {
assert_eq!(
wire.body["document_url"],
"data:application/pdf;base64,YWJj"
);
Ok(WireRequest {
body: json!({"type":"document_url","document_url":"reducto://guarded.pdf"}),
..wire
})
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert!(requests[0].contains("reducto://guarded.pdf"));
}
mod transformation {
use litellm_host::event::{CallEvent, MachineEvent, WireRequest};
use litellm_llms::{
base_llm::ocr::transformation::{BaseOcrConfig, OcrConnection, OcrRequestContext},
reducto::ocr::transformation::*,
};
use rstest::rstest;
use super::*;
use crate::ocr::{
route::LocalOcrHost,
test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request},
};
#[tokio::test]
async fn v3_options_preserve_explicit_null() {
let overrides =
serde_json::from_value(json!({"formatting":null,"settings":{},"unknown":true}))
.unwrap();
let params = ReductoParseV3Config
.map_ocr_params(&overrides, "parse-v3")
.unwrap();
let client = crate::ocr::test_support::ocr_client();
let connection = OcrConnection::default();
let document = serde_json::from_value(
json!({"type":"document_url","document_url":"reducto://ready.pdf"}),
)
.unwrap();
let body = ReductoParseV3Config
.async_transform_ocr_request(
"parse-v3",
document,
&params,
&[],
OcrRequestContext {
client: &client,
connection: &connection,
},
)
.await
.unwrap();
assert_eq!(
serde_json::to_value(body).unwrap(),
json!({
"input":"reducto://ready.pdf", "formatting":null, "settings":{}
})
);
let absent = ReductoParseV3Config
.map_ocr_params(
&litellm_core_utils::call_arguments::CallArguments::default(),
"parse-v3",
)
.unwrap();
assert_eq!(serde_json::to_value(absent).unwrap(), json!({}));
}
#[rstest]
#[case(
"reducto/parse-v3",
json!({
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://already.pdf",
json!({
"input":"reducto://already.pdf",
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[case(
"reducto/parse-legacy",
json!({
"enhance":{"agentic":[{"type":"table"}]},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://legacy.pdf",
json!({
"document_url":"reducto://legacy.pdf",
"options":{"enhance":{"agentic":[{"type":"table"}]}},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[tokio::test]
async fn request_mapping_matches_python(
#[case] model: &str,
#[case] options: Value,
#[case] source: &str,
#[case] expected: Value,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"result":{"chunks":[]}
}))])
.await;
let request =
crate::ocr::test_support::with_source(wire_request(model, &base, options), source);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert_eq!(request_body(&requests[0]), expected);
}
#[rstest]
#[case("parse-v3")]
#[case("parse-legacy")]
#[tokio::test]
async fn data_uri_upload_preserves_multipart_headers(#[case] model: &str) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[{"content":"hello"}]}})),
])
.await;
let mut request = wire_request(&format!("reducto/{model}"), &base, json!({}));
request.transport.extra_headers = vec![
("Content-Type".into(), "application/json".into()),
("X-Trace".into(), "upload-test".into()),
];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("content-type: multipart/form-data; boundary=")
);
assert!(requests[0].contains("x-trace: upload-test"));
assert!(requests[0].contains("application/pdf"));
assert!(requests[0].contains("abc"));
assert!(requests[1].starts_with("POST /parse "));
}
#[tokio::test]
async fn response_received_stays_after_reducto_upload_and_parse() {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let request_count = seen.clone();
let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({})))
.with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
assert_eq!(request_count.lock().unwrap().len(), 2);
assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#);
}
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[rstest]
#[case("https://example.com/a.pdf")]
#[case("reducto://")]
#[case("data:application/pdf;base64")]
#[case("data:application/pdf;base64,INVALID!")]
#[tokio::test]
async fn rejects_invalid_document_sources_before_network(#[case] source: &str) {
let request = crate::ocr::test_support::with_source(
wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})),
source,
);
assert!(perform_ocr(request).await.is_err());
}
#[tokio::test]
async fn facade_omits_native_response_by_default_and_preserves_auth_priority() {
let raw = json!({"job_id":"job-1","result":{"chunks":[]}});
let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await;
let mut request = crate::ocr::test_support::with_source(
wire_request("reducto/parse-v3", &base, json!({})),
"reducto://ready.pdf",
);
request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.provider_native_response, None);
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer existing")
);
}
#[rstest]
#[case("reducto/parse-v3")]
#[case("reducto/parse-legacy")]
#[tokio::test]
async fn guardrail_headers_reach_upload_and_parse(#[case] model: &str) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let mut request = wire_request(model, &base, json!({}));
request.transport.extra_headers = vec![("authorization".into(), "Bearer original".into())];
let host = LocalOcrHost::new(request).with_before_send(|wire, _| {
Ok(WireRequest {
headers: vec![("authorization".into(), "Bearer guarded".into())],
..wire
})
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(requests[1].starts_with("POST /parse "));
for request in requests.iter() {
assert!(request.contains("authorization: Bearer guarded"));
assert!(!request.contains("Bearer original"));
}
}
}

View file

@ -1,143 +0,0 @@
use litellm_auth::InputSource;
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"choices":[{"message":{"content":"recognized"}}],
"usage":{"prompt_tokens":1}
}))])
.await;
let request = wire_request(
"vertex_ai/deepseek-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"temperature":0.1,
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
);
let request = super::test_support::with_source(request, "gs://bucket/document.pdf");
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "recognized");
assert_eq!(
response.usage_info.unwrap().extra_fields["prompt_tokens"],
1
);
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
let body = request_body(&requests[0]);
assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(body["temperature"], 0.1);
assert_eq!(body["future_ocr_option"], true);
assert!(body.get("extra_body").is_none());
assert_eq!(
body["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/document.pdf"})
);
}
#[test]
fn host_registration_selects_deepseek_without_affecting_mistral() {
assert!(crate::ocr::arguments::is_supported_request(
"deepseek-ocr-maas",
Some("vertex_ai")
));
assert!(crate::ocr::arguments::is_supported_request(
"mistral-ocr-maas",
Some("vertex_ai")
));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/deepseek-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.credentials.api_base = Some(litellm_auth::Sourced::new(
"https://caller.example".into(),
InputSource::Request,
));
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}
mod deepseek_transformation {
use serde_json::json;
use super::*;
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
#[tokio::test]
async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"choices":[{"message":{"content":"recognized"}}],
"usage":{"prompt_tokens":1}
}))])
.await;
let request = wire_request(
"vertex_ai/deepseek-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"temperature":0.1,
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
);
let request = crate::ocr::test_support::with_source(request, "gs://bucket/document.pdf");
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "recognized");
assert_eq!(
response.usage_info.unwrap().extra_fields["prompt_tokens"],
1
);
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
let body = request_body(&requests[0]);
assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(body["temperature"], 0.1);
assert_eq!(body["future_ocr_option"], true);
assert_eq!(body["provider_option"], "value");
assert!(body.get("vertex_project").is_none());
assert!(body.get("extra_body").is_none());
assert_eq!(
body["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/document.pdf"})
);
}
}

View file

@ -1,293 +0,0 @@
use litellm_auth::InputSource;
use litellm_llms::base_llm::ocr::{settings::OcrSettings, transformation::OcrResponseFormat};
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, ocr_client, perform_ocr, wire_request};
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_mistral_with_resolved_project_and_location() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let request = wire_request(
"vertex_ai/mistral-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"extract_footer":true
}),
);
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
assert_eq!(
request_body(&requests[0]),
json!({
"model":"mistral-ocr-maas",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"extract_footer":true
})
);
}
#[tokio::test]
async fn configured_project_and_location_apply_when_the_call_sets_neither() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let client = ocr_client().with_settings(OcrSettings {
vertex_project: Some("configured-project".into()),
vertex_location: Some("europe-west4".into()),
..OcrSettings::default()
});
crate::ocr::client::perform(
&client,
wire_request("vertex_ai/mistral-ocr-maas", &base, json!({})),
)
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].starts_with(
"POST /v1/projects/configured-project/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict "
));
}
#[tokio::test]
async fn supplied_authorization_is_forwarded_without_a_static_token() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"vertex_ai/model",
&base,
json!({"vertex_project":"project-1"}),
);
request.credentials.api_key = None;
request.transport.extra_headers = vec![("authorization".into(), "Bearer supplied".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer supplied")
);
}
#[tokio::test]
async fn invalid_credentials_fail_before_provider_http() {
let request = wire_request(
"vertex_ai/model",
"http://127.0.0.1:1",
json!({"vertex_credentials": true}),
);
let error = perform_ocr(request).await.unwrap_err();
assert!(error.to_string().contains("vertex_credentials"));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/mistral-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.credentials.api_base = Some(litellm_auth::Sourced::new(
"https://caller.example".into(),
InputSource::Request,
));
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}
#[tokio::test]
async fn adapters_build_complete_requests_and_share_mistral_normalization() {
use std::time::Duration;
use litellm_llms::{
base_llm::ocr::transformation::BaseOcrConfig,
mistral::ocr::transformation::MistralOcrConfig,
vertex_ai::ocr::transformation::VertexAiOcrConfig,
};
use crate::ocr::test_support::ocr_client;
let client = ocr_client();
let options = json!({
"pages": [0, 2],
"include_image_base64": true,
"vertex_project": "project-1",
"vertex_location": "us-central1",
"unknown": "ignored"
});
let direct = wire_request(
"mistral/mistral-ocr-maas",
"https://mistral.test",
options.clone(),
);
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
let direct = crate::ocr::prepare::prepare_request_for_test(
super::test_support::resolved_request(direct),
);
let vertex = crate::ocr::prepare::prepare_request_for_test(
super::test_support::resolved_request(vertex),
);
let direct_http = MistralOcrConfig
.prepare_request(&direct, &client, &crate::ocr::test_support::NoHooks)
.await
.unwrap();
let vertex_http = VertexAiOcrConfig
.prepare_request(&vertex, &client, &crate::ocr::test_support::NoHooks)
.await
.unwrap();
assert_eq!(direct_http.url(), "https://mistral.test/v1/ocr");
assert_eq!(
vertex_http.url(),
"https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
for http in [&direct_http, &vertex_http] {
assert_eq!(http.header("authorization").unwrap(), "Bearer test-key");
assert_eq!(http.header("content-type").unwrap(), "application/json");
assert_eq!(http.timeout(), Some(Duration::from_secs(2)));
let body: Value = serde_json::from_slice(http.body()).unwrap();
assert_eq!(
body,
json!({
"model": "mistral-ocr-maas",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"pages": [0, 2],
"include_image_base64": true,
"unknown": "ignored"
})
);
}
let payload = json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"});
let raw = serde_json::to_vec(&payload).unwrap();
let direct_response = MistralOcrConfig
.transform_ocr_response(&direct.model, &raw, OcrResponseFormat::Litellm)
.unwrap()
.into_json();
let vertex_response = VertexAiOcrConfig
.transform_ocr_response(&vertex.model, &raw, OcrResponseFormat::Litellm)
.unwrap()
.into_json();
assert_eq!(direct_response, vertex_response);
assert_eq!(direct_response["model"], "mistral-ocr-maas");
assert_eq!(direct_response["object"], "ocr");
assert_eq!(direct_response["extra"], "preserved");
}
mod transformation {
use rstest::rstest;
use serde_json::{Value, json};
use crate::ocr::test_support::wire_request;
#[rstest]
#[case::mistral(false)]
#[case::vertex(true)]
#[tokio::test]
async fn configs_build_complete_requests_and_share_mistral_normalization(
#[case] use_vertex: bool,
) {
use std::time::Duration;
use litellm_llms::{
base_llm::ocr::transformation::BaseOcrConfig,
mistral::ocr::transformation::MistralOcrConfig,
vertex_ai::ocr::transformation::VertexAiOcrConfig,
};
use crate::ocr::test_support::ocr_client;
let client = ocr_client();
let options = json!({
"pages": [0, 2],
"include_image_base64": true,
"vertex_project": "project-1",
"vertex_location": "us-central1",
"unknown": "preserved"
});
let direct = wire_request(
"mistral/mistral-ocr-maas",
"https://mistral.test",
options.clone(),
);
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
let direct = crate::ocr::prepare::prepare_request_for_test(
crate::ocr::test_support::resolved_request(direct),
);
let vertex = crate::ocr::prepare::prepare_request_for_test(
crate::ocr::test_support::resolved_request(vertex),
);
let direct_http = MistralOcrConfig
.prepare_request(&direct, &client, &crate::ocr::test_support::NoHooks)
.await
.unwrap();
let vertex_http = VertexAiOcrConfig
.prepare_request(&vertex, &client, &crate::ocr::test_support::NoHooks)
.await
.unwrap();
assert_eq!(direct_http.url(), "https://mistral.test/v1/ocr");
assert_eq!(
vertex_http.url(),
"https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
let http = if use_vertex {
&vertex_http
} else {
&direct_http
};
assert_eq!(http.header("authorization").unwrap(), "Bearer test-key");
assert_eq!(http.header("content-type").unwrap(), "application/json");
assert_eq!(http.timeout(), Some(Duration::from_secs(2)));
let body: Value = serde_json::from_slice(http.body()).unwrap();
assert_eq!(
body,
json!({
"model": "mistral-ocr-maas",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"pages": [0, 2],
"include_image_base64": true,
"unknown": "preserved"
})
);
let payload = serde_json::to_vec(
&json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"}),
)
.unwrap();
let direct_response = MistralOcrConfig
.transform_ocr_response(&direct.model, &payload, Default::default())
.unwrap()
.into_json();
let vertex_response = VertexAiOcrConfig
.transform_ocr_response(&vertex.model, &payload, Default::default())
.unwrap()
.into_json();
assert_eq!(direct_response, vertex_response);
assert_eq!(direct_response["model"], "mistral-ocr-maas");
assert_eq!(direct_response["object"], "ocr");
assert_eq!(direct_response["extra"], "preserved");
}
}

View file

@ -217,13 +217,25 @@ mod tests {
"Authorization".to_string(),
"Bearer abc".to_string()
)]));
assert!(has_bearer_auth(&[(
"authorization".to_string(),
"bearer abc".to_string()
)]));
assert!(!has_bearer_auth(&[(
"Authorization".to_string(),
"Bearer ".to_string()
)]));
assert!(!has_bearer_auth(&[(
"authorization".to_string(),
String::new()
)]));
assert!(!has_bearer_auth(&[(
"Authorization".to_string(),
"Basic abc".to_string()
)]));
assert!(!has_bearer_auth(&[(
"x-api-key".to_string(),
"abc".to_string()
)]));
}
}

View file

@ -218,7 +218,3 @@ fn anthropic_body(
);
Value::Object(body)
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,270 @@
use litellm_types::llms::anthropic_messages::anthropic_request::{
AnthropicMessage, AnthropicMessagesRequest,
};
use serde_json::{Value, json};
use crate::{
anthropic::common_utils::{
flatten_unencrypted_web_search_results, sanitize_tool_use_ids, strip_empty_content_blocks,
strip_provider_specific_fields,
},
base_llm::chat::transformation::Error,
};
pub fn shape_anthropic_messages_request(
request: AnthropicMessagesRequest,
reasoning_auto_summary: bool,
) -> Result<AnthropicMessagesRequest, Error> {
Ok(AnthropicMessagesRequest {
messages: sanitize_anthropic_messages(request.messages),
metadata: request
.metadata
.as_ref()
.map(validate_anthropic_api_metadata)
.transpose()?,
thinking: with_reasoning_auto_summary(request.thinking, reasoning_auto_summary),
..request
})
}
fn sanitize_anthropic_messages(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
strip_provider_specific_fields(flatten_unencrypted_web_search_results(
sanitize_tool_use_ids(strip_empty_content_blocks(messages)),
))
}
fn validate_anthropic_api_metadata(metadata: &Value) -> Result<Value, Error> {
let Value::Object(fields) = metadata else {
return Err(Error::InvalidRequest(format!(
"metadata must be an object, got {metadata}"
)));
};
match fields.get("user_id") {
None | Some(Value::Null) => Ok(json!({})),
Some(Value::String(user_id)) => Ok(json!({"user_id": user_id})),
Some(other) => Err(Error::InvalidRequest(format!(
"metadata.user_id must be a string, got {other}"
))),
}
}
fn with_reasoning_auto_summary(thinking: Option<Value>, enabled: bool) -> Option<Value> {
let Some(Value::Object(thinking)) = thinking else {
return thinking;
};
if !enabled || thinking.get("type").and_then(Value::as_str) == Some("disabled") {
return Some(Value::Object(thinking));
}
Some(Value::Object(
thinking
.into_iter()
.filter(|(key, _)| key != "display")
.chain([("display".to_string(), json!("summarized"))])
.collect(),
))
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use super::*;
fn messages(value: Value) -> Vec<AnthropicMessage> {
serde_json::from_value(value).unwrap()
}
fn request(body: Value) -> AnthropicMessagesRequest {
serde_json::from_value(body).unwrap()
}
#[rstest]
#[case::empty_text_next_to_a_tool_use(
json!([{"role": "assistant", "content": [
{"type": "text", "text": " "},
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
]}]),
json!([{"role": "assistant", "content": [
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
]}]),
)]
#[case::cross_provider_tool_ids(
json!([
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]}
]),
json!([
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
]),
)]
#[case::replayed_unencrypted_web_search_results(
json!([
{"role": "user", "content": "latest litellm version?"},
{"role": "assistant", "content": [
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "latest litellm version"}},
{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_1", "content": [{
"type": "web_search_result",
"url": "https://github.com/BerriAI/litellm/releases",
"title": "Releases",
"page_age": null,
"encrypted_content": "",
"snippet": "Latest release v1.95.0"
}]}
]},
{"role": "user", "content": "which version?"}
]),
json!([
{"role": "user", "content": "latest litellm version?"},
{"role": "assistant", "content": [{
"type": "text",
"text": "Web search results for 'latest litellm version':\n\nTitle: Releases\nURL: https://github.com/BerriAI/litellm/releases\nSnippet: Latest release v1.95.0"
}]},
{"role": "user", "content": "which version?"}
]),
)]
#[case::replayed_provider_specific_fields(
json!([
{"role": "assistant", "content": [{
"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"},
"provider_specific_fields": {"signature": "sig_abc"}
}]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
]),
json!([
{"role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"}}]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
]),
)]
#[case::ids_are_normalized_before_web_search_results_flatten(
json!([
{"role": "user", "content": "run it"},
{"role": "assistant", "content": [
{"type": "thinking", "thinking": "", "signature": "sig"},
{"type": "text", "text": ""},
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}, "provider_specific_fields": {"x": 1}},
{"type": "server_tool_use", "id": "srv.1", "name": "web_search", "input": {"query": "q"}, "provider_specific_fields": {"x": 2}},
{"type": "web_search_tool_result", "tool_use_id": "srv.1", "provider_specific_fields": {"x": 3}, "content": [
{"type": "web_search_result", "url": "u", "title": "", "encrypted_content": "", "provider_specific_fields": {"x": 4}}
]}
]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]},
{"role": "assistant", "content": [{"type": "text", "text": " "}]}
]),
json!([
{"role": "user", "content": "run it"},
{"role": "assistant", "content": [
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}},
{"type": "server_tool_use", "id": "srv_1", "name": "web_search", "input": {"query": "q"}},
{"type": "text", "text": "Web search results:\n\nURL: u"}
]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
]),
)]
fn sanitize_anthropic_messages_cleans_replayed_history(
#[case] history: Value,
#[case] expected: Value,
) {
assert_eq!(
serde_json::to_value(sanitize_anthropic_messages(messages(history))).unwrap(),
expected
);
}
#[rstest]
#[case::keeps_only_user_id(json!({"user_id": "u-1", "trace_id": "internal"}), Ok(json!({"user_id": "u-1"})))]
#[case::null_user_id(json!({"user_id": null, "trace_id": "internal"}), Ok(json!({})))]
#[case::no_user_id(json!({"trace_id": "internal"}), Ok(json!({})))]
#[case::empty(json!({}), Ok(json!({})))]
#[case::numeric_user_id(
json!({"user_id": 123}),
Err(Error::InvalidRequest("metadata.user_id must be a string, got 123".to_string())),
)]
#[case::boolean_user_id(
json!({"user_id": true}),
Err(Error::InvalidRequest("metadata.user_id must be a string, got true".to_string())),
)]
#[case::not_an_object(
json!(["u-1"]),
Err(Error::InvalidRequest(r#"metadata must be an object, got ["u-1"]"#.to_string())),
)]
fn validate_anthropic_api_metadata_passes_only_a_string_user_id(
#[case] metadata: Value,
#[case] expected: Result<Value, Error>,
) {
assert_eq!(validate_anthropic_api_metadata(&metadata), expected);
}
#[rstest]
#[case::adaptive(
Some(json!({"type": "adaptive", "budget_tokens": 5000})),
true,
Some(json!({"type": "adaptive", "budget_tokens": 5000, "display": "summarized"})),
)]
#[case::enabled(
Some(json!({"type": "enabled", "budget_tokens": 10000})),
true,
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
)]
#[case::no_type(Some(json!({})), true, Some(json!({"display": "summarized"})))]
#[case::display_omitted_is_overridden(
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "omitted"})),
true,
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
)]
#[case::display_summarized_is_kept(
Some(json!({"type": "enabled", "display": "summarized"})),
true,
Some(json!({"type": "enabled", "display": "summarized"})),
)]
#[case::disabled_thinking(Some(json!({"type": "disabled"})), true, Some(json!({"type": "disabled"})))]
#[case::flag_off(
Some(json!({"type": "enabled", "budget_tokens": 10000})),
false,
Some(json!({"type": "enabled", "budget_tokens": 10000})),
)]
#[case::flag_off_keeps_callers_display(
Some(json!({"type": "enabled", "display": "omitted"})),
false,
Some(json!({"type": "enabled", "display": "omitted"})),
)]
#[case::no_thinking(None, true, None)]
#[case::non_object_thinking(Some(json!("enabled")), true, Some(json!("enabled")))]
fn reasoning_auto_summary_marks_active_thinking_as_summarized(
#[case] thinking: Option<Value>,
#[case] enabled: bool,
#[case] expected: Option<Value>,
) {
assert_eq!(with_reasoning_auto_summary(thinking, enabled), expected);
}
#[test]
fn shaping_cleans_messages_metadata_and_thinking() {
let sanitized = shape_anthropic_messages_request(
request(json!({
"model": "m",
"messages": [{"role": "assistant", "content": [
{"type": "text", "text": ""},
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}
]}],
"metadata": {"user_id": "u", "trace_id": "t"},
"thinking": {"type": "enabled", "budget_tokens": 1024},
"safeguards": [{"type": "dangerous_tool_use"}]
})),
true,
)
.unwrap();
assert_eq!(
serde_json::to_value(sanitized).unwrap(),
json!({
"model": "m",
"messages": [{"role": "assistant", "content": [
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}
]}],
"metadata": {"user_id": "u"},
"thinking": {"type": "enabled", "budget_tokens": 1024, "display": "summarized"},
"safeguards": [{"type": "dangerous_tool_use"}]
})
);
}
}

View file

@ -0,0 +1,643 @@
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
use serde_json::Value;
use crate::{
anthropic::{
ANTHROPIC_OAUTH_TOKEN_PREFIX,
common_utils::{
ANTHROPIC_OAUTH_BETA_HEADER, beta, has_advisor_tool, is_anthropic_oauth_key,
is_tool_search_used, join_beta_values, requires_native_compaction_beta,
split_beta_values,
},
},
base_llm::anthropic_messages::transformation::Headers,
};
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
const BETA_HEADER: &str = "anthropic-beta";
const AUTHORIZATION: &str = "authorization";
const API_KEY_HEADER: &str = "x-api-key";
const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access";
fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
headers
.iter()
.find(|(header, _)| header.eq_ignore_ascii_case(name))
.map(|(_, value)| value.as_str())
}
fn without(headers: Headers, names: &[&str]) -> Headers {
headers
.into_iter()
.filter(|(header, _)| !names.iter().any(|name| header.eq_ignore_ascii_case(name)))
.collect()
}
fn existing_betas(headers: &[(String, String)]) -> impl Iterator<Item = String> + '_ {
headers
.iter()
.filter(|(header, _)| header.eq_ignore_ascii_case(BETA_HEADER))
.flat_map(|(_, value)| split_beta_values(Some(value)))
}
fn with_oauth_bearer(headers: Headers, bearer: String) -> Headers {
let beta =
join_beta_values(existing_betas(&headers).chain([ANTHROPIC_OAUTH_BETA_HEADER.to_string()]));
without(headers, &[API_KEY_HEADER, AUTHORIZATION, BETA_HEADER])
.into_iter()
.chain([
(AUTHORIZATION.to_string(), bearer),
(BETA_HEADER.to_string(), beta),
(DIRECT_BROWSER_ACCESS_HEADER.to_string(), "true".to_string()),
])
.collect()
}
fn non_empty(value: Option<&str>) -> Option<&str> {
value.map(str::trim).filter(|value| !value.is_empty())
}
pub fn authenticate(
headers: Headers,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Headers, litellm_auth::Error> {
if let Some(forwarded) = header_value(&headers, AUTHORIZATION)
&& forwarded
.strip_prefix("Bearer ")
.is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX))
{
let bearer = forwarded.to_string();
return Ok(with_oauth_bearer(headers, bearer));
}
if let Some(key) = api_key.filter(|key| key.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) {
return Ok(with_oauth_bearer(headers, format!("Bearer {key}")));
}
if header_value(&headers, API_KEY_HEADER).is_some()
|| header_value(&headers, AUTHORIZATION).is_some()
{
return Ok(headers);
}
let resolved_key = non_empty(api_key)
.map(str::to_string)
.or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty()));
let auth = match resolved_key {
Some(key) if is_anthropic_oauth_key(&key) => {
(AUTHORIZATION.to_string(), format!("Bearer {key}"))
}
Some(key) => (API_KEY_HEADER.to_string(), key),
None => match env_lookup(ANTHROPIC_AUTH_TOKEN_ENV).filter(|value| !value.trim().is_empty())
{
Some(token) => (AUTHORIZATION.to_string(), format!("Bearer {token}")),
None => {
return Err(litellm_auth::Error::MissingApiKey {
provider: "Anthropic",
environment_variable: ANTHROPIC_API_KEY_ENV,
});
}
},
};
Ok(headers.into_iter().chain([auth]).collect())
}
fn context_management_betas(
context_management: Option<&Value>,
) -> impl Iterator<Item = &'static str> {
let edits = context_management
.and_then(|value| value.get("edits"))
.and_then(Value::as_array)
.map(Vec::as_slice)
.unwrap_or(&[]);
let (compact, other) = edits.iter().fold((false, false), |(compact, other), edit| {
match edit.get("type").and_then(Value::as_str) {
Some("compact_20260112") => (true, other),
_ => (compact, true),
}
});
compact
.then_some(beta::COMPACT_2026_01_12)
.into_iter()
.chain(other.then_some(beta::CONTEXT_MANAGEMENT_2025_06_27))
}
fn uses_structured_output(request: &AnthropicMessagesRequest) -> bool {
request.output_format.is_some()
|| request
.output_config
.as_ref()
.and_then(|config| config.get("format"))
.is_some_and(|format| !format.is_null())
}
fn messages_carry_output_config(request: &AnthropicMessagesRequest) -> bool {
request
.messages
.iter()
.any(|message| message.extra.contains_key("output_config"))
}
pub fn feature_betas(request: &AnthropicMessagesRequest) -> Vec<&'static str> {
let tools = request.tools.as_deref();
[
requires_native_compaction_beta(request.compaction.as_ref(), &request.messages)
.then_some(beta::COMPACT_2026_09_04),
uses_structured_output(request).then_some(beta::STRUCTURED_OUTPUT),
(request.speed.as_deref() == Some("fast")).then_some(beta::FAST_MODE_2026_02_01),
messages_carry_output_config(request).then_some(beta::PER_TURN_CONTROL_2026_07_01),
has_advisor_tool(tools).then_some(beta::ADVISOR_TOOL_2026_03_01),
is_tool_search_used(tools).then_some(beta::ADVANCED_TOOL_USE_2025_11_20),
]
.into_iter()
.flatten()
.chain(context_management_betas(
request.context_management.as_ref(),
))
.collect()
}
pub fn with_feature_betas(headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
let existing = existing_betas(&headers).collect::<Vec<_>>();
let features = feature_betas(request);
if existing.is_empty() && features.is_empty() {
return headers;
}
let merged = join_beta_values(
existing
.into_iter()
.chain(features.into_iter().map(str::to_string)),
);
without(headers, &[BETA_HEADER])
.into_iter()
.chain([(BETA_HEADER.to_string(), merged)])
.collect()
}
#[cfg(test)]
mod tests {
use rstest::{fixture, rstest};
use serde_json::json;
use super::*;
const OAUTH_TOKEN: &str = "sk-ant-oat01-token";
const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token";
const REGULAR_KEY: &str = "sk-ant-api03-regular";
const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true");
type Env = &'static [(&'static str, &'static str)];
fn request(fields: Value) -> AnthropicMessagesRequest {
let mut body =
json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]});
body.as_object_mut()
.unwrap()
.extend(fields.as_object().unwrap().clone());
serde_json::from_value(body).unwrap()
}
fn headers(pairs: &[(&str, &str)]) -> Headers {
pairs
.iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect()
}
fn betas(values: &[&str]) -> String {
values.join(",")
}
#[fixture]
fn no_env() -> Env {
&[]
}
#[fixture]
fn full_env() -> Env {
&[
("ANTHROPIC_API_KEY", "sk-env"),
("ANTHROPIC_AUTH_TOKEN", "env-token"),
]
}
fn authenticate_with(
forwarded: &[(&str, &str)],
api_key: Option<&str>,
env: Env,
) -> Result<Headers, litellm_auth::Error> {
let lookup = |name: &str| {
env.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| value.to_string())
};
authenticate(headers(forwarded), api_key, &lookup)
}
#[rstest]
#[case::forwarded_bearer_drops_forwarded_and_deployment_keys(
&[("X-Api-Key", REGULAR_KEY), ("Authorization", OAUTH_BEARER)],
Some(REGULAR_KEY),
OAUTH_BEARER,
&[],
)]
#[case::forwarded_bearer_in_uppercase_authorization_header(
&[("AUTHORIZATION", OAUTH_BEARER)],
None,
OAUTH_BEARER,
&[],
)]
#[case::forwarded_bearer_keeps_unrelated_headers_in_place(
&[("anthropic-version", "2023-06-01"), ("authorization", OAUTH_BEARER)],
None,
OAUTH_BEARER,
&[("anthropic-version", "2023-06-01")],
)]
#[case::forwarded_bearer_wins_over_an_oauth_api_key(
&[("authorization", OAUTH_BEARER)],
Some("sk-ant-oat01-deployment"),
OAUTH_BEARER,
&[],
)]
#[case::api_key_authenticates_as_a_bearer(&[], Some(OAUTH_TOKEN), OAUTH_BEARER, &[])]
#[case::api_key_removes_a_forwarded_x_api_key(
&[("x-api-key", OAUTH_TOKEN)],
Some(OAUTH_TOKEN),
OAUTH_BEARER,
&[],
)]
#[case::api_key_replaces_a_forwarded_non_oauth_bearer(
&[("Authorization", "Bearer some-proxy-token")],
Some(OAUTH_TOKEN),
OAUTH_BEARER,
&[],
)]
fn oauth_token_is_the_whole_credential(
#[case] forwarded: &[(&str, &str)],
#[case] api_key: Option<&str>,
#[case] expected_bearer: &str,
#[case] kept: &[(&str, &str)],
full_env: Env,
) {
let expected = kept
.iter()
.copied()
.chain([
("authorization", expected_bearer),
("anthropic-beta", ANTHROPIC_OAUTH_BETA_HEADER),
BROWSER_ACCESS,
])
.collect::<Vec<_>>();
assert_eq!(
authenticate_with(forwarded, api_key, full_env).unwrap(),
headers(&expected)
);
}
#[rstest]
#[case::forwarded_bearer_merges_a_differently_cased_beta_header(
&[("Anthropic-Beta", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
None,
)]
#[case::forwarded_bearer_dedupes_an_existing_oauth_beta(
&[("anthropic-beta", "web-search-2025-03-05, oauth-2025-04-20"), ("authorization", OAUTH_BEARER)],
None,
)]
#[case::api_key_merges_the_existing_beta_header(
&[("anthropic-beta", " web-search-2025-03-05 ,")],
Some(OAUTH_TOKEN),
)]
#[case::forwarded_bearer_unions_every_beta_header_casing(
&[("anthropic-beta", "oauth-2025-04-20"), ("ANTHROPIC-BETA", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
None,
)]
fn oauth_beta_merges_into_existing_betas(
#[case] forwarded: &[(&str, &str)],
#[case] api_key: Option<&str>,
no_env: Env,
) {
assert_eq!(
authenticate_with(forwarded, api_key, no_env).unwrap(),
headers(&[
("authorization", OAUTH_BEARER),
(
"anthropic-beta",
&betas(&[ANTHROPIC_OAUTH_BETA_HEADER, "web-search-2025-03-05"])
),
BROWSER_ACCESS,
])
);
}
#[rstest]
#[case::x_api_key_over_the_deployment_key(&[("x-api-key", "caller-key")], Some("sk-other"))]
#[case::uppercase_x_api_key(&[("X-API-KEY", "caller-key")], None)]
#[case::non_oauth_bearer(&[("Authorization", "Bearer some-proxy-token")], None)]
#[case::non_oauth_bearer_over_a_regular_api_key(
&[("authorization", "Bearer sk-ant-api03-forwarded")],
Some(REGULAR_KEY),
)]
#[case::oauth_token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)]
#[case::oauth_token_behind_a_lowercase_bearer_scheme(
&[("authorization", "bearer sk-ant-oat01-token")],
None,
)]
fn forwarded_auth_header_is_kept_untouched(
#[case] forwarded: &[(&str, &str)],
#[case] api_key: Option<&str>,
full_env: Env,
) {
assert_eq!(
authenticate_with(forwarded, api_key, full_env).unwrap(),
headers(forwarded)
);
}
#[rstest]
#[case::api_key_param(Some("sk-param"), &[], ("x-api-key", "sk-param"))]
#[case::api_key_param_over_env_key_and_auth_token(
Some("sk-param"),
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
("x-api-key", "sk-param"),
)]
#[case::env_key_without_a_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
#[case::env_key_when_the_param_is_empty(Some(""), &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
#[case::env_key_when_the_param_is_whitespace(
Some(" "),
&[("ANTHROPIC_API_KEY", "sk-env")],
("x-api-key", "sk-env"),
)]
#[case::env_key_over_auth_token(
None,
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
("x-api-key", "sk-env"),
)]
#[case::auth_token_as_a_bearer(
None,
&[("ANTHROPIC_AUTH_TOKEN", "env-token")],
("authorization", "Bearer env-token"),
)]
#[case::auth_token_when_the_env_key_is_whitespace(
None,
&[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
("authorization", "Bearer env-token"),
)]
#[case::oauth_env_key_as_a_plain_bearer(
None,
&[("ANTHROPIC_API_KEY", "sk-ant-oat01-env")],
("authorization", "Bearer sk-ant-oat01-env"),
)]
fn credential_is_resolved_after_the_existing_headers(
#[case] api_key: Option<&str>,
#[case] env: Env,
#[case] expected: (&str, &str),
) {
let forwarded = [("anthropic-beta", "web-search-2025-03-05")];
assert_eq!(
authenticate_with(&forwarded, api_key, env).unwrap(),
headers(&[forwarded[0], expected])
);
}
#[rstest]
#[case::no_credentials(&[], None, &[])]
#[case::empty_api_key(&[], Some(""), &[])]
#[case::whitespace_only_env_values(
&[],
None,
&[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")],
)]
#[case::unrelated_forwarded_headers(&[("anthropic-beta", "web-search-2025-03-05")], None, &[])]
fn missing_credentials_are_an_auth_error(
#[case] forwarded: &[(&str, &str)],
#[case] api_key: Option<&str>,
#[case] env: Env,
) {
assert!(matches!(
authenticate_with(forwarded, api_key, env),
Err(litellm_auth::Error::MissingApiKey {
provider: "Anthropic",
environment_variable: "ANTHROPIC_API_KEY",
})
));
}
#[rstest]
#[case::no_features(json!({}), &[])]
#[case::output_format(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
#[case::null_output_format(json!({"output_format": null}), &[])]
#[case::output_config_format(
json!({"output_config": {"format": {"type": "json_schema"}, "effort": "xhigh"}}),
&[beta::STRUCTURED_OUTPUT]
)]
#[case::null_output_config_format(json!({"output_config": {"format": null}}), &[])]
#[case::top_level_output_config_without_format(json!({"output_config": {"effort": "high"}}), &[])]
#[case::fast_speed(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
#[case::standard_speed(json!({"speed": "standard"}), &[])]
#[case::compaction_param(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
#[case::empty_compaction_param(json!({"compaction": {}}), &[beta::COMPACT_2026_09_04])]
#[case::signed_compaction_block_in_history(
json!({"messages": [
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": "sig"}]},
{"role": "user", "content": "Continue"},
]}),
&[beta::COMPACT_2026_09_04]
)]
#[case::unsigned_compaction_block_in_history(
json!({"messages": [
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": ""}]},
{"role": "user", "content": "Continue"},
]}),
&[]
)]
#[case::advisor_tool(
json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"}]}),
&[beta::ADVISOR_TOOL_2026_03_01]
)]
#[case::no_tools(json!({"tools": []}), &[])]
#[case::regex_tool_search(
json!({"tools": [{"type": "tool_search_tool_regex_20251119"}]}),
&[beta::ADVANCED_TOOL_USE_2025_11_20]
)]
#[case::bm25_tool_search(
json!({"tools": [{"type": "tool_search_tool_bm25_20251119"}]}),
&[beta::ADVANCED_TOOL_USE_2025_11_20]
)]
#[case::unrelated_server_tool(json!({"tools": [{"type": "web_search_20250305", "name": "web_search"}]}), &[])]
#[case::only_compact_edits(
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
&[beta::COMPACT_2026_01_12]
)]
#[case::only_other_edits(
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}}]}}),
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
)]
#[case::compact_and_other_edits(
json!({"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_tool_uses_20250919"}]}}),
&[beta::COMPACT_2026_01_12, beta::CONTEXT_MANAGEMENT_2025_06_27]
)]
#[case::edit_without_a_type(json!({"context_management": {"edits": [{}]}}), &[beta::CONTEXT_MANAGEMENT_2025_06_27])]
#[case::empty_edits(json!({"context_management": {"edits": []}}), &[])]
#[case::context_management_without_edits(json!({"context_management": {}}), &[])]
#[case::per_message_output_config(
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
&[beta::PER_TURN_CONTROL_2026_07_01]
)]
#[case::per_message_null_output_config(
json!({"messages": [{"role": "user", "content": "hi", "output_config": null}]}),
&[beta::PER_TURN_CONTROL_2026_07_01]
)]
fn feature_betas_follow_the_request(#[case] fields: Value, #[case] expected: &[&str]) {
assert_eq!(feature_betas(&request(fields)), expected);
}
#[rstest]
#[case::no_betas(&[("x-api-key", "k"), ("anthropic-version", "2023-06-01")], json!({}))]
#[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], json!({}))]
fn headers_without_any_beta_value_are_untouched(
#[case] input: &[(&str, &str)],
#[case] fields: Value,
) {
assert_eq!(
with_feature_betas(headers(input), &request(fields)),
headers(input)
);
}
#[rstest]
#[case::feature_beta_is_appended(
&[("x-api-key", "k")],
json!({"speed": "fast"}),
&[("x-api-key", "k"), ("anthropic-beta", beta::FAST_MODE_2026_02_01)],
)]
#[case::existing_betas_are_normalized_without_features(
&[("Anthropic-Beta", "web-search-2025-03-05, interleaved-thinking-2025-05-14 ,web-search-2025-03-05"), ("x-api-key", "k")],
json!({}),
&[("x-api-key", "k"), ("anthropic-beta", "interleaved-thinking-2025-05-14,web-search-2025-03-05")],
)]
#[case::existing_advisor_beta_is_kept_without_an_advisor_tool(
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
json!({"tools": []}),
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
)]
#[case::feature_already_sent_is_not_duplicated(
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
json!({"speed": "fast"}),
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
)]
fn feature_betas_merge_into_the_headers(
#[case] input: &[(&str, &str)],
#[case] fields: Value,
#[case] expected: &[(&str, &str)],
) {
assert_eq!(
with_feature_betas(headers(input), &request(fields)),
headers(expected)
);
}
#[test]
fn differently_cased_beta_header_is_replaced_by_one_sorted_header() {
let merged = with_feature_betas(
headers(&[("Anthropic-Beta", "interleaved-thinking-2025-05-14")]),
&request(
json!({"messages": [{"role": "system", "content": "env", "output_config": {"effort": "low"}}]}),
),
);
assert_eq!(
merged,
headers(&[(
"anthropic-beta",
&betas(&[
"interleaved-thinking-2025-05-14",
beta::PER_TURN_CONTROL_2026_07_01
])
)])
);
}
#[test]
fn every_beta_header_casing_is_unioned_into_one_header() {
let merged = with_feature_betas(
headers(&[
("anthropic-beta", "interleaved-thinking-2025-05-14"),
("Anthropic-Beta", "web-search-2025-03-05"),
]),
&request(json!({"speed": "fast"})),
);
assert_eq!(
merged,
headers(&[(
"anthropic-beta",
&betas(&[
beta::FAST_MODE_2026_02_01,
"interleaved-thinking-2025-05-14",
"web-search-2025-03-05"
])
)])
);
}
#[test]
fn unknown_client_betas_survive_alongside_the_added_one() {
let client_betas = [
"claude-code-20250219",
"interleaved-thinking-2025-05-14",
beta::CONTEXT_MANAGEMENT_2025_06_27,
beta::PER_TURN_CONTROL_2026_07_01,
"effort-2025-11-24",
];
let merged = with_feature_betas(
headers(&[("anthropic-beta", &betas(&client_betas))]),
&request(
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
),
);
assert_eq!(
merged,
headers(&[(
"anthropic-beta",
&betas(&[
"claude-code-20250219",
beta::CONTEXT_MANAGEMENT_2025_06_27,
"effort-2025-11-24",
"interleaved-thinking-2025-05-14",
beta::PER_TURN_CONTROL_2026_07_01,
])
)])
);
}
#[test]
fn every_feature_merges_with_the_oauth_beta_sorted_and_last() {
let oauth_headers = authenticate_with(&[], Some(OAUTH_TOKEN), &[]).unwrap();
let all_features = request(json!({
"compaction": {"enabled": true},
"output_format": {"type": "json_schema"},
"speed": "fast",
"tools": [{"type": "advisor_20260301"}, {"type": "tool_search_tool_bm25_20251119"}],
"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_thinking_20251015"}]},
"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}],
}));
assert_eq!(
with_feature_betas(oauth_headers, &all_features),
headers(&[
("authorization", OAUTH_BEARER),
BROWSER_ACCESS,
(
"anthropic-beta",
&betas(&[
beta::ADVANCED_TOOL_USE_2025_11_20,
beta::ADVISOR_TOOL_2026_03_01,
beta::COMPACT_2026_01_12,
beta::COMPACT_2026_09_04,
beta::CONTEXT_MANAGEMENT_2025_06_27,
beta::FAST_MODE_2026_02_01,
ANTHROPIC_OAUTH_BETA_HEADER,
beta::PER_TURN_CONTROL_2026_07_01,
beta::STRUCTURED_OUTPUT,
])
),
])
);
}
}

View file

@ -1,2 +1,5 @@
pub mod handler;
pub mod headers;
pub mod streaming_iterator;
pub mod thinking;
pub mod transformation;

View file

@ -1,9 +1,28 @@
use crate::base_llm::{
anthropic_messages::transformation::BaseAnthropicMessagesConfig, chat::transformation::Error,
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
use serde_json::{Map, Value, json};
use super::{
headers::{authenticate, with_feature_betas},
thinking::{ThinkingBudgets, ThinkingContext, translate_thinking},
};
use crate::{
anthropic::common_utils::{
AnthropicModelCapabilities, has_advisor_tool, strip_advisor_blocks,
strip_encrypted_reasoning_blocks,
},
base_llm::{
anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, Headers, MessagesTransformContext,
},
chat::transformation::Error,
},
};
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE";
const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL";
const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com";
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
@ -11,6 +30,26 @@ pub struct AnthropicMessagesConfig;
pub const ANTHROPIC_MESSAGES_CONFIG: AnthropicMessagesConfig = AnthropicMessagesConfig;
impl MessagesTransformContext {
pub fn new(capabilities: AnthropicModelCapabilities, drop_params: bool) -> Self {
Self::with_lookup(capabilities, drop_params, &ProcessEnvironment)
}
pub fn with_lookup(
capabilities: AnthropicModelCapabilities,
drop_params: bool,
env: &impl Lookup,
) -> Self {
Self {
thinking: ThinkingContext {
capabilities,
budgets: ThinkingBudgets::from_lookup(env),
},
drop_params,
}
}
}
impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
fn get_complete_url(
&self,
@ -21,6 +60,35 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
Ok(complete_anthropic_url(api_base, env_lookup))
}
fn transform_anthropic_messages_request(
&self,
request: AnthropicMessagesRequest,
context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
if request.max_tokens.is_none() {
return Err(Error::InvalidRequest(
"max_tokens is required for Anthropic /v1/messages API".to_string(),
));
}
let request = drop_unsupported_params(request, context)?;
let request = translate_thinking(request, &context.thinking)?;
let context_management = request
.context_management
.as_ref()
.and_then(map_openai_context_management_to_anthropic)
.or_else(|| request.context_management.clone());
let messages = if has_advisor_tool(request.tools.as_deref()) {
request.messages
} else {
strip_advisor_blocks(request.messages)
};
Ok(AnthropicMessagesRequest {
messages: strip_encrypted_reasoning_blocks(messages),
context_management,
..request
})
}
fn resolve_api_key(
&self,
api_key: Option<&str>,
@ -28,6 +96,113 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
) -> Result<String, Error> {
resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from)
}
fn secret_names(&self) -> &'static [&'static str] {
&[
ANTHROPIC_API_KEY_ENV,
ANTHROPIC_AUTH_TOKEN_ENV,
ANTHROPIC_API_BASE_ENV,
ANTHROPIC_BASE_URL_ENV,
]
}
fn authenticate(
&self,
headers: Headers,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Headers, Error> {
authenticate(headers, api_key, env_lookup).map_err(Error::from)
}
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
with_feature_betas(headers, request)
}
}
fn unsupported_param(model: &str, param: &str, value: &str, hint: &str) -> Error {
Error::InvalidRequest(format!(
"{model} does not support {param}={value}. {hint}To drop unsupported params, set `litellm.drop_params = True`."
))
}
fn drop_unsupported_params(
request: AnthropicMessagesRequest,
context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
let capabilities = &context.thinking.capabilities;
let model = request.model.clone();
let reject = |param: &str, value: String, hint: &str| -> Result<(), Error> {
if context.drop_params {
return Ok(());
}
Err(unsupported_param(&model, param, &value, hint))
};
let speed = match request.speed.as_deref() {
Some(speed) if !capabilities.supports_speed => {
reject("speed", format!("'{speed}'"), "")?;
None
}
_ => request.speed.clone(),
};
if capabilities.supports_sampling_params {
return Ok(AnthropicMessagesRequest { speed, ..request });
}
let temperature = match request.temperature {
Some(temperature) if temperature != 1.0 => {
reject(
"temperature",
json!(temperature).to_string(),
"Only temperature=1 is supported. ",
)?;
None
}
temperature => temperature,
};
if let Some(top_p) = request.top_p {
reject("top_p", json!(top_p).to_string(), "")?;
}
if let Some(top_k) = request.top_k {
reject("top_k", json!(top_k).to_string(), "")?;
}
Ok(AnthropicMessagesRequest {
speed,
temperature,
top_p: None,
top_k: None,
..request
})
}
pub fn map_openai_context_management_to_anthropic(context_management: &Value) -> Option<Value> {
match context_management {
Value::Object(edits) if edits.contains_key("edits") => Some(context_management.clone()),
Value::Array(entries) => {
let edits: Vec<Value> = entries
.iter()
.filter_map(Value::as_object)
.filter(|entry| entry.get("type").and_then(Value::as_str) == Some("compaction"))
.map(|entry| {
let trigger = entry.get("compact_threshold").and_then(Value::as_f64).map(
|threshold| json!({"type": "input_tokens", "value": threshold as i64}),
);
let passthrough = entry
.iter()
.filter(|(key, _)| !matches!(key.as_str(), "type" | "compact_threshold"))
.map(|(key, value)| (key.clone(), value.clone()));
Value::Object(
[("type".to_string(), json!("compact_20260112"))]
.into_iter()
.chain(trigger.map(|trigger| ("trigger".to_string(), trigger)))
.chain(passthrough)
.collect::<Map<String, Value>>(),
)
})
.collect();
(!edits.is_empty()).then(|| json!({"edits": edits}))
}
_ => None,
}
}
pub fn non_empty(value: Option<&str>) -> Option<&str> {
@ -64,70 +239,619 @@ pub fn resolve_anthropic_api_base(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> String {
let env = |name: &str| env_lookup(name).filter(|value| !value.trim().is_empty());
non_empty(api_base)
.map(str::to_string)
.or_else(|| env_lookup(ANTHROPIC_API_BASE_ENV).filter(|value| !value.trim().is_empty()))
.or_else(|| env(ANTHROPIC_API_BASE_ENV))
.or_else(|| env(ANTHROPIC_BASE_URL_ENV))
.unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string())
}
#[cfg(test)]
mod tests {
use std::process::Command;
use rstest::{fixture, rstest};
use super::*;
use crate::anthropic::common_utils::{ENCRYPTED_REASONING_SIGNATURE_PREFIX, beta};
#[test]
fn url_defaults_to_public_anthropic_endpoint() {
type Env = &'static [(&'static str, &'static str)];
const BOTH_BASE_ENVS: Env = &[
(ANTHROPIC_API_BASE_ENV, "https://api-base.example.com"),
(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com"),
];
const API_KEY_ENV: Env = &[(ANTHROPIC_API_KEY_ENV, "sk-env")];
const MISSING_API_KEY: &str =
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable";
const LOW_BUDGET_ENV: &str = "DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET";
const PROCESS_ENV_PROBE: &str = "LITELLM_MESSAGES_TRANSFORM_CONTEXT_PROBE";
fn merged(base: Value, fields: Value) -> Value {
Value::Object(
base.as_object()
.unwrap()
.clone()
.into_iter()
.chain(fields.as_object().unwrap().clone())
.collect(),
)
}
fn body(fields: Value) -> Value {
merged(
json!({
"model": "claude",
"max_tokens": 1024,
"messages": [{"role": "user", "content": "Hello"}]
}),
fields,
)
}
fn request(fields: Value) -> AnthropicMessagesRequest {
serde_json::from_value(body(fields)).unwrap()
}
fn no_env(_: &str) -> Option<String> {
None
}
fn env(vars: Env) -> impl Fn(&str) -> Option<String> {
move |name| {
vars.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| value.to_string())
}
}
fn headers(pairs: &[(&str, &str)]) -> Headers {
pairs
.iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect()
}
fn transform(
fields: Value,
capabilities: AnthropicModelCapabilities,
drop_params: bool,
) -> Result<Value, Error> {
ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(
request(fields),
&MessagesTransformContext::with_lookup(capabilities, drop_params, &no_env),
)
.map(|transformed| serde_json::to_value(transformed).unwrap())
}
fn invalid(message: &str) -> Result<Value, Error> {
Err(Error::InvalidRequest(message.to_string()))
}
fn advisor_history() -> Value {
json!([
{"role": "user", "content": "Build a worker pool."},
{"role": "assistant", "content": [
{"type": "text", "text": "Let me consult the advisor."},
{"type": "server_tool_use", "id": "srvtoolu_abc123", "name": "advisor", "input": {}},
{"type": "advisor_tool_result", "tool_use_id": "srvtoolu_abc123", "content": {"type": "advisor_result", "text": "Use channels."}},
{"type": "text", "text": "Here is the implementation."}
]}
])
}
#[fixture]
fn unmapped() -> AnthropicModelCapabilities {
AnthropicModelCapabilities::default()
}
#[fixture]
fn sampling_removed() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
supports_sampling_params: false,
..Default::default()
}
}
#[fixture]
fn fast_mode() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
supports_speed: true,
..Default::default()
}
}
#[rstest]
#[case::alone(json!({"max_tokens": null}))]
#[case::ahead_of_the_param_gate(json!({"max_tokens": null, "speed": "fast"}))]
fn missing_max_tokens_is_rejected(#[case] fields: Value, unmapped: AnthropicModelCapabilities) {
assert_eq!(
complete_anthropic_url(None, &|_| None),
"https://api.anthropic.com/v1/messages"
transform(fields, unmapped, false),
invalid("max_tokens is required for Anthropic /v1/messages API")
);
}
#[rstest]
#[case::sampling_params_on_a_sampling_model(
unmapped(),
false,
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
)]
#[case::sampling_params_on_a_sampling_model_under_drop_params(
unmapped(),
true,
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
)]
#[case::unit_temperature_on_a_sampling_removed_model(
sampling_removed(),
false,
json!({"temperature": 1.0})
)]
#[case::unit_temperature_on_a_sampling_removed_model_under_drop_params(
sampling_removed(),
true,
json!({"temperature": 1.0})
)]
#[case::speed_on_a_fast_mode_model(fast_mode(), false, json!({"speed": "fast"}))]
#[case::speed_on_a_fast_mode_model_under_drop_params(fast_mode(), true, json!({"speed": "fast"}))]
#[case::native_context_management_edits(unmapped(), false, json!({"context_management": {"edits": [{
"type": "clear_tool_uses_20250919",
"trigger": {"type": "input_tokens", "value": 30000},
"keep": {"type": "tool_uses", "value": 3},
"clear_at_least": {"type": "input_tokens", "value": 5000},
"exclude_tools": ["web_search"],
"clear_tool_inputs": false
}]}}))]
#[case::first_party_billing_header_system_block(unmapped(), false, json!({"system": [
{"type": "text", "text": "x-anthropic-billing-header: cc_version=1"},
{"type": "text", "text": "real system prompt"}
]}))]
#[case::anthropic_signed_reasoning_history(unmapped(), false, json!({"messages": [
{"role": "user", "content": "Solve it."},
{"role": "assistant", "content": [
{"type": "thinking", "thinking": "plan", "signature": "EqQBCkYIAxgCIkA_anthropic_signed"},
{"type": "redacted_thinking", "data": "EmwKAhgBEgy_anthropic_minted"},
{"type": "text", "text": "The answer."}
]}
]}))]
#[case::advisor_history_alongside_the_advisor_tool(unmapped(), false, json!({
"messages": advisor_history(),
"tools": [{"type": "advisor_20260301", "name": "advisor"}]
}))]
fn request_is_forwarded_unchanged(
#[case] capabilities: AnthropicModelCapabilities,
#[case] drop_params: bool,
#[case] fields: Value,
) {
assert_eq!(
transform(fields.clone(), capabilities, drop_params),
Ok(body(fields))
);
}
#[rstest]
#[case::temperature(sampling_removed(), json!({"temperature": 0.3}), json!({}))]
#[case::top_p(sampling_removed(), json!({"top_p": 0.9}), json!({}))]
#[case::top_k(sampling_removed(), json!({"top_k": 40}), json!({}))]
#[case::every_sampling_param_keeping_the_rest(
sampling_removed(),
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40, "stream": true}),
json!({"stream": true})
)]
#[case::speed_on_a_sampling_model(
unmapped(),
json!({"speed": "fast", "temperature": 0.5}),
json!({"temperature": 0.5})
)]
#[case::speed_on_a_sampling_removed_model(
sampling_removed(),
json!({"speed": "fast", "temperature": 1.0}),
json!({"temperature": 1.0})
)]
fn removed_params_are_dropped_under_drop_params(
#[case] capabilities: AnthropicModelCapabilities,
#[case] fields: Value,
#[case] expected: Value,
) {
assert_eq!(transform(fields, capabilities, true), Ok(body(expected)));
}
#[rstest]
#[case::temperature(
sampling_removed(),
json!({"temperature": 0.3}),
"claude does not support temperature=0.3. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::temperature_just_below_one(
sampling_removed(),
json!({"temperature": 0.99}),
"claude does not support temperature=0.99. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::whole_number_temperature_keeps_its_decimal(
sampling_removed(),
json!({"temperature": 2.0}),
"claude does not support temperature=2.0. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::top_p(
sampling_removed(),
json!({"top_p": 0.9}),
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::top_k(
sampling_removed(),
json!({"top_k": 5}),
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::top_k_next_to_unit_temperature(
sampling_removed(),
json!({"temperature": 1.0, "top_k": 5}),
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::temperature_ahead_of_top_k(
sampling_removed(),
json!({"temperature": 0.5, "top_k": 5}),
"claude does not support temperature=0.5. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::top_p_ahead_of_top_k(
sampling_removed(),
json!({"top_p": 0.9, "top_k": 5}),
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::speed(
unmapped(),
json!({"speed": "fast"}),
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::speed_ahead_of_sampling_params(
sampling_removed(),
json!({"speed": "fast", "temperature": 0.5}),
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
)]
fn removed_params_are_rejected_without_drop_params(
#[case] capabilities: AnthropicModelCapabilities,
#[case] fields: Value,
#[case] message: &str,
) {
assert_eq!(transform(fields, capabilities, false), invalid(message));
}
#[rstest]
#[case::compaction_threshold(
json!([{"type": "compaction", "compact_threshold": 200000}]),
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]}))
)]
#[case::other_keys_pass_through(
json!([{"type": "compaction", "compact_threshold": 150000, "instructions": "Focus on preserving code snippets"}]),
Some(json!({"edits": [{
"type": "compact_20260112",
"trigger": {"type": "input_tokens", "value": 150000},
"instructions": "Focus on preserving code snippets"
}]}))
)]
#[case::float_threshold_is_truncated(
json!([{"type": "compaction", "compact_threshold": 150000.9}]),
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
)]
#[case::compaction_without_threshold(
json!([{"type": "compaction"}]),
Some(json!({"edits": [{"type": "compact_20260112"}]}))
)]
#[case::non_numeric_threshold_is_dropped(
json!([{"type": "compaction", "compact_threshold": "150000"}]),
Some(json!({"edits": [{"type": "compact_20260112"}]}))
)]
#[case::non_object_entries_are_skipped(
json!([42, "compaction", null, [], {"type": "compaction", "compact_threshold": 1000}]),
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}}]}))
)]
#[case::only_compaction_entries_are_mapped_in_order(
json!([
{"type": "compaction", "compact_threshold": 1000},
{"type": "other", "compact_threshold": 5},
{"type": "compaction", "instructions": "second"}
]),
Some(json!({"edits": [
{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}},
{"type": "compact_20260112", "instructions": "second"}
]}))
)]
#[case::list_without_compaction(json!([{"type": "other"}]), None)]
#[case::empty_list(json!([]), None)]
#[case::anthropic_edits_pass_through(
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}),
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
)]
#[case::object_without_edits(json!({"type": "compaction"}), None)]
#[case::scalar(json!("compaction"), None)]
fn openai_context_management_maps_to_anthropic_edits(
#[case] context_management: Value,
#[case] expected: Option<Value>,
) {
assert_eq!(
map_openai_context_management_to_anthropic(&context_management),
expected
);
}
#[rstest]
#[case::openai_list_is_mapped(
json!([{"type": "compaction", "compact_threshold": 200000}]),
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]})
)]
#[case::unmappable_list_is_kept(json!([{"type": "other"}]), json!([{"type": "other"}]))]
#[case::unmappable_object_is_kept(json!({"type": "other"}), json!({"type": "other"}))]
fn context_management_reaches_the_wire(
#[case] context_management: Value,
#[case] expected: Value,
unmapped: AnthropicModelCapabilities,
) {
assert_eq!(
transform(
json!({"context_management": context_management}),
unmapped,
false
),
Ok(body(json!({"context_management": expected})))
);
}
#[rstest]
#[case::without_tools(json!({}))]
#[case::with_only_other_tools(json!({"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}]}))]
fn advisor_history_is_stripped_without_the_advisor_tool(
#[case] tools: Value,
unmapped: AnthropicModelCapabilities,
) {
let stripped = json!([
{"role": "user", "content": "Build a worker pool."},
{"role": "assistant", "content": [
{"type": "text", "text": "Let me consult the advisor."},
{"type": "text", "text": "Here is the implementation."}
]}
]);
assert_eq!(
transform(
merged(tools.clone(), json!({"messages": advisor_history()})),
unmapped,
false
),
Ok(body(merged(tools, json!({"messages": stripped}))))
);
}
#[rstest]
fn bridge_minted_reasoning_is_stripped_from_the_wire(unmapped: AnthropicModelCapabilities) {
let messages = json!([
{"role": "user", "content": "Solve it."},
{"role": "assistant", "content": [
{"type": "thinking", "thinking": "plan", "signature": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_1")},
{"type": "redacted_thinking", "data": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_2")},
{"type": "text", "text": "The answer."}
]},
{"role": "user", "content": "And the next one?"}
]);
assert_eq!(
transform(json!({"messages": messages}), unmapped, false),
Ok(body(json!({"messages": [
{"role": "user", "content": "Solve it."},
{"role": "assistant", "content": [{"type": "text", "text": "The answer."}]},
{"role": "user", "content": "And the next one?"}
]})))
);
}
#[test]
fn url_appends_messages_suffix_to_custom_base() {
fn thinking_is_translated_with_the_context_budgets() {
let context = MessagesTransformContext::with_lookup(
AnthropicModelCapabilities {
supports_reasoning: true,
..Default::default()
},
false,
&env(&[(LOW_BUDGET_ENV, "2000")]),
);
let transformed = ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(
request(json!({"max_tokens": 4096, "reasoning_effort": "low"})),
&context,
)
.map(|transformed| serde_json::to_value(transformed).unwrap());
assert_eq!(
complete_anthropic_url(Some("https://proxy.internal"), &|_| None),
"https://proxy.internal/v1/messages"
transformed,
Ok(body(json!({
"max_tokens": 4096,
"thinking": {"type": "enabled", "budget_tokens": 2000}
})))
);
}
#[test]
fn url_leaves_complete_messages_endpoint_untouched() {
fn new_reads_thinking_budgets_from_the_process_environment() {
if std::env::var_os(PROCESS_ENV_PROBE).is_some() {
assert_eq!(
MessagesTransformContext::new(sampling_removed(), true),
MessagesTransformContext {
thinking: ThinkingContext {
capabilities: sampling_removed(),
budgets: ThinkingBudgets {
low: 2000,
..ThinkingBudgets::default()
},
},
drop_params: true,
}
);
return;
}
let (_, test_path) = concat!(
module_path!(),
"::new_reads_thinking_budgets_from_the_process_environment"
)
.split_once("::")
.unwrap();
let other_tiers = ["MINIMAL", "MEDIUM", "HIGH", "XHIGH", "MAX"]
.map(|tier| format!("DEFAULT_REASONING_EFFORT_{tier}_THINKING_BUDGET"));
let output = other_tiers
.iter()
.fold(
Command::new(std::env::current_exe().unwrap()),
|mut command, name| {
command.env_remove(name);
command
},
)
.args([test_path, "--exact"])
.env(PROCESS_ENV_PROBE, "1")
.env(LOW_BUDGET_ENV, "2000")
.output()
.unwrap();
let stdout = String::from_utf8_lossy(&output.stdout);
assert!(
output.status.success() && stdout.contains("1 passed"),
"{stdout}{}",
String::from_utf8_lossy(&output.stderr)
);
}
#[rstest]
#[case::public_endpoint_by_default(None, &[], "https://api.anthropic.com")]
#[case::explicit_api_base_beats_env(
Some("https://explicit.example.com"),
BOTH_BASE_ENVS,
"https://explicit.example.com"
)]
#[case::explicit_api_base_is_trimmed(
Some(" https://explicit.example.com "),
&[],
"https://explicit.example.com"
)]
#[case::blank_api_base_falls_back_to_env(
Some(" "),
BOTH_BASE_ENVS,
"https://api-base.example.com"
)]
#[case::api_base_env_beats_base_url_env(None, BOTH_BASE_ENVS, "https://api-base.example.com")]
#[case::base_url_env_without_api_base_env(
None,
&[(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
"https://base-url.example.com"
)]
#[case::blank_api_base_env_falls_back_to_base_url_env(
None,
&[(ANTHROPIC_API_BASE_ENV, " \t "), (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
"https://base-url.example.com"
)]
#[case::blank_envs_fall_back_to_public_endpoint(
None,
&[(ANTHROPIC_API_BASE_ENV, ""), (ANTHROPIC_BASE_URL_ENV, " ")],
"https://api.anthropic.com"
)]
fn api_base_resolution(
#[case] api_base: Option<&str>,
#[case] vars: Env,
#[case] expected: &str,
) {
assert_eq!(resolve_anthropic_api_base(api_base, &env(vars)), expected);
}
#[rstest]
#[case::public_endpoint(None, &[], "https://api.anthropic.com/v1/messages")]
#[case::base_url_env(
None,
&[(ANTHROPIC_BASE_URL_ENV, "https://custom.example.com")],
"https://custom.example.com/v1/messages"
)]
#[case::custom_base(Some("https://proxy.internal"), &[], "https://proxy.internal/v1/messages")]
#[case::trailing_slash(Some("https://proxy.internal/"), &[], "https://proxy.internal/v1/messages")]
#[case::complete_endpoint(
Some("https://proxy.internal/v1/messages"),
&[],
"https://proxy.internal/v1/messages"
)]
#[case::complete_endpoint_with_trailing_slash(
Some("https://proxy.internal/v1/messages/"),
&[],
"https://proxy.internal/v1/messages"
)]
fn complete_url_ends_in_the_messages_path(
#[case] api_base: Option<&str>,
#[case] vars: Env,
#[case] expected: &str,
) {
assert_eq!(
complete_anthropic_url(Some("https://proxy.internal/v1/messages"), &|_| None),
"https://proxy.internal/v1/messages"
ANTHROPIC_MESSAGES_CONFIG.get_complete_url(api_base, "claude", &env(vars)),
Ok(expected.to_string())
);
}
#[rstest]
#[case::param_beats_env(Some("sk-param"), API_KEY_ENV, Ok("sk-param"))]
#[case::param_is_trimmed(Some(" sk-param "), &[], Ok("sk-param"))]
#[case::blank_param_falls_back_to_env(Some(" "), API_KEY_ENV, Ok("sk-env"))]
#[case::env_without_param(None, API_KEY_ENV, Ok("sk-env"))]
#[case::blank_env_is_missing(None, &[(ANTHROPIC_API_KEY_ENV, " ")], Err(MISSING_API_KEY))]
#[case::nothing_is_missing(None, &[], Err(MISSING_API_KEY))]
fn api_key_resolution(
#[case] api_key: Option<&str>,
#[case] vars: Env,
#[case] expected: Result<&str, &str>,
) {
assert_eq!(
resolve_anthropic_api_key(api_key, &env(vars)).map_err(|error| error.to_string()),
expected.map(str::to_string).map_err(str::to_string)
);
}
#[test]
fn url_falls_back_to_env_base() {
let with_env = |key: &str| {
(key == ANTHROPIC_API_BASE_ENV).then(|| "https://env.anthropic".to_string())
};
fn config_reports_a_missing_key_as_an_auth_error() {
assert_eq!(
complete_anthropic_url(Some(" "), &with_env),
"https://env.anthropic/v1/messages"
ANTHROPIC_MESSAGES_CONFIG.resolve_api_key(None, &no_env),
Err(Error::Auth(litellm_auth::Error::MissingApiKey {
provider: "Anthropic",
environment_variable: ANTHROPIC_API_KEY_ENV,
}))
);
}
#[test]
fn api_key_prefers_param_then_env_then_errors() {
fn config_authenticates_with_the_anthropic_auth_token() {
assert_eq!(
resolve_anthropic_api_key(Some("sk-param"), &|_| None).unwrap(),
"sk-param"
ANTHROPIC_MESSAGES_CONFIG.authenticate(
vec![],
None,
&env(&[("ANTHROPIC_AUTH_TOKEN", "auth-token")])
),
Ok(headers(&[("authorization", "Bearer auth-token")]))
);
let with_env = |key: &str| (key == ANTHROPIC_API_KEY_ENV).then(|| "sk-env".to_string());
}
#[test]
fn config_requests_the_betas_the_request_features_need() {
assert_eq!(
resolve_anthropic_api_key(Some(" "), &with_env).unwrap(),
"sk-env"
);
assert_eq!(
resolve_anthropic_api_key(None, &|_| None)
.expect_err("missing key")
.to_string(),
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
ANTHROPIC_MESSAGES_CONFIG.request_headers(
headers(&[("x-api-key", "sk")]),
&request(json!({"speed": "fast"}))
),
headers(&[
("x-api-key", "sk"),
("anthropic-beta", beta::FAST_MODE_2026_02_01)
])
);
}
#[rstest]
#[case::absent(None, None)]
#[case::blank(Some(" \t "), None)]
#[case::padded(Some(" value "), Some("value"))]
fn non_empty_trims_and_drops_blank_values(
#[case] value: Option<&str>,
#[case] expected: Option<&str>,
) {
assert_eq!(non_empty(value), expected);
}
#[test]
fn auth_strategy_and_default_headers_match_anthropic() {
assert_eq!(
@ -142,4 +866,26 @@ mod tests {
]
);
}
#[test]
fn secret_names_cover_every_credential_and_base_lookup() {
let requested = std::cell::RefCell::new(Vec::<String>::new());
let record = |name: &str| -> Option<String> {
requested.borrow_mut().push(name.to_string());
None
};
let _ = ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record);
let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
let requested = requested.into_inner();
assert!(!requested.is_empty());
let undeclared: Vec<&String> = requested
.iter()
.filter(|name| {
!ANTHROPIC_MESSAGES_CONFIG
.secret_names()
.contains(&name.as_str())
})
.collect();
assert_eq!(undeclared, Vec::<&String>::new());
}
}

View file

@ -1,5 +1,6 @@
pub mod batches;
pub mod chat;
pub mod common_utils;
pub mod count_tokens;
pub mod experimental_pass_through;

View file

@ -4,14 +4,15 @@ use litellm_types::llms::anthropic_messages::{
},
anthropic_response::AnthropicMessagesResponse,
};
use serde_json::{Map, Value};
use crate::{
anthropic::experimental_pass_through::messages::transformation::{
ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty,
},
base_llm::{
anthropic_messages::transformation::{BaseAnthropicMessagesConfig, MessagesAuthStrategy},
anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, Headers, MessagesAuthStrategy, MessagesTransformContext,
},
chat::transformation::Error,
},
};
@ -21,7 +22,6 @@ const AZURE_API_BASE_ENV: &str = "AZURE_API_BASE";
const ANTHROPIC_PATH_SEGMENT: &str = "/anthropic";
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
const SYSTEM_ROLE: &str = "system";
const TEXT_BLOCK_TYPE: &str = "text";
pub struct AzureAnthropicMessagesConfig {
anthropic: AnthropicMessagesConfig,
@ -45,6 +45,7 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
fn transform_anthropic_messages_request(
&self,
request: AnthropicMessagesRequest,
context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
let mut request = fold_system_role_messages(request);
if let Some(system) = request.system.as_mut() {
@ -54,7 +55,8 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
.messages
.iter_mut()
.for_each(strip_scope_from_message);
self.anthropic.transform_anthropic_messages_request(request)
self.anthropic
.transform_anthropic_messages_request(request, context)
}
fn transform_anthropic_messages_response(
@ -74,6 +76,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
resolve_azure_api_key(api_key, env_lookup)
}
fn secret_names(&self) -> &'static [&'static str] {
&[AZURE_API_KEY_ENV, AZURE_API_BASE_ENV]
}
fn auth_strategy(&self) -> MessagesAuthStrategy {
self.anthropic.auth_strategy()
}
@ -85,6 +91,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
self.anthropic.default_headers()
}
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
self.anthropic.request_headers(headers, request)
}
}
pub fn resolve_azure_api_key(
@ -143,17 +153,7 @@ fn strip_scope_from_message(message: &mut AnthropicMessage) {
}
fn text_content_block(text: String) -> ContentBlock {
let extra = Map::from_iter([
(
"type".to_string(),
Value::String(TEXT_BLOCK_TYPE.to_string()),
),
("text".to_string(), Value::String(text)),
]);
ContentBlock {
cache_control: None,
extra,
}
ContentBlock::text(text)
}
fn content_into_blocks(content: MessageContent) -> Vec<ContentBlock> {
@ -202,6 +202,7 @@ mod tests {
use serde_json::json;
use super::*;
use crate::anthropic::common_utils::AnthropicModelCapabilities;
fn request_from(value: serde_json::Value) -> AnthropicMessagesRequest {
serde_json::from_value(value).expect("valid request")
@ -346,7 +347,7 @@ mod tests {
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request)
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.expect("request transforms"),
);
@ -373,10 +374,13 @@ mod tests {
"messages": [{"role": "user", "content": "hi"}]
}));
let once = AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request)
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.expect("request transforms");
let twice = AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(once.clone())
.transform_anthropic_messages_request(
once.clone(),
&MessagesTransformContext::default(),
)
.expect("request transforms");
assert_eq!(once, twice);
assert_eq!(to_value(once)["system"], json!("plain string system"));
@ -408,9 +412,21 @@ mod tests {
"inference_geo": "us",
"litellm_metadata": {"trace": "abc"}
});
let context = MessagesTransformContext::with_lookup(
AnthropicModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_legacy_thinking: true,
supports_output_config: true,
supports_speed: true,
..Default::default()
},
false,
&|_: &str| None,
);
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request_from(body.clone()))
.transform_anthropic_messages_request(request_from(body.clone()), &context)
.expect("request transforms"),
);
assert_eq!(transformed, body);
@ -430,7 +446,7 @@ mod tests {
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request)
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.expect("request transforms"),
);
@ -460,7 +476,7 @@ mod tests {
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request)
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.expect("request transforms"),
);
@ -485,9 +501,21 @@ mod tests {
{"role": "assistant", "content": "hello"}
]
});
let context = MessagesTransformContext::with_lookup(
AnthropicModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_legacy_thinking: true,
supports_output_config: true,
supports_speed: true,
..Default::default()
},
false,
&|_: &str| None,
);
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request_from(body.clone()))
.transform_anthropic_messages_request(request_from(body.clone()), &context)
.expect("request transforms"),
);
assert_eq!(transformed, body);
@ -500,6 +528,57 @@ mod tests {
assert!(err.is_data());
}
#[rstest::rstest]
#[case::compact_context_management_edit(
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
&[],
&[("x-api-key", "k"), ("anthropic-beta", "compact-2026-01-12")]
)]
#[case::forwarded_beta_merged_with_structured_output(
json!({"output_config": {"format": {"type": "json_schema"}}}),
&[("anthropic-beta", "web-search-2025-03-05")],
&[("x-api-key", "k"), ("anthropic-beta", "structured-outputs-2025-11-13,web-search-2025-03-05")]
)]
#[case::no_feature_needs_a_beta(json!({}), &[], &[("x-api-key", "k")])]
fn request_headers_carry_the_anthropic_feature_betas(
#[case] fields: serde_json::Value,
#[case] forwarded: &[(&str, &str)],
#[case] expected: &[(&str, &str)],
) {
let pairs = |pairs: &[(&str, &str)]| -> Vec<(String, String)> {
pairs
.iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect()
};
let serde_json::Value::Object(fields) = fields else {
panic!("case fields are an object")
};
let request = request_from(serde_json::Value::Object(
[
("model".to_string(), json!("claude-sonnet")),
("max_tokens".to_string(), json!(16)),
(
"messages".to_string(),
json!([{"role": "user", "content": "hi"}]),
),
]
.into_iter()
.chain(fields)
.collect(),
));
assert_eq!(
AZURE_ANTHROPIC_MESSAGES_CONFIG.request_headers(
pairs(&[("x-api-key", "k")])
.into_iter()
.chain(pairs(forwarded))
.collect(),
&request
),
pairs(expected)
);
}
#[test]
fn transform_response_passes_through() {
let response: AnthropicMessagesResponse = serde_json::from_value(json!({
@ -521,4 +600,26 @@ mod tests {
assert_eq!(value["stop_sequence"], json!(null));
assert_eq!(value["content"][0]["text"], json!("hello"));
}
#[test]
fn secret_names_cover_every_credential_and_base_lookup() {
let requested = std::cell::RefCell::new(Vec::<String>::new());
let record = |name: &str| -> Option<String> {
requested.borrow_mut().push(name.to_string());
None
};
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record);
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
let requested = requested.into_inner();
assert!(!requested.is_empty());
let undeclared: Vec<&String> = requested
.iter()
.filter(|name| {
!AZURE_ANTHROPIC_MESSAGES_CONFIG
.secret_names()
.contains(&name.as_str())
})
.collect();
assert_eq!(undeclared, Vec::<&String>::new());
}
}

View file

@ -1,8 +1,14 @@
use litellm_http::request::{has_bearer_auth, has_header};
use litellm_types::llms::anthropic_messages::{
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
};
use crate::base_llm::chat::transformation::Error;
use crate::{
anthropic::experimental_pass_through::messages::thinking::ThinkingContext,
base_llm::chat::transformation::Error,
};
pub type Headers = Vec<(String, String)>;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessagesAuthStrategy {
@ -19,6 +25,12 @@ impl MessagesAuthStrategy {
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct MessagesTransformContext {
pub thinking: ThinkingContext,
pub drop_params: bool,
}
pub trait BaseAnthropicMessagesConfig: Sync {
fn get_complete_url(
&self,
@ -30,6 +42,7 @@ pub trait BaseAnthropicMessagesConfig: Sync {
fn transform_anthropic_messages_request(
&self,
request: AnthropicMessagesRequest,
_context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
Ok(request)
}
@ -48,6 +61,8 @@ pub trait BaseAnthropicMessagesConfig: Sync {
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error>;
fn secret_names(&self) -> &'static [&'static str];
fn auth_strategy(&self) -> MessagesAuthStrategy {
MessagesAuthStrategy::Header("x-api-key")
}
@ -56,10 +71,225 @@ pub trait BaseAnthropicMessagesConfig: Sync {
false
}
fn authenticate(
&self,
headers: Headers,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Headers, Error> {
let strategy = self.auth_strategy();
if has_header(&headers, strategy.header_name())
|| (self.accepts_bearer_auth() && has_bearer_auth(&headers))
{
return Ok(headers);
}
let api_key = self.resolve_api_key(api_key, env_lookup)?;
let auth_header = match strategy {
MessagesAuthStrategy::Bearer => {
("authorization".to_string(), format!("Bearer {api_key}"))
}
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
};
Ok(headers.into_iter().chain([auth_header]).collect())
}
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
&[
("anthropic-version", "2023-06-01"),
("content-type", "application/json"),
]
}
fn request_headers(&self, headers: Headers, _request: &AnthropicMessagesRequest) -> Headers {
headers
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use super::*;
const X_API_KEY: MessagesAuthStrategy = MessagesAuthStrategy::Header("x-api-key");
struct StubConfig {
strategy: MessagesAuthStrategy,
accepts_bearer: bool,
}
impl BaseAnthropicMessagesConfig for StubConfig {
fn secret_names(&self) -> &'static [&'static str] {
&[]
}
fn get_complete_url(
&self,
_api_base: Option<&str>,
_model: &str,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
Ok(String::new())
}
fn resolve_api_key(
&self,
api_key: Option<&str>,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
api_key
.map(str::to_string)
.ok_or(Error::MissingField("api_key"))
}
fn auth_strategy(&self) -> MessagesAuthStrategy {
self.strategy
}
fn accepts_bearer_auth(&self) -> bool {
self.accepts_bearer
}
}
struct DefaultsConfig;
impl BaseAnthropicMessagesConfig for DefaultsConfig {
fn secret_names(&self) -> &'static [&'static str] {
&[]
}
fn get_complete_url(
&self,
_api_base: Option<&str>,
_model: &str,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
Ok(String::new())
}
fn resolve_api_key(
&self,
api_key: Option<&str>,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
api_key
.map(str::to_string)
.ok_or(Error::MissingField("api_key"))
}
}
#[test]
fn default_config_adds_its_key_next_to_a_forwarded_bearer() {
assert_eq!(
DefaultsConfig.authenticate(
headers(&[("authorization", "Bearer forwarded")]),
Some("sk"),
&|_| None
),
Ok(headers(&[
("authorization", "Bearer forwarded"),
("x-api-key", "sk")
]))
);
}
#[test]
fn default_request_headers_are_the_given_headers() {
let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({
"model": "claude",
"max_tokens": 16,
"speed": "fast",
"messages": [{"role": "user", "content": "hi"}]
}))
.unwrap();
assert_eq!(
DefaultsConfig.request_headers(headers(&[("x-api-key", "sk")]), &request),
headers(&[("x-api-key", "sk")])
);
}
fn headers(pairs: &[(&str, &str)]) -> Headers {
pairs
.iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect()
}
#[rstest]
#[case::own_header_is_kept(
X_API_KEY,
false,
headers(&[("x-api-key", "forwarded")]),
None,
Ok(headers(&[("x-api-key", "forwarded")]))
)]
#[case::own_header_in_any_casing_is_kept(
X_API_KEY,
false,
headers(&[("X-Api-Key", "forwarded")]),
None,
Ok(headers(&[("X-Api-Key", "forwarded")]))
)]
#[case::accepted_bearer_is_kept(
X_API_KEY,
true,
headers(&[("authorization", "Bearer forwarded")]),
None,
Ok(headers(&[("authorization", "Bearer forwarded")]))
)]
#[case::bearer_the_provider_does_not_accept_gets_the_key_too(
X_API_KEY,
false,
headers(&[("authorization", "Bearer forwarded")]),
Some("sk"),
Ok(headers(&[("authorization", "Bearer forwarded"), ("x-api-key", "sk")]))
)]
#[case::blank_bearer_gets_the_key(
X_API_KEY,
true,
headers(&[("authorization", "Bearer ")]),
Some("sk"),
Ok(headers(&[("authorization", "Bearer "), ("x-api-key", "sk")]))
)]
#[case::key_goes_in_the_provider_header(
X_API_KEY,
false,
headers(&[("content-type", "application/json")]),
Some("sk"),
Ok(headers(&[("content-type", "application/json"), ("x-api-key", "sk")]))
)]
#[case::key_goes_in_a_bearer(
MessagesAuthStrategy::Bearer,
false,
headers(&[]),
Some("sk"),
Ok(headers(&[("authorization", "Bearer sk")]))
)]
#[case::bearer_strategy_keeps_a_forwarded_authorization(
MessagesAuthStrategy::Bearer,
false,
headers(&[("authorization", "Bearer forwarded")]),
None,
Ok(headers(&[("authorization", "Bearer forwarded")]))
)]
#[case::missing_key_is_an_error(
X_API_KEY,
false,
headers(&[]),
None,
Err(Error::MissingField("api_key"))
)]
fn default_authenticate_applies_the_key_unless_a_credential_is_forwarded(
#[case] strategy: MessagesAuthStrategy,
#[case] accepts_bearer: bool,
#[case] forwarded: Headers,
#[case] api_key: Option<&str>,
#[case] expected: Result<Headers, Error>,
) {
let config = StubConfig {
strategy,
accepts_bearer,
};
assert_eq!(config.authenticate(forwarded, api_key, &|_| None), expected);
}
}

View file

@ -302,7 +302,3 @@ fn has_blank_text(message: &ChatMessage) -> bool {
}),
}
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;

View file

@ -1,7 +1,11 @@
use serde_json::json;
use super::*;
use crate::base_llm::chat::transformation::Error;
use litellm_llms::{
anthropic::chat::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG,
base_llm::chat::transformation::{
BaseConfig, Error, ProviderChatResponseData, RequestAuth, Unsupported,
},
};
use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse};
use serde_json::{Map, Value, json};
fn messages(value: Value) -> Vec<ChatMessage> {
serde_json::from_value(value).expect("valid messages")
@ -205,7 +209,8 @@ fn declines_tool_calls_tool_results_and_multimodal_content() {
);
assert_eq!(
reason(
json!([{"role": "user", "content": [
json!([
{"role": "user", "content": [
{"type": "image_url", "image_url": {"url": "https://x/y.png"}}
]}]),
json!({})
@ -214,7 +219,8 @@ fn declines_tool_calls_tool_results_and_multimodal_content() {
);
assert_eq!(
reason(
json!([{"role": "user", "content": [
json!([
{"role": "user", "content": [
{"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}}
]}]),
json!({})

View file

@ -1,7 +1,11 @@
use serde_json::json;
use super::*;
use crate::base_llm::chat::transformation::Error;
use litellm_llms::{
base_llm::chat::transformation::{
BaseConfig, Error, ProviderChatResponseData, RequestAuth, Unsupported,
},
bedrock::chat::converse_transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG,
};
use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse};
use serde_json::{Map, Value, json};
fn messages(value: Value) -> Vec<ChatMessage> {
serde_json::from_value(value).expect("valid messages")

View file

@ -2,9 +2,11 @@ use bytes::Bytes;
use litellm_core::messages::{
Error,
route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput},
types::MessagesShaping,
};
use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py};
use litellm_http::transport::Error as TransportError;
use litellm_types::utils::ProviderSpecificHeaders;
use pyo3::{
exceptions::{PyException, PyValueError},
gc::{PyTraverseError, PyVisit},
@ -18,9 +20,10 @@ use crate::{
marshal::{optional_timeout, python_timeout_seconds},
};
/// The Anthropic Messages body fields a caller may pass besides `model` and `messages`,
/// as `AnthropicMessagesRequestOptionalParams` declares them.
const BODY_FIELDS: [&str; 20] = [
const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.messages.route_host";
const REQUEST_ERROR_MARKER: &str = "messages_request_error";
const BODY_FIELDS: [&str; 22] = [
"max_tokens",
"metadata",
"stop_sequences",
@ -35,14 +38,46 @@ const BODY_FIELDS: [&str; 20] = [
"top_p",
"mcp_servers",
"context_management",
"compaction",
"container",
"output_format",
"speed",
"output_config",
"cache_control",
"reasoning_effort",
"safeguards",
];
fn merge_headers(
forwarded: Option<Map<String, Value>>,
extra_headers: Option<Map<String, Value>>,
) -> Option<Map<String, Value>> {
let merged: Map<String, Value> = forwarded
.into_iter()
.flatten()
.chain(extra_headers.into_iter().flatten())
.collect();
(!merged.is_empty()).then_some(merged)
}
fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
match error {
Error::Transport(TransportError::Http { status, body }) => {
let error = RustUpstreamError::new_err((status, body));
error
.value(py)
.setattr("headers", Vec::<(String, String)>::new())?;
Ok(error)
}
Error::InvalidRequest(message) => {
let error = PyValueError::new_err(message);
error.value(py).setattr(REQUEST_ERROR_MARKER, true)?;
Ok(error)
}
other => Ok(messages_error_to_pyerr(other)),
}
}
/// The Python side of the Messages route: projects the prepared arguments and builds the
/// public response, chunks and exceptions.
pub(super) struct MessagesRouteHost {
@ -84,19 +119,65 @@ impl MessagesRouteHost {
.map(|value| python_timeout_seconds(py, value.unbind()))
.transpose()?
.flatten();
let custom_llm_provider = string("custom_llm_provider")?;
let shaping = self.shaping(py, &model, custom_llm_provider.as_deref(), arguments)?;
Ok(MessagesCall {
model,
body,
api_key: string("api_key")?,
api_base: string("api_base")?,
custom_llm_provider: string("custom_llm_provider")?,
extra_headers: argument("extra_headers")?
.map(|value| from_py(&value))
.transpose()?,
extra_headers: self.merged_headers(py, arguments)?,
provider_specific_header: self.provider_specific_header(py, arguments)?,
custom_llm_provider,
timeout: optional_timeout(timeout),
shaping,
})
}
fn merged_headers(
&self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<Option<Map<String, Value>>> {
let request = self.request.bind(py);
let mapping = |name: &str| -> PyResult<Option<Map<String, Value>>> {
lookup(arguments, request, name)?
.filter(|value| !value.is_none())
.map(|value| from_py(&value))
.transpose()
};
Ok(merge_headers(
mapping("headers")?,
mapping("extra_headers")?,
))
}
fn provider_specific_header(
&self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<Option<ProviderSpecificHeaders>> {
lookup(arguments, self.request.bind(py), "provider_specific_header")?
.filter(|value| !value.is_none())
.map(|value| from_py(&value))
.transpose()
}
fn shaping(
&self,
py: Python<'_>,
model: &str,
custom_llm_provider: Option<&str>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<MessagesShaping> {
let projected = py.import(ROUTE_HOST_MODULE)?.getattr("shaping")?.call1((
model,
custom_llm_provider,
arguments,
))?;
from_py(&projected)
}
fn provider(&self, py: Python<'_>) -> String {
self.request
.bind(py)
@ -112,7 +193,7 @@ impl MessagesRouteHost {
return error;
}
let mapped = py
.import("litellm.rust_bridge.messages.route_host")
.import(ROUTE_HOST_MODULE)
.and_then(|module| module.getattr("map_failure"))
.and_then(|map| map.call1((error.value(py), self.request.bind(py), self.provider(py))))
.and_then(|mapped| {
@ -148,7 +229,7 @@ impl RouteHost for MessagesRouteHost {
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {
match response {
MessagesOutput::Message(message) => py
.import("litellm.rust_bridge.messages.route_host")?
.import(ROUTE_HOST_MODULE)?
.getattr("response")?
.call1((to_py(py, message.as_ref())?,))
.map(Bound::unbind),
@ -161,17 +242,12 @@ impl RouteHost for MessagesRouteHost {
}
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
let native = match error {
Error::Transport(TransportError::Http { status, body }) => {
let error = RustUpstreamError::new_err((status, body));
error
.value(py)
.setattr("headers", Vec::<(String, String)>::new())?;
error
}
other => messages_error_to_pyerr(other),
};
Ok(self.map_failure(py, native))
if let Error::Secret(source) = &error
&& let Some(original) = crate::secrets::python_error(py, source.source_error())
{
return Ok(original);
}
Ok(self.map_failure(py, native_error(py, error)?))
}
fn host_error(error: &PyErr) -> Error {
@ -184,3 +260,62 @@ impl RouteHost for MessagesRouteHost {
visit.call(&self.request)
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
fn map(value: Value) -> Map<String, Value> {
serde_json::from_value(value).unwrap()
}
#[rstest]
#[case::extra_over_forwarded(
Some(json!({"X-Priority": "forwarded", "X-Forwarded-Only": "keep"})),
Some(json!({"X-Priority": "extra", "X-Extra-Only": "also-keep"})),
Some(json!({"X-Priority": "extra", "X-Forwarded-Only": "keep", "X-Extra-Only": "also-keep"})),
)]
#[case::only_forwarded(Some(json!({"X-Forwarded": "yes"})), None, Some(json!({"X-Forwarded": "yes"})))]
#[case::only_extra_headers(
None,
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
)]
#[case::nothing(None, Some(json!({})), None)]
fn headers_merge_forwarded_then_extra(
#[case] forwarded: Option<Value>,
#[case] extra_headers: Option<Value>,
#[case] expected: Option<Value>,
) {
assert_eq!(
merge_headers(forwarded.map(map), extra_headers.map(map)),
expected.map(map)
);
}
#[rstest]
#[case::rejected_request(Error::InvalidRequest("does not support top_k=5".into()), true)]
#[case::unresolvable_provider(Error::InvalidProvider("openai".into()), false)]
#[case::upstream_failure(
Error::Transport(TransportError::Http { status: 400, body: "bad".into() }),
false,
)]
fn only_request_rejections_carry_the_request_error_marker(
#[case] error: Error,
#[case] marked: bool,
) {
Python::initialize();
Python::attach(|py| {
let native = native_error(py, error).unwrap();
let marker = native
.value(py)
.getattr_opt(REQUEST_ERROR_MARKER)
.unwrap()
.map(|value| value.extract::<bool>().unwrap());
assert_eq!(marker.unwrap_or(false), marked);
});
}
}

View file

@ -39,11 +39,12 @@ fn run_messages(
"the Rust Messages route does not serve this provider",
));
}
let secrets = crate::secrets::source(py)?;
run_legacy_call(
py,
SURFACE,
PublicCall::capture(&request, &args, &kwargs)?,
crate::logger::LoggedMachine::new(messages_machine()),
crate::logger::LoggedMachine::new(messages_machine(secrets)),
MessagesRouteHost::new(request.unbind()),
asynchronous,
)

View file

@ -8,3 +8,6 @@ repository.workspace = true
[dependencies]
serde.workspace = true
serde_json.workspace = true
[dev-dependencies]
rstest.workspace = true

View file

@ -17,12 +17,48 @@ pub enum MessageContent {
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct ContentBlock {
#[serde(rename = "type", default, skip_serializing_if = "Option::is_none")]
pub block_type: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub text: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub thinking: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub signature: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub data: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_use_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub content: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_specific_fields: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_control: Option<CacheControl>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
impl ContentBlock {
pub fn text(text: impl Into<String>) -> Self {
Self {
block_type: Some("text".to_string()),
text: Some(text.into()),
..Self::default()
}
}
pub fn is_type(&self, block_type: &str) -> bool {
self.block_type.as_deref() == Some(block_type)
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct CacheControl {
#[serde(rename = "type", skip_serializing_if = "Option::is_none")]
@ -85,6 +121,126 @@ pub struct AnthropicMessagesRequest {
pub speed: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub inference_geo: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_effort: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub compaction: Option<Value>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
impl AnthropicMessage {
pub fn blocks(&self) -> &[ContentBlock] {
match &self.content {
MessageContent::Blocks(blocks) => blocks,
MessageContent::Text(_) => &[],
}
}
pub fn with_blocks(self, blocks: Vec<ContentBlock>) -> Self {
Self {
content: MessageContent::Blocks(blocks),
..self
}
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
fn round_trip<T: serde::de::DeserializeOwned + Serialize>(value: &Value) -> Value {
let parsed: T = serde_json::from_value(value.clone()).unwrap();
serde_json::to_value(parsed).unwrap()
}
#[rstest]
#[case::text(json!({"type": "text", "text": "hi"}))]
#[case::text_with_citations_and_cache_control(json!({
"type": "text",
"text": "hi",
"citations": [{"type": "char_location", "cited_text": "x"}],
"cache_control": {"type": "ephemeral", "ttl": "1h", "scope": "global", "future": 1}
}))]
#[case::image(json!({"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "AA=="}}))]
#[case::thinking(json!({"type": "thinking", "thinking": "hmm", "signature": "sig"}))]
#[case::redacted_thinking(json!({"type": "redacted_thinking", "data": "opaque"}))]
#[case::tool_use(json!({"type": "tool_use", "id": "toolu_1", "name": "f", "input": {"q": [1, null]}}))]
#[case::tool_result_with_text(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": "ok", "is_error": false}))]
#[case::tool_result_with_blocks(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": [{"type": "text", "text": "ok"}]}))]
#[case::web_search_result_with_nulls(json!({
"type": "web_search_tool_result",
"tool_use_id": "srvtoolu_1",
"content": [{"type": "web_search_result", "url": "u", "page_age": null, "encrypted_content": ""}]
}))]
#[case::provider_specific_fields(json!({"type": "tool_use", "id": "t", "name": "f", "input": {}, "provider_specific_fields": {"x": 1}}))]
#[case::untyped(json!({"unknown": {"nested": true}}))]
fn content_block_round_trips_unchanged(#[case] block: Value) {
assert_eq!(round_trip::<ContentBlock>(&block), block);
}
#[test]
fn text_constructor_serializes_as_a_text_block() {
assert_eq!(
serde_json::to_value(ContentBlock::text("hello")).unwrap(),
json!({"type": "text", "text": "hello"})
);
}
#[rstest]
#[case::same_type(json!({"type": "tool_use"}), "tool_use", true)]
#[case::other_type(json!({"type": "tool_result"}), "tool_use", false)]
#[case::prefix_of_type(json!({"type": "tool_use"}), "tool", false)]
#[case::no_type(json!({"text": "x"}), "text", false)]
fn is_type_matches_the_exact_block_type(
#[case] block: Value,
#[case] block_type: &str,
#[case] expected: bool,
) {
let block: ContentBlock = serde_json::from_value(block).unwrap();
assert_eq!(block.is_type(block_type), expected);
}
#[rstest]
#[case::string_content(json!({"role": "user", "content": "hi"}), vec![])]
#[case::block_content(
json!({"role": "user", "content": [{"type": "text", "text": "a"}, {"type": "text", "text": "b"}]}),
vec![ContentBlock::text("a"), ContentBlock::text("b")],
)]
fn message_blocks_list_only_block_content(
#[case] message: Value,
#[case] expected: Vec<ContentBlock>,
) {
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
assert_eq!(message.blocks(), expected.as_slice());
}
#[rstest]
#[case::replaces_string_content(json!({"role": "assistant", "content": "old", "name": "kept"}))]
#[case::replaces_block_content(json!({"role": "assistant", "content": [{"type": "text", "text": "old"}], "name": "kept"}))]
fn with_blocks_replaces_content_and_keeps_the_rest(#[case] message: Value) {
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
assert_eq!(
serde_json::to_value(message.with_blocks(vec![ContentBlock::text("new")])).unwrap(),
json!({"role": "assistant", "content": [{"type": "text", "text": "new"}], "name": "kept"})
);
}
#[rstest]
#[case::minimal(json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]}))]
#[case::reasoning_effort_compaction_and_unknown_fields(json!({
"model": "m",
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
"max_tokens": 8,
"reasoning_effort": "high",
"compaction": {"type": "auto"},
"safeguards": [{"type": "dangerous_tool_use", "classifier_context": {"v": 1}}],
"metadata": {"user_id": "u"}
}))]
fn request_round_trips_unchanged(#[case] request: Value) {
assert_eq!(round_trip::<AnthropicMessagesRequest>(&request), request);
}
}

View file

@ -9,8 +9,6 @@ pub struct AnthropicMessagesResponse {
pub role: String,
pub model: String,
pub content: Vec<Value>,
// Anthropic always includes stop_reason / stop_sequence, null until the turn
// ends; serialize them even when None so callers see the same shape as Python.
pub stop_reason: Option<String>,
pub stop_sequence: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
@ -20,3 +18,61 @@ pub struct AnthropicMessagesResponse {
#[serde(flatten)]
pub extra: Map<String, Value>,
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
fn response(
stop_reason: Option<&str>,
stop_sequence: Option<&str>,
usage: Option<Value>,
container: Option<Value>,
) -> AnthropicMessagesResponse {
AnthropicMessagesResponse {
id: "msg_1".to_string(),
message_type: "message".to_string(),
role: "assistant".to_string(),
model: "claude".to_string(),
content: vec![],
stop_reason: stop_reason.map(str::to_string),
stop_sequence: stop_sequence.map(str::to_string),
usage,
container,
extra: Map::new(),
}
}
#[rstest]
#[case::turn_in_progress(None, None, json!(null), json!(null))]
#[case::ended_on_end_turn(Some("end_turn"), None, json!("end_turn"), json!(null))]
#[case::ended_on_stop_sequence(Some("stop_sequence"), Some("###"), json!("stop_sequence"), json!("###"))]
fn stop_fields_are_always_serialized(
#[case] stop_reason: Option<&str>,
#[case] stop_sequence: Option<&str>,
#[case] expected_reason: Value,
#[case] expected_sequence: Value,
) {
let body: Value = serde_json::to_value(response(stop_reason, stop_sequence, None, None))
.expect("serializable");
assert_eq!(body.get("stop_reason"), Some(&expected_reason));
assert_eq!(body.get("stop_sequence"), Some(&expected_sequence));
}
#[rstest]
#[case::absent(None, None)]
#[case::present(Some(json!({"input_tokens": 1})), Some(json!({"id": "c_1"})))]
fn usage_and_container_are_omitted_only_when_none(
#[case] usage: Option<Value>,
#[case] container: Option<Value>,
) {
let body: Value =
serde_json::to_value(response(None, None, usage.clone(), container.clone()))
.expect("serializable");
assert_eq!(body.get("usage").cloned(), usage);
assert_eq!(body.get("container").cloned(), container);
}
}

View file

@ -3,6 +3,21 @@ use serde_json::{Map, Value};
use crate::llms::openai::{ChatCompletionThinkingBlock, ChatCompletionToolCallChunk};
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct ProviderSpecificHeader {
#[serde(default)]
pub custom_llm_provider: String,
#[serde(default)]
pub extra_headers: Map<String, Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ProviderSpecificHeaders {
One(ProviderSpecificHeader),
Many(Vec<ProviderSpecificHeader>),
}
/// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python
/// path reports so cost tracking sees the same numbers on either path.
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
@ -41,7 +56,7 @@ pub struct ChatCompletionsChoice {
///
/// There is deliberately no `id`: Python mints the `chatcmpl-…` id on the
/// `ModelResponse` it already created, and echoing the provider's own id here
/// would change it. Pinned by `response_carries_no_id` in `tests.rs`.
/// would change it. Pinned by `response_carries_no_id` in the Anthropic chat transformation tests.
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ChatCompletionsResponse {
pub created: u64,

View file

@ -11,7 +11,10 @@ from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo
from litellm.llms.bedrock.batches.transformation import titan_embedding_usage_from_batch_output
from litellm.llms.vertex_ai.batches.transformation import vertex_prompt_tokens_details
from litellm.llms.vertex_ai.batches.transformation import (
is_native_vertex_batch_output_row,
native_vertex_batch_row_stats,
)
from litellm.types.llms.openai import Batch
from litellm.types.utils import ModelInfo, Usage
from litellm.utils import token_counter
@ -31,6 +34,20 @@ class BatchCostUsageResult:
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
def _uses_native_vertex_output(
custom_llm_provider: str,
model_name: str | None,
first_row: Mapping[str, object] | None,
) -> bool:
if custom_llm_provider != "vertex_ai":
return False
if model_name and getattr(litellm, "disable_vertex_batch_output_transformation", False):
return True
return first_row is not None and is_native_vertex_batch_output_row(first_row)
_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"})
@ -66,12 +83,9 @@ async def calculate_batch_cost_and_usage(
deployment-specific pricing (e.g. input_cost_per_token_batches)
is used instead of the global cost map.
"""
if (
custom_llm_provider == "vertex_ai"
and model_name
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
):
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
first_row: Final = file_content_dictionary[0] if file_content_dictionary else None
if _uses_native_vertex_output(custom_llm_provider, model_name, first_row):
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name, model_info=model_info)
return _aggregate_batch_cost_usage_models(
entries=file_content_dictionary,
@ -126,11 +140,11 @@ async def _handle_completed_batch(
)
output_file_result: Final = (
calculate_vertex_ai_batch_cost_and_usage(_get_file_content_as_dictionary(file_content), model_name)
if (
custom_llm_provider == "vertex_ai"
and model_name
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
calculate_vertex_ai_batch_cost_and_usage(
_iter_batch_output_entries(file_content), model_name, model_info=model_info
)
if _uses_native_vertex_output(
custom_llm_provider, model_name, next(_iter_batch_output_entries(file_content), None)
)
else _aggregate_batch_cost_usage_models(
entries=_iter_batch_output_entries(file_content),
@ -332,69 +346,36 @@ def _aggregate_batch_cost_usage_models(
def calculate_vertex_ai_batch_cost_and_usage(
vertex_ai_batch_responses: list[dict],
vertex_ai_batch_responses: Iterable[dict],
model_name: str | None = None,
model_info: ModelInfo | None = None,
) -> BatchCostUsageResult:
"""
Calculate both cost and usage from raw Vertex AI batch responses.
Used only when ``litellm.disable_vertex_batch_output_transformation = True``.
In that case the GCS predictions.jsonl is returned as-is, with each line in
the native Vertex format:
{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}}}
usageMetadata contains promptTokenCount, candidatesTokenCount, totalTokenCount.
A row with no ``response`` is counted as failed - the same signal already
used to skip it from cost/usage aggregation, since Vertex batch prediction
output doesn't establish a distinct error shape in this (non-default) path.
Cost and usage of a native Vertex predictions.jsonl, one
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
generateContent row or `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
embedding row per line. `model_name` (the deployment model) prices every row, else each row's own
`modelVersion` does; a row without a usable response counts as failed.
"""
from litellm.cost_calculator import batch_cost_calculator
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
total_prompt_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
total_completion_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
total_tokens = 0
prompt_tokens = 0
completion_tokens = 0
successful_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
failed_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
actual_model_name: Final = model_name or "gemini-2.0-flash-001"
for response in vertex_ai_batch_responses:
response_body = response.get("response")
if response_body is None:
failed_requests += 1
continue
successful_requests += 1
usage_metadata = response_body.get("usageMetadata", {})
_prompt = usage_metadata.get("promptTokenCount", 0) or 0
_completion = usage_metadata.get("candidatesTokenCount", 0) or 0
_total = usage_metadata.get("totalTokenCount", 0) or (_prompt + _completion)
line_usage = Usage(
prompt_tokens=_prompt,
completion_tokens=_completion,
total_tokens=_total,
prompt_tokens_details=vertex_prompt_tokens_details(usage_metadata),
row_stats: Final = tuple(
native_vertex_batch_row_stats(
row,
model_name,
model_info=model_info,
calculate_usage=VertexGeminiConfig._calculate_usage,
cost_calculator=batch_cost_calculator,
)
try:
p_cost, c_cost = batch_cost_calculator(
usage=line_usage,
model=actual_model_name,
custom_llm_provider="vertex_ai",
)
total_prompt_cost += p_cost
total_completion_cost += c_cost
except Exception as e:
verbose_logger.debug("vertex_ai batch cost calculation error for line: %s", str(e))
prompt_tokens += _prompt
completion_tokens += _completion
total_tokens += _total
for row in vertex_ai_batch_responses
)
priced: Final = tuple(stats for stats in row_stats if stats is not None)
total_prompt_cost: Final = sum(stats.prompt_cost for stats in priced)
total_completion_cost: Final = sum(stats.completion_cost for stats in priced)
prompt_tokens: Final = sum(stats.usage.prompt_tokens for stats in priced)
completion_tokens: Final = sum(stats.usage.completion_tokens for stats in priced)
total_tokens: Final = sum(stats.total_tokens for stats in priced)
total_cost: Final = total_prompt_cost + total_completion_cost
verbose_logger.info(
"vertex_ai batch cost: cost=%s, prompt=%d, completion=%d, total=%d, successful=%d, failed=%d",
@ -402,8 +383,8 @@ def calculate_vertex_ai_batch_cost_and_usage(
prompt_tokens,
completion_tokens,
total_tokens,
successful_requests,
failed_requests,
len(priced),
len(row_stats) - len(priced),
)
return BatchCostUsageResult(
@ -413,9 +394,13 @@ def calculate_vertex_ai_batch_cost_and_usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
),
models=[actual_model_name],
successful_requests=successful_requests,
failed_requests=failed_requests,
models=(
[model_name]
if model_name
else list(dict.fromkeys(stats.model for stats in priced if stats.model is not None))
),
successful_requests=len(priced),
failed_requests=len(row_stats) - len(priced),
prompt_cost=total_prompt_cost,
completion_cost=total_completion_cost,
)

View file

@ -176,6 +176,22 @@ def create_file(
if logging_obj is None:
raise ValueError("logging_obj is required")
client: Final = kwargs.get("client")
if litellm_params_dict.get("passthrough") is True and (
custom_llm_provider != "vertex_ai" or purpose != "batch"
):
raise litellm.exceptions.BadRequestError(
message=(
"`passthrough=True` uploads the file bytes unchanged for a native Vertex AI batch, so it needs "
f"custom_llm_provider='vertex_ai' and purpose='batch', got '{custom_llm_provider}' and '{purpose}'."
),
model="n/a",
llm_provider=custom_llm_provider or "n/a",
response=httpx.Response(
status_code=400,
content="passthrough needs a vertex_ai batch",
request=httpx.Request(method="create_file", url="https://github.com/BerriAI/litellm"),
),
)
### TIMEOUT LOGIC ###
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600

View file

@ -118,6 +118,7 @@ class SlackAlerting(CustomBatchLogger):
self.default_webhook_url = default_webhook_url
self.flush_lock = asyncio.Lock()
self.periodic_started = False
self._periodic_flush_task: asyncio.Task[None] | None = None
self.hanging_request_check = AlertingHangingRequestCheck(
slack_alerting_object=self,
)
@ -129,6 +130,12 @@ class SlackAlerting(CustomBatchLogger):
self.digest_lock = asyncio.Lock()
super().__init__(**kwargs, flush_lock=self.flush_lock)
def _ensure_periodic_flush_task(self) -> None:
if self.periodic_started and (self._periodic_flush_task is None or not self._periodic_flush_task.done()):
return
self._periodic_flush_task = asyncio.create_task(self.periodic_flush())
self.periodic_started = True
def update_values(
self,
alerting: list | None = None,
@ -141,17 +148,14 @@ class SlackAlerting(CustomBatchLogger):
):
if alerting is not None:
self.alerting = alerting
asyncio.create_task(self.periodic_flush())
self.periodic_started = True
self._ensure_periodic_flush_task()
if alerting_threshold is not None:
self.alerting_threshold = alerting_threshold
if alert_types is not None:
self.alert_types = alert_types
if alerting_args is not None:
self.alerting_args = SlackAlertingArgs(**alerting_args)
if not self.periodic_started:
asyncio.create_task(self.periodic_flush())
self.periodic_started = True
self._ensure_periodic_flush_task()
if alert_type_config is not None:
for key, val in alert_type_config.items():
self.alert_type_config[key] = AlertTypeConfig(**val) if isinstance(val, dict) else val
@ -1446,9 +1450,8 @@ Model Info:
return
# Start periodic flush if not already started
if not self.periodic_started and self.alerting is not None and len(self.alerting) > 0:
asyncio.create_task(self.periodic_flush())
self.periodic_started = True
if self.alerting is not None and len(self.alerting) > 0:
self._ensure_periodic_flush_task()
if "webhook" in self.alerting and alert_type == "budget_alerts" and user_info is not None:
await self.send_webhook_alert(webhook_event=user_info)

View file

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

View file

@ -62,10 +62,10 @@ if TYPE_CHECKING:
from litellm.proxy.proxy_server import UserAPIKeyAuth as _UserAPIKeyAuth
Span = _Span | Any
Tracer = _Tracer | Any
Context = _Context | Any
SpanExporter = _SpanExporter | Any
UserAPIKeyAuth = _UserAPIKeyAuth | Any
Tracer = _Tracer
Context = _Context
SpanExporter = _SpanExporter
UserAPIKeyAuth = _UserAPIKeyAuth
ManagementEndpointLoggingPayload = _ManagementEndpointLoggingPayload | Any
else:
Span = Any
@ -2730,7 +2730,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
self.handle_callback_failure(callback_name=self.callback_name or "opentelemetry")
verbose_logger.exception("OpenTelemetry logging error in set_attributes %s", str(e))
def _cast_as_primitive_value_type(self, value) -> str | bool | int | float:
def _cast_as_primitive_value_type(self, value: object) -> str | bool | int | float:
"""
Casts the value to a primitive OTEL type if it is not already a primitive type.

View file

@ -2778,6 +2778,13 @@ class PrometheusLogger(CustomLogger):
- increment deployment failure responses metric
- increment deployment total requests metric
Both counters also carry a model_group label. When a deployment was
actually selected, model_group is the router-resolved value and is
trusted as-is. On a pre-routing reject (no deployment selected), it
is caller-supplied via litellm_params.metadata and is bounded with
_bounded_requested_model_label the same way requested_model is, so an
unrecognized value cannot mint unbounded label series.
Args:
request_kwargs: dict
@ -2844,6 +2851,7 @@ class PrometheusLogger(CustomLogger):
label_api_base = api_base
label_api_provider = llm_provider
label_requested_model = model_group or litellm_model_name
label_model_group = model_group
else:
label_litellm_model_name = ""
label_model_id = ""
@ -2852,6 +2860,7 @@ class PrometheusLogger(CustomLogger):
label_requested_model = (
_bounded_requested_model_label(litellm_model_name or model_group, router_originated=True) or ""
)
label_model_group = _bounded_requested_model_label(model_group, router_originated=True)
enum_values: Final = UserAPIKeyLabelValues(
litellm_model_name=label_litellm_model_name,
@ -2861,6 +2870,7 @@ class PrometheusLogger(CustomLogger):
exception_status=exception_status,
exception_class=(self._get_exception_class_name(exception) if exception else None),
requested_model=label_requested_model,
model_group=label_model_group,
hashed_api_key=hashed_api_key,
api_key_alias=api_key_alias,
user_email=user_email,
@ -2912,9 +2922,21 @@ class PrometheusLogger(CustomLogger):
model_id: str | None,
api_base: str | None,
llm_provider: str | None,
model_group: str | None,
):
"""
Set the deployment TPM and RPM limits metrics
Args:
model_info: the deployment's static model_info config (id, tpm, rpm, etc.)
litellm_params: the deployment's litellm_params, as a tpm/rpm fallback source
litellm_model_name: the resolved deployment model name
model_id: the deployment's model_id
api_base: the deployment's api_base
llm_provider: the deployment's custom_llm_provider
model_group: the router-resolved model_group the deployment belongs to,
from the caller's already-resolved enum_values.model_group (trusted,
not caller-supplied at this call site)
"""
tpm: Final = model_info.get("tpm") or litellm_params.get("tpm")
rpm: Final = model_info.get("rpm") or litellm_params.get("rpm")
@ -2927,6 +2949,7 @@ class PrometheusLogger(CustomLogger):
model_id=model_id,
api_base=api_base,
api_provider=llm_provider,
model_group=model_group,
),
)
self.litellm_deployment_tpm_limit.labels(**_labels).set(tpm)
@ -2939,6 +2962,7 @@ class PrometheusLogger(CustomLogger):
model_id=model_id,
api_base=api_base,
api_provider=llm_provider,
model_group=model_group,
),
)
self.litellm_deployment_rpm_limit.labels(**_labels).set(rpm)
@ -3058,6 +3082,7 @@ class PrometheusLogger(CustomLogger):
model_id=model_id,
api_base=api_base,
llm_provider=llm_provider,
model_group=enum_values.model_group,
)
remaining_requests: int | None = None

View file

@ -765,3 +765,30 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo
RESPONSE_COST_HEADER: cost,
}
hidden_params["additional_headers"] = merged
_HIDDEN_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_PROVIDER_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str])
def set_provider_response_headers_in_hidden_params(
response: _CarriesHiddenParams, headers: httpx.Headers | Mapping[str, str]
) -> None:
hidden_params: Final = response._hidden_params # pyright: ignore[reportPrivateUsage] # no public accessor
existing_additional_headers: Final[object] = hidden_params.get("additional_headers")
raw_headers: Final[dict[str, str]] = dict(headers) # mutable-ok: stored as the plain-dict hidden param
additional_headers: Final[dict[str, object]] = { # mutable-ok: assigned into the plain-dict hidden params
**process_response_headers(raw_headers),
**(existing_additional_headers if isinstance(existing_additional_headers, Mapping) else _NO_HEADERS),
}
hidden_params["headers"] = raw_headers
hidden_params["additional_headers"] = additional_headers
def get_provider_response_headers_from_hidden_params(response: object) -> Mapping[str, str] | None:
hidden_params: Final[object] = getattr(response, "_hidden_params", None)
try:
validated: Final = _HIDDEN_PARAMS_ADAPTER.validate_python(hidden_params)
return _PROVIDER_HEADERS_ADAPTER.validate_python(validated.get("headers"))
except ValidationError:
return None

View file

@ -72,6 +72,7 @@ from litellm.litellm_core_utils.classifier_logging import (
is_classifier_call,
)
from litellm.litellm_core_utils.core_helpers import (
get_provider_response_headers_from_hidden_params,
is_expected_client_error,
reconstruct_model_name,
set_response_cost_in_hidden_params,
@ -1566,14 +1567,14 @@ class Logging(LiteLLMLoggingBaseClass):
attr = "debug"
if json_logs:
callattr = getattr(verbose_logger, attr)
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
callattr(
"RAW RESPONSE:\n{}\n\n".format(
self.model_call_details.get("original_response", self.model_call_details)
),
)
else:
callattr = getattr(verbose_logger, attr)
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
callattr(
"RAW RESPONSE:\n{}\n\n".format(
self.model_call_details.get("original_response", self.model_call_details)
@ -2353,6 +2354,15 @@ class Logging(LiteLLMLoggingBaseClass):
)
return logging_result
def _surface_response_headers_from_result(self, logging_result: object) -> None:
existing: Final[object] = self.model_call_details.get("response_headers")
if existing is not None:
return
headers: Final = get_provider_response_headers_from_hidden_params(logging_result)
if headers is None:
return
self.model_call_details["response_headers"] = headers
def _merge_hidden_params_from_response_into_metadata(self, logging_result: object) -> None:
"""
Copy response._hidden_params into litellm_params.metadata['hidden_params'].
@ -2386,6 +2396,7 @@ class Logging(LiteLLMLoggingBaseClass):
build_logging_payload: bool = True,
):
"""Resolve hidden params, compute response cost, and emit the standard logging payload."""
self._surface_response_headers_from_result(logging_result)
hidden_params: Final = getattr(logging_result, "_hidden_params", {})
if hidden_params:
if self.model_call_details.get("litellm_params") is not None:
@ -2788,6 +2799,7 @@ class Logging(LiteLLMLoggingBaseClass):
if complete_streaming_response is not None:
verbose_logger.debug("Logging Details LiteLLM-Success Call streaming complete")
self.model_call_details["complete_streaming_response"] = complete_streaming_response
self._surface_response_headers_from_result(complete_streaming_response)
self.model_call_details["response_cost"] = self._response_cost_calculator(
result=complete_streaming_response
)
@ -3302,6 +3314,7 @@ class Logging(LiteLLMLoggingBaseClass):
print_verbose("Async success callbacks: Got a complete streaming response")
self.model_call_details["async_complete_streaming_response"] = complete_streaming_response
self._surface_response_headers_from_result(complete_streaming_response)
try:
if self.model_call_details.get("cache_hit", False) is True:
@ -5882,7 +5895,7 @@ class StandardLoggingPayloadSetup:
base_model: str | None,
custom_pricing: bool | None,
custom_llm_provider: str | None,
init_response_obj: Any | BaseModel | dict,
init_response_obj: object,
api_base: str | None = None,
) -> StandardLoggingModelInformation:
model_cost_name: Final = _select_model_name_for_cost_calc(
@ -5915,9 +5928,7 @@ class StandardLoggingPayloadSetup:
return model_cost_information
@staticmethod
def get_final_response_obj(
response_obj: dict, init_response_obj: Any | BaseModel | dict, kwargs: dict
) -> dict | str | list | None:
def get_final_response_obj(response_obj: dict, init_response_obj: object, kwargs: dict) -> dict | str | list | None:
"""
Get final response object after redacting the message input/output from logging
"""
@ -6360,16 +6371,19 @@ def _get_status_fields(
def _extract_response_obj_and_hidden_params(
init_response_obj: Any | BaseModel | dict,
init_response_obj: object,
original_exception: Exception | None,
) -> tuple[dict, dict | None]:
"""Extract response_obj and hidden_params from init_response_obj."""
hidden_params: dict | None = None
hidden_params: dict | None = (
getattr(init_response_obj, "_hidden_params", None)
if isinstance(init_response_obj, BaseModel | HttpxBinaryResponseContent)
else None
)
if init_response_obj is None:
response_obj = {}
elif isinstance(init_response_obj, BaseModel):
response_obj = init_response_obj.model_dump()
hidden_params = getattr(init_response_obj, "_hidden_params", None)
elif isinstance(init_response_obj, dict):
response_obj = init_response_obj
elif isinstance(init_response_obj, HttpxBinaryResponseContent):

View file

@ -446,7 +446,7 @@ def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: st
async def _afetch_and_extract_template(
model: str, chat_template: Any | None, get_config_fn, get_template_fn
model: str, chat_template: str | None, get_config_fn, get_template_fn
) -> tuple[str, str, str]:
"""
Async version: Fetch template and tokens from HuggingFace.
@ -500,7 +500,7 @@ async def _afetch_and_extract_template(
def _fetch_and_extract_template(
model: str, chat_template: Any | None, get_config_fn, get_template_fn
model: str, chat_template: str | None, get_config_fn, get_template_fn
) -> tuple[str, str, str]:
"""
Sync version: Fetch template and tokens from HuggingFace.

View file

@ -1218,6 +1218,8 @@ class BedrockModelInfo(BaseLLMModelInfo):
alt_model: Final = BedrockModelInfo.get_non_litellm_routing_model_name(model=model)
if base_model in litellm.bedrock_converse_models or alt_model in litellm.bedrock_converse_models:
return "converse"
if _OPENAI_FAMILY_MODEL_RE.search(base_model):
return "converse"
return "invoke"
@staticmethod

View file

@ -12,6 +12,7 @@ from typing import (
Literal,
NamedTuple,
Optional,
Protocol,
TypedDict,
TypeVar,
Union,
@ -24,6 +25,7 @@ import httpx
from httpx import USE_CLIENT_DEFAULT
from httpx._types import FileContent
from openai.types.file_deleted import FileDeleted
from typing_extensions import ReadOnly
import litellm
import litellm.litellm_core_utils
@ -43,6 +45,7 @@ from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
SUBTITLE_RESPONSE_FORMATS,
synthesize_subtitle_document,
)
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
from litellm.litellm_core_utils.realtime_errors import (
@ -206,6 +209,7 @@ if TYPE_CHECKING:
FakeAnthropicMessagesStreamIterator,
)
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai_evals import (
CancelEvalResponse,
CancelRunResponse,
@ -221,6 +225,21 @@ if TYPE_CHECKING:
else:
LiteLLMLoggingObj = Any
class _RealtimeClientWebSocket(Protocol):
async def send_text(self, data: str) -> None: ...
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
class _ResponsesClientWebSocket(Protocol):
async def send_text(self, data: str) -> None: ...
async def receive_text(self) -> str: ...
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
_ResponseT = TypeVar("_ResponseT")
@ -237,6 +256,17 @@ class _MediaUploadKwargs(TypedDict, total=False):
timeout: float | httpx.Timeout
class _SignedBodyKwargs(TypedDict, total=False):
data: ReadOnly[bytes]
json: ReadOnly[dict[str, object]]
def _signed_body_kwargs(*, signed_body: bytes | None, data: dict[str, object]) -> _SignedBodyKwargs:
if signed_body is not None:
return {"data": signed_body}
return {"json": data}
def _google_genai_streaming_hidden_params(
*,
api_base: str,
@ -318,7 +348,9 @@ def _mask_presigned_request_headers(transformed_request: bytes | str | dict) ->
}
def _aws_signing_overrides(optional_params: Mapping[str, Any], litellm_params: Mapping[str, Any]) -> Mapping[str, Any]:
def _aws_signing_overrides(
optional_params: Mapping[str, object], litellm_params: Mapping[str, object]
) -> Mapping[str, object]:
return MappingProxyType(
{
key: litellm_params[key]
@ -1430,6 +1462,7 @@ class BaseLLMHTTPHandler:
transformed: Final = provider_config.transform_audio_transcription_response(
raw_response=response,
)
set_provider_response_headers_in_hidden_params(transformed, response.headers)
if not provider_config.supports_subtitle_synthesis:
return transformed
requested_format: Final = optional_params.get("response_format")
@ -2739,7 +2772,7 @@ class BaseLLMHTTPHandler:
stream=stream,
fake_stream=fake_stream,
)
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
## LOGGING
logging_obj.pre_call(
@ -2926,7 +2959,7 @@ class BaseLLMHTTPHandler:
stream=stream,
fake_stream=fake_stream,
)
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
## LOGGING
logging_obj.pre_call(
@ -4540,7 +4573,7 @@ class BaseLLMHTTPHandler:
api_key=litellm_params.api_key,
model=model,
)
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
## LOGGING
logging_obj.pre_call(
@ -4634,7 +4667,7 @@ class BaseLLMHTTPHandler:
api_key=litellm_params.api_key,
model=model,
)
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
## LOGGING
logging_obj.pre_call(
@ -6186,6 +6219,7 @@ class BaseLLMHTTPHandler:
"BasePassthroughConfig",
"BaseContainerConfig",
BaseEvalsAPIConfig,
BaseRealtimeHTTPConfig,
],
):
received_status_code: Final = (
@ -6300,7 +6334,7 @@ class BaseLLMHTTPHandler:
async def async_realtime(
self,
model: str,
websocket: Any,
websocket: _RealtimeClientWebSocket,
logging_obj: LiteLLMLoggingObj,
provider_config: BaseRealtimeConfig,
headers: dict,
@ -6308,7 +6342,7 @@ class BaseLLMHTTPHandler:
api_key: str | None = None,
client: Any | None = None,
timeout: float | None = None,
user_api_key_dict: Any | None = None,
user_api_key_dict: object | None = None,
litellm_metadata: dict[str, object] | None = None,
query_params: RealtimeQueryParams | None = None,
):
@ -6483,7 +6517,7 @@ class BaseLLMHTTPHandler:
request_data: dict[str, object],
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout,
provider_config: Any | None = None,
provider_config: BaseRealtimeHTTPConfig | None = None,
model: str | None = None,
extra_headers: dict[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
@ -6555,7 +6589,7 @@ class BaseLLMHTTPHandler:
sdp_body: bytes,
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout,
provider_config: Any | None = None,
provider_config: BaseRealtimeHTTPConfig | None = None,
model: str | None = None,
session_config: dict[str, object] | None = None,
extra_headers: dict[str, object] | None = None,
@ -6633,13 +6667,13 @@ class BaseLLMHTTPHandler:
async def async_responses_websocket(
self,
model: str,
websocket: Any,
websocket: _ResponsesClientWebSocket,
logging_obj: LiteLLMLoggingObj,
responses_api_provider_config: BaseResponsesAPIConfig | None,
api_base: str | None = None,
api_key: str | None = None,
timeout: float | None = None,
user_api_key_dict: Any | None = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
litellm_metadata: dict[str, object] | None = None,
custom_llm_provider: str | None = None,
first_message: str | None = None,
@ -6928,11 +6962,13 @@ class BaseLLMHTTPHandler:
provider_config=image_edit_provider_config,
)
return image_edit_provider_config.transform_image_edit_response(
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
return image_edit_response
async def async_image_edit_handler(
self,
@ -7027,11 +7063,13 @@ class BaseLLMHTTPHandler:
provider_config=image_edit_provider_config,
)
return image_edit_provider_config.transform_image_edit_response(
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
return image_edit_response
def image_generation_handler(
self,
@ -7154,6 +7192,7 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
encoding=None,
)
set_provider_response_headers_in_hidden_params(model_response, response.headers)
return model_response
@ -7261,6 +7300,7 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
encoding=None,
)
set_provider_response_headers_in_hidden_params(model_response, response.headers)
return model_response
@ -7850,7 +7890,7 @@ class BaseLLMHTTPHandler:
def video_create_character_handler(
self,
name: str,
video: Any,
video: FileTypes,
video_provider_config: BaseVideoConfig,
custom_llm_provider: str,
litellm_params,
@ -7934,7 +7974,7 @@ class BaseLLMHTTPHandler:
async def async_video_create_character_handler(
self,
name: str,
video: Any,
video: FileTypes,
video_provider_config: BaseVideoConfig,
custom_llm_provider: str,
litellm_params,
@ -12045,11 +12085,13 @@ class BaseLLMHTTPHandler:
provider_config=text_to_speech_provider_config,
)
return text_to_speech_provider_config.transform_text_to_speech_response(
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
return speech_response
async def async_text_to_speech_handler(
self,
@ -12144,11 +12186,13 @@ class BaseLLMHTTPHandler:
provider_config=text_to_speech_provider_config,
)
return text_to_speech_provider_config.transform_text_to_speech_response(
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
return speech_response
#########################################################
########## SKILLS API HANDLERS ##########################

View file

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

View file

@ -27,6 +27,7 @@ from litellm import LlmProviders
from litellm._logging import verbose_logger
from litellm.constants import DEFAULT_MAX_RETRIES
from litellm.files.types import FileContentStreamingResult
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
@ -1404,7 +1405,6 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
organization: str | None = None,
headers: dict | None = None,
):
response = None
try:
openai_aclient: Final = self._get_openai_client(
is_async=True,
@ -1428,8 +1428,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
)
request_data: Final = {**data, "extra_headers": headers} if headers else data
response = await openai_aclient.images.generate(**request_data, timeout=timeout)
stringified_response: Final = response.model_dump()
raw_response: Final = await openai_aclient.images.with_raw_response.generate(
**request_data, timeout=timeout
)
stringified_response: Final = raw_response.parse().model_dump()
## LOGGING
logging_obj.post_call(
input=prompt,
@ -1437,11 +1439,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
additional_args={"complete_input_dict": data},
original_response=stringified_response,
)
return convert_to_model_response_object(
image_response: Final[ImageResponse] = convert_to_model_response_object(
response_object=stringified_response,
model_response_object=model_response,
response_type="image_generation",
)
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
return image_response
except Exception as e:
## LOGGING
logging_obj.post_call(
@ -1512,9 +1516,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
## COMPLETION CALL
request_data: Final = {**data, "extra_headers": headers} if headers else data
_response: Final = openai_client.images.generate(**request_data, timeout=timeout)
raw_response: Final = openai_client.images.with_raw_response.generate(**request_data, timeout=timeout)
response: Final = _response.model_dump()
response: Final = raw_response.parse().model_dump()
## LOGGING
logging_obj.post_call(
input=prompt,
@ -1522,11 +1526,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
additional_args={"complete_input_dict": data},
original_response=response,
)
return convert_to_model_response_object(
image_response: Final[ImageResponse] = convert_to_model_response_object(
response_object=response,
model_response_object=model_response,
response_type="image_generation",
)
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
return image_response
except OpenAIError as e:
## LOGGING
logging_obj.post_call(
@ -1609,7 +1615,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
input=input,
**optional_params,
)
return HttpxBinaryResponseContent(response=response.response)
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
return speech_response
async def async_audio_speech(
self,
@ -1655,8 +1663,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
input=input,
**optional_params,
)
return HttpxBinaryResponseContent(response=response.response)
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
return speech_response
class OpenAIFilesAPI(BaseLLM):

View file

@ -4,11 +4,10 @@ import httpx
from openai import AsyncOpenAI, OpenAI
from pydantic import BaseModel
import litellm
if TYPE_CHECKING:
from aiohttp import ClientSession
from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_name
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.audio_transcription.transformation import (
BaseAudioTranscriptionConfig,
@ -31,11 +30,6 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
data: dict,
timeout: float | httpx.Timeout,
):
"""
Helper to:
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
- call openai_aclient.audio.transcriptions.create by default
"""
try:
raw_response = await openai_aclient.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
headers: Final = dict(raw_response.headers)
@ -51,20 +45,11 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
data: dict,
timeout: float | httpx.Timeout,
):
"""
Helper to:
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
- call openai_aclient.audio.transcriptions.create by default
"""
try:
if litellm.return_response_headers is True:
raw_response = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
headers: Final = dict(raw_response.headers)
response = raw_response.parse()
return headers, response
else:
response = openai_client.audio.transcriptions.create(**data, timeout=timeout)
return None, response
raw_response: Final = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
headers: Final = dict(raw_response.headers)
response: Final = raw_response.parse()
return headers, response
except Exception as e:
raise e
@ -133,11 +118,12 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
"complete_input_dict": data,
},
)
_, response = self.make_sync_openai_audio_transcriptions_request(
headers, response = self.make_sync_openai_audio_transcriptions_request(
openai_client=openai_client,
data=data,
timeout=timeout,
)
logging_obj.model_call_details["response_headers"] = headers
if isinstance(response, BaseModel):
stringified_response = response.model_dump()
@ -158,6 +144,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
hidden_params=hidden_params,
response_type="audio_transcription",
)
set_provider_response_headers_in_hidden_params(final_response, headers)
return final_response
async def async_audio_transcriptions(
@ -217,12 +204,14 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
actual_model: Final = data.get("model", "whisper-1")
hidden_params: Final = {"model": actual_model, "custom_llm_provider": "openai"}
return convert_to_model_response_object(
final_response: Final[TranscriptionResponse] = convert_to_model_response_object(
response_object=stringified_response,
model_response_object=model_response,
hidden_params=hidden_params,
response_type="audio_transcription",
)
set_provider_response_headers_in_hidden_params(final_response, headers)
return final_response
except Exception as e:
## LOGGING
logging_obj.post_call(

View file

@ -1,7 +1,11 @@
from collections.abc import Mapping
from typing import Any, Final
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Any, Final, Protocol
from urllib.parse import unquote
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.llms.vertex_ai.common_utils import (
VertexAIError,
@ -9,35 +13,128 @@ from litellm.llms.vertex_ai.common_utils import (
)
from litellm.types.llms.openai import BatchJobStatus, CreateBatchRequest
from litellm.types.llms.vertex_ai import *
from litellm.types.utils import LiteLLMBatch, PromptTokensDetailsWrapper
from litellm.types.llms.vertex_ai import GenerateContentResponseBody
from litellm.types.utils import LiteLLMBatch, ModelInfo, Usage
_NATIVE_VERTEX_RESPONSE: Final = TypeAdapter(GenerateContentResponseBody)
def vertex_prompt_tokens_details(
usage_metadata: Mapping[str, object],
) -> PromptTokensDetailsWrapper | None:
raw_details: Final = usage_metadata.get("promptTokensDetails")
if not isinstance(raw_details, list):
return None
def _int_field(mapping: Mapping[str, object], key: str) -> int:
value: Final = mapping.get(key)
if isinstance(value, int):
return value
return int(value) if isinstance(value, str) and value.isdigit() else 0
def _normalize(detail: object) -> tuple[str, int] | None:
if not isinstance(detail, Mapping):
def vertex_embedding_prompt_token_count(vertex_response: Mapping[str, object]) -> int:
"""
Prompt tokens billed for one Vertex Gemini Embedding batch row.
Live rows report usage under `usageMetadata`; the documented `tokenCount` is kept as
a fallback.
"""
usage_metadata: Final = vertex_response.get("usageMetadata")
if isinstance(usage_metadata, Mapping):
return _int_field(usage_metadata, "promptTokenCount")
return _int_field(vertex_response, "tokenCount")
def is_vertex_embedding_batch_output_response(response_body: Mapping[str, object]) -> bool:
return isinstance(response_body.get("embedding"), dict)
def is_native_vertex_batch_output_row(row: Mapping[str, object]) -> bool:
return isinstance(row.get("request"), dict)
class NativeVertexBatchCostCalculator(Protocol):
def __call__(
self,
usage: Usage,
model: str,
custom_llm_provider: str | None = None,
model_info: ModelInfo | None = None,
) -> tuple[float, float]: ...
@dataclass(frozen=True, slots=True)
class NativeVertexBatchRowStats:
usage: Usage
total_tokens: int
model: str | None
prompt_cost: float
completion_cost: float
def _native_vertex_row_usage(
response_body: Mapping[str, object],
calculate_usage: Callable[[GenerateContentResponseBody], Usage],
) -> Usage | None:
if "usageMetadata" not in response_body:
if not is_vertex_embedding_batch_output_response(response_body):
return None
modality: Final = detail.get("modality")
token_count: Final = detail.get("tokenCount")
if not isinstance(modality, str) or not isinstance(token_count, int):
return None
return modality.upper(), token_count
parsed_details: Final = tuple(_normalize(detail) for detail in raw_details)
normalized: Final = tuple(detail for detail in parsed_details if detail is not None)
if len(normalized) != len(parsed_details):
prompt_tokens: Final = vertex_embedding_prompt_token_count(response_body)
return Usage(prompt_tokens=prompt_tokens, completion_tokens=0, total_tokens=prompt_tokens)
try:
completion_response: Final = _NATIVE_VERTEX_RESPONSE.validate_python(response_body)
except ValidationError as e:
verbose_logger.debug("vertex_ai batch row response is not a GenerateContentResponse: %s", str(e))
return None
return calculate_usage(completion_response)
return PromptTokensDetailsWrapper(
text_tokens=sum(token_count for modality, token_count in normalized if modality in ("TEXT", "DOCUMENT")),
audio_tokens=sum(token_count for modality, token_count in normalized if modality == "AUDIO"),
image_tokens=sum(token_count for modality, token_count in normalized if modality == "IMAGE"),
video_tokens=sum(token_count for modality, token_count in normalized if modality == "VIDEO"),
def native_vertex_batch_row_stats(
row: Mapping[str, object],
model_name: str | None,
*,
model_info: ModelInfo | None,
calculate_usage: Callable[[GenerateContentResponseBody], Usage],
cost_calculator: NativeVertexBatchCostCalculator,
) -> NativeVertexBatchRowStats | None:
"""
Usage and cost of one native Vertex predictions.jsonl row, a
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
generateContent object or a `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
embedding object (an embedding row without `usageMetadata` is billed from its documented `tokenCount`).
`model_name` (the deployment model) prices the row unless it is a wildcard, else its own `modelVersion`
does, else the wildcard name so explicit deployment prices still apply; a row without a response, a
generateContent row without `response.usageMetadata`, and a row whose response fails validation are
None (failed).
"""
response_body: Final = row.get("response")
if not isinstance(response_body, dict):
return None
usage: Final = _native_vertex_row_usage(response_body, calculate_usage)
if usage is None:
return None
total_tokens: Final = usage.total_tokens or (usage.prompt_tokens + usage.completion_tokens)
model_version: Final = response_body.get("modelVersion")
deployment_model: Final = model_name if model_name and "*" not in model_name else None
model: Final = deployment_model or (model_version if isinstance(model_version, str) else model_name)
if model is None:
verbose_logger.warning(
"vertex_ai batch output row could not be costed, so it is billed at $0 and the rest of the batch "
"is still billed: the row has no modelVersion and the batch has no deployment model"
)
return NativeVertexBatchRowStats(
usage=usage, total_tokens=total_tokens, model=None, prompt_cost=0.0, completion_cost=0.0
)
try:
prompt_cost, completion_cost = cost_calculator(
usage=usage, model=model, custom_llm_provider="vertex_ai", model_info=model_info
)
except Exception as e: # noqa: BLE001 # one unpriceable row must not abort the batch's cost accounting
verbose_logger.warning(
"vertex_ai batch output row could not be costed, so it is billed at $0 and the rest of the batch "
"is still billed. model=%s error=%s",
model,
str(e),
)
return NativeVertexBatchRowStats(
usage=usage, total_tokens=total_tokens, model=model, prompt_cost=0.0, completion_cost=0.0
)
return NativeVertexBatchRowStats(
usage=usage, total_tokens=total_tokens, model=model, prompt_cost=prompt_cost, completion_cost=completion_cost
)
@ -156,30 +253,15 @@ class VertexAIBatchTransformation:
return uris[0]
@classmethod
def _get_output_file_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str:
def _get_output_file_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str | None:
"""
Gets the output file id from the Vertex AI Batch response
Gets the output file id from the Vertex AI Batch response, None until Vertex reports outputInfo
"""
output_info: Final = response.get("outputInfo") or OutputInfo()
output_file_id: str = output_info.get("gcsOutputDirectory", "")
if output_file_id:
output_file_id = output_file_id.rstrip("/") + "/predictions.jsonl"
if output_file_id and output_file_id != "/predictions.jsonl":
return output_file_id
output_config: Final = response.get("outputConfig")
if output_config is None:
return output_file_id
gcs_destination: Final = output_config.get("gcsDestination")
if gcs_destination is None:
return output_file_id
output_uri_prefix: Final = gcs_destination.get("outputUriPrefix", "")
if output_uri_prefix.endswith("/predictions.jsonl"):
return output_uri_prefix
return output_uri_prefix.rstrip("/") + "/predictions.jsonl"
gcs_output_directory: Final = (output_info.get("gcsOutputDirectory") or "").rstrip("/")
if not gcs_output_directory:
return None
return f"{gcs_output_directory}/predictions.jsonl"
@classmethod
def _get_batch_job_status_from_vertex_ai_batch_response(

View file

@ -9,7 +9,7 @@ from collections.abc import AsyncGenerator, Callable, Iterable, Iterator, Mappin
from contextlib import aclosing
from dataclasses import dataclass
from types import MappingProxyType
from typing import Any, Final, TypedDict
from typing import IO, Any, Final, TypedDict
from urllib.parse import quote, unquote
import httpx
@ -41,6 +41,7 @@ from litellm.llms.base_llm.files.transformation import (
BaseFileUploadStream,
LiteLLMLoggingObj,
)
from litellm.llms.vertex_ai.batches.transformation import vertex_embedding_prompt_token_count
from litellm.llms.vertex_ai.common_utils import (
_convert_vertex_datetime_to_openai_datetime,
get_vertex_ai_fine_tuned_endpoint_id,
@ -56,6 +57,7 @@ from litellm.types.files import StreamingMediaUploadConfig
from litellm.types.llms.openai import (
AllMessageValues,
CreateFileRequest,
FileContent,
FileTypes,
HttpxBinaryResponseContent,
OpenAICreateFileRequestOptionalParams,
@ -87,6 +89,8 @@ _EMBED_REQUEST_FIELD_BY_GEMINI_PARAM: Final = (
_VERTEX_BATCH_FANNED_OUT_KEY_PATTERN: Final = re.compile(r"(?P<custom_id>[^#]*)#(?P<index>\d+)/(?P<total>\d+)")
_JSONL_NEWLINE: Final = b"\n"
_BATCH_OUTPUT_FIRST_ROW_PEEK_LIMIT_BYTES: Final = 32 * 1024 * 1024
_PASSTHROUGH_MANAGED_GCS_PREFIX: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}passthrough/"
_RAW_UPLOAD_CHUNK_BYTES: Final = 1024 * 1024
class _GcsObjectMetadataJson(TypedDict, total=False):
@ -418,19 +422,6 @@ def _split_vertex_batch_key(vertex_output_row: Mapping[str, object]) -> tuple[st
return unquote(match["custom_id"]), int(match["index"]), int(match["total"])
def _embedding_prompt_token_count(vertex_response: _VertexEmbeddingResponse) -> int:
"""
Prompt tokens billed for one Vertex Gemini Embedding batch row.
Live rows report usage under `usageMetadata`; the documented `tokenCount` is kept as
a fallback.
"""
usage_metadata = vertex_response.get("usageMetadata")
if isinstance(usage_metadata, Mapping):
return int(usage_metadata.get("promptTokenCount") or 0)
return int(vertex_response.get("tokenCount") or 0)
def _vertex_embeddings_rows_to_openai_batch_output_row(
custom_id: str,
vertex_output_rows: tuple[_VertexEmbeddingBatchRow, ...],
@ -471,7 +462,7 @@ def _vertex_embeddings_rows_to_openai_batch_output_row(
)
responses = tuple(row["response"] for row in vertex_output_rows)
token_count = sum(_embedding_prompt_token_count(response) for response in responses)
token_count = sum(vertex_embedding_prompt_token_count(response) for response in responses)
body = EmbeddingResponse(
model=model or "",
data=[
@ -528,6 +519,16 @@ def _model_from_managed_gcs_url(url: str) -> str | None:
return match.group(1) if match else None
def is_passthrough_managed_gcs_url(url: str) -> bool:
decoded_url: Final = unquote(url)
managed_prefix_start: Final = decoded_url.find(VERTEX_AI_MANAGED_GCS_PREFIX)
return managed_prefix_start >= 0 and decoded_url.startswith(_PASSTHROUGH_MANAGED_GCS_PREFIX, managed_prefix_start)
def is_passthrough_batch_upload(create_file_data: Mapping[str, object], litellm_params: Mapping[str, object]) -> bool:
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:
"""
Whether an OpenAI batch JSONL line targets the embeddings endpoint.
@ -791,6 +792,58 @@ class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
return self._iter_vertex_jsonl_chunks()
def _read_chunk_as_bytes(handle: IO[bytes]) -> bytes:
chunk: Final[bytes | str] = handle.read(_RAW_UPLOAD_CHUNK_BYTES)
return chunk.encode("utf-8") if isinstance(chunk, str) else bytes(chunk)
def _iter_raw_file_chunks(file_content: FileTypes) -> Iterator[bytes]:
content: Final[FileContent | str] = file_content[1] if isinstance(file_content, tuple) else file_content
if isinstance(content, (bytes, bytearray)):
yield from (
bytes(content[offset : offset + _RAW_UPLOAD_CHUNK_BYTES])
for offset in range(0, len(content), _RAW_UPLOAD_CHUNK_BYTES)
)
return
if isinstance(content, str):
yield content.encode("utf-8")
return
if isinstance(content, PathLike):
with open(str(content), "rb") as handle:
yield from iter(lambda: handle.read(_RAW_UPLOAD_CHUNK_BYTES), b"")
return
if not hasattr(content, "read"):
raise ValueError("Unsupported file content type")
seek: Final = getattr(content, "seek", None)
if seek is None:
raise ValueError(
"Batch upload file handle must be seekable; got a non-seekable "
"stream. Pass bytes, a path, or a seekable handle."
)
seek(0)
yield from iter(lambda: _read_chunk_as_bytes(content), b"")
class _RawFileUploadStream(BaseFileUploadStream):
def __init__(self, file_content: FileTypes) -> None:
self._file_content = file_content
def iter_bytes(self) -> Iterator[bytes]:
return _iter_raw_file_chunks(self._file_content)
def _managed_batch_object_name(raw_model: str, *, passthrough: bool) -> str:
endpoint_id: Final = get_vertex_ai_fine_tuned_endpoint_id(raw_model)
model_path: Final = (
f"endpoints/{endpoint_id}"
if endpoint_id is not None
else (raw_model if "publishers/google/models" in raw_model else f"publishers/google/models/{raw_model}")
)
safe_model_path: Final = sanitize_cloud_object_path(model_path, fallback="model")
prefix: Final = _PASSTHROUGH_MANAGED_GCS_PREFIX if passthrough else VERTEX_AI_MANAGED_GCS_PREFIX
return f"{prefix}{safe_model_path}/{uuid.uuid4()}"
class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
Config for VertexAI Files
@ -848,23 +901,34 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
if deployment_model
else openai_jsonl_content[0].get("body", {}).get("model", "")
)
endpoint_id: Final = get_vertex_ai_fine_tuned_endpoint_id(raw_model)
model_path: Final = (
f"endpoints/{endpoint_id}"
if endpoint_id is not None
else (raw_model if "publishers/google/models" in raw_model else f"publishers/google/models/{raw_model}")
)
safe_model_path: Final = sanitize_cloud_object_path(model_path, fallback="model")
object_name: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
return object_name
return _managed_batch_object_name(raw_model, passthrough=False)
def get_object_name(self, file_data: FileTypes, purpose: str, deployment_model: str | None = None) -> str:
def _get_passthrough_gcs_object_name(self, deployment_model: str | None) -> str:
if not deployment_model:
raise VertexAIError(
status_code=400,
message=(
"Native Vertex batch passthrough uploads need the deployment model to name the GCS object, "
"since native rows carry no model: pass `target_model_names` (proxy) or `model` (SDK)."
),
)
return _managed_batch_object_name(deployment_model.removeprefix("vertex_ai/"), passthrough=True)
def get_object_name(
self,
file_data: FileTypes,
purpose: str,
deployment_model: str | None = None,
passthrough: bool = False,
) -> str:
"""
Get the object name for the request.
Reads only the first JSONL entry (streamed) for batch files, so a large
upload is never materialized just to derive the GCS object name.
"""
if purpose == "batch" and passthrough:
return self._get_passthrough_gcs_object_name(deployment_model)
if purpose == "batch":
## 1. If jsonl, derive the object name from the deployment model (or the first entry's)
first_entry: Final = next(_iter_openai_jsonl_entries(file_data), None)
@ -922,6 +986,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
file_data,
purpose,
deployment_model=configured_model if isinstance(configured_model, str) else None,
passthrough=is_passthrough_batch_upload(data, litellm_params),
)
if object_prefix:
object_name = f"{object_prefix}/{object_name}"
@ -984,6 +1049,14 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
if file_data is None:
raise ValueError("file is required")
if is_passthrough_batch_upload(create_file_data, litellm_params):
return {
"streaming_media_upload": StreamingMediaUploadConfig(
body_stream=_RawFileUploadStream(file_data),
content_type="application/json",
)
}
_, content_type = extract_file_metadata(file_data)
if FilesAPIUtils.is_batch_jsonl_request(
create_file_data=create_file_data,
@ -1164,6 +1237,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
# transformation, e.g. if they consume raw `predictions.jsonl` directly.
if getattr(litellm, "disable_vertex_batch_output_transformation", False):
return HttpxBinaryResponseContent(response=raw_response)
if is_passthrough_managed_gcs_url(str(raw_response.request.url)):
return HttpxBinaryResponseContent(response=raw_response)
# Try to transform batch output if it's a JSONL file
content: Final = raw_response.content
@ -1209,7 +1284,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
Everything else is passed through unchanged, including a row that fails to
transform mid-stream.
"""
if litellm.disable_vertex_batch_output_transformation:
if litellm.disable_vertex_batch_output_transformation or is_passthrough_managed_gcs_url(request_url):
return FileContentStreamingResult(stream_iterator=stream_iterator, headers=headers)
first_line, buffered = await _peek_first_jsonl_line(

View file

@ -57,6 +57,7 @@ from litellm.utils import (
# Logging is imported lazily when needed to avoid loading litellm_logging at import time
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.router import Router
from litellm.types.utils import TokenCountResponse
from litellm.constants import (
@ -351,7 +352,7 @@ class LiteLLM:
class Chat:
def __init__(self, params, router_obj: Any | None):
def __init__(self, params, router_obj: "Router | None"):
self.params = params
if self.params.get("acompletion", False) is True:
self.params.pop("acompletion")
@ -361,7 +362,7 @@ class Chat:
class Completions:
def __init__(self, params, router_obj: Any | None):
def __init__(self, params, router_obj: "Router | None"):
self.params = params
self.router_obj = router_obj
@ -377,7 +378,7 @@ class Completions:
class AsyncCompletions:
def __init__(self, params, router_obj: Any | None):
def __init__(self, params, router_obj: "Router | None"):
self.params = params
self.router_obj = router_obj

View file

@ -5442,7 +5442,7 @@
"supports_web_search": false
},
"azure/gpt-4.1-nano": {
"deprecation_date": "2026-10-14",
"deprecation_date": "2027-04-14",
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_token": 1e-07,
"input_cost_per_token_batches": 5e-08,
@ -5476,7 +5476,7 @@
"supports_vision": true
},
"azure/gpt-4.1-nano-2025-04-14": {
"deprecation_date": "2026-10-14",
"deprecation_date": "2027-04-14",
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_token": 1e-07,
"input_cost_per_token_batches": 5e-08,
@ -5603,6 +5603,39 @@
"supports_tool_choice": true,
"supports_vision": true
},
"azure/gpt-audio": {
"deprecation_date": "2027-03-02",
"input_cost_per_audio_token": 4e-05,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "azure",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 1e-05,
"supported_endpoints": [
"/v1/chat/completions"
],
"supported_modalities": [
"text",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": false,
"supports_reasoning": false,
"supports_response_schema": false,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false,
"source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure"
},
"azure/gpt-audio-2025-08-28": {
"deprecation_date": "2027-03-02",
"input_cost_per_audio_token": 4e-05,
@ -5635,6 +5668,39 @@
"supports_tool_choice": true,
"supports_vision": false
},
"azure/gpt-audio-1.5": {
"deprecation_date": "2027-08-24",
"input_cost_per_audio_token": 4e-05,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "azure",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 1e-05,
"supported_endpoints": [
"/v1/chat/completions"
],
"supported_modalities": [
"text",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": false,
"supports_reasoning": false,
"supports_response_schema": false,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false,
"source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure"
},
"azure/gpt-audio-1.5-2026-02-23": {
"deprecation_date": "2027-08-24",
"input_cost_per_audio_token": 4e-05,
@ -5668,7 +5734,7 @@
"supports_vision": false
},
"azure/gpt-audio-mini": {
"deprecation_date": "2027-04-06",
"deprecation_date": "2027-06-15",
"input_cost_per_audio_token": 1e-05,
"input_cost_per_token": 6e-07,
"litellm_provider": "azure",
@ -5849,6 +5915,41 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"azure/gpt-realtime": {
"cache_creation_input_audio_token_cost": 4e-06,
"cache_read_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
"deprecation_date": "2027-03-02",
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_image_token": 5e-06,
"input_cost_per_token": 4e-06,
"litellm_provider": "azure",
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"image",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure"
},
"azure/gpt-realtime-2025-08-28": {
"cache_creation_input_audio_token_cost": 4e-06,
"cache_read_input_audio_token_cost": 4e-07,
@ -5883,6 +5984,41 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"azure/gpt-realtime-1.5": {
"cache_creation_input_audio_token_cost": 4e-06,
"cache_read_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
"deprecation_date": "2027-08-24",
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_image_token": 5e-06,
"input_cost_per_token": 4e-06,
"litellm_provider": "azure",
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"image",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure"
},
"azure/gpt-realtime-1.5-2026-02-23": {
"cache_creation_input_audio_token_cost": 4e-06,
"cache_read_input_audio_token_cost": 4e-07,
@ -5921,7 +6057,7 @@
"cache_creation_input_audio_token_cost": 4e-07,
"cache_read_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
"deprecation_date": "2027-06-25",
"deprecation_date": "2027-07-31",
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_image_token": 5e-06,
"input_cost_per_token": 4e-06,
@ -5956,7 +6092,7 @@
"cache_creation_input_audio_token_cost": 3e-07,
"cache_read_input_audio_token_cost": 3e-07,
"cache_read_input_token_cost": 6e-08,
"deprecation_date": "2027-06-25",
"deprecation_date": "2027-07-31",
"input_cost_per_audio_token": 1e-05,
"input_cost_per_image_token": 8e-07,
"input_cost_per_token": 6e-07,
@ -6133,7 +6269,7 @@
"supports_tool_choice": true
},
"azure/gpt-4o-transcribe": {
"deprecation_date": "2026-10-15",
"deprecation_date": "2026-12-31",
"input_cost_per_audio_token": 2.5e-06,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "azure",
@ -6632,7 +6768,7 @@
},
"azure/gpt-5-chat": {
"cache_read_input_token_cost": 1.25e-07,
"deprecation_date": "2026-05-13",
"deprecation_date": "2026-06-29",
"input_cost_per_token": 1.25e-06,
"litellm_provider": "azure",
"max_input_tokens": 128000,
@ -10660,7 +10796,7 @@
"supports_web_search": false
},
"azure/us/gpt-4.1-nano-2025-04-14": {
"deprecation_date": "2026-10-14",
"deprecation_date": "2027-04-14",
"cache_read_input_token_cost": 2.8e-08,
"input_cost_per_token": 1.1e-07,
"input_cost_per_token_batches": 5.5e-08,
@ -26063,6 +26199,7 @@
"input_cost_per_token_batches": 1.5e-07,
"input_cost_per_token_flex": 1.5e-07,
"input_cost_per_token_priority": 5.4e-07,
"input_cost_per_audio_token_priority": 1.8e-06,
"output_cost_per_token_batches": 1.25e-06,
"output_cost_per_token_flex": 1.25e-06,
"output_cost_per_token_priority": 4.5e-06,
@ -26075,7 +26212,6 @@
"input_cost_per_token": 3e-07,
"input_cost_per_token_batches": 1.5e-07,
"input_cost_per_token_flex": 1.5e-07,
"input_cost_per_token_priority": 5.4e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
@ -26175,7 +26311,11 @@
},
"gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -26185,6 +26325,7 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_batches": 6e-06,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
@ -26261,7 +26402,10 @@
},
"gemini-3.1-flash-image-preview": {
"input_cost_per_image": 0.00056,
"cache_read_input_token_cost": 5e-08,
"cache_read_input_token_cost_batches": 2.5e-08,
"input_cost_per_token": 5e-07,
"input_cost_per_token_batches": 2.5e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
"max_output_tokens": 32768,
@ -26270,6 +26414,7 @@
"output_cost_per_image": 0.0672,
"output_cost_per_image_token": 6e-05,
"output_cost_per_token": 3e-06,
"output_cost_per_token_batches": 1.5e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"supported_endpoints": [
"/v1/chat/completions",
@ -26400,6 +26545,7 @@
"input_cost_per_token_batches": 1.25e-07,
"input_cost_per_token_flex": 1.25e-07,
"input_cost_per_token_priority": 4.5e-07,
"input_cost_per_audio_token_priority": 9e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
@ -26595,6 +26741,7 @@
"input_cost_per_token_batches": 5e-08,
"input_cost_per_token_flex": 5e-08,
"input_cost_per_token_priority": 1.8e-07,
"input_cost_per_audio_token_priority": 5.4e-07,
"output_cost_per_token_batches": 2e-07,
"output_cost_per_token_flex": 2e-07,
"output_cost_per_token_priority": 7.2e-07,
@ -41242,34 +41389,33 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro": {
"input_cost_per_token": 9.24462e-07,
"input_cost_per_token_cache_hit": 4.4e-08,
"input_cost_per_token": 9.1263e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 1.848924e-06,
"output_cost_per_token": 1.82526e-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.70385e-08,
"cache_read_input_token_cost": 7.60525e-08,
"supports_audio_input": false,
"supports_pdf_input": false,
"supports_vision": false,
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4.1-flash": {
"input_cost_per_token": 1.4e-07,
"output_cost_per_token": 4.2e-07,
"cache_read_input_token_cost": 4.2e-09,
"input_cost_per_token": 3e-07,
"output_cost_per_token": 1.2e-06,
"cache_read_input_token_cost": 6e-09,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 943718,
"max_tokens": 943718,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"00:00-01:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"04:00-06:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"10:00-00:00"}],"input_cost_per_token":1.5e-7,"output_cost_per_token":6e-7,"cache_read_input_token_cost":3e-9},
"source": "https://openrouter.ai/api/v1/models",
@ -49269,7 +49415,6 @@
"input_cost_per_token": 3e-07,
"input_cost_per_token_batches": 1.5e-07,
"input_cost_per_token_flex": 1.5e-07,
"input_cost_per_token_priority": 5.4e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
@ -49346,7 +49491,11 @@
},
"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_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -49356,6 +49505,7 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-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"
@ -49384,7 +49534,10 @@
},
"vertex_ai/gemini-3.1-flash-image-preview": {
"input_cost_per_image": 0.00056,
"cache_read_input_token_cost": 5e-08,
"cache_read_input_token_cost_batches": 2.5e-08,
"input_cost_per_token": 5e-07,
"input_cost_per_token_batches": 2.5e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
"max_output_tokens": 32768,
@ -49393,6 +49546,7 @@
"output_cost_per_image": 0.0672,
"output_cost_per_image_token": 6e-05,
"output_cost_per_token": 3e-06,
"output_cost_per_token_batches": 1.5e-06,
"supports_reasoning": false,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models"
},
@ -49499,6 +49653,7 @@
"input_cost_per_token_batches": 1.25e-07,
"input_cost_per_token_flex": 1.25e-07,
"input_cost_per_token_priority": 4.5e-07,
"input_cost_per_audio_token_priority": 9e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
@ -57189,12 +57344,12 @@
]
},
"global.openai.gpt-5.4": {
"input_cost_per_token": 2.75e-06,
"input_cost_per_token_above_272k_tokens": 5.5e-06,
"cache_read_input_token_cost": 2.75e-07,
"cache_read_input_token_cost_above_272k_tokens": 5.5e-07,
"output_cost_per_token": 1.65e-05,
"output_cost_per_token_above_272k_tokens": 2.475e-05,
"input_cost_per_token": 2.5e-06,
"input_cost_per_token_above_272k_tokens": 5e-06,
"cache_read_input_token_cost": 2.5e-07,
"cache_read_input_token_cost_above_272k_tokens": 5e-07,
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_above_272k_tokens": 2.25e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
@ -57251,12 +57406,12 @@
]
},
"global.openai.gpt-5.5": {
"input_cost_per_token": 5.5e-06,
"input_cost_per_token_above_272k_tokens": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
"output_cost_per_token": 3.3e-05,
"output_cost_per_token_above_272k_tokens": 4.95e-05,
"input_cost_per_token": 5e-06,
"input_cost_per_token_above_272k_tokens": 1e-05,
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
"output_cost_per_token": 3e-05,
"output_cost_per_token_above_272k_tokens": 4.5e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
@ -63384,6 +63539,73 @@
"supports_tool_choice": true,
"supports_vision": true
},
"azure_ai/deepseek-v4.1-flash": {
"cache_read_input_token_cost": 8e-09,
"input_cost_per_token": 3.75e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 1.5e-06,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"deprecation_date": "2026-12-15",
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
},
"azure_ai/muse-spark-1.3": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.25e-06,
"source": "https://ai.developer.meta.com/docs/pricing-rate-limits",
"supported_endpoints": [
"/v1/chat/completions"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"azure_ai/MAI-Image-2.6": {
"input_cost_per_image_token": 8e-06,
"input_cost_per_token": 5e-06,
"litellm_provider": "azure_ai",
"mode": "image_generation",
"output_cost_per_image_token": 3.8e-05,
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'",
"supported_endpoints": [
"/v1/images/generations",
"/v1/images/edits"
]
},
"azure_ai/MAI-Image-2.6-Flash": {
"input_cost_per_image_token": 2.5e-06,
"input_cost_per_token": 1.75e-06,
"litellm_provider": "azure_ai",
"mode": "image_generation",
"output_cost_per_image_token": 1.9e-05,
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'",
"supported_endpoints": [
"/v1/images/generations",
"/v1/images/edits"
]
},
"azure_ai/FW-DeepSeek-V4.1-Flash": {
"cache_read_input_token_cost": 8e-09,
"input_cost_per_token": 3.75e-07,
@ -63445,6 +63667,7 @@
"supports_tool_choice": true
},
"azure_ai/FW-GPT-OSS-120B": {
"deprecation_date": "2027-07-01",
"cache_read_input_token_cost": 8.2e-08,
"input_cost_per_token": 1.65e-07,
"litellm_provider": "azure_ai",
@ -63462,6 +63685,7 @@
"supports_tool_choice": true
},
"azure_ai/Cohere-command-a-plus-05-2026": {
"deprecation_date": "2026-10-16",
"input_cost_per_token": 8e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 128000,
@ -66422,9 +66646,9 @@
"supports_web_search": true
},
"openrouter/deepseek/deepseek-v4-flash": {
"input_cost_per_token": 8.554e-08,
"output_cost_per_token": 1.7108e-07,
"cache_read_input_token_cost": 1.7108e-08,
"input_cost_per_token": 8.4e-08,
"output_cost_per_token": 1.68e-07,
"cache_read_input_token_cost": 1.68e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
@ -68339,6 +68563,7 @@
"source": "https://api.together.ai/v1/models"
},
"vertex_ai/gemini-2.5-flash-native-audio": {
"deprecation_date": "2026-12-13",
"input_cost_per_audio_token": 3e-06,
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai",
@ -68385,7 +68610,9 @@
"input_cost_per_token": 1.5e-06,
"litellm_provider": "vertex_ai",
"mode": "chat",
"output_cost_per_reasoning_token": 9e-06,
"output_cost_per_token": 9e-06,
"output_cost_per_video_token": 1.75e-05,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
},
"vertex_ai/gemini-omni-1.1-flash-preview": {
@ -68891,7 +69118,7 @@
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
},
"azure/eu/gpt-4.1-nano": {
"deprecation_date": "2026-10-14",
"deprecation_date": "2027-04-14",
"cache_read_input_token_cost": 2.8e-08,
"input_cost_per_token": 1.1e-07,
"input_cost_per_token_batches": 5.5e-08,
@ -69079,7 +69306,8 @@
"mode": "chat",
"output_cost_per_token": 5.5e-05,
"output_cost_per_token_above_272k_tokens": 8.25e-05,
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'",
"supports_reasoning": true
},
"azure/eu/gpt-6-luna": {
"deprecation_date": "2028-03-11",
@ -69326,7 +69554,7 @@
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
},
"azure/us/gpt-4.1-nano": {
"deprecation_date": "2026-10-14",
"deprecation_date": "2027-04-14",
"cache_read_input_token_cost": 2.8e-08,
"input_cost_per_token": 1.1e-07,
"input_cost_per_token_batches": 5.5e-08,
@ -69546,6 +69774,7 @@
"cache_creation_input_audio_token_cost": 4e-07,
"cache_read_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
"deprecation_date": "2026-10-31",
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_image_token": 5e-06,
"input_cost_per_token": 4e-06,
@ -69577,6 +69806,7 @@
"supports_tool_choice": true
},
"azure/gpt-live-1": {
"deprecation_date": "2027-09-10",
"input_cost_per_second": 0.000833333333333,
"litellm_provider": "azure",
"mode": "realtime",
@ -69594,6 +69824,7 @@
"supports_function_calling": true
},
"azure/gpt-live-transcribe": {
"deprecation_date": "2028-02-01",
"input_cost_per_second": 0.000283333333333,
"litellm_provider": "azure",
"max_input_tokens": 32000,
@ -69615,6 +69846,7 @@
"supports_audio_input": true
},
"azure/gpt-transcribe": {
"deprecation_date": "2028-02-01",
"input_cost_per_second": 7.5e-05,
"litellm_provider": "azure",
"mode": "audio_transcription",
@ -69633,6 +69865,7 @@
"supports_audio_input": true
},
"azure/gpt-realtime-translate": {
"deprecation_date": "2027-05-06",
"input_cost_per_second": 0.000566666666667,
"litellm_provider": "azure",
"max_input_tokens": 32000,

View file

@ -111,6 +111,13 @@ _MCP_DESTINATIONS_SCOPE_KEY: Final = "litellm_otel_request_destinations"
_MCP_PROTOCOL_VERSION_HEADER: Final = b"mcp-protocol-version"
def reject_disallowed_mcp_origin(request: StarletteRequest) -> None:
from litellm.proxy.proxy_server import origins # noqa: PLC0415 # proxy imports this module during startup
if "*" not in origins and any(origin not in origins for origin in request.headers.getlist("origin")):
raise HTTPException(status_code=403, detail="Invalid Origin header")
def unsupported_protocol_version(scope: Scope) -> str | None:
"""Return the unsupported ``MCP-Protocol-Version`` header value, if any.
@ -1931,6 +1938,7 @@ if MCP_AVAILABLE:
async def handle_streamable_http_mcp(scope: Scope, receive: Receive, send: Send) -> None:
"""Handle MCP requests through StreamableHTTP."""
try:
reject_disallowed_mcp_origin(StarletteRequest(scope))
bad_version: Final = unsupported_protocol_version(scope)
if bad_version is not None:
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
@ -2275,6 +2283,7 @@ if MCP_AVAILABLE:
async def handle_sse_mcp(scope: Scope, receive: Receive, send: Send) -> None:
"""Handle MCP requests through SSE."""
try:
reject_disallowed_mcp_origin(StarletteRequest(scope))
bad_version: Final = unsupported_protocol_version(scope)
if bad_version is not None:
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))

View file

@ -307,6 +307,8 @@ class KeyManagementRoutes(str, enum.Enum):
# team usage routes
TEAM_DAILY_ACTIVITY = "/team/daily/activity"
TEAM_DAILY_ACTIVITY_AGGREGATED = "/team/daily/activity/aggregated"
TEAM_DAILY_ACTIVITY_EXPORT = "/team/daily/activity/export"
TEAM_DAILY_ACTIVITY_AGGREGATED_SEARCH = "/team/daily/activity/aggregated/search"
# team spend-log viewing
SPEND_LOGS = "/spend/logs"
@ -673,6 +675,7 @@ class LiteLLMRoutes(enum.Enum):
KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value,
KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value,
KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED.value,
KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED_SEARCH.value,
KeyManagementRoutes.SPEND_LOGS.value,
KeyManagementRoutes.SPEND_LOGS_V2.value,
KeyManagementRoutes.KEY_RESET_SPEND.value,
@ -699,6 +702,7 @@ class LiteLLMRoutes(enum.Enum):
"/user/list",
"/user/daily/activity",
"/user/daily/activity/aggregated",
"/user/daily/activity/aggregated/search",
# team
"/team/new",
"/team/update",
@ -716,6 +720,8 @@ class LiteLLMRoutes(enum.Enum):
"/team/permissions_bulk_update",
"/team/daily/activity",
"/team/daily/activity/aggregated",
"/team/daily/activity/export",
"/team/daily/activity/aggregated/search",
"/team/spend/by_user",
# gateway request counts (SGR); deployment-wide, admin-only
"/gateway/daily/activity",
@ -886,6 +892,8 @@ class LiteLLMRoutes(enum.Enum):
"/team/permissions_update",
"/team/daily/activity",
"/team/daily/activity/aggregated",
"/team/daily/activity/export",
"/team/daily/activity/aggregated/search",
"/team/spend/by_user",
"/team/{team_id}/members/me",
# POST/GET the team's logging callbacks, and DELETE one of them. Every
@ -901,6 +909,7 @@ class LiteLLMRoutes(enum.Enum):
"/model/delete",
"/user/daily/activity",
"/user/daily/activity/aggregated",
"/user/daily/activity/aggregated/search",
# Endpoint restricts results to organizations the caller is ORG_ADMIN
# of; a caller who administers none gets an empty result set.
"/organization/daily/activity",
@ -984,6 +993,8 @@ class LiteLLMRoutes(enum.Enum):
"/user/daily/activity",
"/team/daily/activity",
"/team/daily/activity/aggregated",
"/team/daily/activity/export",
"/team/daily/activity/aggregated/search",
"/tag/daily/activity",
"/tag/list",
"/audit",

View file

@ -218,7 +218,7 @@ ProxyRouteType: TypeAlias = Literal[
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
# Type alias for streaming chunk serializer (chunk after hooks + cost injection -> wire format)
StreamChunkSerializer = Callable[[Any], str]
StreamChunkSerializer = Callable[[object], str]
# Type alias for streaming error serializer (ProxyException -> wire format)
StreamErrorSerializer = Callable[[ProxyException], str]
@ -459,7 +459,7 @@ async def _bill_partial_streamed_spend_on_disconnect(request_data: dict, respons
return True
async def _cancel_pending_gather_tasks(tasks: list["asyncio.Task[Any]"]) -> None:
async def _cancel_pending_gather_tasks(tasks: Sequence["asyncio.Task[object]"]) -> None:
pending_tasks: Final = [task for task in tasks if not task.done()]
for task in pending_tasks:
task.cancel()
@ -3145,7 +3145,7 @@ class ProxyBaseLLMRequestProcessing:
logging_obj._on_detached_stream_failure = _on_detached_stream_failure
def _is_streaming_response(self, response: Any) -> bool:
def _is_streaming_response(self, response: object) -> bool:
"""
Check if the response object is actually a streaming response by inspecting its type.
@ -3259,7 +3259,7 @@ class ProxyBaseLLMRequestProcessing:
async def _handle_non_streaming_allm_passthrough_route(
self,
response: Any,
response: _UpstreamHttpResponse,
proxy_logging_obj: "ProxyLogging",
user_api_key_dict: "UserAPIKeyAuth",
custom_headers: Mapping[str, str],
@ -3852,7 +3852,7 @@ class ProxyBaseLLMRequestProcessing:
@staticmethod
async def async_streaming_data_generator(
response: Any,
response: object,
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
proxy_logging_obj: ProxyLogging,
@ -3993,7 +3993,7 @@ class ProxyBaseLLMRequestProcessing:
@staticmethod
def async_sse_data_generator(
response: Any,
response: object,
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
proxy_logging_obj: ProxyLogging,

View file

@ -124,6 +124,12 @@ class _SpendIncrement(TypedDict):
increment: ReadOnly[float]
class _MemberSpendRow(TypedDict):
user_id: ReadOnly[str]
team_id: ReadOnly[str]
cost: ReadOnly[float]
class _SpendBatch(Protocol):
litellm_usertable: BatchTable
litellm_verificationtoken: BatchTable
@ -351,17 +357,22 @@ _TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS
# One statement adds every member's cost to their membership row. A missing row is created only
# while the user is still on the team's roster, so a spend flush landing after a removal never
# recreates the member.
# recreates the member. The rows travel as one JSON document, not as a numeric array: Prisma
# types a raw array parameter from the first batch a connection sees, so after an all-$0 batch
# (integers) every later fractional batch on that connection failed with "improper binary format".
_TEAM_MEMBER_SPEND_SQL: Final = """
INSERT INTO "LiteLLM_TeamMembership" (user_id, team_id, spend, total_spend)
SELECT p.user_id, p.team_id, p.cost, p.cost
FROM unnest($1::text[], $2::text[], $3::float8[]) AS p(user_id, team_id, cost)
SELECT member.user_id, member.team_id, member.cost, member.cost
FROM jsonb_to_recordset($1::jsonb) AS member(user_id text, team_id text, cost float8)
WHERE EXISTS (
SELECT 1 FROM "LiteLLM_TeamTable" t
WHERE t.team_id = p.team_id
AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))
WHERE t.team_id = member.team_id
AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', member.user_id))
)
OR EXISTS (
SELECT 1 FROM "LiteLLM_TeamMembership" m
WHERE m.user_id = member.user_id AND m.team_id = member.team_id
)
OR EXISTS (SELECT 1 FROM "LiteLLM_TeamMembership" m WHERE m.user_id = p.user_id AND m.team_id = p.team_id)
ON CONFLICT (user_id, team_id) DO UPDATE
SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend,
total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend
@ -371,15 +382,12 @@ SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend,
async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None:
# key is "team_id::<value>::user_id::<value>"; locks are taken in sorted team_id order like the team endpoints
rows: Final = sorted((key.split("::")[1], key.split("::")[3], cost) for key, cost in spend_by_member_key.items())
team_ids: Final = tuple(team_id for team_id, _user_id, _cost in rows)
for team_id in dict.fromkeys(team_ids):
for team_id in dict.fromkeys(team_id for team_id, _user_id, _cost in rows):
_ = await transaction.execute_raw(_TEAM_ADVISORY_LOCK_SQL, team_id)
_ = await transaction.execute_raw(
_TEAM_MEMBER_SPEND_SQL,
tuple(user_id for _team_id, user_id, _cost in rows),
team_ids,
tuple(cost for _team_id, _user_id, cost in rows),
members: Final = tuple(
_MemberSpendRow(user_id=user_id, team_id=team_id, cost=cost) for team_id, user_id, cost in rows
)
_ = await transaction.execute_raw(_TEAM_MEMBER_SPEND_SQL, json.dumps(members))
def get_llm_router():

View file

@ -1,4 +1,6 @@
import asyncio
import dataclasses
import itertools
from collections.abc import Awaitable, Callable, Mapping, Sequence
from collections.abc import Set as AbstractSet
from datetime import datetime, timedelta, timezone
@ -35,6 +37,10 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
SpendMetrics,
)
from litellm.types.proxy.management_endpoints.team_endpoints import (
TeamDailyActivityExportRow,
TeamDailyActivityExportType,
)
if TYPE_CHECKING:
from prisma.models import (
@ -198,7 +204,7 @@ class _AggregatedQueryKwargs(TypedDict):
include_current_utc_day: ReadOnly[bool]
_SqlQuery = tuple[str, list[str]]
_SqlQuery = tuple[str, Sequence[str]]
async def _query_raw_optional(
@ -974,6 +980,291 @@ def _build_entity_rollup_sql_query(
return sql_query, sql_params
def _build_export_sql_query(
*,
table_name: str,
entity_id_field: str,
entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
start_date: str,
end_date: str,
api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
exclude_entity_ids: list[str] | None, # mutable-ok: filter union shared with the paginated path
timezone_offset_minutes: int | None,
export_type: TeamDailyActivityExportType,
) -> tuple[str, tuple[str, ...]]:
"""One unbounded rollup for the export route, on the aggregated path's WHERE clause.
No LIMIT anywhere: the export exists so a caller can reach keys past
USAGE_TOP_API_KEYS_LIMIT. PTU sentinel rows stay in `daily` so per-team
totals match breakdown.entities, and are excluded from the key, user and
model exports where the flat-cost row has no meaning.
"""
pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name)
if pg_table is None:
raise ValueError(f"Unknown table name: {table_name}")
adjusted_start, adjusted_end = _adjust_dates_for_timezone(start_date, end_date, timezone_offset_minutes)
where_clause, where_params = _build_aggregated_where_clause(
entity_id_field=entity_id_field,
entity_id=entity_id,
adjusted_start=adjusted_start,
adjusted_end=adjusted_end,
model=None,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
)
keyed: Final = export_type in ("daily_with_keys", "daily_with_users")
by_model: Final = export_type == "daily_with_models"
group_extras: Final = tuple(field for field in ("api_key" if keyed else "", "model" if by_model else "") if field)
group_by: Final = f'date, "{entity_id_field}"' + "".join(f", {field}" for field in group_extras)
sentinel_clause: Final = f" AND api_key <> ${len(where_params) + 1}" if (keyed or by_model) else ""
sentinel_params: Final = (PTU_SENTINEL_API_KEY,) if (keyed or by_model) else ()
sql_query: Final = f"""
SELECT
date,
"{entity_id_field}" AS entity_id,
{"api_key" if keyed else "NULL::text AS api_key"},
{"model" if by_model else "NULL::text AS model"},{_rollup_metric_select(table_name)}
FROM "{pg_table}"
WHERE {where_clause}{sentinel_clause}
GROUP BY {group_by}
ORDER BY {group_by}
"""
return sql_query, (*where_params, *sentinel_params)
class _ExportRow(_RollupMetricsRow):
entity_id: str | None
model: str | None
def _export_team_alias(entity_metadata_field: Mapping[str, dict[str, object]] | None, entity_id: str) -> str | None:
alias: Final = _entity_metadata(entity_metadata_field, entity_id).get("team_alias")
return alias if isinstance(alias, str) else None
@dataclasses.dataclass(frozen=True, slots=True)
class _ExportMetrics:
spend: float
api_requests: int
successful_requests: int
failed_requests: int
total_tokens: int
prompt_tokens: int
completion_tokens: int
cache_read_input_tokens: int
cache_creation_input_tokens: int
@classmethod
def from_record(cls, record: _RollupMetricsRow) -> "_ExportMetrics":
prompt_tokens: Final = record.prompt_tokens or 0
completion_tokens: Final = record.completion_tokens or 0
return cls(
spend=record.spend or 0.0,
api_requests=record.api_requests or 0,
successful_requests=record.successful_requests or 0,
failed_requests=record.failed_requests or 0,
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
cache_read_input_tokens=record.cache_read_input_tokens or 0,
cache_creation_input_tokens=record.cache_creation_input_tokens or 0,
)
@classmethod
def zero(cls) -> "_ExportMetrics":
return cls(
spend=0.0,
api_requests=0,
successful_requests=0,
failed_requests=0,
total_tokens=0,
prompt_tokens=0,
completion_tokens=0,
cache_read_input_tokens=0,
cache_creation_input_tokens=0,
)
def __add__(self, other: "_ExportMetrics") -> "_ExportMetrics":
return _ExportMetrics(
spend=self.spend + other.spend,
api_requests=self.api_requests + other.api_requests,
successful_requests=self.successful_requests + other.successful_requests,
failed_requests=self.failed_requests + other.failed_requests,
total_tokens=self.total_tokens + other.total_tokens,
prompt_tokens=self.prompt_tokens + other.prompt_tokens,
completion_tokens=self.completion_tokens + other.completion_tokens,
cache_read_input_tokens=self.cache_read_input_tokens + other.cache_read_input_tokens,
cache_creation_input_tokens=self.cache_creation_input_tokens + other.cache_creation_input_tokens,
)
def _export_base_row(
record: _ExportRow,
entity_metadata_field: Mapping[str, dict[str, object]] | None,
) -> TeamDailyActivityExportRow:
entity_id: Final = record.entity_id or "Unassigned"
metrics: Final = _ExportMetrics.from_record(record)
return TeamDailyActivityExportRow(
date=record.date,
team_id=entity_id,
team_alias=_export_team_alias(entity_metadata_field, entity_id),
model=record.model,
spend=metrics.spend,
flat_cost=_reported_flat_cost(record),
api_requests=metrics.api_requests,
successful_requests=metrics.successful_requests,
failed_requests=metrics.failed_requests,
total_tokens=metrics.total_tokens,
prompt_tokens=metrics.prompt_tokens,
completion_tokens=metrics.completion_tokens,
cache_read_input_tokens=metrics.cache_read_input_tokens,
cache_creation_input_tokens=metrics.cache_creation_input_tokens,
)
def _export_key_row(
record: _ExportRow,
entity_metadata_field: Mapping[str, dict[str, object]] | None,
api_key_metadata: Mapping[str, _KeyMetadataDict],
) -> TeamDailyActivityExportRow:
entity_id: Final = record.entity_id or "Unassigned"
metadata: Final = _key_metadata(api_key_metadata, record.api_key or "")
metrics: Final = _ExportMetrics.from_record(record)
return TeamDailyActivityExportRow(
date=record.date,
team_id=entity_id,
team_alias=_export_team_alias(entity_metadata_field, entity_id),
api_key=record.api_key,
key_alias=metadata.key_alias,
user_id=metadata.user_id,
user_email=metadata.user_email,
spend=metrics.spend,
api_requests=metrics.api_requests,
successful_requests=metrics.successful_requests,
failed_requests=metrics.failed_requests,
total_tokens=metrics.total_tokens,
prompt_tokens=metrics.prompt_tokens,
completion_tokens=metrics.completion_tokens,
cache_read_input_tokens=metrics.cache_read_input_tokens,
cache_creation_input_tokens=metrics.cache_creation_input_tokens,
)
def _fold_export_users(
records: Sequence[_ExportRow],
entity_metadata_field: Mapping[str, dict[str, object]] | None,
api_key_metadata: Mapping[str, _KeyMetadataDict],
) -> tuple[TeamDailyActivityExportRow, ...]:
"""Fold (date, team, api_key) rows into (date, team, user) rows."""
def bucket_of(record: _ExportRow) -> tuple[str, str, str]:
return (
record.date,
record.entity_id or "Unassigned",
_key_metadata(api_key_metadata, record.api_key or "").user_id or "Unassigned",
)
key_sets: Final = MappingProxyType(
{
bucket: frozenset(record.api_key or "" for record in group)
for bucket, group in itertools.groupby(sorted(records, key=bucket_of), key=bucket_of)
}
)
sums: Final[dict[tuple[str, str, str], _ExportMetrics]] = {} # mutable-ok: local fold accumulator
emails: Final[dict[tuple[str, str, str], str | None]] = {} # mutable-ok: local fold accumulator
for record in records:
metadata = _key_metadata(api_key_metadata, record.api_key or "")
bucket_key = bucket_of(record)
sums[bucket_key] = sums.get(bucket_key, _ExportMetrics.zero()) + _ExportMetrics.from_record(record)
emails.setdefault(bucket_key, metadata.user_email)
if emails[bucket_key] is None and metadata.user_email is not None:
emails[bucket_key] = metadata.user_email
return tuple(
_export_folded_user_row(
bucket_key, sums[bucket_key], emails[bucket_key], len(key_sets[bucket_key]), entity_metadata_field
)
for bucket_key in sorted(sums)
)
def _export_folded_user_row(
bucket_key: tuple[str, str, str],
metrics: _ExportMetrics,
user_email: str | None,
keys: int,
entity_metadata_field: Mapping[str, dict[str, object]] | None,
) -> TeamDailyActivityExportRow:
date, entity_id, user_id = bucket_key
return TeamDailyActivityExportRow(
date=date,
team_id=entity_id,
team_alias=_export_team_alias(entity_metadata_field, entity_id),
user_id=user_id if user_id != "Unassigned" else None,
user_email=user_email,
keys=keys,
spend=metrics.spend,
api_requests=metrics.api_requests,
successful_requests=metrics.successful_requests,
failed_requests=metrics.failed_requests,
total_tokens=metrics.total_tokens,
prompt_tokens=metrics.prompt_tokens,
completion_tokens=metrics.completion_tokens,
cache_read_input_tokens=metrics.cache_read_input_tokens,
cache_creation_input_tokens=metrics.cache_creation_input_tokens,
)
async def get_daily_activity_export_rows(
*,
prisma_client: PrismaClient,
table_name: str,
entity_id_field: str,
entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
entity_metadata_field: Mapping[str, dict[str, object]] | None,
start_date: str,
end_date: str,
api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
exclude_entity_ids: list[str] | None, # mutable-ok: filter union shared with the paginated path
timezone_offset_minutes: int | None,
export_type: TeamDailyActivityExportType,
) -> tuple[TeamDailyActivityExportRow, ...]:
"""Every (date, entity[, api_key|model]) rollup row in the range, uncapped."""
sql_query, sql_params = _build_export_sql_query(
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
start_date=start_date,
end_date=end_date,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
timezone_offset_minutes=timezone_offset_minutes,
export_type=export_type,
)
raw_rows: Final = await _query_raw_optional(prisma_client, (sql_query, sql_params))
records: Final = tuple(_ExportRow(**row) for row in (raw_rows or ()))
if export_type in ("daily", "daily_with_models"):
return await asyncio.to_thread(
lambda: tuple(_export_base_row(record, entity_metadata_field) for record in records)
)
api_keys: Final = frozenset(record.api_key for record in records if record.api_key)
api_key_metadata: Final = (
await get_api_key_metadata(prisma_client, api_keys, _spend_logs_window(frozenset(r.date for r in records)))
if api_keys
else _EMPTY_KEY_METADATA
)
if export_type == "daily_with_keys":
return await asyncio.to_thread(
lambda: tuple(_export_key_row(record, entity_metadata_field, api_key_metadata) for record in records)
)
return await asyncio.to_thread(_fold_export_users, records, entity_metadata_field, api_key_metadata)
def _aggregate_spend_records_sync(
*,
records: Sequence[DailySpendRecord],

View file

@ -28,6 +28,7 @@ from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import *
from litellm.proxy.auth.auth_checks import (
@ -87,11 +88,13 @@ from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendMetadata,
SpendAnalyticsPaginatedResponse,
)
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
BulkUpdateUserResponse,
KeyActivitySearchWhere,
UserListResponse,
UserSearchWhere,
UserUpdateResult,
@ -2991,6 +2994,27 @@ async def get_user_daily_activity(
)
def _resolve_user_daily_activity_entity_id(
user_api_key_dict: UserAPIKeyAuth,
user_id: str | None,
) -> str | None:
is_admin: Final = _user_has_admin_view(user_api_key_dict)
if is_admin:
return user_id
caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict)
effective_user_id: Final = user_id if user_id is not None else caller_user_id
if effective_user_id != caller_user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={ # mutable-ok: FastAPI detail payload shape
"error": "Non-admin users can only view their own spend data."
},
)
return effective_user_id
@router.get(
"/user/daily/activity/aggregated",
tags=["Budget & Spend Tracking", "Internal User management"],
@ -3057,20 +3081,7 @@ async def get_user_daily_activity_aggregated(
)
try:
is_admin: Final = _user_has_admin_view(user_api_key_dict)
if is_admin:
entity_id = user_id # None means global view, otherwise filter by user
else:
caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict)
if user_id is None:
user_id = caller_user_id
if user_id != caller_user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": "Non-admin users can only view their own spend data."},
)
entity_id = user_id
entity_id: Final = _resolve_user_daily_activity_entity_id(user_api_key_dict, user_id)
return await get_daily_activity_aggregated(
prisma_client=prisma_client,
@ -3094,3 +3105,117 @@ async def get_user_daily_activity_aggregated(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Failed to fetch analytics: {e}"},
)
@router.get(
"/user/daily/activity/aggregated/search",
tags=["Budget & Spend Tracking", "Internal User management"], # mutable-ok: FastAPI route tags shape
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI route dependencies shape
response_model=SpendAnalyticsPaginatedResponse,
)
@management_endpoint_wrapper
async def search_user_daily_activity_keys(
search: str = fastapi.Query(
...,
min_length=1,
description="Matches keys whose hash equals the value, or whose key alias or user ID contains it (case-insensitive)",
),
start_date: str | None = fastapi.Query(
default=None,
description="Start date in YYYY-MM-DD format",
),
end_date: str | None = fastapi.Query(
default=None,
description="End date in YYYY-MM-DD format",
),
user_id: str | None = fastapi.Query(
default=None,
description="Filter by specific user ID. Admins can filter by any user or omit for global view. Non-admins must provide their own user_id.",
),
timezone: int | None = fastapi.Query(
default=None,
description="Timezone offset in minutes from UTC (e.g., 480 for PST). "
"Matches JavaScript's Date.getTimezoneOffset() convention.",
),
include_current_utc_day: bool = fastapi.Query(
default=False,
description="When the range ends on the caller's current local day, extend it to "
"today's UTC bucket so spend written after the caller's local midnight (in UTC "
"terms) is included. Requires the timezone parameter. Historical ranges are "
"never extended.",
),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI dependency injection
) -> SpendAnalyticsPaginatedResponse:
"""
Search verification tokens by exact token hash or by a case-insensitive substring of
the key alias or owning user ID, then return the aggregated daily activity for the
matches. Lets the Usage page surface keys that fell outside the top-spend subset
the aggregated endpoint loads.
"""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={ # mutable-ok: FastAPI detail payload shape
"error": CommonProxyErrors.db_not_connected_error.value
},
)
if start_date is None or end_date is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Please provide start_date and end_date"}, # mutable-ok: FastAPI detail payload shape
)
try:
entity_id: Final = _resolve_user_daily_activity_entity_id(user_api_key_dict, user_id)
search_or: Final = (
{"token": search}, # mutable-ok: prisma serializes where clauses, keep plain dicts
{"key_alias": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
{"user_id": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
)
where: Final[KeyActivitySearchWhere] = (
{"OR": search_or} # mutable-ok: prisma where clause root
if entity_id is None
else {"user_id": entity_id, "OR": search_or} # mutable-ok: prisma where clause root
)
matched_keys: Final = await VerificationTokenRepository(prisma_client).table.find_many(
where=where,
take=USAGE_TOP_API_KEYS_LIMIT,
order={"spend": "desc"}, # mutable-ok: prisma serializes order, keep it a plain dict
)
tokens: Final = [key.token for key in matched_keys] # mutable-ok: api_key filter union expects a list
if not tokens:
return SpendAnalyticsPaginatedResponse(
results=[], # mutable-ok: response model field shape
metadata=DailySpendMetadata(
api_key_limit=USAGE_TOP_API_KEYS_LIMIT,
total_api_keys=0,
),
)
return await get_daily_activity_aggregated(
prisma_client=prisma_client,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=entity_id,
entity_metadata_field=None,
start_date=start_date,
end_date=end_date,
model=None,
api_key=tokens,
timezone_offset_minutes=timezone,
include_current_utc_day=include_current_utc_day,
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception("/user/daily/activity/aggregated/search: Exception occured - %s", e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Failed to fetch analytics: {e}"}, # mutable-ok: FastAPI detail payload shape
)

View file

@ -453,7 +453,7 @@ def _regenerate_request_as_update_request(key: str, data: RegenerateKeyRequest)
)
if not changed_fields:
return None
return UpdateKeyRequest(key=key, **changed_fields)
return UpdateKeyRequest.model_validate(MappingProxyType({"key": key, **changed_fields}))
class _LegacyDumpable(Protocol):

View file

@ -11,6 +11,8 @@ All /team management endpoints
import asyncio
import copy
import csv
import io
import json
import math
import traceback
@ -33,13 +35,15 @@ from typing import (
)
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, Response, status
from fastapi.responses import JSONResponse
from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict, assert_never
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
from litellm.integrations.prometheus import PrometheusLogger
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import (
@ -125,6 +129,7 @@ from litellm.proxy.hooks.model_max_budget_limiter import (
)
from litellm.proxy.management_endpoints.common_daily_activity import (
get_daily_activity_aggregated,
get_daily_activity_export_rows,
)
from litellm.proxy.management_endpoints.common_utils import (
_check_disable_global_guardrails_caller_permission,
@ -196,6 +201,7 @@ from litellm.repositories.verification_token_repository import (
from litellm.router import Router
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendMetadata,
SpendAnalyticsPaginatedResponse,
)
from litellm.types.proxy.management_endpoints.team_endpoints import (
@ -204,7 +210,14 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
BulkUpdateTeamMemberPermissionsRequest,
BulkUpdateTeamMemberPermissionsResponse,
GetTeamMemberPermissionsResponse,
TeamDailyActivityExportFormat,
TeamDailyActivityExportMetadata,
TeamDailyActivityExportResponse,
TeamDailyActivityExportRow,
TeamDailyActivityExportType,
TeamIdSearchFilter,
TeamIdSearchMatch,
TeamKeyActivitySearchWhere,
TeamListItem,
TeamListResponse,
TeamMemberAddResult,
@ -6805,6 +6818,283 @@ async def get_team_daily_activity_aggregated(
)
_EXPORT_CSV_METRIC_HEADERS: Final = (
"Spend ($)",
"Requests",
"Successful Requests",
"Failed Requests",
"Total Tokens",
"Prompt Tokens",
"Completion Tokens",
"Cache Read Input Tokens",
"Cache Creation Input Tokens",
)
def _export_csv_headers(export_type: TeamDailyActivityExportType) -> tuple[str, ...]:
base: Final = ("Date", "Team", "Team ID")
if export_type == "daily_with_keys":
return (*base, "Key Alias", "Key ID", "User ID", "User Email", *_EXPORT_CSV_METRIC_HEADERS)
if export_type == "daily_with_users":
return (*base, "User ID", "User Email", "Keys", *_EXPORT_CSV_METRIC_HEADERS)
if export_type == "daily_with_models":
return (
*base,
"Model",
"Spend ($)",
"Requests",
"Successful",
"Failed",
"Total Tokens",
"Prompt Tokens",
"Completion Tokens",
"Cache Read Input Tokens",
"Cache Creation Input Tokens",
)
return (*base, *_EXPORT_CSV_METRIC_HEADERS)
def _csv_safe(value: str) -> str:
return "'" + value if value[:1] in ("=", "+", "-", "@", "\t", "\r") else value
def _export_csv_record(row: TeamDailyActivityExportRow) -> dict[str, object]:
return { # mutable-ok: csv.DictWriter consumes a plain mapping per row
"Date": row.date,
"Team": _csv_safe(row.team_alias) if row.team_alias else "-",
"Team ID": row.team_id,
"Key Alias": _csv_safe(row.key_alias) if row.key_alias else "-",
"Key ID": row.api_key or "-",
"User ID": _csv_safe(row.user_id) if row.user_id else "-",
"User Email": _csv_safe(row.user_email) if row.user_email else "-",
"Keys": row.keys,
"Model": _csv_safe(row.model) if row.model else "-",
"Spend ($)": f"{row.spend:.4f}",
"Flat Cost ($)": f"{row.flat_cost:.4f}",
"Total Cost ($)": f"{row.spend + row.flat_cost:.4f}",
"Requests": row.api_requests,
"Successful Requests": row.successful_requests,
"Failed Requests": row.failed_requests,
"Successful": row.successful_requests,
"Failed": row.failed_requests,
"Total Tokens": row.total_tokens,
"Prompt Tokens": row.prompt_tokens,
"Completion Tokens": row.completion_tokens,
"Cache Read Input Tokens": row.cache_read_input_tokens,
"Cache Creation Input Tokens": row.cache_creation_input_tokens,
}
def _team_export_csv(export_type: TeamDailyActivityExportType, rows: Sequence[TeamDailyActivityExportRow]) -> str:
base_headers: Final = _export_csv_headers(export_type)
spend_index: Final = base_headers.index("Spend ($)") + 1
headers: Final = (
(*base_headers[:spend_index], "Flat Cost ($)", "Total Cost ($)", *base_headers[spend_index:])
if sum(row.flat_cost for row in rows) > 0
else base_headers
)
buffer: Final = io.StringIO()
writer: Final = csv.DictWriter(buffer, fieldnames=headers, extrasaction="ignore")
writer.writeheader()
writer.writerows(_export_csv_record(row) for row in rows)
return buffer.getvalue()
@router.get(
"/team/daily/activity/export",
response_model=TeamDailyActivityExportResponse,
responses={200: {"content": {"text/csv": {}, "application/json": {}}}}, # mutable-ok: OpenAPI content map
tags=["team management"], # mutable-ok: fastapi's decorator signature types tags as a list
)
async def get_team_daily_activity_export(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
start_date: str | None = None,
end_date: str | None = None,
export_type: TeamDailyActivityExportType = "daily",
format: TeamDailyActivityExportFormat = "csv",
team_id: str | None = None,
exclude_team_ids: str | None = None,
timezone_offset: Annotated[int | None, Query(alias="timezone")] = None,
) -> Response:
"""
Server-side Team Usage export, not subject to USAGE_TOP_API_KEYS_LIMIT.
Same scoping as /team/daily/activity/aggregated, answered by one unbounded
rollup query, returned as CSV or JSON. For daily_with_keys,
daily_with_users and daily_with_models the PTU sentinel flat-cost rows are
excluded, so metadata totals under those export types cover request spend
only; the plain daily export includes them.
"""
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if prisma_client is None:
raise _daily_activity_error(status_code=500, message=CommonProxyErrors.db_not_connected_error.value)
range_error: Final = _aggregated_date_range_error(start_date, end_date)
if range_error is not None or start_date is None or end_date is None:
raise _daily_activity_error(status_code=400, message=range_error or "Please provide start_date and end_date")
scope: Final = await _resolve_team_daily_activity_scope(
team_ids=team_id,
exclude_team_ids=exclude_team_ids,
api_key=None,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
rows: Final = await get_daily_activity_export_rows(
prisma_client=prisma_client,
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=scope.team_ids,
entity_metadata_field=scope.team_alias_metadata,
start_date=start_date,
end_date=end_date,
api_key=scope.api_key_filter,
exclude_entity_ids=scope.exclude_team_ids,
timezone_offset_minutes=timezone_offset,
export_type=export_type,
)
now: Final = datetime.now(timezone.utc)
metadata: Final = TeamDailyActivityExportMetadata(
export_date=now.isoformat(),
export_type=export_type,
start_date=start_date,
end_date=end_date,
team_ids=list(scope.team_ids) if scope.team_ids else None, # mutable-ok: response model field type
total_spend=sum(row.spend for row in rows),
total_flat_cost=sum(row.flat_cost for row in rows),
total_api_requests=sum(row.api_requests for row in rows),
total_successful_requests=sum(row.successful_requests for row in rows),
total_failed_requests=sum(row.failed_requests for row in rows),
total_tokens=sum(row.total_tokens for row in rows),
)
if format == "json":
return JSONResponse(
content=TeamDailyActivityExportResponse(metadata=metadata, data=rows).model_dump(mode="json")
)
return Response(
content=_team_export_csv(export_type, rows),
media_type="text/csv; charset=utf-8",
headers={ # mutable-ok: starlette Response headers is a dict
"Content-Disposition": f'attachment; filename="team_usage_{export_type}_{now.date().isoformat()}.csv"'
},
)
def _team_key_search_where(*, search: str, scope: _TeamDailyActivityScope) -> TeamKeyActivitySearchWhere:
"""Caller scoping lives inside the same Prisma where as the search term so `take`
never trims visible matches in favour of keys the caller is not allowed to see."""
search_or: Final = (
{"token": search}, # mutable-ok: prisma where clause leaf
{"key_alias": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
{"user_id": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
)
own_keys: Final = tuple(scope.api_key_filter) if isinstance(scope.api_key_filter, list) else None
team_filter: Final[TeamIdSearchFilter | None] = (
{ # mutable-ok: prisma where clause leaf
"in": tuple(scope.team_ids),
"notIn": tuple(scope.exclude_team_ids),
}
if scope.team_ids is not None and scope.exclude_team_ids is not None
else {"in": tuple(scope.team_ids)} # mutable-ok: prisma where clause leaf
if scope.team_ids is not None
else {"notIn": tuple(scope.exclude_team_ids)} # mutable-ok: prisma where clause leaf
if scope.exclude_team_ids is not None
else None
)
if team_filter is None and own_keys is None:
return {"OR": search_or} # mutable-ok: prisma where clause root
if team_filter is None and own_keys is not None:
return {"token": {"in": own_keys}, "OR": search_or} # mutable-ok: prisma where clause root
if team_filter is not None and own_keys is None:
return {"team_id": team_filter, "OR": search_or} # mutable-ok: prisma where clause root
assert team_filter is not None and own_keys is not None
return { # mutable-ok: prisma where clause root
"team_id": team_filter,
"token": {"in": own_keys}, # mutable-ok: prisma where clause leaf
"OR": search_or,
}
@router.get(
"/team/daily/activity/aggregated/search",
response_model=SpendAnalyticsPaginatedResponse,
tags=["team management"], # mutable-ok: FastAPI route tags shape
)
async def search_team_daily_activity_keys(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
search: str = fastapi.Query(
...,
min_length=1,
description="Exact token hash, or a case-insensitive substring of the key alias or owning user id",
),
team_ids: str | None = None,
start_date: str | None = None,
end_date: str | None = None,
exclude_team_ids: str | None = None,
timezone: int | None = None,
) -> SpendAnalyticsPaginatedResponse:
"""Aggregated daily team activity for the keys matching `search`, across every key the caller may
see rather than only the top USAGE_TOP_API_KEYS_LIMIT keys by spend."""
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if prisma_client is None:
raise _daily_activity_error(status_code=500, message=CommonProxyErrors.db_not_connected_error.value)
range_error: Final = _aggregated_date_range_error(start_date, end_date)
if range_error is not None:
raise _daily_activity_error(status_code=400, message=range_error)
scope: Final = await _resolve_team_daily_activity_scope(
team_ids=team_ids,
exclude_team_ids=exclude_team_ids,
api_key=None,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
matched_keys: Final = await _tokens_db(prisma_client).find_many(
where=_team_key_search_where(search=search, scope=scope),
take=USAGE_TOP_API_KEYS_LIMIT,
order={"spend": "desc"}, # mutable-ok: prisma serializes order, keep it a plain dict
)
tokens: Final = [key.token for key in matched_keys] # mutable-ok: get_daily_activity_aggregated takes list[str]
if not tokens:
return SpendAnalyticsPaginatedResponse(
results=[], # mutable-ok: response model field shape
metadata=DailySpendMetadata(api_key_limit=USAGE_TOP_API_KEYS_LIMIT, total_api_keys=0),
)
return await get_daily_activity_aggregated(
prisma_client=prisma_client,
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=scope.team_ids,
entity_metadata_field=scope.team_alias_metadata,
start_date=start_date,
end_date=end_date,
model=None,
api_key=tokens,
exclude_entity_ids=scope.exclude_team_ids,
timezone_offset_minutes=timezone,
include_entity_breakdown=True,
)
def _team_user_spend_sql(*, team_count: int, restrict_to_user: bool) -> str:
team_placeholders: Final = ", ".join(f"${i}" for i in range(3, 3 + team_count))
user_clause: Final = f' AND sl."user" = ${3 + team_count}' if restrict_to_user else ""

View file

@ -8,10 +8,25 @@ from typing_extensions import assert_never
from litellm.proxy._types import ProxyException
BATCH_LINE_REQUIRED_KEYS: Final = ("custom_id", "method", "url", "body")
_MB: Final = 1024 * 1024
@dataclass(frozen=True, slots=True)
class BatchLineShape:
required_keys: tuple[str, ...]
hint: str
BATCH_LINE_SHAPE: Final = BatchLineShape(
required_keys=("custom_id", "method", "url", "body"),
hint="Each line must be a JSON object with keys custom_id, method, url, body",
)
PASSTHROUGH_BATCH_LINE_SHAPE: Final = BatchLineShape(
required_keys=("request",),
hint="A passthrough upload takes native Vertex batch rows, so each line must be a JSON object with a request key",
)
@dataclass(frozen=True, slots=True)
class BatchFileTooLarge:
size_bytes: int
@ -42,6 +57,7 @@ class BatchFileLineNotObject:
class BatchFileMissingLineKey:
line_number: int
key: str
line_shape: BatchLineShape = BATCH_LINE_SHAPE
BatchFileValidationFailure = (
@ -70,20 +86,20 @@ def _iter_lines(file_source: bytes | BinaryIO) -> Iterator[bytes]:
return iter(file_source)
def _check_line(line_number: int, raw_line: bytes) -> BatchFileValidationFailure | None:
def _check_line(line_number: int, raw_line: bytes, line_shape: BatchLineShape) -> BatchFileValidationFailure | None:
try:
parsed: Final = json.loads(raw_line)
except (json.JSONDecodeError, UnicodeDecodeError):
return BatchFileInvalidJsonLine(line_number=line_number)
if not isinstance(parsed, dict):
return BatchFileLineNotObject(line_number=line_number)
missing: Final = next((key for key in BATCH_LINE_REQUIRED_KEYS if key not in parsed), None)
missing: Final = next((key for key in line_shape.required_keys if key not in parsed), None)
if missing is None:
return None
return BatchFileMissingLineKey(line_number=line_number, key=missing)
return BatchFileMissingLineKey(line_number=line_number, key=missing, line_shape=line_shape)
def _scan_lines(file_source: bytes | BinaryIO) -> BatchFileValidationFailure | None:
def _scan_lines(file_source: bytes | BinaryIO, line_shape: BatchLineShape) -> BatchFileValidationFailure | None:
content_lines: Final = (
(line_number, raw_line)
for line_number, raw_line in enumerate(_iter_lines(file_source), start=1)
@ -96,7 +112,7 @@ def _scan_lines(file_source: bytes | BinaryIO) -> BatchFileValidationFailure | N
(
failure
for line_number, raw_line in chain((first_line,), content_lines)
for failure in (_check_line(line_number, raw_line),)
for failure in (_check_line(line_number, raw_line, line_shape),)
if failure is not None
),
None,
@ -107,6 +123,7 @@ def check_batch_file_upload(
filename: str | None,
file_source: bytes | BinaryIO,
max_batch_file_size_mb: int | None,
line_shape: BatchLineShape = BATCH_LINE_SHAPE,
) -> BatchFileValidationFailure | None:
if filename is None or not filename.lower().endswith(".jsonl"):
return BatchFileWrongExtension(filename=filename or "")
@ -114,7 +131,7 @@ def check_batch_file_upload(
size_bytes: Final = _file_size_bytes(file_source)
if size_bytes > max_batch_file_size_mb * _MB:
return BatchFileTooLarge(size_bytes=size_bytes, limit_mb=max_batch_file_size_mb)
scan_failure: Final = _scan_lines(file_source)
scan_failure: Final = _scan_lines(file_source, line_shape)
if not isinstance(file_source, bytes):
file_source.seek(0)
return scan_failure
@ -169,11 +186,11 @@ def raise_batch_file_validation_failure(failure: BatchFileValidationFailure) ->
param="file",
code=400,
)
case BatchFileMissingLineKey(line_number=line_number, key=key):
case BatchFileMissingLineKey(line_number=line_number, key=key, line_shape=line_shape):
raise ProxyException(
message=(
f"Missing required parameter: '{key}' (batch input file line {line_number}). "
f"Each line must be a JSON object with keys {', '.join(BATCH_LINE_REQUIRED_KEYS)}. "
f"{line_shape.hint}. "
"The file was not forwarded to the provider."
),
type="invalid_request_error",

View file

@ -57,6 +57,8 @@ from litellm.proxy.common_utils.openai_error_payload import (
openai_error_type,
)
from litellm.proxy.openai_files_endpoints.batch_file_validation import (
BATCH_LINE_SHAPE,
PASSTHROUGH_BATCH_LINE_SHAPE,
check_batch_file_upload,
raise_batch_file_validation_failure,
)
@ -207,10 +209,91 @@ def get_files_provider_config(
return None
def _deployment_provider(llm_router: Router, model_id: str, team_id: str | None) -> str | None:
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model_id, team_id=team_id)
return None if credentials is None else credentials.get("custom_llm_provider")
def _resolves_to_vertex_deployments_only(llm_router: Router | None, model_name: str, team_id: str | None) -> bool:
if llm_router is None or _deployment_provider(llm_router, model_name, team_id) != "vertex_ai":
return False
return all(
_deployment_provider(llm_router, str(deployment["model_info"]["id"]), team_id) == "vertex_ai"
for deployment in llm_router.get_model_list(model_name=model_name, team_id=team_id) or ()
if "id" in deployment.get("model_info", {})
)
def _validate_passthrough_upload(
*,
purpose: str,
target_model_names: Sequence[str],
model: str | None,
target_storage: str | None,
llm_router: Router | None,
team_id: str | None,
) -> None:
if purpose != "batch":
raise ProxyException(
message=(
"`passthrough` uploads the file bytes unchanged for a native Vertex batch, "
f"so purpose must be 'batch', got '{purpose}'."
),
type="invalid_request_error",
param="passthrough",
code=400,
)
if target_storage and target_storage != "default":
raise ProxyException(
message=(
"`passthrough` writes the native batch file to the Vertex AI deployment's GCS bucket, "
f"so it cannot be combined with target_storage='{target_storage}'."
),
type="invalid_request_error",
param="target_storage",
code=400,
)
named_deployments: Final = (
*(("target_model_names", name) for name in target_model_names),
*((("model", model),) if model else ()),
)
if not named_deployments:
raise ProxyException(
message=(
"`passthrough` needs the Vertex AI deployment that will run the batch, "
"since native rows carry no model: pass `target_model_names` or `model`."
),
type="invalid_request_error",
param="target_model_names",
code=400,
)
offending: Final = next(
(
(param, name)
for param, name in named_deployments
if not _resolves_to_vertex_deployments_only(llm_router, name, team_id)
),
None,
)
if offending is None:
return
param, name = offending
raise ProxyException(
message=(
f"`passthrough` is only supported for Vertex AI deployments; '{name}' does not resolve "
"to vertex_ai deployments only."
),
type="invalid_request_error",
param=param,
code=400,
)
async def _scan_batch_upload(
*,
file_source: bytes | BinaryIO,
purpose: str,
passthrough: bool,
request_metadata: Mapping[str, object],
user_api_key_dict: UserAPIKeyAuth,
proxy_logging_obj: ProxyLogging,
@ -222,6 +305,17 @@ async def _scan_batch_upload(
or not proxy_logging_obj.has_pre_call_guardrails(request_metadata)
):
return None
if passthrough:
raise ProxyException(
message=(
"Batch guardrails cannot scan native Vertex batch rows, so a `passthrough` upload is refused "
"when the key, team, or request has pre-call guardrails configured. "
"The file was not forwarded to the provider."
),
type="invalid_request_error",
param="passthrough",
code=400,
)
outcome: Final = await scan_batch_input_file(
file_source=file_source,
request_metadata=request_metadata,
@ -458,6 +552,7 @@ async def create_file(
custom_llm_provider: str = Form(default="openai"),
file: UploadFile = File(...),
litellm_metadata: str | None = Form(default=None),
passthrough: bool = Form(default=False),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
@ -560,17 +655,28 @@ async def create_file(
if blocked_extension_failure is not None:
raise_upload_validation_failure(blocked_extension_failure)
if passthrough:
_validate_passthrough_upload(
purpose=purpose,
target_model_names=target_model_names_list,
model=model_param,
target_storage=target_storage,
llm_router=llm_router,
team_id=user_api_key_dict.team_id,
)
if purpose == "batch":
batch_file_failure: Final = await asyncio.to_thread(
check_batch_file_upload,
file.filename,
file_source,
_MAX_BATCH_FILE_SIZE_MB_ADAPTER.validate_python(general_settings.get("max_batch_file_size_mb")),
PASSTHROUGH_BATCH_LINE_SHAPE if passthrough else BATCH_LINE_SHAPE,
)
if batch_file_failure is not None:
raise_batch_file_validation_failure(batch_file_failure)
data = {}
data = {"passthrough": True} if passthrough else {}
# Parse expires_after if provided
expires_after: FileExpiresAfter | None = None
@ -673,6 +779,7 @@ async def create_file(
scan_result: Final = await _scan_batch_upload(
file_source=file_source,
purpose=purpose,
passthrough=passthrough,
request_metadata=request_metadata,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,

View file

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

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