mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Merge remote-tracking branch 'origin/main' into litellm_psycopg_no_binary_wheel
This commit is contained in:
commit
185a56ef65
443 changed files with 25977 additions and 8539 deletions
|
|
@ -31,7 +31,7 @@ while IFS= read -r file || [ -n "$file" ]; do
|
|||
case "$file" in
|
||||
model_prices_and_context_window.json | litellm/model_prices_and_context_window_backup.json | model_prices_and_context_window.schema.json)
|
||||
has_cost_map=true ;;
|
||||
tests/test_litellm/* | tests/proxy_unit_tests/*) : ;;
|
||||
tests/test_litellm/* | tests/proxy_unit_tests/* | tests/unit/proxy/*) : ;;
|
||||
*) outside_cost_map_set=true ;;
|
||||
esac
|
||||
done
|
||||
|
|
|
|||
140
.circleci/scripts/unit_selection.sh
Executable file
140
.circleci/scripts/unit_selection.sh
Executable file
|
|
@ -0,0 +1,140 @@
|
|||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
flag="${1:?usage: unit_selection.sh <codecov flag>}"
|
||||
|
||||
legacy_flags=(
|
||||
caching-local
|
||||
enterprise-package
|
||||
enterprise-routing
|
||||
mcp-integration
|
||||
proxy-db-auth-checks
|
||||
proxy-db-budgets
|
||||
proxy-db-custom-logging
|
||||
proxy-db-db-and-spend
|
||||
proxy-db-endpoints-and-responses
|
||||
proxy-db-guardrails-hooks
|
||||
proxy-db-jwt-and-keys
|
||||
proxy-db-key-generation
|
||||
proxy-db-logging-misc
|
||||
proxy-db-proxy-runtime
|
||||
proxy-db-proxy-server-core
|
||||
proxy-db-proxy-utils
|
||||
proxy-extras
|
||||
proxy-infra
|
||||
)
|
||||
|
||||
legacy_paths() {
|
||||
case "$1" in
|
||||
caching-local) echo tests/unit/caching ;;
|
||||
enterprise-package)
|
||||
echo tests/unit/enterprise/integrations
|
||||
echo tests/unit/enterprise/proxy/auth
|
||||
echo tests/unit/enterprise/proxy/guardrails
|
||||
echo tests/unit/enterprise/proxy/hooks
|
||||
echo tests/unit/enterprise/proxy/management_endpoints
|
||||
echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
|
||||
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
|
||||
enterprise-routing)
|
||||
echo tests/unit/enterprise/enterprise_callbacks/send_emails
|
||||
echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py
|
||||
echo tests/unit/enterprise/proxy/test_enterprise_routes.py
|
||||
echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_access_check.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;;
|
||||
mcp-integration)
|
||||
echo tests/unit/proxy/_experimental/mcp_server
|
||||
echo tests/unit/responses/mcp
|
||||
echo tests/mcp_tests/test_proxy_mcp_e2e.py ;;
|
||||
proxy-db-auth-checks)
|
||||
echo tests/unit/proxy/auth/test_auth_checks.py
|
||||
echo tests/unit/proxy/auth/test_user_api_key_auth.py
|
||||
echo tests/unit/proxy/test_deprecated_key_grace_period.py ;;
|
||||
proxy-db-budgets)
|
||||
echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py
|
||||
echo tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py
|
||||
echo tests/unit/proxy/test_zero_cost_model_budget_bypass.py ;;
|
||||
proxy-db-custom-logging)
|
||||
echo tests/unit/proxy/test_custom_callback_input.py
|
||||
echo tests/unit/proxy/test_custom_logger_s3_gcs.py ;;
|
||||
proxy-db-db-and-spend)
|
||||
echo tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py
|
||||
echo tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py
|
||||
echo tests/unit/proxy/db/test_update_daily_tag_spend.py
|
||||
echo tests/unit/proxy/test_db_schema_changes.py
|
||||
echo tests/unit/proxy/test_prisma_client_backoff_retry.py
|
||||
echo tests/unit/proxy/test_update_spend.py
|
||||
echo tests/unit/skills/test_skills_db.py ;;
|
||||
proxy-db-endpoints-and-responses)
|
||||
echo tests/unit/proxy/auth/test_models_fallback_endpoint.py
|
||||
echo tests/unit/proxy/common_utils/test_check_batch_cost.py
|
||||
echo tests/unit/proxy/common_utils/test_check_responses_cost.py
|
||||
echo tests/unit/proxy/common_utils/test_realtime_cache.py
|
||||
echo tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py
|
||||
echo tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py
|
||||
echo tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py
|
||||
echo tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py
|
||||
echo tests/unit/proxy/response_polling/test_response_polling_handler.py
|
||||
echo tests/unit/proxy/test_custom_tokenizer_bug.py
|
||||
echo tests/unit/proxy/test_get_favicon.py
|
||||
echo tests/unit/proxy/test_get_image.py
|
||||
echo tests/unit/proxy/test_prompt_test_endpoint.py
|
||||
echo tests/unit/proxy/test_reducto_ocr_route.py
|
||||
echo tests/unit/proxy/test_response_polling_pre_call_checks.py
|
||||
echo tests/unit/proxy/test_ui_path_detection.py ;;
|
||||
proxy-db-guardrails-hooks)
|
||||
echo tests/unit/proxy/hooks/test_banned_keyword_list.py
|
||||
echo tests/unit/proxy/test_proxy_setting_guardrails.py
|
||||
echo tests/unit/proxy/test_unit_test_proxy_hooks.py ;;
|
||||
proxy-db-jwt-and-keys)
|
||||
echo tests/unit/proxy/auth/test_jwt.py
|
||||
echo tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py
|
||||
echo tests/unit/proxy/test_proxy_custom_auth.py ;;
|
||||
proxy-db-key-generation) echo tests/unit/proxy/management_endpoints/test_key_generate_prisma.py ;;
|
||||
proxy-db-logging-misc)
|
||||
echo tests/unit/proxy/management_helpers/test_audit_logs_proxy.py
|
||||
echo tests/unit/proxy/spend_tracking/test_search_api_logging.py
|
||||
echo tests/unit/proxy/test_proxy_reject_logging.py ;;
|
||||
proxy-db-proxy-runtime)
|
||||
echo tests/unit/proxy/auth/test_multipart_bypass_repro.py
|
||||
echo tests/unit/proxy/auth/test_proxy_routes.py
|
||||
echo tests/unit/proxy/middleware/test_request_size_limit_middleware.py
|
||||
echo tests/unit/proxy/test_proxy_config_unit_test.py
|
||||
echo tests/unit/proxy/test_proxy_token_counter.py
|
||||
echo tests/unit/proxy/test_server_root_path.py ;;
|
||||
proxy-db-proxy-server-core)
|
||||
echo tests/unit/proxy/test_aproxy_startup.py
|
||||
echo tests/unit/proxy/test_proxy_server.py ;;
|
||||
proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;;
|
||||
proxy-extras) echo tests/unit/litellm_proxy_extras ;;
|
||||
proxy-infra) echo tests/unit/gateway ;;
|
||||
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
|
||||
esac
|
||||
}
|
||||
|
||||
expand() {
|
||||
while read -r path; do
|
||||
if [ -d "$path" ]; then
|
||||
find "$path" -name 'test_*.py'
|
||||
elif [ -f "$path" ]; then
|
||||
echo "$path"
|
||||
else
|
||||
echo "unit_selection.sh: $path does not exist" >&2
|
||||
exit 1
|
||||
fi
|
||||
done
|
||||
}
|
||||
|
||||
if [ "$flag" = unit ]; then
|
||||
comm -23 \
|
||||
<(find tests/unit -name 'test_*.py' | sort) \
|
||||
<(for legacy in "${legacy_flags[@]}"; do legacy_paths "$legacy"; done | expand | sort)
|
||||
exit 0
|
||||
fi
|
||||
|
||||
legacy_paths "$flag" | expand | sort
|
||||
|
|
@ -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 >>
|
||||
|
|
|
|||
6
.github/pull_request_template.md
vendored
6
.github/pull_request_template.md
vendored
|
|
@ -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:
|
||||
|
||||
|
|
|
|||
13
.github/scripts/assert_ci_coverage.py
vendored
13
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -34,7 +34,6 @@ GLOB_CHARS = frozenset("*?")
|
|||
# tests has to be named by some shard or it runs nowhere. A child listed here is
|
||||
# itself decomposed one level deeper and is checked through its own entry.
|
||||
SHARDED_ROOTS: tuple[str, ...] = (
|
||||
"tests/proxy_unit_tests",
|
||||
"tests/test_litellm",
|
||||
"tests/test_litellm/proxy",
|
||||
)
|
||||
|
|
@ -120,6 +119,13 @@ def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
|||
)
|
||||
|
||||
|
||||
def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]:
|
||||
script: Final = repo_root / ".circleci/scripts/unit_selection.sh"
|
||||
if not script.is_file():
|
||||
return frozenset()
|
||||
return frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(_uncommented(script.read_text())))
|
||||
|
||||
|
||||
def _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
||||
return frozenset(
|
||||
match.group(0)
|
||||
|
|
@ -611,7 +617,10 @@ def main() -> int:
|
|||
scalars = _all_scalars()
|
||||
|
||||
integration_paths, ownership_findings = _integration_ownership()
|
||||
test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars) | integration_paths) + ownership_findings
|
||||
test_findings = (
|
||||
_uncovered_tests(allowlist, _invoked_test_tokens(scalars) | _unit_selection_tokens() | integration_paths)
|
||||
+ ownership_findings
|
||||
)
|
||||
dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars))
|
||||
stale_findings = _stale_allowlist_paths(allowlist, test_files=_test_files(), dockerfiles=_dockerfiles())
|
||||
|
||||
|
|
|
|||
33
.github/workflows/_test-unit-base.yml
vendored
33
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -13,6 +13,15 @@ on:
|
|||
have its path existence-checked like any other token.
|
||||
required: true
|
||||
type: string
|
||||
fork-flag:
|
||||
description: >-
|
||||
Codecov flag of the `.circleci/tests.yml` job that now owns part of
|
||||
this shard. CircleCI does not run on pull requests from forks, so on
|
||||
those events this shard also runs the files
|
||||
`.circleci/scripts/unit_selection.sh` lists for the flag.
|
||||
required: false
|
||||
type: string
|
||||
default: ""
|
||||
workers:
|
||||
description: "Number of pytest-xdist workers"
|
||||
required: false
|
||||
|
|
@ -92,6 +101,7 @@ jobs:
|
|||
pull-requests: read
|
||||
outputs:
|
||||
decision: ${{ steps.changes.outputs.decision }}
|
||||
has-coverage: ${{ steps.tests.outputs.has-coverage }}
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
|
|
@ -160,10 +170,13 @@ jobs:
|
|||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Run tests
|
||||
id: tests
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
timeout-minutes: ${{ inputs.timeout-minutes }}
|
||||
env:
|
||||
TEST_PATH: ${{ inputs.test-path }}
|
||||
FORK_FLAG: ${{ inputs.fork-flag }}
|
||||
IS_FORK: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository }}
|
||||
MAX_FAILURES: ${{ inputs.max-failures }}
|
||||
WORKERS: ${{ inputs.workers }}
|
||||
RERUNS: ${{ inputs.reruns }}
|
||||
|
|
@ -171,9 +184,18 @@ jobs:
|
|||
DIST: ${{ inputs.dist }}
|
||||
COVERAGE_CORE: sysmon
|
||||
run: |
|
||||
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
|
||||
selection="${TEST_PATH}"
|
||||
if [ "${IS_FORK}" = "true" ] && [ -n "${FORK_FLAG}" ]; then
|
||||
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${FORK_FLAG}" | tr '\n' ' ')"
|
||||
fi
|
||||
if [ -z "${selection// /}" ]; then
|
||||
echo "shard selection is empty on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
pytest_args=()
|
||||
existing_paths=0
|
||||
for token in ${TEST_PATH:?}; do
|
||||
for token in ${selection}; do
|
||||
case "${token}" in
|
||||
-*) pytest_args+=("${token}") ;;
|
||||
*)
|
||||
|
|
@ -187,7 +209,7 @@ jobs:
|
|||
esac
|
||||
done
|
||||
if [ "${existing_paths}" -eq 0 ]; then
|
||||
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
|
||||
echo "No path in the selection exists (${selection}); nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
xdist_args=()
|
||||
|
|
@ -209,8 +231,11 @@ jobs:
|
|||
--cov-config=pyproject.toml
|
||||
status=$?
|
||||
set -e
|
||||
if [ -f coverage.xml ]; then
|
||||
echo "has-coverage=true" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
if [ "$status" -eq 5 ]; then
|
||||
echo "pytest collected no tests from ${TEST_PATH}; passing"
|
||||
echo "pytest collected no tests from ${selection}; passing"
|
||||
exit 0
|
||||
fi
|
||||
exit "$status"
|
||||
|
|
@ -226,7 +251,7 @@ jobs:
|
|||
upload-coverage:
|
||||
name: Upload coverage to Codecov
|
||||
needs: run
|
||||
if: always() && needs.run.outputs.decision != 'skip'
|
||||
if: always() && needs.run.outputs.decision != 'skip' && needs.run.outputs.has-coverage == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
13
.github/workflows/compat-matrix-image.yml
vendored
13
.github/workflows/compat-matrix-image.yml
vendored
|
|
@ -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
|
||||
'
|
||||
|
|
|
|||
97
.github/workflows/test-unit-proxy-db.yml
vendored
97
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -20,6 +20,12 @@ concurrency:
|
|||
# rather than alphabetical letter ranges. Adding a new test file means adding it
|
||||
# to whichever group it belongs to, not reshuffling slices.
|
||||
#
|
||||
# `.circleci/tests.yml` runs each group's files on same-repo events under the
|
||||
# `proxy-db-<group>` Codecov flag; `.circleci/scripts/unit_selection.sh` holds
|
||||
# the file lists. CircleCI does not build pull requests from forks, so `fork-flag`
|
||||
# makes the shard run that list there. `test-path` keeps the files that still
|
||||
# reach real providers and never left tests/proxy_unit_tests.
|
||||
#
|
||||
# Design targets:
|
||||
# * Every shard runs in <= 7 minutes of wall-clock on the default runner.
|
||||
# Most of a shard's time is pytest plugin load + xdist worker imports +
|
||||
|
|
@ -58,7 +64,7 @@ jobs:
|
|||
proxy-db:
|
||||
needs: assert-shard-coverage
|
||||
# Display only the semantic shard name in the checks UI instead of GHA's
|
||||
# default "proxy-db (key-generation, tests/proxy_unit_tests/…, 0, loadscope, 20)"
|
||||
# default "proxy-db (key-generation, tests/unit/proxy/…, 0, loadscope, 20)"
|
||||
# which includes every matrix field and gets truncated past the test-path.
|
||||
name: ${{ matrix.test-group }}
|
||||
permissions:
|
||||
|
|
@ -71,132 +77,93 @@ jobs:
|
|||
include:
|
||||
# Must run serially — event-loop conflict with the logging worker.
|
||||
- test-group: key-generation
|
||||
test-path: "tests/proxy_unit_tests/test_key_generate_prisma.py"
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-key-generation
|
||||
workers: 0
|
||||
dist: loadscope
|
||||
timeout: 20
|
||||
|
||||
# ---- auth: split into 2 shards ----
|
||||
- test-group: auth-checks
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_auth_checks.py
|
||||
tests/proxy_unit_tests/test_user_api_key_auth.py
|
||||
tests/proxy_unit_tests/test_deprecated_key_grace_period.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-auth-checks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: jwt-and-keys
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_jwt.py
|
||||
tests/proxy_unit_tests/test_jwt_key_mapping.py
|
||||
tests/proxy_unit_tests/test_proxy_custom_auth.py
|
||||
tests/proxy_unit_tests/test_key_generate_dynamodb.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-jwt-and-keys
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
|
||||
- test-group: proxy-utils
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_utils.py"
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-utils
|
||||
workers: 4
|
||||
dist: worksteal
|
||||
timeout: 15
|
||||
|
||||
# ---- proxy server: split into 2 shards ----
|
||||
- test-group: proxy-server-core
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_proxy_server.py
|
||||
tests/proxy_unit_tests/test_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 }}
|
||||
|
|
|
|||
32
.github/workflows/test-unit.yml
vendored
32
.github/workflows/test-unit.yml
vendored
|
|
@ -31,10 +31,14 @@ concurrency:
|
|||
# number, so a partially-specified entry would fail the call rather than fall
|
||||
# back to the default.
|
||||
#
|
||||
# tests/proxy_unit_tests keeps its own caller (test-unit-proxy-db.yml): it is
|
||||
# already a matrix and carries a shard-coverage guard that reads that file by
|
||||
# name. Folding it in here is a follow-up, together with generalising that guard
|
||||
# into assert_ci_coverage.py.
|
||||
# tests/unit/proxy keeps its own caller (test-unit-proxy-db.yml): it is already
|
||||
# a matrix and carries a shard-coverage guard that reads that file by name.
|
||||
# Folding it in here is a follow-up, together with generalising that guard into
|
||||
# assert_ci_coverage.py.
|
||||
#
|
||||
# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the
|
||||
# shard under the same Codecov flag. CircleCI does not build pull requests from
|
||||
# forks, so the shard still runs those files there and skips them elsewhere.
|
||||
jobs:
|
||||
unit:
|
||||
name: ${{ matrix.shard }}
|
||||
|
|
@ -49,6 +53,7 @@ jobs:
|
|||
- shard: mcp-integration
|
||||
artifact-name: mcp-integration
|
||||
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
|
||||
fork-flag: mcp-integration
|
||||
workers: 2
|
||||
reruns: 0
|
||||
timeout-minutes: 20
|
||||
|
|
@ -65,10 +70,10 @@ jobs:
|
|||
- shard: enterprise-routing
|
||||
artifact-name: enterprise-routing
|
||||
test-path: >-
|
||||
tests/test_litellm/enterprise
|
||||
tests/test_litellm/google_genai
|
||||
tests/test_litellm/router_utils
|
||||
tests/test_litellm/router_strategy
|
||||
fork-flag: enterprise-routing
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -108,6 +113,7 @@ jobs:
|
|||
tests/test_litellm/completion_extras
|
||||
tests/test_litellm/containers
|
||||
tests/test_litellm/endpoints
|
||||
tests/test_litellm/files
|
||||
tests/test_litellm/images
|
||||
tests/test_litellm/interactions
|
||||
tests/test_litellm/messages
|
||||
|
|
@ -199,7 +205,7 @@ jobs:
|
|||
tests/test_litellm/proxy/types_utils
|
||||
tests/test_litellm/proxy/logging_endpoints
|
||||
tests/test_litellm/proxy/test_*.py
|
||||
tests/test_gateway
|
||||
fork-flag: proxy-infra
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -207,11 +213,8 @@ jobs:
|
|||
|
||||
- shard: caching-local
|
||||
artifact-name: caching-local
|
||||
test-path: >-
|
||||
tests/local_testing/test_cache_preset_key.py
|
||||
tests/local_testing/test_caching_handler.py
|
||||
tests/local_testing/test_responses_stream_cache_keys.py
|
||||
tests/local_testing/test_unit_test_caching.py
|
||||
test-path: ""
|
||||
fork-flag: caching-local
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -219,7 +222,8 @@ jobs:
|
|||
|
||||
- shard: proxy-extras
|
||||
artifact-name: proxy-extras
|
||||
test-path: "tests/litellm-proxy-extras"
|
||||
test-path: ""
|
||||
fork-flag: proxy-extras
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -227,7 +231,8 @@ jobs:
|
|||
|
||||
- shard: enterprise-package
|
||||
artifact-name: enterprise-package
|
||||
test-path: "tests/enterprise"
|
||||
test-path: ""
|
||||
fork-flag: enterprise-package
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -246,6 +251,7 @@ jobs:
|
|||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: ${{ matrix.test-path }}
|
||||
fork-flag: ${{ matrix.fork-flag || '' }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: ${{ matrix.reruns }}
|
||||
timeout-minutes: ${{ matrix.timeout-minutes }}
|
||||
|
|
|
|||
12
Makefile
12
Makefile
|
|
@ -51,8 +51,8 @@ help:
|
|||
@echo " make test-unit-core-utils - Run core utils tests (~32 files)"
|
||||
@echo " make test-unit-other - Run other tests (caching, responses, etc., ~69 files)"
|
||||
@echo " make test-unit-root - Run root-level tests (~34 files)"
|
||||
@echo " make test-proxy-unit-a - Run proxy_unit_tests (a-o, ~20 files)"
|
||||
@echo " make test-proxy-unit-b - Run proxy_unit_tests (p-z, ~28 files)"
|
||||
@echo " make test-proxy-unit-a - Run tests/unit/proxy (a-o)"
|
||||
@echo " make test-proxy-unit-b - Run tests/unit/proxy (p-z)"
|
||||
@echo " make test-integration - Run integration tests"
|
||||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
|
||||
|
|
@ -332,17 +332,17 @@ test-unit-core-utils: install-test-deps
|
|||
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-unit-other: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/test_litellm/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/unit/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-root: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
# Proxy unit tests (tests/proxy_unit_tests split alphabetically)
|
||||
# Proxy unit tests (tests/unit/proxy split alphabetically)
|
||||
test-proxy-unit-a: install-test-deps
|
||||
$(UV_RUN) pytest tests/proxy_unit_tests/test_[a-o]*.py --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/proxy --ignore-glob='tests/unit/proxy/test_[p-z]*.py' --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-proxy-unit-b: install-test-deps
|
||||
$(UV_RUN) pytest tests/proxy_unit_tests/test_[p-z]*.py --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/proxy/test_[p-z]*.py tests/unit/skills --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-integration: install-test-deps
|
||||
$(UV_RUN) pytest tests/ -k "not test_litellm"
|
||||
|
|
|
|||
10
litellm-rust/AGENTS.md
Normal file
10
litellm-rust/AGENTS.md
Normal 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
|
||||
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -3414,6 +3414,7 @@ dependencies = [
|
|||
name = "litellm-types"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
",
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
",
|
||||
);
|
||||
});
|
||||
}
|
||||
|
|
@ -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")));
|
||||
});
|
||||
}
|
||||
|
|
@ -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())]);
|
||||
}
|
||||
}
|
||||
|
|
@ -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,
|
||||
)
|
||||
}
|
||||
|
|
@ -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",
|
||||
);
|
||||
});
|
||||
}
|
||||
274
litellm-rust/crates/core-utils/src/dot_notation_indexing.rs
Normal file
274
litellm-rust/crates/core-utils/src/dot_notation_indexing.rs
Normal file
|
|
@ -0,0 +1,274 @@
|
|||
use serde_json::Value;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
enum Segment {
|
||||
Field(String),
|
||||
Every,
|
||||
Index(usize),
|
||||
}
|
||||
|
||||
fn parse_segments(path: &str) -> Option<Vec<Segment>> {
|
||||
let mut segments = Vec::new();
|
||||
let mut rest = path;
|
||||
while !rest.is_empty() {
|
||||
if let Some(after_open) = rest.strip_prefix('[') {
|
||||
let (inside, after) = after_open.split_once(']')?;
|
||||
segments.push(match inside {
|
||||
"*" => Segment::Every,
|
||||
index => Segment::Index(index.trim().parse().ok()?),
|
||||
});
|
||||
rest = after.strip_prefix('.').unwrap_or(after);
|
||||
continue;
|
||||
}
|
||||
let end = rest.find(['.', '[']).unwrap_or(rest.len());
|
||||
let (field, after) = rest.split_at(end);
|
||||
if !field.is_empty() {
|
||||
segments.push(Segment::Field(field.to_string()));
|
||||
}
|
||||
rest = after.strip_prefix('.').unwrap_or(after);
|
||||
}
|
||||
Some(segments)
|
||||
}
|
||||
|
||||
fn without_path(value: Value, segments: &[Segment]) -> Value {
|
||||
let Some((segment, tail)) = segments.split_first() else {
|
||||
return value;
|
||||
};
|
||||
match (segment, value) {
|
||||
(Segment::Field(name), Value::Object(object)) => Value::Object(
|
||||
object
|
||||
.into_iter()
|
||||
.filter_map(|(key, item)| {
|
||||
if key != *name {
|
||||
return Some((key, item));
|
||||
}
|
||||
(!tail.is_empty()).then(|| (key, without_path(item, tail)))
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
(Segment::Every, Value::Array(items)) => Value::Array(
|
||||
items
|
||||
.into_iter()
|
||||
.map(|item| without_path(item, tail))
|
||||
.collect(),
|
||||
),
|
||||
(Segment::Index(index), Value::Array(items)) => Value::Array(
|
||||
items
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(position, item)| {
|
||||
if position == *index {
|
||||
without_path(item, tail)
|
||||
} else {
|
||||
item
|
||||
}
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
(_, value) => value,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn delete_nested_value(value: Value, path: &str) -> Value {
|
||||
match parse_segments(path) {
|
||||
Some(segments) => without_path(value, &segments),
|
||||
None => value,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[fixture]
|
||||
fn body() -> Value {
|
||||
json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
})
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::top_level_field("top", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}}
|
||||
}))]
|
||||
#[case::whole_object("meta", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::nested_field("meta.inner.drop", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::trailing_dot("meta.inner.drop.", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::leading_and_doubled_dots(".meta..inner.drop", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::field_in_every_element("tools[*].examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::whole_array_field_in_every_element("tools[*].arr", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"]},
|
||||
{"name": "t1", "examples": ["b"]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::field_in_indexed_element("tools[1].examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::padded_index("tools[ 1 ].examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::field_right_after_bracket("tools[0]examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::index_then_wildcard("tools[0].arr[*].f", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::wildcard_then_index_only_where_it_exists("tools[*].arr[1].f", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::nested_wildcards("tools[*].arr[*].f", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
fn deletes_the_addressed_field(body: Value, #[case] path: &str, #[case] expected: Value) {
|
||||
assert_eq!(delete_nested_value(body, path), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty_path("")]
|
||||
#[case::missing_field("missing")]
|
||||
#[case::missing_parent("missing.field")]
|
||||
#[case::field_through_a_scalar("top.value")]
|
||||
#[case::field_on_an_array("tools.name")]
|
||||
#[case::index_on_an_object("meta[0].user")]
|
||||
#[case::wildcard_on_an_object("meta[*].user")]
|
||||
#[case::wildcard_over_scalars("tools[*].examples[*].name")]
|
||||
#[case::index_out_of_range("tools[5].name")]
|
||||
#[case::every_element_itself("tools[*]")]
|
||||
#[case::indexed_element_itself("tools[0]")]
|
||||
#[case::nested_element_itself("tools[*].arr[0]")]
|
||||
#[case::negative_index("tools[-1].name")]
|
||||
#[case::non_numeric_index("tools[x].name")]
|
||||
#[case::empty_index("tools[].name")]
|
||||
#[case::unclosed_bracket("top[0")]
|
||||
fn leaves_the_value_untouched(body: Value, #[case] path: &str) {
|
||||
assert_eq!(delete_nested_value(body.clone(), path), body);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::wildcards_indices_and_nesting(
|
||||
json!({"tools": [
|
||||
{"name": "t0", "configs": [{"id": "c0", "remove_me": 1, "keep": 1}, {"id": "c1", "remove_me": 2, "keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
|
||||
{"name": "t1", "configs": [{"id": "c0", "remove_me": 3, "keep": 3}, {"id": "c1", "remove_me": 4, "keep": 4}], "metadata": {"drop_this": 2, "preserve": 2}},
|
||||
{"name": "t2", "configs": [{"id": "c0", "remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
|
||||
]}),
|
||||
&["tools[*].configs[1].remove_me", "tools[1].metadata.drop_this", "tools[*].configs[*].id"],
|
||||
json!({"tools": [
|
||||
{"name": "t0", "configs": [{"remove_me": 1, "keep": 1}, {"keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
|
||||
{"name": "t1", "configs": [{"remove_me": 3, "keep": 3}, {"keep": 4}], "metadata": {"preserve": 2}},
|
||||
{"name": "t2", "configs": [{"remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
|
||||
]}),
|
||||
)]
|
||||
#[case::simple_and_wildcard_nesting(
|
||||
json!({
|
||||
"tools": [{"name": "t1", "simple_nested": {"remove": 1, "keep": 2}, "complex": [{"nested": {"remove": 3, "keep": 4}}]}],
|
||||
"top_level_remove": "should_go",
|
||||
"top_level_keep": "should_stay"
|
||||
}),
|
||||
&["tools[*].simple_nested.remove", "tools[*].complex[*].nested.remove"],
|
||||
json!({
|
||||
"tools": [{"name": "t1", "simple_nested": {"keep": 2}, "complex": [{"nested": {"keep": 4}}]}],
|
||||
"top_level_remove": "should_go",
|
||||
"top_level_keep": "should_stay"
|
||||
}),
|
||||
)]
|
||||
#[case::triple_nested_wildcards(
|
||||
json!({"tools": [{"name": "t1", "arr1": [
|
||||
{"arr2": [{"field": 1, "keep": 1}, {"field": 2, "keep": 2}]},
|
||||
{"arr2": [{"field": 3, "keep": 3}]}
|
||||
]}]}),
|
||||
&["tools[*].arr1[*].arr2[*].field"],
|
||||
json!({"tools": [{"name": "t1", "arr1": [
|
||||
{"arr2": [{"keep": 1}, {"keep": 2}]},
|
||||
{"arr2": [{"keep": 3}]}
|
||||
]}]}),
|
||||
)]
|
||||
fn applies_paths_in_sequence(
|
||||
#[case] value: Value,
|
||||
#[case] paths: &[&str],
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let deleted = paths
|
||||
.iter()
|
||||
.fold(value, |value, path| delete_nested_value(value, path));
|
||||
assert_eq!(deleted, expected);
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,93 @@
|
|||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub fn get_provider_specific_headers(
|
||||
provider_specific_header: Option<&ProviderSpecificHeaders>,
|
||||
custom_llm_provider: &str,
|
||||
) -> Map<String, Value> {
|
||||
let entries: &[ProviderSpecificHeader] = match provider_specific_header {
|
||||
None => &[],
|
||||
Some(ProviderSpecificHeaders::One(entry)) => std::slice::from_ref(entry),
|
||||
Some(ProviderSpecificHeaders::Many(entries)) => entries,
|
||||
};
|
||||
entries
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
entry
|
||||
.custom_llm_provider
|
||||
.split(',')
|
||||
.any(|scoped| scoped.trim() == custom_llm_provider)
|
||||
})
|
||||
.flat_map(|entry| entry.extra_headers.clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::single_entry_for_the_provider(
|
||||
json!({"custom_llm_provider": "anthropic", "extra_headers": {"Authorization": "Bearer t", "Custom-Header": "v"}}),
|
||||
json!({"Authorization": "Bearer t", "Custom-Header": "v"}),
|
||||
)]
|
||||
#[case::single_entry_for_another_provider(
|
||||
json!({"custom_llm_provider": "openai", "extra_headers": {"Authorization": "Bearer t"}}),
|
||||
json!({}),
|
||||
)]
|
||||
#[case::provider_in_a_comma_separated_scope(
|
||||
json!({"custom_llm_provider": "bedrock,anthropic,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}}),
|
||||
json!({"anthropic-beta": "context-1m-2025-08-07"}),
|
||||
)]
|
||||
#[case::provider_missing_from_a_comma_separated_scope(
|
||||
json!({"custom_llm_provider": "bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
|
||||
json!({}),
|
||||
)]
|
||||
#[case::scope_with_spaces(
|
||||
json!({"custom_llm_provider": "bedrock, anthropic , vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
|
||||
json!({"anthropic-beta": "test"}),
|
||||
)]
|
||||
#[case::scope_names_must_match_exactly(
|
||||
json!({"custom_llm_provider": "anthropic_text", "extra_headers": {"anthropic-beta": "test"}}),
|
||||
json!({}),
|
||||
)]
|
||||
#[case::entries_scope_independently(
|
||||
json!([
|
||||
{"custom_llm_provider": "anthropic,bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}},
|
||||
{"custom_llm_provider": "bedrock", "extra_headers": {"x-bedrock-only": "no"}},
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"authorization": "Bearer sk-ant-oat01-fake-token"}}
|
||||
]),
|
||||
json!({"anthropic-beta": "context-1m-2025-08-07", "authorization": "Bearer sk-ant-oat01-fake-token"}),
|
||||
)]
|
||||
#[case::later_entries_win(
|
||||
json!([
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "first"}},
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "second"}}
|
||||
]),
|
||||
json!({"x-scoped": "second"}),
|
||||
)]
|
||||
#[case::empty_list(json!([]), json!({}))]
|
||||
#[case::entry_without_scope(json!({"extra_headers": {"x-scoped": "yes"}}), json!({}))]
|
||||
#[case::entry_without_headers(json!({"custom_llm_provider": "anthropic"}), json!({}))]
|
||||
fn provider_specific_headers_match_the_scoped_provider(
|
||||
#[case] configured: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let configured: ProviderSpecificHeaders = serde_json::from_value(configured).unwrap();
|
||||
assert_eq!(
|
||||
Value::Object(get_provider_specific_headers(
|
||||
Some(&configured),
|
||||
"anthropic"
|
||||
)),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn no_configured_headers_match_nothing() {
|
||||
assert_eq!(get_provider_specific_headers(None, "anthropic"), Map::new());
|
||||
}
|
||||
}
|
||||
|
|
@ -1,7 +1,9 @@
|
|||
pub mod call_arguments;
|
||||
pub mod core_helpers;
|
||||
pub mod dot_notation_indexing;
|
||||
pub mod exception_mapping_utils;
|
||||
pub mod get_llm_provider_logic;
|
||||
pub mod get_provider_specific_headers;
|
||||
pub mod params;
|
||||
pub mod prompt_templates;
|
||||
pub mod secret_redaction;
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ version = "0.1.0"
|
|||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
autotests = false
|
||||
|
||||
[dependencies]
|
||||
litellm-secrets.workspace = true
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -51,6 +51,3 @@ pub fn chat_completions_decline_reason(
|
|||
.unsupported_reason(&messages, optional_params)
|
||||
.map(|reason| reason.0)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -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, ¶ms)
|
||||
}
|
||||
|
||||
#[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, .. })
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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, ¶ms)
|
||||
}
|
||||
|
||||
#[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, .. })
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
use litellm_http::request::string_headers as shared_string_headers;
|
||||
pub(super) use litellm_http::request::{has_bearer_auth, has_header, truncate_error_body};
|
||||
pub(super) use litellm_http::request::truncate_error_body;
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
|
|
@ -26,3 +26,186 @@ pub(super) fn string_headers(
|
|||
) -> Result<Vec<(String, String)>, Error> {
|
||||
shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(Error::from)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_secrets::{SecretValue, source::SecretSource};
|
||||
use serde_json::{Value, json};
|
||||
use tokio::{
|
||||
io::{AsyncReadExt, AsyncWriteExt},
|
||||
net::{TcpListener, TcpStream},
|
||||
};
|
||||
|
||||
use super::{messages_provider_config, string_headers, truncate_error_body};
|
||||
use crate::messages::{
|
||||
Error,
|
||||
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
|
||||
types::MessagesShaping,
|
||||
};
|
||||
|
||||
struct RecordingSecrets {
|
||||
values: Vec<(&'static str, String)>,
|
||||
requested: std::sync::Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl SecretSource for RecordingSecrets {
|
||||
fn get_secret_str<'a>(
|
||||
&'a self,
|
||||
name: &'a str,
|
||||
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
|
||||
Box::pin(async move {
|
||||
self.requested.lock().unwrap().push(name.to_string());
|
||||
Ok(self
|
||||
.values
|
||||
.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| SecretValue::new(value.clone())))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn secrets_call() -> MessagesCall {
|
||||
let Value::Object(body) = json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}) else {
|
||||
unreachable!("literal object")
|
||||
};
|
||||
MessagesCall {
|
||||
model: "claude-sonnet-4-5".into(),
|
||||
body,
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn read_http_request(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
let mut buffer = [0_u8; 1024];
|
||||
let header_end = loop {
|
||||
let n = socket.read(&mut buffer).await.expect("reads request");
|
||||
if n == 0 {
|
||||
break request.len();
|
||||
}
|
||||
request.extend_from_slice(&buffer[..n]);
|
||||
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") {
|
||||
break position + 4;
|
||||
}
|
||||
};
|
||||
let headers = String::from_utf8_lossy(&request[..header_end]);
|
||||
let content_length = headers
|
||||
.lines()
|
||||
.find_map(|line| {
|
||||
let (name, value) = line.split_once(':')?;
|
||||
name.eq_ignore_ascii_case("content-length")
|
||||
.then(|| value.trim().parse::<usize>().ok())
|
||||
.flatten()
|
||||
})
|
||||
.unwrap_or(0);
|
||||
while request.len().saturating_sub(header_end) < content_length {
|
||||
let n = socket.read(&mut buffer).await.expect("reads body");
|
||||
if n == 0 {
|
||||
break;
|
||||
}
|
||||
request.extend_from_slice(&buffer[..n]);
|
||||
}
|
||||
String::from_utf8(request).expect("request is utf8")
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn route_reads_the_provider_credential_and_base_from_the_secret_source() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
let addr = listener.local_addr().expect("addr");
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
||||
let request = read_http_request(&mut socket).await;
|
||||
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#;
|
||||
let response = format!(
|
||||
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
|
||||
response_body.len(),
|
||||
response_body
|
||||
);
|
||||
socket
|
||||
.write_all(response.as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
let secrets = Arc::new(RecordingSecrets {
|
||||
values: vec![
|
||||
("ANTHROPIC_API_KEY", "sk-from-manager".to_string()),
|
||||
("ANTHROPIC_BASE_URL", format!("http://{addr}")),
|
||||
],
|
||||
requested: std::sync::Mutex::new(Vec::new()),
|
||||
});
|
||||
|
||||
let output = litellm_host::run::run(
|
||||
messages_machine(secrets.clone()),
|
||||
&LocalMessagesHost::new(secrets_call()),
|
||||
)
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
let request = server.await.expect("server task completes");
|
||||
assert!(
|
||||
request
|
||||
.to_ascii_lowercase()
|
||||
.contains("x-api-key: sk-from-manager"),
|
||||
"{request}"
|
||||
);
|
||||
let requested = secrets.requested.lock().unwrap().clone();
|
||||
assert_eq!(
|
||||
requested,
|
||||
messages_provider_config("anthropic")
|
||||
.unwrap()
|
||||
.secret_names()
|
||||
.iter()
|
||||
.map(ToString::to_string)
|
||||
.collect::<Vec<_>>()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_resolves_anthropic_and_azure_ai() {
|
||||
assert!(messages_provider_config("anthropic").is_some());
|
||||
assert!(messages_provider_config("azure_ai").is_some());
|
||||
assert!(messages_provider_config("openai").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncate_error_body_caps_long_payloads() {
|
||||
let body = "x".repeat(400);
|
||||
let truncated = truncate_error_body(&body);
|
||||
assert!(truncated.ends_with("... (truncated)"));
|
||||
let prefix_chars = truncated
|
||||
.strip_suffix("... (truncated)")
|
||||
.expect("truncated marker present")
|
||||
.chars()
|
||||
.count();
|
||||
assert_eq!(prefix_chars, 256);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn string_headers_rejects_non_string_values() {
|
||||
let headers = json!({"x-count": 3}).as_object().unwrap().clone();
|
||||
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
|
||||
assert_eq!(
|
||||
err,
|
||||
Error::Headers(litellm_http::request::HeaderError {
|
||||
context: "messages",
|
||||
name: "x-count".to_string(),
|
||||
actual: "number",
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
|
|
@ -18,8 +20,34 @@ pub enum Error {
|
|||
Transport(#[from] litellm_http::transport::Error),
|
||||
#[error(transparent)]
|
||||
Headers(#[from] litellm_http::request::HeaderError),
|
||||
#[error(transparent)]
|
||||
Secret(#[from] SecretError),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, thiserror::Error)]
|
||||
#[error(transparent)]
|
||||
pub struct SecretError(Arc<litellm_secrets::Error>);
|
||||
|
||||
impl SecretError {
|
||||
pub fn source_error(&self) -> &litellm_secrets::Error {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl From<litellm_secrets::Error> for Error {
|
||||
fn from(error: litellm_secrets::Error) -> Self {
|
||||
Self::Secret(SecretError(Arc::new(error)))
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for SecretError {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
Arc::ptr_eq(&self.0, &other.0)
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for SecretError {}
|
||||
|
||||
impl From<LlmError> for Error {
|
||||
fn from(error: LlmError) -> Self {
|
||||
match error {
|
||||
|
|
|
|||
|
|
@ -12,6 +12,9 @@ mod common_utils;
|
|||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_secrets::source::EnvironmentSecrets;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
|
||||
use serde_json::Value;
|
||||
|
|
@ -31,15 +34,15 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
|
|||
api_base: request.api_base.map(Into::into),
|
||||
custom_llm_provider: request.custom_llm_provider.map(Into::into),
|
||||
extra_headers: request.extra_headers,
|
||||
provider_specific_header: request.provider_specific_header,
|
||||
timeout: request.timeout,
|
||||
shaping: request.shaping,
|
||||
};
|
||||
match litellm_host::run::run(messages_machine(), &LocalMessagesHost::new(call)).await? {
|
||||
let secrets = Arc::new(EnvironmentSecrets::python_compatible());
|
||||
match litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await? {
|
||||
MessagesOutput::Message(message) => Ok(*message),
|
||||
MessagesOutput::Streamed => Err(Error::Unsupported(
|
||||
"streamed responses need a streaming host",
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -1,51 +1,102 @@
|
|||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_llms::base_llm::anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, MessagesAuthStrategy,
|
||||
use litellm_core_utils::{
|
||||
dot_notation_indexing::delete_nested_value,
|
||||
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
|
||||
get_provider_specific_headers::get_provider_specific_headers,
|
||||
settings::Lookup,
|
||||
};
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request,
|
||||
base_llm::anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, MessagesTransformContext,
|
||||
},
|
||||
};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers},
|
||||
common_utils::{messages_provider_config, string_headers},
|
||||
};
|
||||
use crate::messages::types::{MessagesRequest, ProviderMessagesRequest};
|
||||
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: MessagesRequest<'_>,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
|
||||
pub(super) struct ResolvedProvider<'a> {
|
||||
pub(super) model: &'a str,
|
||||
pub(super) provider: &'a str,
|
||||
pub(super) config: &'static dyn BaseAnthropicMessagesConfig,
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<ResolvedProvider<'a>, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
} = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
request
|
||||
.custom_llm_provider
|
||||
.map(|provider| CustomLlmProvider {
|
||||
model: request.model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for messages request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let model = provider_info.model.to_string();
|
||||
let provider = provider_info.custom_llm_provider;
|
||||
|
||||
let config = messages_provider_config(provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
Ok(ResolvedProvider {
|
||||
model,
|
||||
provider,
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
||||
let headers =
|
||||
validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?;
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: MessagesRequest<'_>,
|
||||
resolved: ResolvedProvider<'_>,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let ResolvedProvider {
|
||||
model,
|
||||
provider,
|
||||
config,
|
||||
} = resolved;
|
||||
let model = model.to_string();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
|
||||
let typed_request: AnthropicMessagesRequest =
|
||||
serde_json::from_value(request.body).map_err(|err| {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
})?;
|
||||
let transformed = config.transform_anthropic_messages_request(AnthropicMessagesRequest {
|
||||
model: model.clone(),
|
||||
..typed_request
|
||||
})?;
|
||||
serde_json::from_value(request.body).map_err(invalid_request)?;
|
||||
let sanitized = shape_anthropic_messages_request(
|
||||
AnthropicMessagesRequest {
|
||||
model: model.clone(),
|
||||
..typed_request
|
||||
},
|
||||
request.shaping.reasoning_auto_summary,
|
||||
)?;
|
||||
let trimmed =
|
||||
without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?;
|
||||
let transformed = config.transform_anthropic_messages_request(
|
||||
trimmed,
|
||||
&MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params),
|
||||
)?;
|
||||
|
||||
let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider);
|
||||
let forwarded = string_headers(Some(
|
||||
request
|
||||
.extra_headers
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(scoped)
|
||||
.collect(),
|
||||
))?;
|
||||
let authenticated = config.authenticate(forwarded, request.api_key, &env_lookup)?;
|
||||
let headers = config.request_headers(
|
||||
with_default_headers(authenticated, config.default_headers()),
|
||||
&transformed,
|
||||
);
|
||||
|
||||
let body = serde_json::to_value(transformed).map_err(|err| {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
|
|
@ -65,33 +116,371 @@ pub(super) fn prepare_provider_request(
|
|||
})
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
config: &dyn BaseAnthropicMessagesConfig,
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
let mut headers = string_headers(extra_headers)?;
|
||||
|
||||
let auth_strategy = config.auth_strategy();
|
||||
let already_authorized = has_header(&headers, auth_strategy.header_name())
|
||||
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
|
||||
if !already_authorized {
|
||||
let api_key = config.resolve_api_key(api_key, env_lookup)?;
|
||||
let auth_header = match auth_strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
("authorization".to_string(), format!("Bearer {api_key}"))
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
headers.push(auth_header);
|
||||
}
|
||||
|
||||
for (name, value) in config.default_headers() {
|
||||
if !has_header(&headers, name) {
|
||||
headers.push((name.to_string(), value.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(headers)
|
||||
fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
}
|
||||
|
||||
fn without_additional_drop_params(
|
||||
request: AnthropicMessagesRequest,
|
||||
paths: &[String],
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
if paths.is_empty() {
|
||||
return Ok(request);
|
||||
}
|
||||
let Value::Object(fields) = serde_json::to_value(request).map_err(invalid_request)? else {
|
||||
return Err(Error::InvalidRequest(
|
||||
"Anthropic messages request did not serialize to an object".to_string(),
|
||||
));
|
||||
};
|
||||
let (required, optional): (Map<String, Value>, Map<String, Value>) = fields
|
||||
.into_iter()
|
||||
.partition(|(key, _)| matches!(key.as_str(), "model" | "messages"));
|
||||
let trimmed = paths.iter().fold(Value::Object(optional), |body, path| {
|
||||
delete_nested_value(body, path)
|
||||
});
|
||||
let merged: Map<String, Value> = required
|
||||
.into_iter()
|
||||
.chain(trimmed.as_object().cloned().unwrap_or_default())
|
||||
.collect();
|
||||
serde_json::from_value(Value::Object(merged)).map_err(invalid_request)
|
||||
}
|
||||
|
||||
fn with_default_headers(
|
||||
headers: Vec<(String, String)>,
|
||||
defaults: &[(&str, &str)],
|
||||
) -> Vec<(String, String)> {
|
||||
let missing: Vec<(String, String)> = defaults
|
||||
.iter()
|
||||
.filter(|(name, _)| {
|
||||
!headers
|
||||
.iter()
|
||||
.any(|(header, _)| header.eq_ignore_ascii_case(name))
|
||||
})
|
||||
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
|
||||
.collect();
|
||||
headers.into_iter().chain(missing).collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
use crate::messages::types::MessagesShaping;
|
||||
|
||||
#[fixture]
|
||||
fn shaping() -> MessagesShaping {
|
||||
MessagesShaping::default()
|
||||
}
|
||||
|
||||
fn prepare(request: MessagesRequest<'_>) -> Result<ProviderMessagesRequest, Error> {
|
||||
prepare_with_secrets(request, &|_: &str| None)
|
||||
}
|
||||
|
||||
fn prepare_with_secrets(
|
||||
request: MessagesRequest<'_>,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let resolved = resolve_provider(request.model, request.custom_llm_provider)?;
|
||||
prepare_provider_request(request, resolved, secrets)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key(
|
||||
&[("ANTHROPIC_API_KEY", "sk-secret")],
|
||||
&[("x-api-key", "sk-secret")],
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
)]
|
||||
#[case::auth_token(
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "token")],
|
||||
&[("authorization", "Bearer token")],
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
)]
|
||||
#[case::api_base(
|
||||
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_API_BASE", "https://gateway.test")],
|
||||
&[("x-api-key", "sk-secret")],
|
||||
"https://gateway.test/v1/messages"
|
||||
)]
|
||||
#[case::sdk_base_url(
|
||||
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_BASE_URL", "https://sdk.test")],
|
||||
&[("x-api-key", "sk-secret")],
|
||||
"https://sdk.test/v1/messages"
|
||||
)]
|
||||
fn credentials_and_base_come_from_the_resolved_secrets(
|
||||
shaping: MessagesShaping,
|
||||
#[case] secrets: &[(&str, &str)],
|
||||
#[case] expected_auth: &[(&str, &str)],
|
||||
#[case] expected_url: &str,
|
||||
) {
|
||||
let lookup = |name: &str| {
|
||||
secrets
|
||||
.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
};
|
||||
let prepared = prepare_with_secrets(
|
||||
MessagesRequest {
|
||||
model: "claude-test",
|
||||
body: json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: None,
|
||||
shaping,
|
||||
},
|
||||
&lookup,
|
||||
)
|
||||
.unwrap();
|
||||
let auth: Vec<(&str, &str)> = prepared
|
||||
.upstream_headers
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization"))
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
.collect();
|
||||
assert_eq!(
|
||||
(auth.as_slice(), prepared.url.as_str()),
|
||||
(expected_auth, expected_url)
|
||||
);
|
||||
}
|
||||
|
||||
fn prepared_body(body: Value, shaping: MessagesShaping) -> Result<Value, Error> {
|
||||
prepare(MessagesRequest {
|
||||
model: "anthropic/claude-test",
|
||||
body,
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some("https://anthropic.test"),
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: None,
|
||||
shaping,
|
||||
})
|
||||
.map(|prepared| prepared.body)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::nothing_forwarded(
|
||||
&[],
|
||||
&[("x-version", "1"), ("content-type", "application/json")],
|
||||
&[("x-version", "1"), ("content-type", "application/json")],
|
||||
)]
|
||||
#[case::forwarded_header_wins_in_any_case(
|
||||
&[("X-Version", "custom"), ("x-api-key", "k")],
|
||||
&[("x-version", "1"), ("content-type", "application/json")],
|
||||
&[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")],
|
||||
)]
|
||||
#[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])]
|
||||
fn default_headers_fill_only_missing_names(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] defaults: &[(&str, &str)],
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
let owned = |headers: &[(&str, &str)]| -> Vec<(String, String)> {
|
||||
headers
|
||||
.iter()
|
||||
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
|
||||
.collect()
|
||||
};
|
||||
assert_eq!(
|
||||
with_default_headers(owned(forwarded), defaults),
|
||||
owned(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::top_level_and_nested_paths(
|
||||
json!({
|
||||
"max_tokens": 1024,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048},
|
||||
"context_management": {"edits": [{"type": "clear_thinking_20251015"}]},
|
||||
"metadata": {"user_id": "u1"},
|
||||
"tools": [{"name": "lookup", "input_schema": {"type": "object"}, "input_examples": [{"q": "x"}]}]
|
||||
}),
|
||||
&["thinking", "context_management", "tools[*].input_examples"],
|
||||
json!({
|
||||
"max_tokens": 1024,
|
||||
"metadata": {"user_id": "u1"},
|
||||
"tools": [{"name": "lookup", "input_schema": {"type": "object"}}]
|
||||
}),
|
||||
)]
|
||||
#[case::no_paths(
|
||||
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
|
||||
&[],
|
||||
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
|
||||
)]
|
||||
#[case::model_and_messages_are_never_dropped(
|
||||
json!({"max_tokens": 16}),
|
||||
&["model", "messages", "messages[0].content"],
|
||||
json!({"max_tokens": 16}),
|
||||
)]
|
||||
fn prepared_body_drops_configured_paths(
|
||||
shaping: MessagesShaping,
|
||||
#[case] fields: Value,
|
||||
#[case] additional_drop_params: &[&str],
|
||||
#[case] expected_fields: Value,
|
||||
) {
|
||||
let with_messages = |fields: Value| -> Value {
|
||||
let Value::Object(fields) = fields else {
|
||||
unreachable!()
|
||||
};
|
||||
Value::Object(
|
||||
[
|
||||
("model".to_string(), json!("claude-test")),
|
||||
(
|
||||
"messages".to_string(),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
),
|
||||
]
|
||||
.into_iter()
|
||||
.chain(fields)
|
||||
.collect(),
|
||||
)
|
||||
};
|
||||
let shaping = MessagesShaping {
|
||||
additional_drop_params: additional_drop_params
|
||||
.iter()
|
||||
.map(ToString::to_string)
|
||||
.collect(),
|
||||
..shaping
|
||||
};
|
||||
assert_eq!(
|
||||
prepared_body(with_messages(fields), shaping),
|
||||
Ok(with_messages(expected_fields))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::model_prefix_picks_the_provider(
|
||||
"azure_ai/claude-test",
|
||||
None,
|
||||
&[("x-priority", "extra"), ("x-scoped", "azure_ai")]
|
||||
)]
|
||||
#[case::explicit_provider(
|
||||
"claude-test",
|
||||
Some("anthropic"),
|
||||
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
|
||||
)]
|
||||
#[case::provider_prefix_on_an_anthropic_model(
|
||||
"anthropic/claude-test",
|
||||
None,
|
||||
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
|
||||
)]
|
||||
fn provider_specific_headers_follow_the_resolved_provider(
|
||||
shaping: MessagesShaping,
|
||||
#[case] model: &str,
|
||||
#[case] custom_llm_provider: Option<&str>,
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
let configured: ProviderSpecificHeaders = serde_json::from_value(json!([
|
||||
{"custom_llm_provider": "azure_ai", "extra_headers": {"x-scoped": "azure_ai"}},
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "anthropic", "x-priority": "scoped"}}
|
||||
]))
|
||||
.unwrap();
|
||||
let prepared = prepare(MessagesRequest {
|
||||
model,
|
||||
body: json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some("https://resource.services.ai.azure.com"),
|
||||
custom_llm_provider,
|
||||
extra_headers: Some(serde_json::from_value(json!({"x-priority": "extra"})).unwrap()),
|
||||
provider_specific_header: Some(configured),
|
||||
timeout: None,
|
||||
shaping,
|
||||
})
|
||||
.unwrap();
|
||||
let caller_headers: Vec<(&str, &str)> = prepared
|
||||
.upstream_headers
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped"))
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
.collect();
|
||||
assert_eq!(caller_headers, expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn prepared_body_carries_the_provider_stripped_model(shaping: MessagesShaping) {
|
||||
assert_eq!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "anthropic/claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Ok(json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn dropped_thinking_display_is_not_restored_by_auto_summary(shaping: MessagesShaping) {
|
||||
let shaping = MessagesShaping {
|
||||
reasoning_auto_summary: true,
|
||||
additional_drop_params: vec!["thinking.display".to_string()],
|
||||
..shaping
|
||||
};
|
||||
assert_eq!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 4096,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048}
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Ok(json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 4096,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048}
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn dropping_an_invalid_metadata_user_id_does_not_skip_its_validation(shaping: MessagesShaping) {
|
||||
let shaping = MessagesShaping {
|
||||
additional_drop_params: vec!["metadata.user_id".to_string()],
|
||||
..shaping
|
||||
};
|
||||
assert!(matches!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16,
|
||||
"metadata": {"user_id": 123}
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Err(Error::InvalidRequest(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn prepared_body_rejects_invalid_metadata_before_the_call(shaping: MessagesShaping) {
|
||||
assert_eq!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16,
|
||||
"metadata": {"user_id": 123}
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Err(Error::InvalidRequest(
|
||||
"metadata.user_id must be a string, got 123".to_string()
|
||||
))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,7 @@
|
|||
use std::{sync::Mutex, time::Duration};
|
||||
use std::{
|
||||
sync::{Arc, Mutex},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_auth::SecretValue;
|
||||
|
|
@ -9,15 +12,19 @@ use litellm_host::{
|
|||
machine::{HostChannel, MachineFault, RouteMachine},
|
||||
route::Route,
|
||||
};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
|
||||
utils::ProviderSpecificHeaders,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::messages_provider_config,
|
||||
handler::{decode_response, network, provider_error, send},
|
||||
prepare::prepare_provider_request,
|
||||
types::MessagesRequest,
|
||||
prepare::{prepare_provider_request, resolve_provider},
|
||||
types::{MessagesRequest, MessagesShaping},
|
||||
};
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
|
||||
|
|
@ -38,7 +45,9 @@ pub struct MessagesCall {
|
|||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub provider_specific_header: Option<ProviderSpecificHeaders>,
|
||||
pub timeout: Option<Duration>,
|
||||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
||||
impl MessagesCall {
|
||||
|
|
@ -120,22 +129,33 @@ impl Host<Messages> for LocalMessagesHost {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn messages_machine() -> MessagesMachine {
|
||||
RouteMachine::new(|host| Box::pin(execute(host)))
|
||||
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
|
||||
RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
|
||||
}
|
||||
|
||||
async fn execute(host: MessagesHost) -> Result<MessagesOutput, Error> {
|
||||
async fn execute(
|
||||
host: MessagesHost,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?;
|
||||
let stream = call.streams();
|
||||
let request = prepare_provider_request(MessagesRequest {
|
||||
model: &call.model,
|
||||
body: Value::Object(call.body.clone()),
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers.clone(),
|
||||
timeout: call.timeout,
|
||||
})?;
|
||||
let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
|
||||
let secrets = secrets.resolve(resolved.config.secret_names()).await?;
|
||||
let request = prepare_provider_request(
|
||||
MessagesRequest {
|
||||
model: &call.model,
|
||||
body: Value::Object(call.body.clone()),
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers.clone(),
|
||||
provider_specific_header: call.provider_specific_header.clone(),
|
||||
timeout: call.timeout,
|
||||
shaping: call.shaping.clone(),
|
||||
},
|
||||
resolved,
|
||||
secrets.as_ref(),
|
||||
)?;
|
||||
if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER {
|
||||
return Err(Error::Unsupported("streaming messages for this provider"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,25 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig;
|
||||
use litellm_llms::{
|
||||
anthropic::common_utils::AnthropicModelCapabilities,
|
||||
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
|
||||
};
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct MessagesShaping {
|
||||
#[serde(default)]
|
||||
pub capabilities: AnthropicModelCapabilities,
|
||||
#[serde(default)]
|
||||
pub drop_params: bool,
|
||||
#[serde(default)]
|
||||
pub reasoning_auto_summary: bool,
|
||||
#[serde(default)]
|
||||
pub additional_drop_params: Vec<String>,
|
||||
}
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub body: Value,
|
||||
|
|
@ -10,7 +27,9 @@ pub struct MessagesRequest<'a> {
|
|||
pub api_base: Option<&'a str>,
|
||||
pub custom_llm_provider: Option<&'a str>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub provider_specific_header: Option<ProviderSpecificHeaders>,
|
||||
pub timeout: Option<Duration>,
|
||||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
||||
pub struct ProviderMessagesRequest {
|
||||
|
|
@ -22,3 +41,86 @@ pub struct ProviderMessagesRequest {
|
|||
pub upstream_headers: Vec<(String, String)>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::nothing_projected(json!({}), MessagesShaping::default())]
|
||||
#[case::only_drop_params(
|
||||
json!({"drop_params": true}),
|
||||
MessagesShaping { drop_params: true, ..MessagesShaping::default() },
|
||||
)]
|
||||
#[case::only_reasoning_auto_summary(
|
||||
json!({"reasoning_auto_summary": true}),
|
||||
MessagesShaping { reasoning_auto_summary: true, ..MessagesShaping::default() },
|
||||
)]
|
||||
#[case::only_additional_drop_params(
|
||||
json!({"additional_drop_params": ["tools[*].input_examples"]}),
|
||||
MessagesShaping {
|
||||
additional_drop_params: vec!["tools[*].input_examples".to_string()],
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
)]
|
||||
#[case::partial_capabilities(
|
||||
json!({"capabilities": {"supports_reasoning": true}}),
|
||||
MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
)]
|
||||
#[case::everything_the_python_host_projects(
|
||||
json!({
|
||||
"capabilities": {
|
||||
"supports_reasoning": true,
|
||||
"supports_adaptive_thinking": true,
|
||||
"thinking_always_on": false,
|
||||
"supports_legacy_thinking": false,
|
||||
"supports_output_config": true,
|
||||
"supports_sampling_params": false,
|
||||
"supports_speed": true,
|
||||
"effort_tiers": {"minimal": false, "low": true, "medium": true, "high": true, "xhigh": true, "max": false}
|
||||
},
|
||||
"drop_params": true,
|
||||
"reasoning_auto_summary": true,
|
||||
"additional_drop_params": ["metadata.user_id", "thinking"]
|
||||
}),
|
||||
MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
thinking_always_on: false,
|
||||
supports_legacy_thinking: false,
|
||||
supports_output_config: true,
|
||||
supports_sampling_params: false,
|
||||
supports_speed: true,
|
||||
effort_tiers: SupportedEffortTiers {
|
||||
minimal: false,
|
||||
low: true,
|
||||
medium: true,
|
||||
high: true,
|
||||
xhigh: true,
|
||||
max: false,
|
||||
},
|
||||
},
|
||||
drop_params: true,
|
||||
reasoning_auto_summary: true,
|
||||
additional_drop_params: vec!["metadata.user_id".to_string(), "thinking".to_string()],
|
||||
},
|
||||
)]
|
||||
fn shaping_deserializes_with_defaults_for_absent_fields(
|
||||
#[case] projected: Value,
|
||||
#[case] expected: MessagesShaping,
|
||||
) {
|
||||
let shaping: MessagesShaping = serde_json::from_value(projected).unwrap();
|
||||
assert_eq!(shaping, expected);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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");
|
||||
|
|
@ -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");
|
||||
}
|
||||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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()),
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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(), ¶ms, &[])
|
||||
.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());
|
||||
}
|
||||
}
|
||||
|
|
@ -1,19 +1,89 @@
|
|||
use std::time::Duration;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_core::messages::{
|
||||
Error, messages,
|
||||
route::{LocalMessagesHost, MessagesCall, messages_machine},
|
||||
types::{MessagesRequest, MessagesShaping},
|
||||
};
|
||||
use litellm_secrets::{SecretValue, source::SecretSource};
|
||||
use serde_json::{Map, Value, json};
|
||||
use tokio::{
|
||||
io::{AsyncReadExt, AsyncWriteExt},
|
||||
net::{TcpListener, TcpStream},
|
||||
};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{
|
||||
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
|
||||
},
|
||||
messages,
|
||||
};
|
||||
use crate::messages::types::MessagesRequest;
|
||||
struct RecordingSecrets {
|
||||
values: Vec<(&'static str, String)>,
|
||||
fails: bool,
|
||||
requested: std::sync::Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl RecordingSecrets {
|
||||
fn new(values: Vec<(&'static str, String)>, fails: bool) -> Self {
|
||||
Self {
|
||||
values,
|
||||
fails,
|
||||
requested: std::sync::Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl SecretSource for RecordingSecrets {
|
||||
fn get_secret_str<'a>(
|
||||
&'a self,
|
||||
name: &'a str,
|
||||
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
|
||||
Box::pin(async move {
|
||||
self.requested.lock().unwrap().push(name.to_string());
|
||||
if self.fails {
|
||||
return Err(litellm_secrets::Error::ManagedSecretMissing);
|
||||
}
|
||||
Ok(self
|
||||
.values
|
||||
.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| SecretValue::new(value.clone())))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn secrets_call() -> MessagesCall {
|
||||
let Value::Object(body) = json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}) else {
|
||||
unreachable!("literal object")
|
||||
};
|
||||
MessagesCall {
|
||||
model: "claude-sonnet-4-5".into(),
|
||||
body,
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn route_surfaces_a_secret_manager_failure_before_the_call() {
|
||||
let Err(error) = litellm_host::run::run(
|
||||
messages_machine(Arc::new(RecordingSecrets::new(Vec::new(), true))),
|
||||
&LocalMessagesHost::new(secrets_call()),
|
||||
)
|
||||
.await
|
||||
else {
|
||||
panic!("a secret manager failure fails the call");
|
||||
};
|
||||
assert!(
|
||||
matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)),
|
||||
"{error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
async fn read_http_request(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
|
|
@ -56,75 +126,6 @@ fn write_response(body: &str) -> String {
|
|||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_resolves_anthropic_and_azure_ai() {
|
||||
assert!(messages_provider_config("anthropic").is_some());
|
||||
assert!(messages_provider_config("azure_ai").is_some());
|
||||
assert!(messages_provider_config("openai").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncate_error_body_caps_long_payloads() {
|
||||
let body = "x".repeat(400);
|
||||
let truncated = truncate_error_body(&body);
|
||||
assert!(truncated.ends_with("... (truncated)"));
|
||||
let prefix_chars = truncated
|
||||
.strip_suffix("... (truncated)")
|
||||
.expect("truncated marker present")
|
||||
.chars()
|
||||
.count();
|
||||
assert_eq!(prefix_chars, 256);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn string_headers_rejects_non_string_values() {
|
||||
let headers = json!({"x-count": 3}).as_object().unwrap().clone();
|
||||
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
|
||||
assert_eq!(
|
||||
err,
|
||||
Error::Headers(litellm_http::request::HeaderError {
|
||||
context: "messages",
|
||||
name: "x-count".to_string(),
|
||||
actual: "number",
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn has_header_is_case_insensitive() {
|
||||
let headers = vec![("X-Api-Key".to_string(), "secret".to_string())];
|
||||
assert!(has_header(&headers, "x-api-key"));
|
||||
assert!(!has_header(&headers, "authorization"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn has_bearer_auth_requires_a_nonempty_bearer_token() {
|
||||
assert!(has_bearer_auth(&[(
|
||||
"Authorization".to_string(),
|
||||
"Bearer tok".to_string()
|
||||
)]));
|
||||
assert!(has_bearer_auth(&[(
|
||||
"authorization".to_string(),
|
||||
"bearer tok".to_string()
|
||||
)]));
|
||||
assert!(!has_bearer_auth(&[(
|
||||
"authorization".to_string(),
|
||||
"Bearer ".to_string()
|
||||
)]));
|
||||
assert!(!has_bearer_auth(&[(
|
||||
"authorization".to_string(),
|
||||
String::new()
|
||||
)]));
|
||||
assert!(!has_bearer_auth(&[(
|
||||
"authorization".to_string(),
|
||||
"Basic abc".to_string()
|
||||
)]));
|
||||
assert!(!has_bearer_auth(&[(
|
||||
"x-api-key".to_string(),
|
||||
"sk".to_string()
|
||||
)]));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_round_trip_builds_azure_request_and_passes_response_through() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
|
|
@ -159,7 +160,9 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
|
@ -215,7 +218,9 @@ async fn messages_round_trip_builds_native_anthropic_request() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
|
@ -268,7 +273,9 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
|
@ -322,7 +329,9 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("entra id request succeeds without api key");
|
||||
|
|
@ -346,7 +355,9 @@ async fn messages_requires_auth_when_no_key_and_no_header() {
|
|||
api_base: Some("http://127.0.0.1:1"),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_millis(50)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect_err("missing auth errors");
|
||||
|
|
@ -384,7 +395,9 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("falls back to api key");
|
||||
|
|
@ -425,7 +438,9 @@ async fn messages_maps_provider_error_status_to_http_error() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect_err("provider error propagates");
|
||||
|
|
@ -445,7 +460,9 @@ async fn messages_rejects_unsupported_provider() {
|
|||
api_base: Some("http://127.0.0.1:1"),
|
||||
custom_llm_provider: Some("openai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_millis(50)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect_err("unsupported provider errors");
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -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)
|
||||
);
|
||||
}
|
||||
|
|
@ -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())
|
||||
})
|
||||
}
|
||||
|
|
@ -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,
|
||||
¶ms,
|
||||
&[],
|
||||
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"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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"})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
@ -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()
|
||||
)]));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -218,7 +218,3 @@ fn anthropic_body(
|
|||
);
|
||||
Value::Object(body)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "tests.rs"]
|
||||
mod tests;
|
||||
|
|
|
|||
1568
litellm-rust/crates/llms/src/anthropic/common_utils.rs
Normal file
1568
litellm-rust/crates/llms/src/anthropic/common_utils.rs
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -0,0 +1,270 @@
|
|||
use litellm_types::llms::anthropic_messages::anthropic_request::{
|
||||
AnthropicMessage, AnthropicMessagesRequest,
|
||||
};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::{
|
||||
anthropic::common_utils::{
|
||||
flatten_unencrypted_web_search_results, sanitize_tool_use_ids, strip_empty_content_blocks,
|
||||
strip_provider_specific_fields,
|
||||
},
|
||||
base_llm::chat::transformation::Error,
|
||||
};
|
||||
|
||||
pub fn shape_anthropic_messages_request(
|
||||
request: AnthropicMessagesRequest,
|
||||
reasoning_auto_summary: bool,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
Ok(AnthropicMessagesRequest {
|
||||
messages: sanitize_anthropic_messages(request.messages),
|
||||
metadata: request
|
||||
.metadata
|
||||
.as_ref()
|
||||
.map(validate_anthropic_api_metadata)
|
||||
.transpose()?,
|
||||
thinking: with_reasoning_auto_summary(request.thinking, reasoning_auto_summary),
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
fn sanitize_anthropic_messages(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
|
||||
strip_provider_specific_fields(flatten_unencrypted_web_search_results(
|
||||
sanitize_tool_use_ids(strip_empty_content_blocks(messages)),
|
||||
))
|
||||
}
|
||||
|
||||
fn validate_anthropic_api_metadata(metadata: &Value) -> Result<Value, Error> {
|
||||
let Value::Object(fields) = metadata else {
|
||||
return Err(Error::InvalidRequest(format!(
|
||||
"metadata must be an object, got {metadata}"
|
||||
)));
|
||||
};
|
||||
match fields.get("user_id") {
|
||||
None | Some(Value::Null) => Ok(json!({})),
|
||||
Some(Value::String(user_id)) => Ok(json!({"user_id": user_id})),
|
||||
Some(other) => Err(Error::InvalidRequest(format!(
|
||||
"metadata.user_id must be a string, got {other}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn with_reasoning_auto_summary(thinking: Option<Value>, enabled: bool) -> Option<Value> {
|
||||
let Some(Value::Object(thinking)) = thinking else {
|
||||
return thinking;
|
||||
};
|
||||
if !enabled || thinking.get("type").and_then(Value::as_str) == Some("disabled") {
|
||||
return Some(Value::Object(thinking));
|
||||
}
|
||||
Some(Value::Object(
|
||||
thinking
|
||||
.into_iter()
|
||||
.filter(|(key, _)| key != "display")
|
||||
.chain([("display".to_string(), json!("summarized"))])
|
||||
.collect(),
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn messages(value: Value) -> Vec<AnthropicMessage> {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
fn request(body: Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(body).unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty_text_next_to_a_tool_use(
|
||||
json!([{"role": "assistant", "content": [
|
||||
{"type": "text", "text": " "},
|
||||
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
|
||||
]}]),
|
||||
json!([{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
|
||||
]}]),
|
||||
)]
|
||||
#[case::cross_provider_tool_ids(
|
||||
json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
|
||||
]),
|
||||
)]
|
||||
#[case::replayed_unencrypted_web_search_results(
|
||||
json!([
|
||||
{"role": "user", "content": "latest litellm version?"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "latest litellm version"}},
|
||||
{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_1", "content": [{
|
||||
"type": "web_search_result",
|
||||
"url": "https://github.com/BerriAI/litellm/releases",
|
||||
"title": "Releases",
|
||||
"page_age": null,
|
||||
"encrypted_content": "",
|
||||
"snippet": "Latest release v1.95.0"
|
||||
}]}
|
||||
]},
|
||||
{"role": "user", "content": "which version?"}
|
||||
]),
|
||||
json!([
|
||||
{"role": "user", "content": "latest litellm version?"},
|
||||
{"role": "assistant", "content": [{
|
||||
"type": "text",
|
||||
"text": "Web search results for 'latest litellm version':\n\nTitle: Releases\nURL: https://github.com/BerriAI/litellm/releases\nSnippet: Latest release v1.95.0"
|
||||
}]},
|
||||
{"role": "user", "content": "which version?"}
|
||||
]),
|
||||
)]
|
||||
#[case::replayed_provider_specific_fields(
|
||||
json!([
|
||||
{"role": "assistant", "content": [{
|
||||
"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"},
|
||||
"provider_specific_fields": {"signature": "sig_abc"}
|
||||
}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"}}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
|
||||
]),
|
||||
)]
|
||||
#[case::ids_are_normalized_before_web_search_results_flatten(
|
||||
json!([
|
||||
{"role": "user", "content": "run it"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "", "signature": "sig"},
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}, "provider_specific_fields": {"x": 1}},
|
||||
{"type": "server_tool_use", "id": "srv.1", "name": "web_search", "input": {"query": "q"}, "provider_specific_fields": {"x": 2}},
|
||||
{"type": "web_search_tool_result", "tool_use_id": "srv.1", "provider_specific_fields": {"x": 3}, "content": [
|
||||
{"type": "web_search_result", "url": "u", "title": "", "encrypted_content": "", "provider_specific_fields": {"x": 4}}
|
||||
]}
|
||||
]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": " "}]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "user", "content": "run it"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}},
|
||||
{"type": "server_tool_use", "id": "srv_1", "name": "web_search", "input": {"query": "q"}},
|
||||
{"type": "text", "text": "Web search results:\n\nURL: u"}
|
||||
]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
|
||||
]),
|
||||
)]
|
||||
fn sanitize_anthropic_messages_cleans_replayed_history(
|
||||
#[case] history: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
serde_json::to_value(sanitize_anthropic_messages(messages(history))).unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::keeps_only_user_id(json!({"user_id": "u-1", "trace_id": "internal"}), Ok(json!({"user_id": "u-1"})))]
|
||||
#[case::null_user_id(json!({"user_id": null, "trace_id": "internal"}), Ok(json!({})))]
|
||||
#[case::no_user_id(json!({"trace_id": "internal"}), Ok(json!({})))]
|
||||
#[case::empty(json!({}), Ok(json!({})))]
|
||||
#[case::numeric_user_id(
|
||||
json!({"user_id": 123}),
|
||||
Err(Error::InvalidRequest("metadata.user_id must be a string, got 123".to_string())),
|
||||
)]
|
||||
#[case::boolean_user_id(
|
||||
json!({"user_id": true}),
|
||||
Err(Error::InvalidRequest("metadata.user_id must be a string, got true".to_string())),
|
||||
)]
|
||||
#[case::not_an_object(
|
||||
json!(["u-1"]),
|
||||
Err(Error::InvalidRequest(r#"metadata must be an object, got ["u-1"]"#.to_string())),
|
||||
)]
|
||||
fn validate_anthropic_api_metadata_passes_only_a_string_user_id(
|
||||
#[case] metadata: Value,
|
||||
#[case] expected: Result<Value, Error>,
|
||||
) {
|
||||
assert_eq!(validate_anthropic_api_metadata(&metadata), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::adaptive(
|
||||
Some(json!({"type": "adaptive", "budget_tokens": 5000})),
|
||||
true,
|
||||
Some(json!({"type": "adaptive", "budget_tokens": 5000, "display": "summarized"})),
|
||||
)]
|
||||
#[case::enabled(
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000})),
|
||||
true,
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
|
||||
)]
|
||||
#[case::no_type(Some(json!({})), true, Some(json!({"display": "summarized"})))]
|
||||
#[case::display_omitted_is_overridden(
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "omitted"})),
|
||||
true,
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
|
||||
)]
|
||||
#[case::display_summarized_is_kept(
|
||||
Some(json!({"type": "enabled", "display": "summarized"})),
|
||||
true,
|
||||
Some(json!({"type": "enabled", "display": "summarized"})),
|
||||
)]
|
||||
#[case::disabled_thinking(Some(json!({"type": "disabled"})), true, Some(json!({"type": "disabled"})))]
|
||||
#[case::flag_off(
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000})),
|
||||
false,
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000})),
|
||||
)]
|
||||
#[case::flag_off_keeps_callers_display(
|
||||
Some(json!({"type": "enabled", "display": "omitted"})),
|
||||
false,
|
||||
Some(json!({"type": "enabled", "display": "omitted"})),
|
||||
)]
|
||||
#[case::no_thinking(None, true, None)]
|
||||
#[case::non_object_thinking(Some(json!("enabled")), true, Some(json!("enabled")))]
|
||||
fn reasoning_auto_summary_marks_active_thinking_as_summarized(
|
||||
#[case] thinking: Option<Value>,
|
||||
#[case] enabled: bool,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
assert_eq!(with_reasoning_auto_summary(thinking, enabled), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn shaping_cleans_messages_metadata_and_thinking() {
|
||||
let sanitized = shape_anthropic_messages_request(
|
||||
request(json!({
|
||||
"model": "m",
|
||||
"messages": [{"role": "assistant", "content": [
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}
|
||||
]}],
|
||||
"metadata": {"user_id": "u", "trace_id": "t"},
|
||||
"thinking": {"type": "enabled", "budget_tokens": 1024},
|
||||
"safeguards": [{"type": "dangerous_tool_use"}]
|
||||
})),
|
||||
true,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(sanitized).unwrap(),
|
||||
json!({
|
||||
"model": "m",
|
||||
"messages": [{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}
|
||||
]}],
|
||||
"metadata": {"user_id": "u"},
|
||||
"thinking": {"type": "enabled", "budget_tokens": 1024, "display": "summarized"},
|
||||
"safeguards": [{"type": "dangerous_tool_use"}]
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,643 @@
|
|||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
anthropic::{
|
||||
ANTHROPIC_OAUTH_TOKEN_PREFIX,
|
||||
common_utils::{
|
||||
ANTHROPIC_OAUTH_BETA_HEADER, beta, has_advisor_tool, is_anthropic_oauth_key,
|
||||
is_tool_search_used, join_beta_values, requires_native_compaction_beta,
|
||||
split_beta_values,
|
||||
},
|
||||
},
|
||||
base_llm::anthropic_messages::transformation::Headers,
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
|
||||
const BETA_HEADER: &str = "anthropic-beta";
|
||||
const AUTHORIZATION: &str = "authorization";
|
||||
const API_KEY_HEADER: &str = "x-api-key";
|
||||
const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access";
|
||||
|
||||
fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
|
||||
headers
|
||||
.iter()
|
||||
.find(|(header, _)| header.eq_ignore_ascii_case(name))
|
||||
.map(|(_, value)| value.as_str())
|
||||
}
|
||||
|
||||
fn without(headers: Headers, names: &[&str]) -> Headers {
|
||||
headers
|
||||
.into_iter()
|
||||
.filter(|(header, _)| !names.iter().any(|name| header.eq_ignore_ascii_case(name)))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn existing_betas(headers: &[(String, String)]) -> impl Iterator<Item = String> + '_ {
|
||||
headers
|
||||
.iter()
|
||||
.filter(|(header, _)| header.eq_ignore_ascii_case(BETA_HEADER))
|
||||
.flat_map(|(_, value)| split_beta_values(Some(value)))
|
||||
}
|
||||
|
||||
fn with_oauth_bearer(headers: Headers, bearer: String) -> Headers {
|
||||
let beta =
|
||||
join_beta_values(existing_betas(&headers).chain([ANTHROPIC_OAUTH_BETA_HEADER.to_string()]));
|
||||
without(headers, &[API_KEY_HEADER, AUTHORIZATION, BETA_HEADER])
|
||||
.into_iter()
|
||||
.chain([
|
||||
(AUTHORIZATION.to_string(), bearer),
|
||||
(BETA_HEADER.to_string(), beta),
|
||||
(DIRECT_BROWSER_ACCESS_HEADER.to_string(), "true".to_string()),
|
||||
])
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
value.map(str::trim).filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
pub fn authenticate(
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Headers, litellm_auth::Error> {
|
||||
if let Some(forwarded) = header_value(&headers, AUTHORIZATION)
|
||||
&& forwarded
|
||||
.strip_prefix("Bearer ")
|
||||
.is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX))
|
||||
{
|
||||
let bearer = forwarded.to_string();
|
||||
return Ok(with_oauth_bearer(headers, bearer));
|
||||
}
|
||||
if let Some(key) = api_key.filter(|key| key.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) {
|
||||
return Ok(with_oauth_bearer(headers, format!("Bearer {key}")));
|
||||
}
|
||||
if header_value(&headers, API_KEY_HEADER).is_some()
|
||||
|| header_value(&headers, AUTHORIZATION).is_some()
|
||||
{
|
||||
return Ok(headers);
|
||||
}
|
||||
let resolved_key = non_empty(api_key)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty()));
|
||||
let auth = match resolved_key {
|
||||
Some(key) if is_anthropic_oauth_key(&key) => {
|
||||
(AUTHORIZATION.to_string(), format!("Bearer {key}"))
|
||||
}
|
||||
Some(key) => (API_KEY_HEADER.to_string(), key),
|
||||
None => match env_lookup(ANTHROPIC_AUTH_TOKEN_ENV).filter(|value| !value.trim().is_empty())
|
||||
{
|
||||
Some(token) => (AUTHORIZATION.to_string(), format!("Bearer {token}")),
|
||||
None => {
|
||||
return Err(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
});
|
||||
}
|
||||
},
|
||||
};
|
||||
Ok(headers.into_iter().chain([auth]).collect())
|
||||
}
|
||||
|
||||
fn context_management_betas(
|
||||
context_management: Option<&Value>,
|
||||
) -> impl Iterator<Item = &'static str> {
|
||||
let edits = context_management
|
||||
.and_then(|value| value.get("edits"))
|
||||
.and_then(Value::as_array)
|
||||
.map(Vec::as_slice)
|
||||
.unwrap_or(&[]);
|
||||
let (compact, other) = edits.iter().fold((false, false), |(compact, other), edit| {
|
||||
match edit.get("type").and_then(Value::as_str) {
|
||||
Some("compact_20260112") => (true, other),
|
||||
_ => (compact, true),
|
||||
}
|
||||
});
|
||||
compact
|
||||
.then_some(beta::COMPACT_2026_01_12)
|
||||
.into_iter()
|
||||
.chain(other.then_some(beta::CONTEXT_MANAGEMENT_2025_06_27))
|
||||
}
|
||||
|
||||
fn uses_structured_output(request: &AnthropicMessagesRequest) -> bool {
|
||||
request.output_format.is_some()
|
||||
|| request
|
||||
.output_config
|
||||
.as_ref()
|
||||
.and_then(|config| config.get("format"))
|
||||
.is_some_and(|format| !format.is_null())
|
||||
}
|
||||
|
||||
fn messages_carry_output_config(request: &AnthropicMessagesRequest) -> bool {
|
||||
request
|
||||
.messages
|
||||
.iter()
|
||||
.any(|message| message.extra.contains_key("output_config"))
|
||||
}
|
||||
|
||||
pub fn feature_betas(request: &AnthropicMessagesRequest) -> Vec<&'static str> {
|
||||
let tools = request.tools.as_deref();
|
||||
[
|
||||
requires_native_compaction_beta(request.compaction.as_ref(), &request.messages)
|
||||
.then_some(beta::COMPACT_2026_09_04),
|
||||
uses_structured_output(request).then_some(beta::STRUCTURED_OUTPUT),
|
||||
(request.speed.as_deref() == Some("fast")).then_some(beta::FAST_MODE_2026_02_01),
|
||||
messages_carry_output_config(request).then_some(beta::PER_TURN_CONTROL_2026_07_01),
|
||||
has_advisor_tool(tools).then_some(beta::ADVISOR_TOOL_2026_03_01),
|
||||
is_tool_search_used(tools).then_some(beta::ADVANCED_TOOL_USE_2025_11_20),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(context_management_betas(
|
||||
request.context_management.as_ref(),
|
||||
))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn with_feature_betas(headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
let existing = existing_betas(&headers).collect::<Vec<_>>();
|
||||
let features = feature_betas(request);
|
||||
if existing.is_empty() && features.is_empty() {
|
||||
return headers;
|
||||
}
|
||||
let merged = join_beta_values(
|
||||
existing
|
||||
.into_iter()
|
||||
.chain(features.into_iter().map(str::to_string)),
|
||||
);
|
||||
without(headers, &[BETA_HEADER])
|
||||
.into_iter()
|
||||
.chain([(BETA_HEADER.to_string(), merged)])
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
const OAUTH_TOKEN: &str = "sk-ant-oat01-token";
|
||||
const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token";
|
||||
const REGULAR_KEY: &str = "sk-ant-api03-regular";
|
||||
const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true");
|
||||
|
||||
type Env = &'static [(&'static str, &'static str)];
|
||||
|
||||
fn request(fields: Value) -> AnthropicMessagesRequest {
|
||||
let mut body =
|
||||
json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]});
|
||||
body.as_object_mut()
|
||||
.unwrap()
|
||||
.extend(fields.as_object().unwrap().clone());
|
||||
serde_json::from_value(body).unwrap()
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn betas(values: &[&str]) -> String {
|
||||
values.join(",")
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn no_env() -> Env {
|
||||
&[]
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn full_env() -> Env {
|
||||
&[
|
||||
("ANTHROPIC_API_KEY", "sk-env"),
|
||||
("ANTHROPIC_AUTH_TOKEN", "env-token"),
|
||||
]
|
||||
}
|
||||
|
||||
fn authenticate_with(
|
||||
forwarded: &[(&str, &str)],
|
||||
api_key: Option<&str>,
|
||||
env: Env,
|
||||
) -> Result<Headers, litellm_auth::Error> {
|
||||
let lookup = |name: &str| {
|
||||
env.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
};
|
||||
authenticate(headers(forwarded), api_key, &lookup)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_bearer_drops_forwarded_and_deployment_keys(
|
||||
&[("X-Api-Key", REGULAR_KEY), ("Authorization", OAUTH_BEARER)],
|
||||
Some(REGULAR_KEY),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::forwarded_bearer_in_uppercase_authorization_header(
|
||||
&[("AUTHORIZATION", OAUTH_BEARER)],
|
||||
None,
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::forwarded_bearer_keeps_unrelated_headers_in_place(
|
||||
&[("anthropic-version", "2023-06-01"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
OAUTH_BEARER,
|
||||
&[("anthropic-version", "2023-06-01")],
|
||||
)]
|
||||
#[case::forwarded_bearer_wins_over_an_oauth_api_key(
|
||||
&[("authorization", OAUTH_BEARER)],
|
||||
Some("sk-ant-oat01-deployment"),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::api_key_authenticates_as_a_bearer(&[], Some(OAUTH_TOKEN), OAUTH_BEARER, &[])]
|
||||
#[case::api_key_removes_a_forwarded_x_api_key(
|
||||
&[("x-api-key", OAUTH_TOKEN)],
|
||||
Some(OAUTH_TOKEN),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::api_key_replaces_a_forwarded_non_oauth_bearer(
|
||||
&[("Authorization", "Bearer some-proxy-token")],
|
||||
Some(OAUTH_TOKEN),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
fn oauth_token_is_the_whole_credential(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] expected_bearer: &str,
|
||||
#[case] kept: &[(&str, &str)],
|
||||
full_env: Env,
|
||||
) {
|
||||
let expected = kept
|
||||
.iter()
|
||||
.copied()
|
||||
.chain([
|
||||
("authorization", expected_bearer),
|
||||
("anthropic-beta", ANTHROPIC_OAUTH_BETA_HEADER),
|
||||
BROWSER_ACCESS,
|
||||
])
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, full_env).unwrap(),
|
||||
headers(&expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_bearer_merges_a_differently_cased_beta_header(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::forwarded_bearer_dedupes_an_existing_oauth_beta(
|
||||
&[("anthropic-beta", "web-search-2025-03-05, oauth-2025-04-20"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::api_key_merges_the_existing_beta_header(
|
||||
&[("anthropic-beta", " web-search-2025-03-05 ,")],
|
||||
Some(OAUTH_TOKEN),
|
||||
)]
|
||||
#[case::forwarded_bearer_unions_every_beta_header_casing(
|
||||
&[("anthropic-beta", "oauth-2025-04-20"), ("ANTHROPIC-BETA", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
fn oauth_beta_merges_into_existing_betas(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
no_env: Env,
|
||||
) {
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, no_env).unwrap(),
|
||||
headers(&[
|
||||
("authorization", OAUTH_BEARER),
|
||||
(
|
||||
"anthropic-beta",
|
||||
&betas(&[ANTHROPIC_OAUTH_BETA_HEADER, "web-search-2025-03-05"])
|
||||
),
|
||||
BROWSER_ACCESS,
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::x_api_key_over_the_deployment_key(&[("x-api-key", "caller-key")], Some("sk-other"))]
|
||||
#[case::uppercase_x_api_key(&[("X-API-KEY", "caller-key")], None)]
|
||||
#[case::non_oauth_bearer(&[("Authorization", "Bearer some-proxy-token")], None)]
|
||||
#[case::non_oauth_bearer_over_a_regular_api_key(
|
||||
&[("authorization", "Bearer sk-ant-api03-forwarded")],
|
||||
Some(REGULAR_KEY),
|
||||
)]
|
||||
#[case::oauth_token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)]
|
||||
#[case::oauth_token_behind_a_lowercase_bearer_scheme(
|
||||
&[("authorization", "bearer sk-ant-oat01-token")],
|
||||
None,
|
||||
)]
|
||||
fn forwarded_auth_header_is_kept_untouched(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
full_env: Env,
|
||||
) {
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, full_env).unwrap(),
|
||||
headers(forwarded)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key_param(Some("sk-param"), &[], ("x-api-key", "sk-param"))]
|
||||
#[case::api_key_param_over_env_key_and_auth_token(
|
||||
Some("sk-param"),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("x-api-key", "sk-param"),
|
||||
)]
|
||||
#[case::env_key_without_a_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
|
||||
#[case::env_key_when_the_param_is_empty(Some(""), &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
|
||||
#[case::env_key_when_the_param_is_whitespace(
|
||||
Some(" "),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env")],
|
||||
("x-api-key", "sk-env"),
|
||||
)]
|
||||
#[case::env_key_over_auth_token(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("x-api-key", "sk-env"),
|
||||
)]
|
||||
#[case::auth_token_as_a_bearer(
|
||||
None,
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("authorization", "Bearer env-token"),
|
||||
)]
|
||||
#[case::auth_token_when_the_env_key_is_whitespace(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("authorization", "Bearer env-token"),
|
||||
)]
|
||||
#[case::oauth_env_key_as_a_plain_bearer(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-ant-oat01-env")],
|
||||
("authorization", "Bearer sk-ant-oat01-env"),
|
||||
)]
|
||||
fn credential_is_resolved_after_the_existing_headers(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] env: Env,
|
||||
#[case] expected: (&str, &str),
|
||||
) {
|
||||
let forwarded = [("anthropic-beta", "web-search-2025-03-05")];
|
||||
assert_eq!(
|
||||
authenticate_with(&forwarded, api_key, env).unwrap(),
|
||||
headers(&[forwarded[0], expected])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_credentials(&[], None, &[])]
|
||||
#[case::empty_api_key(&[], Some(""), &[])]
|
||||
#[case::whitespace_only_env_values(
|
||||
&[],
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")],
|
||||
)]
|
||||
#[case::unrelated_forwarded_headers(&[("anthropic-beta", "web-search-2025-03-05")], None, &[])]
|
||||
fn missing_credentials_are_an_auth_error(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] env: Env,
|
||||
) {
|
||||
assert!(matches!(
|
||||
authenticate_with(forwarded, api_key, env),
|
||||
Err(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
})
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_features(json!({}), &[])]
|
||||
#[case::output_format(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
|
||||
#[case::null_output_format(json!({"output_format": null}), &[])]
|
||||
#[case::output_config_format(
|
||||
json!({"output_config": {"format": {"type": "json_schema"}, "effort": "xhigh"}}),
|
||||
&[beta::STRUCTURED_OUTPUT]
|
||||
)]
|
||||
#[case::null_output_config_format(json!({"output_config": {"format": null}}), &[])]
|
||||
#[case::top_level_output_config_without_format(json!({"output_config": {"effort": "high"}}), &[])]
|
||||
#[case::fast_speed(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
|
||||
#[case::standard_speed(json!({"speed": "standard"}), &[])]
|
||||
#[case::compaction_param(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::empty_compaction_param(json!({"compaction": {}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::signed_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": "sig"}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&[beta::COMPACT_2026_09_04]
|
||||
)]
|
||||
#[case::unsigned_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": ""}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&[]
|
||||
)]
|
||||
#[case::advisor_tool(
|
||||
json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"}]}),
|
||||
&[beta::ADVISOR_TOOL_2026_03_01]
|
||||
)]
|
||||
#[case::no_tools(json!({"tools": []}), &[])]
|
||||
#[case::regex_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_regex_20251119"}]}),
|
||||
&[beta::ADVANCED_TOOL_USE_2025_11_20]
|
||||
)]
|
||||
#[case::bm25_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_bm25_20251119"}]}),
|
||||
&[beta::ADVANCED_TOOL_USE_2025_11_20]
|
||||
)]
|
||||
#[case::unrelated_server_tool(json!({"tools": [{"type": "web_search_20250305", "name": "web_search"}]}), &[])]
|
||||
#[case::only_compact_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
|
||||
&[beta::COMPACT_2026_01_12]
|
||||
)]
|
||||
#[case::only_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}}]}}),
|
||||
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
)]
|
||||
#[case::compact_and_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_tool_uses_20250919"}]}}),
|
||||
&[beta::COMPACT_2026_01_12, beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
)]
|
||||
#[case::edit_without_a_type(json!({"context_management": {"edits": [{}]}}), &[beta::CONTEXT_MANAGEMENT_2025_06_27])]
|
||||
#[case::empty_edits(json!({"context_management": {"edits": []}}), &[])]
|
||||
#[case::context_management_without_edits(json!({"context_management": {}}), &[])]
|
||||
#[case::per_message_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
)]
|
||||
#[case::per_message_null_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": null}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
)]
|
||||
fn feature_betas_follow_the_request(#[case] fields: Value, #[case] expected: &[&str]) {
|
||||
assert_eq!(feature_betas(&request(fields)), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_betas(&[("x-api-key", "k"), ("anthropic-version", "2023-06-01")], json!({}))]
|
||||
#[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], json!({}))]
|
||||
fn headers_without_any_beta_value_are_untouched(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] fields: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
with_feature_betas(headers(input), &request(fields)),
|
||||
headers(input)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::feature_beta_is_appended(
|
||||
&[("x-api-key", "k")],
|
||||
json!({"speed": "fast"}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
)]
|
||||
#[case::existing_betas_are_normalized_without_features(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05, interleaved-thinking-2025-05-14 ,web-search-2025-03-05"), ("x-api-key", "k")],
|
||||
json!({}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "interleaved-thinking-2025-05-14,web-search-2025-03-05")],
|
||||
)]
|
||||
#[case::existing_advisor_beta_is_kept_without_an_advisor_tool(
|
||||
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
|
||||
json!({"tools": []}),
|
||||
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
|
||||
)]
|
||||
#[case::feature_already_sent_is_not_duplicated(
|
||||
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
json!({"speed": "fast"}),
|
||||
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
)]
|
||||
fn feature_betas_merge_into_the_headers(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] fields: Value,
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
assert_eq!(
|
||||
with_feature_betas(headers(input), &request(fields)),
|
||||
headers(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn differently_cased_beta_header_is_replaced_by_one_sorted_header() {
|
||||
let merged = with_feature_betas(
|
||||
headers(&[("Anthropic-Beta", "interleaved-thinking-2025-05-14")]),
|
||||
&request(
|
||||
json!({"messages": [{"role": "system", "content": "env", "output_config": {"effort": "low"}}]}),
|
||||
),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::PER_TURN_CONTROL_2026_07_01
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_beta_header_casing_is_unioned_into_one_header() {
|
||||
let merged = with_feature_betas(
|
||||
headers(&[
|
||||
("anthropic-beta", "interleaved-thinking-2025-05-14"),
|
||||
("Anthropic-Beta", "web-search-2025-03-05"),
|
||||
]),
|
||||
&request(json!({"speed": "fast"})),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
beta::FAST_MODE_2026_02_01,
|
||||
"interleaved-thinking-2025-05-14",
|
||||
"web-search-2025-03-05"
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_client_betas_survive_alongside_the_added_one() {
|
||||
let client_betas = [
|
||||
"claude-code-20250219",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
"effort-2025-11-24",
|
||||
];
|
||||
let merged = with_feature_betas(
|
||||
headers(&[("anthropic-beta", &betas(&client_betas))]),
|
||||
&request(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
"claude-code-20250219",
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
"effort-2025-11-24",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_feature_merges_with_the_oauth_beta_sorted_and_last() {
|
||||
let oauth_headers = authenticate_with(&[], Some(OAUTH_TOKEN), &[]).unwrap();
|
||||
let all_features = request(json!({
|
||||
"compaction": {"enabled": true},
|
||||
"output_format": {"type": "json_schema"},
|
||||
"speed": "fast",
|
||||
"tools": [{"type": "advisor_20260301"}, {"type": "tool_search_tool_bm25_20251119"}],
|
||||
"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_thinking_20251015"}]},
|
||||
"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}],
|
||||
}));
|
||||
assert_eq!(
|
||||
with_feature_betas(oauth_headers, &all_features),
|
||||
headers(&[
|
||||
("authorization", OAUTH_BEARER),
|
||||
BROWSER_ACCESS,
|
||||
(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
beta::ADVANCED_TOOL_USE_2025_11_20,
|
||||
beta::ADVISOR_TOOL_2026_03_01,
|
||||
beta::COMPACT_2026_01_12,
|
||||
beta::COMPACT_2026_09_04,
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
beta::FAST_MODE_2026_02_01,
|
||||
ANTHROPIC_OAUTH_BETA_HEADER,
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
beta::STRUCTURED_OUTPUT,
|
||||
])
|
||||
),
|
||||
])
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,2 +1,5 @@
|
|||
pub mod handler;
|
||||
pub mod headers;
|
||||
pub mod streaming_iterator;
|
||||
pub mod thinking;
|
||||
pub mod transformation;
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,9 +1,28 @@
|
|||
use crate::base_llm::{
|
||||
anthropic_messages::transformation::BaseAnthropicMessagesConfig, chat::transformation::Error,
|
||||
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::{
|
||||
headers::{authenticate, with_feature_betas},
|
||||
thinking::{ThinkingBudgets, ThinkingContext, translate_thinking},
|
||||
};
|
||||
use crate::{
|
||||
anthropic::common_utils::{
|
||||
AnthropicModelCapabilities, has_advisor_tool, strip_advisor_blocks,
|
||||
strip_encrypted_reasoning_blocks,
|
||||
},
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, Headers, MessagesTransformContext,
|
||||
},
|
||||
chat::transformation::Error,
|
||||
},
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
|
||||
const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE";
|
||||
const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL";
|
||||
const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com";
|
||||
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
|
||||
|
||||
|
|
@ -11,6 +30,26 @@ pub struct AnthropicMessagesConfig;
|
|||
|
||||
pub const ANTHROPIC_MESSAGES_CONFIG: AnthropicMessagesConfig = AnthropicMessagesConfig;
|
||||
|
||||
impl MessagesTransformContext {
|
||||
pub fn new(capabilities: AnthropicModelCapabilities, drop_params: bool) -> Self {
|
||||
Self::with_lookup(capabilities, drop_params, &ProcessEnvironment)
|
||||
}
|
||||
|
||||
pub fn with_lookup(
|
||||
capabilities: AnthropicModelCapabilities,
|
||||
drop_params: bool,
|
||||
env: &impl Lookup,
|
||||
) -> Self {
|
||||
Self {
|
||||
thinking: ThinkingContext {
|
||||
capabilities,
|
||||
budgets: ThinkingBudgets::from_lookup(env),
|
||||
},
|
||||
drop_params,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
|
|
@ -21,6 +60,35 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
|||
Ok(complete_anthropic_url(api_base, env_lookup))
|
||||
}
|
||||
|
||||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
if request.max_tokens.is_none() {
|
||||
return Err(Error::InvalidRequest(
|
||||
"max_tokens is required for Anthropic /v1/messages API".to_string(),
|
||||
));
|
||||
}
|
||||
let request = drop_unsupported_params(request, context)?;
|
||||
let request = translate_thinking(request, &context.thinking)?;
|
||||
let context_management = request
|
||||
.context_management
|
||||
.as_ref()
|
||||
.and_then(map_openai_context_management_to_anthropic)
|
||||
.or_else(|| request.context_management.clone());
|
||||
let messages = if has_advisor_tool(request.tools.as_deref()) {
|
||||
request.messages
|
||||
} else {
|
||||
strip_advisor_blocks(request.messages)
|
||||
};
|
||||
Ok(AnthropicMessagesRequest {
|
||||
messages: strip_encrypted_reasoning_blocks(messages),
|
||||
context_management,
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
|
|
@ -28,6 +96,113 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
|||
) -> Result<String, Error> {
|
||||
resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from)
|
||||
}
|
||||
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[
|
||||
ANTHROPIC_API_KEY_ENV,
|
||||
ANTHROPIC_AUTH_TOKEN_ENV,
|
||||
ANTHROPIC_API_BASE_ENV,
|
||||
ANTHROPIC_BASE_URL_ENV,
|
||||
]
|
||||
}
|
||||
|
||||
fn authenticate(
|
||||
&self,
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Headers, Error> {
|
||||
authenticate(headers, api_key, env_lookup).map_err(Error::from)
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
with_feature_betas(headers, request)
|
||||
}
|
||||
}
|
||||
|
||||
fn unsupported_param(model: &str, param: &str, value: &str, hint: &str) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"{model} does not support {param}={value}. {hint}To drop unsupported params, set `litellm.drop_params = True`."
|
||||
))
|
||||
}
|
||||
|
||||
fn drop_unsupported_params(
|
||||
request: AnthropicMessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
let capabilities = &context.thinking.capabilities;
|
||||
let model = request.model.clone();
|
||||
let reject = |param: &str, value: String, hint: &str| -> Result<(), Error> {
|
||||
if context.drop_params {
|
||||
return Ok(());
|
||||
}
|
||||
Err(unsupported_param(&model, param, &value, hint))
|
||||
};
|
||||
let speed = match request.speed.as_deref() {
|
||||
Some(speed) if !capabilities.supports_speed => {
|
||||
reject("speed", format!("'{speed}'"), "")?;
|
||||
None
|
||||
}
|
||||
_ => request.speed.clone(),
|
||||
};
|
||||
if capabilities.supports_sampling_params {
|
||||
return Ok(AnthropicMessagesRequest { speed, ..request });
|
||||
}
|
||||
let temperature = match request.temperature {
|
||||
Some(temperature) if temperature != 1.0 => {
|
||||
reject(
|
||||
"temperature",
|
||||
json!(temperature).to_string(),
|
||||
"Only temperature=1 is supported. ",
|
||||
)?;
|
||||
None
|
||||
}
|
||||
temperature => temperature,
|
||||
};
|
||||
if let Some(top_p) = request.top_p {
|
||||
reject("top_p", json!(top_p).to_string(), "")?;
|
||||
}
|
||||
if let Some(top_k) = request.top_k {
|
||||
reject("top_k", json!(top_k).to_string(), "")?;
|
||||
}
|
||||
Ok(AnthropicMessagesRequest {
|
||||
speed,
|
||||
temperature,
|
||||
top_p: None,
|
||||
top_k: None,
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
pub fn map_openai_context_management_to_anthropic(context_management: &Value) -> Option<Value> {
|
||||
match context_management {
|
||||
Value::Object(edits) if edits.contains_key("edits") => Some(context_management.clone()),
|
||||
Value::Array(entries) => {
|
||||
let edits: Vec<Value> = entries
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.filter(|entry| entry.get("type").and_then(Value::as_str) == Some("compaction"))
|
||||
.map(|entry| {
|
||||
let trigger = entry.get("compact_threshold").and_then(Value::as_f64).map(
|
||||
|threshold| json!({"type": "input_tokens", "value": threshold as i64}),
|
||||
);
|
||||
let passthrough = entry
|
||||
.iter()
|
||||
.filter(|(key, _)| !matches!(key.as_str(), "type" | "compact_threshold"))
|
||||
.map(|(key, value)| (key.clone(), value.clone()));
|
||||
Value::Object(
|
||||
[("type".to_string(), json!("compact_20260112"))]
|
||||
.into_iter()
|
||||
.chain(trigger.map(|trigger| ("trigger".to_string(), trigger)))
|
||||
.chain(passthrough)
|
||||
.collect::<Map<String, Value>>(),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
(!edits.is_empty()).then(|| json!({"edits": edits}))
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
|
|
@ -64,70 +239,619 @@ pub fn resolve_anthropic_api_base(
|
|||
api_base: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
let env = |name: &str| env_lookup(name).filter(|value| !value.trim().is_empty());
|
||||
non_empty(api_base)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(ANTHROPIC_API_BASE_ENV).filter(|value| !value.trim().is_empty()))
|
||||
.or_else(|| env(ANTHROPIC_API_BASE_ENV))
|
||||
.or_else(|| env(ANTHROPIC_BASE_URL_ENV))
|
||||
.unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::process::Command;
|
||||
|
||||
use rstest::{fixture, rstest};
|
||||
|
||||
use super::*;
|
||||
use crate::anthropic::common_utils::{ENCRYPTED_REASONING_SIGNATURE_PREFIX, beta};
|
||||
|
||||
#[test]
|
||||
fn url_defaults_to_public_anthropic_endpoint() {
|
||||
type Env = &'static [(&'static str, &'static str)];
|
||||
|
||||
const BOTH_BASE_ENVS: Env = &[
|
||||
(ANTHROPIC_API_BASE_ENV, "https://api-base.example.com"),
|
||||
(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com"),
|
||||
];
|
||||
const API_KEY_ENV: Env = &[(ANTHROPIC_API_KEY_ENV, "sk-env")];
|
||||
const MISSING_API_KEY: &str =
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable";
|
||||
const LOW_BUDGET_ENV: &str = "DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET";
|
||||
const PROCESS_ENV_PROBE: &str = "LITELLM_MESSAGES_TRANSFORM_CONTEXT_PROBE";
|
||||
|
||||
fn merged(base: Value, fields: Value) -> Value {
|
||||
Value::Object(
|
||||
base.as_object()
|
||||
.unwrap()
|
||||
.clone()
|
||||
.into_iter()
|
||||
.chain(fields.as_object().unwrap().clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
fn body(fields: Value) -> Value {
|
||||
merged(
|
||||
json!({
|
||||
"model": "claude",
|
||||
"max_tokens": 1024,
|
||||
"messages": [{"role": "user", "content": "Hello"}]
|
||||
}),
|
||||
fields,
|
||||
)
|
||||
}
|
||||
|
||||
fn request(fields: Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(body(fields)).unwrap()
|
||||
}
|
||||
|
||||
fn no_env(_: &str) -> Option<String> {
|
||||
None
|
||||
}
|
||||
|
||||
fn env(vars: Env) -> impl Fn(&str) -> Option<String> {
|
||||
move |name| {
|
||||
vars.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn transform(
|
||||
fields: Value,
|
||||
capabilities: AnthropicModelCapabilities,
|
||||
drop_params: bool,
|
||||
) -> Result<Value, Error> {
|
||||
ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(
|
||||
request(fields),
|
||||
&MessagesTransformContext::with_lookup(capabilities, drop_params, &no_env),
|
||||
)
|
||||
.map(|transformed| serde_json::to_value(transformed).unwrap())
|
||||
}
|
||||
|
||||
fn invalid(message: &str) -> Result<Value, Error> {
|
||||
Err(Error::InvalidRequest(message.to_string()))
|
||||
}
|
||||
|
||||
fn advisor_history() -> Value {
|
||||
json!([
|
||||
{"role": "user", "content": "Build a worker pool."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "text", "text": "Let me consult the advisor."},
|
||||
{"type": "server_tool_use", "id": "srvtoolu_abc123", "name": "advisor", "input": {}},
|
||||
{"type": "advisor_tool_result", "tool_use_id": "srvtoolu_abc123", "content": {"type": "advisor_result", "text": "Use channels."}},
|
||||
{"type": "text", "text": "Here is the implementation."}
|
||||
]}
|
||||
])
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn unmapped() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities::default()
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn sampling_removed() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities {
|
||||
supports_sampling_params: false,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn fast_mode() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities {
|
||||
supports_speed: true,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::alone(json!({"max_tokens": null}))]
|
||||
#[case::ahead_of_the_param_gate(json!({"max_tokens": null, "speed": "fast"}))]
|
||||
fn missing_max_tokens_is_rejected(#[case] fields: Value, unmapped: AnthropicModelCapabilities) {
|
||||
assert_eq!(
|
||||
complete_anthropic_url(None, &|_| None),
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
transform(fields, unmapped, false),
|
||||
invalid("max_tokens is required for Anthropic /v1/messages API")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::sampling_params_on_a_sampling_model(
|
||||
unmapped(),
|
||||
false,
|
||||
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
|
||||
)]
|
||||
#[case::sampling_params_on_a_sampling_model_under_drop_params(
|
||||
unmapped(),
|
||||
true,
|
||||
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
|
||||
)]
|
||||
#[case::unit_temperature_on_a_sampling_removed_model(
|
||||
sampling_removed(),
|
||||
false,
|
||||
json!({"temperature": 1.0})
|
||||
)]
|
||||
#[case::unit_temperature_on_a_sampling_removed_model_under_drop_params(
|
||||
sampling_removed(),
|
||||
true,
|
||||
json!({"temperature": 1.0})
|
||||
)]
|
||||
#[case::speed_on_a_fast_mode_model(fast_mode(), false, json!({"speed": "fast"}))]
|
||||
#[case::speed_on_a_fast_mode_model_under_drop_params(fast_mode(), true, json!({"speed": "fast"}))]
|
||||
#[case::native_context_management_edits(unmapped(), false, json!({"context_management": {"edits": [{
|
||||
"type": "clear_tool_uses_20250919",
|
||||
"trigger": {"type": "input_tokens", "value": 30000},
|
||||
"keep": {"type": "tool_uses", "value": 3},
|
||||
"clear_at_least": {"type": "input_tokens", "value": 5000},
|
||||
"exclude_tools": ["web_search"],
|
||||
"clear_tool_inputs": false
|
||||
}]}}))]
|
||||
#[case::first_party_billing_header_system_block(unmapped(), false, json!({"system": [
|
||||
{"type": "text", "text": "x-anthropic-billing-header: cc_version=1"},
|
||||
{"type": "text", "text": "real system prompt"}
|
||||
]}))]
|
||||
#[case::anthropic_signed_reasoning_history(unmapped(), false, json!({"messages": [
|
||||
{"role": "user", "content": "Solve it."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": "EqQBCkYIAxgCIkA_anthropic_signed"},
|
||||
{"type": "redacted_thinking", "data": "EmwKAhgBEgy_anthropic_minted"},
|
||||
{"type": "text", "text": "The answer."}
|
||||
]}
|
||||
]}))]
|
||||
#[case::advisor_history_alongside_the_advisor_tool(unmapped(), false, json!({
|
||||
"messages": advisor_history(),
|
||||
"tools": [{"type": "advisor_20260301", "name": "advisor"}]
|
||||
}))]
|
||||
fn request_is_forwarded_unchanged(
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] drop_params: bool,
|
||||
#[case] fields: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
transform(fields.clone(), capabilities, drop_params),
|
||||
Ok(body(fields))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::temperature(sampling_removed(), json!({"temperature": 0.3}), json!({}))]
|
||||
#[case::top_p(sampling_removed(), json!({"top_p": 0.9}), json!({}))]
|
||||
#[case::top_k(sampling_removed(), json!({"top_k": 40}), json!({}))]
|
||||
#[case::every_sampling_param_keeping_the_rest(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40, "stream": true}),
|
||||
json!({"stream": true})
|
||||
)]
|
||||
#[case::speed_on_a_sampling_model(
|
||||
unmapped(),
|
||||
json!({"speed": "fast", "temperature": 0.5}),
|
||||
json!({"temperature": 0.5})
|
||||
)]
|
||||
#[case::speed_on_a_sampling_removed_model(
|
||||
sampling_removed(),
|
||||
json!({"speed": "fast", "temperature": 1.0}),
|
||||
json!({"temperature": 1.0})
|
||||
)]
|
||||
fn removed_params_are_dropped_under_drop_params(
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
assert_eq!(transform(fields, capabilities, true), Ok(body(expected)));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::temperature(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.3}),
|
||||
"claude does not support temperature=0.3. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::temperature_just_below_one(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.99}),
|
||||
"claude does not support temperature=0.99. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::whole_number_temperature_keeps_its_decimal(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 2.0}),
|
||||
"claude does not support temperature=2.0. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_p(
|
||||
sampling_removed(),
|
||||
json!({"top_p": 0.9}),
|
||||
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_k(
|
||||
sampling_removed(),
|
||||
json!({"top_k": 5}),
|
||||
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_k_next_to_unit_temperature(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 1.0, "top_k": 5}),
|
||||
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::temperature_ahead_of_top_k(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.5, "top_k": 5}),
|
||||
"claude does not support temperature=0.5. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_p_ahead_of_top_k(
|
||||
sampling_removed(),
|
||||
json!({"top_p": 0.9, "top_k": 5}),
|
||||
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::speed(
|
||||
unmapped(),
|
||||
json!({"speed": "fast"}),
|
||||
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::speed_ahead_of_sampling_params(
|
||||
sampling_removed(),
|
||||
json!({"speed": "fast", "temperature": 0.5}),
|
||||
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
fn removed_params_are_rejected_without_drop_params(
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] message: &str,
|
||||
) {
|
||||
assert_eq!(transform(fields, capabilities, false), invalid(message));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::compaction_threshold(
|
||||
json!([{"type": "compaction", "compact_threshold": 200000}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]}))
|
||||
)]
|
||||
#[case::other_keys_pass_through(
|
||||
json!([{"type": "compaction", "compact_threshold": 150000, "instructions": "Focus on preserving code snippets"}]),
|
||||
Some(json!({"edits": [{
|
||||
"type": "compact_20260112",
|
||||
"trigger": {"type": "input_tokens", "value": 150000},
|
||||
"instructions": "Focus on preserving code snippets"
|
||||
}]}))
|
||||
)]
|
||||
#[case::float_threshold_is_truncated(
|
||||
json!([{"type": "compaction", "compact_threshold": 150000.9}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
|
||||
)]
|
||||
#[case::compaction_without_threshold(
|
||||
json!([{"type": "compaction"}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112"}]}))
|
||||
)]
|
||||
#[case::non_numeric_threshold_is_dropped(
|
||||
json!([{"type": "compaction", "compact_threshold": "150000"}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112"}]}))
|
||||
)]
|
||||
#[case::non_object_entries_are_skipped(
|
||||
json!([42, "compaction", null, [], {"type": "compaction", "compact_threshold": 1000}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}}]}))
|
||||
)]
|
||||
#[case::only_compaction_entries_are_mapped_in_order(
|
||||
json!([
|
||||
{"type": "compaction", "compact_threshold": 1000},
|
||||
{"type": "other", "compact_threshold": 5},
|
||||
{"type": "compaction", "instructions": "second"}
|
||||
]),
|
||||
Some(json!({"edits": [
|
||||
{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}},
|
||||
{"type": "compact_20260112", "instructions": "second"}
|
||||
]}))
|
||||
)]
|
||||
#[case::list_without_compaction(json!([{"type": "other"}]), None)]
|
||||
#[case::empty_list(json!([]), None)]
|
||||
#[case::anthropic_edits_pass_through(
|
||||
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
|
||||
)]
|
||||
#[case::object_without_edits(json!({"type": "compaction"}), None)]
|
||||
#[case::scalar(json!("compaction"), None)]
|
||||
fn openai_context_management_maps_to_anthropic_edits(
|
||||
#[case] context_management: Value,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
assert_eq!(
|
||||
map_openai_context_management_to_anthropic(&context_management),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::openai_list_is_mapped(
|
||||
json!([{"type": "compaction", "compact_threshold": 200000}]),
|
||||
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]})
|
||||
)]
|
||||
#[case::unmappable_list_is_kept(json!([{"type": "other"}]), json!([{"type": "other"}]))]
|
||||
#[case::unmappable_object_is_kept(json!({"type": "other"}), json!({"type": "other"}))]
|
||||
fn context_management_reaches_the_wire(
|
||||
#[case] context_management: Value,
|
||||
#[case] expected: Value,
|
||||
unmapped: AnthropicModelCapabilities,
|
||||
) {
|
||||
assert_eq!(
|
||||
transform(
|
||||
json!({"context_management": context_management}),
|
||||
unmapped,
|
||||
false
|
||||
),
|
||||
Ok(body(json!({"context_management": expected})))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::without_tools(json!({}))]
|
||||
#[case::with_only_other_tools(json!({"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}]}))]
|
||||
fn advisor_history_is_stripped_without_the_advisor_tool(
|
||||
#[case] tools: Value,
|
||||
unmapped: AnthropicModelCapabilities,
|
||||
) {
|
||||
let stripped = json!([
|
||||
{"role": "user", "content": "Build a worker pool."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "text", "text": "Let me consult the advisor."},
|
||||
{"type": "text", "text": "Here is the implementation."}
|
||||
]}
|
||||
]);
|
||||
assert_eq!(
|
||||
transform(
|
||||
merged(tools.clone(), json!({"messages": advisor_history()})),
|
||||
unmapped,
|
||||
false
|
||||
),
|
||||
Ok(body(merged(tools, json!({"messages": stripped}))))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn bridge_minted_reasoning_is_stripped_from_the_wire(unmapped: AnthropicModelCapabilities) {
|
||||
let messages = json!([
|
||||
{"role": "user", "content": "Solve it."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_1")},
|
||||
{"type": "redacted_thinking", "data": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_2")},
|
||||
{"type": "text", "text": "The answer."}
|
||||
]},
|
||||
{"role": "user", "content": "And the next one?"}
|
||||
]);
|
||||
assert_eq!(
|
||||
transform(json!({"messages": messages}), unmapped, false),
|
||||
Ok(body(json!({"messages": [
|
||||
{"role": "user", "content": "Solve it."},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "The answer."}]},
|
||||
{"role": "user", "content": "And the next one?"}
|
||||
]})))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_appends_messages_suffix_to_custom_base() {
|
||||
fn thinking_is_translated_with_the_context_budgets() {
|
||||
let context = MessagesTransformContext::with_lookup(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
&env(&[(LOW_BUDGET_ENV, "2000")]),
|
||||
);
|
||||
let transformed = ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(
|
||||
request(json!({"max_tokens": 4096, "reasoning_effort": "low"})),
|
||||
&context,
|
||||
)
|
||||
.map(|transformed| serde_json::to_value(transformed).unwrap());
|
||||
assert_eq!(
|
||||
complete_anthropic_url(Some("https://proxy.internal"), &|_| None),
|
||||
"https://proxy.internal/v1/messages"
|
||||
transformed,
|
||||
Ok(body(json!({
|
||||
"max_tokens": 4096,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2000}
|
||||
})))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_leaves_complete_messages_endpoint_untouched() {
|
||||
fn new_reads_thinking_budgets_from_the_process_environment() {
|
||||
if std::env::var_os(PROCESS_ENV_PROBE).is_some() {
|
||||
assert_eq!(
|
||||
MessagesTransformContext::new(sampling_removed(), true),
|
||||
MessagesTransformContext {
|
||||
thinking: ThinkingContext {
|
||||
capabilities: sampling_removed(),
|
||||
budgets: ThinkingBudgets {
|
||||
low: 2000,
|
||||
..ThinkingBudgets::default()
|
||||
},
|
||||
},
|
||||
drop_params: true,
|
||||
}
|
||||
);
|
||||
return;
|
||||
}
|
||||
let (_, test_path) = concat!(
|
||||
module_path!(),
|
||||
"::new_reads_thinking_budgets_from_the_process_environment"
|
||||
)
|
||||
.split_once("::")
|
||||
.unwrap();
|
||||
let other_tiers = ["MINIMAL", "MEDIUM", "HIGH", "XHIGH", "MAX"]
|
||||
.map(|tier| format!("DEFAULT_REASONING_EFFORT_{tier}_THINKING_BUDGET"));
|
||||
let output = other_tiers
|
||||
.iter()
|
||||
.fold(
|
||||
Command::new(std::env::current_exe().unwrap()),
|
||||
|mut command, name| {
|
||||
command.env_remove(name);
|
||||
command
|
||||
},
|
||||
)
|
||||
.args([test_path, "--exact"])
|
||||
.env(PROCESS_ENV_PROBE, "1")
|
||||
.env(LOW_BUDGET_ENV, "2000")
|
||||
.output()
|
||||
.unwrap();
|
||||
let stdout = String::from_utf8_lossy(&output.stdout);
|
||||
assert!(
|
||||
output.status.success() && stdout.contains("1 passed"),
|
||||
"{stdout}{}",
|
||||
String::from_utf8_lossy(&output.stderr)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::public_endpoint_by_default(None, &[], "https://api.anthropic.com")]
|
||||
#[case::explicit_api_base_beats_env(
|
||||
Some("https://explicit.example.com"),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::explicit_api_base_is_trimmed(
|
||||
Some(" https://explicit.example.com "),
|
||||
&[],
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_falls_back_to_env(
|
||||
Some(" "),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://api-base.example.com"
|
||||
)]
|
||||
#[case::api_base_env_beats_base_url_env(None, BOTH_BASE_ENVS, "https://api-base.example.com")]
|
||||
#[case::base_url_env_without_api_base_env(
|
||||
None,
|
||||
&[(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_env_falls_back_to_base_url_env(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, " \t "), (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_envs_fall_back_to_public_endpoint(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, ""), (ANTHROPIC_BASE_URL_ENV, " ")],
|
||||
"https://api.anthropic.com"
|
||||
)]
|
||||
fn api_base_resolution(
|
||||
#[case] api_base: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(resolve_anthropic_api_base(api_base, &env(vars)), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::public_endpoint(None, &[], "https://api.anthropic.com/v1/messages")]
|
||||
#[case::base_url_env(
|
||||
None,
|
||||
&[(ANTHROPIC_BASE_URL_ENV, "https://custom.example.com")],
|
||||
"https://custom.example.com/v1/messages"
|
||||
)]
|
||||
#[case::custom_base(Some("https://proxy.internal"), &[], "https://proxy.internal/v1/messages")]
|
||||
#[case::trailing_slash(Some("https://proxy.internal/"), &[], "https://proxy.internal/v1/messages")]
|
||||
#[case::complete_endpoint(
|
||||
Some("https://proxy.internal/v1/messages"),
|
||||
&[],
|
||||
"https://proxy.internal/v1/messages"
|
||||
)]
|
||||
#[case::complete_endpoint_with_trailing_slash(
|
||||
Some("https://proxy.internal/v1/messages/"),
|
||||
&[],
|
||||
"https://proxy.internal/v1/messages"
|
||||
)]
|
||||
fn complete_url_ends_in_the_messages_path(
|
||||
#[case] api_base: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(
|
||||
complete_anthropic_url(Some("https://proxy.internal/v1/messages"), &|_| None),
|
||||
"https://proxy.internal/v1/messages"
|
||||
ANTHROPIC_MESSAGES_CONFIG.get_complete_url(api_base, "claude", &env(vars)),
|
||||
Ok(expected.to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::param_beats_env(Some("sk-param"), API_KEY_ENV, Ok("sk-param"))]
|
||||
#[case::param_is_trimmed(Some(" sk-param "), &[], Ok("sk-param"))]
|
||||
#[case::blank_param_falls_back_to_env(Some(" "), API_KEY_ENV, Ok("sk-env"))]
|
||||
#[case::env_without_param(None, API_KEY_ENV, Ok("sk-env"))]
|
||||
#[case::blank_env_is_missing(None, &[(ANTHROPIC_API_KEY_ENV, " ")], Err(MISSING_API_KEY))]
|
||||
#[case::nothing_is_missing(None, &[], Err(MISSING_API_KEY))]
|
||||
fn api_key_resolution(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: Result<&str, &str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(api_key, &env(vars)).map_err(|error| error.to_string()),
|
||||
expected.map(str::to_string).map_err(str::to_string)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_falls_back_to_env_base() {
|
||||
let with_env = |key: &str| {
|
||||
(key == ANTHROPIC_API_BASE_ENV).then(|| "https://env.anthropic".to_string())
|
||||
};
|
||||
fn config_reports_a_missing_key_as_an_auth_error() {
|
||||
assert_eq!(
|
||||
complete_anthropic_url(Some(" "), &with_env),
|
||||
"https://env.anthropic/v1/messages"
|
||||
ANTHROPIC_MESSAGES_CONFIG.resolve_api_key(None, &no_env),
|
||||
Err(Error::Auth(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_key_prefers_param_then_env_then_errors() {
|
||||
fn config_authenticates_with_the_anthropic_auth_token() {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(Some("sk-param"), &|_| None).unwrap(),
|
||||
"sk-param"
|
||||
ANTHROPIC_MESSAGES_CONFIG.authenticate(
|
||||
vec![],
|
||||
None,
|
||||
&env(&[("ANTHROPIC_AUTH_TOKEN", "auth-token")])
|
||||
),
|
||||
Ok(headers(&[("authorization", "Bearer auth-token")]))
|
||||
);
|
||||
let with_env = |key: &str| (key == ANTHROPIC_API_KEY_ENV).then(|| "sk-env".to_string());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_requests_the_betas_the_request_features_need() {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(Some(" "), &with_env).unwrap(),
|
||||
"sk-env"
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(None, &|_| None)
|
||||
.expect_err("missing key")
|
||||
.to_string(),
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
|
||||
ANTHROPIC_MESSAGES_CONFIG.request_headers(
|
||||
headers(&[("x-api-key", "sk")]),
|
||||
&request(json!({"speed": "fast"}))
|
||||
),
|
||||
headers(&[
|
||||
("x-api-key", "sk"),
|
||||
("anthropic-beta", beta::FAST_MODE_2026_02_01)
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::absent(None, None)]
|
||||
#[case::blank(Some(" \t "), None)]
|
||||
#[case::padded(Some(" value "), Some("value"))]
|
||||
fn non_empty_trims_and_drops_blank_values(
|
||||
#[case] value: Option<&str>,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
assert_eq!(non_empty(value), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_strategy_and_default_headers_match_anthropic() {
|
||||
assert_eq!(
|
||||
|
|
@ -142,4 +866,26 @@ mod tests {
|
|||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn secret_names_cover_every_credential_and_base_lookup() {
|
||||
let requested = std::cell::RefCell::new(Vec::<String>::new());
|
||||
let record = |name: &str| -> Option<String> {
|
||||
requested.borrow_mut().push(name.to_string());
|
||||
None
|
||||
};
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record);
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
|
||||
let requested = requested.into_inner();
|
||||
assert!(!requested.is_empty());
|
||||
let undeclared: Vec<&String> = requested
|
||||
.iter()
|
||||
.filter(|name| {
|
||||
!ANTHROPIC_MESSAGES_CONFIG
|
||||
.secret_names()
|
||||
.contains(&name.as_str())
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(undeclared, Vec::<&String>::new());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
pub mod batches;
|
||||
pub mod chat;
|
||||
pub mod common_utils;
|
||||
pub mod count_tokens;
|
||||
pub mod experimental_pass_through;
|
||||
|
||||
|
|
|
|||
|
|
@ -4,14 +4,15 @@ use litellm_types::llms::anthropic_messages::{
|
|||
},
|
||||
anthropic_response::AnthropicMessagesResponse,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::transformation::{
|
||||
ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty,
|
||||
},
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::{BaseAnthropicMessagesConfig, MessagesAuthStrategy},
|
||||
anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, Headers, MessagesAuthStrategy, MessagesTransformContext,
|
||||
},
|
||||
chat::transformation::Error,
|
||||
},
|
||||
};
|
||||
|
|
@ -21,7 +22,6 @@ const AZURE_API_BASE_ENV: &str = "AZURE_API_BASE";
|
|||
const ANTHROPIC_PATH_SEGMENT: &str = "/anthropic";
|
||||
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
|
||||
const SYSTEM_ROLE: &str = "system";
|
||||
const TEXT_BLOCK_TYPE: &str = "text";
|
||||
|
||||
pub struct AzureAnthropicMessagesConfig {
|
||||
anthropic: AnthropicMessagesConfig,
|
||||
|
|
@ -45,6 +45,7 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
let mut request = fold_system_role_messages(request);
|
||||
if let Some(system) = request.system.as_mut() {
|
||||
|
|
@ -54,7 +55,8 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
.messages
|
||||
.iter_mut()
|
||||
.for_each(strip_scope_from_message);
|
||||
self.anthropic.transform_anthropic_messages_request(request)
|
||||
self.anthropic
|
||||
.transform_anthropic_messages_request(request, context)
|
||||
}
|
||||
|
||||
fn transform_anthropic_messages_response(
|
||||
|
|
@ -74,6 +76,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
resolve_azure_api_key(api_key, env_lookup)
|
||||
}
|
||||
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[AZURE_API_KEY_ENV, AZURE_API_BASE_ENV]
|
||||
}
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
self.anthropic.auth_strategy()
|
||||
}
|
||||
|
|
@ -85,6 +91,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
self.anthropic.default_headers()
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
self.anthropic.request_headers(headers, request)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_azure_api_key(
|
||||
|
|
@ -143,17 +153,7 @@ fn strip_scope_from_message(message: &mut AnthropicMessage) {
|
|||
}
|
||||
|
||||
fn text_content_block(text: String) -> ContentBlock {
|
||||
let extra = Map::from_iter([
|
||||
(
|
||||
"type".to_string(),
|
||||
Value::String(TEXT_BLOCK_TYPE.to_string()),
|
||||
),
|
||||
("text".to_string(), Value::String(text)),
|
||||
]);
|
||||
ContentBlock {
|
||||
cache_control: None,
|
||||
extra,
|
||||
}
|
||||
ContentBlock::text(text)
|
||||
}
|
||||
|
||||
fn content_into_blocks(content: MessageContent) -> Vec<ContentBlock> {
|
||||
|
|
@ -202,6 +202,7 @@ mod tests {
|
|||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
use crate::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
|
||||
fn request_from(value: serde_json::Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(value).expect("valid request")
|
||||
|
|
@ -346,7 +347,7 @@ mod tests {
|
|||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
|
|
@ -373,10 +374,13 @@ mod tests {
|
|||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}));
|
||||
let once = AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms");
|
||||
let twice = AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(once.clone())
|
||||
.transform_anthropic_messages_request(
|
||||
once.clone(),
|
||||
&MessagesTransformContext::default(),
|
||||
)
|
||||
.expect("request transforms");
|
||||
assert_eq!(once, twice);
|
||||
assert_eq!(to_value(once)["system"], json!("plain string system"));
|
||||
|
|
@ -408,9 +412,21 @@ mod tests {
|
|||
"inference_geo": "us",
|
||||
"litellm_metadata": {"trace": "abc"}
|
||||
});
|
||||
let context = MessagesTransformContext::with_lookup(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
supports_legacy_thinking: true,
|
||||
supports_output_config: true,
|
||||
supports_speed: true,
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
&|_: &str| None,
|
||||
);
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request_from(body.clone()))
|
||||
.transform_anthropic_messages_request(request_from(body.clone()), &context)
|
||||
.expect("request transforms"),
|
||||
);
|
||||
assert_eq!(transformed, body);
|
||||
|
|
@ -430,7 +446,7 @@ mod tests {
|
|||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
|
|
@ -460,7 +476,7 @@ mod tests {
|
|||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
|
|
@ -485,9 +501,21 @@ mod tests {
|
|||
{"role": "assistant", "content": "hello"}
|
||||
]
|
||||
});
|
||||
let context = MessagesTransformContext::with_lookup(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
supports_legacy_thinking: true,
|
||||
supports_output_config: true,
|
||||
supports_speed: true,
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
&|_: &str| None,
|
||||
);
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request_from(body.clone()))
|
||||
.transform_anthropic_messages_request(request_from(body.clone()), &context)
|
||||
.expect("request transforms"),
|
||||
);
|
||||
assert_eq!(transformed, body);
|
||||
|
|
@ -500,6 +528,57 @@ mod tests {
|
|||
assert!(err.is_data());
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::compact_context_management_edit(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
|
||||
&[],
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "compact-2026-01-12")]
|
||||
)]
|
||||
#[case::forwarded_beta_merged_with_structured_output(
|
||||
json!({"output_config": {"format": {"type": "json_schema"}}}),
|
||||
&[("anthropic-beta", "web-search-2025-03-05")],
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "structured-outputs-2025-11-13,web-search-2025-03-05")]
|
||||
)]
|
||||
#[case::no_feature_needs_a_beta(json!({}), &[], &[("x-api-key", "k")])]
|
||||
fn request_headers_carry_the_anthropic_feature_betas(
|
||||
#[case] fields: serde_json::Value,
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
let pairs = |pairs: &[(&str, &str)]| -> Vec<(String, String)> {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
};
|
||||
let serde_json::Value::Object(fields) = fields else {
|
||||
panic!("case fields are an object")
|
||||
};
|
||||
let request = request_from(serde_json::Value::Object(
|
||||
[
|
||||
("model".to_string(), json!("claude-sonnet")),
|
||||
("max_tokens".to_string(), json!(16)),
|
||||
(
|
||||
"messages".to_string(),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
),
|
||||
]
|
||||
.into_iter()
|
||||
.chain(fields)
|
||||
.collect(),
|
||||
));
|
||||
assert_eq!(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG.request_headers(
|
||||
pairs(&[("x-api-key", "k")])
|
||||
.into_iter()
|
||||
.chain(pairs(forwarded))
|
||||
.collect(),
|
||||
&request
|
||||
),
|
||||
pairs(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_response_passes_through() {
|
||||
let response: AnthropicMessagesResponse = serde_json::from_value(json!({
|
||||
|
|
@ -521,4 +600,26 @@ mod tests {
|
|||
assert_eq!(value["stop_sequence"], json!(null));
|
||||
assert_eq!(value["content"][0]["text"], json!("hello"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn secret_names_cover_every_credential_and_base_lookup() {
|
||||
let requested = std::cell::RefCell::new(Vec::<String>::new());
|
||||
let record = |name: &str| -> Option<String> {
|
||||
requested.borrow_mut().push(name.to_string());
|
||||
None
|
||||
};
|
||||
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record);
|
||||
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
|
||||
let requested = requested.into_inner();
|
||||
assert!(!requested.is_empty());
|
||||
let undeclared: Vec<&String> = requested
|
||||
.iter()
|
||||
.filter(|name| {
|
||||
!AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.secret_names()
|
||||
.contains(&name.as_str())
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(undeclared, Vec::<&String>::new());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,14 @@
|
|||
use litellm_http::request::{has_bearer_auth, has_header};
|
||||
use litellm_types::llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
};
|
||||
|
||||
use crate::base_llm::chat::transformation::Error;
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::thinking::ThinkingContext,
|
||||
base_llm::chat::transformation::Error,
|
||||
};
|
||||
|
||||
pub type Headers = Vec<(String, String)>;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum MessagesAuthStrategy {
|
||||
|
|
@ -19,6 +25,12 @@ impl MessagesAuthStrategy {
|
|||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
|
||||
pub struct MessagesTransformContext {
|
||||
pub thinking: ThinkingContext,
|
||||
pub drop_params: bool,
|
||||
}
|
||||
|
||||
pub trait BaseAnthropicMessagesConfig: Sync {
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
|
|
@ -30,6 +42,7 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
_context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
Ok(request)
|
||||
}
|
||||
|
|
@ -48,6 +61,8 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error>;
|
||||
|
||||
fn secret_names(&self) -> &'static [&'static str];
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
MessagesAuthStrategy::Header("x-api-key")
|
||||
}
|
||||
|
|
@ -56,10 +71,225 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
false
|
||||
}
|
||||
|
||||
fn authenticate(
|
||||
&self,
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Headers, Error> {
|
||||
let strategy = self.auth_strategy();
|
||||
if has_header(&headers, strategy.header_name())
|
||||
|| (self.accepts_bearer_auth() && has_bearer_auth(&headers))
|
||||
{
|
||||
return Ok(headers);
|
||||
}
|
||||
let api_key = self.resolve_api_key(api_key, env_lookup)?;
|
||||
let auth_header = match strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
("authorization".to_string(), format!("Bearer {api_key}"))
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
Ok(headers.into_iter().chain([auth_header]).collect())
|
||||
}
|
||||
|
||||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
&[
|
||||
("anthropic-version", "2023-06-01"),
|
||||
("content-type", "application/json"),
|
||||
]
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, _request: &AnthropicMessagesRequest) -> Headers {
|
||||
headers
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
const X_API_KEY: MessagesAuthStrategy = MessagesAuthStrategy::Header("x-api-key");
|
||||
|
||||
struct StubConfig {
|
||||
strategy: MessagesAuthStrategy,
|
||||
accepts_bearer: bool,
|
||||
}
|
||||
|
||||
impl BaseAnthropicMessagesConfig for StubConfig {
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[]
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
_api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
api_key
|
||||
.map(str::to_string)
|
||||
.ok_or(Error::MissingField("api_key"))
|
||||
}
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
self.strategy
|
||||
}
|
||||
|
||||
fn accepts_bearer_auth(&self) -> bool {
|
||||
self.accepts_bearer
|
||||
}
|
||||
}
|
||||
|
||||
struct DefaultsConfig;
|
||||
|
||||
impl BaseAnthropicMessagesConfig for DefaultsConfig {
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[]
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
_api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
api_key
|
||||
.map(str::to_string)
|
||||
.ok_or(Error::MissingField("api_key"))
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_config_adds_its_key_next_to_a_forwarded_bearer() {
|
||||
assert_eq!(
|
||||
DefaultsConfig.authenticate(
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
Some("sk"),
|
||||
&|_| None
|
||||
),
|
||||
Ok(headers(&[
|
||||
("authorization", "Bearer forwarded"),
|
||||
("x-api-key", "sk")
|
||||
]))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_request_headers_are_the_given_headers() {
|
||||
let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "claude",
|
||||
"max_tokens": 16,
|
||||
"speed": "fast",
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
DefaultsConfig.request_headers(headers(&[("x-api-key", "sk")]), &request),
|
||||
headers(&[("x-api-key", "sk")])
|
||||
);
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::own_header_is_kept(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("x-api-key", "forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("x-api-key", "forwarded")]))
|
||||
)]
|
||||
#[case::own_header_in_any_casing_is_kept(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("X-Api-Key", "forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("X-Api-Key", "forwarded")]))
|
||||
)]
|
||||
#[case::accepted_bearer_is_kept(
|
||||
X_API_KEY,
|
||||
true,
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("authorization", "Bearer forwarded")]))
|
||||
)]
|
||||
#[case::bearer_the_provider_does_not_accept_gets_the_key_too(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("authorization", "Bearer forwarded"), ("x-api-key", "sk")]))
|
||||
)]
|
||||
#[case::blank_bearer_gets_the_key(
|
||||
X_API_KEY,
|
||||
true,
|
||||
headers(&[("authorization", "Bearer ")]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("authorization", "Bearer "), ("x-api-key", "sk")]))
|
||||
)]
|
||||
#[case::key_goes_in_the_provider_header(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("content-type", "application/json")]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("content-type", "application/json"), ("x-api-key", "sk")]))
|
||||
)]
|
||||
#[case::key_goes_in_a_bearer(
|
||||
MessagesAuthStrategy::Bearer,
|
||||
false,
|
||||
headers(&[]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("authorization", "Bearer sk")]))
|
||||
)]
|
||||
#[case::bearer_strategy_keeps_a_forwarded_authorization(
|
||||
MessagesAuthStrategy::Bearer,
|
||||
false,
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("authorization", "Bearer forwarded")]))
|
||||
)]
|
||||
#[case::missing_key_is_an_error(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[]),
|
||||
None,
|
||||
Err(Error::MissingField("api_key"))
|
||||
)]
|
||||
fn default_authenticate_applies_the_key_unless_a_credential_is_forwarded(
|
||||
#[case] strategy: MessagesAuthStrategy,
|
||||
#[case] accepts_bearer: bool,
|
||||
#[case] forwarded: Headers,
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] expected: Result<Headers, Error>,
|
||||
) {
|
||||
let config = StubConfig {
|
||||
strategy,
|
||||
accepts_bearer,
|
||||
};
|
||||
assert_eq!(config.authenticate(forwarded, api_key, &|_| None), expected);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -302,7 +302,3 @@ fn has_blank_text(message: &ChatMessage) -> bool {
|
|||
}),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "tests.rs"]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -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!({})
|
||||
|
|
@ -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")
|
||||
|
|
@ -2,9 +2,11 @@ use bytes::Bytes;
|
|||
use litellm_core::messages::{
|
||||
Error,
|
||||
route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput},
|
||||
types::MessagesShaping,
|
||||
};
|
||||
use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use pyo3::{
|
||||
exceptions::{PyException, PyValueError},
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
|
|
@ -18,9 +20,10 @@ use crate::{
|
|||
marshal::{optional_timeout, python_timeout_seconds},
|
||||
};
|
||||
|
||||
/// The Anthropic Messages body fields a caller may pass besides `model` and `messages`,
|
||||
/// as `AnthropicMessagesRequestOptionalParams` declares them.
|
||||
const BODY_FIELDS: [&str; 20] = [
|
||||
const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.messages.route_host";
|
||||
const REQUEST_ERROR_MARKER: &str = "messages_request_error";
|
||||
|
||||
const BODY_FIELDS: [&str; 22] = [
|
||||
"max_tokens",
|
||||
"metadata",
|
||||
"stop_sequences",
|
||||
|
|
@ -35,14 +38,46 @@ const BODY_FIELDS: [&str; 20] = [
|
|||
"top_p",
|
||||
"mcp_servers",
|
||||
"context_management",
|
||||
"compaction",
|
||||
"container",
|
||||
"output_format",
|
||||
"speed",
|
||||
"output_config",
|
||||
"cache_control",
|
||||
"reasoning_effort",
|
||||
"safeguards",
|
||||
];
|
||||
|
||||
fn merge_headers(
|
||||
forwarded: Option<Map<String, Value>>,
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
) -> Option<Map<String, Value>> {
|
||||
let merged: Map<String, Value> = forwarded
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(extra_headers.into_iter().flatten())
|
||||
.collect();
|
||||
(!merged.is_empty()).then_some(merged)
|
||||
}
|
||||
|
||||
fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
||||
match error {
|
||||
Error::Transport(TransportError::Http { status, body }) => {
|
||||
let error = RustUpstreamError::new_err((status, body));
|
||||
error
|
||||
.value(py)
|
||||
.setattr("headers", Vec::<(String, String)>::new())?;
|
||||
Ok(error)
|
||||
}
|
||||
Error::InvalidRequest(message) => {
|
||||
let error = PyValueError::new_err(message);
|
||||
error.value(py).setattr(REQUEST_ERROR_MARKER, true)?;
|
||||
Ok(error)
|
||||
}
|
||||
other => Ok(messages_error_to_pyerr(other)),
|
||||
}
|
||||
}
|
||||
|
||||
/// The Python side of the Messages route: projects the prepared arguments and builds the
|
||||
/// public response, chunks and exceptions.
|
||||
pub(super) struct MessagesRouteHost {
|
||||
|
|
@ -84,19 +119,65 @@ impl MessagesRouteHost {
|
|||
.map(|value| python_timeout_seconds(py, value.unbind()))
|
||||
.transpose()?
|
||||
.flatten();
|
||||
let custom_llm_provider = string("custom_llm_provider")?;
|
||||
let shaping = self.shaping(py, &model, custom_llm_provider.as_deref(), arguments)?;
|
||||
Ok(MessagesCall {
|
||||
model,
|
||||
body,
|
||||
api_key: string("api_key")?,
|
||||
api_base: string("api_base")?,
|
||||
custom_llm_provider: string("custom_llm_provider")?,
|
||||
extra_headers: argument("extra_headers")?
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()?,
|
||||
extra_headers: self.merged_headers(py, arguments)?,
|
||||
provider_specific_header: self.provider_specific_header(py, arguments)?,
|
||||
custom_llm_provider,
|
||||
timeout: optional_timeout(timeout),
|
||||
shaping,
|
||||
})
|
||||
}
|
||||
|
||||
fn merged_headers(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<Option<Map<String, Value>>> {
|
||||
let request = self.request.bind(py);
|
||||
let mapping = |name: &str| -> PyResult<Option<Map<String, Value>>> {
|
||||
lookup(arguments, request, name)?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()
|
||||
};
|
||||
Ok(merge_headers(
|
||||
mapping("headers")?,
|
||||
mapping("extra_headers")?,
|
||||
))
|
||||
}
|
||||
|
||||
fn provider_specific_header(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<Option<ProviderSpecificHeaders>> {
|
||||
lookup(arguments, self.request.bind(py), "provider_specific_header")?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn shaping(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<MessagesShaping> {
|
||||
let projected = py.import(ROUTE_HOST_MODULE)?.getattr("shaping")?.call1((
|
||||
model,
|
||||
custom_llm_provider,
|
||||
arguments,
|
||||
))?;
|
||||
from_py(&projected)
|
||||
}
|
||||
|
||||
fn provider(&self, py: Python<'_>) -> String {
|
||||
self.request
|
||||
.bind(py)
|
||||
|
|
@ -112,7 +193,7 @@ impl MessagesRouteHost {
|
|||
return error;
|
||||
}
|
||||
let mapped = py
|
||||
.import("litellm.rust_bridge.messages.route_host")
|
||||
.import(ROUTE_HOST_MODULE)
|
||||
.and_then(|module| module.getattr("map_failure"))
|
||||
.and_then(|map| map.call1((error.value(py), self.request.bind(py), self.provider(py))))
|
||||
.and_then(|mapped| {
|
||||
|
|
@ -148,7 +229,7 @@ impl RouteHost for MessagesRouteHost {
|
|||
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {
|
||||
match response {
|
||||
MessagesOutput::Message(message) => py
|
||||
.import("litellm.rust_bridge.messages.route_host")?
|
||||
.import(ROUTE_HOST_MODULE)?
|
||||
.getattr("response")?
|
||||
.call1((to_py(py, message.as_ref())?,))
|
||||
.map(Bound::unbind),
|
||||
|
|
@ -161,17 +242,12 @@ impl RouteHost for MessagesRouteHost {
|
|||
}
|
||||
|
||||
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
||||
let native = match error {
|
||||
Error::Transport(TransportError::Http { status, body }) => {
|
||||
let error = RustUpstreamError::new_err((status, body));
|
||||
error
|
||||
.value(py)
|
||||
.setattr("headers", Vec::<(String, String)>::new())?;
|
||||
error
|
||||
}
|
||||
other => messages_error_to_pyerr(other),
|
||||
};
|
||||
Ok(self.map_failure(py, native))
|
||||
if let Error::Secret(source) = &error
|
||||
&& let Some(original) = crate::secrets::python_error(py, source.source_error())
|
||||
{
|
||||
return Ok(original);
|
||||
}
|
||||
Ok(self.map_failure(py, native_error(py, error)?))
|
||||
}
|
||||
|
||||
fn host_error(error: &PyErr) -> Error {
|
||||
|
|
@ -184,3 +260,62 @@ impl RouteHost for MessagesRouteHost {
|
|||
visit.call(&self.request)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn map(value: Value) -> Map<String, Value> {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::extra_over_forwarded(
|
||||
Some(json!({"X-Priority": "forwarded", "X-Forwarded-Only": "keep"})),
|
||||
Some(json!({"X-Priority": "extra", "X-Extra-Only": "also-keep"})),
|
||||
Some(json!({"X-Priority": "extra", "X-Forwarded-Only": "keep", "X-Extra-Only": "also-keep"})),
|
||||
)]
|
||||
#[case::only_forwarded(Some(json!({"X-Forwarded": "yes"})), None, Some(json!({"X-Forwarded": "yes"})))]
|
||||
#[case::only_extra_headers(
|
||||
None,
|
||||
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
|
||||
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
|
||||
)]
|
||||
#[case::nothing(None, Some(json!({})), None)]
|
||||
fn headers_merge_forwarded_then_extra(
|
||||
#[case] forwarded: Option<Value>,
|
||||
#[case] extra_headers: Option<Value>,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
assert_eq!(
|
||||
merge_headers(forwarded.map(map), extra_headers.map(map)),
|
||||
expected.map(map)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::rejected_request(Error::InvalidRequest("does not support top_k=5".into()), true)]
|
||||
#[case::unresolvable_provider(Error::InvalidProvider("openai".into()), false)]
|
||||
#[case::upstream_failure(
|
||||
Error::Transport(TransportError::Http { status: 400, body: "bad".into() }),
|
||||
false,
|
||||
)]
|
||||
fn only_request_rejections_carry_the_request_error_marker(
|
||||
#[case] error: Error,
|
||||
#[case] marked: bool,
|
||||
) {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let native = native_error(py, error).unwrap();
|
||||
let marker = native
|
||||
.value(py)
|
||||
.getattr_opt(REQUEST_ERROR_MARKER)
|
||||
.unwrap()
|
||||
.map(|value| value.extract::<bool>().unwrap());
|
||||
assert_eq!(marker.unwrap_or(false), marked);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -39,11 +39,12 @@ fn run_messages(
|
|||
"the Rust Messages route does not serve this provider",
|
||||
));
|
||||
}
|
||||
let secrets = crate::secrets::source(py)?;
|
||||
run_legacy_call(
|
||||
py,
|
||||
SURFACE,
|
||||
PublicCall::capture(&request, &args, &kwargs)?,
|
||||
crate::logger::LoggedMachine::new(messages_machine()),
|
||||
crate::logger::LoggedMachine::new(messages_machine(secrets)),
|
||||
MessagesRouteHost::new(request.unbind()),
|
||||
asynchronous,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,3 +8,6 @@ repository.workspace = true
|
|||
[dependencies]
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -17,12 +17,48 @@ pub enum MessageContent {
|
|||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ContentBlock {
|
||||
#[serde(rename = "type", default, skip_serializing_if = "Option::is_none")]
|
||||
pub block_type: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub text: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub thinking: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub signature: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub data: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub name: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tool_use_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub content: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub provider_specific_fields: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_control: Option<CacheControl>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
impl ContentBlock {
|
||||
pub fn text(text: impl Into<String>) -> Self {
|
||||
Self {
|
||||
block_type: Some("text".to_string()),
|
||||
text: Some(text.into()),
|
||||
..Self::default()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_type(&self, block_type: &str) -> bool {
|
||||
self.block_type.as_deref() == Some(block_type)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct CacheControl {
|
||||
#[serde(rename = "type", skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -85,6 +121,126 @@ pub struct AnthropicMessagesRequest {
|
|||
pub speed: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub inference_geo: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning_effort: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub compaction: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
impl AnthropicMessage {
|
||||
pub fn blocks(&self) -> &[ContentBlock] {
|
||||
match &self.content {
|
||||
MessageContent::Blocks(blocks) => blocks,
|
||||
MessageContent::Text(_) => &[],
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_blocks(self, blocks: Vec<ContentBlock>) -> Self {
|
||||
Self {
|
||||
content: MessageContent::Blocks(blocks),
|
||||
..self
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn round_trip<T: serde::de::DeserializeOwned + Serialize>(value: &Value) -> Value {
|
||||
let parsed: T = serde_json::from_value(value.clone()).unwrap();
|
||||
serde_json::to_value(parsed).unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::text(json!({"type": "text", "text": "hi"}))]
|
||||
#[case::text_with_citations_and_cache_control(json!({
|
||||
"type": "text",
|
||||
"text": "hi",
|
||||
"citations": [{"type": "char_location", "cited_text": "x"}],
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h", "scope": "global", "future": 1}
|
||||
}))]
|
||||
#[case::image(json!({"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "AA=="}}))]
|
||||
#[case::thinking(json!({"type": "thinking", "thinking": "hmm", "signature": "sig"}))]
|
||||
#[case::redacted_thinking(json!({"type": "redacted_thinking", "data": "opaque"}))]
|
||||
#[case::tool_use(json!({"type": "tool_use", "id": "toolu_1", "name": "f", "input": {"q": [1, null]}}))]
|
||||
#[case::tool_result_with_text(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": "ok", "is_error": false}))]
|
||||
#[case::tool_result_with_blocks(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": [{"type": "text", "text": "ok"}]}))]
|
||||
#[case::web_search_result_with_nulls(json!({
|
||||
"type": "web_search_tool_result",
|
||||
"tool_use_id": "srvtoolu_1",
|
||||
"content": [{"type": "web_search_result", "url": "u", "page_age": null, "encrypted_content": ""}]
|
||||
}))]
|
||||
#[case::provider_specific_fields(json!({"type": "tool_use", "id": "t", "name": "f", "input": {}, "provider_specific_fields": {"x": 1}}))]
|
||||
#[case::untyped(json!({"unknown": {"nested": true}}))]
|
||||
fn content_block_round_trips_unchanged(#[case] block: Value) {
|
||||
assert_eq!(round_trip::<ContentBlock>(&block), block);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn text_constructor_serializes_as_a_text_block() {
|
||||
assert_eq!(
|
||||
serde_json::to_value(ContentBlock::text("hello")).unwrap(),
|
||||
json!({"type": "text", "text": "hello"})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::same_type(json!({"type": "tool_use"}), "tool_use", true)]
|
||||
#[case::other_type(json!({"type": "tool_result"}), "tool_use", false)]
|
||||
#[case::prefix_of_type(json!({"type": "tool_use"}), "tool", false)]
|
||||
#[case::no_type(json!({"text": "x"}), "text", false)]
|
||||
fn is_type_matches_the_exact_block_type(
|
||||
#[case] block: Value,
|
||||
#[case] block_type: &str,
|
||||
#[case] expected: bool,
|
||||
) {
|
||||
let block: ContentBlock = serde_json::from_value(block).unwrap();
|
||||
assert_eq!(block.is_type(block_type), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::string_content(json!({"role": "user", "content": "hi"}), vec![])]
|
||||
#[case::block_content(
|
||||
json!({"role": "user", "content": [{"type": "text", "text": "a"}, {"type": "text", "text": "b"}]}),
|
||||
vec![ContentBlock::text("a"), ContentBlock::text("b")],
|
||||
)]
|
||||
fn message_blocks_list_only_block_content(
|
||||
#[case] message: Value,
|
||||
#[case] expected: Vec<ContentBlock>,
|
||||
) {
|
||||
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
|
||||
assert_eq!(message.blocks(), expected.as_slice());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::replaces_string_content(json!({"role": "assistant", "content": "old", "name": "kept"}))]
|
||||
#[case::replaces_block_content(json!({"role": "assistant", "content": [{"type": "text", "text": "old"}], "name": "kept"}))]
|
||||
fn with_blocks_replaces_content_and_keeps_the_rest(#[case] message: Value) {
|
||||
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(message.with_blocks(vec![ContentBlock::text("new")])).unwrap(),
|
||||
json!({"role": "assistant", "content": [{"type": "text", "text": "new"}], "name": "kept"})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::minimal(json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]}))]
|
||||
#[case::reasoning_effort_compaction_and_unknown_fields(json!({
|
||||
"model": "m",
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
|
||||
"max_tokens": 8,
|
||||
"reasoning_effort": "high",
|
||||
"compaction": {"type": "auto"},
|
||||
"safeguards": [{"type": "dangerous_tool_use", "classifier_context": {"v": 1}}],
|
||||
"metadata": {"user_id": "u"}
|
||||
}))]
|
||||
fn request_round_trips_unchanged(#[case] request: Value) {
|
||||
assert_eq!(round_trip::<AnthropicMessagesRequest>(&request), request);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,8 +9,6 @@ pub struct AnthropicMessagesResponse {
|
|||
pub role: String,
|
||||
pub model: String,
|
||||
pub content: Vec<Value>,
|
||||
// Anthropic always includes stop_reason / stop_sequence, null until the turn
|
||||
// ends; serialize them even when None so callers see the same shape as Python.
|
||||
pub stop_reason: Option<String>,
|
||||
pub stop_sequence: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -20,3 +18,61 @@ pub struct AnthropicMessagesResponse {
|
|||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn response(
|
||||
stop_reason: Option<&str>,
|
||||
stop_sequence: Option<&str>,
|
||||
usage: Option<Value>,
|
||||
container: Option<Value>,
|
||||
) -> AnthropicMessagesResponse {
|
||||
AnthropicMessagesResponse {
|
||||
id: "msg_1".to_string(),
|
||||
message_type: "message".to_string(),
|
||||
role: "assistant".to_string(),
|
||||
model: "claude".to_string(),
|
||||
content: vec![],
|
||||
stop_reason: stop_reason.map(str::to_string),
|
||||
stop_sequence: stop_sequence.map(str::to_string),
|
||||
usage,
|
||||
container,
|
||||
extra: Map::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::turn_in_progress(None, None, json!(null), json!(null))]
|
||||
#[case::ended_on_end_turn(Some("end_turn"), None, json!("end_turn"), json!(null))]
|
||||
#[case::ended_on_stop_sequence(Some("stop_sequence"), Some("###"), json!("stop_sequence"), json!("###"))]
|
||||
fn stop_fields_are_always_serialized(
|
||||
#[case] stop_reason: Option<&str>,
|
||||
#[case] stop_sequence: Option<&str>,
|
||||
#[case] expected_reason: Value,
|
||||
#[case] expected_sequence: Value,
|
||||
) {
|
||||
let body: Value = serde_json::to_value(response(stop_reason, stop_sequence, None, None))
|
||||
.expect("serializable");
|
||||
assert_eq!(body.get("stop_reason"), Some(&expected_reason));
|
||||
assert_eq!(body.get("stop_sequence"), Some(&expected_sequence));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::absent(None, None)]
|
||||
#[case::present(Some(json!({"input_tokens": 1})), Some(json!({"id": "c_1"})))]
|
||||
fn usage_and_container_are_omitted_only_when_none(
|
||||
#[case] usage: Option<Value>,
|
||||
#[case] container: Option<Value>,
|
||||
) {
|
||||
let body: Value =
|
||||
serde_json::to_value(response(None, None, usage.clone(), container.clone()))
|
||||
.expect("serializable");
|
||||
assert_eq!(body.get("usage").cloned(), usage);
|
||||
assert_eq!(body.get("container").cloned(), container);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,21 @@ use serde_json::{Map, Value};
|
|||
|
||||
use crate::llms::openai::{ChatCompletionThinkingBlock, ChatCompletionToolCallChunk};
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ProviderSpecificHeader {
|
||||
#[serde(default)]
|
||||
pub custom_llm_provider: String,
|
||||
#[serde(default)]
|
||||
pub extra_headers: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum ProviderSpecificHeaders {
|
||||
One(ProviderSpecificHeader),
|
||||
Many(Vec<ProviderSpecificHeader>),
|
||||
}
|
||||
|
||||
/// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python
|
||||
/// path reports so cost tracking sees the same numbers on either path.
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
|
|
@ -41,7 +56,7 @@ pub struct ChatCompletionsChoice {
|
|||
///
|
||||
/// There is deliberately no `id`: Python mints the `chatcmpl-…` id on the
|
||||
/// `ModelResponse` it already created, and echoing the provider's own id here
|
||||
/// would change it. Pinned by `response_carries_no_id` in `tests.rs`.
|
||||
/// would change it. Pinned by `response_carries_no_id` in the Anthropic chat transformation tests.
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ChatCompletionsResponse {
|
||||
pub created: u64,
|
||||
|
|
|
|||
|
|
@ -11,7 +11,10 @@ from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_
|
|||
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo
|
||||
from litellm.llms.bedrock.batches.transformation import titan_embedding_usage_from_batch_output
|
||||
from litellm.llms.vertex_ai.batches.transformation import vertex_prompt_tokens_details
|
||||
from litellm.llms.vertex_ai.batches.transformation import (
|
||||
is_native_vertex_batch_output_row,
|
||||
native_vertex_batch_row_stats,
|
||||
)
|
||||
from litellm.types.llms.openai import Batch
|
||||
from litellm.types.utils import ModelInfo, Usage
|
||||
from litellm.utils import token_counter
|
||||
|
|
@ -31,6 +34,20 @@ class BatchCostUsageResult:
|
|||
|
||||
|
||||
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
|
||||
|
||||
|
||||
def _uses_native_vertex_output(
|
||||
custom_llm_provider: str,
|
||||
model_name: str | None,
|
||||
first_row: Mapping[str, object] | None,
|
||||
) -> bool:
|
||||
if custom_llm_provider != "vertex_ai":
|
||||
return False
|
||||
if model_name and getattr(litellm, "disable_vertex_batch_output_transformation", False):
|
||||
return True
|
||||
return first_row is not None and is_native_vertex_batch_output_row(first_row)
|
||||
|
||||
|
||||
_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"})
|
||||
|
||||
|
||||
|
|
@ -66,12 +83,9 @@ async def calculate_batch_cost_and_usage(
|
|||
deployment-specific pricing (e.g. input_cost_per_token_batches)
|
||||
is used instead of the global cost map.
|
||||
"""
|
||||
if (
|
||||
custom_llm_provider == "vertex_ai"
|
||||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
):
|
||||
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
|
||||
first_row: Final = file_content_dictionary[0] if file_content_dictionary else None
|
||||
if _uses_native_vertex_output(custom_llm_provider, model_name, first_row):
|
||||
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name, model_info=model_info)
|
||||
|
||||
return _aggregate_batch_cost_usage_models(
|
||||
entries=file_content_dictionary,
|
||||
|
|
@ -126,11 +140,11 @@ async def _handle_completed_batch(
|
|||
)
|
||||
|
||||
output_file_result: Final = (
|
||||
calculate_vertex_ai_batch_cost_and_usage(_get_file_content_as_dictionary(file_content), model_name)
|
||||
if (
|
||||
custom_llm_provider == "vertex_ai"
|
||||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
calculate_vertex_ai_batch_cost_and_usage(
|
||||
_iter_batch_output_entries(file_content), model_name, model_info=model_info
|
||||
)
|
||||
if _uses_native_vertex_output(
|
||||
custom_llm_provider, model_name, next(_iter_batch_output_entries(file_content), None)
|
||||
)
|
||||
else _aggregate_batch_cost_usage_models(
|
||||
entries=_iter_batch_output_entries(file_content),
|
||||
|
|
@ -332,69 +346,36 @@ def _aggregate_batch_cost_usage_models(
|
|||
|
||||
|
||||
def calculate_vertex_ai_batch_cost_and_usage(
|
||||
vertex_ai_batch_responses: list[dict],
|
||||
vertex_ai_batch_responses: Iterable[dict],
|
||||
model_name: str | None = None,
|
||||
model_info: ModelInfo | None = None,
|
||||
) -> BatchCostUsageResult:
|
||||
"""
|
||||
Calculate both cost and usage from raw Vertex AI batch responses.
|
||||
|
||||
Used only when ``litellm.disable_vertex_batch_output_transformation = True``.
|
||||
In that case the GCS predictions.jsonl is returned as-is, with each line in
|
||||
the native Vertex format:
|
||||
|
||||
{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}}}
|
||||
|
||||
usageMetadata contains promptTokenCount, candidatesTokenCount, totalTokenCount.
|
||||
|
||||
A row with no ``response`` is counted as failed - the same signal already
|
||||
used to skip it from cost/usage aggregation, since Vertex batch prediction
|
||||
output doesn't establish a distinct error shape in this (non-default) path.
|
||||
Cost and usage of a native Vertex predictions.jsonl, one
|
||||
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
|
||||
generateContent row or `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
|
||||
embedding row per line. `model_name` (the deployment model) prices every row, else each row's own
|
||||
`modelVersion` does; a row without a usable response counts as failed.
|
||||
"""
|
||||
from litellm.cost_calculator import batch_cost_calculator
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
|
||||
total_prompt_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
|
||||
total_completion_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
|
||||
total_tokens = 0
|
||||
prompt_tokens = 0
|
||||
completion_tokens = 0
|
||||
successful_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
|
||||
failed_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
|
||||
actual_model_name: Final = model_name or "gemini-2.0-flash-001"
|
||||
|
||||
for response in vertex_ai_batch_responses:
|
||||
response_body = response.get("response")
|
||||
if response_body is None:
|
||||
failed_requests += 1
|
||||
continue
|
||||
successful_requests += 1
|
||||
|
||||
usage_metadata = response_body.get("usageMetadata", {})
|
||||
_prompt = usage_metadata.get("promptTokenCount", 0) or 0
|
||||
_completion = usage_metadata.get("candidatesTokenCount", 0) or 0
|
||||
_total = usage_metadata.get("totalTokenCount", 0) or (_prompt + _completion)
|
||||
|
||||
line_usage = Usage(
|
||||
prompt_tokens=_prompt,
|
||||
completion_tokens=_completion,
|
||||
total_tokens=_total,
|
||||
prompt_tokens_details=vertex_prompt_tokens_details(usage_metadata),
|
||||
row_stats: Final = tuple(
|
||||
native_vertex_batch_row_stats(
|
||||
row,
|
||||
model_name,
|
||||
model_info=model_info,
|
||||
calculate_usage=VertexGeminiConfig._calculate_usage,
|
||||
cost_calculator=batch_cost_calculator,
|
||||
)
|
||||
|
||||
try:
|
||||
p_cost, c_cost = batch_cost_calculator(
|
||||
usage=line_usage,
|
||||
model=actual_model_name,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
total_prompt_cost += p_cost
|
||||
total_completion_cost += c_cost
|
||||
except Exception as e:
|
||||
verbose_logger.debug("vertex_ai batch cost calculation error for line: %s", str(e))
|
||||
|
||||
prompt_tokens += _prompt
|
||||
completion_tokens += _completion
|
||||
total_tokens += _total
|
||||
|
||||
for row in vertex_ai_batch_responses
|
||||
)
|
||||
priced: Final = tuple(stats for stats in row_stats if stats is not None)
|
||||
total_prompt_cost: Final = sum(stats.prompt_cost for stats in priced)
|
||||
total_completion_cost: Final = sum(stats.completion_cost for stats in priced)
|
||||
prompt_tokens: Final = sum(stats.usage.prompt_tokens for stats in priced)
|
||||
completion_tokens: Final = sum(stats.usage.completion_tokens for stats in priced)
|
||||
total_tokens: Final = sum(stats.total_tokens for stats in priced)
|
||||
total_cost: Final = total_prompt_cost + total_completion_cost
|
||||
verbose_logger.info(
|
||||
"vertex_ai batch cost: cost=%s, prompt=%d, completion=%d, total=%d, successful=%d, failed=%d",
|
||||
|
|
@ -402,8 +383,8 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
prompt_tokens,
|
||||
completion_tokens,
|
||||
total_tokens,
|
||||
successful_requests,
|
||||
failed_requests,
|
||||
len(priced),
|
||||
len(row_stats) - len(priced),
|
||||
)
|
||||
|
||||
return BatchCostUsageResult(
|
||||
|
|
@ -413,9 +394,13 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
),
|
||||
models=[actual_model_name],
|
||||
successful_requests=successful_requests,
|
||||
failed_requests=failed_requests,
|
||||
models=(
|
||||
[model_name]
|
||||
if model_name
|
||||
else list(dict.fromkeys(stats.model for stats in priced if stats.model is not None))
|
||||
),
|
||||
successful_requests=len(priced),
|
||||
failed_requests=len(row_stats) - len(priced),
|
||||
prompt_cost=total_prompt_cost,
|
||||
completion_cost=total_completion_cost,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1772,6 +1772,8 @@ RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS: Final = max(1, int(os.getenv("RESPONSES_S
|
|||
RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL: Final = float(os.getenv("RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL", "0.2"))
|
||||
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
|
||||
PROXY_DB_LOOKUP_MAX_CONCURRENCY: Final = max(1, int(os.getenv("PROXY_DB_LOOKUP_MAX_CONCURRENCY", "25")))
|
||||
PROXY_DB_LOOKUP_DEADLINE_SECONDS: Final = max(0.1, float(os.getenv("PROXY_DB_LOOKUP_DEADLINE_SECONDS", "10")))
|
||||
PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS: Final = max(0.0, float(os.getenv("PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS", "30")))
|
||||
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
|
||||
RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500")))
|
||||
|
|
|
|||
|
|
@ -2017,8 +2017,8 @@ def _deployment_model_info(
|
|||
return cast(ModelInfo, registered_deployment_info) # cast-ok: router registers deployment prices under its id
|
||||
if litellm_logging_obj is None:
|
||||
return None
|
||||
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None)
|
||||
if litellm_params is None:
|
||||
litellm_params: Final = litellm_logging_obj.litellm_params
|
||||
if not litellm_params:
|
||||
return None
|
||||
return next(
|
||||
(
|
||||
|
|
@ -2036,7 +2036,9 @@ def _ocr_model_info(
|
|||
router_model_id: str | None,
|
||||
) -> OCRPricing | None:
|
||||
deployment_info: Final = _deployment_model_info(litellm_logging_obj, custom_pricing, router_model_id)
|
||||
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None) if custom_pricing else None
|
||||
litellm_params: Final = (
|
||||
litellm_logging_obj.litellm_params if custom_pricing and litellm_logging_obj is not None else None
|
||||
)
|
||||
if litellm_params is None:
|
||||
return deployment_info
|
||||
return _layered_ocr_pricing(litellm_params, deployment_info)
|
||||
|
|
|
|||
|
|
@ -129,7 +129,7 @@ async def list_tools_with_pagination(
|
|||
)
|
||||
tools.extend(result.tools)
|
||||
|
||||
next_cursor = getattr(result, "next_cursor", None)
|
||||
next_cursor = result.next_cursor
|
||||
if not isinstance(next_cursor, str) or not next_cursor:
|
||||
return tools
|
||||
if next_cursor in seen_cursors:
|
||||
|
|
|
|||
|
|
@ -176,6 +176,22 @@ def create_file(
|
|||
if logging_obj is None:
|
||||
raise ValueError("logging_obj is required")
|
||||
client: Final = kwargs.get("client")
|
||||
if litellm_params_dict.get("passthrough") is True and (
|
||||
custom_llm_provider != "vertex_ai" or purpose != "batch"
|
||||
):
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message=(
|
||||
"`passthrough=True` uploads the file bytes unchanged for a native Vertex AI batch, so it needs "
|
||||
f"custom_llm_provider='vertex_ai' and purpose='batch', got '{custom_llm_provider}' and '{purpose}'."
|
||||
),
|
||||
model="n/a",
|
||||
llm_provider=custom_llm_provider or "n/a",
|
||||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="passthrough needs a vertex_ai batch",
|
||||
request=httpx.Request(method="create_file", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -112,7 +112,7 @@ class ArizeLogger(OpenTelemetry):
|
|||
if value is None or value in ("", "None"):
|
||||
return None
|
||||
try:
|
||||
rate = float(value)
|
||||
rate: Final = float(value)
|
||||
except (TypeError, ValueError):
|
||||
verbose_logger.warning(
|
||||
"ArizeLogger: %s value %r is not a number; exporting the request",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ from __future__ import annotations
|
|||
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -35,17 +35,6 @@ else:
|
|||
AsyncIOScheduler = Any
|
||||
|
||||
|
||||
class _PodLockManager(Protocol):
|
||||
"""The subset of PodLockManager this logger drives to serialize the export across pods."""
|
||||
|
||||
@property
|
||||
def redis_cache(self) -> object: ...
|
||||
|
||||
async def acquire_lock(self, cronjob_id: str) -> bool | None: ...
|
||||
|
||||
async def release_lock(self, cronjob_id: str) -> None: ...
|
||||
|
||||
|
||||
def _parse_metrics_marker(
|
||||
marker: object | None,
|
||||
) -> datetime | None:
|
||||
|
|
@ -237,13 +226,10 @@ class MavvrikFocusLogger(FocusLogger):
|
|||
"""Scheduler entry point — uses Mavvrik-specific pod-lock key."""
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj # noqa: PLC0415
|
||||
|
||||
pod_lock_manager: _PodLockManager | None = None
|
||||
if proxy_logging_obj is not None:
|
||||
writer: Final[object] = getattr(proxy_logging_obj, "db_spend_update_writer", None)
|
||||
if writer is not None:
|
||||
pod_lock_manager = getattr(writer, "pod_lock_manager", None)
|
||||
|
||||
if pod_lock_manager and pod_lock_manager.redis_cache:
|
||||
pod_lock_manager: Final = (
|
||||
proxy_logging_obj.db_spend_update_writer.pod_lock_manager if proxy_logging_obj is not None else None
|
||||
)
|
||||
if pod_lock_manager is not None and pod_lock_manager.redis_cache:
|
||||
acquired: Final = await pod_lock_manager.acquire_lock(cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME)
|
||||
if not acquired:
|
||||
verbose_proxy_logger.debug("Mavvrik FOCUS export: unable to acquire pod lock")
|
||||
|
|
|
|||
|
|
@ -62,10 +62,10 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy.proxy_server import UserAPIKeyAuth as _UserAPIKeyAuth
|
||||
|
||||
Span = _Span | Any
|
||||
Tracer = _Tracer | Any
|
||||
Context = _Context | Any
|
||||
SpanExporter = _SpanExporter | Any
|
||||
UserAPIKeyAuth = _UserAPIKeyAuth | Any
|
||||
Tracer = _Tracer
|
||||
Context = _Context
|
||||
SpanExporter = _SpanExporter
|
||||
UserAPIKeyAuth = _UserAPIKeyAuth
|
||||
ManagementEndpointLoggingPayload = _ManagementEndpointLoggingPayload | Any
|
||||
else:
|
||||
Span = Any
|
||||
|
|
@ -2730,7 +2730,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
self.handle_callback_failure(callback_name=self.callback_name or "opentelemetry")
|
||||
verbose_logger.exception("OpenTelemetry logging error in set_attributes %s", str(e))
|
||||
|
||||
def _cast_as_primitive_value_type(self, value) -> str | bool | int | float:
|
||||
def _cast_as_primitive_value_type(self, value: object) -> str | bool | int | float:
|
||||
"""
|
||||
Casts the value to a primitive OTEL type if it is not already a primitive type.
|
||||
|
||||
|
|
|
|||
|
|
@ -2778,6 +2778,13 @@ class PrometheusLogger(CustomLogger):
|
|||
- increment deployment failure responses metric
|
||||
- increment deployment total requests metric
|
||||
|
||||
Both counters also carry a model_group label. When a deployment was
|
||||
actually selected, model_group is the router-resolved value and is
|
||||
trusted as-is. On a pre-routing reject (no deployment selected), it
|
||||
is caller-supplied via litellm_params.metadata and is bounded with
|
||||
_bounded_requested_model_label the same way requested_model is, so an
|
||||
unrecognized value cannot mint unbounded label series.
|
||||
|
||||
Args:
|
||||
request_kwargs: dict
|
||||
|
||||
|
|
@ -2844,6 +2851,7 @@ class PrometheusLogger(CustomLogger):
|
|||
label_api_base = api_base
|
||||
label_api_provider = llm_provider
|
||||
label_requested_model = model_group or litellm_model_name
|
||||
label_model_group = model_group
|
||||
else:
|
||||
label_litellm_model_name = ""
|
||||
label_model_id = ""
|
||||
|
|
@ -2852,6 +2860,7 @@ class PrometheusLogger(CustomLogger):
|
|||
label_requested_model = (
|
||||
_bounded_requested_model_label(litellm_model_name or model_group, router_originated=True) or ""
|
||||
)
|
||||
label_model_group = _bounded_requested_model_label(model_group, router_originated=True)
|
||||
|
||||
enum_values: Final = UserAPIKeyLabelValues(
|
||||
litellm_model_name=label_litellm_model_name,
|
||||
|
|
@ -2861,6 +2870,7 @@ class PrometheusLogger(CustomLogger):
|
|||
exception_status=exception_status,
|
||||
exception_class=(self._get_exception_class_name(exception) if exception else None),
|
||||
requested_model=label_requested_model,
|
||||
model_group=label_model_group,
|
||||
hashed_api_key=hashed_api_key,
|
||||
api_key_alias=api_key_alias,
|
||||
user_email=user_email,
|
||||
|
|
@ -2912,9 +2922,21 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id: str | None,
|
||||
api_base: str | None,
|
||||
llm_provider: str | None,
|
||||
model_group: str | None,
|
||||
):
|
||||
"""
|
||||
Set the deployment TPM and RPM limits metrics
|
||||
|
||||
Args:
|
||||
model_info: the deployment's static model_info config (id, tpm, rpm, etc.)
|
||||
litellm_params: the deployment's litellm_params, as a tpm/rpm fallback source
|
||||
litellm_model_name: the resolved deployment model name
|
||||
model_id: the deployment's model_id
|
||||
api_base: the deployment's api_base
|
||||
llm_provider: the deployment's custom_llm_provider
|
||||
model_group: the router-resolved model_group the deployment belongs to,
|
||||
from the caller's already-resolved enum_values.model_group (trusted,
|
||||
not caller-supplied at this call site)
|
||||
"""
|
||||
tpm: Final = model_info.get("tpm") or litellm_params.get("tpm")
|
||||
rpm: Final = model_info.get("rpm") or litellm_params.get("rpm")
|
||||
|
|
@ -2927,6 +2949,7 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
api_provider=llm_provider,
|
||||
model_group=model_group,
|
||||
),
|
||||
)
|
||||
self.litellm_deployment_tpm_limit.labels(**_labels).set(tpm)
|
||||
|
|
@ -2939,6 +2962,7 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
api_provider=llm_provider,
|
||||
model_group=model_group,
|
||||
),
|
||||
)
|
||||
self.litellm_deployment_rpm_limit.labels(**_labels).set(rpm)
|
||||
|
|
@ -3058,6 +3082,7 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
llm_provider=llm_provider,
|
||||
model_group=enum_values.model_group,
|
||||
)
|
||||
|
||||
remaining_requests: int | None = None
|
||||
|
|
|
|||
|
|
@ -1849,7 +1849,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
for tool_call in tool_calls:
|
||||
# Handle both Anthropic-style input and OpenAI-style function.arguments
|
||||
query = None
|
||||
tool_args: dict | None = None # mutable-ok: the tool call's own arguments dict
|
||||
tool_args: dict[str, object] | None = None # mutable-ok: the tool call's own arguments dict
|
||||
if "input" in tool_call and isinstance(tool_call["input"], dict):
|
||||
tool_args = tool_call["input"]
|
||||
query = tool_args.get("query")
|
||||
|
|
|
|||
|
|
@ -365,7 +365,7 @@ def _budget_reservation_on_auth_object(user_api_key_auth: object) -> object:
|
|||
return getattr(user_api_key_auth, "budget_reservation", None)
|
||||
|
||||
|
||||
def budget_reservation_from_metadata(metadata: Mapping[str, object]) -> dict | None:
|
||||
def budget_reservation_from_metadata(metadata: Mapping[str, object]) -> dict[str, object] | None:
|
||||
stamped: Final = metadata.get("user_api_key_budget_reservation")
|
||||
if isinstance(stamped, dict):
|
||||
return stamped
|
||||
|
|
@ -765,3 +765,30 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo
|
|||
RESPONSE_COST_HEADER: cost,
|
||||
}
|
||||
hidden_params["additional_headers"] = merged
|
||||
|
||||
|
||||
_HIDDEN_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
_PROVIDER_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str])
|
||||
|
||||
|
||||
def set_provider_response_headers_in_hidden_params(
|
||||
response: _CarriesHiddenParams, headers: httpx.Headers | Mapping[str, str]
|
||||
) -> None:
|
||||
hidden_params: Final = response._hidden_params # pyright: ignore[reportPrivateUsage] # no public accessor
|
||||
existing_additional_headers: Final[object] = hidden_params.get("additional_headers")
|
||||
raw_headers: Final[dict[str, str]] = dict(headers) # mutable-ok: stored as the plain-dict hidden param
|
||||
additional_headers: Final[dict[str, object]] = { # mutable-ok: assigned into the plain-dict hidden params
|
||||
**process_response_headers(raw_headers),
|
||||
**(existing_additional_headers if isinstance(existing_additional_headers, Mapping) else _NO_HEADERS),
|
||||
}
|
||||
hidden_params["headers"] = raw_headers
|
||||
hidden_params["additional_headers"] = additional_headers
|
||||
|
||||
|
||||
def get_provider_response_headers_from_hidden_params(response: object) -> Mapping[str, str] | None:
|
||||
hidden_params: Final[object] = getattr(response, "_hidden_params", None)
|
||||
try:
|
||||
validated: Final = _HIDDEN_PARAMS_ADAPTER.validate_python(hidden_params)
|
||||
return _PROVIDER_HEADERS_ADAPTER.validate_python(validated.get("headers"))
|
||||
except ValidationError:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -72,6 +72,7 @@ from litellm.litellm_core_utils.classifier_logging import (
|
|||
is_classifier_call,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_provider_response_headers_from_hidden_params,
|
||||
is_expected_client_error,
|
||||
reconstruct_model_name,
|
||||
set_response_cost_in_hidden_params,
|
||||
|
|
@ -1566,14 +1567,14 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
attr = "debug"
|
||||
|
||||
if json_logs:
|
||||
callattr = getattr(verbose_logger, attr)
|
||||
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
|
||||
callattr(
|
||||
"RAW RESPONSE:\n{}\n\n".format(
|
||||
self.model_call_details.get("original_response", self.model_call_details)
|
||||
),
|
||||
)
|
||||
else:
|
||||
callattr = getattr(verbose_logger, attr)
|
||||
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
|
||||
callattr(
|
||||
"RAW RESPONSE:\n{}\n\n".format(
|
||||
self.model_call_details.get("original_response", self.model_call_details)
|
||||
|
|
@ -2353,6 +2354,15 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
return logging_result
|
||||
|
||||
def _surface_response_headers_from_result(self, logging_result: object) -> None:
|
||||
existing: Final[object] = self.model_call_details.get("response_headers")
|
||||
if existing is not None:
|
||||
return
|
||||
headers: Final = get_provider_response_headers_from_hidden_params(logging_result)
|
||||
if headers is None:
|
||||
return
|
||||
self.model_call_details["response_headers"] = headers
|
||||
|
||||
def _merge_hidden_params_from_response_into_metadata(self, logging_result: object) -> None:
|
||||
"""
|
||||
Copy response._hidden_params into litellm_params.metadata['hidden_params'].
|
||||
|
|
@ -2386,6 +2396,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
build_logging_payload: bool = True,
|
||||
):
|
||||
"""Resolve hidden params, compute response cost, and emit the standard logging payload."""
|
||||
self._surface_response_headers_from_result(logging_result)
|
||||
hidden_params: Final = getattr(logging_result, "_hidden_params", {})
|
||||
if hidden_params:
|
||||
if self.model_call_details.get("litellm_params") is not None:
|
||||
|
|
@ -2788,6 +2799,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if complete_streaming_response is not None:
|
||||
verbose_logger.debug("Logging Details LiteLLM-Success Call streaming complete")
|
||||
self.model_call_details["complete_streaming_response"] = complete_streaming_response
|
||||
self._surface_response_headers_from_result(complete_streaming_response)
|
||||
self.model_call_details["response_cost"] = self._response_cost_calculator(
|
||||
result=complete_streaming_response
|
||||
)
|
||||
|
|
@ -3302,6 +3314,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
print_verbose("Async success callbacks: Got a complete streaming response")
|
||||
|
||||
self.model_call_details["async_complete_streaming_response"] = complete_streaming_response
|
||||
self._surface_response_headers_from_result(complete_streaming_response)
|
||||
|
||||
try:
|
||||
if self.model_call_details.get("cache_hit", False) is True:
|
||||
|
|
@ -5191,7 +5204,7 @@ def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list[Custom
|
|||
for callback in _in_memory_loggers:
|
||||
if (
|
||||
isinstance(callback, OpenTelemetryV2)
|
||||
and getattr(callback, "callback_name", None) == callback_name
|
||||
and callback.callback_name == callback_name
|
||||
and (serves_a_destination or not _exports_nowhere(callback.config))
|
||||
):
|
||||
return callback
|
||||
|
|
@ -5882,7 +5895,7 @@ class StandardLoggingPayloadSetup:
|
|||
base_model: str | None,
|
||||
custom_pricing: bool | None,
|
||||
custom_llm_provider: str | None,
|
||||
init_response_obj: Any | BaseModel | dict,
|
||||
init_response_obj: object,
|
||||
api_base: str | None = None,
|
||||
) -> StandardLoggingModelInformation:
|
||||
model_cost_name: Final = _select_model_name_for_cost_calc(
|
||||
|
|
@ -5915,9 +5928,7 @@ class StandardLoggingPayloadSetup:
|
|||
return model_cost_information
|
||||
|
||||
@staticmethod
|
||||
def get_final_response_obj(
|
||||
response_obj: dict, init_response_obj: Any | BaseModel | dict, kwargs: dict
|
||||
) -> dict | str | list | None:
|
||||
def get_final_response_obj(response_obj: dict, init_response_obj: object, kwargs: dict) -> dict | str | list | None:
|
||||
"""
|
||||
Get final response object after redacting the message input/output from logging
|
||||
"""
|
||||
|
|
@ -6360,16 +6371,19 @@ def _get_status_fields(
|
|||
|
||||
|
||||
def _extract_response_obj_and_hidden_params(
|
||||
init_response_obj: Any | BaseModel | dict,
|
||||
init_response_obj: object,
|
||||
original_exception: Exception | None,
|
||||
) -> tuple[dict, dict | None]:
|
||||
"""Extract response_obj and hidden_params from init_response_obj."""
|
||||
hidden_params: dict | None = None
|
||||
hidden_params: dict | None = (
|
||||
getattr(init_response_obj, "_hidden_params", None)
|
||||
if isinstance(init_response_obj, BaseModel | HttpxBinaryResponseContent)
|
||||
else None
|
||||
)
|
||||
if init_response_obj is None:
|
||||
response_obj = {}
|
||||
elif isinstance(init_response_obj, BaseModel):
|
||||
response_obj = init_response_obj.model_dump()
|
||||
hidden_params = getattr(init_response_obj, "_hidden_params", None)
|
||||
elif isinstance(init_response_obj, dict):
|
||||
response_obj = init_response_obj
|
||||
elif isinstance(init_response_obj, HttpxBinaryResponseContent):
|
||||
|
|
@ -6663,7 +6677,7 @@ def get_standard_logging_object_payload(
|
|||
cost_breakdown=request_cost_breakdown,
|
||||
autorouter_savings=autorouter_savings,
|
||||
autorouter_savings_estimate=(
|
||||
{
|
||||
{ # mutable-ok: spend-log JSON serialization requires plain mappings
|
||||
"version": 3,
|
||||
"status": "unknown",
|
||||
"reason": "pending_projection",
|
||||
|
|
|
|||
|
|
@ -5,7 +5,12 @@ from typing import Final
|
|||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.types.utils import StandardLoggingZeroCostDiagnostic, Usage
|
||||
from litellm.types.utils import (
|
||||
CompletionTokensDetailsWrapper,
|
||||
PromptTokensDetailsWrapper,
|
||||
StandardLoggingZeroCostDiagnostic,
|
||||
Usage,
|
||||
)
|
||||
|
||||
ZERO_COST_COUNTER_NAME: Final = "litellm_zero_cost_requests_total"
|
||||
|
||||
|
|
@ -18,8 +23,8 @@ _NESTED_PRICING: Final = TypeAdapter(Mapping[str, object] | tuple[object, ...])
|
|||
_MAX_PRICING_DEPTH: Final = 4
|
||||
|
||||
|
||||
def _audio_tokens(details: object) -> int:
|
||||
audio_tokens: Final = getattr(details, "audio_tokens", None)
|
||||
def _audio_tokens(details: PromptTokensDetailsWrapper | CompletionTokensDetailsWrapper | None) -> int:
|
||||
audio_tokens: Final = details.audio_tokens if details is not None else None
|
||||
return audio_tokens if isinstance(audio_tokens, int) and audio_tokens > 0 else 0
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2003,11 +2003,11 @@ def strip_encrypted_reasoning_from_messages(messages: object) -> None:
|
|||
"""
|
||||
if not isinstance(messages, list):
|
||||
return
|
||||
for content in _anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
|
||||
for content in anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
|
||||
_strip_encrypted_reasoning_from_blocks(content)
|
||||
|
||||
|
||||
def _anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
|
||||
def anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
|
||||
return (
|
||||
cast(list[object], content) # cast-ok: narrowed by isinstance
|
||||
for message in messages
|
||||
|
|
|
|||
|
|
@ -446,7 +446,7 @@ def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: st
|
|||
|
||||
|
||||
async def _afetch_and_extract_template(
|
||||
model: str, chat_template: Any | None, get_config_fn, get_template_fn
|
||||
model: str, chat_template: str | None, get_config_fn, get_template_fn
|
||||
) -> tuple[str, str, str]:
|
||||
"""
|
||||
Async version: Fetch template and tokens from HuggingFace.
|
||||
|
|
@ -500,7 +500,7 @@ async def _afetch_and_extract_template(
|
|||
|
||||
|
||||
def _fetch_and_extract_template(
|
||||
model: str, chat_template: Any | None, get_config_fn, get_template_fn
|
||||
model: str, chat_template: str | None, get_config_fn, get_template_fn
|
||||
) -> tuple[str, str, str]:
|
||||
"""
|
||||
Sync version: Fetch template and tokens from HuggingFace.
|
||||
|
|
|
|||
|
|
@ -1329,7 +1329,7 @@ class CustomStreamWrapper:
|
|||
"is_finished": chunk_finish_reason is not None,
|
||||
"finish_reason": chunk_finish_reason,
|
||||
"original_chunk": cached_chunk,
|
||||
"tool_calls": (getattr(cached_choice.delta, "tool_calls", None) if cached_choice is not None else None),
|
||||
"tool_calls": cached_choice.delta.tool_calls if cached_choice is not None else None,
|
||||
}
|
||||
|
||||
completion_obj["content"] = response_obj["text"]
|
||||
|
|
|
|||
|
|
@ -48,7 +48,7 @@ def _registry_api_key(agent_litellm_params: Mapping[str, object]) -> str | None:
|
|||
return configured_api_key if isinstance(configured_api_key, str) else None
|
||||
|
||||
|
||||
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, Any] | None:
|
||||
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, object] | None:
|
||||
stored_headers: Final = agent_litellm_params.get("headers")
|
||||
if not isinstance(stored_headers, Mapping):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -685,7 +685,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def _hoisted_top_level_system_message(self, data: dict) -> AllMessageValues | None:
|
||||
def _hoisted_top_level_system_message(self, data: Mapping[str, object]) -> AllMessageValues | None:
|
||||
"""Return the system message produced by translating the top-level prompt."""
|
||||
system: Final = data.get("system")
|
||||
if not system:
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ from typing import (
|
|||
Final,
|
||||
Literal,
|
||||
Protocol,
|
||||
cast, # noqa: TID251 # rebuilt message_delta dict spans the ContentBlockDelta/MessageBlockDelta union
|
||||
get_args,
|
||||
)
|
||||
|
||||
|
|
@ -27,6 +26,7 @@ from litellm.types.llms.anthropic import (
|
|||
ContentBlockDelta,
|
||||
ContextManagementResponse,
|
||||
MessageBlockDelta,
|
||||
MessageDelta,
|
||||
StreamingContentBlockDeltaType,
|
||||
UsageDelta,
|
||||
UsageIteration,
|
||||
|
|
@ -1028,26 +1028,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
self,
|
||||
processed_chunk: ContentBlockDelta | MessageBlockDelta,
|
||||
) -> ContentBlockDelta | MessageBlockDelta:
|
||||
if processed_chunk.get("type") != "message_delta" or not self._refusal_text:
|
||||
if processed_chunk["type"] != "message_delta" or not self._refusal_text:
|
||||
return processed_chunk
|
||||
delta: Final = cast(Mapping[str, object], processed_chunk["delta"]) # cast-ok: keys checked before use
|
||||
delta: Final = processed_chunk["delta"]
|
||||
if delta.get("stop_reason") == "max_tokens":
|
||||
return processed_chunk
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
refusal_stop_details,
|
||||
)
|
||||
|
||||
return cast( # cast-ok: rebuilt dict matches the message_delta TypedDict shape for this branch
|
||||
ContentBlockDelta | MessageBlockDelta,
|
||||
{ # mutable-ok: fresh translation payload; never mutated after construction
|
||||
**processed_chunk,
|
||||
"delta": { # mutable-ok: fresh message_delta payload; never mutated after construction
|
||||
**delta,
|
||||
"stop_reason": "refusal",
|
||||
"stop_details": refusal_stop_details(self._refusal_text),
|
||||
},
|
||||
},
|
||||
)
|
||||
refusal_delta: Final[MessageDelta] = {
|
||||
**delta,
|
||||
"stop_reason": "refusal",
|
||||
"stop_details": refusal_stop_details(self._refusal_text),
|
||||
}
|
||||
refusal_chunk: Final[MessageBlockDelta] = {**processed_chunk, "delta": refusal_delta}
|
||||
return refusal_chunk
|
||||
|
||||
@staticmethod
|
||||
def _delta_has_content(processed_chunk: Mapping[str, object]) -> bool:
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ def _mapping_field(container: object, key: str) -> object | None:
|
|||
"""One key of a raw provider payload, or None when the payload is not a mapping."""
|
||||
if not isinstance(container, Mapping):
|
||||
return None
|
||||
return cast(Mapping[str, object], container).get(key) # cast-ok: raw payload, callers re-check every value
|
||||
return container.get(key)
|
||||
|
||||
|
||||
def _mapping_str_field(container: object, key: str) -> str | None:
|
||||
|
|
|
|||
|
|
@ -169,7 +169,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
cls,
|
||||
summary: Iterable[object],
|
||||
encrypted_content: object,
|
||||
) -> dict[str, Any] | None: # mutable-ok: API message payload
|
||||
) -> dict[str, object] | None: # mutable-ok: API message payload
|
||||
"""The one Anthropic block for a Responses reasoning item.
|
||||
|
||||
The item's encrypted reasoning rides the block's opaque field (`signature`, or
|
||||
|
|
@ -198,7 +198,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
@classmethod
|
||||
def _assistant_group_to_input_items(
|
||||
cls, group: tuple[Mapping[str, object], ...]
|
||||
) -> tuple[dict[str, Any], ...]: # mutable-ok: API message payload
|
||||
) -> tuple[dict[str, object], ...]: # mutable-ok: API message payload
|
||||
first: Final = group[0]
|
||||
btype: Final = first.get("type")
|
||||
if btype in ("thinking", "redacted_thinking"):
|
||||
|
|
|
|||
|
|
@ -1218,6 +1218,8 @@ class BedrockModelInfo(BaseLLMModelInfo):
|
|||
alt_model: Final = BedrockModelInfo.get_non_litellm_routing_model_name(model=model)
|
||||
if base_model in litellm.bedrock_converse_models or alt_model in litellm.bedrock_converse_models:
|
||||
return "converse"
|
||||
if _OPENAI_FAMILY_MODEL_RE.search(base_model):
|
||||
return "converse"
|
||||
return "invoke"
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing import (
|
|||
Literal,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypedDict,
|
||||
TypeVar,
|
||||
Union,
|
||||
|
|
@ -24,6 +25,7 @@ import httpx
|
|||
from httpx import USE_CLIENT_DEFAULT
|
||||
from httpx._types import FileContent
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
import litellm.litellm_core_utils
|
||||
|
|
@ -43,6 +45,7 @@ from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
|||
SUBTITLE_RESPONSE_FORMATS,
|
||||
synthesize_subtitle_document,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
|
||||
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
|
||||
from litellm.litellm_core_utils.realtime_errors import (
|
||||
|
|
@ -206,6 +209,7 @@ if TYPE_CHECKING:
|
|||
FakeAnthropicMessagesStreamIterator,
|
||||
)
|
||||
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.llms.openai_evals import (
|
||||
CancelEvalResponse,
|
||||
CancelRunResponse,
|
||||
|
|
@ -221,6 +225,21 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class _RealtimeClientWebSocket(Protocol):
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
|
||||
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
|
||||
|
||||
|
||||
class _ResponsesClientWebSocket(Protocol):
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
|
||||
async def receive_text(self) -> str: ...
|
||||
|
||||
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
|
||||
|
||||
|
||||
_ResponseT = TypeVar("_ResponseT")
|
||||
|
||||
|
||||
|
|
@ -237,6 +256,17 @@ class _MediaUploadKwargs(TypedDict, total=False):
|
|||
timeout: float | httpx.Timeout
|
||||
|
||||
|
||||
class _SignedBodyKwargs(TypedDict, total=False):
|
||||
data: ReadOnly[bytes]
|
||||
json: ReadOnly[dict[str, object]]
|
||||
|
||||
|
||||
def _signed_body_kwargs(*, signed_body: bytes | None, data: dict[str, object]) -> _SignedBodyKwargs:
|
||||
if signed_body is not None:
|
||||
return {"data": signed_body}
|
||||
return {"json": data}
|
||||
|
||||
|
||||
def _google_genai_streaming_hidden_params(
|
||||
*,
|
||||
api_base: str,
|
||||
|
|
@ -318,7 +348,9 @@ def _mask_presigned_request_headers(transformed_request: bytes | str | dict) ->
|
|||
}
|
||||
|
||||
|
||||
def _aws_signing_overrides(optional_params: Mapping[str, Any], litellm_params: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
def _aws_signing_overrides(
|
||||
optional_params: Mapping[str, object], litellm_params: Mapping[str, object]
|
||||
) -> Mapping[str, object]:
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: litellm_params[key]
|
||||
|
|
@ -1430,6 +1462,7 @@ class BaseLLMHTTPHandler:
|
|||
transformed: Final = provider_config.transform_audio_transcription_response(
|
||||
raw_response=response,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(transformed, response.headers)
|
||||
if not provider_config.supports_subtitle_synthesis:
|
||||
return transformed
|
||||
requested_format: Final = optional_params.get("response_format")
|
||||
|
|
@ -2739,7 +2772,7 @@ class BaseLLMHTTPHandler:
|
|||
stream=stream,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
|
||||
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -2926,7 +2959,7 @@ class BaseLLMHTTPHandler:
|
|||
stream=stream,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
|
||||
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -4540,7 +4573,7 @@ class BaseLLMHTTPHandler:
|
|||
api_key=litellm_params.api_key,
|
||||
model=model,
|
||||
)
|
||||
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
|
||||
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -4634,7 +4667,7 @@ class BaseLLMHTTPHandler:
|
|||
api_key=litellm_params.api_key,
|
||||
model=model,
|
||||
)
|
||||
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
|
||||
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -6186,6 +6219,7 @@ class BaseLLMHTTPHandler:
|
|||
"BasePassthroughConfig",
|
||||
"BaseContainerConfig",
|
||||
BaseEvalsAPIConfig,
|
||||
BaseRealtimeHTTPConfig,
|
||||
],
|
||||
):
|
||||
received_status_code: Final = (
|
||||
|
|
@ -6300,7 +6334,7 @@ class BaseLLMHTTPHandler:
|
|||
async def async_realtime(
|
||||
self,
|
||||
model: str,
|
||||
websocket: Any,
|
||||
websocket: _RealtimeClientWebSocket,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
provider_config: BaseRealtimeConfig,
|
||||
headers: dict,
|
||||
|
|
@ -6308,7 +6342,7 @@ class BaseLLMHTTPHandler:
|
|||
api_key: str | None = None,
|
||||
client: Any | None = None,
|
||||
timeout: float | None = None,
|
||||
user_api_key_dict: Any | None = None,
|
||||
user_api_key_dict: object | None = None,
|
||||
litellm_metadata: dict[str, object] | None = None,
|
||||
query_params: RealtimeQueryParams | None = None,
|
||||
):
|
||||
|
|
@ -6483,7 +6517,7 @@ class BaseLLMHTTPHandler:
|
|||
request_data: dict[str, object],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
timeout: float | httpx.Timeout,
|
||||
provider_config: Any | None = None,
|
||||
provider_config: BaseRealtimeHTTPConfig | None = None,
|
||||
model: str | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
client: HTTPHandler | AsyncHTTPHandler | None = None,
|
||||
|
|
@ -6555,7 +6589,7 @@ class BaseLLMHTTPHandler:
|
|||
sdp_body: bytes,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
timeout: float | httpx.Timeout,
|
||||
provider_config: Any | None = None,
|
||||
provider_config: BaseRealtimeHTTPConfig | None = None,
|
||||
model: str | None = None,
|
||||
session_config: dict[str, object] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
|
|
@ -6633,13 +6667,13 @@ class BaseLLMHTTPHandler:
|
|||
async def async_responses_websocket(
|
||||
self,
|
||||
model: str,
|
||||
websocket: Any,
|
||||
websocket: _ResponsesClientWebSocket,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
responses_api_provider_config: BaseResponsesAPIConfig | None,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
timeout: float | None = None,
|
||||
user_api_key_dict: Any | None = None,
|
||||
user_api_key_dict: "UserAPIKeyAuth | None" = None,
|
||||
litellm_metadata: dict[str, object] | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
first_message: str | None = None,
|
||||
|
|
@ -6928,11 +6962,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=image_edit_provider_config,
|
||||
)
|
||||
|
||||
return image_edit_provider_config.transform_image_edit_response(
|
||||
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
|
||||
return image_edit_response
|
||||
|
||||
async def async_image_edit_handler(
|
||||
self,
|
||||
|
|
@ -7027,11 +7063,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=image_edit_provider_config,
|
||||
)
|
||||
|
||||
return image_edit_provider_config.transform_image_edit_response(
|
||||
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
|
||||
return image_edit_response
|
||||
|
||||
def image_generation_handler(
|
||||
self,
|
||||
|
|
@ -7154,6 +7192,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
encoding=None,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(model_response, response.headers)
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -7261,6 +7300,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
encoding=None,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(model_response, response.headers)
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -7850,7 +7890,7 @@ class BaseLLMHTTPHandler:
|
|||
def video_create_character_handler(
|
||||
self,
|
||||
name: str,
|
||||
video: Any,
|
||||
video: FileTypes,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
|
|
@ -7934,7 +7974,7 @@ class BaseLLMHTTPHandler:
|
|||
async def async_video_create_character_handler(
|
||||
self,
|
||||
name: str,
|
||||
video: Any,
|
||||
video: FileTypes,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
|
|
@ -12045,11 +12085,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=text_to_speech_provider_config,
|
||||
)
|
||||
|
||||
return text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
|
||||
return speech_response
|
||||
|
||||
async def async_text_to_speech_handler(
|
||||
self,
|
||||
|
|
@ -12144,11 +12186,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=text_to_speech_provider_config,
|
||||
)
|
||||
|
||||
return text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
|
||||
return speech_response
|
||||
|
||||
#########################################################
|
||||
########## SKILLS API HANDLERS ##########################
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm import LlmProviders
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import DEFAULT_MAX_RETRIES
|
||||
from litellm.files.types import FileContentStreamingResult
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
|
|
@ -1404,7 +1405,6 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
organization: str | None = None,
|
||||
headers: dict | None = None,
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
openai_aclient: Final = self._get_openai_client(
|
||||
is_async=True,
|
||||
|
|
@ -1428,8 +1428,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
)
|
||||
|
||||
request_data: Final = {**data, "extra_headers": headers} if headers else data
|
||||
response = await openai_aclient.images.generate(**request_data, timeout=timeout)
|
||||
stringified_response: Final = response.model_dump()
|
||||
raw_response: Final = await openai_aclient.images.with_raw_response.generate(
|
||||
**request_data, timeout=timeout
|
||||
)
|
||||
stringified_response: Final = raw_response.parse().model_dump()
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
input=prompt,
|
||||
|
|
@ -1437,11 +1439,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=stringified_response,
|
||||
)
|
||||
return convert_to_model_response_object(
|
||||
image_response: Final[ImageResponse] = convert_to_model_response_object(
|
||||
response_object=stringified_response,
|
||||
model_response_object=model_response,
|
||||
response_type="image_generation",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
|
||||
return image_response
|
||||
except Exception as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -1512,9 +1516,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
|
||||
## COMPLETION CALL
|
||||
request_data: Final = {**data, "extra_headers": headers} if headers else data
|
||||
_response: Final = openai_client.images.generate(**request_data, timeout=timeout)
|
||||
raw_response: Final = openai_client.images.with_raw_response.generate(**request_data, timeout=timeout)
|
||||
|
||||
response: Final = _response.model_dump()
|
||||
response: Final = raw_response.parse().model_dump()
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
input=prompt,
|
||||
|
|
@ -1522,11 +1526,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=response,
|
||||
)
|
||||
return convert_to_model_response_object(
|
||||
image_response: Final[ImageResponse] = convert_to_model_response_object(
|
||||
response_object=response,
|
||||
model_response_object=model_response,
|
||||
response_type="image_generation",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
|
||||
return image_response
|
||||
except OpenAIError as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -1609,7 +1615,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
input=input,
|
||||
**optional_params,
|
||||
)
|
||||
return HttpxBinaryResponseContent(response=response.response)
|
||||
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
|
||||
return speech_response
|
||||
|
||||
async def async_audio_speech(
|
||||
self,
|
||||
|
|
@ -1655,8 +1663,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
input=input,
|
||||
**optional_params,
|
||||
)
|
||||
|
||||
return HttpxBinaryResponseContent(response=response.response)
|
||||
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
|
||||
return speech_response
|
||||
|
||||
|
||||
class OpenAIFilesAPI(BaseLLM):
|
||||
|
|
|
|||
|
|
@ -994,7 +994,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
def _spread_text_rewrite_over_stream_events(
|
||||
self,
|
||||
stream_events: Sequence[Any],
|
||||
stream_events: Sequence[object],
|
||||
rewritten_text: str,
|
||||
guardrail_name: str,
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -4,11 +4,10 @@ import httpx
|
|||
from openai import AsyncOpenAI, OpenAI
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from aiohttp import ClientSession
|
||||
from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_name
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
BaseAudioTranscriptionConfig,
|
||||
|
|
@ -31,11 +30,6 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
data: dict,
|
||||
timeout: float | httpx.Timeout,
|
||||
):
|
||||
"""
|
||||
Helper to:
|
||||
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
|
||||
- call openai_aclient.audio.transcriptions.create by default
|
||||
"""
|
||||
try:
|
||||
raw_response = await openai_aclient.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
|
||||
headers: Final = dict(raw_response.headers)
|
||||
|
|
@ -51,20 +45,11 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
data: dict,
|
||||
timeout: float | httpx.Timeout,
|
||||
):
|
||||
"""
|
||||
Helper to:
|
||||
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
|
||||
- call openai_aclient.audio.transcriptions.create by default
|
||||
"""
|
||||
try:
|
||||
if litellm.return_response_headers is True:
|
||||
raw_response = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
|
||||
headers: Final = dict(raw_response.headers)
|
||||
response = raw_response.parse()
|
||||
return headers, response
|
||||
else:
|
||||
response = openai_client.audio.transcriptions.create(**data, timeout=timeout)
|
||||
return None, response
|
||||
raw_response: Final = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
|
||||
headers: Final = dict(raw_response.headers)
|
||||
response: Final = raw_response.parse()
|
||||
return headers, response
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
|
@ -133,11 +118,12 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
"complete_input_dict": data,
|
||||
},
|
||||
)
|
||||
_, response = self.make_sync_openai_audio_transcriptions_request(
|
||||
headers, response = self.make_sync_openai_audio_transcriptions_request(
|
||||
openai_client=openai_client,
|
||||
data=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
logging_obj.model_call_details["response_headers"] = headers
|
||||
|
||||
if isinstance(response, BaseModel):
|
||||
stringified_response = response.model_dump()
|
||||
|
|
@ -158,6 +144,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
hidden_params=hidden_params,
|
||||
response_type="audio_transcription",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(final_response, headers)
|
||||
return final_response
|
||||
|
||||
async def async_audio_transcriptions(
|
||||
|
|
@ -217,12 +204,14 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
actual_model: Final = data.get("model", "whisper-1")
|
||||
hidden_params: Final = {"model": actual_model, "custom_llm_provider": "openai"}
|
||||
|
||||
return convert_to_model_response_object(
|
||||
final_response: Final[TranscriptionResponse] = convert_to_model_response_object(
|
||||
response_object=stringified_response,
|
||||
model_response_object=model_response,
|
||||
hidden_params=hidden_params,
|
||||
response_type="audio_transcription",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(final_response, headers)
|
||||
return final_response
|
||||
except Exception as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue