chore: merge main into litellm_batch_jsonl_line_item_callbacks
Some checks failed
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Modules / fmt, validate, test (aws) (push) Has been cancelled
Terraform Modules / fmt, validate, test (gcp) (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-24 23:39:19 +00:00
commit dd9d4fd559
323 changed files with 14434 additions and 7577 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
@ -200,7 +205,7 @@ jobs:
tests/test_litellm/proxy/types_utils
tests/test_litellm/proxy/logging_endpoints
tests/test_litellm/proxy/test_*.py
tests/test_gateway
fork-flag: proxy-infra
workers: 4
reruns: 2
timeout-minutes: 20
@ -208,11 +213,8 @@ jobs:
- shard: caching-local
artifact-name: caching-local
test-path: >-
tests/local_testing/test_cache_preset_key.py
tests/local_testing/test_caching_handler.py
tests/local_testing/test_responses_stream_cache_keys.py
tests/local_testing/test_unit_test_caching.py
test-path: ""
fork-flag: caching-local
workers: 2
reruns: 2
timeout-minutes: 20
@ -220,7 +222,8 @@ jobs:
- shard: proxy-extras
artifact-name: proxy-extras
test-path: "tests/litellm-proxy-extras"
test-path: ""
fork-flag: proxy-extras
workers: 2
reruns: 2
timeout-minutes: 20
@ -228,7 +231,8 @@ jobs:
- shard: enterprise-package
artifact-name: enterprise-package
test-path: "tests/enterprise"
test-path: ""
fork-flag: enterprise-package
workers: 4
reruns: 2
timeout-minutes: 20
@ -247,6 +251,7 @@ jobs:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
fork-flag: ${{ matrix.fork-flag || '' }}
workers: ${{ matrix.workers }}
reruns: ${{ matrix.reruns }}
timeout-minutes: ${{ matrix.timeout-minutes }}

View file

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

View file

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

View file

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

View file

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

View file

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

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

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

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

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

@ -46,6 +46,3 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
)),
}
}
#[cfg(test)]
mod tests;

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,7 +1,11 @@
use std::{sync::Arc, time::Duration};
use futures_util::future::BoxFuture;
use litellm_http::request::{has_bearer_auth, has_header};
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::{
@ -9,14 +13,6 @@ use tokio::{
net::{TcpListener, TcpStream},
};
use super::{
Error,
common_utils::{messages_provider_config, string_headers, truncate_error_body},
messages,
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
};
use crate::messages::types::{MessagesRequest, MessagesShaping};
struct RecordingSecrets {
values: Vec<(&'static str, String)>,
fails: bool,
@ -73,55 +69,6 @@ fn secrets_call() -> MessagesCall {
}
}
#[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}}"#;
socket
.write_all(write_response(response_body).as_bytes())
.await
.expect("writes response");
request
});
let secrets = Arc::new(RecordingSecrets::new(
vec![
("ANTHROPIC_API_KEY", "sk-from-manager".to_string()),
("ANTHROPIC_BASE_URL", format!("http://{addr}")),
],
false,
));
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<_>>()
);
}
#[tokio::test]
async fn route_surfaces_a_secret_manager_failure_before_the_call() {
let Err(error) = litellm_host::run::run(
@ -179,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");

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;

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

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

@ -557,6 +557,8 @@ SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS: Final[float] = float(
request_timeout: float = float(os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS))))
request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ
DEFAULT_A2A_AGENT_TIMEOUT: Final[float] = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes
AGENT_KILL_SWITCH_TIMEOUT_SECONDS: Final = 10.0
AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS: Final = 2000
# Patterns that indicate a localhost/internal URL in A2A agent cards that should be
# replaced with the original base_url. This is a common misconfiguration where
# developers deploy agents with development URLs in their agent cards.

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

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

View file

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

View file

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

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

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

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

View file

@ -6092,7 +6092,7 @@
"cache_creation_input_audio_token_cost": 3e-07,
"cache_read_input_audio_token_cost": 3e-07,
"cache_read_input_token_cost": 6e-08,
"deprecation_date": "2027-06-25",
"deprecation_date": "2027-07-31",
"input_cost_per_audio_token": 1e-05,
"input_cost_per_image_token": 8e-07,
"input_cost_per_token": 6e-07,
@ -26212,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,
@ -26313,10 +26312,14 @@
"gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_priority": 3.6e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_priority": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -26326,7 +26329,9 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_priority": 2.16e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
@ -27849,6 +27854,7 @@
"gemini-embedding-001": {
"deprecation_date": "2028-05-20",
"input_cost_per_token": 1.5e-07,
"input_cost_per_token_batches": 1.2e-07,
"litellm_provider": "vertex_ai-embedding-models",
"max_input_tokens": 2048,
"max_tokens": 2048,
@ -49416,7 +49422,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,
@ -49494,10 +49499,14 @@
"vertex_ai/gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_priority": 3.6e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_priority": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -49507,7 +49516,9 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_priority": 2.16e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"supports_reasoning": false,
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
@ -68612,7 +68623,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": {
@ -69473,6 +69486,7 @@
"input_cost_per_audio_token": 3e-06,
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -69493,6 +69507,7 @@
"input_cost_per_audio_token": 3e-06,
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -69774,6 +69789,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,
@ -69805,6 +69821,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",
@ -69822,6 +69839,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,
@ -69843,6 +69861,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",
@ -69861,6 +69880,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

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

View file

@ -238,6 +238,7 @@ class LitellmTableNames(str, enum.Enum):
CONFIG_TABLE_NAME = "LiteLLM_Config"
SSO_CONFIG_TABLE_NAME = "LiteLLM_SSOConfig"
UI_SETTINGS_TABLE_NAME = "LiteLLM_UISettings"
AGENT_TABLE_NAME = "LiteLLM_AgentsTable"
class Litellm_EntityType(enum.Enum):
@ -307,6 +308,7 @@ 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
@ -577,6 +579,7 @@ class LiteLLMRoutes(enum.Enum):
"/v1/agents/{agent_id}",
"/v1/agents/make_public",
"/v1/agents/{agent_id}/make_public",
"/v1/agents/{agent_id}/kill_switch",
)
# Backwards-compat union — virtual keys may be configured with
@ -719,6 +722,7 @@ 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
@ -890,6 +894,7 @@ 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",
@ -990,6 +995,7 @@ 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",
@ -3688,7 +3694,7 @@ from litellm.models.spend_logs import ( # noqa: E402
)
from litellm.models.tag import LiteLLM_TagTable as LiteLLM_TagTable # noqa: E402
AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "unblocked", "rotated"]
AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "unblocked", "rotated", "kill_switch_fired"]
class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase):

View file

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

View file

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

View file

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

View file

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

@ -11,6 +11,8 @@ All /team management endpoints
import asyncio
import copy
import csv
import io
import json
import math
import traceback
@ -33,7 +35,8 @@ 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
@ -126,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,
@ -206,6 +210,11 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
BulkUpdateTeamMemberPermissionsRequest,
BulkUpdateTeamMemberPermissionsResponse,
GetTeamMemberPermissionsResponse,
TeamDailyActivityExportFormat,
TeamDailyActivityExportMetadata,
TeamDailyActivityExportResponse,
TeamDailyActivityExportRow,
TeamDailyActivityExportType,
TeamIdSearchFilter,
TeamIdSearchMatch,
TeamKeyActivitySearchWhere,
@ -6809,6 +6818,178 @@ 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."""

View file

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

View file

@ -152,6 +152,7 @@ from litellm.router_utils.auto_router_tuning_baseline import (
snapshot_tuning_baselines,
tuning_limit_violation,
)
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.router_utils.routing_groups import parse_routing_groups
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.utils import (
@ -11117,6 +11118,40 @@ class ProxyStartupEvent:
#### API ENDPOINTS ####
async def _names_hidden_by_listing_callbacks(
user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
) -> frozenset[str]:
hidden: Final = await proxy_logging_obj.hidden_by_listing_callbacks(user_api_key_dict, model_names)
if not hidden or llm_router is None:
return hidden
aliases: Final = llm_router.model_group_alias
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings)
return hidden | frozenset(
alias
for alias in aliases
if (target := resolve_model_group_alias(aliases, alias)) is not None
and internal_to_public.get(target, target) in hidden
)
async def _entries_kept_by_listing_callbacks(
entries: Sequence[tuple[str, str]], user_api_key_dict: UserAPIKeyAuth
) -> tuple[tuple[str, str], ...]:
hidden: Final = await _names_hidden_by_listing_callbacks(
user_api_key_dict, tuple(response_id for response_id, _ in entries)
)
if not hidden:
return tuple(entries)
return tuple(entry for entry in entries if entry[0] not in hidden)
async def _deployment_hidden_by_listing_callbacks(deployment: Deployment, user_api_key_dict: UserAPIKeyAuth) -> bool:
listed_name: Final = _translate_model_name_for_response(deployment.model_dump(exclude_none=True)).get("model_name")
if not isinstance(listed_name, str):
return False
return listed_name in await _names_hidden_by_listing_callbacks(user_api_key_dict, (listed_name,))
@router.get("/v1/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"])
@router.get(
"/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"]
@ -11267,7 +11302,9 @@ async def model_list(
# The internal routing key drives the metadata/fallback lookup, while the
# public name is what the client sees as the model id.
model_data = []
admin_entries: Final = TeamModelNameTranslator.listing_entries(all_models, llm_router, settings)
admin_entries: Final = await _entries_kept_by_listing_callbacks(
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict
)
for response_id, lookup_id in admin_entries:
model_info = create_model_info_response(
model_id=lookup_id,
@ -11323,7 +11360,10 @@ async def model_list(
# public name is what the client sees as the model id.
model_data = []
entries: Final = alias_listing_entries(
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), caller_aliases
await _entries_kept_by_listing_callbacks(
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict
),
caller_aliases,
)
for response_id, lookup_id in entries:
model_info = create_model_info_response(
@ -11417,13 +11457,24 @@ async def model_info(
llm_router=llm_router,
)
hidden_names: Final = blocked_names | unhealthy_names
if hidden_names:
all_models = [m for m in all_models if m not in hidden_names]
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings)
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(
user_api_key_dict,
tuple(
response_id
for response_id, _ in TeamModelNameTranslator.listing_entries(
tuple(m for m in all_models if m not in hidden_names), llm_router, settings
)
),
)
if hidden_names or callback_hidden_names:
all_models = [
m for m in all_models if m not in hidden_names and internal_to_public.get(m, m) not in callback_hidden_names
]
undiscoverable_names: Final = undiscoverable_model_names(
all_models, llm_router, user_api_key_dict, team_id or user_api_key_dict.team_id
)
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings)
aliased_model_id: Final = alias_target(
model_id,
caller_alias_maps(
@ -13728,7 +13779,7 @@ async def transform_request(request: TransformRequestBody):
except ValueError as e:
raise HTTPException(status_code=400, detail={"error": str(e)})
return return_raw_request(endpoint=request.call_type, kwargs=request.request_body)
return await asyncio.to_thread(return_raw_request, request.call_type, request.request_body)
async def _check_if_model_is_user_added(
@ -15743,7 +15794,7 @@ async def model_info_v1(
if litellm_model_id is not None:
# user is trying to get specific model from litellm router
deployment_info: Final = llm_router.get_deployment(model_id=litellm_model_id)
if deployment_info is None:
if deployment_info is None or await _deployment_hidden_by_listing_callbacks(deployment_info, user_api_key_dict):
raise HTTPException(
status_code=400,
detail={"error": f"Model id = {litellm_model_id} not found on litellm proxy"},
@ -15832,10 +15883,17 @@ async def model_info_v1(
general_settings=general_settings,
llm_router=llm_router,
)
visible_models: Final = discoverable_rows(
servable_rows: Final = discoverable_rows(
(model for model in all_models if model.get("model_name") not in hidden_names),
user_api_key_dict,
)
listed_names: Final = tuple(
dict.fromkeys(name for model in servable_rows if isinstance(name := model.get("model_name"), str))
)
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(user_api_key_dict, listed_names)
visible_models: Final = tuple(
model for model in servable_rows if model.get("model_name") not in callback_hidden_names
)
verbose_proxy_logger.debug("all_models: %s", visible_models)
return _model_info_json_response(visible_models)
@ -15884,7 +15942,7 @@ async def model_deprecations(
def _get_model_group_info(
llm_router: Router, all_models_str: list[str], model_group: str | None
llm_router: Router, all_models_str: Sequence[str], model_group: str | None
) -> list[ModelGroupInfoProxy]:
model_groups: Final[list[ModelGroupInfoProxy]] = []
@ -16117,23 +16175,34 @@ async def model_group_info(
undiscoverable_group_names: Final = undiscoverable_model_names(
all_models_str, llm_router, user_api_key_dict, user_api_key_dict.team_id
)
model_groups: list[ModelGroupInfoProxy] = _get_model_group_info(
llm_router=llm_router,
all_models_str=[name for name in all_models_str if name not in undiscoverable_group_names],
model_group=model_group,
)
listed_group_names: Final = tuple(name for name in all_models_str if name not in undiscoverable_group_names)
# Append A2A agents to model groups
from litellm.proxy.agent_endpoints.model_list_helpers import (
append_agents_to_model_group,
)
model_groups = await append_agents_to_model_group(
model_groups=model_groups,
model_groups: Final = await append_agents_to_model_group(
model_groups=_get_model_group_info(
llm_router=llm_router, all_models_str=listed_group_names, model_group=model_group
),
user_api_key_dict=user_api_key_dict,
)
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings)
public_group_names: Final = tuple(
internal_to_public.get(group.model_group, group.model_group) for group in model_groups
)
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(
user_api_key_dict, tuple(dict.fromkeys(public_group_names))
)
return {"data": model_groups}
return {
"data": [
group
for group, public_name in zip(model_groups, public_group_names, strict=True)
if public_name not in callback_hidden_names
]
}
@router.get(

View file

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

View file

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

View file

@ -118,6 +118,7 @@ from litellm.llms.openai_like.model_info import (
MODEL_INFO_REFRESH_SECONDS,
get_openai_compatible_model_info,
)
from litellm.router_strategy.base_routing_strategy import BaseRoutingStrategy
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.router_strategy.complexity_router.context_compaction import (
arm_compaction,
@ -1377,6 +1378,9 @@ class Router:
`_init_routing_groups`) so repeated `update_settings` calls don't
accumulate dead selectors that keep receiving callback events.
"""
for selector in selectors:
if isinstance(selector, BaseRoutingStrategy):
selector.retire()
selector_ids: Final = {id(s) for s in selectors if s is not None}
if not selector_ids:
return
@ -12117,7 +12121,7 @@ class Router:
)
rebuild_routing_groups = True
elif var == "routing_strategy_args":
routing_args_updated = True
routing_args_updated = value != self.routing_strategy_args
setattr(self, var, value)
else:
verbose_router_logger.debug("Setting %s is not allowed", var)

View file

@ -40,10 +40,24 @@ class BaseRoutingStrategy(ABC):
self.periodic_sync_in_memory_spend_with_redis(default_sync_interval=default_sync_interval)
)
def cancel_sync_task(self) -> None:
if self._sync_task is not None:
self._sync_task.cancel()
def retire(self) -> None:
self.cancel_sync_task()
if not self.redis_increment_operation_queue:
return
try:
loop: Final = asyncio.get_running_loop()
except RuntimeError:
return
loop.create_task(self._push_in_memory_increments_to_redis())
async def cleanup(self):
"""Cleanup method to be called when shutting down"""
if self._sync_task is not None:
self._sync_task.cancel()
self.cancel_sync_task()
try:
await self._sync_task
except asyncio.CancelledError:

View file

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

View file

@ -26,6 +26,7 @@ class httpxSpecialProvider(str, Enum):
RAG = "rag"
A2AProvider = "a2a_provider"
AgentHealthCheck = "agent_health_check"
AgentKillSwitch = "agent_kill_switch"
A2A = "a2a"
PromptManagement = "prompt_management"
UI = "ui"

View file

@ -277,3 +277,48 @@ class TeamUserSpendResponse(BaseModel):
start_date: str
end_date: str
results: tuple[TeamUserSpendRow, ...]
TeamDailyActivityExportType = Literal["daily", "daily_with_keys", "daily_with_users", "daily_with_models"]
TeamDailyActivityExportFormat = Literal["csv", "json"]
class TeamDailyActivityExportRow(BaseModel):
date: str
team_id: str
team_alias: str | None = None
api_key: str | None = None
key_alias: str | None = None
user_id: str | None = None
user_email: str | None = None
keys: int | None = None
model: str | None = None
spend: float
flat_cost: float = 0.0
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
class TeamDailyActivityExportMetadata(BaseModel):
export_date: str
export_type: TeamDailyActivityExportType
start_date: str
end_date: str
team_ids: list[str] | None
total_spend: float
total_flat_cost: float = 0.0
total_api_requests: int
total_successful_requests: int
total_failed_requests: int
total_tokens: int
class TeamDailyActivityExportResponse(BaseModel):
metadata: TeamDailyActivityExportMetadata
data: list[TeamDailyActivityExportRow]

View file

@ -10273,7 +10273,7 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
"""
from datetime import datetime
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.litellm_logging import Logging, RawRequestCaptured
litellm_logging_obj: Final = Logging(
model="gpt-3.5-turbo",
@ -10284,6 +10284,7 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
start_time=datetime.now(),
function_id="1234",
log_raw_request_response=True,
raw_request_only=True,
)
llm_api_endpoint: Final = getattr(litellm, endpoint.value)
@ -10294,7 +10295,11 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
llm_api_endpoint(
**kwargs,
litellm_logging_obj=litellm_logging_obj,
api_key="my-fake-api-key", # 👈 ensure the request fails
api_key="my-fake-api-key",
)
except RawRequestCaptured:
received_exception = (
"raw request was not captured before the provider call; check the proxy logs for the pre_call error"
)
except Exception as e:
received_exception = str(e)

View file

@ -6092,7 +6092,7 @@
"cache_creation_input_audio_token_cost": 3e-07,
"cache_read_input_audio_token_cost": 3e-07,
"cache_read_input_token_cost": 6e-08,
"deprecation_date": "2027-06-25",
"deprecation_date": "2027-07-31",
"input_cost_per_audio_token": 1e-05,
"input_cost_per_image_token": 8e-07,
"input_cost_per_token": 6e-07,
@ -26212,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,
@ -26313,10 +26312,14 @@
"gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_priority": 3.6e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_priority": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -26326,7 +26329,9 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_priority": 2.16e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
@ -27849,6 +27854,7 @@
"gemini-embedding-001": {
"deprecation_date": "2028-05-20",
"input_cost_per_token": 1.5e-07,
"input_cost_per_token_batches": 1.2e-07,
"litellm_provider": "vertex_ai-embedding-models",
"max_input_tokens": 2048,
"max_tokens": 2048,
@ -49416,7 +49422,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,
@ -49494,10 +49499,14 @@
"vertex_ai/gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_priority": 3.6e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
"cache_read_input_token_cost_batches": 1e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_priority": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
"input_cost_per_token_batches": 1e-06,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 65536,
@ -49507,7 +49516,9 @@
"output_cost_per_image": 0.134,
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_priority": 2.16e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"supports_reasoning": false,
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
@ -68612,7 +68623,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": {
@ -69473,6 +69486,7 @@
"input_cost_per_audio_token": 3e-06,
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -69493,6 +69507,7 @@
"input_cost_per_audio_token": 3e-06,
"input_cost_per_image_token": 1e-06,
"input_cost_per_token": 7.5e-07,
"input_cost_per_video_token": 1e-06,
"litellm_provider": "gemini",
"max_input_tokens": 131072,
"max_output_tokens": 65536,
@ -69774,6 +69789,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,
@ -69805,6 +69821,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",
@ -69822,6 +69839,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,
@ -69843,6 +69861,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",
@ -69861,6 +69880,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

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

View file

@ -16,6 +16,7 @@ longer signal it.
### Added
- **model**: Optional `display_name` argument on `litellm_model`, sent as `model_info.display_name` and returned as `display_name` by `/v1/models`, so client model pickers show a readable name; changes are persisted through `/model/{id}/update` since `/model/update` ignores `model_info`; also exported by the `litellm_model` and `litellm_models` data sources
- **key**: Computed `server_metadata` attribute on `litellm_key` exposing every metadata entry the proxy stores, so metadata created outside Terraform is visible in state and drift on it shows on refresh, while `metadata` keeps tracking only the declared entries and updates keep preserving undeclared ones
- **team_member_add**: `tpm_limit`, `rpm_limit`, `budget_duration`, and `allowed_models` attributes on `litellm_team_member_add`, applied to every member of the resource; `budget_duration` and `allowed_models` ride on `/team/member_add`, while the limits are sent through `/team/member_update`, which is where the proxy accepts them
- **team**: Optional `team_id` argument on `litellm_team`, so teams can be created with a stable, human-readable ID instead of a provider-generated UUID; changing it forces replacement
@ -48,6 +49,7 @@ longer signal it.
### Fixed
- **model**: `litellm_model` refresh now reads the `{"data": [...]}` envelope `/model/info` returns, so `model_info` fields changed outside Terraform show up as drift instead of silently keeping the previous state
- **key**: An update that changes `team_id` and fails because the key was already cascade-deleted along with its previous team now recovers by recreating the key under the new team, instead of aborting the apply. The key's absence is confirmed against the proxy first, so an unrelated failure still errors out, and a `team_id` change between two teams that both still exist stays a plain in-place update
- **credential**: create now reports a `credential_name` collision as a clear error naming the `terraform import` command that adopts the existing credential, instead of surfacing the proxy's raw 500 with a Prisma `Unique constraint failed` message. New `adopt_existing` argument (default `false`) opts into taking the existing credential over during create, which makes `apply` idempotent again once state loses track of a credential that still exists on the proxy. Requires a proxy that answers 409 on the collision; older proxies are still detected by their 500 message
- **credential**: credential names and `model_id` are now percent-encoded in request URLs, so a name containing `/`, `?`, `#` or spaces reaches the proxy intact instead of being cut at the first reserved character and read, updated or deleted as a different credential

View file

@ -43,6 +43,7 @@ In addition to all arguments above, the following attributes are exported:
* `tier` - Model tier (`free` or `paid`).
* `mode` - Model mode, e.g. `chat` or `embedding`.
* `team_id` - Team the deployment is scoped to, if any.
* `display_name` - Human-readable name returned by `/v1/models`, if configured.
* `db_model` - Whether the deployment is stored in the database (as opposed to config).
## Security Note

View file

@ -41,4 +41,5 @@ In addition to all arguments above, the following attributes are exported:
* `tier` - Model tier (`free` or `paid`).
* `mode` - Model mode, e.g. `chat` or `embedding`.
* `team_id` - Team the deployment is scoped to, if any.
* `display_name` - Human-readable name returned by `/v1/models`, if configured.
* `db_model` - Whether the deployment is stored in the database.

View file

@ -126,6 +126,8 @@ The following arguments are supported:
* `team_id` - (Optional) string. Associate the model with a specific team.
* `display_name` - (Optional) string. Human-readable name stored in `model_info.display_name` and returned as `display_name` by `/v1/models`, so clients such as Claude Code and Claude Desktop show it in their model picker instead of `model_name`. When unset, clients fall back to `model_name`.
* `mode` - (Optional) string. The intended use of the model. Valid values are:
* `completion`
* `embedding`

View file

@ -23,14 +23,15 @@ type modelInfoParams struct {
}
type modelInfoMeta struct {
ID string `json:"id"`
DBModel bool `json:"db_model"`
BaseModel string `json:"base_model"`
Tier string `json:"tier"`
Mode string `json:"mode"`
TeamID string `json:"team_id"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
ID string `json:"id"`
DBModel bool `json:"db_model"`
BaseModel string `json:"base_model"`
Tier string `json:"tier"`
Mode string `json:"mode"`
TeamID string `json:"team_id"`
DisplayName string `json:"display_name"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
}
type modelInfoEntry struct {
@ -112,6 +113,10 @@ func dataSourceLiteLLMModel() *schema.Resource {
Type: schema.TypeString,
Computed: true,
},
"display_name": {
Type: schema.TypeString,
Computed: true,
},
"db_model": {
Type: schema.TypeBool,
Computed: true,
@ -161,6 +166,7 @@ func dataSourceLiteLLMModelRead(d *schema.ResourceData, m interface{}) error {
d.Set("tier", entry.ModelInfo.Tier)
d.Set("mode", entry.ModelInfo.Mode)
d.Set("team_id", entry.ModelInfo.TeamID)
d.Set("display_name", entry.ModelInfo.DisplayName)
d.Set("db_model", entry.ModelInfo.DBModel)
log.Printf("[INFO] Successfully read model with ID: %s", modelID)
@ -197,6 +203,7 @@ func dataSourceLiteLLMModels() *schema.Resource {
"tier": {Type: schema.TypeString, Computed: true},
"mode": {Type: schema.TypeString, Computed: true},
"team_id": {Type: schema.TypeString, Computed: true},
"display_name": {Type: schema.TypeString, Computed: true},
"db_model": {Type: schema.TypeBool, Computed: true},
},
},
@ -247,6 +254,7 @@ func dataSourceLiteLLMModelsRead(d *schema.ResourceData, m interface{}) error {
"tier": entry.ModelInfo.Tier,
"mode": entry.ModelInfo.Mode,
"team_id": entry.ModelInfo.TeamID,
"display_name": entry.ModelInfo.DisplayName,
"db_model": entry.ModelInfo.DBModel,
})
}

View file

@ -35,7 +35,8 @@ func TestDataSourceModelReadSingleObject(t *testing.T) {
"base_model": "gpt-4o",
"tier": "paid",
"mode": "chat",
"team_id": "team-1"
"team_id": "team-1",
"display_name": "GPT-4o"
}
}
}`))
@ -66,6 +67,7 @@ func TestDataSourceModelReadSingleObject(t *testing.T) {
"tier": "paid",
"mode": "chat",
"team_id": "team-1",
"display_name": "GPT-4o",
"db_model": true,
}
for attr, want := range checks {
@ -115,7 +117,7 @@ func TestDataSourceModelsRead(t *testing.T) {
w.Header().Set("Content-Type", "application/json")
w.Write([]byte(`{
"data": [
{"model_name": "a", "litellm_params": {"model": "openai/a", "custom_llm_provider": "openai"}, "model_info": {"id": "id-1", "db_model": true}},
{"model_name": "a", "litellm_params": {"model": "openai/a", "custom_llm_provider": "openai"}, "model_info": {"id": "id-1", "db_model": true, "display_name": "Model A"}},
{"model_name": "b", "litellm_params": {"model": "anthropic/b", "custom_llm_provider": "anthropic"}, "model_info": {"id": "id-2"}}
]
}`))
@ -143,7 +145,11 @@ func TestDataSourceModelsRead(t *testing.T) {
t.Fatalf("expected 2 models, got %d", len(models))
}
first := models[0].(map[string]interface{})
if first["model_name"] != "a" || first["custom_llm_provider"] != "openai" || first["db_model"] != true {
if first["model_name"] != "a" || first["custom_llm_provider"] != "openai" || first["db_model"] != true || first["display_name"] != "Model A" {
t.Errorf("unexpected first model: %v", first)
}
second := models[1].(map[string]interface{})
if second["display_name"] != "" {
t.Errorf("expected empty display_name for model without one, got %v", second["display_name"])
}
}

View file

@ -93,6 +93,11 @@ func resourceLiteLLMModel() *schema.Resource {
Type: schema.TypeString,
Optional: true,
},
"display_name": {
Type: schema.TypeString,
Optional: true,
Description: "Human-readable name returned as display_name by /v1/models, shown in client model pickers instead of model_name",
},
"mode": {
Type: schema.TypeString,
Optional: true,

View file

@ -4,6 +4,7 @@ import (
"encoding/json"
"fmt"
"log"
"net/url"
"strconv"
"strings"
"time"
@ -53,6 +54,7 @@ func retryModelRead(d *schema.ResourceData, m interface{}, maxRetries int) error
const (
endpointModelNew = "/model/new"
endpointModelUpdate = "/model/update"
endpointModelPatch = "/model/%s/update"
endpointModelInfo = "/model/info"
endpointModelDelete = "/model/delete"
)
@ -246,12 +248,13 @@ func createOrUpdateModel(d *schema.ResourceData, m interface{}, isUpdate bool) e
ModelName: d.Get("model_name").(string),
LiteLLMParams: litellmParams,
ModelInfo: ModelInfo{
ID: modelID,
DBModel: true,
BaseModel: pricingBaseModel,
Tier: d.Get("tier").(string),
Mode: d.Get("mode").(string),
TeamID: d.Get("team_id").(string),
ID: modelID,
DBModel: true,
BaseModel: pricingBaseModel,
Tier: d.Get("tier").(string),
Mode: d.Get("mode").(string),
TeamID: d.Get("team_id").(string),
DisplayName: d.Get("display_name").(string),
},
Additional: make(map[string]interface{}),
}
@ -275,6 +278,12 @@ func createOrUpdateModel(d *schema.ResourceData, m interface{}, isUpdate bool) e
return fmt.Errorf("failed to %s model: %w", map[bool]string{true: "update", false: "create"}[isUpdate], err)
}
if isUpdate && d.HasChange("display_name") {
if err := patchModelDisplayName(client, modelID, d.Get("display_name").(string)); err != nil {
return fmt.Errorf("failed to update model display_name: %w", err)
}
}
d.SetId(modelID)
log.Printf("[INFO] Model created with ID %s. Starting retry mechanism to read the model...", modelID)
@ -282,6 +291,19 @@ func createOrUpdateModel(d *schema.ResourceData, m interface{}, isUpdate bool) e
return retryModelRead(d, m, 5)
}
// /model/update only merges litellm_params, so model_info changes go through the PATCH endpoint.
func patchModelDisplayName(client *Client, modelID, displayName string) error {
resp, err := MakeRequest(client, "PATCH", fmt.Sprintf(endpointModelPatch, url.PathEscape(modelID)), ModelInfoPatch{
ModelInfo: ModelInfoPatchFields{ID: modelID, DisplayName: displayName},
})
if err != nil {
return err
}
defer resp.Body.Close()
_, err = handleAPIResponse(resp, nil, client)
return err
}
func resourceLiteLLMModelCreate(d *schema.ResourceData, m interface{}) error {
return createOrUpdateModel(d, m, false)
}
@ -327,6 +349,7 @@ func resourceLiteLLMModelRead(d *schema.ResourceData, m interface{}) error {
d.Set("tier", GetStringValue(modelResp.ModelInfo.Tier, d.Get("tier").(string)))
d.Set("mode", GetStringValue(modelResp.ModelInfo.Mode, d.Get("mode").(string)))
d.Set("team_id", GetStringValue(modelResp.ModelInfo.TeamID, d.Get("team_id").(string)))
d.Set("display_name", modelResp.ModelInfo.DisplayName)
// Preserve credential name from state since it might not be returned by API
d.Set("litellm_credential_name", d.Get("litellm_credential_name").(string))

View file

@ -0,0 +1,244 @@
package litellm
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema"
"github.com/hashicorp/terraform-plugin-sdk/v2/terraform"
)
func modelInfoBody(displayName string) string {
modelInfo := map[string]interface{}{
"id": "model-123",
"db_model": true,
"base_model": "claude-sonnet-4-5",
"tier": "free",
"mode": "chat",
}
if displayName != "" {
modelInfo["display_name"] = displayName
}
body, _ := json.Marshal(map[string]interface{}{
"model_name": "sonnet-4-5-anthropic",
"litellm_params": map[string]interface{}{"model": "anthropic/claude-sonnet-4-5", "custom_llm_provider": "anthropic"},
"model_info": modelInfo,
})
return string(body)
}
func modelInfoDataEnvelope(displayName string) string {
return `{"data": [` + modelInfoBody(displayName) + `]}`
}
func TestResourceLiteLLMModelCreateSendsDisplayName(t *testing.T) {
var createPayload map[string]interface{}
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/model/new":
if err := json.NewDecoder(r.Body).Decode(&createPayload); err != nil {
t.Errorf("failed to decode create payload: %v", err)
}
w.Write([]byte(modelInfoBody("Claude Sonnet 4.5")))
case "/model/info":
w.Write([]byte(modelInfoBody("Claude Sonnet 4.5")))
default:
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
w.WriteHeader(http.StatusNotFound)
}
}))
defer srv.Close()
d := schema.TestResourceDataRaw(t, resourceLiteLLMModel().Schema, map[string]interface{}{
"model_name": "sonnet-4-5-anthropic",
"custom_llm_provider": "anthropic",
"base_model": "claude-sonnet-4-5",
"model_api_key": "sk-ant-test",
"mode": "chat",
"display_name": "Claude Sonnet 4.5",
})
if err := resourceLiteLLMModelCreate(d, NewClient(srv.URL, "test-key", true)); err != nil {
t.Fatalf("create failed: %v", err)
}
modelInfo, ok := createPayload["model_info"].(map[string]interface{})
if !ok {
t.Fatalf("expected model_info object in create payload, got %v", createPayload["model_info"])
}
if modelInfo["display_name"] != "Claude Sonnet 4.5" {
t.Errorf("expected model_info.display_name 'Claude Sonnet 4.5', got %v", modelInfo["display_name"])
}
if got := d.Get("display_name").(string); got != "Claude Sonnet 4.5" {
t.Errorf("expected state display_name 'Claude Sonnet 4.5', got %q", got)
}
}
func TestResourceLiteLLMModelCreateOmitsUnsetDisplayName(t *testing.T) {
var createPayload map[string]interface{}
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/model/new":
if err := json.NewDecoder(r.Body).Decode(&createPayload); err != nil {
t.Errorf("failed to decode create payload: %v", err)
}
w.Write([]byte(modelInfoBody("")))
case "/model/info":
w.Write([]byte(modelInfoBody("")))
default:
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
w.WriteHeader(http.StatusNotFound)
}
}))
defer srv.Close()
d := schema.TestResourceDataRaw(t, resourceLiteLLMModel().Schema, map[string]interface{}{
"model_name": "sonnet-4-5-anthropic",
"custom_llm_provider": "anthropic",
"base_model": "claude-sonnet-4-5",
"model_api_key": "sk-ant-test",
})
if err := resourceLiteLLMModelCreate(d, NewClient(srv.URL, "test-key", true)); err != nil {
t.Fatalf("create failed: %v", err)
}
modelInfo := createPayload["model_info"].(map[string]interface{})
if _, present := modelInfo["display_name"]; present {
t.Errorf("expected display_name to be omitted from model_info when unset, got %v", modelInfo["display_name"])
}
if got := d.Get("display_name").(string); got != "" {
t.Errorf("expected empty state display_name, got %q", got)
}
}
func TestResourceLiteLLMModelReadDisplayName(t *testing.T) {
cases := map[string]struct {
serverBody string
want string
}{
"server value wins inside data envelope": {serverBody: modelInfoDataEnvelope("Renamed In Admin UI"), want: "Renamed In Admin UI"},
"server value wins unwrapped": {serverBody: modelInfoBody("Renamed In Admin UI"), want: "Renamed In Admin UI"},
"external removal clears state": {serverBody: modelInfoDataEnvelope(""), want: ""},
}
for name, tc := range cases {
t.Run(name, func(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/model/info" {
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
}
w.Write([]byte(tc.serverBody))
}))
defer srv.Close()
d := schema.TestResourceDataRaw(t, resourceLiteLLMModel().Schema, map[string]interface{}{
"model_name": "sonnet-4-5-anthropic",
"custom_llm_provider": "anthropic",
"base_model": "claude-sonnet-4-5",
"display_name": "Claude Sonnet 4.5",
})
d.SetId("model-123")
if err := resourceLiteLLMModelRead(d, NewClient(srv.URL, "test-key", true)); err != nil {
t.Fatalf("read failed: %v", err)
}
if got := d.Get("display_name").(string); got != tc.want {
t.Errorf("expected display_name %q, got %q", tc.want, got)
}
})
}
}
func updateResourceData(t *testing.T, oldDisplayName, newDisplayName string) *schema.ResourceData {
t.Helper()
res := resourceLiteLLMModel()
attrs := map[string]string{
"model_name": "sonnet-4-5-anthropic",
"custom_llm_provider": "anthropic",
"base_model": "claude-sonnet-4-5",
}
if oldDisplayName != "" {
attrs["display_name"] = oldDisplayName
}
state := &terraform.InstanceState{ID: "model-123", Attributes: attrs}
diff, err := res.Diff(context.Background(), state, &terraform.ResourceConfig{Config: map[string]interface{}{
"model_name": "sonnet-4-5-anthropic",
"custom_llm_provider": "anthropic",
"base_model": "claude-sonnet-4-5",
"display_name": newDisplayName,
}}, nil)
if err != nil {
t.Fatalf("diff failed: %v", err)
}
d, err := schema.InternalMap(res.Schema).Data(state, diff)
if err != nil {
t.Fatalf("data failed: %v", err)
}
return d
}
func TestResourceLiteLLMModelUpdatePatchesDisplayName(t *testing.T) {
cases := map[string]struct {
newName string
}{
"changed name is patched": {newName: "Claude Sonnet 4.5 v2"},
"cleared name is patched": {newName: ""},
}
for name, tc := range cases {
t.Run(name, func(t *testing.T) {
var patchPayload map[string]interface{}
var patchPath string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch {
case r.Method == http.MethodPost && r.URL.Path == "/model/update":
w.Write([]byte(modelInfoBody("Claude Sonnet 4.5")))
case r.Method == http.MethodPatch:
patchPath = r.URL.Path
if err := json.NewDecoder(r.Body).Decode(&patchPayload); err != nil {
t.Errorf("failed to decode patch payload: %v", err)
}
w.Write([]byte(modelInfoBody(tc.newName)))
case r.URL.Path == "/model/info":
w.Write([]byte(modelInfoDataEnvelope(tc.newName)))
default:
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
w.WriteHeader(http.StatusNotFound)
}
}))
defer srv.Close()
d := updateResourceData(t, "Claude Sonnet 4.5", tc.newName)
if err := resourceLiteLLMModelUpdate(d, NewClient(srv.URL, "test-key", true)); err != nil {
t.Fatalf("update failed: %v", err)
}
if patchPath != "/model/model-123/update" {
t.Fatalf("expected PATCH /model/model-123/update, got %q", patchPath)
}
modelInfo := patchPayload["model_info"].(map[string]interface{})
if modelInfo["display_name"] != tc.newName {
t.Errorf("expected patched display_name %q, got %v", tc.newName, modelInfo["display_name"])
}
if got := d.Get("display_name").(string); got != tc.newName {
t.Errorf("expected state display_name %q, got %q", tc.newName, got)
}
})
}
}
func TestResourceLiteLLMModelUpdateSkipsPatchWhenDisplayNameUnchanged(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method == http.MethodPatch {
t.Errorf("unexpected PATCH %s", r.URL.Path)
}
w.Write([]byte(modelInfoDataEnvelope("Claude Sonnet 4.5")))
}))
defer srv.Close()
d := updateResourceData(t, "Claude Sonnet 4.5", "Claude Sonnet 4.5")
if err := resourceLiteLLMModelUpdate(d, NewClient(srv.URL, "test-key", true)); err != nil {
t.Fatalf("update failed: %v", err)
}
}

View file

@ -25,6 +25,16 @@ type ModelResponse struct {
Additional map[string]interface{} `json:"additional"`
}
// ModelInfoPatch is the body for PATCH /model/{id}/update; display_name is sent even when empty so it can be cleared.
type ModelInfoPatch struct {
ModelInfo ModelInfoPatchFields `json:"model_info"`
}
type ModelInfoPatchFields struct {
ID string `json:"id"`
DisplayName string `json:"display_name"`
}
// ModelRequest represents a request to create or update a model.
type ModelRequest struct {
ModelName string `json:"model_name"`
@ -108,12 +118,13 @@ type LiteLLMParams struct {
// ModelInfo represents information about a model.
type ModelInfo struct {
ID string `json:"id"`
DBModel bool `json:"db_model"`
BaseModel string `json:"base_model"`
Tier string `json:"tier"`
Mode string `json:"mode"`
TeamID string `json:"team_id,omitempty"`
ID string `json:"id"`
DBModel bool `json:"db_model"`
BaseModel string `json:"base_model"`
Tier string `json:"tier"`
Mode string `json:"mode"`
TeamID string `json:"team_id,omitempty"`
DisplayName string `json:"display_name,omitempty"`
}
// Key represents a LiteLLM API key.

View file

@ -55,6 +55,13 @@ func handleAPIResponse(resp *http.Response, reqBody interface{}, client *Client)
resp.Status, client.redactSensitiveData(string(bodyBytes)), client.redactSensitiveData(string(reqBodyBytes)))
}
var envelope struct {
Data []json.RawMessage `json:"data"`
}
if err := json.Unmarshal(bodyBytes, &envelope); err == nil && len(envelope.Data) > 0 {
bodyBytes = envelope.Data[0]
}
var modelResp ModelResponse
if err := json.Unmarshal(bodyBytes, &modelResp); err != nil {
return nil, fmt.Errorf("failed to parse response: %v", err)

View file

@ -28,6 +28,7 @@ GET /tag/user-agent/per-user-analytics
GET /tag/wau
GET /team/daily/activity
GET /team/daily/activity/aggregated
GET /team/daily/activity/export
GET /team/daily/activity/aggregated/search
GET /team/spend/by_user
GET /team/spend/report
@ -96,7 +97,6 @@ GET /guardrails/{guardrail_id}
GET /prompts/{prompt_id}
GET /prompts/{prompt_id}/versions
PATCH /guardrails/{guardrail_id}
PATCH /model/{model_id}/update
PATCH /prompts/{prompt_id}
PATCH /team/{team_id}
POST /team/model/add

View file

@ -31,11 +31,11 @@ def get_function_names_from_file(file_path):
def get_all_functions_called_in_tests(base_dir):
"""
Returns a set of function names that are called in test functions
inside 'local_testing' and 'proxy_unit_tests' directories,
inside 'local_testing' and 'unit/proxy' directories,
specifically in files containing the word 'router'.
"""
called_functions = set()
test_dirs = ["local_testing", "proxy_unit_tests"]
test_dirs = ["local_testing", "unit/proxy"]
for test_dir in test_dirs:
dir_path = os.path.join(base_dir, test_dir)

View file

@ -0,0 +1,60 @@
"""Unit tests for the Claude Code PR-gate version resolver.
Markerless harness tests: they feed the resolver a hand-built packument and a
fixed clock, so they run without a proxy, never reach the npm registry, and
carry no `e2e` marker.
"""
from __future__ import annotations
from datetime import datetime, timezone
from typing import Final, Mapping
import pytest
from claude_code.pr_gate_version_resolver import NoEligibleVersionError, resolve_pr_gate_version
NOW: Final = datetime(2026, 4, 25, 12, 0, tzinfo=timezone.utc)
INSIDE_THE_2_1_88_WINDOW: Final = datetime(2026, 4, 3, 12, 0, tzinfo=timezone.utc)
def _packument(times: Mapping[str, str], unpublished: frozenset[str] = frozenset()) -> dict[str, object]:
return {
"name": "@anthropic-ai/claude-code",
"time": {"created": "2024-01-01T00:00:00.000Z", "modified": "2026-04-25T00:00:00.000Z", **times},
"versions": {version: {"version": version} for version in times if version not in unpublished},
}
def test_skips_a_version_npm_has_unpublished() -> None:
metadata: Final = _packument(
{
"2.1.87": "2026-03-28T20:00:00.000Z",
"2.1.88": "2026-03-30T22:36:48.424Z",
"2.1.89": "2026-03-31T23:32:40.000Z",
},
unpublished=frozenset({"2.1.88"}),
)
assert resolve_pr_gate_version(metadata=metadata, as_of=INSIDE_THE_2_1_88_WINDOW) == "2.1.87"
def test_raises_when_the_only_old_enough_version_is_unpublished() -> None:
metadata: Final = _packument(
{"2.1.88": "2026-03-30T22:36:48.424Z", "2.1.89": "2026-03-31T23:32:40.000Z"},
unpublished=frozenset({"2.1.88"}),
)
with pytest.raises(NoEligibleVersionError):
resolve_pr_gate_version(metadata=metadata, as_of=INSIDE_THE_2_1_88_WINDOW)
def test_picks_the_newest_published_version_at_least_min_age_old() -> None:
metadata: Final = _packument(
{
"2.1.118": "2026-04-15T10:00:00.000Z",
"2.1.119": "2026-04-21T10:00:00.000Z",
"2.2.0-rc.1": "2026-04-22T10:00:00.000Z",
"2.1.120": "2026-04-23T10:00:00.000Z",
"2.1.121": "2026-04-25T11:00:00.000Z",
}
)
assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.1.119"

View file

@ -297,6 +297,7 @@ def run_claude(
}
env["ANTHROPIC_BASE_URL"] = base_url
env["ANTHROPIC_AUTH_TOKEN"] = api_key
env["DISABLE_AUTOUPDATER"] = "1"
# Hand the CLI a fresh empty HOME so a compromised claude package
# or a model-directed Read tool call can't see the runtime user's
# real dotfiles. Created here, removed in the `finally` below

View file

@ -4,8 +4,6 @@ ARG GH_VERSION=2.101.0
ARG GH_SHA256=9bca2d1c16825f109907a23307628a2f0698fbf99662b73a5cf0b020293072b8
ARG UV_VERSION=0.10.9
ARG UV_SHA256=20d79708222611fa540b5c9ed84f352bcd3937740e51aacc0f8b15b271c57594
ARG CLAUDE_CODE_VERSION=2.1.228
ARG CLAUDE_CODE_SHA256=d535985e6941a3eb00179ccd7f52ceb0c6623a0305a518ebc4e6514f84a94c99
SHELL ["/bin/bash", "-o", "pipefail", "-c"]
@ -23,11 +21,6 @@ RUN curl -fsSLo /tmp/uv.tar.gz "https://github.com/astral-sh/uv/releases/downloa
&& tar -xzf /tmp/uv.tar.gz -C /usr/local/bin --strip-components=1 uv-x86_64-unknown-linux-gnu/uv \
&& rm /tmp/uv.tar.gz
RUN curl -fsSLo /tmp/claude "https://downloads.claude.ai/claude-code-releases/${CLAUDE_CODE_VERSION}/linux-x64/claude" \
&& echo "${CLAUDE_CODE_SHA256} /tmp/claude" | sha256sum -c - \
&& install -m 0755 /tmp/claude /usr/local/bin/claude \
&& rm /tmp/claude
RUN groupadd --gid 1000 populator && useradd --uid 1000 --gid 1000 --create-home populator
ENV HOME=/home/populator \

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