mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Merge remote-tracking branch 'origin/main' into litellm_lit8140_reference_pixel_billing
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled
# Conflicts: # litellm/llms/custom_httpx/llm_http_handler.py
This commit is contained in:
commit
f91b601d8c
678 changed files with 57885 additions and 9924 deletions
|
|
@ -12,6 +12,9 @@ parameters:
|
|||
migration_source_sha:
|
||||
type: string
|
||||
default: ""
|
||||
routing_parity_base:
|
||||
type: string
|
||||
default: ""
|
||||
orbs:
|
||||
codecov: codecov/codecov@4.0.1
|
||||
node: circleci/node@5.1.0 # Add this line to declare the node orb
|
||||
|
|
@ -176,6 +179,9 @@ commands:
|
|||
image:
|
||||
type: string
|
||||
default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
|
||||
server_args:
|
||||
type: string
|
||||
default: ""
|
||||
steps:
|
||||
- run:
|
||||
name: Start PostgreSQL
|
||||
|
|
@ -186,7 +192,7 @@ commands:
|
|||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=<< parameters.db_name >> \
|
||||
-p 5432:5432 \
|
||||
<< parameters.image >>
|
||||
<< parameters.image >> << parameters.server_args >>
|
||||
- wait_for_service:
|
||||
url: tcp://localhost:5432
|
||||
timeout: "60"
|
||||
|
|
@ -3108,6 +3114,10 @@ jobs:
|
|||
parameters:
|
||||
suite:
|
||||
type: string
|
||||
mode:
|
||||
type: enum
|
||||
enum: [standard, replica]
|
||||
default: standard
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
resource_class: large
|
||||
|
|
@ -3142,18 +3152,19 @@ jobs:
|
|||
command: cd ui/litellm-dashboard && NEXT_TELEMETRY_DISABLED=1 npm run build
|
||||
- start_postgres:
|
||||
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
|
||||
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
|
||||
- start_redis
|
||||
- run:
|
||||
name: Run owned integration contracts
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> << parameters.mode >>
|
||||
no_output_timeout: 15m
|
||||
- run:
|
||||
name: Stop owned database and Redis
|
||||
when: always
|
||||
command: |
|
||||
mkdir -p test-results/integration-<< parameters.suite >>
|
||||
docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true
|
||||
docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true
|
||||
mkdir -p test-results/services-<< parameters.suite >>-<< parameters.mode >>
|
||||
docker logs postgres-db > test-results/services-<< parameters.suite >>-<< parameters.mode >>/postgres.log 2>&1 || true
|
||||
docker logs redis-cache > test-results/services-<< parameters.suite >>-<< parameters.mode >>/redis.log 2>&1 || true
|
||||
docker rm -f postgres-db redis-cache
|
||||
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
|
||||
- store_test_results:
|
||||
|
|
@ -3161,6 +3172,76 @@ jobs:
|
|||
- store_artifacts:
|
||||
path: test-results
|
||||
|
||||
routing_parity:
|
||||
parameters:
|
||||
suite:
|
||||
type: string
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
steps:
|
||||
- setup_litellm_test_deps
|
||||
- run:
|
||||
name: Check out base product code
|
||||
environment:
|
||||
ROUTING_PARITY_BASE: << pipeline.parameters.routing_parity_base >>
|
||||
command: |
|
||||
[[ "$ROUTING_PARITY_BASE" =~ ^[0-9a-f]{40}$ ]] || exit 1
|
||||
git fetch --depth 1 origin "$ROUTING_PARITY_BASE"
|
||||
git rm -r -f --quiet litellm enterprise litellm-proxy-extras
|
||||
git checkout "$ROUTING_PARITY_BASE" -- litellm enterprise litellm-proxy-extras
|
||||
git reset --quiet
|
||||
test -f litellm/rust_bridge/_native.abi3.so
|
||||
- start_postgres:
|
||||
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
|
||||
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
|
||||
- start_redis
|
||||
- run:
|
||||
name: Run base side
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> parity base
|
||||
no_output_timeout: 15m
|
||||
- run:
|
||||
name: Stop base database and Redis
|
||||
when: always
|
||||
command: |
|
||||
mkdir -p test-results/services-<< parameters.suite >>-parity-base
|
||||
docker logs postgres-db > test-results/services-<< parameters.suite >>-parity-base/postgres.log 2>&1 || true
|
||||
docker logs redis-cache > test-results/services-<< parameters.suite >>-parity-base/redis.log 2>&1 || true
|
||||
docker rm -f postgres-db redis-cache
|
||||
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
|
||||
- run:
|
||||
name: Check out head product code
|
||||
command: |
|
||||
git rm -r -f --quiet litellm enterprise litellm-proxy-extras
|
||||
git checkout "$CIRCLE_SHA1" -- litellm enterprise litellm-proxy-extras
|
||||
git reset --quiet
|
||||
test -f litellm/rust_bridge/_native.abi3.so
|
||||
- start_postgres:
|
||||
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
|
||||
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
|
||||
- start_redis
|
||||
- run:
|
||||
name: Run head side
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> parity head
|
||||
no_output_timeout: 15m
|
||||
- run:
|
||||
name: Stop head database and Redis
|
||||
when: always
|
||||
command: |
|
||||
mkdir -p test-results/services-<< parameters.suite >>-parity-head
|
||||
docker logs postgres-db > test-results/services-<< parameters.suite >>-parity-head/postgres.log 2>&1 || true
|
||||
docker logs redis-cache > test-results/services-<< parameters.suite >>-parity-head/redis.log 2>&1 || true
|
||||
docker rm -f postgres-db redis-cache
|
||||
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
|
||||
- run:
|
||||
name: Compare routing parity
|
||||
command: PYTHONPATH="$PWD/tests" .venv/bin/python -m integration._support.routing check test-results/parity-<< parameters.suite >>/base test-results/parity-<< parameters.suite >>/head
|
||||
- store_test_results:
|
||||
path: test-results
|
||||
- store_artifacts:
|
||||
path: test-results
|
||||
|
||||
unit:
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
|
|
@ -3224,8 +3305,22 @@ workflows:
|
|||
branches:
|
||||
only: main
|
||||
jobs: *migration_jobs
|
||||
routing_parity:
|
||||
when:
|
||||
not:
|
||||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- routing_parity:
|
||||
name: routing-parity-<< matrix.suite >>
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, accounting, database, providers, extensions, cost, mcp]
|
||||
integration:
|
||||
unless: << pipeline.parameters.run_migration_tests >>
|
||||
unless:
|
||||
or:
|
||||
- << pipeline.parameters.run_migration_tests >>
|
||||
- not:
|
||||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- integration_contracts:
|
||||
name: integration-<< matrix.suite >>
|
||||
|
|
@ -3237,8 +3332,23 @@ workflows:
|
|||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- integration_contracts:
|
||||
name: integration-<< matrix.suite >>-replica
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, database]
|
||||
mode: [replica]
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
build_and_test:
|
||||
unless: << pipeline.parameters.run_migration_tests >>
|
||||
unless:
|
||||
or:
|
||||
- << pipeline.parameters.run_migration_tests >>
|
||||
- not:
|
||||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- using_litellm_on_windows:
|
||||
filters: &main_branches
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
52
.circleci/scripts/prepare_replica_roles.py
Normal file
52
.circleci/scripts/prepare_replica_roles.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
import psycopg
|
||||
|
||||
DATABASE_URL: Final = os.environ["DATABASE_URL"]
|
||||
|
||||
|
||||
def postgres_url() -> str:
|
||||
parsed: Final = urlsplit(DATABASE_URL)
|
||||
return urlunsplit(parsed._replace(path="/postgres"))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
with psycopg.connect(postgres_url(), autocommit=True) as admin:
|
||||
admin.execute("CREATE EXTENSION IF NOT EXISTS pg_stat_statements")
|
||||
admin.execute("CREATE ROLE litellm_writer LOGIN PASSWORD 'litellm-writer' NOSUPERUSER")
|
||||
admin.execute("CREATE ROLE litellm_reader LOGIN PASSWORD 'litellm-reader' NOSUPERUSER NOINHERIT")
|
||||
admin.execute("ALTER ROLE litellm_reader SET default_transaction_read_only = on")
|
||||
admin.execute("ALTER DATABASE circle_test OWNER TO litellm_writer")
|
||||
admin.execute("GRANT CONNECT ON DATABASE circle_test TO litellm_reader")
|
||||
with psycopg.connect(DATABASE_URL, autocommit=True) as admin:
|
||||
admin.execute("GRANT USAGE ON SCHEMA public TO litellm_reader")
|
||||
admin.execute(
|
||||
"ALTER DEFAULT PRIVILEGES FOR ROLE litellm_writer IN SCHEMA public GRANT SELECT ON TABLES TO litellm_reader"
|
||||
)
|
||||
admin.execute("GRANT SELECT ON ALL TABLES IN SCHEMA public TO litellm_reader")
|
||||
|
||||
parsed: Final = urlsplit(DATABASE_URL)
|
||||
reader_url: Final = urlunsplit(
|
||||
parsed._replace(netloc=f"litellm_reader:litellm-reader@{parsed.hostname}:{parsed.port}")
|
||||
)
|
||||
writer_url: Final = urlunsplit(
|
||||
parsed._replace(netloc=f"litellm_writer:litellm-writer@{parsed.hostname}:{parsed.port}")
|
||||
)
|
||||
with psycopg.connect(reader_url, autocommit=True) as reader:
|
||||
assert reader.execute("SHOW transaction_read_only").fetchone() == ("on",)
|
||||
try:
|
||||
reader.execute("CREATE TABLE integration_readonly_probe (id int)")
|
||||
except psycopg.errors.ReadOnlySqlTransaction:
|
||||
pass
|
||||
else:
|
||||
raise AssertionError("litellm_reader executed a write statement")
|
||||
with psycopg.connect(writer_url, autocommit=True) as writer:
|
||||
assert writer.execute("SELECT current_user").fetchone() == ("litellm_writer",)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -7,7 +7,15 @@ if [ "${GITHUB_ACTIONS:-}" = true ]; then
|
|||
fi
|
||||
|
||||
suite="${1:?integration suite required}"
|
||||
results="test-results/integration-${suite}"
|
||||
mode="${2:-standard}"
|
||||
side="${3:-}"
|
||||
if [ "$mode" = replica ]; then
|
||||
results="test-results/integration-${suite}-replica"
|
||||
elif [ "$mode" = parity ]; then
|
||||
results="test-results/parity-${suite}/${side:?parity side required}"
|
||||
else
|
||||
results="test-results/integration-${suite}"
|
||||
fi
|
||||
mkdir -p "$results"
|
||||
integration_identity="$(.venv/bin/python -c 'import uuid; print(uuid.uuid4().hex)')"
|
||||
upstream_pid=""
|
||||
|
|
@ -80,6 +88,18 @@ export INTEGRATION_ORDER_SEED="$INTEGRATION_SEED"
|
|||
|
||||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma > "$results/prisma-generate.log" 2>&1
|
||||
|
||||
export INTEGRATION_PROXY_DATABASE_URL=""
|
||||
export INTEGRATION_PROXY_READ_REPLICA_URL=""
|
||||
export INTEGRATION_ROUTING=""
|
||||
if [ "$mode" = replica ] || [ "$mode" = parity ]; then
|
||||
.venv/bin/python .circleci/scripts/prepare_replica_roles.py > "$results/prepare-replica-roles.log" 2>&1
|
||||
export INTEGRATION_PROXY_DATABASE_URL="postgresql://litellm_writer:litellm-writer@127.0.0.1:5432/circle_test"
|
||||
export INTEGRATION_PROXY_READ_REPLICA_URL="postgresql://litellm_reader:litellm-reader@127.0.0.1:5432/circle_test"
|
||||
fi
|
||||
if [ "$mode" = parity ]; then
|
||||
export INTEGRATION_ROUTING=capture
|
||||
fi
|
||||
|
||||
sudo iptables -N integration_only
|
||||
guard_created=true
|
||||
sudo iptables -A integration_only -o lo -j ACCEPT
|
||||
|
|
@ -137,8 +157,12 @@ start_proxy() {
|
|||
else
|
||||
cost_map_env=("LITELLM_LOCAL_MODEL_COST_MAP=True")
|
||||
fi
|
||||
local -a database_env=("DATABASE_URL=${INTEGRATION_PROXY_DATABASE_URL:-$DATABASE_URL}")
|
||||
if [ -n "$INTEGRATION_PROXY_READ_REPLICA_URL" ]; then
|
||||
database_env+=("DATABASE_URL_READ_REPLICA=$INTEGRATION_PROXY_READ_REPLICA_URL")
|
||||
fi
|
||||
setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \
|
||||
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
|
||||
"${database_env[@]}" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
|
||||
INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \
|
||||
LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" LITELLM_UI_PATH="$LITELLM_UI_PATH" PROXY_BASE_URL="http://127.0.0.1:$port" \
|
||||
LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True "${cost_map_env[@]}" \
|
||||
|
|
@ -195,6 +219,9 @@ env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
|
|||
INTEGRATION_SEED="$INTEGRATION_SEED" \
|
||||
INTEGRATION_ORDER_SEED="$INTEGRATION_ORDER_SEED" \
|
||||
LITELLM_LOCAL_MODEL_COST_MAP=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
|
||||
INTEGRATION_PROXY_DATABASE_URL="$INTEGRATION_PROXY_DATABASE_URL" \
|
||||
INTEGRATION_PROXY_READ_REPLICA_URL="$INTEGRATION_PROXY_READ_REPLICA_URL" \
|
||||
INTEGRATION_ROUTING="$INTEGRATION_ROUTING" \
|
||||
.venv/bin/python tests/integration/run.py "$suite" --results "$results"
|
||||
|
||||
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then
|
||||
|
|
|
|||
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 >>
|
||||
|
|
|
|||
166
.github/pull_request_template.md
vendored
166
.github/pull_request_template.md
vendored
|
|
@ -1,27 +1,167 @@
|
|||
<!-- Plain English please. Describe the change the way you would explain it to a teammate who has not seen the code: what it does and why, not which functions, files, or tables it touches -->
|
||||
<!-- The whole description's target audience is humans, not AI agents: write it in plain, simple,
|
||||
everyday engineering language, extremely parsable and readable at a glance. This goes double for
|
||||
the TLDR, User Flow, and Caveats sections
|
||||
Drop every section you have nothing to put in, heading included: a bare "## Relevant issues" or
|
||||
"## Affected release" with nothing under it must not appear in the final description -->
|
||||
|
||||
## What's the problem?
|
||||
## TLDR
|
||||
|
||||
## What's the solution?
|
||||
<!-- 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 -->
|
||||
|
||||
<!-- The approach, in a sentence or two -->
|
||||
Problem this solves:
|
||||
|
||||
## How does it fix it?
|
||||
- <blah>
|
||||
- ...
|
||||
|
||||
<!-- What actually changes so the problem can't happen anymore -->
|
||||
How it solves it:
|
||||
|
||||
## How does the product experience change?
|
||||
- <blah>
|
||||
- ...
|
||||
|
||||
<!-- What a user could do or see before, and what they can do or see after. If nothing user-facing changes, say so -->
|
||||
## User Flow
|
||||
|
||||
## What caveats are there, if any?
|
||||
<!-- Two ordered lists, Before and After, walking the same end user through the same task, written strictly from that user's seat
|
||||
Read the linked issue, ticket, or customer thread first so the flow reflects the real application and the routes its users actually hit; don't invent a generic scenario
|
||||
Lead each list with one plain sentence saying where the flow fails (Before) or succeeds (After), then number the steps
|
||||
Keep it tight: aim for 3 to 5 steps per list, one line each, roughly 20 words max, and never pad a shorter flow with filler steps to hit the count. Cover the one path the PR changes and fold variants (case, other field, second endpoint) into a clause on the step they belong to rather than their own steps. The example below is the target length
|
||||
Every step is something the user does or observes: the HTTP method and full URL they hit, what they sent, and what visibly came back (status code, error text, the shape of an ID). UI steps name the page URL and what is on screen
|
||||
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, 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
|
||||
|
||||
<!-- Major ones only: behavior that breaks on purpose, migrations that lock tables, auth changes, known gaps you did not fix. Write "None" if there are none -->
|
||||
Example:
|
||||
|
||||
Before: a developer whose app streams chat completions gets no token counts back, so their cost dashboard reads zero
|
||||
|
||||
1. They send POST https://litellm-domain/v1/chat/completions with `"stream": true` and no `stream_options`
|
||||
2. The last SSE chunk arrives with `"usage": null`, so their app records 0 prompt and 0 completion tokens
|
||||
3. They open https://litellm-domain/ui/?page=logs and see the request logged at $0 spend
|
||||
|
||||
After: the same request comes back with real token counts, so the dashboard shows real spend
|
||||
|
||||
1. The proxy admin sets `always_include_stream_usage: true` and restarts the proxy
|
||||
2. The developer sends the same POST https://litellm-domain/v1/chat/completions with `"stream": true` and no `stream_options`
|
||||
3. The last SSE chunk now carries a `usage` object with real prompt and completion token counts
|
||||
4. https://litellm-domain/ui/?page=logs shows that request at non-zero spend
|
||||
-->
|
||||
|
||||
## Relevant issues
|
||||
|
||||
<!-- e.g., "Fixes #000". Drop the section if there is none -->
|
||||
|
||||
## Affected release
|
||||
|
||||
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1", and add the `backport-stable` label so the fix is cherry-picked onto the rc line before the stable is tagged. Drop the section otherwise -->
|
||||
|
||||
## Linear ticket
|
||||
|
||||
<!-- Internal contributors: Resolves LIT-1234. Otherwise leave blank -->
|
||||
<!-- if you are an internal contributor, add "Resolves " followed by the Linear ticket e.g., "Resolves LIT-1234" to link the Linear ticket to the GitHub PR. If you don't have one, drop the section rather than guessing -->
|
||||
|
||||
## How did you test this?
|
||||
## Pre-Submission checklist
|
||||
|
||||
**Please complete all items before asking a LiteLLM maintainer to review your PR**
|
||||
|
||||
- [ ] I have added meaningful tests
|
||||
- [ ] The handful of test files covering my change pass locally, e.g. `uv run pytest tests/test_litellm/<your_test_file>.py -v`. Leave the suites (`make test-unit-*`, `make test-unit`) to CI: it finishes in ~15 minutes where a laptop takes an hour or more
|
||||
- [ ] My PR passes all required CI/CD checks (e.g., lint, schema.d.ts sync check, etc.)
|
||||
- [ ] My PR's scope is as isolated as possible; it only solves 1 specific problem
|
||||
- [ ] I have received a Greptile **Confidence Score of at least 4/5** before requesting a maintainer review (Greptile reviews automatically once the PR is opened; only comment `@greptileai` to re-request a review after pushing changes)
|
||||
|
||||
## Delays in PR merge?
|
||||
|
||||
If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slack (#pr-review)](https://join.slack.com/t/litellmossslack/shared_invite/zt-3o7nkuyfr-p_kbNJj8taRfXGgQI1~YyA).
|
||||
|
||||
## Screenshots / Proof of Fix
|
||||
|
||||
<!-- Include screenshots, screen recordings, or command (e.g., curl) + output demonstrating that your changes work as expected
|
||||
The proof must be completely e2e with no mocks, using actual LLM calls costing real $$$ if applicable. `pytest` commands are not enough
|
||||
Show ONLY the latest run: capture Before at the merge base and After at the PR's current tip, and when new commits change behavior, replace this whole section with the fresh run instead of stacking it on top of older ones. The run must be up to date. As soon as a new commit is made and it makes this PR description's after sha stale (it's no longer tip of PR), you must re-run the QA
|
||||
Structure the section exactly as below: Before and After one heading level below this section, each naming the commit hash it was captured at, one lower-level heading per case inside each, the same case names in the same order on both sides, and numbered steps (command, observed output) under every case, never loose prose; shared setup (config, payloads) goes above Before, and with a single case, drop the case headings and number the steps directly
|
||||
|
||||
### Before (<hash>)
|
||||
|
||||
#### <case 1>
|
||||
|
||||
1. ...
|
||||
2. ...
|
||||
|
||||
#### <case 2>
|
||||
|
||||
1. ...
|
||||
|
||||
### After (<hash>)
|
||||
|
||||
#### <case 1>
|
||||
|
||||
1. ...
|
||||
2. ...
|
||||
|
||||
#### <case 2>
|
||||
|
||||
1. ...
|
||||
|
||||
For bug fixes: Before shows the reproduction, After shows the same steps passing
|
||||
For new features: Before shows the capability missing, After shows it working end-to-end
|
||||
If the change applies to all three LLM endpoints (/v1/responses, /v1/chat/completions, /v1/messages), make each endpoint its own case, not just one
|
||||
For UI changes: before/after screenshots under the same headings
|
||||
If the main use case runs through a coding tool like Claude Code or Codex, drive that tool interactively the way the user does (never `claude -p`, `codex exec`, or curl on its own) and embed before/after screenshots of its pane under the same headings; curl replays and headless runs can follow as extra cases, never as the only proof -->
|
||||
|
||||
## Type
|
||||
|
||||
<!-- Select the type of Pull Request -->
|
||||
<!-- Keep only the necessary ones -->
|
||||
|
||||
🆕 New Feature
|
||||
🐛 Bug Fix
|
||||
🧹 Refactoring
|
||||
📖 Documentation
|
||||
🚄 Infrastructure
|
||||
✅ Test
|
||||
|
||||
## Caveats (if any)
|
||||
|
||||
<!-- Group caveats under severity subheadings (### Severe, ### High, ### Medium, ### Low), with
|
||||
short bullet points inside each, just like the TLDR: one line per bullet, roughly 10 words max
|
||||
Call out known limitations, follow-up work, or anything a reviewer should watch out for
|
||||
Include only the tiers that have caveats; drop the empty ones
|
||||
- Severe: inherent to what the PR deliberately ships, there even when the code works as intended:
|
||||
it can degrade or take down a running deployment (e.g. a slow or table-locking boot migration),
|
||||
rewrite data by design, break an existing workflow on purpose, or change auth behavior. An
|
||||
operator must plan around it before rollout
|
||||
- High: an unintended hole: a correctness, security, data-loss, or backward-compatibility bug,
|
||||
unsafe to ship as is
|
||||
- Medium: a real gap someone can hit, but with a workaround or a narrow blast radius
|
||||
- Low: anything else worth noting: naming, cleanup, an edge case nobody hits
|
||||
Nest bullets as deep as helps: hierarchy beats one long line when it makes things clearer to a
|
||||
human reader
|
||||
If you assumed something instead of testing it, e.g. "only reproduces with X on" or "no
|
||||
user-observable behavior difference", list it here too with what breaks if it is wrong
|
||||
Drop this section if there are none -->
|
||||
|
||||
## QA runbook
|
||||
|
||||
<!-- Only needed when your PR edits tests/e2e; delete this section otherwise
|
||||
|
||||
For each e2e test you added or changed, list the manual steps a reviewer can follow to reproduce it by hand against a live proxy, mapping 1:1 to what the test asserts: one top-level bullet per test giving its pytest node id followed by what it proves in plain words, then a nested "- [ ]" checklist where each item is a concrete action (route, request body, expected response) and the final item is the sanity-check step shown in the examples. Note environment prerequisites (provider credentials, config flags) and any nuances a manual run will hit. See PRs #32914 and #32963 for full examples
|
||||
|
||||
Example checklists:
|
||||
|
||||
- tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py::TestKeyRateLimits::test_rpm_limit_blocks_over_limit - a key allowed 2 requests a minute serves exactly 2 and refuses the 3rd
|
||||
- [ ] Generate a limited key: curl -X POST http://localhost:4000/key/generate -H "Authorization: Bearer sk-1234" -d '{"rpm_limit": 2}'
|
||||
- [ ] Send three /v1/chat/completions requests with that key inside one minute
|
||||
- [ ] Expect the first two to return 200 and the third to return 429 naming the rpm limit
|
||||
- [ ] Sanity check: this test makes sense to add and is not hand-wavey (e.g., assert actual expected spend instead of just spend > 0) or potentially flaky
|
||||
|
||||
- tests/e2e/management/test_management_e2e.py::TestModelRoutes::test_model_create_appears_in_ui - a deployment created through the API shows up on the Admin UI models page
|
||||
- [ ] POST /model/new with the master key, a bedrock model, and aws_region_name (needs STORE_MODEL_IN_DB=True and AWS credentials)
|
||||
- [ ] Open http://localhost:4000/ui/?page=models and expect a deployment row showing the returned model id
|
||||
- [ ] Sanity check: this test makes sense to add and is not hand-wavey (e.g., assert actual expected spend instead of just spend > 0) or potentially flaky
|
||||
-->
|
||||
|
||||
## Final Attestation
|
||||
|
||||
- [ ] The tests check the right things, including the edge cases, and regressions in the respective real-world customer use-cases are not possible after this PR
|
||||
|
||||
<!-- What you ran against a live proxy and what came back. Commands with output or before/after screenshots work well. Unit tests alone are not enough -->
|
||||
|
|
|
|||
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 }}
|
||||
|
|
|
|||
|
|
@ -33,13 +33,13 @@ End-to-end tests belong in `tests/e2e/` and must follow the harness conventions
|
|||
|
||||
When creating PRs, target the repository's current default branch for both internal and external / OSS contributions. Check it with `python3 scripts/default_branch.py --branch` instead of assuming a branch name or relying on cached `origin/HEAD`
|
||||
|
||||
When writing a PR body, treat the comments and imperative instructions inside .github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule
|
||||
When writing a PR body, treat the comments and imperative instructions inside .github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule. A section you have nothing to put in (Relevant issues, Affected release, Linear ticket, Caveats, QA runbook, and so on) is removed entirely, heading included, never left as an empty title
|
||||
|
||||
Same applies for filing bug reports and feature requests, with .github/ISSUE_TEMPLATE/bug_report.yml and .github/ISSUE_TEMPLATE/feature_request.yml, respectively
|
||||
|
||||
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank
|
||||
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just drop the section
|
||||
|
||||
Never use `pytest` commands or the like as the answer to "How did you test this?". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it
|
||||
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it
|
||||
|
||||
If you ever write any human-facing text (pull requests, issues, commit messages, discussion posts, github comments, release notes, docs, etc.), always follow these guidelines to sound less AI-y:
|
||||
- don't use emojis
|
||||
|
|
|
|||
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"
|
||||
|
|
|
|||
43
db_scripts/backfill_key_total_spend.sql
Normal file
43
db_scripts/backfill_key_total_spend.sql
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
-- One-shot backfill of LiteLLM_VerificationToken.total_spend (lifetime spend)
|
||||
-- for keys created before the column was introduced in LiteLLM v1.103.0.
|
||||
--
|
||||
-- The column was added with DEFAULT 0 and no backfill, so keys that predate
|
||||
-- the upgrade report lifetime spend below their current period spend. New
|
||||
-- deployments do not need this script: total_spend is updated at request
|
||||
-- time from the moment the release is deployed. Run it only if you want
|
||||
-- pre-upgrade keys to show their historical lifetime spend. It sets lifetime
|
||||
-- spend to at least the current spend on every key, active and archived,
|
||||
-- because current period spend is a valid lower bound on lifetime spend.
|
||||
-- For keys with no budget reset that is already the exact lifetime value;
|
||||
-- for resetting keys it only recovers the current period. It is idempotent:
|
||||
-- it only touches rows where total_spend is below spend, so re-running is a
|
||||
-- no-op. It touches no spend logs and runs in seconds.
|
||||
--
|
||||
-- IMPORTANT caveats before running:
|
||||
--
|
||||
-- 1. Take a backup of the affected tables first:
|
||||
-- pg_dump "$DATABASE_URL" -t '"LiteLLM_VerificationToken"' -t '"LiteLLM_DeletedVerificationToken"' > key_total_spend_backup.sql
|
||||
--
|
||||
-- 2. A key "resets" when its own budget_duration IS NOT NULL, or when its
|
||||
-- budget_id links to a LiteLLM_BudgetTable row whose budget_duration IS
|
||||
-- NOT NULL (a linked budget resets the key's spend each period too). For
|
||||
-- those keys this script only recovers the current period;
|
||||
-- db_scripts/backfill_key_total_spend_from_spend_logs.sql is an optional
|
||||
-- follow-up that rebuilds the earlier periods from LiteLLM_SpendLogs.
|
||||
--
|
||||
-- 3. No proxy restart is needed. The proxy picks up the corrected values on
|
||||
-- its next read of each key.
|
||||
--
|
||||
-- Usage:
|
||||
-- psql "$DATABASE_URL" -f db_scripts/backfill_key_total_spend.sql
|
||||
|
||||
UPDATE "LiteLLM_VerificationToken"
|
||||
SET total_spend = spend
|
||||
WHERE total_spend < spend;
|
||||
|
||||
UPDATE "LiteLLM_DeletedVerificationToken"
|
||||
SET total_spend = spend
|
||||
WHERE total_spend < spend;
|
||||
|
||||
-- Verify: this should return 0.
|
||||
-- SELECT count(*) FROM "LiteLLM_VerificationToken" WHERE total_spend < spend;
|
||||
89
db_scripts/backfill_key_total_spend_from_spend_logs.sql
Normal file
89
db_scripts/backfill_key_total_spend_from_spend_logs.sql
Normal file
|
|
@ -0,0 +1,89 @@
|
|||
-- Optional follow-up to db_scripts/backfill_key_total_spend.sql. Run that
|
||||
-- script first; this one rebuilds earlier budget periods for the keys it
|
||||
-- can only partially fix: keys whose spend resets each period, because their own
|
||||
-- budget_duration IS NOT NULL or because their budget_id links to a
|
||||
-- LiteLLM_BudgetTable row whose budget_duration IS NOT NULL.
|
||||
--
|
||||
-- For those keys the "spend" column only covers the current period, so
|
||||
-- lifetime spend is reconstructed from LiteLLM_SpendLogs. The join matches
|
||||
-- l.api_key against both the stored token and its second sha256
|
||||
-- (encode(sha256(convert_to(token, 'UTF8')), 'hex')), because spend logs
|
||||
-- written by older paths recorded the re-hashed digest instead of the
|
||||
-- token. It is idempotent and never lowers a value: every statement only
|
||||
-- touches rows where total_spend is below the rebuilt sum, so re-running is
|
||||
-- a no-op, and a key whose log history is shorter than its current period
|
||||
-- keeps the value backfill_key_total_spend.sql already gave it.
|
||||
--
|
||||
-- IMPORTANT caveats before running:
|
||||
--
|
||||
-- 1. Take a backup of the affected tables first:
|
||||
-- pg_dump "$DATABASE_URL" -t '"LiteLLM_VerificationToken"' -t '"LiteLLM_DeletedVerificationToken"' > key_total_spend_backup.sql
|
||||
--
|
||||
-- 2. It requires spend logs to have been enabled, and coverage is bounded
|
||||
-- by maximum_spend_logs_retention_period: spend older than the retention
|
||||
-- window is already gone and cannot be recovered.
|
||||
--
|
||||
-- 3. On a large SpendLogs table the join scan is slow, so run it off peak.
|
||||
--
|
||||
-- 4. Run it while the proxy is idle (or with traffic paused). The proxy
|
||||
-- flushes spend logs in batches, so a request that already raised
|
||||
-- total_spend but whose log is still queued is missing from the sum, and
|
||||
-- the rebuilt value would be short by that in-flight amount.
|
||||
--
|
||||
-- 5. A custom token can be deleted and recreated, so the archived table can
|
||||
-- hold several lifetimes of one token. The update only rewrites archived
|
||||
-- rows that reset, and the log sum covers every lifetime of that token.
|
||||
--
|
||||
-- 6. No proxy restart is needed. The proxy picks up the corrected values on
|
||||
-- its next read of each key.
|
||||
--
|
||||
-- Usage:
|
||||
-- psql "$DATABASE_URL" -f db_scripts/backfill_key_total_spend_from_spend_logs.sql
|
||||
|
||||
-- Active keys whose spend resets (own budget_duration, or a linked
|
||||
-- LiteLLM_BudgetTable row with one). Rebuild from LiteLLM_SpendLogs,
|
||||
-- matching api_key against the stored token and its second sha256 digest.
|
||||
UPDATE "LiteLLM_VerificationToken" k
|
||||
SET total_spend = s.sum_spend
|
||||
FROM (
|
||||
SELECT k2.token, SUM(l.spend) AS sum_spend
|
||||
FROM "LiteLLM_VerificationToken" k2
|
||||
JOIN "LiteLLM_SpendLogs" l
|
||||
ON l.api_key IN (k2.token, encode(sha256(convert_to(k2.token, 'UTF8')), 'hex'))
|
||||
WHERE k2.budget_duration IS NOT NULL
|
||||
OR k2.budget_id IN (
|
||||
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
|
||||
)
|
||||
GROUP BY k2.token
|
||||
) s
|
||||
WHERE k.token = s.token
|
||||
AND k.total_spend < s.sum_spend;
|
||||
|
||||
-- Archived tokens are not unique, so collapse them to one row per token
|
||||
-- before joining spend logs; the update then hits every resetting archived
|
||||
-- row.
|
||||
UPDATE "LiteLLM_DeletedVerificationToken" k
|
||||
SET total_spend = s.sum_spend
|
||||
FROM (
|
||||
SELECT k2.token, SUM(l.spend) AS sum_spend
|
||||
FROM (
|
||||
SELECT DISTINCT token
|
||||
FROM "LiteLLM_DeletedVerificationToken"
|
||||
WHERE budget_duration IS NOT NULL
|
||||
OR budget_id IN (
|
||||
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
|
||||
)
|
||||
) k2
|
||||
JOIN "LiteLLM_SpendLogs" l
|
||||
ON l.api_key IN (k2.token, encode(sha256(convert_to(k2.token, 'UTF8')), 'hex'))
|
||||
GROUP BY k2.token
|
||||
) s
|
||||
WHERE k.token = s.token
|
||||
AND k.total_spend < s.sum_spend
|
||||
AND (k.budget_duration IS NOT NULL
|
||||
OR k.budget_id IN (
|
||||
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
|
||||
));
|
||||
|
||||
-- Verify: this should return 0.
|
||||
-- SELECT count(*) FROM "LiteLLM_VerificationToken" WHERE total_spend < spend;
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "kill_switch" JSONB;
|
||||
|
|
@ -72,6 +72,7 @@ model LiteLLM_AgentsTable {
|
|||
agent_card_params Json
|
||||
static_headers Json? @default("{}")
|
||||
extra_headers String[] @default([])
|
||||
kill_switch Json?
|
||||
agent_access_groups String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
object_permission_id String?
|
||||
|
|
|
|||
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")
|
||||
|
|
@ -245,6 +245,8 @@ pub struct ModelInfo {
|
|||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_272k_tokens_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_32k_tokens: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -283,6 +285,8 @@ pub struct ModelInfo {
|
|||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_272k_tokens_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_32k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_512k_tokens: Option<f64>,
|
||||
|
|
@ -377,6 +381,8 @@ pub struct ModelInfo {
|
|||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_272k_tokens_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_32k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_512k_tokens: Option<f64>,
|
||||
|
|
@ -498,6 +504,8 @@ pub struct ModelInfo {
|
|||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_272k_tokens_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_32k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_512k_tokens: Option<f64>,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -631,9 +631,9 @@ class LevelRoutingStreamHandler(logging.StreamHandler):
|
|||
)
|
||||
preferred: Final = sys.stdout if is_stdout_record else sys.stderr
|
||||
if preferred is None or getattr(preferred, "closed", False):
|
||||
self.stream = sys.stderr # rebind-ok: fall back to the pre-fix stream rather than raising per record
|
||||
self.stream = sys.stderr
|
||||
else:
|
||||
self.stream = preferred # rebind-ok: StreamHandler.emit writes self.stream under the handler lock
|
||||
self.stream = preferred
|
||||
super().emit(record)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -429,6 +429,7 @@ class LLMCachingHandler:
|
|||
kwargs=kwargs,
|
||||
cached_result=cached_result,
|
||||
is_async=False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
if not _should_defer_streaming_cache_hit_callbacks(cached_result=cached_result):
|
||||
|
|
|
|||
|
|
@ -191,8 +191,6 @@ def _as_chat_reasoning_items(
|
|||
) -> list[ChatCompletionReasoningItem] | None:
|
||||
if not reasoning_items:
|
||||
return None
|
||||
# cast-ok: _BuiltReasoningItem is the structural shape ChatCompletionReasoningItem
|
||||
# describes, and TypedDict invariance is what stops the two from unifying here.
|
||||
return cast(list[ChatCompletionReasoningItem], list(reasoning_items))
|
||||
|
||||
|
||||
|
|
@ -1370,7 +1368,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
if tool_call_index_map is None:
|
||||
return output_index
|
||||
if output_index not in tool_call_index_map:
|
||||
tool_call_index_map[output_index] = len(tool_call_index_map) # mutable-ok: per-stream accumulator state
|
||||
tool_call_index_map[output_index] = len(tool_call_index_map)
|
||||
return tool_call_index_map[output_index]
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
|
|||
DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
|
||||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))
|
||||
DEFAULT_S3_BATCH_SIZE: Final = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512))
|
||||
DEFAULT_S3_MAX_CONCURRENT_UPLOADS: Final = int(os.getenv("DEFAULT_S3_MAX_CONCURRENT_UPLOADS", "16"))
|
||||
# https://docs.aws.amazon.com/AmazonS3/latest/userguide/object-keys.html
|
||||
MAX_S3_OBJECT_KEY_BYTES: Final = 1024
|
||||
S3_BOUNDED_OBJECT_KEY_HEAD_BYTES: Final = 64
|
||||
|
|
@ -556,6 +557,8 @@ SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS: Final[float] = float(
|
|||
request_timeout: float = float(os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS))))
|
||||
request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ
|
||||
DEFAULT_A2A_AGENT_TIMEOUT: Final[float] = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes
|
||||
AGENT_KILL_SWITCH_TIMEOUT_SECONDS: Final = 10.0
|
||||
AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS: Final = 2000
|
||||
# Patterns that indicate a localhost/internal URL in A2A agent cards that should be
|
||||
# replaced with the original base_url. This is a common misconfiguration where
|
||||
# developers deploy agents with development URLs in their agent cards.
|
||||
|
|
@ -1518,6 +1521,7 @@ PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES: Final = int(
|
|||
CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60))
|
||||
MCP_TOOL_NAME_PREFIX: Final = "mcp_tool"
|
||||
MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))
|
||||
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: Final = 4096
|
||||
|
||||
# Headers to control callbacks
|
||||
X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"
|
||||
|
|
@ -1770,6 +1774,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")))
|
||||
|
|
|
|||
|
|
@ -2018,8 +2018,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
|
||||
declared: Final = MappingProxyType(
|
||||
{key: value for key in _DEPLOYMENT_PRICING_KEYS if (value := litellm_params.get(key)) is not None}
|
||||
|
|
@ -2042,7 +2042,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)
|
||||
|
|
|
|||
|
|
@ -764,7 +764,7 @@ class MCPClient:
|
|||
follow_redirects=True,
|
||||
event_hooks=MappingProxyType(
|
||||
{"response": [capture_upstream_error_response], "request": [guard] if guard else []}
|
||||
), # mutable-ok: httpx types require lists of hooks
|
||||
),
|
||||
)
|
||||
|
||||
return factory
|
||||
|
|
@ -921,9 +921,7 @@ class MCPClient:
|
|||
with anyio.fail_after(max(self.timeout, MCP_TOOL_LISTING_TIMEOUT)):
|
||||
for page_index in range(MCP_TOOL_LISTING_MAX_PAGES):
|
||||
try:
|
||||
page = await fetch_page( # rebind-ok: each SDK page replaces the previous one
|
||||
None if cursor is None else PaginatedRequestParams(cursor=cursor)
|
||||
)
|
||||
page = await fetch_page(None if cursor is None else PaginatedRequestParams(cursor=cursor))
|
||||
except MCPError as error:
|
||||
if page_index > 0 and error.error.code == METHOD_NOT_FOUND:
|
||||
raise RuntimeError("MCP list operation became unavailable during pagination") from error
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -1641,5 +1641,5 @@ def log_guardrail_information(func):
|
|||
return async_wrapper(*args, **kwargs)
|
||||
return sync_wrapper(*args, **kwargs)
|
||||
|
||||
vars(wrapper)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True # rebind-ok: stamps the wrapper this call just built
|
||||
vars(wrapper)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True
|
||||
return wrapper
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -366,9 +366,7 @@ class NewRelicMetricsLogger(CustomBatchLogger):
|
|||
error to keep the client-error path (drop) distinct from 5xx (retry)."""
|
||||
payload: Final = build_metric_payload(records=batch, window_start=window_start, now=time.time())
|
||||
try:
|
||||
status = (
|
||||
await self.async_send_compressed_data(payload)
|
||||
).status_code # rebind-ok: reassigned from the raised HTTPStatusError below
|
||||
status = (await self.async_send_compressed_data(payload)).status_code
|
||||
except HTTPStatusError as e:
|
||||
status = e.response.status_code
|
||||
except Exception as e: # noqa: BLE001 # transport/network failure re-queues the batch
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -353,7 +353,7 @@ class _DrainPool:
|
|||
|
||||
def _drain_until_closed(self) -> None:
|
||||
while True:
|
||||
processor: SpanProcessor | None = self._pending.get() # rebind-ok: loop variable
|
||||
processor: SpanProcessor | None = self._pending.get()
|
||||
if processor is None:
|
||||
return
|
||||
_shutdown_quietly(processor)
|
||||
|
|
@ -572,7 +572,7 @@ class TenantFanOutSpanProcessor(SpanProcessor):
|
|||
span, destination.span_scope
|
||||
):
|
||||
continue
|
||||
processor = self._acquire(destination) # rebind-ok: loop variable; pyright forbids Final in a loop
|
||||
processor = self._acquire(destination)
|
||||
if processor is None:
|
||||
continue
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -151,7 +151,7 @@ def destination_for(
|
|||
endpoint, protocol = resolved
|
||||
return OtelDestination(
|
||||
endpoint=endpoint,
|
||||
headers=MappingProxyType(dict(headers)), # mutable-ok: MappingProxyType needs a concrete mapping to wrap
|
||||
headers=MappingProxyType(dict(headers)),
|
||||
resource_attributes=MappingProxyType({"service.name": service_name}) if service_name else _NO_ATTRS,
|
||||
callback_name=callback_name,
|
||||
protocol=protocol,
|
||||
|
|
|
|||
|
|
@ -131,7 +131,7 @@ def _paginated_table(repository: BaseRepository[_TableRowT]) -> _PaginatedPrisma
|
|||
"""View a repository's prisma table through the pagination surface budget metrics need."""
|
||||
return cast(
|
||||
_PaginatedPrismaTable[_TableRowT],
|
||||
repository.table, # cast-ok: prisma rows carry the budget columns the domain model declares
|
||||
repository.table,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -20,7 +20,8 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
_S3_LOG_PROMPTS_ONLY: Final = TypeAdapter(bool)
|
||||
_S3_BOOL: Final = TypeAdapter(bool)
|
||||
_UPLOAD_BOUND: Final = TypeAdapter(int)
|
||||
|
||||
|
||||
def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] | None = None) -> bool:
|
||||
|
|
@ -29,12 +30,42 @@ def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] |
|
|||
if raw is None or raw == "":
|
||||
return False
|
||||
try:
|
||||
return _S3_LOG_PROMPTS_ONLY.validate_python(raw.strip() if isinstance(raw, str) else raw)
|
||||
return _S3_BOOL.validate_python(raw.strip() if isinstance(raw, str) else raw)
|
||||
except ValidationError:
|
||||
verbose_logger.warning("s3 logging: s3_log_prompts_only=%r is not a boolean, logging prompts only", raw)
|
||||
return True
|
||||
|
||||
|
||||
def resolve_s3_max_concurrent_uploads(configured: object, fallback: int) -> int:
|
||||
if configured is None or configured == "":
|
||||
return fallback
|
||||
try:
|
||||
bound: Final = _UPLOAD_BOUND.validate_python(configured.strip() if isinstance(configured, str) else configured)
|
||||
except ValidationError:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_max_concurrent_uploads=%r is not an integer, using %s", configured, fallback
|
||||
)
|
||||
return fallback
|
||||
if bound < 1:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_max_concurrent_uploads=%r must be at least 1, using %s", configured, fallback
|
||||
)
|
||||
return fallback
|
||||
return bound
|
||||
|
||||
|
||||
def resolve_s3_batch_file_upload(configured: object) -> bool:
|
||||
if configured is None or configured == "":
|
||||
return False
|
||||
try:
|
||||
return _S3_BOOL.validate_python(configured.strip() if isinstance(configured, str) else configured)
|
||||
except ValidationError:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_batch_file_upload=%r is not a boolean, keeping per-request objects", configured
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def prompts_only_payload(payload: StandardLoggingPayload) -> StandardLoggingPayload:
|
||||
return {**payload, "response": None}
|
||||
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue