mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge upstream main into fix/responses-reasoning-effort-capability
This commit is contained in:
commit
acfbdc2cc0
1674 changed files with 39545 additions and 19219 deletions
|
|
@ -323,6 +323,7 @@ jobs:
|
|||
CHOCOLATEY_CONFIRM_ALL: "true"
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
no_output_timeout: 30m
|
||||
environment:
|
||||
UV_HTTP_TIMEOUT: "300"
|
||||
command: |
|
||||
|
|
@ -381,6 +382,7 @@ jobs:
|
|||
uv run --no-sync python -m pytest tests/windows_tests/ -v
|
||||
- run:
|
||||
name: Guard against MAX_PATH-busting packaged wheel paths
|
||||
no_output_timeout: 30m
|
||||
environment:
|
||||
UV_HTTP_TIMEOUT: "300"
|
||||
command: |
|
||||
|
|
@ -3327,22 +3329,12 @@ workflows:
|
|||
matrix:
|
||||
parameters:
|
||||
suite: [management, accounting, database, providers, extensions, mcp, sdk, cost, browser]
|
||||
filters:
|
||||
branches:
|
||||
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:
|
||||
or:
|
||||
|
|
@ -3350,101 +3342,60 @@ workflows:
|
|||
- not:
|
||||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- using_litellm_on_windows:
|
||||
filters: &main_branches
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- unit:
|
||||
filters: *main_branches
|
||||
- using_litellm_on_windows
|
||||
- unit
|
||||
- provider_replay_harness
|
||||
- base_sdk_install:
|
||||
filters: *main_branches
|
||||
- local_testing_part1:
|
||||
filters: *main_branches
|
||||
- local_testing_part2:
|
||||
filters: *main_branches
|
||||
- langfuse_logging_unit_tests:
|
||||
filters: *main_branches
|
||||
- litellm_assistants_api_testing:
|
||||
filters: *main_branches
|
||||
- litellm_router_testing:
|
||||
filters: *main_branches
|
||||
- litellm_router_unit_testing:
|
||||
filters: *main_branches
|
||||
- auth_ui_unit_tests:
|
||||
filters: *main_branches
|
||||
- build_docker_database_image:
|
||||
filters: *main_branches
|
||||
- e2e_ui_testing:
|
||||
filters: *main_branches
|
||||
- e2e_ui_testing_server_root_path:
|
||||
filters: *main_branches
|
||||
- base_sdk_install
|
||||
- local_testing_part1
|
||||
- local_testing_part2
|
||||
- langfuse_logging_unit_tests
|
||||
- litellm_assistants_api_testing
|
||||
- litellm_router_testing
|
||||
- litellm_router_unit_testing
|
||||
- auth_ui_unit_tests
|
||||
- build_docker_database_image
|
||||
- e2e_ui_testing
|
||||
- e2e_ui_testing_server_root_path
|
||||
- build_and_test:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- e2e_openai_endpoints:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_logging_guardrails_model_info_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_spend_accuracy_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_multi_instance_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_store_model_in_db_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_build_from_pip_tests:
|
||||
filters: *main_branches
|
||||
- proxy_build_from_pip_tests
|
||||
- proxy_pass_through_endpoint_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_e2e_anthropic_messages_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- llm_translation_testing:
|
||||
filters: *main_branches
|
||||
- realtime_translation_testing:
|
||||
filters: *main_branches
|
||||
- agent_testing:
|
||||
filters: *main_branches
|
||||
- guardrails_testing:
|
||||
filters: *main_branches
|
||||
- google_generate_content_endpoint_testing:
|
||||
filters: *main_branches
|
||||
- llm_responses_api_testing:
|
||||
filters: *main_branches
|
||||
- ocr_testing:
|
||||
filters: *main_branches
|
||||
- search_testing:
|
||||
filters: *main_branches
|
||||
- batches_testing:
|
||||
filters: *main_branches
|
||||
- litellm_utils_testing:
|
||||
filters: *main_branches
|
||||
- pass_through_unit_testing:
|
||||
filters: *main_branches
|
||||
- image_gen_testing:
|
||||
filters: *main_branches
|
||||
- logging_testing:
|
||||
filters: *main_branches
|
||||
- audio_testing:
|
||||
filters: *main_branches
|
||||
- redis_caching_unit_tests:
|
||||
filters: *main_branches
|
||||
- llm_translation_testing
|
||||
- realtime_translation_testing
|
||||
- agent_testing
|
||||
- guardrails_testing
|
||||
- google_generate_content_endpoint_testing
|
||||
- llm_responses_api_testing
|
||||
- ocr_testing
|
||||
- search_testing
|
||||
- batches_testing
|
||||
- litellm_utils_testing
|
||||
- pass_through_unit_testing
|
||||
- image_gen_testing
|
||||
- logging_testing
|
||||
- audio_testing
|
||||
- redis_caching_unit_tests
|
||||
- upload-coverage:
|
||||
requires:
|
||||
- realtime_translation_testing
|
||||
|
|
@ -3469,18 +3420,12 @@ workflows:
|
|||
- db_migration_disable_update_check:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python_3_13:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python_v2_migration_resolver:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python
|
||||
- installing_litellm_on_python_3_13
|
||||
- installing_litellm_on_python_v2_migration_resolver
|
||||
- helm_chart_testing:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- test_bad_database_url:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
|
|
|
|||
|
|
@ -14,12 +14,12 @@ while IFS= read -r file || [ -n "$file" ]; do
|
|||
[ -n "$file" ] || continue
|
||||
case "$file" in
|
||||
*.md | *.mdx) : ;;
|
||||
pyproject.toml | */pyproject.toml | uv.lock | uv.toml | .python-version | rust-toolchain.toml | litellm-rust/* | litellm/__init__.py | litellm/proxy/proxy_server.py | litellm/*mcp* | tests/*mcp* | litellm/integrations/arize/* | tests/base_sdk_tests/* | scripts/check_mcp_sdk_install.py | .github/workflows/test-mcp-dependency-resolution.yml | .github/actions/detect-changes/* | .github/actions/setup-uv-with-retries/* | .github/actions/cache-cargo-build/* | .github/scripts/detect_changes.sh | .github/scripts/uv_sync_with_retries.sh | .circleci/scripts/classify_changes.sh | tests/test_litellm/test_circleci_path_filter.py | tests/test_litellm/test_detect_changes.py)
|
||||
pyproject.toml | */pyproject.toml | uv.lock | uv.toml | .python-version | rust-toolchain.toml | litellm-rust/* | litellm/__init__.py | litellm/proxy/proxy_server.py | litellm/*mcp* | tests/*mcp* | litellm/integrations/arize/* | tests/base_sdk_tests/* | scripts/check_mcp_sdk_install.py | .github/workflows/test-mcp-dependency-resolution.yml | .github/actions/detect-changes/* | .github/actions/setup-uv-with-retries/* | .github/actions/cache-cargo-build/* | .github/scripts/detect_changes.sh | .github/scripts/uv_sync_with_retries.sh | .circleci/scripts/classify_changes.sh | tests/unit/test_circleci_path_filter.py | tests/unit/test_detect_changes.py)
|
||||
has_mcp_dependencies=true ;;
|
||||
esac
|
||||
case "$file" in
|
||||
tests/e2e/*/*.py) : ;;
|
||||
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/test_litellm/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
|
||||
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/unit/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
|
||||
has_provider_harness=true ;;
|
||||
esac
|
||||
case "$file" in
|
||||
|
|
|
|||
|
|
@ -5,9 +5,14 @@ flag="${1:?usage: unit_selection.sh <codecov flag>}"
|
|||
|
||||
legacy_flags=(
|
||||
caching-local
|
||||
core-utils
|
||||
enterprise-package
|
||||
enterprise-routing
|
||||
integrations
|
||||
llm-other-providers
|
||||
llm-vertex-ai
|
||||
mcp-integration
|
||||
misc
|
||||
proxy-db-auth-checks
|
||||
proxy-db-budgets
|
||||
proxy-db-custom-logging
|
||||
|
|
@ -22,11 +27,13 @@ legacy_flags=(
|
|||
proxy-db-proxy-utils
|
||||
proxy-extras
|
||||
proxy-infra
|
||||
responses-caching-types
|
||||
)
|
||||
|
||||
legacy_paths() {
|
||||
case "$1" in
|
||||
caching-local) echo tests/unit/caching ;;
|
||||
core-utils) echo tests/unit/litellm_core_utils ;;
|
||||
enterprise-package)
|
||||
echo tests/unit/enterprise/integrations
|
||||
echo tests/unit/enterprise/proxy/auth
|
||||
|
|
@ -36,6 +43,9 @@ legacy_paths() {
|
|||
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/google_genai
|
||||
echo tests/unit/router_strategy
|
||||
echo tests/unit/router_utils
|
||||
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
|
||||
|
|
@ -47,10 +57,34 @@ legacy_paths() {
|
|||
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 ;;
|
||||
integrations) echo tests/unit/integrations ;;
|
||||
llm-other-providers) find tests/unit/llms -name 'test_*.py' -not -path 'tests/unit/llms/vertex_ai/*' ;;
|
||||
llm-vertex-ai) echo tests/unit/llms/vertex_ai ;;
|
||||
mcp-integration)
|
||||
echo tests/unit/experimental_mcp_client
|
||||
echo tests/unit/proxy/_experimental/mcp_server
|
||||
echo tests/unit/responses/mcp
|
||||
echo tests/mcp_tests/test_proxy_mcp_e2e.py ;;
|
||||
misc)
|
||||
find tests/unit -maxdepth 1 -name 'test_*.py'
|
||||
echo tests/unit/test_router
|
||||
echo tests/unit/a2a_protocol
|
||||
echo tests/unit/batches
|
||||
echo tests/unit/chat_completions
|
||||
echo tests/unit/completion_extras
|
||||
echo tests/unit/containers
|
||||
echo tests/unit/embeddings
|
||||
echo tests/unit/endpoints
|
||||
echo tests/unit/files
|
||||
echo tests/unit/images
|
||||
echo tests/unit/interactions
|
||||
echo tests/unit/messages
|
||||
echo tests/unit/rag
|
||||
echo tests/unit/rerank_api
|
||||
echo tests/unit/rust_bridge
|
||||
echo tests/unit/secret_managers
|
||||
echo tests/unit/vector_stores
|
||||
echo tests/unit/videos ;;
|
||||
proxy-db-auth-checks)
|
||||
echo tests/unit/proxy/auth/test_auth_checks.py
|
||||
echo tests/unit/proxy/auth/test_user_api_key_auth.py
|
||||
|
|
@ -113,6 +147,9 @@ legacy_paths() {
|
|||
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 ;;
|
||||
responses-caching-types)
|
||||
find tests/unit/responses -name 'test_*.py' -not -path 'tests/unit/responses/mcp/*'
|
||||
echo tests/unit/types ;;
|
||||
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
|
||||
esac
|
||||
}
|
||||
|
|
|
|||
|
|
@ -341,6 +341,7 @@ workflows:
|
|||
flag:
|
||||
- enterprise-package
|
||||
- proxy-infra
|
||||
- responses-caching-types
|
||||
- proxy-db-auth-checks
|
||||
- proxy-db-jwt-and-keys
|
||||
- proxy-db-proxy-server-core
|
||||
|
|
@ -353,6 +354,42 @@ workflows:
|
|||
- 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-llm-vertex-ai
|
||||
flag: llm-vertex-ai
|
||||
shards: 2
|
||||
workers: 1
|
||||
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 "" >>
|
||||
- unit:
|
||||
name: unit-llm-other-providers
|
||||
flag: llm-other-providers
|
||||
shards: 3
|
||||
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 "" >>
|
||||
- unit:
|
||||
name: unit-core-utils
|
||||
flag: core-utils
|
||||
shards: 2
|
||||
reruns: 1
|
||||
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-integrations
|
||||
flag: integrations
|
||||
shards: 2
|
||||
reruns: 3
|
||||
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-misc
|
||||
flag: misc
|
||||
shards: 2
|
||||
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 "" >>
|
||||
- unit:
|
||||
name: unit-proxy-db-proxy-utils
|
||||
flag: proxy-db-proxy-utils
|
||||
|
|
|
|||
18
.github/merge-smoke-tests.json
vendored
18
.github/merge-smoke-tests.json
vendored
|
|
@ -1,15 +1,15 @@
|
|||
{
|
||||
"cases": {
|
||||
"CHAT-JSON": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
|
||||
"CHAT-TEXT-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
|
||||
"CHAT-TOOL-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
|
||||
"CHAT-JSON": "tests/unit/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
|
||||
"CHAT-TEXT-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
|
||||
"CHAT-TOOL-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
|
||||
"MODEL-ALLOW": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_allows_listed_model_for_key",
|
||||
"MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]",
|
||||
"COST-EXPLICIT": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
|
||||
"COST-ZERO": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
|
||||
"LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
|
||||
"LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
|
||||
"CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
|
||||
"CALLBACK-FAILURE": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
|
||||
"COST-EXPLICIT": "tests/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
|
||||
"COST-ZERO": "tests/unit/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
|
||||
"LOG-CONTENT-ON": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
|
||||
"LOG-CONTENT-OFF": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
|
||||
"CALLBACK-SUCCESS": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
|
||||
"CALLBACK-FAILURE": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
|
||||
}
|
||||
}
|
||||
|
|
|
|||
2
.github/pull_request_template.md
vendored
2
.github/pull_request_template.md
vendored
|
|
@ -65,7 +65,7 @@ After: the same request comes back with real token counts, so the dashboard show
|
|||
**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
|
||||
- [ ] The handful of test files covering my change pass locally, e.g. `uv run pytest tests/unit/<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)
|
||||
|
|
|
|||
18
.github/workflows/_test-unit-base.yml
vendored
18
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -13,12 +13,13 @@ on:
|
|||
have its path existence-checked like any other token.
|
||||
required: true
|
||||
type: string
|
||||
fork-flag:
|
||||
unit-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.
|
||||
this shard. The shard also runs the files
|
||||
`.circleci/scripts/unit_selection.sh` lists for the flag, on every
|
||||
event, because the CircleCI pipeline is manual-only while the tests
|
||||
migrate.
|
||||
required: false
|
||||
type: string
|
||||
default: ""
|
||||
|
|
@ -175,8 +176,7 @@ jobs:
|
|||
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 }}
|
||||
UNIT_FLAG: ${{ inputs.unit-flag }}
|
||||
MAX_FAILURES: ${{ inputs.max-failures }}
|
||||
WORKERS: ${{ inputs.workers }}
|
||||
RERUNS: ${{ inputs.reruns }}
|
||||
|
|
@ -186,11 +186,11 @@ jobs:
|
|||
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' ' ')"
|
||||
if [ -n "${UNIT_FLAG}" ]; then
|
||||
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${UNIT_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"
|
||||
echo "shard selection is empty; nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
pytest_args=()
|
||||
|
|
|
|||
14
.github/workflows/test-redis-compat.yml
vendored
14
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -8,9 +8,13 @@ on:
|
|||
paths:
|
||||
- "litellm/_redis.py"
|
||||
- "litellm/_redis_credential_provider.py"
|
||||
- "tests/test_litellm/test_redis.py"
|
||||
- "litellm/caching/redis_cache.py"
|
||||
- "litellm/caching/evicted_client_closer.py"
|
||||
- "tests/unit/test_redis.py"
|
||||
- "tests/local_testing/test_caching.py"
|
||||
- "tests/test_litellm/caching/test_redis_connection_pool.py"
|
||||
- "tests/unit/caching/test_redis_connection_pool.py"
|
||||
- "tests/unit/caching/test_redis_cluster_cache.py"
|
||||
- "tests/unit/caching/test_evicted_client_closer.py"
|
||||
- ".github/workflows/test-redis-compat.yml"
|
||||
- "pyproject.toml"
|
||||
- "uv.lock"
|
||||
|
|
@ -80,8 +84,10 @@ jobs:
|
|||
run: |
|
||||
redis-server --version
|
||||
uv run --no-sync pytest \
|
||||
tests/test_litellm/test_redis.py \
|
||||
tests/test_litellm/caching/test_redis_connection_pool.py \
|
||||
tests/unit/test_redis.py \
|
||||
tests/unit/caching/test_redis_connection_pool.py \
|
||||
tests/unit/caching/test_redis_cluster_cache.py \
|
||||
tests/unit/caching/test_evicted_client_closer.py \
|
||||
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
|
||||
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
|
||||
--tb=short -vv \
|
||||
|
|
|
|||
8
.github/workflows/test-rust.yml
vendored
8
.github/workflows/test-rust.yml
vendored
|
|
@ -14,7 +14,6 @@ on:
|
|||
- "litellm/ocr/**"
|
||||
- "litellm/llms/base_llm/ocr/**"
|
||||
- "litellm/llms/custom_httpx/llm_http_handler.py"
|
||||
- "tests/test_litellm/ocr/**"
|
||||
- "tests/test_litellm/conftest.py"
|
||||
- "Makefile"
|
||||
- ".cargo/**"
|
||||
|
|
@ -24,7 +23,7 @@ on:
|
|||
- ".github/actions/setup-uv-with-retries/**"
|
||||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- "tests/unit/rust_bridge/native_route_wheel_test.py"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
pull_request:
|
||||
branches:
|
||||
|
|
@ -42,7 +41,6 @@ on:
|
|||
- "litellm/ocr/**"
|
||||
- "litellm/llms/base_llm/ocr/**"
|
||||
- "litellm/llms/custom_httpx/llm_http_handler.py"
|
||||
- "tests/test_litellm/ocr/**"
|
||||
- "tests/test_litellm/conftest.py"
|
||||
- "Makefile"
|
||||
- ".cargo/**"
|
||||
|
|
@ -52,7 +50,7 @@ on:
|
|||
- ".github/actions/setup-uv-with-retries/**"
|
||||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- "tests/unit/rust_bridge/native_route_wheel_test.py"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
|
||||
permissions:
|
||||
|
|
@ -171,7 +169,7 @@ jobs:
|
|||
env:
|
||||
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
|
||||
|
||||
- run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
- run: python tests/unit/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
|
||||
- name: Run pytest tests/test_litellm_rust with the compiled extension
|
||||
run: make test-rust-extension
|
||||
|
|
|
|||
33
.github/workflows/test-unit-proxy-db.yml
vendored
33
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -22,9 +22,10 @@ concurrency:
|
|||
#
|
||||
# `.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.
|
||||
# the file lists. That pipeline is manual-only while the tests migrate, so
|
||||
# `unit-flag` makes the shard run that list on every event. `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.
|
||||
|
|
@ -78,7 +79,7 @@ jobs:
|
|||
# Must run serially — event-loop conflict with the logging worker.
|
||||
- test-group: key-generation
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-key-generation
|
||||
unit-flag: proxy-db-key-generation
|
||||
workers: 0
|
||||
dist: loadscope
|
||||
timeout: 20
|
||||
|
|
@ -86,13 +87,13 @@ jobs:
|
|||
# ---- auth: split into 2 shards ----
|
||||
- test-group: auth-checks
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-auth-checks
|
||||
unit-flag: proxy-db-auth-checks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: jwt-and-keys
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-jwt-and-keys
|
||||
unit-flag: proxy-db-jwt-and-keys
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
@ -100,7 +101,7 @@ jobs:
|
|||
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
|
||||
- test-group: proxy-utils
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-utils
|
||||
unit-flag: proxy-db-proxy-utils
|
||||
workers: 4
|
||||
dist: worksteal
|
||||
timeout: 15
|
||||
|
|
@ -108,13 +109,13 @@ jobs:
|
|||
# ---- proxy server: split into 2 shards ----
|
||||
- test-group: proxy-server-core
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py"
|
||||
fork-flag: proxy-db-proxy-server-core
|
||||
unit-flag: proxy-db-proxy-server-core
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: proxy-runtime
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-runtime
|
||||
unit-flag: proxy-db-proxy-runtime
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
@ -122,20 +123,20 @@ jobs:
|
|||
# ---- logging: split into 2 shards ----
|
||||
- test-group: custom-logging
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py"
|
||||
fork-flag: proxy-db-custom-logging
|
||||
unit-flag: proxy-db-custom-logging
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: logging-misc
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-logging-misc
|
||||
unit-flag: proxy-db-logging-misc
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: db-and-spend
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-db-and-spend
|
||||
unit-flag: proxy-db-db-and-spend
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
@ -143,27 +144,27 @@ jobs:
|
|||
# ---- guardrails + budget + hooks: split into 2 ----
|
||||
- test-group: guardrails-hooks
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-guardrails-hooks
|
||||
unit-flag: proxy-db-guardrails-hooks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: budgets
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-budgets
|
||||
unit-flag: proxy-db-budgets
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: endpoints-and-responses
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py"
|
||||
fork-flag: proxy-db-endpoints-and-responses
|
||||
unit-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 }}
|
||||
unit-flag: ${{ matrix.unit-flag }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: 2
|
||||
timeout-minutes: ${{ matrix.timeout }}
|
||||
|
|
|
|||
65
.github/workflows/test-unit.yml
vendored
65
.github/workflows/test-unit.yml
vendored
|
|
@ -36,9 +36,9 @@ concurrency:
|
|||
# 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.
|
||||
# `unit-flag` names the `.circleci/tests.yml` job that now runs part of the
|
||||
# shard under the same Codecov flag. That pipeline is manual-only while the
|
||||
# tests migrate, so the shard also runs those files on every event.
|
||||
jobs:
|
||||
unit:
|
||||
name: ${{ matrix.shard }}
|
||||
|
|
@ -52,8 +52,8 @@ jobs:
|
|||
include:
|
||||
- shard: mcp-integration
|
||||
artifact-name: mcp-integration
|
||||
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
|
||||
fork-flag: mcp-integration
|
||||
test-path: "tests/mcp_tests"
|
||||
unit-flag: mcp-integration
|
||||
workers: 2
|
||||
reruns: 0
|
||||
timeout-minutes: 20
|
||||
|
|
@ -61,7 +61,8 @@ jobs:
|
|||
|
||||
- shard: core-utils
|
||||
artifact-name: core-utils
|
||||
test-path: "tests/test_litellm/litellm_core_utils"
|
||||
test-path: ""
|
||||
unit-flag: core-utils
|
||||
workers: 2
|
||||
reruns: 1
|
||||
timeout-minutes: 20
|
||||
|
|
@ -69,11 +70,8 @@ jobs:
|
|||
|
||||
- shard: enterprise-routing
|
||||
artifact-name: enterprise-routing
|
||||
test-path: >-
|
||||
tests/test_litellm/google_genai
|
||||
tests/test_litellm/router_utils
|
||||
tests/test_litellm/router_strategy
|
||||
fork-flag: enterprise-routing
|
||||
test-path: ""
|
||||
unit-flag: enterprise-routing
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -81,7 +79,8 @@ jobs:
|
|||
|
||||
- shard: integrations
|
||||
artifact-name: integrations
|
||||
test-path: "tests/test_litellm/integrations"
|
||||
test-path: ""
|
||||
unit-flag: integrations
|
||||
workers: 2
|
||||
reruns: 3
|
||||
timeout-minutes: 20
|
||||
|
|
@ -89,7 +88,8 @@ jobs:
|
|||
|
||||
- shard: Vertex AI
|
||||
artifact-name: llm-vertex-ai
|
||||
test-path: "tests/test_litellm/llms/vertex_ai"
|
||||
test-path: ""
|
||||
unit-flag: llm-vertex-ai
|
||||
workers: 1
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -97,7 +97,8 @@ jobs:
|
|||
|
||||
- shard: All Other Providers
|
||||
artifact-name: llm-other-providers
|
||||
test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai"
|
||||
test-path: ""
|
||||
unit-flag: llm-other-providers
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -106,26 +107,8 @@ jobs:
|
|||
- shard: misc
|
||||
artifact-name: misc
|
||||
test-path: >-
|
||||
tests/test_litellm/batches
|
||||
tests/test_litellm/secret_managers
|
||||
tests/test_litellm/a2a_protocol
|
||||
tests/test_litellm/chat_completions
|
||||
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
|
||||
tests/test_litellm/embeddings
|
||||
tests/test_litellm/ocr
|
||||
tests/test_litellm/passthrough
|
||||
tests/test_litellm/rag
|
||||
tests/test_litellm/rerank_api
|
||||
tests/test_litellm/rust_bridge
|
||||
tests/test_litellm/vector_stores
|
||||
tests/test_litellm/videos
|
||||
tests/test_litellm/test_*.py
|
||||
unit-flag: misc
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -205,7 +188,7 @@ jobs:
|
|||
tests/test_litellm/proxy/types_utils
|
||||
tests/test_litellm/proxy/logging_endpoints
|
||||
tests/test_litellm/proxy/test_*.py
|
||||
fork-flag: proxy-infra
|
||||
unit-flag: proxy-infra
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -214,7 +197,7 @@ jobs:
|
|||
- shard: caching-local
|
||||
artifact-name: caching-local
|
||||
test-path: ""
|
||||
fork-flag: caching-local
|
||||
unit-flag: caching-local
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -223,7 +206,7 @@ jobs:
|
|||
- shard: proxy-extras
|
||||
artifact-name: proxy-extras
|
||||
test-path: ""
|
||||
fork-flag: proxy-extras
|
||||
unit-flag: proxy-extras
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -232,7 +215,7 @@ jobs:
|
|||
- shard: enterprise-package
|
||||
artifact-name: enterprise-package
|
||||
test-path: ""
|
||||
fork-flag: enterprise-package
|
||||
unit-flag: enterprise-package
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -240,10 +223,8 @@ jobs:
|
|||
|
||||
- shard: responses-caching-types
|
||||
artifact-name: responses-caching-types
|
||||
test-path: >-
|
||||
tests/test_litellm/responses
|
||||
tests/test_litellm/caching
|
||||
tests/test_litellm/types
|
||||
test-path: ""
|
||||
unit-flag: responses-caching-types
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -251,7 +232,7 @@ jobs:
|
|||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: ${{ matrix.test-path }}
|
||||
fork-flag: ${{ matrix.fork-flag || '' }}
|
||||
unit-flag: ${{ matrix.unit-flag || '' }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: ${{ matrix.reruns }}
|
||||
timeout-minutes: ${{ matrix.timeout-minutes }}
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ Never test structure of code only function of it
|
|||
|
||||
A test must only fail when litellm code changes. Never pin facts we don't own (a vendor's price, a third party's field, an upstream default, today's date) as literals or as "X must be absent"; assert the invariant our code guarantees instead, e.g. two rows agree, a value is within range, a field is derived from another. If an outside fact is truly load-bearing, cite its source and date next to the assertion so a reader can tell stale from broken
|
||||
|
||||
`tests/test_litellm/` mirrors `litellm/` in a parallel path (see `tests/test_litellm/readme.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
|
||||
`tests/unit/` mirrors `litellm/` in a parallel path (see `tests/unit/AGENTS.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
|
||||
|
||||
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `AGENTS.md`
|
||||
|
||||
|
|
|
|||
|
|
@ -255,7 +255,7 @@ Conventions to follow when touching this layer:
|
|||
| Column vs. field names | Where a model field differs from its DB column (for example `org_id` maps to the `organization_id` column), the repository translates in both directions rather than relying on Pydantic to guess. |
|
||||
| Array mutations | Adds use Prisma's atomic `push` (`add_member`, `add_admin`, `add_models`) to avoid read-modify-write races. Removals fall back to read-modify-write because Prisma has no atomic array remove. |
|
||||
|
||||
To add a new entity, define the model under `litellm/models/`, re-export it from `proxy/_types.py` if existing code imports it from there, and add a repository under `litellm/repositories/` (subclass `BaseRepository` for plain CRUD, or add bespoke methods when the entity needs encryption, archiving, or atomic array updates). Mirror the tests in `tests/test_litellm/repositories/`.
|
||||
To add a new entity, define the model under `litellm/models/`, re-export it from `proxy/_types.py` if existing code imports it from there, and add a repository under `litellm/repositories/` (subclass `BaseRepository` for plain CRUD, or add bespoke methods when the entity needs encryption, archiving, or atomic array updates). Mirror the tests in `tests/unit/repositories/`.
|
||||
|
||||
---
|
||||
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ Here are the core requirements for any PR submitted to LiteLLM:
|
|||
- [ ] **Add testing** - Adding at least 1 test is a hard requirement - [see details](#adding-testing)
|
||||
- [ ] **Ensure your PR passes all checks**:
|
||||
- [ ] [Linting / Formatting](#running-linting-and-formatting-checks) - `make lint`
|
||||
- [ ] [The tests covering your change](#running-unit-tests) pass, e.g. `uv run pytest tests/test_litellm/<your_test_file>.py -v`. CI runs the full unit test matrix, so you don't need to run the whole suite locally
|
||||
- [ ] [The tests covering your change](#running-unit-tests) pass, e.g. `uv run pytest tests/unit/<your_test_file>.py -v`. CI runs the full unit test matrix, so you don't need to run the whole suite locally
|
||||
|
||||
#### UI PRs
|
||||
|
||||
|
|
@ -72,7 +72,7 @@ make format
|
|||
make lint
|
||||
|
||||
# Run the tests covering your change (CI runs the full suite)
|
||||
uv run pytest tests/test_litellm/<your_test_file>.py -v
|
||||
uv run pytest tests/unit/<your_test_file>.py -v
|
||||
|
||||
# Commit your changes (must follow Conventional Commits — see above)
|
||||
git add .
|
||||
|
|
@ -88,7 +88,7 @@ git push origin feature/your-feature
|
|||
|
||||
### Where to Add Tests
|
||||
|
||||
Add your tests to the [`tests/test_litellm/` directory](https://github.com/BerriAI/litellm/tree/main/tests/test_litellm).
|
||||
Add your tests to the [`tests/unit/` directory](https://github.com/BerriAI/litellm/tree/main/tests/unit).
|
||||
|
||||
- This directory mirrors the structure of the `litellm/` directory
|
||||
- **Only add mocked tests** - no real LLM API calls in this directory
|
||||
|
|
@ -96,10 +96,10 @@ Add your tests to the [`tests/test_litellm/` directory](https://github.com/Berri
|
|||
|
||||
### File Naming Convention
|
||||
|
||||
The `tests/test_litellm/` directory follows the same structure as `litellm/`:
|
||||
The `tests/unit/` directory follows the same structure as `litellm/`:
|
||||
|
||||
- `litellm/proxy/caching_routes.py` → `tests/test_litellm/proxy/test_caching_routes.py`
|
||||
- `litellm/utils.py` → `tests/test_litellm/test_utils.py`
|
||||
- `litellm/utils.py` → `tests/unit/test_utils.py`
|
||||
|
||||
### Example Test
|
||||
|
||||
|
|
@ -125,10 +125,10 @@ def test_your_feature():
|
|||
|
||||
Run the tests covering your change:
|
||||
```bash
|
||||
uv run pytest tests/test_litellm/test_your_file.py -v
|
||||
uv run pytest tests/unit/test_your_file.py -v
|
||||
```
|
||||
|
||||
`tests/test_litellm` holds thousands of tests, so running all of it locally takes a long time. CI runs it as a parallel matrix (`make test-unit-llms`, `make test-unit-proxy-core`, and the other `test-unit-*` targets) on beefier boxes, so if, for whatever reason, you must run the whole suite, it's better to rely on CI to do that.
|
||||
`tests/unit` holds thousands of tests, so running all of it locally takes a long time. CI runs it as a parallel matrix (`make test-unit-llms`, `make test-unit-proxy-core`, and the other `test-unit-*` targets) on beefier boxes, so if, for whatever reason, you must run the whole suite, it's better to rely on CI to do that.
|
||||
|
||||
If you're running broader test suites, proxy tests, or anything that touches PostgreSQL-backed fixtures/plugins, install the full local test environment first:
|
||||
|
||||
|
|
|
|||
16
Makefile
16
Makefile
|
|
@ -42,7 +42,7 @@ help:
|
|||
@echo " make check-circular-imports - Check for circular imports"
|
||||
@echo " make check-import-safety - Check import safety"
|
||||
@echo " make test - Run all tests"
|
||||
@echo " make test-unit - Run unit tests (tests/test_litellm)"
|
||||
@echo " make test-unit - Run unit tests (tests/unit and tests/test_litellm)"
|
||||
@echo " make test-unit-llms - Run LLM provider tests (~225 files)"
|
||||
@echo " make test-unit-proxy-guardrails - Run proxy guardrails+mgmt tests (~51 files)"
|
||||
@echo " make test-unit-proxy-core - Run proxy auth+client+db+hooks tests (~52 files)"
|
||||
|
|
@ -301,7 +301,7 @@ test-rust-extension:
|
|||
UV_PROJECT_ENVIRONMENT="$$temporary/venv" $(UV) sync --python 3.12 --frozen --no-install-project --all-groups --all-extras && \
|
||||
$(UV) pip install --python "$$temporary/venv/bin/python" --no-deps "$$1" && \
|
||||
"$$temporary/venv/bin/python" -I -m mypy.stubtest \
|
||||
--mypy-config-file tests/test_litellm/rust_bridge/stubtest.ini \
|
||||
--mypy-config-file tests/unit/rust_bridge/stubtest.ini \
|
||||
litellm.rust_bridge._native && \
|
||||
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
|
||||
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust
|
||||
|
|
@ -310,11 +310,11 @@ test: install-test-deps
|
|||
$(UV_RUN) pytest tests/
|
||||
|
||||
test-unit: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm -x -vv -n 4
|
||||
$(UV_RUN) pytest tests/unit tests/test_litellm -x -vv -n 4
|
||||
|
||||
# Matrix test targets (matching CI workflow groups)
|
||||
test-unit-llms: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/llms --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/llms --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-proxy-guardrails: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers --tb=short -vv -n 4 --durations=20
|
||||
|
|
@ -326,16 +326,16 @@ test-unit-proxy-misc: install-test-deps
|
|||
$(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-integrations: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-core-utils: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/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/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
|
||||
$(UV_RUN) pytest tests/unit/caching tests/unit/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/unit/router_strategy tests/unit/router_utils tests/unit/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
|
||||
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
# Proxy unit tests (tests/unit/proxy split alphabetically)
|
||||
test-proxy-unit-a: install-test-deps
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
|
||||
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/unit/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
|
||||
|
||||
The first eleven rows need only `callbacks: ["prometheus"]`. The last three rows and the `litellm_admission_*` panels are emitted by other subsystems and stay empty until those are on: the service callback row needs `service_callback: ["prometheus_system"]` in `litellm_settings`, the circuit breaker row needs a Redis cache, the cleanup row needs spend log retention, and admission control needs its middleware enabled. Within the base rows, many panels only fill in once the matching feature is in use: budgets need keys, teams, users or orgs with `max_budget` set, cache panels need caching on, guardrail and MCP panels need those features configured, deployment health needs the router with more than one deployment or a failure to record, and `litellm_in_flight_requests` needs traffic at scrape time. An empty panel for a feature you do not use is expected
|
||||
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from litellm.integrations.email_templates.key_rotated_email import (
|
|||
from litellm.integrations.email_templates.templates import (
|
||||
MAX_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
SOFT_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
TEAM_SOFT_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
)
|
||||
from litellm.integrations.email_templates.user_invitation_email import (
|
||||
|
|
@ -48,6 +49,12 @@ from litellm.secret_managers.main import get_secret_bool
|
|||
from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL
|
||||
|
||||
|
||||
def _max_budget_alert_id(user_info: CallInfo) -> str:
|
||||
if user_info.event_group == Litellm_EntityType.TEAM_MEMBER:
|
||||
return f"team_member:{user_info.user_id}:{user_info.team_id}"
|
||||
return user_info.token or user_info.user_id or "default_id"
|
||||
|
||||
|
||||
def _parse_email_list(raw) -> List[str]:
|
||||
"""Parse emails from a list or comma-separated string."""
|
||||
if isinstance(raw, list):
|
||||
|
|
@ -373,17 +380,31 @@ class BaseEmailLogger(CustomLogger):
|
|||
greeting = html.escape(
|
||||
event.user_email or event.key_alias or event.token or ""
|
||||
)
|
||||
email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format(
|
||||
email_logo_url=email_params.logo_url,
|
||||
recipient_email=greeting,
|
||||
percentage=percentage,
|
||||
spend=spend_str,
|
||||
max_budget=max_budget_str,
|
||||
alert_threshold=alert_threshold_str,
|
||||
base_url=email_params.base_url,
|
||||
email_support_contact=email_params.support_contact,
|
||||
email_footer=email_params.signature,
|
||||
)
|
||||
if event.event_group == Litellm_EntityType.TEAM_MEMBER:
|
||||
email_html_content = TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format(
|
||||
email_logo_url=email_params.logo_url,
|
||||
member=html.escape(event.user_email or event.user_id or ""),
|
||||
team_alias=html.escape(event.team_alias or event.team_id or ""),
|
||||
percentage=percentage,
|
||||
spend=spend_str,
|
||||
max_budget=max_budget_str,
|
||||
alert_threshold=alert_threshold_str,
|
||||
base_url=email_params.base_url,
|
||||
email_support_contact=email_params.support_contact,
|
||||
email_footer=email_params.signature,
|
||||
)
|
||||
else:
|
||||
email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format(
|
||||
email_logo_url=email_params.logo_url,
|
||||
recipient_email=greeting,
|
||||
percentage=percentage,
|
||||
spend=spend_str,
|
||||
max_budget=max_budget_str,
|
||||
alert_threshold=alert_threshold_str,
|
||||
base_url=email_params.base_url,
|
||||
email_support_contact=email_params.support_contact,
|
||||
email_footer=email_params.signature,
|
||||
)
|
||||
await self.send_email(
|
||||
from_email=self.DEFAULT_LITELLM_EMAIL,
|
||||
to_email=recipient_emails,
|
||||
|
|
@ -607,7 +628,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
if user_info.spend < threshold_amount:
|
||||
continue
|
||||
|
||||
_id = user_info.token or user_info.user_id or "default_id"
|
||||
_id = _max_budget_alert_id(user_info)
|
||||
_cache_key = (
|
||||
f"email_budget_alerts:max_budget_alert:{threshold_pct}:{_id}"
|
||||
)
|
||||
|
|
@ -618,7 +639,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
emails.append(user_info.user_email)
|
||||
if not emails:
|
||||
verbose_proxy_logger.warning(
|
||||
"No recipients for %d%% threshold on key %s, skipping alert",
|
||||
"No recipients for %d%% threshold on %s, skipping alert",
|
||||
threshold_pct,
|
||||
_id,
|
||||
)
|
||||
|
|
@ -633,7 +654,11 @@ class BaseEmailLogger(CustomLogger):
|
|||
if send_count is not None and send_count > 1:
|
||||
continue
|
||||
|
||||
event_message = f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
|
||||
event_message = (
|
||||
f"Team Member Budget Alert - {threshold_pct}% of Team Member Budget Reached"
|
||||
if user_info.event_group == Litellm_EntityType.TEAM_MEMBER
|
||||
else f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
|
||||
)
|
||||
webhook_event = WebhookEvent(
|
||||
event="max_budget_alert",
|
||||
event_message=event_message,
|
||||
|
|
|
|||
|
|
@ -9,10 +9,14 @@
|
|||
- 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
|
||||
|
||||
## Test fixtures and cases
|
||||
|
||||
Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new and updated tests and [`#[fixture]`](https://docs.rs/rstest/latest/rstest/attr.fixture.html) for reusable setup, injected through typed test arguments. Express input variations as named `#[case::name(...)]` cases instead of loops or duplicated tests so each failure identifies its case. Keep behavior assertions in the test body and fixtures focused on setup. Use the workspace `rstest` dependency
|
||||
|
||||
## Error definitions
|
||||
|
||||
- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs`
|
||||
- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each
|
||||
- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each. A failure mode is something a caller handles differently (phase, status code, retry, a message Python parity pins exactly); failures no caller tells apart share one variant and differ only in its message
|
||||
- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string
|
||||
- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return
|
||||
- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it
|
||||
|
|
|
|||
542
litellm-rust/Cargo.lock
generated
542
litellm-rust/Cargo.lock
generated
|
|
@ -73,6 +73,15 @@ version = "1.0.104"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470"
|
||||
|
||||
[[package]]
|
||||
name = "arbitrary"
|
||||
version = "1.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1"
|
||||
dependencies = [
|
||||
"derive_arbitrary",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "arc-swap"
|
||||
version = "1.9.2"
|
||||
|
|
@ -190,7 +199,7 @@ checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec"
|
|||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 3.0.0",
|
||||
"syn 3.0.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -701,14 +710,20 @@ dependencies = [
|
|||
"http 1.4.2",
|
||||
"http-body 1.1.0",
|
||||
"http-body-util",
|
||||
"hyper 1.10.1",
|
||||
"hyper-util",
|
||||
"itoa",
|
||||
"matchit",
|
||||
"memchr",
|
||||
"mime",
|
||||
"multer",
|
||||
"percent-encoding",
|
||||
"pin-project-lite",
|
||||
"serde_core",
|
||||
"serde_json",
|
||||
"serde_path_to_error",
|
||||
"sync_wrapper",
|
||||
"tokio",
|
||||
"tower",
|
||||
"tower-layer",
|
||||
"tower-service",
|
||||
|
|
@ -897,6 +912,12 @@ dependencies = [
|
|||
"hybrid-array",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "borrow-or-share"
|
||||
version = "0.2.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "dc0b364ead1874514c8c2855ab558056ebfeb775653e7ae45ff72f28f8f3166c"
|
||||
|
||||
[[package]]
|
||||
name = "bstr"
|
||||
version = "1.13.1"
|
||||
|
|
@ -914,6 +935,12 @@ version = "3.20.3"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649"
|
||||
|
||||
[[package]]
|
||||
name = "bytecount"
|
||||
version = "0.6.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "175812e0be2bccb6abe50bb8d566126198344f707e304f45c648fd8f2cc0365e"
|
||||
|
||||
[[package]]
|
||||
name = "byteorder"
|
||||
version = "1.5.0"
|
||||
|
|
@ -1032,18 +1059,18 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "clap"
|
||||
version = "4.6.6"
|
||||
version = "4.6.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "473c7e07f409a8d772161724aa8db6a765a2532a70f9667eeb7b49d3d02fbdca"
|
||||
checksum = "aa8876b300ab35ba921adea3dfd70157a46249b33f95c9084ae5709785478946"
|
||||
dependencies = [
|
||||
"clap_builder",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "clap_builder"
|
||||
version = "4.6.6"
|
||||
version = "4.6.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7b48fea5a88e9ae728a2dcbedbfc0e730f7d60da42e1cb049a83c9fb8b789889"
|
||||
checksum = "ec0797fb7aeb1406c84efac526901f7ec3ead2124f946b494e72879d4b54704d"
|
||||
dependencies = [
|
||||
"anstyle",
|
||||
"clap_lex",
|
||||
|
|
@ -1159,7 +1186,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "e75b2483e97a5a7da73ac68a05b629f9c53cff58d8ed1c77866079e18b00dba5"
|
||||
dependencies = [
|
||||
"digest 0.10.7",
|
||||
"spin",
|
||||
"spin 0.10.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -1458,6 +1485,17 @@ dependencies = [
|
|||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "derive_arbitrary"
|
||||
version = "1.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1e567bd82dcff979e4b03460c307b3cdc9e96fde3d73bed1496d2bc75d9dd62a"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "derive_builder"
|
||||
version = "0.20.2"
|
||||
|
|
@ -1540,6 +1578,24 @@ version = "1.16.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e"
|
||||
|
||||
[[package]]
|
||||
name = "email_address"
|
||||
version = "0.2.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e079f19b08ca6239f47f8ba8509c11cf3ea30095831f7fed61441475edd8c449"
|
||||
dependencies = [
|
||||
"serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "encoding_rs"
|
||||
version = "0.8.35"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "equivalent"
|
||||
version = "1.0.2"
|
||||
|
|
@ -1622,6 +1678,16 @@ version = "2.5.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223"
|
||||
|
||||
[[package]]
|
||||
name = "filetime"
|
||||
version = "0.2.29"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5c287a33c7f0a620c38e641e7f60827713987b3c0f26e8ddc9462cc69cf75759"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "find-msvc-tools"
|
||||
version = "0.1.9"
|
||||
|
|
@ -1639,6 +1705,17 @@ dependencies = [
|
|||
"zlib-rs",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "fluent-uri"
|
||||
version = "0.4.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "bc74ac4d8359ae70623506d512209619e5cf8f347124910440dbc221714b328e"
|
||||
dependencies = [
|
||||
"borrow-or-share",
|
||||
"ref-cast",
|
||||
"serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "fnv"
|
||||
version = "1.0.7"
|
||||
|
|
@ -1660,6 +1737,16 @@ dependencies = [
|
|||
"percent-encoding",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "fraction"
|
||||
version = "0.17.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e246562084dde8ebbcc943b261c406ce4f68e5032ec28029a251a47d6a295500"
|
||||
dependencies = [
|
||||
"num",
|
||||
"num-bigint 0.4.8",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "fs_extra"
|
||||
version = "1.3.0"
|
||||
|
|
@ -1817,9 +1904,11 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"js-sys",
|
||||
"libc",
|
||||
"r-efi 5.3.0",
|
||||
"wasip2",
|
||||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -2660,6 +2749,59 @@ dependencies = [
|
|||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jsonschema"
|
||||
version = "0.55.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b68339c3d874e48151d74ffe256d93a58cffa240983cb0967d3cbaea083a44fe"
|
||||
dependencies = [
|
||||
"ahash",
|
||||
"bytecount",
|
||||
"data-encoding",
|
||||
"email_address",
|
||||
"fancy-regex 0.19.2",
|
||||
"fraction",
|
||||
"getrandom 0.3.4",
|
||||
"itoa",
|
||||
"jsonschema-regex",
|
||||
"jsonschema-value",
|
||||
"num-cmp",
|
||||
"num-traits",
|
||||
"percent-encoding",
|
||||
"referencing",
|
||||
"regex",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"strum",
|
||||
"unicode-general-category",
|
||||
"uuid-simd",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jsonschema-regex"
|
||||
version = "0.55.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6307b5b51216ec9b941b52244c74043fa0b1d6b657b56199f57cb1416d3641c5"
|
||||
dependencies = [
|
||||
"regex-syntax",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jsonschema-value"
|
||||
version = "0.55.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0230ac05e09c6111e96c147b75c390579f5cbd45654b980c68ac60fe17b3f129"
|
||||
dependencies = [
|
||||
"ahash",
|
||||
"bytecount",
|
||||
"fraction",
|
||||
"getrandom 0.3.4",
|
||||
"num-cmp",
|
||||
"num-traits",
|
||||
"serde_json",
|
||||
"zmij",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "lazy_static"
|
||||
version = "1.5.0"
|
||||
|
|
@ -2689,6 +2831,10 @@ version = "0.12.1"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53"
|
||||
|
||||
[[package]]
|
||||
name = "litellm"
|
||||
version = "0.0.1"
|
||||
|
||||
[[package]]
|
||||
name = "litellm-auth"
|
||||
version = "0.1.0"
|
||||
|
|
@ -2713,6 +2859,7 @@ dependencies = [
|
|||
"litellm-http",
|
||||
"moka",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
"sha2 0.10.9",
|
||||
"thiserror 2.0.19",
|
||||
|
|
@ -2784,6 +2931,7 @@ dependencies = [
|
|||
"litellm-cache",
|
||||
"litellm-cache-response",
|
||||
"litellm-cache-testing",
|
||||
"litellm-http",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
|
|
@ -2817,6 +2965,7 @@ dependencies = [
|
|||
"litellm-auth-types",
|
||||
"litellm-cache",
|
||||
"litellm-cache-testing",
|
||||
"litellm-http",
|
||||
"percent-encoding",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
|
|
@ -2843,6 +2992,7 @@ dependencies = [
|
|||
"futures-util",
|
||||
"litellm-cache",
|
||||
"litellm-cache-testing",
|
||||
"litellm-http",
|
||||
"qdrant-client",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
|
|
@ -2916,6 +3066,7 @@ dependencies = [
|
|||
"litellm-auth-aws",
|
||||
"litellm-cache",
|
||||
"litellm-cache-testing",
|
||||
"litellm-http",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
|
|
@ -2960,6 +3111,18 @@ dependencies = [
|
|||
"strum",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-config"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"litellm-auth-types",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_yaml_ng",
|
||||
"tempfile",
|
||||
"thiserror 2.0.19",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-core"
|
||||
version = "0.1.0"
|
||||
|
|
@ -2975,6 +3138,7 @@ dependencies = [
|
|||
"litellm-http",
|
||||
"litellm-llms",
|
||||
"litellm-secrets",
|
||||
"litellm-tracing",
|
||||
"litellm-types",
|
||||
"mime_guess",
|
||||
"moka",
|
||||
|
|
@ -2995,6 +3159,7 @@ dependencies = [
|
|||
"tokio-tungstenite",
|
||||
"url",
|
||||
"veil",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3009,6 +3174,7 @@ dependencies = [
|
|||
"serde_json",
|
||||
"serde_path_to_error",
|
||||
"serde_with",
|
||||
"strum",
|
||||
"thiserror 2.0.19",
|
||||
"url",
|
||||
]
|
||||
|
|
@ -3038,10 +3204,74 @@ dependencies = [
|
|||
"aws-smithy-types",
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"proptest",
|
||||
"rstest",
|
||||
"sse-stream",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-gateway"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"axum",
|
||||
"futures-util",
|
||||
"http-body-util",
|
||||
"litellm-config",
|
||||
"litellm-core",
|
||||
"litellm-gateway-auth",
|
||||
"litellm-gateway-inference",
|
||||
"litellm-http",
|
||||
"litellm-llms",
|
||||
"litellm-secrets",
|
||||
"litellm-tracing",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"tower",
|
||||
"tracing",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-gateway-auth"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"axum",
|
||||
"futures-util",
|
||||
"litellm-auth-types",
|
||||
"litellm-config",
|
||||
"litellm-secrets",
|
||||
"rstest",
|
||||
"sha2 0.10.9",
|
||||
"subtle",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tower",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-gateway-inference"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"axum",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"litellm-auth",
|
||||
"litellm-core",
|
||||
"litellm-http",
|
||||
"litellm-llms",
|
||||
"litellm-router",
|
||||
"litellm-secrets",
|
||||
"litellm-types",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tower",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3078,11 +3308,13 @@ dependencies = [
|
|||
"http 1.4.2",
|
||||
"hyper-util",
|
||||
"litellm-core-utils",
|
||||
"rcgen",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"rustls 0.23.42",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tempfile",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"veil",
|
||||
|
|
@ -3107,6 +3339,7 @@ dependencies = [
|
|||
"litellm-framing",
|
||||
"litellm-host",
|
||||
"litellm-http",
|
||||
"litellm-python-compat",
|
||||
"litellm-secrets",
|
||||
"litellm-types",
|
||||
"reqwest 0.12.28",
|
||||
|
|
@ -3126,14 +3359,15 @@ dependencies = [
|
|||
name = "litellm-model-catalog"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"criterion",
|
||||
"indexmap 2.14.0",
|
||||
"litellm-model-catalog",
|
||||
"jsonschema",
|
||||
"litellm-types",
|
||||
"rstest",
|
||||
"schemars 1.2.2",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"thiserror 2.0.19",
|
||||
"time",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3147,7 +3381,6 @@ dependencies = [
|
|||
"futures-util",
|
||||
"litellm-auth",
|
||||
"litellm-auth-aws",
|
||||
"litellm-auth-gcp",
|
||||
"litellm-cache",
|
||||
"litellm-cache-azure-blob",
|
||||
"litellm-cache-disk",
|
||||
|
|
@ -3182,6 +3415,7 @@ dependencies = [
|
|||
"serde_json",
|
||||
"serde_with",
|
||||
"sha2 0.10.9",
|
||||
"strum",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
|
|
@ -3205,6 +3439,15 @@ dependencies = [
|
|||
"thiserror 2.0.19",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-router"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"litellm-config",
|
||||
"litellm-core",
|
||||
"rstest",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-secrets"
|
||||
version = "0.1.0"
|
||||
|
|
@ -3215,6 +3458,7 @@ dependencies = [
|
|||
"google-cloud-auth",
|
||||
"google-cloud-kms-v1",
|
||||
"litellm-core-utils",
|
||||
"litellm-http",
|
||||
"litellm-python-compat",
|
||||
"litellm-secrets-aws",
|
||||
"litellm-secrets-azure",
|
||||
|
|
@ -3262,6 +3506,7 @@ dependencies = [
|
|||
"litellm-auth-azure",
|
||||
"litellm-auth-types",
|
||||
"litellm-core-utils",
|
||||
"litellm-http",
|
||||
"litellm-secrets-types",
|
||||
"percent-encoding",
|
||||
"reqwest 0.12.28",
|
||||
|
|
@ -3281,6 +3526,7 @@ version = "0.1.0"
|
|||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"litellm-core-utils",
|
||||
"litellm-http",
|
||||
"litellm-secrets-types",
|
||||
"litellm-tracing",
|
||||
"moka",
|
||||
|
|
@ -3309,6 +3555,7 @@ dependencies = [
|
|||
"litellm-auth-gcp",
|
||||
"litellm-auth-types",
|
||||
"litellm-core-utils",
|
||||
"litellm-http",
|
||||
"litellm-secrets-types",
|
||||
"moka",
|
||||
"percent-encoding",
|
||||
|
|
@ -3350,11 +3597,33 @@ dependencies = [
|
|||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"strum",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"veil",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-testkit"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"flate2",
|
||||
"futures-util",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"semver",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2 0.10.9",
|
||||
"tar",
|
||||
"target-lexicon",
|
||||
"tempfile",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"toml",
|
||||
"zip",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-token-counter"
|
||||
version = "0.1.0"
|
||||
|
|
@ -3412,6 +3681,7 @@ dependencies = [
|
|||
name = "litellm-tracing"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"fancy-regex 0.19.2",
|
||||
"percent-encoding",
|
||||
"rstest",
|
||||
|
|
@ -3426,8 +3696,10 @@ name = "litellm-types"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"rstest",
|
||||
"schemars 1.2.2",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"strum",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3504,6 +3776,12 @@ version = "2.8.3"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98"
|
||||
|
||||
[[package]]
|
||||
name = "micromap"
|
||||
version = "0.3.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c2a86d3146ed3995b5913c414f6664344b9617457320782e64f0bb44afd49d74"
|
||||
|
||||
[[package]]
|
||||
name = "mime"
|
||||
version = "0.3.17"
|
||||
|
|
@ -3589,6 +3867,23 @@ dependencies = [
|
|||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "multer"
|
||||
version = "3.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "83e87776546dc87511aa5ee218730c92b666d7264ab6ed41f9d215af9cd5224b"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"encoding_rs",
|
||||
"futures-util",
|
||||
"http 1.4.2",
|
||||
"httparse",
|
||||
"memchr",
|
||||
"mime",
|
||||
"spin 0.9.9",
|
||||
"version_check",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "nom"
|
||||
version = "7.1.3"
|
||||
|
|
@ -3599,6 +3894,20 @@ dependencies = [
|
|||
"minimal-lexical",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num"
|
||||
version = "0.4.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "35bd024e8b2ff75562e5f34e7f4905839deb4b22955ef5e73d2fea1b9813cb23"
|
||||
dependencies = [
|
||||
"num-bigint 0.4.8",
|
||||
"num-complex",
|
||||
"num-integer",
|
||||
"num-iter",
|
||||
"num-rational",
|
||||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-bigint"
|
||||
version = "0.4.8"
|
||||
|
|
@ -3619,6 +3928,12 @@ dependencies = [
|
|||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-cmp"
|
||||
version = "0.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "63335b2e2c34fae2fb0aa2cecfd9f0832a1e24b3b32ecec612c3426d46dc8aaa"
|
||||
|
||||
[[package]]
|
||||
name = "num-complex"
|
||||
version = "0.4.6"
|
||||
|
|
@ -3643,6 +3958,27 @@ dependencies = [
|
|||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-iter"
|
||||
version = "0.1.46"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c92800bd69a1eac91786bcfe9da64a897eb72911b8dc3095decbd07429e8048b"
|
||||
dependencies = [
|
||||
"num-integer",
|
||||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-rational"
|
||||
version = "0.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f83d14da390562dca69fc84082e73e548e1ad308d24accdedd2720017cb37824"
|
||||
dependencies = [
|
||||
"num-bigint 0.4.8",
|
||||
"num-integer",
|
||||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-traits"
|
||||
version = "0.2.19"
|
||||
|
|
@ -4455,7 +4791,24 @@ checksum = "92ecd8964f8453721699a1ed72037b0db49ce2f5a5138486ee89bed6f67cdf3a"
|
|||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 3.0.0",
|
||||
"syn 3.0.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "referencing"
|
||||
version = "0.55.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a196a5b4a8a12f46b6353174df865a05d41a6055aff212ec30877492788618b6"
|
||||
dependencies = [
|
||||
"ahash",
|
||||
"fluent-uri",
|
||||
"getrandom 0.3.4",
|
||||
"hashbrown 0.17.1",
|
||||
"itoa",
|
||||
"micromap",
|
||||
"parking_lot",
|
||||
"percent-encoding",
|
||||
"serde_json",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -4917,7 +5270,7 @@ dependencies = [
|
|||
"proc-macro2",
|
||||
"quote",
|
||||
"serde_derive_internals",
|
||||
"syn 3.0.0",
|
||||
"syn 3.0.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5005,7 +5358,7 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348"
|
|||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 3.0.0",
|
||||
"syn 3.0.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5016,7 +5369,7 @@ checksum = "f852137cce035d6a4df67ccce505ff6b3e9fd3a10e3e52b24dc71e650bb1a9bd"
|
|||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 3.0.0",
|
||||
"syn 3.0.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5044,6 +5397,15 @@ dependencies = [
|
|||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "serde_spanned"
|
||||
version = "1.1.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6662b5879511e06e8999a8a235d848113e942c9124f211511b16466ee2995f26"
|
||||
dependencies = [
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "serde_urlencoded"
|
||||
version = "0.7.1"
|
||||
|
|
@ -5087,6 +5449,19 @@ dependencies = [
|
|||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "serde_yaml_ng"
|
||||
version = "0.10.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7b4db627b98b36d4203a7b458cf3573730f2bb591b28871d916dfa9efabfd41f"
|
||||
dependencies = [
|
||||
"indexmap 2.14.0",
|
||||
"itoa",
|
||||
"ryu",
|
||||
"serde",
|
||||
"unsafe-libyaml",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sha1"
|
||||
version = "0.10.7"
|
||||
|
|
@ -5216,6 +5591,12 @@ dependencies = [
|
|||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "spin"
|
||||
version = "0.9.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e"
|
||||
|
||||
[[package]]
|
||||
name = "spin"
|
||||
version = "0.10.1"
|
||||
|
|
@ -5246,19 +5627,6 @@ dependencies = [
|
|||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sse-stream"
|
||||
version = "0.2.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c25ac7aff0abd1dbc474536e40416e1102c7dd9bfba0b9861c6d357f835dcfb4"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"http-body 1.1.0",
|
||||
"http-body-util",
|
||||
"pin-project-lite",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "stable_deref_trait"
|
||||
version = "1.2.1"
|
||||
|
|
@ -5328,9 +5696,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "syn"
|
||||
version = "3.0.0"
|
||||
version = "3.0.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f2fac314a64dc9a36e61a9eb4261a5e9bbfbc922b27e518af97bc32b926cf967"
|
||||
checksum = "8593e8e72159ed2257d083c7a454a85cbf854f37a0966d8d483aff8c8a3ebcee"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
|
|
@ -5375,6 +5743,17 @@ version = "0.2.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417"
|
||||
|
||||
[[package]]
|
||||
name = "tar"
|
||||
version = "0.4.46"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840"
|
||||
dependencies = [
|
||||
"filetime",
|
||||
"libc",
|
||||
"xattr",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "target-lexicon"
|
||||
version = "0.13.5"
|
||||
|
|
@ -5431,7 +5810,7 @@ checksum = "43cbfe0cf76104d42a574802844187e84a305e531ed54455f11fbde0f10541cd"
|
|||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 3.0.0",
|
||||
"syn 3.0.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5644,6 +6023,30 @@ dependencies = [
|
|||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml"
|
||||
version = "0.9.12+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "cf92845e79fc2e2def6a5d828f0801e29a2f8acc037becc5ab08595c7d5e9863"
|
||||
dependencies = [
|
||||
"indexmap 2.14.0",
|
||||
"serde_core",
|
||||
"serde_spanned",
|
||||
"toml_datetime 0.7.5+spec-1.1.0",
|
||||
"toml_parser",
|
||||
"toml_writer",
|
||||
"winnow 0.7.15",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml_datetime"
|
||||
version = "0.7.5+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "92e1cfed4a3038bc5a127e35a2d360f145e1f4b971b551a2ba5fd7aedf7e1347"
|
||||
dependencies = [
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml_datetime"
|
||||
version = "1.1.1+spec-1.1.0"
|
||||
|
|
@ -5660,9 +6063,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "6975367e4d2ef766d86af01ffad14b622fecc8d4357a998fbc4deb6e9bacaf9b"
|
||||
dependencies = [
|
||||
"indexmap 2.14.0",
|
||||
"toml_datetime",
|
||||
"toml_datetime 1.1.1+spec-1.1.0",
|
||||
"toml_parser",
|
||||
"winnow",
|
||||
"winnow 1.0.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5671,9 +6074,15 @@ version = "1.1.3+spec-1.1.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1d38ac1cf9b95face32296c0a3ede1fdc270627c9d9c02a7274dd6d960dc4d56"
|
||||
dependencies = [
|
||||
"winnow",
|
||||
"winnow 1.0.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml_writer"
|
||||
version = "1.1.2+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7d56353a2a665ad0f41a421187180aab746c8c325620617ad883a99a1cbe66d2"
|
||||
|
||||
[[package]]
|
||||
name = "tonic"
|
||||
version = "0.14.6"
|
||||
|
|
@ -5941,6 +6350,12 @@ version = "2.9.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142"
|
||||
|
||||
[[package]]
|
||||
name = "unicode-general-category"
|
||||
version = "1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0b993bddc193ae5bd0d623b49ec06ac3e9312875fdae725a975c51db1cc1677f"
|
||||
|
||||
[[package]]
|
||||
name = "unicode-ident"
|
||||
version = "1.0.24"
|
||||
|
|
@ -5974,6 +6389,12 @@ version = "0.1.1"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e"
|
||||
|
||||
[[package]]
|
||||
name = "unsafe-libyaml"
|
||||
version = "0.2.11"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "673aac59facbab8a9007c7f6108d11f63b603f7cabff99fabf650fea5c32b861"
|
||||
|
||||
[[package]]
|
||||
name = "untrusted"
|
||||
version = "0.9.0"
|
||||
|
|
@ -6021,6 +6442,16 @@ dependencies = [
|
|||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "uuid-simd"
|
||||
version = "0.8.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "23b082222b4f6619906941c17eb2297fff4c2fb96cb60164170522942a200bd8"
|
||||
dependencies = [
|
||||
"outref",
|
||||
"vsimd",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "valuable"
|
||||
version = "0.1.1"
|
||||
|
|
@ -6419,6 +6850,12 @@ version = "0.52.6"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec"
|
||||
|
||||
[[package]]
|
||||
name = "winnow"
|
||||
version = "0.7.15"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945"
|
||||
|
||||
[[package]]
|
||||
name = "winnow"
|
||||
version = "1.0.4"
|
||||
|
|
@ -6481,6 +6918,16 @@ dependencies = [
|
|||
"time",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "xattr"
|
||||
version = "1.6.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"rustix",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "xmlparser"
|
||||
version = "0.13.6"
|
||||
|
|
@ -6606,6 +7053,23 @@ dependencies = [
|
|||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "zip"
|
||||
version = "2.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fabe6324e908f85a1c52063ce7aa26b68dcb7eb6dbc83a2d148403c9bc3eba50"
|
||||
dependencies = [
|
||||
"arbitrary",
|
||||
"crc32fast",
|
||||
"crossbeam-utils",
|
||||
"displaydoc",
|
||||
"flate2",
|
||||
"indexmap 2.14.0",
|
||||
"memchr",
|
||||
"thiserror 2.0.19",
|
||||
"zopfli",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "zlib-rs"
|
||||
version = "0.6.7"
|
||||
|
|
@ -6617,3 +7081,15 @@ name = "zmij"
|
|||
version = "1.0.23"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b"
|
||||
|
||||
[[package]]
|
||||
name = "zopfli"
|
||||
version = "0.8.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f05cd8797d63865425ff89b5c4a48804f35ba0ce8d125800027ad6017d2b5249"
|
||||
dependencies = [
|
||||
"bumpalo",
|
||||
"crc32fast",
|
||||
"log",
|
||||
"simd-adler32",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -9,9 +9,13 @@ license = "MIT"
|
|||
repository = "https://github.com/BerriAI/litellm"
|
||||
|
||||
[workspace.dependencies]
|
||||
litellm-config = { path = "crates/config" }
|
||||
litellm-router = { path = "crates/router" }
|
||||
litellm-tracing = { path = "crates/tracing" }
|
||||
tracing = "0.1"
|
||||
litellm-core = { path = "crates/core" }
|
||||
litellm-gateway = { path = "crates/gateway" }
|
||||
litellm-gateway-inference = { path = "crates/gateway-inference" }
|
||||
litellm-gateway-auth = { path = "crates/gateway-auth" }
|
||||
litellm-coroutine = { path = "crates/coroutine" }
|
||||
litellm-host = { path = "crates/host" }
|
||||
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
|
||||
|
|
@ -48,7 +52,10 @@ litellm-token-counter-fast = { path = "crates/token-counter-fast" }
|
|||
litellm-token-counter-huggingface = { path = "crates/token-counter-huggingface" }
|
||||
litellm-token-counter-tiktoken = { path = "crates/token-counter-tiktoken" }
|
||||
litellm-host-python = { path = "crates/host-python" }
|
||||
litellm-python-compat = { path = "crates/python-compat" }
|
||||
|
||||
tracing = "0.1"
|
||||
axum = { version = "0.8.9", default-features = false, features = ["http1", "tokio", "multipart"] }
|
||||
bytes = "1"
|
||||
http = "1"
|
||||
google-cloud-auth = { version = "1.16.0", default-features = false }
|
||||
|
|
@ -57,8 +64,8 @@ hyper-util = { version = "0.1.20", default-features = false, features = ["client
|
|||
proptest = "1.7.0"
|
||||
pyo3 = "0.29.2"
|
||||
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
|
||||
pythonize = "0.29.0"
|
||||
rand = "0.8"
|
||||
schemars = "1"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls", "http2", "stream"] }
|
||||
qdrant-client = { version = "1.19.0", default-features = false }
|
||||
uuid = { version = "1", features = ["v4"] }
|
||||
|
|
@ -81,6 +88,12 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"]
|
|||
tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
|
||||
futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
|
||||
base64 = "0.22"
|
||||
flate2 = "1"
|
||||
semver = "1"
|
||||
tar = "0.4"
|
||||
target-lexicon = "0.13.5"
|
||||
tempfile = "3"
|
||||
zip = { version = "2", default-features = false, features = ["deflate"] }
|
||||
moka = { version = "0.12.16", features = ["future"] }
|
||||
strum = { version = "0.28.0", features = ["derive"] }
|
||||
url = "2.5.8"
|
||||
|
|
|
|||
|
|
@ -7,4 +7,16 @@ disallowed-methods = [
|
|||
{ path = "pyo3_async_runtimes::tokio::local_future_into_py", reason = "use litellm_host_python::run_async / run_async_value" },
|
||||
{ path = "pyo3_async_runtimes::tokio::run", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
{ path = "pyo3_async_runtimes::tokio::run_until_complete", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
{ path = "reqwest::Client::new", reason = "take litellm_http::Client from HttpClientPool" },
|
||||
{ path = "reqwest::Client::builder", reason = "HttpClientConfig owns client construction" },
|
||||
{ path = "reqwest::ClientBuilder::danger_accept_invalid_certs", reason = "set HttpClientConfig::verify instead" },
|
||||
{ path = "reqwest::ClientBuilder::identity", reason = "set HttpClientConfig::client_certificate instead" },
|
||||
{ path = "reqwest::ClientBuilder::use_preconfigured_tls", reason = "HttpClientConfig owns the TLS configuration" },
|
||||
]
|
||||
|
||||
# Every outbound client comes from litellm_http::HttpClientPool so it honors the host's TLS,
|
||||
# proxy and timeout settings. Only crates/http builds one.
|
||||
disallowed-types = [
|
||||
{ path = "reqwest::Client", reason = "take litellm_http::Client from HttpClientPool; only crates/http builds one" },
|
||||
{ path = "reqwest::ClientBuilder", reason = "HttpClientConfig owns client construction" },
|
||||
]
|
||||
|
|
|
|||
|
|
@ -22,5 +22,7 @@ aws-types = "1.4.0"
|
|||
aws-smithy-runtime-api = "1.13.0"
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
reqwest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::sync::OnceLock;
|
||||
use std::time::Duration;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
|
|
@ -26,8 +25,26 @@ use super::constants::{
|
|||
const STATIC_CREDENTIALS_TTL: Duration = Duration::from_secs(3600 - 60);
|
||||
const AMBIENT_CREDENTIALS_TTL: Duration = Duration::from_secs(600);
|
||||
|
||||
static STATIC_CREDENTIALS_CACHE: OnceLock<Cache<String, Credentials>> = OnceLock::new();
|
||||
static AMBIENT_CREDENTIALS_CACHE: OnceLock<Cache<String, Credentials>> = OnceLock::new();
|
||||
#[derive(Clone)]
|
||||
pub struct AwsAuthService {
|
||||
static_credentials: Cache<String, Credentials>,
|
||||
ambient_credentials: Cache<String, Credentials>,
|
||||
}
|
||||
|
||||
impl Default for AwsAuthService {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
static_credentials: Cache::builder()
|
||||
.max_capacity(200)
|
||||
.time_to_live(STATIC_CREDENTIALS_TTL)
|
||||
.build(),
|
||||
ambient_credentials: Cache::builder()
|
||||
.max_capacity(200)
|
||||
.time_to_live(AMBIENT_CREDENTIALS_TTL)
|
||||
.build(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn credential_cache_ttl(flow: &AwsAuthFlow) -> Option<Duration> {
|
||||
match flow {
|
||||
|
|
@ -108,35 +125,19 @@ fn cache_key(config: &AwsAuthConfig, flow: &AwsAuthFlow) -> String {
|
|||
format!("{:x}", hasher.finalize())
|
||||
}
|
||||
|
||||
fn static_credentials_cache() -> &'static Cache<String, Credentials> {
|
||||
STATIC_CREDENTIALS_CACHE.get_or_init(|| {
|
||||
Cache::builder()
|
||||
.max_capacity(200)
|
||||
.time_to_live(STATIC_CREDENTIALS_TTL)
|
||||
.build()
|
||||
})
|
||||
}
|
||||
impl AwsAuthService {
|
||||
fn get_cached_credentials(&self, key: &str) -> Option<Credentials> {
|
||||
self.static_credentials
|
||||
.get(key)
|
||||
.or_else(|| self.ambient_credentials.get(key))
|
||||
}
|
||||
|
||||
fn ambient_credentials_cache() -> &'static Cache<String, Credentials> {
|
||||
AMBIENT_CREDENTIALS_CACHE.get_or_init(|| {
|
||||
Cache::builder()
|
||||
.max_capacity(200)
|
||||
.time_to_live(AMBIENT_CREDENTIALS_TTL)
|
||||
.build()
|
||||
})
|
||||
}
|
||||
|
||||
fn get_cached_credentials(key: &str) -> Option<Credentials> {
|
||||
static_credentials_cache()
|
||||
.get(key)
|
||||
.or_else(|| ambient_credentials_cache().get(key))
|
||||
}
|
||||
|
||||
fn set_cached_credentials(key: String, credentials: Credentials, ttl: Duration) {
|
||||
if ttl == STATIC_CREDENTIALS_TTL {
|
||||
static_credentials_cache().insert(key, credentials);
|
||||
} else {
|
||||
ambient_credentials_cache().insert(key, credentials);
|
||||
fn set_cached_credentials(&self, key: String, credentials: Credentials, ttl: Duration) {
|
||||
if ttl == STATIC_CREDENTIALS_TTL {
|
||||
self.static_credentials.insert(key, credentials);
|
||||
} else {
|
||||
self.ambient_credentials.insert(key, credentials);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -214,66 +215,157 @@ pub fn classify_auth(
|
|||
AwsAuthFlow::DefaultChain
|
||||
}
|
||||
|
||||
pub async fn resolve_credentials(
|
||||
config: AwsAuthConfig,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Credentials, Error> {
|
||||
let resolved = config.clone().with_environment(env_lookup);
|
||||
let flow = classify_auth(config, env_lookup);
|
||||
match flow {
|
||||
AwsAuthFlow::SessionToken {
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
session_token,
|
||||
} => Ok(Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
Some(session_token),
|
||||
None,
|
||||
"litellm-static-session",
|
||||
)),
|
||||
AwsAuthFlow::StaticKeys {
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
region_name,
|
||||
} => {
|
||||
let flow = AwsAuthFlow::StaticKeys {
|
||||
access_key_id: access_key_id.clone(),
|
||||
secret_access_key: secret_access_key.clone(),
|
||||
region_name,
|
||||
};
|
||||
let key = cache_key(&resolved, &flow);
|
||||
if let Some(credentials) = get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let credentials = Credentials::new(
|
||||
impl AwsAuthService {
|
||||
pub async fn resolve_credentials(
|
||||
&self,
|
||||
config: AwsAuthConfig,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Credentials, Error> {
|
||||
let resolved = config.clone().with_environment(env_lookup);
|
||||
let flow = classify_auth(config, env_lookup);
|
||||
match flow {
|
||||
AwsAuthFlow::SessionToken {
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
session_token,
|
||||
} => Ok(Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
Some(session_token),
|
||||
None,
|
||||
None,
|
||||
"litellm-static",
|
||||
);
|
||||
set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&flow).unwrap_or(STATIC_CREDENTIALS_TTL),
|
||||
);
|
||||
Ok(credentials)
|
||||
}
|
||||
AwsAuthFlow::Profile { name } => {
|
||||
let provider = aws_config::profile::ProfileFileCredentialsProvider::builder()
|
||||
.profile_name(name)
|
||||
.build();
|
||||
provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsProfile(error.to_string()))
|
||||
}
|
||||
AwsAuthFlow::AssumeRole { role, session_name } => {
|
||||
if is_already_running_as_role(&role, &resolved).await? {
|
||||
let ambient_flow = AwsAuthFlow::DefaultChain;
|
||||
let key = cache_key(&resolved, &ambient_flow);
|
||||
if let Some(credentials) = get_cached_credentials(&key) {
|
||||
"litellm-static-session",
|
||||
)),
|
||||
AwsAuthFlow::StaticKeys {
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
region_name,
|
||||
} => {
|
||||
let flow = AwsAuthFlow::StaticKeys {
|
||||
access_key_id: access_key_id.clone(),
|
||||
secret_access_key: secret_access_key.clone(),
|
||||
region_name,
|
||||
};
|
||||
let key = cache_key(&resolved, &flow);
|
||||
if let Some(credentials) = self.get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let credentials = Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
None,
|
||||
None,
|
||||
"litellm-static",
|
||||
);
|
||||
self.set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&flow).unwrap_or(STATIC_CREDENTIALS_TTL),
|
||||
);
|
||||
Ok(credentials)
|
||||
}
|
||||
AwsAuthFlow::Profile { name } => {
|
||||
let provider = aws_config::profile::ProfileFileCredentialsProvider::builder()
|
||||
.profile_name(name)
|
||||
.build();
|
||||
provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsProfile(error.to_string()))
|
||||
}
|
||||
AwsAuthFlow::AssumeRole { role, session_name } => {
|
||||
if is_already_running_as_role(&role, &resolved).await? {
|
||||
let ambient_flow = AwsAuthFlow::DefaultChain;
|
||||
let key = cache_key(&resolved, &ambient_flow);
|
||||
if let Some(credentials) = self.get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let provider =
|
||||
aws_config::default_provider::credentials::DefaultCredentialsChain::builder()
|
||||
.build()
|
||||
.await;
|
||||
let credentials = provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsDefaultChain(error.to_string()))?;
|
||||
self.set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL),
|
||||
);
|
||||
return Ok(credentials);
|
||||
}
|
||||
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
|
||||
if let Some(region) = resolved.region_name.clone() {
|
||||
loader = loader.region(aws_types::region::Region::new(region));
|
||||
}
|
||||
if let Some(endpoint) = resolved.sts_endpoint.clone() {
|
||||
loader = loader.endpoint_url(endpoint);
|
||||
}
|
||||
if let (Some(access_key_id), Some(secret_access_key)) =
|
||||
(resolved.access_key_id, resolved.secret_access_key)
|
||||
{
|
||||
loader = loader.credentials_provider(Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
resolved.session_token,
|
||||
None,
|
||||
"litellm-role-source",
|
||||
));
|
||||
}
|
||||
let sdk_config = loader.load().await;
|
||||
let builder = aws_config::sts::AssumeRoleProvider::builder(role);
|
||||
let builder = match session_name {
|
||||
Some(name) => builder.session_name(name),
|
||||
None => builder.session_name(default_session_name()),
|
||||
};
|
||||
let builder = match resolved.external_id {
|
||||
Some(id) => builder.external_id(id),
|
||||
None => builder,
|
||||
};
|
||||
let provider = builder.configure(&sdk_config).build().await;
|
||||
provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsAssumeRole(error.to_string()))
|
||||
}
|
||||
AwsAuthFlow::WebIdentity {
|
||||
token,
|
||||
role,
|
||||
session_name,
|
||||
} => {
|
||||
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
|
||||
if let Some(region) = resolved.region_name {
|
||||
loader = loader.region(aws_types::region::Region::new(region));
|
||||
}
|
||||
if let Some(endpoint) = resolved.sts_endpoint {
|
||||
loader = loader.endpoint_url(endpoint);
|
||||
}
|
||||
let sdk_config = loader.load().await;
|
||||
let client = aws_sdk_sts::Client::new(&sdk_config);
|
||||
let response = client
|
||||
.assume_role_with_web_identity()
|
||||
.role_arn(role)
|
||||
.role_session_name(session_name)
|
||||
.web_identity_token(token)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| Error::AwsWebIdentity(error.to_string()))?;
|
||||
let credentials = response
|
||||
.credentials()
|
||||
.ok_or(Error::AwsMissingWebIdentityCredentials)?;
|
||||
let expiration = SystemTime::try_from(*credentials.expiration())
|
||||
.map_err(|error| Error::AwsWebIdentityExpiration(error.to_string()))?;
|
||||
Ok(Credentials::new(
|
||||
credentials.access_key_id(),
|
||||
credentials.secret_access_key(),
|
||||
Some(credentials.session_token().to_string()),
|
||||
Some(expiration),
|
||||
"litellm-web-identity",
|
||||
))
|
||||
}
|
||||
AwsAuthFlow::DefaultChain => {
|
||||
let key = cache_key(&resolved, &AwsAuthFlow::DefaultChain);
|
||||
if let Some(credentials) = self.get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let provider =
|
||||
|
|
@ -284,101 +376,14 @@ pub async fn resolve_credentials(
|
|||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsDefaultChain(error.to_string()))?;
|
||||
set_cached_credentials(
|
||||
self.set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL),
|
||||
credential_cache_ttl(&AwsAuthFlow::DefaultChain)
|
||||
.unwrap_or(AMBIENT_CREDENTIALS_TTL),
|
||||
);
|
||||
return Ok(credentials);
|
||||
Ok(credentials)
|
||||
}
|
||||
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
|
||||
if let Some(region) = resolved.region_name.clone() {
|
||||
loader = loader.region(aws_types::region::Region::new(region));
|
||||
}
|
||||
if let Some(endpoint) = resolved.sts_endpoint.clone() {
|
||||
loader = loader.endpoint_url(endpoint);
|
||||
}
|
||||
if let (Some(access_key_id), Some(secret_access_key)) =
|
||||
(resolved.access_key_id, resolved.secret_access_key)
|
||||
{
|
||||
loader = loader.credentials_provider(Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
resolved.session_token,
|
||||
None,
|
||||
"litellm-role-source",
|
||||
));
|
||||
}
|
||||
let sdk_config = loader.load().await;
|
||||
let builder = aws_config::sts::AssumeRoleProvider::builder(role);
|
||||
let builder = match session_name {
|
||||
Some(name) => builder.session_name(name),
|
||||
None => builder.session_name(default_session_name()),
|
||||
};
|
||||
let builder = match resolved.external_id {
|
||||
Some(id) => builder.external_id(id),
|
||||
None => builder,
|
||||
};
|
||||
let provider = builder.configure(&sdk_config).build().await;
|
||||
provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsAssumeRole(error.to_string()))
|
||||
}
|
||||
AwsAuthFlow::WebIdentity {
|
||||
token,
|
||||
role,
|
||||
session_name,
|
||||
} => {
|
||||
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
|
||||
if let Some(region) = resolved.region_name {
|
||||
loader = loader.region(aws_types::region::Region::new(region));
|
||||
}
|
||||
if let Some(endpoint) = resolved.sts_endpoint {
|
||||
loader = loader.endpoint_url(endpoint);
|
||||
}
|
||||
let sdk_config = loader.load().await;
|
||||
let client = aws_sdk_sts::Client::new(&sdk_config);
|
||||
let response = client
|
||||
.assume_role_with_web_identity()
|
||||
.role_arn(role)
|
||||
.role_session_name(session_name)
|
||||
.web_identity_token(token)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| Error::AwsWebIdentity(error.to_string()))?;
|
||||
let credentials = response
|
||||
.credentials()
|
||||
.ok_or(Error::AwsMissingWebIdentityCredentials)?;
|
||||
let expiration = SystemTime::try_from(*credentials.expiration())
|
||||
.map_err(|error| Error::AwsWebIdentityExpiration(error.to_string()))?;
|
||||
Ok(Credentials::new(
|
||||
credentials.access_key_id(),
|
||||
credentials.secret_access_key(),
|
||||
Some(credentials.session_token().to_string()),
|
||||
Some(expiration),
|
||||
"litellm-web-identity",
|
||||
))
|
||||
}
|
||||
AwsAuthFlow::DefaultChain => {
|
||||
let key = cache_key(&resolved, &AwsAuthFlow::DefaultChain);
|
||||
if let Some(credentials) = get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let provider =
|
||||
aws_config::default_provider::credentials::DefaultCredentialsChain::builder()
|
||||
.build()
|
||||
.await;
|
||||
let credentials = provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsDefaultChain(error.to_string()))?;
|
||||
set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&AwsAuthFlow::DefaultChain).unwrap_or(AMBIENT_CREDENTIALS_TTL),
|
||||
);
|
||||
Ok(credentials)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -585,6 +590,37 @@ pub fn aws_auth_config(
|
|||
}
|
||||
}
|
||||
|
||||
/// Where the credentials that sign a request come from, decided when the request is
|
||||
/// prepared and resolved when it is sent.
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub enum AwsCredentialSource {
|
||||
HostSupplied(Credentials),
|
||||
Chain(AwsAuthConfig),
|
||||
}
|
||||
|
||||
impl AwsCredentialSource {
|
||||
pub fn from_params(
|
||||
optional_params: &Map<String, Value>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Self {
|
||||
match host_supplied_credentials(optional_params) {
|
||||
Some(credentials) => Self::HostSupplied(credentials),
|
||||
None => Self::Chain(aws_auth_config(optional_params, env_lookup)),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn resolve(
|
||||
self,
|
||||
auth: &AwsAuthService,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Credentials, Error> {
|
||||
match self {
|
||||
Self::HostSupplied(credentials) => Ok(credentials),
|
||||
Self::Chain(config) => auth.resolve_credentials(config, env_lookup).await,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Credentials a host resolved through its own chain and handed down verbatim.
|
||||
///
|
||||
/// A host with its own resolution (LiteLLM's Python `BaseAWSLLM`, which reads
|
||||
|
|
@ -747,17 +783,18 @@ mod tests {
|
|||
|
||||
#[tokio::test]
|
||||
async fn static_credentials_do_not_use_network() {
|
||||
let credentials = resolve_credentials(
|
||||
AwsAuthConfig {
|
||||
access_key_id: Some("ak".into()),
|
||||
secret_access_key: Some("sk".into()),
|
||||
region_name: Some("us-east-1".into()),
|
||||
..Default::default()
|
||||
},
|
||||
&no_env,
|
||||
)
|
||||
.await
|
||||
.expect("static credentials");
|
||||
let credentials = AwsAuthService::default()
|
||||
.resolve_credentials(
|
||||
AwsAuthConfig {
|
||||
access_key_id: Some("ak".into()),
|
||||
secret_access_key: Some("sk".into()),
|
||||
region_name: Some("us-east-1".into()),
|
||||
..Default::default()
|
||||
},
|
||||
&no_env,
|
||||
)
|
||||
.await
|
||||
.expect("static credentials");
|
||||
assert_eq!(credentials.access_key_id(), "ak");
|
||||
assert_eq!(credentials.session_token(), None);
|
||||
}
|
||||
|
|
@ -807,17 +844,67 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn cache_round_trip_preserves_credentials() {
|
||||
let auth = AwsAuthService::default();
|
||||
let key = format!("cache-test-{}", std::process::id());
|
||||
let credentials = Credentials::new("cache-ak", "cache-sk", None, None, "test");
|
||||
set_cached_credentials(key.clone(), credentials.clone(), STATIC_CREDENTIALS_TTL);
|
||||
auth.set_cached_credentials(key.clone(), credentials.clone(), STATIC_CREDENTIALS_TTL);
|
||||
assert_eq!(
|
||||
get_cached_credentials(&key).map(|value| value.access_key_id().to_string()),
|
||||
auth.get_cached_credentials(&key)
|
||||
.map(|value| value.access_key_id().to_string()),
|
||||
Some("cache-ak".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn cloned_services_reuse_credentials_but_independent_services_do_not() {
|
||||
let auth = AwsAuthService::default();
|
||||
let config = AwsAuthConfig {
|
||||
access_key_id: Some("configured-key".into()),
|
||||
secret_access_key: Some("configured-secret".into()),
|
||||
region_name: Some("us-east-1".into()),
|
||||
..AwsAuthConfig::default()
|
||||
};
|
||||
let flow = classify_auth(config.clone(), &no_env);
|
||||
let cached = Credentials::new("cached-key", "cached-secret", None, None, "test");
|
||||
auth.set_cached_credentials(
|
||||
cache_key(&config, &flow),
|
||||
cached.clone(),
|
||||
STATIC_CREDENTIALS_TTL,
|
||||
);
|
||||
|
||||
let reused = auth
|
||||
.clone()
|
||||
.resolve_credentials(config.clone(), &no_env)
|
||||
.await
|
||||
.unwrap();
|
||||
let independent = AwsAuthService::default()
|
||||
.resolve_credentials(config.clone(), &no_env)
|
||||
.await
|
||||
.unwrap();
|
||||
let different = AwsAuthConfig {
|
||||
access_key_id: Some("different-key".into()),
|
||||
..config.clone()
|
||||
};
|
||||
let other_identity = auth
|
||||
.resolve_credentials(different.clone(), &no_env)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(reused.access_key_id(), cached.access_key_id());
|
||||
assert_eq!(reused.secret_access_key(), cached.secret_access_key());
|
||||
assert_eq!(
|
||||
Some(independent.access_key_id()),
|
||||
config.access_key_id.as_deref()
|
||||
);
|
||||
assert_eq!(
|
||||
Some(other_identity.access_key_id()),
|
||||
different.access_key_id.as_deref()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn same_role_comparison_matches_partition_account_and_role() {
|
||||
assert!(same_role_arns(
|
||||
|
|
@ -952,17 +1039,18 @@ mod tests {
|
|||
let body = br#"{"anthropic_version":"bedrock-2023-05-31","max_tokens":1,"messages":[{"role":"user","content":[{"type":"text","text":"ping"}]}]}"#.to_vec();
|
||||
let headers =
|
||||
BTreeMap::from([("Content-Type".to_string(), "application/json".to_string())]);
|
||||
let credentials = resolve_credentials(
|
||||
AwsAuthConfig {
|
||||
access_key_id: Some(access_key_id),
|
||||
secret_access_key: Some(secret_access_key),
|
||||
region_name: Some("us-west-2".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
&no_env,
|
||||
)
|
||||
.await?;
|
||||
let client = reqwest::Client::new();
|
||||
let credentials = AwsAuthService::default()
|
||||
.resolve_credentials(
|
||||
AwsAuthConfig {
|
||||
access_key_id: Some(access_key_id),
|
||||
secret_access_key: Some(secret_access_key),
|
||||
region_name: Some("us-west-2".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
&no_env,
|
||||
)
|
||||
.await?;
|
||||
let client = litellm_http::Client::plain_for_test();
|
||||
let mut failures = Vec::new();
|
||||
|
||||
for region in ["us-west-2", "us-east-1"] {
|
||||
|
|
|
|||
|
|
@ -1,13 +1,11 @@
|
|||
use std::{collections::BTreeMap, time::SystemTime};
|
||||
|
||||
use crate::{
|
||||
AwsAuthService, AwsCredentialSource, Error, aws_signature_headers, is_sigv4_computed_header,
|
||||
sign_post,
|
||||
};
|
||||
use aws_credential_types::Credentials;
|
||||
use litellm_http::outbound::{RequestSigner, UnsignedRequest};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{
|
||||
Error, aws_auth_config, aws_signature_headers, host_supplied_credentials,
|
||||
is_sigv4_computed_header, resolve_credentials, sign_post,
|
||||
};
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct SigV4Signer {
|
||||
|
|
@ -32,19 +30,17 @@ impl SigV4Signer {
|
|||
}
|
||||
|
||||
pub async fn resolve(
|
||||
auth: &AwsAuthService,
|
||||
region: String,
|
||||
service: &'static str,
|
||||
optional_params: &Map<String, Value>,
|
||||
credentials: AwsCredentialSource,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Self, Error> {
|
||||
let credentials = match host_supplied_credentials(optional_params) {
|
||||
Some(credentials) => credentials,
|
||||
None => {
|
||||
resolve_credentials(aws_auth_config(optional_params, env_lookup), env_lookup)
|
||||
.await?
|
||||
}
|
||||
};
|
||||
Ok(Self::new(region, service, credentials))
|
||||
Ok(Self::new(
|
||||
region,
|
||||
service,
|
||||
credentials.resolve(auth, env_lookup).await?,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -80,7 +76,7 @@ mod tests {
|
|||
use std::time::{Duration, UNIX_EPOCH};
|
||||
|
||||
use litellm_http::outbound::OutboundRequest;
|
||||
use serde_json::json;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
|
|
|
|||
|
|
@ -131,7 +131,7 @@ impl Default for VertexAuth {
|
|||
}
|
||||
|
||||
impl VertexAuth {
|
||||
fn new(loader: Arc<dyn VertexProviderLoader>) -> Self {
|
||||
pub fn new(loader: Arc<dyn VertexProviderLoader>) -> Self {
|
||||
Self {
|
||||
providers: Cache::builder().max_capacity(64).build(),
|
||||
loader,
|
||||
|
|
@ -220,16 +220,16 @@ impl VertexAuth {
|
|||
}
|
||||
}
|
||||
|
||||
trait VertexTokenSource: Send + Sync {
|
||||
pub trait VertexTokenSource: Send + Sync {
|
||||
fn project_id(&self) -> VertexAuthFuture<'_, String>;
|
||||
fn token(&self) -> VertexAuthFuture<'_, String>;
|
||||
}
|
||||
|
||||
trait VertexProviderLoader: Send + Sync {
|
||||
pub trait VertexProviderLoader: Send + Sync {
|
||||
fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc<dyn VertexTokenSource>>;
|
||||
}
|
||||
|
||||
type VertexAuthFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
|
||||
pub type VertexAuthFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
|
||||
|
||||
struct GcpTokenSource(Arc<dyn TokenProvider>);
|
||||
|
||||
|
|
@ -305,7 +305,7 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> {
|
|||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
enum CredentialSource {
|
||||
pub enum CredentialSource {
|
||||
Inline(SecretValue),
|
||||
Trusted(SecretValue),
|
||||
ApplicationCredentials(String),
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ pub enum CredentialPlacement {
|
|||
}
|
||||
|
||||
impl CredentialPlacement {
|
||||
pub fn header_name(self) -> &'static str {
|
||||
pub const fn header_name(self) -> &'static str {
|
||||
match self {
|
||||
Self::Bearer => "Authorization",
|
||||
Self::Header(name) => name,
|
||||
|
|
@ -40,21 +40,6 @@ pub fn apply_credential(
|
|||
)
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum RequestAuth {
|
||||
Header {
|
||||
name: &'static str,
|
||||
value: String,
|
||||
},
|
||||
Bearer {
|
||||
token: String,
|
||||
},
|
||||
AwsSigV4 {
|
||||
region: String,
|
||||
service: &'static str,
|
||||
},
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{CredentialPlacement, apply_credential};
|
||||
|
|
|
|||
|
|
@ -51,7 +51,7 @@ pub use credential::{
|
|||
CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle,
|
||||
};
|
||||
pub use error::Error;
|
||||
pub use http::{CredentialPlacement, RequestAuth};
|
||||
pub use http::CredentialPlacement;
|
||||
pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy};
|
||||
pub use secret::SecretValue;
|
||||
pub use token::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};
|
||||
|
|
|
|||
|
|
@ -2,6 +2,9 @@
|
|||
|
||||
pub use litellm_auth_types::*;
|
||||
|
||||
mod services;
|
||||
pub use services::AuthServices;
|
||||
|
||||
#[cfg(feature = "aws")]
|
||||
pub use litellm_auth_aws as aws;
|
||||
#[cfg(feature = "azure")]
|
||||
|
|
|
|||
9
litellm-rust/crates/auth/src/services.rs
Normal file
9
litellm-rust/crates/auth/src/services.rs
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
#[derive(Default)]
|
||||
pub struct AuthServices {
|
||||
#[cfg(feature = "aws")]
|
||||
pub aws: litellm_auth_aws::AwsAuthService,
|
||||
#[cfg(feature = "azure")]
|
||||
pub azure: litellm_auth_azure::AzureAuthService,
|
||||
#[cfg(feature = "gcp")]
|
||||
pub gcp: litellm_auth_gcp::VertexAuth,
|
||||
}
|
||||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
litellm-auth-azure.workspace = true
|
||||
litellm-auth-types.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
|
|
@ -19,6 +20,7 @@ tokio.workspace = true
|
|||
url.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-response.workspace = true
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ impl<C: CacheCodec> AzureBlobCache<C> {
|
|||
pub async fn connect(
|
||||
account_url: &str,
|
||||
container: &str,
|
||||
http: reqwest::Client,
|
||||
http: litellm_http::Client,
|
||||
codec: C,
|
||||
runtime: Handle,
|
||||
) -> Result<Self, Error> {
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ use azure_core::{
|
|||
use futures_util::TryStreamExt;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct ReqwestTransport(pub reqwest::Client);
|
||||
pub struct ReqwestTransport(pub litellm_http::Client);
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl HttpClient for ReqwestTransport {
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ async fn connect(server: &MockServer) -> AzureBlobCache<JsonCodec<Value>> {
|
|||
None,
|
||||
ClientOptions {
|
||||
transport: Some(Transport::new(Arc::new(ReqwestTransport(
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
)))),
|
||||
..ClientOptions::default()
|
||||
},
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
futures-util.workspace = true
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-auth-types.workspace = true
|
||||
|
|
@ -15,6 +16,7 @@ reqwest.workspace = true
|
|||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ use litellm_cache::{
|
|||
BaseCache, BatchCache, BatchEntry, CacheCodec, DisconnectCache, Error, ExactCacheContext,
|
||||
FlushCache,
|
||||
};
|
||||
use litellm_http::Client;
|
||||
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, percent_encode};
|
||||
use reqwest::Client;
|
||||
|
||||
use crate::{GcpTokenSource, TokenSource};
|
||||
|
||||
|
|
|
|||
|
|
@ -96,7 +96,7 @@ async fn cache_exposes_its_configuration(#[future(awt)] server: MockServer) {
|
|||
path_service_account: Some("/secrets/sa.json".into()),
|
||||
..support::config(&server, Some("folder"))
|
||||
},
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
litellm_cache::JsonCodec::<Value>::new(),
|
||||
);
|
||||
assert_eq!(cache.bucket_name(), "bucket");
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ pub fn cache_with_token(
|
|||
) -> JsonGcsCache {
|
||||
GcsCache::with_token_source(
|
||||
config(server, gcs_path),
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
JsonCodec::new(),
|
||||
token,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
futures-util.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
qdrant-client = { workspace = true, features = ["serde"] }
|
||||
|
|
@ -17,6 +18,7 @@ tokio.workspace = true
|
|||
uuid.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
futures-executor = "0.3"
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_cache::{Error, semantic::Embedder};
|
||||
use reqwest::Client;
|
||||
use litellm_http::Client;
|
||||
use serde_json::Value;
|
||||
|
||||
pub struct OpenAiEmbedder {
|
||||
|
|
|
|||
|
|
@ -5,6 +5,10 @@ use std::{
|
|||
|
||||
use litellm_cache::{Error, semantic::Embedder};
|
||||
use litellm_cache_qdrant_semantic::{OpenAiEmbedder, OpenAiEmbedderConfig};
|
||||
use litellm_http::{
|
||||
ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, Resolution,
|
||||
media::PublicDnsResolver,
|
||||
};
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
use tokio::{
|
||||
|
|
@ -104,7 +108,7 @@ fn config(base: String, timeout: Option<Duration>) -> OpenAiEmbedderConfig {
|
|||
async fn posts_embeddings_request_and_parses_vector() {
|
||||
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
|
||||
let embedder = OpenAiEmbedder::new(
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
config(
|
||||
format!("{}/", server.base_url()),
|
||||
Some(Duration::from_secs(1)),
|
||||
|
|
@ -156,14 +160,17 @@ async fn status_timeout_and_body_errors_are_unavailable(
|
|||
) {
|
||||
let server =
|
||||
TestHttpServer::response_after(status, body, Duration::from_millis(delay_ms)).await;
|
||||
let embedder = OpenAiEmbedder::new(reqwest::Client::new(), config(server.base_url(), timeout));
|
||||
let embedder = OpenAiEmbedder::new(
|
||||
litellm_http::Client::plain_for_test(),
|
||||
config(server.base_url(), timeout),
|
||||
);
|
||||
assert_eq!(embedder.async_embed("hello", None).await, expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn sync_embedding_is_unsupported() {
|
||||
let embedder = OpenAiEmbedder::new(
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
config("http://127.0.0.1:9".to_owned(), None),
|
||||
);
|
||||
assert_eq!(
|
||||
|
|
@ -176,9 +183,12 @@ fn sync_embedding_is_unsupported() {
|
|||
#[tokio::test]
|
||||
async fn uses_the_injected_client() {
|
||||
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
|
||||
let client = reqwest::Client::builder()
|
||||
.user_agent("litellm-embedder-test")
|
||||
.build()
|
||||
let config_with_agent = HttpClientConfig {
|
||||
user_agent: Some("litellm-embedder-test".into()),
|
||||
..Resolution::from(&HttpSettings::default()).config
|
||||
};
|
||||
let client = HttpClientPool::new(Arc::new(PublicDnsResolver))
|
||||
.client(&config_with_agent, ClientVariant::Provider)
|
||||
.unwrap();
|
||||
let embedder = OpenAiEmbedder::new(client, config(server.base_url(), None));
|
||||
assert_eq!(
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
litellm-auth-aws.workspace = true
|
||||
aws-sdk-s3 = { version = "1.146.1", default-features = false, features = ["rustls", "rt-tokio"] }
|
||||
|
|
@ -19,6 +20,7 @@ reqwest.workspace = true
|
|||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
|
|
|
|||
|
|
@ -2,10 +2,11 @@ use aws_credential_types::{
|
|||
Credentials as AwsCredentials,
|
||||
provider::{ProvideCredentials, error::CredentialsError, future},
|
||||
};
|
||||
use litellm_auth_aws::{AwsAuthConfig, resolve_credentials};
|
||||
use litellm_auth_aws::{AwsAuthConfig, AwsAuthService};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct S3Credentials {
|
||||
auth: AwsAuthService,
|
||||
config: AwsAuthConfig,
|
||||
env: fn(&str) -> Option<String>,
|
||||
}
|
||||
|
|
@ -16,7 +17,11 @@ impl S3Credentials {
|
|||
}
|
||||
|
||||
pub fn with_env(config: AwsAuthConfig, env: fn(&str) -> Option<String>) -> Self {
|
||||
Self { config, env }
|
||||
Self {
|
||||
auth: AwsAuthService::default(),
|
||||
config,
|
||||
env,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -38,7 +43,8 @@ impl ProvideCredentials for S3Credentials {
|
|||
"litellm-s3-cache",
|
||||
));
|
||||
}
|
||||
resolve_credentials(self.config.clone(), &self.env)
|
||||
self.auth
|
||||
.resolve_credentials(self.config.clone(), &self.env)
|
||||
.await
|
||||
.map_err(|_| CredentialsError::provider_error("S3 cache authentication failed"))
|
||||
})
|
||||
|
|
|
|||
|
|
@ -42,7 +42,12 @@ pub struct S3Cache<C: CacheCodec> {
|
|||
}
|
||||
|
||||
impl<C: CacheCodec> S3Cache<C> {
|
||||
pub fn new(config: S3CacheConfig, http: reqwest::Client, codec: C, runtime: Handle) -> Self {
|
||||
pub fn new(
|
||||
config: S3CacheConfig,
|
||||
http: litellm_http::Client,
|
||||
codec: C,
|
||||
runtime: Handle,
|
||||
) -> Self {
|
||||
let endpoint_url: Option<String> = config.endpoint.map(|endpoint| endpoint.url);
|
||||
let base = aws_sdk_s3::Config::builder()
|
||||
.behavior_version(BehaviorVersion::latest())
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ use aws_smithy_runtime_api::client::{
|
|||
use aws_smithy_types::body::SdkBody;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct ReqwestHttpClient(pub(crate) reqwest::Client);
|
||||
pub(crate) struct ReqwestHttpClient(pub(crate) litellm_http::Client);
|
||||
|
||||
impl HttpClient for ReqwestHttpClient {
|
||||
fn http_connector(
|
||||
|
|
|
|||
|
|
@ -32,7 +32,12 @@ pub fn config(endpoint: &str) -> S3CacheConfig {
|
|||
}
|
||||
|
||||
pub fn cache_with(config: S3CacheConfig, runtime: Handle) -> JsonS3Cache {
|
||||
S3Cache::new(config, reqwest::Client::new(), JsonCodec::new(), runtime)
|
||||
S3Cache::new(
|
||||
config,
|
||||
litellm_http::Client::plain_for_test(),
|
||||
JsonCodec::new(),
|
||||
runtime,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn cache(endpoint: &str) -> JsonS3Cache {
|
||||
|
|
|
|||
|
|
@ -1,19 +1,19 @@
|
|||
- Target invariants, not completion claims
|
||||
- Keep this crate the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on) plus the kwargs rewrites the wrapper makes on the way in (credential-name inheritance, the budget and retry-count limits)
|
||||
- This crate is the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on)
|
||||
- Smell test: if a future callback host (`callbacks-v1-python`, WASM, in-process Rust) could share a piece of this crate, it does not belong here
|
||||
- SDK request policy (credential inheritance, the budget and retry-count limits) is the driver's preflight, supplied by `python-bridge`; this crate only adopts the keyword view it produces
|
||||
- The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call
|
||||
- Rust drives the call; every litellm Python internal it still borrows is a variant of `LegacyPython`, grouped by subsystem (`Wrapper`, `Logging`, `DeploymentHooks`)
|
||||
- Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json`
|
||||
- The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it
|
||||
- Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython`
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the legacy path rewrites it (setup, deployment hook, prepare) and the bound request object whose attributes back keywords the caller omitted; routes hand it over through `run_legacy_call` and keep no copy
|
||||
- `setup` reuses a `Logging` the caller passed as `litellm_logging_obj` (the proxy and Router are the live cases) and otherwise builds one through `function_setup`, as `@client` does
|
||||
- Either way every phase calls the same `Logging` method the Python path calls; which callbacks run is `Logging`'s decision, never this crate's
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; routes hand it over through `run_legacy_call` and keep no copy
|
||||
- `setup` reuses a `Logging` passed as `litellm_logging_obj` (the proxy and Router) and otherwise builds one through `function_setup`; which callbacks run is `Logging`'s decision, never this crate's
|
||||
- Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation
|
||||
- Retain complete boundary arguments, opaque unknown values, aliases, omitted/default distinctions and deliberate copies; preserve the established deployment-hook kwargs view
|
||||
- Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object; this crate compares the two itself, and the argument is resolved by `litellm_host_python::lookup`
|
||||
- Retain independently captured body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only
|
||||
- A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-host`, `litellm-host-python` and the bridge; the only fact that crosses from the route is the prepared keyword view
|
||||
- Success and failure handlers receive the exact selected public response or exception; logging projections, redaction and snapshots keep their own copy contracts
|
||||
- Ordinary failure-handler errors cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch
|
||||
- Dispatch errors never replay provider work or trigger the opposite outcome; the proxy's acceptance or rejection releases deferred success at most once
|
||||
- Retain complete boundary arguments, opaque values, aliases, omitted/default distinctions and deliberate copies; preserve the deployment-hook kwargs view
|
||||
- Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object, resolved through `litellm_host_python::lookup`
|
||||
- Retain body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only
|
||||
- Success and failure handlers receive the exact selected public response or exception
|
||||
- A failure-handler error cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch
|
||||
- Dispatch errors never replay provider work or trigger the opposite outcome; the proxy releases deferred success at most once
|
||||
- Delivery follows the registry, not the callable's type: direct, awaited, executor-submitted, logging-worker and deferred paths stay distinct
|
||||
- Traverse every retained Python edge; `close` is idempotent and restores the correlation context once
|
||||
|
|
|
|||
|
|
@ -6,9 +6,6 @@
|
|||
"start_time",
|
||||
"asynchronous"
|
||||
],
|
||||
"check_limits": [
|
||||
"kwargs"
|
||||
],
|
||||
"finalize": [
|
||||
"response",
|
||||
"logger",
|
||||
|
|
@ -76,11 +73,6 @@
|
|||
],
|
||||
"custom_pricing_fields": [],
|
||||
"is_internal_call": [],
|
||||
"credential_list": [],
|
||||
"warn_unknown_credential": [
|
||||
"name",
|
||||
"loaded"
|
||||
],
|
||||
"before_deployment_call": [
|
||||
"kwargs",
|
||||
"call_type"
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ use serde_json::Value;
|
|||
use crate::{
|
||||
DeploymentHooks, LegacyCallbacks, PublicCall, PythonLogger,
|
||||
deferred::{PendingLogging, PendingSuccess},
|
||||
finalize, is_internal_call, prepare,
|
||||
finalize, is_internal_call,
|
||||
python::Streaming,
|
||||
setup,
|
||||
};
|
||||
|
|
@ -117,9 +117,13 @@ impl LegacyLogging {
|
|||
})
|
||||
}
|
||||
|
||||
/// The keyword view the rest of the call reads: a copy, so the deployment hook's own
|
||||
/// dict is left as the hook returned it, carrying the logger as `@client` injects it.
|
||||
/// The driver's preflight rewrites this same dict before the host projects from it.
|
||||
fn prepare(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
|
||||
let prepared = prepare(py, self.call.kwargs().bind(py), self.logger()?)?.unbind();
|
||||
self.call.set_kwargs(prepared);
|
||||
let prepared = self.call.kwargs().bind(py).copy()?;
|
||||
prepared.set_item("litellm_logging_obj", self.logger()?.object(py))?;
|
||||
self.call.set_kwargs(prepared.unbind());
|
||||
Ok(LifecycleStep::Arguments(self.call.kwargs().clone_ref(py)))
|
||||
}
|
||||
|
||||
|
|
@ -465,6 +469,7 @@ impl PythonLifecycle for LegacyLogging {
|
|||
error.write_unraisable(py, None);
|
||||
}
|
||||
self.body = None;
|
||||
self.headers = None;
|
||||
self.context = None;
|
||||
self.stream = None;
|
||||
}
|
||||
|
|
@ -482,7 +487,8 @@ impl PythonLifecycle for LegacyLogging {
|
|||
visit.call(&stream.chunks)?;
|
||||
visit.call(&stream.first_chunk)?;
|
||||
}
|
||||
visit.call(&self.body)
|
||||
visit.call(&self.body)?;
|
||||
visit.call(&self.headers)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -580,8 +586,6 @@ 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
|
||||
",
|
||||
);
|
||||
});
|
||||
|
|
@ -616,8 +620,6 @@ kwargs = {'logger': logger, 'vendor_extension': opaque}
|
|||
&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
|
||||
",
|
||||
);
|
||||
|
|
@ -733,45 +735,6 @@ assert all(value is failure for name, value in logger.calls if name.endswith('_h
|
|||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[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")));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -782,6 +745,7 @@ mod payload_tests {
|
|||
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py};
|
||||
use proptest::prelude::*;
|
||||
use pyo3::gc::{PyTraverseError, PyVisit};
|
||||
use pyo3::prelude::*;
|
||||
use rstest::rstest;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
|
@ -871,16 +835,7 @@ check = lambda: None
|
|||
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
|
||||
));
|
||||
let (_, step) = send_and_receive(py, &mut logging, wire, &context);
|
||||
run(py, &locals, c"check()");
|
||||
let LifecycleStep::Wire(wire) = step else {
|
||||
panic!("before_send did not hand back the wire request");
|
||||
|
|
@ -889,6 +844,134 @@ check = lambda: None
|
|||
})
|
||||
}
|
||||
|
||||
/// `before_send` over `wire`, then the provider's raw response the way the driver
|
||||
/// delivers it, so `pre_call` and `post_call` have both seen the retained payload.
|
||||
fn send_and_receive<'a>(
|
||||
py: Python<'_>,
|
||||
logging: &'a mut LegacyLogging,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
) -> (&'a mut LegacyLogging, LifecycleStep) {
|
||||
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
|
||||
));
|
||||
(logging, step)
|
||||
}
|
||||
|
||||
fn route_context() -> RequestContext {
|
||||
RequestContext {
|
||||
model: "model".into(),
|
||||
custom_llm_provider: "provider".into(),
|
||||
optional_params: json!({}),
|
||||
secret_fields: vec![],
|
||||
api_key: Some(SecretValue::new("route-key")),
|
||||
}
|
||||
}
|
||||
|
||||
fn route_wire() -> WireRequest {
|
||||
WireRequest {
|
||||
url: "https://provider.invalid/ocr".into(),
|
||||
headers: vec![("x-route".into(), "route".into())],
|
||||
body: json!({}),
|
||||
}
|
||||
}
|
||||
|
||||
/// A Python object owning one `LegacyLogging`, so the interpreter's collector sees the
|
||||
/// edges the adapter reports and clears them the way the driver's `Execution` does.
|
||||
#[pyclass(weakref)]
|
||||
struct Retained {
|
||||
logging: Option<LegacyLogging>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl Retained {
|
||||
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
match &self.logging {
|
||||
Some(logging) => logging.traverse(&visit),
|
||||
None => Ok(()),
|
||||
}
|
||||
}
|
||||
|
||||
fn __clear__(slf: &Bound<'_, Self>) {
|
||||
drop(slf.borrow_mut().logging.take());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_cycle_through_the_retained_headers_is_collected() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = namespace(py, PAYLOAD_LOGGER);
|
||||
let mut logging = LegacyLogging {
|
||||
logger: Some(PythonLogger::new(local(&locals, "logger").unbind())),
|
||||
..legacy_call(py, &locals, false)
|
||||
};
|
||||
send_and_receive(py, &mut logging, route_wire(), &route_context());
|
||||
let retained = Py::new(
|
||||
py,
|
||||
Retained {
|
||||
logging: Some(logging),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
locals.set_item("retained", retained).unwrap();
|
||||
run(
|
||||
py,
|
||||
&locals,
|
||||
c"
|
||||
import gc
|
||||
import weakref
|
||||
|
||||
logger.post[2]['headers']['owner'] = retained
|
||||
logger.pre = logger.post = None
|
||||
reference = weakref.ref(retained)
|
||||
del retained
|
||||
gc.collect()
|
||||
assert reference() is None
|
||||
",
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn close_releases_the_retained_headers() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = namespace(py, PAYLOAD_LOGGER);
|
||||
let mut logging = LegacyLogging {
|
||||
logger: Some(PythonLogger::new(local(&locals, "logger").unbind())),
|
||||
..legacy_call(py, &locals, false)
|
||||
};
|
||||
send_and_receive(py, &mut logging, route_wire(), &route_context());
|
||||
run(
|
||||
py,
|
||||
&locals,
|
||||
c"
|
||||
import weakref
|
||||
|
||||
class Sentinel:
|
||||
pass
|
||||
|
||||
sentinel = Sentinel()
|
||||
logger.post[2]['headers']['sentinel'] = sentinel
|
||||
logger.pre = logger.post = None
|
||||
reference = weakref.ref(sentinel)
|
||||
del sentinel
|
||||
assert reference() is not None
|
||||
",
|
||||
);
|
||||
logging.close(py);
|
||||
run(py, &locals, c"assert reference() is None");
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::caller_keyword(c"
|
||||
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@
|
|||
//! this crate holds them.
|
||||
|
||||
use litellm_host::{machine::Machine, protocol::Protocol};
|
||||
use litellm_host_python::{ProtocolHost, lookup, run_call};
|
||||
use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call};
|
||||
use pyo3::{
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
|
|
@ -39,7 +39,8 @@ impl PublicCall {
|
|||
}
|
||||
|
||||
/// The keyword view the legacy path currently reads: the caller's copy until
|
||||
/// `function_setup`, then each rewrite (setup, deployment hook, prepare) in turn.
|
||||
/// `function_setup`, then each rewrite (setup, deployment hook, the driver's preflight)
|
||||
/// in turn.
|
||||
pub(crate) fn kwargs(&self) -> &Py<PyDict> {
|
||||
&self.kwargs
|
||||
}
|
||||
|
|
@ -64,13 +65,15 @@ impl PublicCall {
|
|||
}
|
||||
|
||||
/// Runs one native call under the legacy `Logging` contract: the protocol host projects from
|
||||
/// the keyword view the contract prepares, and the contract observes the call.
|
||||
/// the keyword view the contract prepares and `preflight` rewrites, and the contract
|
||||
/// observes the call.
|
||||
pub fn run_legacy_call<H, M>(
|
||||
py: Python<'_>,
|
||||
surface: LegacySurface,
|
||||
call: PublicCall,
|
||||
machine: M,
|
||||
host: H,
|
||||
preflight: Preflight,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
|
|
@ -83,6 +86,7 @@ where
|
|||
machine,
|
||||
host,
|
||||
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
|
||||
preflight,
|
||||
arguments,
|
||||
asynchronous,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
//! The legacy `@client` wrapper as the native call sees it: litellm's `Logging` object, the
|
||||
//! sync and async callback registries it fans out to, the deployment hooks, the deferred
|
||||
//! proxy release, and the kwargs rewrites the wrapper makes on the way in (credential-name
|
||||
//! inheritance, budget and retry-count limits). All of it sits behind one
|
||||
//! sync and async callback registries it fans out to, the deployment hooks and the deferred
|
||||
//! proxy release. All of it sits behind one
|
||||
//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and
|
||||
//! core never learn which Python object is on the other end.
|
||||
//! core never learn which Python object is on the other end. The SDK's own request policy
|
||||
//! (credential inheritance, the budget and retry limits) is the driver's preflight, not this
|
||||
//! crate's.
|
||||
//!
|
||||
//! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`]
|
||||
//! is where those objects live, and [`run_legacy_call`] is how a route hands them over
|
||||
|
|
@ -14,220 +15,12 @@ mod call;
|
|||
mod callbacks;
|
||||
mod deferred;
|
||||
mod logger;
|
||||
mod preparation;
|
||||
mod python;
|
||||
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,
|
||||
)
|
||||
}
|
||||
}
|
||||
mod test_support;
|
||||
|
|
|
|||
|
|
@ -19,18 +19,12 @@ pub(crate) enum LegacyPython {
|
|||
Streaming(Streaming),
|
||||
}
|
||||
|
||||
/// The `@client` wrapper around the call: `function_setup`, limits, credentials,
|
||||
/// response metadata and the correlation context.
|
||||
/// The `@client` wrapper around the call: `function_setup`, response metadata and the
|
||||
/// correlation context.
|
||||
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)]
|
||||
pub(crate) enum Wrapper {
|
||||
#[strum(serialize = "setup")]
|
||||
Setup,
|
||||
#[strum(serialize = "check_limits")]
|
||||
CheckLimits,
|
||||
#[strum(serialize = "credential_list")]
|
||||
CredentialList,
|
||||
#[strum(serialize = "warn_unknown_credential")]
|
||||
WarnUnknownCredential,
|
||||
#[strum(serialize = "is_internal_call")]
|
||||
IsInternalCall,
|
||||
#[strum(serialize = "finalize")]
|
||||
|
|
|
|||
199
litellm-rust/crates/callbacks-legacy-python/src/test_support.rs
Normal file
199
litellm-rust/crates/callbacks-legacy-python/src/test_support.rs
Normal file
|
|
@ -0,0 +1,199 @@
|
|||
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/unit/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,
|
||||
),
|
||||
'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(),
|
||||
'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 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,
|
||||
)
|
||||
}
|
||||
16
litellm-rust/crates/config/Cargo.toml
Normal file
16
litellm-rust/crates/config/Cargo.toml
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
[package]
|
||||
name = "litellm-config"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-auth-types.workspace = true
|
||||
serde.workspace = true
|
||||
serde_yaml_ng = "0.10.0"
|
||||
thiserror.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
tempfile.workspace = true
|
||||
7
litellm-rust/crates/config/src/error.rs
Normal file
7
litellm-rust/crates/config/src/error.rs
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("could not read config")]
|
||||
Read(#[from] std::io::Error),
|
||||
#[error("invalid YAML config")]
|
||||
Parse(#[from] serde_yaml_ng::Error),
|
||||
}
|
||||
48
litellm-rust/crates/config/src/lib.rs
Normal file
48
litellm-rust/crates/config/src/lib.rs
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
mod error;
|
||||
|
||||
use std::path::Path;
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
pub use error::Error;
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct Config {
|
||||
pub model_list: Box<[Model]>,
|
||||
#[serde(default)]
|
||||
pub general_settings: GeneralSettings,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct GeneralSettings {
|
||||
pub master_key: Option<SecretValue>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct Model {
|
||||
pub model_name: String,
|
||||
pub litellm_params: LiteLlmParams,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct LiteLlmParams {
|
||||
pub model: String,
|
||||
pub api_key: Option<SecretValue>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
}
|
||||
|
||||
impl Config {
|
||||
pub fn from_yaml(yaml: &str) -> Result<Self, Error> {
|
||||
Ok(serde_yaml_ng::from_str(yaml)?)
|
||||
}
|
||||
|
||||
pub fn load(path: impl AsRef<Path>) -> Result<Self, Error> {
|
||||
Self::from_yaml(&std::fs::read_to_string(path)?)
|
||||
}
|
||||
}
|
||||
119
litellm-rust/crates/config/tests/config.rs
Normal file
119
litellm-rust/crates/config/tests/config.rs
Normal file
|
|
@ -0,0 +1,119 @@
|
|||
use litellm_config::{Config, Error};
|
||||
use rstest::{fixture, rstest};
|
||||
use tempfile::TempDir;
|
||||
|
||||
#[fixture]
|
||||
fn directory() -> TempDir {
|
||||
tempfile::tempdir().unwrap()
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn model_list_yaml() -> &'static str {
|
||||
r#"
|
||||
model_list:
|
||||
- model_name: assistant
|
||||
litellm_params:
|
||||
model: anthropic/test-model
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
- model_name: local
|
||||
litellm_params:
|
||||
model: test-model
|
||||
api_base: http://localhost:8000/v1
|
||||
custom_llm_provider: openai
|
||||
"#
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn loads_model_list_from_file(directory: TempDir, model_list_yaml: &str) {
|
||||
let path = directory.path().join("config.yaml");
|
||||
std::fs::write(&path, model_list_yaml).unwrap();
|
||||
|
||||
let config = Config::load(path).unwrap();
|
||||
assert_eq!(config.model_list.len(), 2);
|
||||
let anthropic = &config.model_list[0];
|
||||
assert_eq!(anthropic.model_name, "assistant");
|
||||
assert_eq!(anthropic.litellm_params.model, "anthropic/test-model");
|
||||
assert_eq!(
|
||||
anthropic.litellm_params.api_key.as_ref().unwrap().expose(),
|
||||
"os.environ/ANTHROPIC_API_KEY"
|
||||
);
|
||||
assert!(anthropic.litellm_params.api_base.is_none());
|
||||
assert!(anthropic.litellm_params.custom_llm_provider.is_none());
|
||||
let local = &config.model_list[1];
|
||||
assert_eq!(local.model_name, "local");
|
||||
assert_eq!(local.litellm_params.model, "test-model");
|
||||
assert!(local.litellm_params.api_key.is_none());
|
||||
assert_eq!(
|
||||
local.litellm_params.api_base.as_deref(),
|
||||
Some("http://localhost:8000/v1")
|
||||
);
|
||||
assert_eq!(
|
||||
local.litellm_params.custom_llm_provider.as_deref(),
|
||||
Some("openai")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn config_debug_redacts_api_keys() {
|
||||
let config = Config::from_yaml(
|
||||
"model_list: [{model_name: assistant, litellm_params: {model: anthropic/test-model, api_key: secret-value}}]",
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
config.model_list[0]
|
||||
.litellm_params
|
||||
.api_key
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"secret-value"
|
||||
);
|
||||
assert!(!format!("{config:?}").contains("secret-value"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::malformed_yaml("model_list: [")]
|
||||
#[case::missing_model_list("{}")]
|
||||
#[case::missing_params("model_list: [{model_name: assistant}]")]
|
||||
#[case::missing_model("model_list: [{model_name: assistant, litellm_params: {api_key: key}}]")]
|
||||
#[case::unsupported_settings("model_list: []\ngeneral_settings: {unknown: true}")]
|
||||
#[case::misspelled_param(
|
||||
"model_list: [{model_name: assistant, litellm_params: {model: test, api_bsae: url}}]"
|
||||
)]
|
||||
fn rejects_malformed_incomplete_and_unsupported_config(#[case] yaml: &str) {
|
||||
assert!(matches!(Config::from_yaml(yaml), Err(Error::Parse(_))));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn distinguishes_read_errors_from_parse_errors(directory: TempDir) {
|
||||
assert!(matches!(
|
||||
Config::load(directory.path().join("missing.yaml")),
|
||||
Err(Error::Read(error)) if error.kind() == std::io::ErrorKind::NotFound
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::literal("secret-master-key")]
|
||||
#[case::reference("os.environ/LITELLM_MASTER_KEY")]
|
||||
fn loads_and_redacts_the_master_key(#[case] key: &str) {
|
||||
let config = Config::from_yaml(&format!(
|
||||
"model_list: []\ngeneral_settings:\n master_key: {key}\n"
|
||||
))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
config
|
||||
.general_settings
|
||||
.master_key
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
key
|
||||
);
|
||||
assert!(!format!("{config:?}").contains(key));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn missing_general_settings_has_no_master_key() {
|
||||
let config = Config::from_yaml("model_list: []").unwrap();
|
||||
assert!(config.general_settings.master_key.is_none());
|
||||
}
|
||||
|
|
@ -13,6 +13,7 @@ serde.workspace = true
|
|||
serde_json.workspace = true
|
||||
serde_path_to_error = "0.1"
|
||||
serde_with.workspace = true
|
||||
strum.workspace = true
|
||||
thiserror.workspace = true
|
||||
url.workspace = true
|
||||
|
||||
|
|
|
|||
|
|
@ -11,11 +11,13 @@
|
|||
//! accepts; anything richer is declined upstream by the capability gate.
|
||||
|
||||
use litellm_types::llms::openai::{ChatMessage, ChatMessageContent};
|
||||
use strum::IntoStaticStr;
|
||||
|
||||
pub const EMPTY_TEXT_PLACEHOLDER: &str =
|
||||
"[System: Empty message content sanitised to satisfy protocol]";
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub enum TurnRole {
|
||||
User,
|
||||
Assistant,
|
||||
|
|
@ -23,10 +25,7 @@ pub enum TurnRole {
|
|||
|
||||
impl TurnRole {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::User => "user",
|
||||
Self::Assistant => "assistant",
|
||||
}
|
||||
self.into()
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back.
|
||||
litellm-core is the LiteLLM SDK in Rust. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route. `messages::messages()` returns `MessagesResponse::Message` for a completed response or `MessagesResponse::Stream { headers, chunks }` when the request sets `stream: true`. The chunks are Anthropic SSE bytes in a `Stream<Item = Result<Bytes, Error>>`. Dropping the stream cancels the call. The Python bridge drives `messages::route::messages_machine()` instead, because Python has to answer the call's operations on its own thread; the gateway and the Rust SDK call the plain entrypoint
|
||||
|
||||
A route module has the same five pieces, in the order Python runs them. `types.rs` holds the call, the provider request, and the response. `prepare.rs` resolves the provider and credentials and shapes the request (Python's `validate_environment`, `get_complete_url`, `transform_request`). `handler.rs` resolves auth, offers the wire request to `litellm_host::hooks::RouteHooks::before_send`, sends it, reports the raw response through `emit`, and normalizes the response or stream (`pre_call`, `post`, `post_call`, `transform_response`). `mod.rs` exposes the entrypoint that runs prepare then handler with no hooks (`()`). `route.rs`, where a host needs it, wraps the same two calls in a `CallMachine` whose `HostChannel` is the hooks, and pumps a stream through `open` and `deliver`. A handler takes `&impl RouteHooks<Error>` and never a `HostChannel` directly, so it runs without a coroutine. Keep provider transport and transformation details out of the machine driver
|
||||
|
||||
## Crate layering
|
||||
|
||||
|
|
@ -10,6 +12,16 @@ Each crate mirrors one top-level Python package, so a Rust path reads as its Pyt
|
|||
- `litellm-llms` mirrors `litellm/llms/`: `base_llm/<api>/transformation.rs`, `<provider>/<api>/transformation.rs`, and `base_llm/ocr/handler.rs` (the OCR request handler)
|
||||
- `litellm-core` mirrors the route packages (`litellm/ocr/`, `litellm/messages/`, ...): entrypoints, route request types, provider dispatch, the route machine, and hooks
|
||||
|
||||
A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, which each route implements over its host. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
|
||||
A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, the provider-level hooks OCR implements over its host until it folds into `litellm_host::hooks::RouteHooks`. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
|
||||
|
||||
## Error placement
|
||||
|
||||
The workspace `Error definitions` rules shape each crate's error; this section decides which crate and module a failure belongs to
|
||||
|
||||
A failure is declared once, by the lowest crate that raises it. Every crate above nests that error unchanged (`#[error(transparent)] Auth(#[from] litellm_auth::Error)`) or maps it once at its boundary, as `src/error.rs` does for `litellm_llms::Error`. `RouteError` collects route failures and never re-declares a variant a lower crate raises
|
||||
|
||||
Scope follows the concept, not the first caller. An error type under `litellm-llms`'s `<provider>/` is private to that provider: no other provider and nothing in `base_llm` may import it. A failure two providers or two routes can hit, such as wire framing, stream event decoding, or a malformed provider response, belongs to the crate that owns the concept: `litellm-framing` for framing, `litellm_llms::Error` for the transformation layer
|
||||
|
||||
`litellm_llms::Error` (`crates/llms/src/error.rs`) is the one transformation error for every provider and API. `base_llm/ocr/error.rs` is the recorded exception until OCR folds into it
|
||||
|
||||
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback execution of any kind. Core runs each route as a machine that yields host operations and call events; which integrations consume those events is the host's business.
|
||||
|
|
|
|||
|
|
@ -13,10 +13,11 @@ litellm-host.workspace = true
|
|||
bytes.workspace = true
|
||||
futures-util.workspace = true
|
||||
base64.workspace = true
|
||||
litellm-auth.workspace = true
|
||||
litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] }
|
||||
litellm-auth-aws.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
moka.workspace = true
|
||||
mime_guess = "2.0.5"
|
||||
rand.workspace = true
|
||||
|
|
@ -36,7 +37,9 @@ url.workspace = true
|
|||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-llms = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
rstest_reuse.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
|
|
|
|||
|
|
@ -1,13 +0,0 @@
|
|||
use std::{sync::OnceLock, time::Duration};
|
||||
|
||||
use crate::constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS;
|
||||
|
||||
pub(super) fn http_client() -> &'static reqwest::Client {
|
||||
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
|
||||
CLIENT.get_or_init(|| {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS))
|
||||
.build()
|
||||
.unwrap_or_else(|_| reqwest::Client::new())
|
||||
})
|
||||
}
|
||||
|
|
@ -1,43 +0,0 @@
|
|||
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("expected {expected}, got {actual}")]
|
||||
InvalidType {
|
||||
expected: &'static str,
|
||||
actual: &'static str,
|
||||
},
|
||||
#[error("missing required field: {0}")]
|
||||
MissingField(&'static str),
|
||||
#[error("invalid provider: {0}")]
|
||||
InvalidProvider(String),
|
||||
#[error("invalid request: {0}")]
|
||||
InvalidRequest(String),
|
||||
#[error("invalid response: {0}")]
|
||||
InvalidResponse(String),
|
||||
#[error("unsupported by the rust path: {0}")]
|
||||
Unsupported(&'static str),
|
||||
#[error(transparent)]
|
||||
Auth(#[from] litellm_auth::Error),
|
||||
#[error(transparent)]
|
||||
Transport(#[from] litellm_http::transport::Error),
|
||||
#[error(transparent)]
|
||||
Headers(#[from] litellm_http::request::HeaderError),
|
||||
#[error(transparent)]
|
||||
Http(#[from] litellm_http::Error),
|
||||
#[error(transparent)]
|
||||
Aws(#[from] litellm_auth_aws::Error),
|
||||
}
|
||||
|
||||
impl From<LlmError> for Error {
|
||||
fn from(error: LlmError) -> Self {
|
||||
match error {
|
||||
LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual },
|
||||
LlmError::MissingField(field) => Self::MissingField(field),
|
||||
LlmError::InvalidRequest(message) => Self::InvalidRequest(message),
|
||||
LlmError::InvalidResponse(message) => Self::InvalidResponse(message),
|
||||
LlmError::Unsupported(reason) => Self::Unsupported(reason),
|
||||
LlmError::Auth(error) => Self::Auth(error),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1,22 +1,33 @@
|
|||
use litellm_http::request::truncate_error_body;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_http::{Client, request::truncate_error_body};
|
||||
use litellm_llms::base_llm::auth::resolve_auth;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{Error, client::http_client};
|
||||
use crate::audio_transcription::types::ProviderAudioTranscriptionRequest;
|
||||
use super::Error;
|
||||
use crate::{
|
||||
audio_transcription::types::ProviderAudioTranscriptionRequest,
|
||||
constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS,
|
||||
};
|
||||
|
||||
pub async fn execute_audio_transcription_provider_call(
|
||||
http: &Client,
|
||||
auth: &litellm_auth::AuthServices,
|
||||
request: ProviderAudioTranscriptionRequest,
|
||||
) -> Result<Value, Error> {
|
||||
let response = crate::outbound::outbound_request::<Error>(
|
||||
&request.auth,
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?;
|
||||
let response = crate::outbound::outbound_request(
|
||||
authenticated,
|
||||
request.url.clone(),
|
||||
request.upstream_headers.clone(),
|
||||
&request.body,
|
||||
request.timeout,
|
||||
&request.optional_params,
|
||||
)
|
||||
.await?
|
||||
.send(http_client())
|
||||
Some(
|
||||
request
|
||||
.timeout
|
||||
.unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)),
|
||||
),
|
||||
)?
|
||||
.send(http)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
|
|
|
|||
|
|
@ -1,16 +1,20 @@
|
|||
mod error;
|
||||
pub mod types;
|
||||
pub use error::Error;
|
||||
mod client;
|
||||
pub use crate::error::RouteError as Error;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub use handler::execute_audio_transcription_provider_call;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
pub use prepare::prepare_audio_transcription_provider_call;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::audio_transcription::types::AudioTranscriptionRequest;
|
||||
|
||||
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
|
||||
.await
|
||||
pub async fn audio_transcription(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
request: AudioTranscriptionRequest<'_>,
|
||||
) -> Result<Value, Error> {
|
||||
let request = prepare_audio_transcription_provider_call(request)?;
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
execute_audio_transcription_provider_call(&http, &resources.auth, request).await
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,10 @@
|
|||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_http::request::{has_header, string_headers};
|
||||
use litellm_http::request::string_headers;
|
||||
use litellm_llms::{
|
||||
base_llm::audio_transcription::transformation::{BaseAudioTranscriptionConfig, RequestAuth},
|
||||
base_llm::{
|
||||
audio_transcription::transformation::BaseAudioTranscriptionConfig,
|
||||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
},
|
||||
bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG,
|
||||
};
|
||||
|
||||
|
|
@ -39,20 +42,13 @@ pub fn prepare_audio_transcription_provider_call(
|
|||
let config = provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let mut headers = string_headers("audio transcription", request.extra_headers)?;
|
||||
let auth = config.auth_strategy(&model, &request.optional_params, &env_lookup)?;
|
||||
match &auth {
|
||||
RequestAuth::Bearer { token } if !has_header(&headers, "authorization") => {
|
||||
headers.push(("Authorization".to_string(), format!("Bearer {token}")));
|
||||
}
|
||||
RequestAuth::Header { name, value } if !has_header(&headers, name) => {
|
||||
headers.push(((*name).to_string(), value.clone()));
|
||||
}
|
||||
RequestAuth::Bearer { .. } | RequestAuth::Header { .. } | RequestAuth::AwsSigV4 { .. } => {}
|
||||
}
|
||||
if !has_header(&headers, "content-type") {
|
||||
headers.push(("Content-Type".to_string(), "application/json".to_string()));
|
||||
}
|
||||
let forwarded = string_headers("audio transcription", request.extra_headers)?;
|
||||
let validated =
|
||||
config.validate_environment(forwarded, &model, &request.optional_params, &env_lookup)?;
|
||||
let environment = ValidatedEnvironment {
|
||||
headers: with_default_headers(validated.headers, &[("Content-Type", "application/json")]),
|
||||
auth: validated.auth,
|
||||
};
|
||||
let url = config.get_complete_url(
|
||||
request.api_base,
|
||||
&model,
|
||||
|
|
@ -68,9 +64,7 @@ pub fn prepare_audio_transcription_provider_call(
|
|||
config,
|
||||
url,
|
||||
body: transformed.body,
|
||||
upstream_headers: headers,
|
||||
auth,
|
||||
optional_params: request.optional_params,
|
||||
environment,
|
||||
timeout: request.timeout,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_llms::base_llm::audio_transcription::transformation::{
|
||||
BaseAudioTranscriptionConfig, RequestAuth,
|
||||
use litellm_llms::base_llm::{
|
||||
audio_transcription::transformation::BaseAudioTranscriptionConfig, auth::ValidatedEnvironment,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
|
|
@ -23,9 +23,7 @@ pub struct ProviderAudioTranscriptionRequest {
|
|||
pub config: &'static dyn BaseAudioTranscriptionConfig,
|
||||
pub url: String,
|
||||
pub body: Value,
|
||||
pub upstream_headers: Vec<(String, String)>,
|
||||
pub auth: RequestAuth,
|
||||
pub optional_params: Map<String, Value>,
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,14 +0,0 @@
|
|||
use std::{sync::OnceLock, time::Duration};
|
||||
|
||||
use crate::constants::{CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS, CHAT_COMPLETIONS_TIMEOUT_SECS};
|
||||
|
||||
pub(super) fn http_client() -> &'static reqwest::Client {
|
||||
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
|
||||
CLIENT.get_or_init(|| {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS))
|
||||
.connect_timeout(Duration::from_secs(CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS))
|
||||
.build()
|
||||
.unwrap_or_else(|_| reqwest::Client::new())
|
||||
})
|
||||
}
|
||||
|
|
@ -1,43 +0,0 @@
|
|||
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("expected {expected}, got {actual}")]
|
||||
InvalidType {
|
||||
expected: &'static str,
|
||||
actual: &'static str,
|
||||
},
|
||||
#[error("missing required field: {0}")]
|
||||
MissingField(&'static str),
|
||||
#[error("invalid provider: {0}")]
|
||||
InvalidProvider(String),
|
||||
#[error("invalid request: {0}")]
|
||||
InvalidRequest(String),
|
||||
#[error("invalid response: {0}")]
|
||||
InvalidResponse(String),
|
||||
#[error("unsupported by the rust path: {0}")]
|
||||
Unsupported(&'static str),
|
||||
#[error(transparent)]
|
||||
Auth(#[from] litellm_auth::Error),
|
||||
#[error(transparent)]
|
||||
Transport(#[from] litellm_http::transport::Error),
|
||||
#[error(transparent)]
|
||||
Headers(#[from] litellm_http::request::HeaderError),
|
||||
#[error(transparent)]
|
||||
Http(#[from] litellm_http::Error),
|
||||
#[error(transparent)]
|
||||
Aws(#[from] litellm_auth_aws::Error),
|
||||
}
|
||||
|
||||
impl From<LlmError> for Error {
|
||||
fn from(error: LlmError) -> Self {
|
||||
match error {
|
||||
LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual },
|
||||
LlmError::MissingField(field) => Self::MissingField(field),
|
||||
LlmError::InvalidRequest(message) => Self::InvalidRequest(message),
|
||||
LlmError::InvalidResponse(message) => Self::InvalidResponse(message),
|
||||
LlmError::Unsupported(reason) => Self::Unsupported(reason),
|
||||
LlmError::Auth(error) => Self::Auth(error),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1,20 +1,70 @@
|
|||
use litellm_http::{outbound::OutboundRequest, request::truncate_error_body};
|
||||
use litellm_llms::base_llm::chat::transformation::ProviderChatResponseData;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body};
|
||||
use litellm_llms::base_llm::{
|
||||
auth::{Authenticated, resolve_auth},
|
||||
chat::transformation::ProviderChatResponseData,
|
||||
};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{Error, client::http_client, prepare::prepare_provider_request};
|
||||
use crate::chat_completions::types::{
|
||||
ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
|
||||
use super::Error;
|
||||
use crate::{
|
||||
chat_completions::types::ProviderChatCompletionsRequest,
|
||||
constants::CHAT_COMPLETIONS_TIMEOUT_SECS,
|
||||
};
|
||||
|
||||
pub(super) async fn execute_chat_completions_provider_call(
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
pub(super) async fn execute(
|
||||
http: &Client,
|
||||
auth: &AuthServices,
|
||||
request: ProviderChatCompletionsRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let request = prepare_provider_request(request)?;
|
||||
let outbound = outbound_request(&request).await?;
|
||||
let ProviderChatCompletionsRequest {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
config,
|
||||
url,
|
||||
body,
|
||||
optional_params,
|
||||
environment,
|
||||
timeout,
|
||||
api_key,
|
||||
} = request;
|
||||
let context = RequestContext {
|
||||
model: model.clone(),
|
||||
custom_llm_provider,
|
||||
optional_params: Value::Object(optional_params),
|
||||
secret_fields: Vec::new(),
|
||||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
body,
|
||||
},
|
||||
context,
|
||||
)
|
||||
.await?;
|
||||
let outbound = outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)?;
|
||||
|
||||
let response = outbound.send(http_client()).await.map_err(|err| {
|
||||
let response = outbound.send(http).await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
// so the host can still serve it. Everything else here, a timeout
|
||||
// above all, may have reached the provider and been answered.
|
||||
|
|
@ -36,13 +86,17 @@ pub(super) async fn execute_chat_completions_provider_call(
|
|||
body: truncate_error_body(&text),
|
||||
}));
|
||||
}
|
||||
hooks
|
||||
.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
|
||||
let body: Value = serde_json::from_str(&text).map_err(|err| {
|
||||
Error::InvalidResponse(format!("invalid chat completions response JSON: {err}"))
|
||||
})?;
|
||||
request
|
||||
.config
|
||||
.transform_response(&request.model, ProviderChatResponseData { body })
|
||||
config
|
||||
.transform_response(&model, ProviderChatResponseData { body })
|
||||
.map_err(Error::from)
|
||||
.map_err(as_response_error)
|
||||
}
|
||||
|
|
@ -64,24 +118,176 @@ pub(super) fn as_response_error(err: Error) -> Error {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) async fn outbound_request(
|
||||
request: &ProviderChatCompletionsRequest,
|
||||
pub(super) fn outbound_request(
|
||||
authenticated: Authenticated,
|
||||
url: String,
|
||||
body: &Value,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<OutboundRequest, Error> {
|
||||
crate::outbound::outbound_request(
|
||||
&request.auth,
|
||||
request.url.clone(),
|
||||
request.upstream_headers.clone(),
|
||||
&request.body,
|
||||
request.timeout,
|
||||
&request.optional_params,
|
||||
authenticated,
|
||||
url,
|
||||
body,
|
||||
Some(timeout.unwrap_or(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS))),
|
||||
)
|
||||
.await
|
||||
.map_err(|error| match error {
|
||||
// Python drops the caller's copy and prefers a forwarded Authorization
|
||||
// over the signature, so leave the request to it.
|
||||
Error::Http(litellm_http::Error::ComputedHeader(_)) => {
|
||||
litellm_http::Error::ComputedHeader(_) => {
|
||||
Error::Unsupported("request forwards a header AWS SigV4 computes")
|
||||
}
|
||||
other => other,
|
||||
other => Error::Http(other),
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Mutex;
|
||||
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
|
||||
|
||||
use super::*;
|
||||
use crate::chat_completions::{
|
||||
prepare::{prepare_provider_request, resolve_request},
|
||||
types::ChatCompletionsRequest,
|
||||
};
|
||||
|
||||
const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
|
||||
|
||||
/// Rewrites the outgoing request and records what the call reports back.
|
||||
#[derive(Default)]
|
||||
struct RecordingHooks {
|
||||
contexts: Mutex<Vec<RequestContext>>,
|
||||
raw: Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl RouteHooks<Error> for RecordingHooks {
|
||||
async fn before_send(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
self.contexts.lock().unwrap().push(context);
|
||||
let mut body = wire.body;
|
||||
body["system"] = json!("added by the host");
|
||||
Ok(WireRequest {
|
||||
headers: wire
|
||||
.headers
|
||||
.into_iter()
|
||||
.chain([("x-host".to_string(), "seen".to_string())])
|
||||
.collect(),
|
||||
body,
|
||||
..wire
|
||||
})
|
||||
}
|
||||
|
||||
async fn emit(&self, event: MachineEvent) -> Result<(), Error> {
|
||||
let MachineEvent::ResponseReceived { raw } = event;
|
||||
self.raw.lock().unwrap().push(raw.body);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn prepared(api_base: &str) -> ProviderChatCompletionsRequest {
|
||||
prepare_provider_request(
|
||||
resolve_request(ChatCompletionsRequest {
|
||||
model: "anthropic/claude-sonnet-4-5",
|
||||
messages: json!([{"role": "user", "content": "hi"}]),
|
||||
optional_params: json!({"max_tokens": 16}).as_object().unwrap().clone(),
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some(api_base),
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
})
|
||||
.unwrap(),
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_hooks_rewrite_the_wire_request_and_see_the_raw_response() {
|
||||
let upstream = MockServer::start().await;
|
||||
Mock::given(any())
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_raw(ANTHROPIC_MESSAGE, "application/json"),
|
||||
)
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let hooks = RecordingHooks::default();
|
||||
|
||||
execute(
|
||||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
&hooks,
|
||||
)
|
||||
.await
|
||||
.expect("chat completions call succeeds");
|
||||
|
||||
let [request] = <[Request; 1]>::try_from(upstream.received_requests().await.unwrap())
|
||||
.unwrap_or_else(|requests| panic!("one request, saw {}", requests.len()));
|
||||
let sent: Value = serde_json::from_slice(&request.body).unwrap();
|
||||
assert_eq!(sent["system"], "added by the host");
|
||||
assert_eq!(request.headers["x-host"], "seen");
|
||||
assert_eq!(request.headers["x-api-key"], "sk-test");
|
||||
let [context] = <[RequestContext; 1]>::try_from(hooks.contexts.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
|
||||
assert_eq!(
|
||||
(context.model.as_str(), context.custom_llm_provider.as_str()),
|
||||
("claude-sonnet-4-5", "anthropic")
|
||||
);
|
||||
assert_eq!(context.optional_params, json!({"max_tokens": 16}));
|
||||
assert_eq!(hooks.raw.into_inner().unwrap(), [ANTHROPIC_MESSAGE]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn an_upstream_failure_is_not_reported_as_a_received_response() {
|
||||
let upstream = MockServer::start().await;
|
||||
Mock::given(any())
|
||||
.respond_with(ResponseTemplate::new(500).set_body_string("boom"))
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let hooks = RecordingHooks::default();
|
||||
|
||||
let error = execute(
|
||||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
&hooks,
|
||||
)
|
||||
.await
|
||||
.expect_err("the upstream failure fails the call");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
|
||||
));
|
||||
assert!(hooks.raw.into_inner().unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
|
||||
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"
|
||||
);
|
||||
}
|
||||
let upstream = Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: 500,
|
||||
body: "boom".to_string(),
|
||||
});
|
||||
assert_eq!(as_response_error(upstream.clone()), upstream);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,24 +6,26 @@
|
|||
//! credentials, and it resolves the provider, translates the conversation,
|
||||
//! calls the provider, and returns a typed OpenAI-shaped response.
|
||||
|
||||
mod error;
|
||||
pub mod types;
|
||||
pub use error::Error;
|
||||
mod client;
|
||||
pub use crate::error::RouteError as Error;
|
||||
mod common_utils;
|
||||
pub(crate) mod handler;
|
||||
mod prepare;
|
||||
use handler::execute_chat_completions_provider_call;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use prepare::{parse_messages, resolve_provider_config, resolve_request};
|
||||
use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::chat_completions::types::ChatCompletionsRequest;
|
||||
|
||||
pub async fn chat_completions(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
execute_chat_completions_provider_call(resolve_request(request)?).await
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let request = prepare_provider_request(resolve_request(request)?)?;
|
||||
handler::execute(&http, &resources.auth, request, &()).await
|
||||
}
|
||||
|
||||
/// Whether the core would accept this request, without resolving credentials or
|
||||
|
|
@ -38,9 +40,10 @@ pub fn chat_completions_decline_reason(
|
|||
messages: Value,
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Option<&'static str> {
|
||||
let Ok((_, config)) = resolve_provider_config(model, custom_llm_provider) else {
|
||||
let Ok(resolved) = resolve_provider_config(model, custom_llm_provider) else {
|
||||
return Some("provider is not on the rust chat completions path");
|
||||
};
|
||||
let config = resolved.config;
|
||||
let Ok(messages) = parse_messages(messages) else {
|
||||
return Some("unreadable message list");
|
||||
};
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_http::request::has_header;
|
||||
use litellm_llms::base_llm::chat::transformation::{BaseConfig, RequestAuth};
|
||||
use litellm_llms::base_llm::{
|
||||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
chat::transformation::BaseConfig,
|
||||
};
|
||||
use litellm_types::llms::openai::ChatMessage;
|
||||
use serde_json::Value;
|
||||
|
||||
|
|
@ -12,10 +15,16 @@ use crate::chat_completions::types::{
|
|||
ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
|
||||
};
|
||||
|
||||
pub(super) struct ResolvedProvider {
|
||||
pub(super) model: String,
|
||||
pub(super) custom_llm_provider: String,
|
||||
pub(super) config: &'static dyn BaseConfig,
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider_config<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<(String, &'static dyn BaseConfig), Error> {
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
let provider_info = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
|
|
@ -30,7 +39,11 @@ pub(super) fn resolve_provider_config<'a>(
|
|||
})?;
|
||||
let config = chat_completions_provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
Ok((provider_info.model.to_string(), config))
|
||||
Ok(ResolvedProvider {
|
||||
model: provider_info.model.to_string(),
|
||||
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error> {
|
||||
|
|
@ -41,7 +54,11 @@ pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error>
|
|||
pub(super) fn resolve_request(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ResolvedChatCompletionsRequest<'_>, Error> {
|
||||
let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?;
|
||||
let ResolvedProvider {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
config,
|
||||
} = resolve_provider_config(request.model, request.custom_llm_provider)?;
|
||||
let messages = parse_messages(request.messages)?;
|
||||
if messages.is_empty() {
|
||||
return Err(Error::InvalidRequest(
|
||||
|
|
@ -53,6 +70,7 @@ pub(super) fn resolve_request(
|
|||
}
|
||||
Ok(ResolvedChatCompletionsRequest {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
config,
|
||||
messages,
|
||||
optional_params: request.optional_params,
|
||||
|
|
@ -67,59 +85,26 @@ fn validate_environment(
|
|||
request: &ResolvedChatCompletionsRequest<'_>,
|
||||
model: &str,
|
||||
config: &dyn BaseConfig,
|
||||
) -> Result<(Vec<(String, String)>, RequestAuth), Error> {
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let mut headers = string_headers(request.extra_headers.clone())?;
|
||||
let auth = config.auth(
|
||||
let forwarded = string_headers(request.extra_headers.clone())?;
|
||||
let validated = config.validate_environment(
|
||||
forwarded,
|
||||
request.api_key,
|
||||
model,
|
||||
&request.optional_params,
|
||||
&env_lookup,
|
||||
)?;
|
||||
match &auth {
|
||||
RequestAuth::Header { name, value } => {
|
||||
// The deployment's credential replaces whatever the caller forwarded
|
||||
// under the same name, mirroring Python's
|
||||
// `{**headers, **anthropic_headers}`: letting a request header win
|
||||
// would let its sender choose the principal the call bills to.
|
||||
//
|
||||
// The exception is a scheme the provider hands off to entirely, such
|
||||
// as an Anthropic OAuth bearer, where Python drops `x-api-key`
|
||||
// instead of resolving one. Re-adding it there would put the
|
||||
// credential into a header the host removed on purpose.
|
||||
if !config.defers_to_forwarded_auth(&headers) {
|
||||
headers.retain(|(header, _)| !header.eq_ignore_ascii_case(name));
|
||||
headers.push(((*name).to_string(), value.clone()));
|
||||
}
|
||||
}
|
||||
RequestAuth::Bearer { token } => {
|
||||
// Bedrock's `get_request_headers` assigns `headers["Authorization"]`
|
||||
// unconditionally once a bearer token resolves, so the deployment's
|
||||
// identity outranks whatever the caller forwarded. Keeping the
|
||||
// caller's would bill and authorize the call as a different
|
||||
// principal than the same deployment uses on Python.
|
||||
//
|
||||
// The `Header` arm below keeps the opposite precedence on purpose:
|
||||
// Anthropic's transform honours a forwarded OAuth bearer.
|
||||
headers.retain(|(name, _)| !name.eq_ignore_ascii_case("authorization"));
|
||||
headers.push(("authorization".to_string(), format!("Bearer {token}")));
|
||||
}
|
||||
// SigV4 signs the serialized body, so the handler adds its headers.
|
||||
RequestAuth::AwsSigV4 { .. } => {}
|
||||
}
|
||||
|
||||
for (name, value) in config.default_headers() {
|
||||
if !has_header(&headers, name) {
|
||||
headers.push(((*name).to_string(), (*value).to_string()));
|
||||
}
|
||||
}
|
||||
Ok((headers, auth))
|
||||
Ok(ValidatedEnvironment {
|
||||
headers: with_default_headers(validated.headers, config.default_headers()),
|
||||
auth: validated.auth,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
let (headers, auth) = validate_environment(&request, &request.model, request.config)?;
|
||||
let environment = validate_environment(&request, &request.model, request.config)?;
|
||||
let model = request.model;
|
||||
let config = request.config;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
|
|
@ -134,19 +119,21 @@ pub(super) fn prepare_provider_request(
|
|||
|
||||
Ok(ProviderChatCompletionsRequest {
|
||||
model,
|
||||
custom_llm_provider: request.custom_llm_provider,
|
||||
config,
|
||||
url,
|
||||
body: transformed.body,
|
||||
upstream_headers: headers,
|
||||
auth,
|
||||
optional_params: request.optional_params,
|
||||
environment,
|
||||
timeout: request.timeout,
|
||||
api_key: request.api_key.map(|key| SecretValue::new(key.to_string())),
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::base_llm::chat::transformation::RequestAuth;
|
||||
use litellm_auth::CredentialPlacement;
|
||||
use litellm_llms::base_llm::auth::{AuthScheme, resolve_auth};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::{prepare_provider_request, resolve_request};
|
||||
|
|
@ -161,6 +148,20 @@ mod tests {
|
|||
prepare_provider_request(resolve_request(request)?)
|
||||
}
|
||||
|
||||
/// The headers as they go on the wire, credential applied.
|
||||
fn wire_headers(prepared: &ProviderChatCompletionsRequest) -> Vec<(String, String)> {
|
||||
tokio::runtime::Builder::new_current_thread()
|
||||
.build()
|
||||
.unwrap()
|
||||
.block_on(resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
prepared.environment.clone(),
|
||||
&|_| None,
|
||||
))
|
||||
.unwrap()
|
||||
.headers
|
||||
}
|
||||
|
||||
fn request<'a>(
|
||||
model: &'a str,
|
||||
provider: Option<&'a str>,
|
||||
|
|
@ -227,19 +228,16 @@ mod tests {
|
|||
))
|
||||
.expect("prepares");
|
||||
assert!(
|
||||
prepared
|
||||
.upstream_headers
|
||||
.contains(&("x-api-key".to_string(), "sk-test".to_string()))
|
||||
wire_headers(&prepared).contains(&("x-api-key".to_string(), "sk-test".to_string()))
|
||||
);
|
||||
assert!(
|
||||
prepared
|
||||
.upstream_headers
|
||||
wire_headers(&prepared)
|
||||
.contains(&("anthropic-version".to_string(), "2023-06-01".to_string()))
|
||||
);
|
||||
assert!(matches!(
|
||||
prepared.auth,
|
||||
RequestAuth::Header {
|
||||
name: "x-api-key",
|
||||
prepared.environment.auth,
|
||||
AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Header("x-api-key"),
|
||||
..
|
||||
}
|
||||
));
|
||||
|
|
@ -261,12 +259,12 @@ mod tests {
|
|||
json!("sk-caller"),
|
||||
)]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let keys: Vec<_> = prepared
|
||||
.upstream_headers
|
||||
let headers = wire_headers(&prepared);
|
||||
let keys: Vec<_> = 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.len(), 1, "got {:?}", headers);
|
||||
assert_eq!(keys[0].1, "sk-test");
|
||||
}
|
||||
|
||||
|
|
@ -290,16 +288,14 @@ mod tests {
|
|||
]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
assert!(
|
||||
!prepared
|
||||
.upstream_headers
|
||||
!wire_headers(&prepared)
|
||||
.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
|
||||
wire_headers(&prepared)
|
||||
);
|
||||
assert!(
|
||||
prepared
|
||||
.upstream_headers
|
||||
wire_headers(&prepared)
|
||||
.iter()
|
||||
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
|
||||
&& value == "Bearer sk-ant-oat01-token")
|
||||
|
|
@ -322,21 +318,20 @@ mod tests {
|
|||
("X-Api-Key".to_string(), json!("sk-caller")),
|
||||
]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let keys: Vec<_> = prepared
|
||||
.upstream_headers
|
||||
let headers = wire_headers(&prepared);
|
||||
let keys: Vec<_> = 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.len(), 1, "got {:?}", headers);
|
||||
assert_eq!(keys[0].1, "sk-test");
|
||||
assert!(
|
||||
prepared
|
||||
.upstream_headers
|
||||
wire_headers(&prepared)
|
||||
.iter()
|
||||
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
|
||||
&& value == "Bearer unrelated"),
|
||||
"the unrelated authorization must survive, got {:?}",
|
||||
prepared.upstream_headers
|
||||
wire_headers(&prepared)
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -435,18 +430,16 @@ mod tests {
|
|||
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",
|
||||
}
|
||||
);
|
||||
assert!(matches!(
|
||||
&prepared.environment.auth,
|
||||
AuthScheme::AwsSigV4 { region, service: "bedrock", .. } if region == "us-east-1"
|
||||
));
|
||||
// SigV4 signs the serialized body, so prepare must not have added an
|
||||
// Authorization header; the handler does it.
|
||||
// Authorization header; the signer does it over the bytes sent.
|
||||
assert!(
|
||||
!prepared
|
||||
.upstream_headers
|
||||
.environment
|
||||
.headers
|
||||
.iter()
|
||||
.any(|(name, _)| name.eq_ignore_ascii_case("authorization"))
|
||||
);
|
||||
|
|
@ -475,9 +468,20 @@ mod tests {
|
|||
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 authenticated = resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
prepared.environment,
|
||||
&|_| None,
|
||||
)
|
||||
.await
|
||||
.expect("resolves");
|
||||
let signed = crate::chat_completions::handler::outbound_request(
|
||||
authenticated,
|
||||
prepared.url,
|
||||
&prepared.body,
|
||||
prepared.timeout,
|
||||
)
|
||||
.expect("signs");
|
||||
|
||||
let authorization = signed
|
||||
.header("authorization")
|
||||
|
|
@ -525,9 +529,20 @@ mod tests {
|
|||
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");
|
||||
let authenticated = resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
prepared.environment,
|
||||
&|_| None,
|
||||
)
|
||||
.await
|
||||
.expect("resolves");
|
||||
let error = crate::chat_completions::handler::outbound_request(
|
||||
authenticated,
|
||||
prepared.url,
|
||||
&prepared.body,
|
||||
prepared.timeout,
|
||||
)
|
||||
.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"
|
||||
|
|
@ -552,8 +567,8 @@ mod tests {
|
|||
json!("Bearer caller-supplied"),
|
||||
)]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let authorizations: Vec<_> = prepared
|
||||
.upstream_headers
|
||||
let headers = wire_headers(&prepared);
|
||||
let authorizations: Vec<_> = headers
|
||||
.iter()
|
||||
.filter(|(name, _)| name.eq_ignore_ascii_case("authorization"))
|
||||
.map(|(_, value)| value.as_str())
|
||||
|
|
@ -585,16 +600,15 @@ mod tests {
|
|||
json!("Bearer sk-ant-oat01-forwarded"),
|
||||
)]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let keys: Vec<_> = prepared
|
||||
.upstream_headers
|
||||
let headers = wire_headers(&prepared);
|
||||
let keys: Vec<_> = 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!(keys.is_empty(), "got {:?}", headers);
|
||||
assert!(
|
||||
prepared
|
||||
.upstream_headers
|
||||
wire_headers(&prepared)
|
||||
.iter()
|
||||
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
|
||||
&& value == "Bearer sk-ant-oat01-forwarded")
|
||||
|
|
@ -613,15 +627,13 @@ mod tests {
|
|||
json!({"maxTokens": 16}),
|
||||
))
|
||||
.expect("prepares");
|
||||
assert_eq!(
|
||||
prepared.auth,
|
||||
RequestAuth::Bearer {
|
||||
token: "sk-test".to_string()
|
||||
}
|
||||
);
|
||||
assert!(matches!(
|
||||
&prepared.environment.auth,
|
||||
AuthScheme::Credential { placement: CredentialPlacement::Bearer, secret }
|
||||
if secret.expose() == "sk-test"
|
||||
));
|
||||
assert!(
|
||||
prepared
|
||||
.upstream_headers
|
||||
wire_headers(&prepared)
|
||||
.iter()
|
||||
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
|
||||
&& value == "Bearer sk-test"),
|
||||
|
|
@ -736,248 +748,4 @@ mod tests {
|
|||
.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,6 +1,7 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_llms::base_llm::chat::transformation::{BaseConfig, RequestAuth};
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig};
|
||||
use litellm_types::llms::openai::ChatMessage;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
|
|
@ -23,6 +24,7 @@ pub struct ChatCompletionsRequest<'a> {
|
|||
|
||||
pub struct ResolvedChatCompletionsRequest<'a> {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
pub config: &'static dyn BaseConfig,
|
||||
pub messages: Vec<ChatMessage>,
|
||||
pub optional_params: Map<String, Value>,
|
||||
|
|
@ -34,11 +36,16 @@ pub struct ResolvedChatCompletionsRequest<'a> {
|
|||
|
||||
pub struct ProviderChatCompletionsRequest {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
pub config: &'static dyn BaseConfig,
|
||||
pub url: String,
|
||||
pub body: Value,
|
||||
pub upstream_headers: Vec<(String, String)>,
|
||||
pub auth: RequestAuth,
|
||||
/// The route's parameters before the provider transformation, reported to the host
|
||||
/// beside the wire request.
|
||||
pub optional_params: Map<String, Value>,
|
||||
/// The forwarded and default headers plus how the call authenticates; the credential
|
||||
/// itself is applied when the request is sent.
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub timeout: Option<Duration>,
|
||||
pub api_key: Option<SecretValue>,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,20 +5,10 @@ pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com";
|
|||
/// timeout from the caller still overrides this on the request builder.
|
||||
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// Connect timeout for Anthropic Messages provider calls, in seconds.
|
||||
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
|
||||
/// Provider name used for Anthropic Messages when a deployment's provider model
|
||||
/// does not carry an explicit provider prefix.
|
||||
pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
|
||||
|
||||
/// Full-request timeout ceiling for chat completions provider calls, in
|
||||
/// seconds. Mirrors the Python chat completions default.
|
||||
pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// Connect timeout for chat completions provider calls, in seconds.
|
||||
pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
|
||||
pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// `object` field every non-streaming chat completion response carries.
|
||||
|
|
|
|||
|
|
@ -1,15 +1,166 @@
|
|||
use litellm_llms::base_llm::ocr::error::Error as OcrError;
|
||||
//! One error for every route in this crate. OCR still carries its own, richer enum.
|
||||
//!
|
||||
//! A variant is declared by the layer that produces it and nested here as is:
|
||||
//! credentials by `litellm_auth` (AWS folds into it at that crate's boundary), the wire by
|
||||
//! `litellm_http`, secrets by `litellm_secrets`. The transformation layer's [`LlmError`]
|
||||
//! maps onto the same-named variants once, here, so no route re-declares them.
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_llms::Error as LlmError;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
pub enum RouteError {
|
||||
#[error("expected {expected}, got {actual}")]
|
||||
InvalidType {
|
||||
expected: &'static str,
|
||||
actual: &'static str,
|
||||
},
|
||||
#[error("missing required field: {0}")]
|
||||
MissingField(&'static str),
|
||||
#[error("invalid provider: {0}")]
|
||||
InvalidProvider(String),
|
||||
#[error("invalid request: {0}")]
|
||||
InvalidRequest(String),
|
||||
#[error("invalid response: {0}")]
|
||||
InvalidResponse(String),
|
||||
#[error("unsupported by the rust path: {0}")]
|
||||
Unsupported(&'static str),
|
||||
#[error(transparent)]
|
||||
Ocr(#[from] OcrError),
|
||||
Auth(#[from] litellm_auth::Error),
|
||||
#[error(transparent)]
|
||||
Messages(#[from] crate::messages::Error),
|
||||
Transport(#[from] TransportError),
|
||||
#[error(transparent)]
|
||||
ChatCompletions(#[from] crate::chat_completions::Error),
|
||||
Headers(#[from] litellm_http::request::HeaderError),
|
||||
#[error(transparent)]
|
||||
AudioTranscription(#[from] crate::audio_transcription::Error),
|
||||
Http(#[from] litellm_http::Error),
|
||||
#[error(transparent)]
|
||||
Responses(#[from] crate::responses::Error),
|
||||
Secret(#[from] SecretError),
|
||||
}
|
||||
|
||||
/// Whether the provider had already been called when the route failed. Before the send, a
|
||||
/// host may retry on another path; after it, the provider has done the work and billed for it.
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Phase {
|
||||
BeforeSend,
|
||||
AfterSend,
|
||||
}
|
||||
|
||||
impl RouteError {
|
||||
pub fn phase(&self) -> Phase {
|
||||
match self {
|
||||
Self::InvalidResponse(_)
|
||||
| Self::Transport(TransportError::Http { .. } | TransportError::Network(_)) => {
|
||||
Phase::AfterSend
|
||||
}
|
||||
Self::Transport(TransportError::Connect(_))
|
||||
| Self::InvalidType { .. }
|
||||
| Self::MissingField(_)
|
||||
| Self::InvalidProvider(_)
|
||||
| Self::InvalidRequest(_)
|
||||
| Self::Unsupported(_)
|
||||
| Self::Auth(_)
|
||||
| Self::Headers(_)
|
||||
| Self::Http(_)
|
||||
| Self::Secret(_) => Phase::BeforeSend,
|
||||
}
|
||||
}
|
||||
|
||||
/// The caller's request is what is wrong, as opposed to the environment, the wire, or
|
||||
/// the provider's answer.
|
||||
pub fn is_request(&self) -> bool {
|
||||
match self {
|
||||
Self::InvalidType { .. }
|
||||
| Self::MissingField(_)
|
||||
| Self::InvalidProvider(_)
|
||||
| Self::InvalidRequest(_)
|
||||
| Self::Unsupported(_)
|
||||
| Self::Headers(_) => true,
|
||||
Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }),
|
||||
Self::InvalidResponse(_) | Self::Transport(_) | Self::Http(_) | Self::Secret(_) => {
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<LlmError> for RouteError {
|
||||
fn from(error: LlmError) -> Self {
|
||||
match error {
|
||||
LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual },
|
||||
LlmError::MissingField(field) => Self::MissingField(field),
|
||||
LlmError::InvalidRequest(message) => Self::InvalidRequest(message),
|
||||
LlmError::InvalidResponse(message) => Self::InvalidResponse(message),
|
||||
LlmError::Unsupported(reason) => Self::Unsupported(reason),
|
||||
LlmError::Auth(error) => Self::Auth(error),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[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 RouteError {
|
||||
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 {}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{Phase, RouteError};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
|
||||
#[test]
|
||||
fn only_a_provider_answer_or_a_lost_connection_counts_as_after_send() {
|
||||
let after = [
|
||||
RouteError::InvalidResponse("bad json".into()),
|
||||
RouteError::Transport(TransportError::Http {
|
||||
status: 500,
|
||||
body: "boom".into(),
|
||||
}),
|
||||
RouteError::Transport(TransportError::Network("reset".into())),
|
||||
];
|
||||
for error in after {
|
||||
assert_eq!(error.phase(), Phase::AfterSend, "{error:?}");
|
||||
}
|
||||
let before = [
|
||||
RouteError::Transport(TransportError::Connect("refused".into())),
|
||||
RouteError::Unsupported("streaming"),
|
||||
RouteError::Auth(litellm_auth::Error::InvalidHeader),
|
||||
];
|
||||
for error in before {
|
||||
assert_eq!(error.phase(), Phase::BeforeSend, "{error:?}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_missing_api_key_is_the_environment_not_the_request() {
|
||||
assert!(
|
||||
!RouteError::Auth(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
})
|
||||
.is_request()
|
||||
);
|
||||
assert!(RouteError::Auth(litellm_auth::Error::InvalidHeader).is_request());
|
||||
assert!(RouteError::InvalidRequest("top_k".into()).is_request());
|
||||
assert!(!RouteError::InvalidResponse("bad json".into()).is_request());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ pub mod error;
|
|||
pub mod messages;
|
||||
pub mod ocr;
|
||||
mod outbound;
|
||||
pub mod resources;
|
||||
pub mod responses;
|
||||
|
||||
pub use error::Error;
|
||||
pub use error::{Phase, RouteError};
|
||||
|
|
|
|||
|
|
@ -1,14 +0,0 @@
|
|||
use std::{sync::OnceLock, time::Duration};
|
||||
|
||||
use crate::constants::{MESSAGES_CONNECT_TIMEOUT_SECS, MESSAGES_TIMEOUT_SECS};
|
||||
|
||||
pub(super) fn http_client() -> &'static reqwest::Client {
|
||||
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
|
||||
CLIENT.get_or_init(|| {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(MESSAGES_TIMEOUT_SECS))
|
||||
.connect_timeout(Duration::from_secs(MESSAGES_CONNECT_TIMEOUT_SECS))
|
||||
.build()
|
||||
.unwrap_or_else(|_| reqwest::Client::new())
|
||||
})
|
||||
}
|
||||
|
|
@ -1,23 +1,37 @@
|
|||
use litellm_http::request::string_headers as shared_string_headers;
|
||||
pub(super) use litellm_http::request::truncate_error_body;
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
|
||||
bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
use strum::{EnumString, IntoStaticStr};
|
||||
|
||||
use super::Error;
|
||||
|
||||
const HEADER_CONTEXT: &str = "messages";
|
||||
|
||||
pub(super) fn messages_provider_config(
|
||||
provider: &str,
|
||||
) -> Option<&'static dyn BaseAnthropicMessagesConfig> {
|
||||
match provider {
|
||||
"anthropic" => Some(&ANTHROPIC_MESSAGES_CONFIG),
|
||||
"azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG),
|
||||
_ => None,
|
||||
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub(crate) enum MessagesProvider {
|
||||
Anthropic,
|
||||
AzureAi,
|
||||
Bedrock,
|
||||
}
|
||||
|
||||
impl MessagesProvider {
|
||||
pub(crate) fn as_str(self) -> &'static str {
|
||||
self.into()
|
||||
}
|
||||
|
||||
pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig {
|
||||
match self {
|
||||
Self::Anthropic => &ANTHROPIC_MESSAGES_CONFIG,
|
||||
Self::AzureAi => &AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
Self::Bedrock => &BEDROCK_ANTHROPIC_MESSAGES_CONFIG,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -29,157 +43,28 @@ pub(super) fn string_headers(
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{sync::Arc, time::Duration};
|
||||
use serde_json::json;
|
||||
|
||||
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 rstest::rstest;
|
||||
|
||||
use super::{messages_provider_config, string_headers, truncate_error_body};
|
||||
use crate::messages::{
|
||||
Error,
|
||||
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
|
||||
types::MessagesShaping,
|
||||
};
|
||||
use super::{MessagesProvider, string_headers, truncate_error_body};
|
||||
use crate::messages::Error;
|
||||
|
||||
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<_>>()
|
||||
);
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic", MessagesProvider::Anthropic)]
|
||||
#[case::azure_ai("azure_ai", MessagesProvider::AzureAi)]
|
||||
#[case::bedrock("bedrock", MessagesProvider::Bedrock)]
|
||||
fn provider_round_trips_through_its_python_name(
|
||||
#[case] name: &str,
|
||||
#[case] provider: MessagesProvider,
|
||||
) {
|
||||
assert_eq!(name.parse::<MessagesProvider>(), Ok(provider));
|
||||
assert_eq!(provider.as_str(), name);
|
||||
}
|
||||
|
||||
#[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());
|
||||
fn provider_without_a_messages_config_is_rejected() {
|
||||
assert!("openai".parse::<MessagesProvider>().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -1,80 +0,0 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("invalid provider: {0}")]
|
||||
InvalidProvider(String),
|
||||
#[error("missing required field: {0}")]
|
||||
MissingField(&'static str),
|
||||
#[error("invalid request: {0}")]
|
||||
InvalidRequest(String),
|
||||
#[error("invalid response: {0}")]
|
||||
InvalidResponse(String),
|
||||
#[error("unsupported by the Rust messages route: {0}")]
|
||||
Unsupported(&'static str),
|
||||
#[error(transparent)]
|
||||
Auth(#[from] litellm_auth::Error),
|
||||
#[error(transparent)]
|
||||
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 {
|
||||
error @ LlmError::InvalidType { .. } => Self::InvalidRequest(error.to_string()),
|
||||
LlmError::MissingField(field) => Self::MissingField(field),
|
||||
LlmError::InvalidRequest(message) => Self::InvalidRequest(message),
|
||||
LlmError::InvalidResponse(message) => Self::InvalidResponse(message),
|
||||
LlmError::Unsupported(reason) => Self::Unsupported(reason),
|
||||
LlmError::Auth(error) => Self::Auth(error),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Error {
|
||||
pub fn is_request(&self) -> bool {
|
||||
match self {
|
||||
Self::InvalidProvider(_)
|
||||
| Self::MissingField(_)
|
||||
| Self::InvalidRequest(_)
|
||||
| Self::Unsupported(_)
|
||||
| Self::Headers(_) => true,
|
||||
Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_response(&self) -> bool {
|
||||
matches!(self, Self::InvalidResponse(_))
|
||||
}
|
||||
}
|
||||
|
|
@ -1,45 +1,143 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_http::{request::http_request, transport::Error as TransportError};
|
||||
use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig;
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt, stream::BoxStream};
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_llms::base_llm::{
|
||||
anthropic_messages::{
|
||||
streaming::{ByteStream, StreamDecoder, encode_anthropic_sse},
|
||||
transformation::BaseAnthropicMessagesConfig,
|
||||
},
|
||||
auth::{Authenticated, resolve_auth},
|
||||
};
|
||||
use litellm_tracing::{ByteChunk, debug};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{Error, client::http_client, common_utils::truncate_error_body};
|
||||
use super::{
|
||||
Error, MessagesResponse, common_utils::truncate_error_body, prepare::ProviderMessagesRequest,
|
||||
};
|
||||
use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request};
|
||||
|
||||
pub(super) fn network(error: reqwest::Error) -> Error {
|
||||
pub(super) async fn execute(
|
||||
http: &litellm_http::Client,
|
||||
auth: &AuthServices,
|
||||
request: ProviderMessagesRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let ProviderMessagesRequest {
|
||||
provider,
|
||||
url,
|
||||
body,
|
||||
environment,
|
||||
timeout,
|
||||
api_key,
|
||||
} = request;
|
||||
let stream = body.params.stream == Some(true);
|
||||
let context = RequestContext {
|
||||
model: body.model.clone(),
|
||||
custom_llm_provider: provider.as_str().to_string(),
|
||||
optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?,
|
||||
secret_fields: Vec::new(),
|
||||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
body: serde_json::to_value(&body).map_err(serialize_failure)?,
|
||||
},
|
||||
context,
|
||||
)
|
||||
.await?;
|
||||
let provider_name = provider.as_str();
|
||||
debug!(provider = provider_name, stream, body = %wire.body, "provider request");
|
||||
let response = send(
|
||||
http,
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
&wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)
|
||||
.await?;
|
||||
debug!(
|
||||
provider = provider_name,
|
||||
status = response.status().as_u16(),
|
||||
"provider response headers"
|
||||
);
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
let config = provider.config();
|
||||
if stream {
|
||||
return Ok(streaming_response(
|
||||
response,
|
||||
config.stream_decoder(),
|
||||
provider_name,
|
||||
));
|
||||
}
|
||||
let text = response.text().await.map_err(network)?;
|
||||
debug!(body = text.as_str(), "provider response body");
|
||||
hooks
|
||||
.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
decode_response(config, &body.model, &text)
|
||||
.map(|message| MessagesResponse::Message(Box::new(message)))
|
||||
}
|
||||
|
||||
fn serialize_failure(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
))
|
||||
}
|
||||
|
||||
fn network(error: reqwest::Error) -> Error {
|
||||
Error::Transport(TransportError::Network(error.to_string()))
|
||||
}
|
||||
|
||||
pub(super) async fn send(
|
||||
async fn send(
|
||||
http: &litellm_http::Client,
|
||||
authenticated: Authenticated,
|
||||
url: &str,
|
||||
headers: &[(String, String)],
|
||||
body: &Value,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<reqwest::Response, Error> {
|
||||
let builder = headers.iter().fold(
|
||||
http_client().post(url).json(body),
|
||||
|builder, (key, value)| builder.header(key, value),
|
||||
);
|
||||
let builder = match timeout {
|
||||
Some(duration) => builder.timeout(duration),
|
||||
None => builder,
|
||||
};
|
||||
http_request(builder).await.map_err(network)
|
||||
let request = outbound_request(
|
||||
authenticated,
|
||||
url.to_string(),
|
||||
body,
|
||||
Some(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))),
|
||||
)?;
|
||||
request.send(http).await.map_err(network)
|
||||
}
|
||||
|
||||
pub(super) async fn provider_error(response: reqwest::Response) -> Error {
|
||||
async fn provider_error(response: reqwest::Response) -> Error {
|
||||
let status = response.status().as_u16();
|
||||
match response.text().await {
|
||||
Ok(text) => Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: truncate_error_body(&text),
|
||||
}),
|
||||
Ok(text) => {
|
||||
litellm_tracing::debug!(status, body = text.as_str(), "provider error body");
|
||||
Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: truncate_error_body(&text),
|
||||
})
|
||||
}
|
||||
Err(error) => network(error),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn decode_response(
|
||||
fn decode_response(
|
||||
config: &dyn BaseAnthropicMessagesConfig,
|
||||
model: &str,
|
||||
text: &str,
|
||||
|
|
@ -50,3 +148,97 @@ pub(super) fn decode_response(
|
|||
.transform_anthropic_messages_response(model, response)
|
||||
.map_err(Error::from)
|
||||
}
|
||||
|
||||
fn streaming_response(
|
||||
response: reqwest::Response,
|
||||
decoder: Option<StreamDecoder>,
|
||||
provider: &'static str,
|
||||
) -> MessagesResponse {
|
||||
let headers = response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string())))
|
||||
.collect();
|
||||
let chunks = match decoder {
|
||||
None => futures_util::stream::try_unfold(response, move |mut response| async move {
|
||||
let chunk = response.chunk().await.map_err(network)?;
|
||||
Ok(chunk.map(|chunk| {
|
||||
log_chunk(provider, "provider_response", &chunk);
|
||||
(chunk, response)
|
||||
}))
|
||||
})
|
||||
.boxed(),
|
||||
Some(decode) => decoded_chunks(response, decode, provider),
|
||||
};
|
||||
MessagesResponse::Stream { headers, chunks }
|
||||
}
|
||||
|
||||
fn decoded_chunks(
|
||||
response: reqwest::Response,
|
||||
decode: StreamDecoder,
|
||||
provider: &'static str,
|
||||
) -> BoxStream<'static, Result<Bytes, Error>> {
|
||||
let bytes: ByteStream = response
|
||||
.bytes_stream()
|
||||
.inspect_ok(move |chunk| log_chunk(provider, "provider_response", chunk))
|
||||
.map_err(std::io::Error::other)
|
||||
.boxed();
|
||||
futures_util::stream::try_unfold(decode(bytes), move |mut events| async move {
|
||||
let Some(event) = events.try_next().await? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let chunk = encode_anthropic_sse(&event)?;
|
||||
log_chunk(provider, "client_response", &chunk);
|
||||
Ok(Some((chunk, events)))
|
||||
})
|
||||
.boxed()
|
||||
}
|
||||
|
||||
fn log_chunk(provider: &str, stage: &str, data: &Bytes) {
|
||||
let chunk = ByteChunk::new(data);
|
||||
debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::base_llm::anthropic_messages::streaming::anthropic_sse_event_stream;
|
||||
use rstest::rstest;
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::any};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::event(
|
||||
"data: {\"type\":\"ping\"}\n\n",
|
||||
Some("event: ping\ndata: {\"type\":\"ping\"}\n\n")
|
||||
)]
|
||||
#[case::invalid_event("data: invalid\n\ndata: {\"type\":\"ping\"}\n\n", None)]
|
||||
#[tokio::test]
|
||||
async fn decoded_streams_encode_events_and_stop_at_the_first_error(
|
||||
#[case] body: &'static str,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
let upstream = MockServer::start().await;
|
||||
Mock::given(any())
|
||||
.respond_with(ResponseTemplate::new(200).set_body_raw(body, "text/event-stream"))
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let response = litellm_http::Client::plain_for_test()
|
||||
.get(upstream.uri())
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
let MessagesResponse::Stream { mut chunks, .. } =
|
||||
streaming_response(response, Some(anthropic_sse_event_stream), "test")
|
||||
else {
|
||||
panic!("a streaming response returns chunks");
|
||||
};
|
||||
|
||||
let chunk = chunks.next().await.unwrap();
|
||||
match expected {
|
||||
Some(expected) => assert_eq!(chunk.unwrap().as_ref(), expected.as_bytes()),
|
||||
None => assert!(matches!(chunk, Err(Error::InvalidResponse(_))), "{chunk:?}"),
|
||||
}
|
||||
assert!(chunks.next().await.is_none());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,48 +1,27 @@
|
|||
//! The Anthropic Messages call, the Rust equivalent of Python's
|
||||
//! `litellm.messages()`.
|
||||
//! The Anthropic Messages call, the Rust equivalent of Python's `litellm.messages()`.
|
||||
//!
|
||||
//! [`route`] is the call as a machine a host drives, streaming or not. [`messages`] runs
|
||||
//! it in process for a caller that already holds the request and wants the message.
|
||||
//! [`messages`] prepares the provider request and sends it in process. [`route`] runs the
|
||||
//! same two steps as a machine for a host that answers the call's operations itself.
|
||||
|
||||
mod error;
|
||||
pub mod types;
|
||||
pub use error::Error;
|
||||
mod client;
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
use std::sync::Arc;
|
||||
mod types;
|
||||
|
||||
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;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
||||
use crate::messages::types::MessagesRequest;
|
||||
pub use crate::error::RouteError as Error;
|
||||
pub use types::{MessagesCall, MessagesResponse, MessagesShaping, messages_body};
|
||||
|
||||
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
|
||||
let Value::Object(body) = request.body else {
|
||||
return Err(Error::InvalidRequest(
|
||||
"messages body must be an object".into(),
|
||||
));
|
||||
};
|
||||
let call = MessagesCall {
|
||||
model: request.model.into(),
|
||||
body,
|
||||
api_key: request.api_key.map(Into::into),
|
||||
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,
|
||||
};
|
||||
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",
|
||||
)),
|
||||
}
|
||||
pub async fn messages(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
secrets: &dyn SecretSource,
|
||||
call: MessagesCall,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let request = prepare::prepare(call, secrets).await?;
|
||||
handler::execute(&http, &resources.auth, request, &()).await
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,6 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::{
|
||||
dot_notation_indexing::delete_nested_value,
|
||||
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
|
||||
|
|
@ -5,30 +8,51 @@ use litellm_core_utils::{
|
|||
settings::Lookup,
|
||||
};
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request,
|
||||
base_llm::anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, MessagesTransformContext,
|
||||
anthropic::messages::handler::shape_anthropic_messages_request,
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::MessagesTransformContext,
|
||||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
},
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{messages_provider_config, string_headers},
|
||||
Error, MessagesCall,
|
||||
common_utils::{MessagesProvider, string_headers},
|
||||
types::invalid_request,
|
||||
};
|
||||
use crate::messages::types::{MessagesRequest, ProviderMessagesRequest};
|
||||
|
||||
pub(super) struct ResolvedProvider<'a> {
|
||||
pub(super) model: &'a str,
|
||||
pub(super) provider: &'a str,
|
||||
pub(super) config: &'static dyn BaseAnthropicMessagesConfig,
|
||||
struct ResolvedProvider {
|
||||
model: String,
|
||||
provider: MessagesProvider,
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<ResolvedProvider<'a>, Error> {
|
||||
pub(super) struct ProviderMessagesRequest {
|
||||
pub(super) provider: MessagesProvider,
|
||||
pub(super) url: String,
|
||||
pub(super) body: AnthropicMessagesRequest,
|
||||
pub(super) environment: ValidatedEnvironment,
|
||||
pub(super) timeout: Option<Duration>,
|
||||
/// The caller's own credential, reported to the host beside the wire request.
|
||||
pub(super) api_key: Option<SecretValue>,
|
||||
}
|
||||
|
||||
pub(super) async fn prepare(
|
||||
call: MessagesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?;
|
||||
let secrets = secrets
|
||||
.resolve(resolved.provider.config().secret_names())
|
||||
.await?;
|
||||
prepare_provider_request(call, resolved, secrets.as_ref())
|
||||
}
|
||||
|
||||
fn resolve_provider(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
|
|
@ -44,82 +68,79 @@ pub(super) fn resolve_provider<'a>(
|
|||
"unable to resolve custom_llm_provider for messages request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let config = messages_provider_config(provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
|
||||
let provider = provider
|
||||
.parse()
|
||||
.map_err(|_| Error::InvalidProvider(provider.to_string()))?;
|
||||
Ok(ResolvedProvider {
|
||||
model,
|
||||
model: model.to_string(),
|
||||
provider,
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: MessagesRequest<'_>,
|
||||
resolved: ResolvedProvider<'_>,
|
||||
fn prepare_provider_request(
|
||||
call: MessagesCall,
|
||||
resolved: ResolvedProvider,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let ResolvedProvider {
|
||||
model,
|
||||
provider,
|
||||
config,
|
||||
} = resolved;
|
||||
let model = model.to_string();
|
||||
let ResolvedProvider { model, provider } = resolved;
|
||||
let MessagesCall {
|
||||
body,
|
||||
api_key,
|
||||
api_base,
|
||||
extra_headers,
|
||||
provider_specific_header,
|
||||
timeout,
|
||||
shaping,
|
||||
..
|
||||
} = call;
|
||||
let config = provider.config();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
|
||||
let typed_request: AnthropicMessagesRequest =
|
||||
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,
|
||||
AnthropicMessagesRequest { model, ..body },
|
||||
shaping.reasoning_auto_summary,
|
||||
)?;
|
||||
let trimmed =
|
||||
without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?;
|
||||
let trimmed = without_additional_drop_params(sanitized, &shaping.additional_drop_params)?;
|
||||
let transformed = config.transform_anthropic_messages_request(
|
||||
trimmed,
|
||||
&MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params),
|
||||
&MessagesTransformContext::new(shaping.capabilities, shaping.drop_params),
|
||||
)?;
|
||||
|
||||
let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider);
|
||||
let scoped =
|
||||
get_provider_specific_headers(provider_specific_header.as_ref(), provider.as_str());
|
||||
let forwarded = string_headers(Some(
|
||||
request
|
||||
.extra_headers
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(scoped)
|
||||
.collect(),
|
||||
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 validated = config.validate_environment(
|
||||
forwarded,
|
||||
api_key.as_deref(),
|
||||
&transformed.model,
|
||||
&env_lookup,
|
||||
)?;
|
||||
let environment = ValidatedEnvironment {
|
||||
headers: config.request_headers(
|
||||
with_default_headers(validated.headers, config.default_headers()),
|
||||
&transformed,
|
||||
),
|
||||
auth: validated.auth,
|
||||
};
|
||||
|
||||
let body = serde_json::to_value(transformed).map_err(|err| {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
))
|
||||
})?;
|
||||
|
||||
let url = config.get_complete_url(request.api_base, &model, &env_lookup)?;
|
||||
let url = if transformed.params.stream == Some(true) {
|
||||
config.complete_stream_url(api_base.as_deref(), &transformed.model, &env_lookup)?
|
||||
} else {
|
||||
config.get_complete_url(api_base.as_deref(), &transformed.model, &env_lookup)?
|
||||
};
|
||||
|
||||
Ok(ProviderMessagesRequest {
|
||||
provider: provider.to_string(),
|
||||
model,
|
||||
config,
|
||||
provider,
|
||||
url,
|
||||
body,
|
||||
upstream_headers: headers,
|
||||
timeout: request.timeout,
|
||||
body: transformed,
|
||||
environment,
|
||||
timeout,
|
||||
api_key: api_key.map(SecretValue::new),
|
||||
})
|
||||
}
|
||||
|
||||
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],
|
||||
|
|
@ -127,64 +148,59 @@ fn without_additional_drop_params(
|
|||
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
|
||||
let params = serde_json::to_value(request.params).map_err(invalid_request)?;
|
||||
let trimmed = paths
|
||||
.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()
|
||||
.fold(params, |params, path| delete_nested_value(params, path));
|
||||
Ok(AnthropicMessagesRequest {
|
||||
params: serde_json::from_value(trimmed).map_err(invalid_request)?,
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::base_llm::auth::resolve_auth;
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::*;
|
||||
use crate::messages::types::MessagesShaping;
|
||||
use crate::messages::MessagesShaping;
|
||||
|
||||
#[fixture]
|
||||
fn shaping() -> MessagesShaping {
|
||||
MessagesShaping::default()
|
||||
}
|
||||
|
||||
fn prepare(request: MessagesRequest<'_>) -> Result<ProviderMessagesRequest, Error> {
|
||||
prepare_with_secrets(request, &|_: &str| None)
|
||||
fn body(value: Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
fn prepare(call: MessagesCall) -> Result<ProviderMessagesRequest, Error> {
|
||||
prepare_with_secrets(call, &|_: &str| None)
|
||||
}
|
||||
|
||||
fn prepare_with_secrets(
|
||||
request: MessagesRequest<'_>,
|
||||
call: MessagesCall,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let resolved = resolve_provider(request.model, request.custom_llm_provider)?;
|
||||
prepare_provider_request(request, resolved, secrets)
|
||||
let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?;
|
||||
prepare_provider_request(call, resolved, secrets)
|
||||
}
|
||||
|
||||
/// The headers as they go on the wire, credential applied.
|
||||
fn wire_headers(prepared: &ProviderMessagesRequest) -> Vec<(String, String)> {
|
||||
tokio::runtime::Builder::new_current_thread()
|
||||
.build()
|
||||
.unwrap()
|
||||
.block_on(resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
prepared.environment.clone(),
|
||||
&|_| None,
|
||||
))
|
||||
.unwrap()
|
||||
.headers
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -221,12 +237,13 @@ mod tests {
|
|||
.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}),
|
||||
MessagesCall {
|
||||
body: body(
|
||||
json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
|
||||
),
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: None,
|
||||
|
|
@ -235,8 +252,8 @@ mod tests {
|
|||
&lookup,
|
||||
)
|
||||
.unwrap();
|
||||
let auth: Vec<(&str, &str)> = prepared
|
||||
.upstream_headers
|
||||
let headers = wire_headers(&prepared);
|
||||
let auth: Vec<(&str, &str)> = headers
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization"))
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
|
|
@ -247,48 +264,18 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
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"),
|
||||
fn prepared_body(fields: Value, shaping: MessagesShaping) -> Result<Value, Error> {
|
||||
prepare(MessagesCall {
|
||||
body: body(fields),
|
||||
api_key: Some("sk-test".into()),
|
||||
api_base: Some("https://anthropic.test".into()),
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
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)
|
||||
);
|
||||
.map(|prepared| serde_json::to_value(prepared.body).unwrap())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -380,20 +367,22 @@ mod tests {
|
|||
{"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()),
|
||||
let prepared = prepare(MessagesCall {
|
||||
body: body(
|
||||
json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
|
||||
),
|
||||
api_key: Some("sk-test".into()),
|
||||
api_base: Some("https://resource.services.ai.azure.com".into()),
|
||||
custom_llm_provider: custom_llm_provider.map(Into::into),
|
||||
extra_headers: Some(Map::from_iter([("x-priority".into(), json!("extra"))])),
|
||||
provider_specific_header: Some(configured),
|
||||
timeout: None,
|
||||
shaping,
|
||||
})
|
||||
.unwrap();
|
||||
let caller_headers: Vec<(&str, &str)> = prepared
|
||||
.upstream_headers
|
||||
.environment
|
||||
.headers
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped"))
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
|
|
|
|||
|
|
@ -1,52 +1,20 @@
|
|||
use std::{
|
||||
convert::Infallible,
|
||||
sync::{Arc, Mutex},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider;
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
host::{Demand, Host},
|
||||
machine::{CallMachine, HostChannel, MachineFault},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_http::{Client, ClientVariant, HttpClientConfig};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
|
||||
utils::ProviderSpecificHeaders,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::messages_provider_config,
|
||||
handler::{decode_response, network, provider_error, send},
|
||||
prepare::{prepare_provider_request, resolve_provider},
|
||||
types::{MessagesRequest, MessagesShaping},
|
||||
};
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
|
||||
/// The caller's request as the host projects it.
|
||||
pub struct MessagesCall {
|
||||
pub model: String,
|
||||
pub body: Map<String, Value>,
|
||||
pub api_key: Option<String>,
|
||||
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 {
|
||||
fn streams(&self) -> bool {
|
||||
self.body.get("stream").and_then(Value::as_bool) == Some(true)
|
||||
}
|
||||
}
|
||||
use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare};
|
||||
|
||||
pub enum MessagesOutput {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
|
|
@ -54,6 +22,11 @@ pub enum MessagesOutput {
|
|||
Streamed,
|
||||
}
|
||||
|
||||
/// The upstream response as the caller sees it at stream hand-off, before any chunk.
|
||||
pub struct MessagesStreamHead {
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
pub struct Messages;
|
||||
|
||||
impl Protocol for Messages {
|
||||
|
|
@ -62,7 +35,7 @@ impl Protocol for Messages {
|
|||
type Projection = MessagesCall;
|
||||
type Op = Infallible;
|
||||
type Chunk = Bytes;
|
||||
type StreamHead = ();
|
||||
type StreamHead = MessagesStreamHead;
|
||||
}
|
||||
|
||||
impl From<MachineFault> for Error {
|
||||
|
|
@ -77,19 +50,6 @@ impl From<MachineFault> for Error {
|
|||
pub type MessagesHost = HostChannel<Messages>;
|
||||
pub type MessagesMachine = CallMachine<Messages>;
|
||||
|
||||
/// Whether this route serves the request, decided before any callback runs so a host
|
||||
/// can still run its own path.
|
||||
pub fn supports(model: &str, custom_llm_provider: Option<&str>, stream: bool) -> bool {
|
||||
let provider = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.map(|resolved| resolved.custom_llm_provider)
|
||||
.or(custom_llm_provider);
|
||||
match provider {
|
||||
Some(ANTHROPIC_MESSAGES_PROVIDER) => true,
|
||||
Some(provider) => !stream && messages_provider_config(provider).is_some(),
|
||||
None => false,
|
||||
}
|
||||
}
|
||||
|
||||
/// The in-process host for a request already in hand. It answers projection once and
|
||||
/// observes nothing.
|
||||
pub struct LocalMessagesHost {
|
||||
|
|
@ -118,88 +78,43 @@ impl Host<Messages> for LocalMessagesHost {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
|
||||
CallMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
|
||||
pub fn messages_machine(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesMachine, litellm_http::Error> {
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let auth = resources.auth.clone();
|
||||
Ok(CallMachine::new(move |host| {
|
||||
Box::pin(drive(host, http, auth, secrets))
|
||||
}))
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
/// The call as its host sees it: projection first, then the same prepare and execute as
|
||||
/// [`super::messages`], with each chunk of a stream handed over as it arrives.
|
||||
async fn drive(
|
||||
host: MessagesHost,
|
||||
http: Client,
|
||||
auth: Arc<litellm_auth::AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let call = host.project().await?;
|
||||
let stream = call.streams();
|
||||
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"));
|
||||
}
|
||||
let context = RequestContext {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: request.provider.clone(),
|
||||
optional_params: Value::Object(
|
||||
call.body
|
||||
.iter()
|
||||
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages"))
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect(),
|
||||
),
|
||||
secret_fields: Vec::new(),
|
||||
api_key: call.api_key.clone().map(SecretValue::new),
|
||||
};
|
||||
let wire = host
|
||||
.before_send(
|
||||
WireRequest {
|
||||
url: request.url,
|
||||
headers: request.upstream_headers,
|
||||
body: request.body,
|
||||
},
|
||||
context,
|
||||
)
|
||||
.await?;
|
||||
let response = send(&wire.url, &wire.headers, &wire.body, request.timeout).await?;
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
if stream {
|
||||
return relay(&host, response).await;
|
||||
}
|
||||
let text = response.text().await.map_err(network)?;
|
||||
host.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
decode_response(request.config, &request.model, &text)
|
||||
.map(|message| MessagesOutput::Message(Box::new(message)))
|
||||
}
|
||||
|
||||
/// Hands each upstream chunk to the caller as it arrives. A caller that stops reading
|
||||
/// ends the upstream read, and the call completes with what it delivered.
|
||||
async fn relay(
|
||||
host: &MessagesHost,
|
||||
mut response: reqwest::Response,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
if host.open(()).await? == Demand::Detached {
|
||||
return Ok(MessagesOutput::Streamed);
|
||||
}
|
||||
while let Some(chunk) = response.chunk().await.map_err(network)? {
|
||||
if host.deliver(chunk).await? == Demand::Detached {
|
||||
break;
|
||||
let request = prepare(call, secrets.as_ref()).await?;
|
||||
match execute(&http, &auth, request, &host).await? {
|
||||
MessagesResponse::Message(message) => Ok(MessagesOutput::Message(message)),
|
||||
MessagesResponse::Stream {
|
||||
headers,
|
||||
mut chunks,
|
||||
} => {
|
||||
if host.open(MessagesStreamHead { headers }).await? == Demand::Detached {
|
||||
return Ok(MessagesOutput::Streamed);
|
||||
}
|
||||
while let Some(chunk) = chunks.try_next().await? {
|
||||
if host.deliver(chunk).await? == Demand::Detached {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(MessagesOutput::Streamed)
|
||||
}
|
||||
}
|
||||
Ok(MessagesOutput::Streamed)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,13 +1,46 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_llms::{
|
||||
anthropic::common_utils::AnthropicModelCapabilities,
|
||||
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
|
||||
use bytes::Bytes;
|
||||
use futures_util::stream::BoxStream;
|
||||
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
},
|
||||
utils::ProviderSpecificHeaders,
|
||||
};
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::Error;
|
||||
|
||||
pub struct MessagesCall {
|
||||
pub body: AnthropicMessagesRequest,
|
||||
pub api_key: Option<String>,
|
||||
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,
|
||||
}
|
||||
|
||||
pub fn messages_body(body: Map<String, Value>) -> Result<AnthropicMessagesRequest, Error> {
|
||||
serde_json::from_value(Value::Object(body)).map_err(invalid_request)
|
||||
}
|
||||
|
||||
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
}
|
||||
|
||||
pub enum MessagesResponse {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
Stream {
|
||||
headers: Vec<(String, String)>,
|
||||
chunks: BoxStream<'static, Result<Bytes, Error>>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct MessagesShaping {
|
||||
#[serde(default)]
|
||||
|
|
@ -20,33 +53,11 @@ pub struct MessagesShaping {
|
|||
pub additional_drop_params: Vec<String>,
|
||||
}
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub body: Value,
|
||||
pub api_key: Option<&'a str>,
|
||||
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 {
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
pub config: &'static dyn BaseAnthropicMessagesConfig,
|
||||
pub url: String,
|
||||
pub body: Value,
|
||||
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 serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
|
|
|
|||
|
|
@ -244,160 +244,3 @@ 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)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,212 +7,3 @@ pub mod provider_config;
|
|||
pub mod route;
|
||||
pub mod types;
|
||||
pub mod wire;
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) mod test_support {
|
||||
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())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -70,20 +70,206 @@ pub(crate) fn prepare_request(
|
|||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn prepare_request_for_test(request: ResolvedOcrRequest) -> PreparedOcrRequest {
|
||||
prepare_request(
|
||||
request,
|
||||
true,
|
||||
&OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new()),
|
||||
std::sync::Arc::new(litellm_core_utils::settings::ProcessEnvironment),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_core_utils::call_arguments::{CallArguments, compose_body, parse_options};
|
||||
use serde_json::json;
|
||||
use litellm_host::event::WireRequest;
|
||||
use litellm_llms::{
|
||||
base_llm::ocr::{
|
||||
error::Error,
|
||||
handler::{CallHooks, OcrClient},
|
||||
transformation::{BaseOcrConfig, OcrResponseFormat},
|
||||
},
|
||||
cohere::ocr::transformation::CohereParseConfig,
|
||||
mistral::ocr::transformation::MistralOcrConfig,
|
||||
vertex_ai::ocr::transformation::VertexAiOcrConfig,
|
||||
};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
use crate::ocr::{
|
||||
document::prepare_document,
|
||||
types::LiteLLMOcrRequest,
|
||||
wire::{OcrWireRequest, decode_request},
|
||||
};
|
||||
|
||||
/// Stands in for a host with no hooks registered.
|
||||
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(()) })
|
||||
}
|
||||
}
|
||||
|
||||
fn client() -> OcrClient {
|
||||
OcrClient::for_test(
|
||||
litellm_http::Client::plain_for_test(),
|
||||
litellm_http::Client::no_redirect_for_test(),
|
||||
)
|
||||
}
|
||||
|
||||
fn request(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()
|
||||
}
|
||||
|
||||
fn prepared(request: LiteLLMOcrRequest) -> PreparedOcrRequest {
|
||||
prepare_request(
|
||||
request.map_document(prepare_document).unwrap(),
|
||||
true,
|
||||
&client(),
|
||||
std::sync::Arc::new(litellm_core_utils::settings::ProcessEnvironment),
|
||||
)
|
||||
}
|
||||
|
||||
fn image(url: &str) -> Value {
|
||||
json!({"type": "image_url", "image_url": url})
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cohere_body_keeps_native_document_fields_and_untyped_overrides() {
|
||||
let request = request(
|
||||
"cohere/parse",
|
||||
"https://example.com",
|
||||
image("https://example.com/original.png"),
|
||||
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 http = CohereParseConfig
|
||||
.prepare_request(&prepared(request), &client(), &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 = request(
|
||||
"cohere/parse",
|
||||
"https://example.com",
|
||||
image("https://example.com/a.png"),
|
||||
json!({"output_format": null, "req_format": null}),
|
||||
);
|
||||
assert_eq!(
|
||||
request.response_format().unwrap(),
|
||||
OcrResponseFormat::Litellm
|
||||
);
|
||||
|
||||
let http = CohereParseConfig
|
||||
.prepare_request(&prepared(request), &client(), &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());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn direct_and_vertex_mistral_build_the_same_request_and_share_normalization() {
|
||||
let options = json!({
|
||||
"pages": [0, 2],
|
||||
"include_image_base64": true,
|
||||
"vertex_project": "project-1",
|
||||
"vertex_location": "us-central1",
|
||||
"unknown": "preserved"
|
||||
});
|
||||
let document =
|
||||
json!({"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"});
|
||||
let direct = prepared(request(
|
||||
"mistral/mistral-ocr-maas",
|
||||
"https://mistral.test",
|
||||
document.clone(),
|
||||
options.clone(),
|
||||
));
|
||||
let vertex = prepared(request(
|
||||
"vertex_ai/mistral-ocr-maas",
|
||||
"https://vertex.test",
|
||||
document,
|
||||
options,
|
||||
));
|
||||
|
||||
let direct_http = MistralOcrConfig
|
||||
.prepare_request(&direct, &client(), &NoHooks)
|
||||
.await
|
||||
.unwrap();
|
||||
let vertex_http = VertexAiOcrConfig
|
||||
.prepare_request(&vertex, &client(), &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": "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, OcrResponseFormat::Litellm)
|
||||
.unwrap()
|
||||
.into_json();
|
||||
let vertex_response = VertexAiOcrConfig
|
||||
.transform_ocr_response(&vertex.model, &payload, 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");
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct KnownParams {
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,30 +1,20 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::RequestAuth;
|
||||
use litellm_auth_aws::SigV4Signer;
|
||||
use litellm_http::outbound::OutboundRequest;
|
||||
use serde_json::{Map, Value};
|
||||
use litellm_llms::base_llm::auth::Authenticated;
|
||||
use serde_json::Value;
|
||||
|
||||
/// Header credentials are already in `headers`; SigV4 is applied here, over the
|
||||
/// bytes that are sent.
|
||||
pub(crate) async fn outbound_request<E>(
|
||||
auth: &RequestAuth,
|
||||
pub(crate) fn outbound_request(
|
||||
authenticated: Authenticated,
|
||||
url: String,
|
||||
headers: Vec<(String, String)>,
|
||||
body: &Value,
|
||||
timeout: Option<Duration>,
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Result<OutboundRequest, E>
|
||||
where
|
||||
E: From<litellm_http::Error> + From<litellm_auth_aws::Error>,
|
||||
{
|
||||
let RequestAuth::AwsSigV4 { region, service } = auth else {
|
||||
return Ok(OutboundRequest::json(url, headers, body, timeout)?);
|
||||
};
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let signer =
|
||||
SigV4Signer::resolve(region.clone(), service, optional_params, &env_lookup).await?;
|
||||
Ok(OutboundRequest::signed_json(
|
||||
url, headers, body, timeout, &signer,
|
||||
)?)
|
||||
) -> Result<OutboundRequest, litellm_http::Error> {
|
||||
let Authenticated { headers, signer } = authenticated;
|
||||
match signer {
|
||||
None => OutboundRequest::json(url, headers, body, timeout),
|
||||
Some(signer) => OutboundRequest::signed_json(url, headers, body, timeout, &signer),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
38
litellm-rust/crates/core/src/resources.rs
Normal file
38
litellm-rust/crates/core/src/resources.rs
Normal file
|
|
@ -0,0 +1,38 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_http::{HttpClientConfig, HttpClientPool, media::UrlPolicy};
|
||||
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CoreResources {
|
||||
pub pool: Arc<HttpClientPool>,
|
||||
pub auth: Arc<AuthServices>,
|
||||
}
|
||||
|
||||
impl CoreResources {
|
||||
pub fn new(pool: Arc<HttpClientPool>) -> Self {
|
||||
Self {
|
||||
pool,
|
||||
auth: Arc::new(AuthServices::default()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn ocr_client(
|
||||
&self,
|
||||
config: &HttpClientConfig,
|
||||
url_policy: UrlPolicy,
|
||||
settings: OcrSettings,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<OcrClient, litellm_http::Error> {
|
||||
OcrClient::new(
|
||||
&self.pool,
|
||||
config,
|
||||
url_policy,
|
||||
self.auth.clone(),
|
||||
settings,
|
||||
secrets,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,17 +0,0 @@
|
|||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("invalid provider: {0}")]
|
||||
InvalidProvider(String),
|
||||
#[error("invalid request: {0}")]
|
||||
InvalidRequest(String),
|
||||
#[error("invalid response: {0}")]
|
||||
InvalidResponse(String),
|
||||
#[error("routing error: {0}")]
|
||||
Routing(String),
|
||||
#[error(transparent)]
|
||||
Auth(#[from] litellm_auth::Error),
|
||||
#[error(transparent)]
|
||||
Transport(#[from] litellm_http::transport::Error),
|
||||
#[error(transparent)]
|
||||
Headers(#[from] litellm_http::request::HeaderError),
|
||||
}
|
||||
|
|
@ -1,3 +1,2 @@
|
|||
mod error;
|
||||
pub use error::Error;
|
||||
pub use crate::error::RouteError as Error;
|
||||
pub mod websocket;
|
||||
|
|
|
|||
|
|
@ -1,50 +1,254 @@
|
|||
use std::{
|
||||
io::{Read, Write},
|
||||
net::TcpListener,
|
||||
thread,
|
||||
use litellm_core::audio_transcription::{
|
||||
Error, audio_transcription, types::AudioTranscriptionRequest,
|
||||
};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
|
||||
use litellm_core::audio_transcription::{audio_transcription, types::AudioTranscriptionRequest};
|
||||
use serde_json::{Map, json};
|
||||
mod support;
|
||||
use support::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn bedrock_request_is_signed_and_contains_audio() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").expect("listener");
|
||||
let address = listener.local_addr().expect("address");
|
||||
let server = thread::spawn(move || {
|
||||
let (mut stream, _) = listener.accept().expect("connection");
|
||||
let mut request = Vec::new();
|
||||
let mut buffer = [0_u8; 16_384];
|
||||
let count = stream.read(&mut buffer).expect("request");
|
||||
request.extend_from_slice(&buffer[..count]);
|
||||
let request = String::from_utf8_lossy(&request);
|
||||
assert!(request.contains("POST /model/mistral.voxtral-mini-3b-2507/converse"));
|
||||
assert!(request.contains("authorization: AWS4-HMAC-SHA256"));
|
||||
assert!(request.contains("x-amz-date:"));
|
||||
assert!(request.contains("\"bytes\":\"AQI=\""));
|
||||
assert!(request.contains("Transcribe the audio. Respond with only the transcript."));
|
||||
let response = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 53\r\nConnection: close\r\n\r\n{\"output\":{\"message\":{\"content\":[{\"text\":\"hello\"}]}}}";
|
||||
stream.write_all(response).expect("response");
|
||||
});
|
||||
const MODEL: &str = "mistral.voxtral-mini-3b-2507";
|
||||
|
||||
let optional_params = Map::from_iter([
|
||||
async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
audio_transcription(&support::resources(), &http_config(), request).await
|
||||
}
|
||||
|
||||
fn transcript_response(text: &str) -> ResponseTemplate {
|
||||
json_response(json!({"output": {"message": {"content": [{"text": text}]}}}))
|
||||
}
|
||||
|
||||
fn aws_params(region: &str) -> Map<String, Value> {
|
||||
Map::from_iter([
|
||||
("aws_access_key_id".to_string(), json!("access-key")),
|
||||
("aws_secret_access_key".to_string(), json!("secret-key")),
|
||||
("aws_region_name".to_string(), json!("us-east-1")),
|
||||
]);
|
||||
let api_base = format!("http://{address}");
|
||||
let response = audio_transcription(AudioTranscriptionRequest {
|
||||
model: "mistral.voxtral-mini-3b-2507",
|
||||
("aws_region_name".to_string(), json!(region)),
|
||||
])
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn request() -> AudioTranscriptionRequest<'static> {
|
||||
AudioTranscriptionRequest {
|
||||
model: MODEL,
|
||||
audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}),
|
||||
api_key: None,
|
||||
api_base: Some(&api_base),
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("bedrock"),
|
||||
extra_headers: None,
|
||||
optional_params,
|
||||
optional_params: aws_params("us-east-1"),
|
||||
timeout: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::us_east_1("us-east-1")]
|
||||
#[case::eu_west_1("eu-west-1")]
|
||||
#[tokio::test]
|
||||
async fn bedrock_converse_request_is_signed_for_the_requested_region(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
#[case] region: &str,
|
||||
) {
|
||||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let response = transcribe(AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
optional_params: aws_params(region),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect("transcription");
|
||||
|
||||
assert_eq!(response, json!({"text": "hello"}));
|
||||
server.join().expect("server");
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.method.as_str(), "POST");
|
||||
assert_eq!(sent.url.path(), format!("/model/{MODEL}/converse"));
|
||||
let authorization = sent.header("authorization").expect("request is signed");
|
||||
assert!(
|
||||
authorization.starts_with("AWS4-HMAC-SHA256 Credential=access-key/"),
|
||||
"{authorization}"
|
||||
);
|
||||
assert!(
|
||||
authorization.contains(&format!("/{region}/bedrock/aws4_request")),
|
||||
"{authorization}"
|
||||
);
|
||||
assert!(sent.header("x-amz-date").is_some());
|
||||
assert!(!sent.body_text().contains("secret-key"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_provider_can_come_from_the_model_prefix(request: AudioTranscriptionRequest<'static>) {
|
||||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
let model = format!("bedrock/{MODEL}");
|
||||
|
||||
transcribe(AudioTranscriptionRequest {
|
||||
model: &model,
|
||||
custom_llm_provider: None,
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect("transcription");
|
||||
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.url.path(),
|
||||
format!("/model/{MODEL}/converse")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn audio_and_transcription_params_reach_the_converse_body(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
#[values("wav", "mp3", "flac", "ogg")] format: &str,
|
||||
) {
|
||||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
let optional_params = aws_params("us-east-1")
|
||||
.into_iter()
|
||||
.chain([
|
||||
("language".to_string(), json!("fr")),
|
||||
("temperature".to_string(), json!(0.2)),
|
||||
])
|
||||
.collect();
|
||||
|
||||
transcribe(AudioTranscriptionRequest {
|
||||
audio: json!({"data": "AQI=", "format": format}),
|
||||
api_base: Some(&base),
|
||||
optional_params,
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect("transcription");
|
||||
|
||||
let body = only_request(&upstream).await.json();
|
||||
let content = &body["messages"][0]["content"];
|
||||
assert_eq!(
|
||||
content[0],
|
||||
json!({"audio": {"format": format, "source": {"bytes": "AQI="}}})
|
||||
);
|
||||
let instruction = content[1]["text"].as_str().expect("instruction text");
|
||||
assert!(instruction.contains("fr"), "{instruction}");
|
||||
assert_eq!(body["inferenceConfig"]["temperature"], 0.2);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::unknown_format(json!({"data": "AQI=", "format": "aac"}))]
|
||||
#[case::missing_data(json!({"format": "wav"}))]
|
||||
#[case::not_an_object(json!("AQI="))]
|
||||
#[tokio::test]
|
||||
async fn invalid_audio_is_rejected_before_sending(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
#[case] audio: Value,
|
||||
) {
|
||||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
audio,
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("invalid audio is rejected");
|
||||
|
||||
assert!(
|
||||
matches!(
|
||||
error,
|
||||
Error::InvalidRequest(_) | Error::MissingField(_) | Error::InvalidType { .. }
|
||||
),
|
||||
"{error:?}"
|
||||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::unknown_provider(MODEL, Some("openai"), "openai")]
|
||||
#[case::unresolvable_model(
|
||||
"no-such-model",
|
||||
None,
|
||||
"unable to resolve custom_llm_provider for audio transcription request"
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn unsupported_providers_are_rejected_before_sending(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
#[case] model: &'static str,
|
||||
#[case] provider: Option<&'static str>,
|
||||
#[case] reported: &str,
|
||||
) {
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
api_base: Some(UNREACHABLE_BASE),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("unsupported provider errors");
|
||||
|
||||
assert_eq!(error, Error::InvalidProvider(reported.into()));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_non_string_extra_header_is_rejected(request: AudioTranscriptionRequest<'static>) {
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
extra_headers: Some(Map::from_iter([("x-count".to_string(), json!(3))])),
|
||||
api_base: Some(UNREACHABLE_BASE),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("a non-string header is rejected");
|
||||
|
||||
assert!(matches!(error, Error::Headers(_)), "{error:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::throttled(429)]
|
||||
#[case::server_error(500)]
|
||||
#[tokio::test]
|
||||
async fn an_upstream_error_keeps_its_status_and_body(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
#[case] status: u16,
|
||||
) {
|
||||
let upstream =
|
||||
upstream([ResponseTemplate::new(status).set_body_string("upstream said no")]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status,
|
||||
body: "upstream said no".into()
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::not_json(ResponseTemplate::new(200).set_body_string("not json"))]
|
||||
#[case::no_output(json_response(json!({"unexpected": true})))]
|
||||
#[tokio::test]
|
||||
async fn an_unreadable_success_body_is_an_invalid_response(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
#[case] response: ResponseTemplate,
|
||||
) {
|
||||
let upstream = upstream([response]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("an unreadable body fails");
|
||||
|
||||
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
|
||||
}
|
||||
|
|
|
|||
325
litellm-rust/crates/core/tests/chat_completions.rs
Normal file
325
litellm-rust/crates/core/tests/chat_completions.rs
Normal file
|
|
@ -0,0 +1,325 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_core::chat_completions::{
|
||||
Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest,
|
||||
};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
|
||||
mod support;
|
||||
use support::*;
|
||||
|
||||
const ANTHROPIC_MESSAGE: &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}}"#;
|
||||
|
||||
async fn complete(request: ChatCompletionsRequest<'_>) -> Result<ChatCompletionsResponse, Error> {
|
||||
chat_completions(&support::resources(), &http_config(), request).await
|
||||
}
|
||||
|
||||
fn object(value: Value) -> Map<String, Value> {
|
||||
let Value::Object(map) = value else {
|
||||
panic!("expected a json object, got {value}");
|
||||
};
|
||||
map
|
||||
}
|
||||
|
||||
fn anthropic_response(body: &str) -> ResponseTemplate {
|
||||
ResponseTemplate::new(200).set_body_raw(body, "application/json")
|
||||
}
|
||||
|
||||
fn hi() -> Value {
|
||||
json!([{"role": "user", "content": "hi"}])
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn request() -> ChatCompletionsRequest<'static> {
|
||||
ChatCompletionsRequest {
|
||||
model: "anthropic/claude-sonnet-4-5",
|
||||
messages: hi(),
|
||||
optional_params: object(json!({"max_tokens": 16})),
|
||||
api_key: Some("sk-test"),
|
||||
api_base: None,
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: Some(Duration::from_secs(10)),
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn anthropic_round_trip_translates_the_conversation_and_normalizes_the_response(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let response = complete(ChatCompletionsRequest {
|
||||
messages: json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]),
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect("call succeeds");
|
||||
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.url.path(), "/v1/messages");
|
||||
assert_eq!(sent.header_values("x-api-key"), ["sk-test"]);
|
||||
let body = sent.json();
|
||||
assert_eq!(body["model"], "claude-sonnet-4-5");
|
||||
assert_eq!(
|
||||
body["messages"],
|
||||
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}])
|
||||
);
|
||||
assert_eq!(
|
||||
body["system"],
|
||||
json!([{"type": "text", "text": "be terse"}])
|
||||
);
|
||||
assert_eq!(body["max_tokens"], 16);
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("hello")
|
||||
);
|
||||
assert_eq!(response.usage.total_tokens, 15);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_deployment_key_replaces_a_caller_supplied_x_api_key(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
extra_headers: Some(object(
|
||||
json!({"x-api-key": "caller-key", "x-trace": "kept"}),
|
||||
)),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect("call succeeds");
|
||||
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.header_values("x-api-key"), ["sk-test"]);
|
||||
assert_eq!(sent.header("x-trace"), Some("kept"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsRequest<'static>) {
|
||||
let upstream = upstream([json_response(json!({
|
||||
"output": {"message": {"role": "assistant", "content": [{"text": "hello"}]}},
|
||||
"stopReason": "end_turn",
|
||||
"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}
|
||||
}))])
|
||||
.await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let response = complete(ChatCompletionsRequest {
|
||||
model: "bedrock/anthropic.claude-sonnet-4-5",
|
||||
optional_params: object(json!({
|
||||
"aws_access_key_id": "access-key",
|
||||
"aws_secret_access_key": "secret-key",
|
||||
"aws_region_name": "eu-west-1"
|
||||
})),
|
||||
api_key: None,
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect("call succeeds");
|
||||
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(
|
||||
sent.url.path(),
|
||||
"/model/anthropic.claude-sonnet-4-5/converse"
|
||||
);
|
||||
let authorization = sent.header("authorization").expect("request is signed");
|
||||
assert!(
|
||||
authorization.contains("/eu-west-1/bedrock/aws4_request"),
|
||||
"{authorization}"
|
||||
);
|
||||
assert_eq!(
|
||||
sent.json()["messages"],
|
||||
json!([{"role": "user", "content": [{"text": "hi"}]}])
|
||||
);
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("hello")
|
||||
);
|
||||
assert_eq!(response.usage.total_tokens, 15);
|
||||
}
|
||||
|
||||
/// The provider already answered and billed these, so the host must not retry them on
|
||||
/// its own path: they surface as `InvalidResponse`, never as a pre-send decline.
|
||||
#[rstest]
|
||||
#[case::missing_usage(
|
||||
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#
|
||||
)]
|
||||
#[case::tool_use_block(r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#)]
|
||||
#[case::not_json("not json")]
|
||||
#[tokio::test]
|
||||
async fn a_response_it_cannot_normalize_is_reported_as_already_sent(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
#[case] body: &str,
|
||||
) {
|
||||
let upstream = upstream([anthropic_response(body)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("response cannot be normalized");
|
||||
|
||||
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::rate_limited(429)]
|
||||
#[case::server_error(500)]
|
||||
#[tokio::test]
|
||||
async fn an_upstream_error_status_keeps_its_code_and_body(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
#[case] status: u16,
|
||||
) {
|
||||
let upstream = upstream([ResponseTemplate::new(status).set_body_string("slow down")]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("upstream rejects");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: "slow down".into()
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
/// Nothing was sent, so nothing was billed and the host can still serve the request.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_connection_that_is_never_established_declines_instead_of_failing(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
api_base: Some(UNREACHABLE_BASE),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("nothing is listening");
|
||||
|
||||
assert!(
|
||||
matches!(error, Error::Transport(TransportError::Connect(_))),
|
||||
"{error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_timeout_after_sending_is_not_a_pre_send_decline(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let upstream =
|
||||
upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
timeout: Some(Duration::from_millis(100)),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("the call times out");
|
||||
|
||||
assert!(
|
||||
matches!(error, Error::Transport(TransportError::Network(_))),
|
||||
"{error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::accepted("anthropic/claude-sonnet-4-5", None, hi(), json!({"max_tokens": 16}), None)]
|
||||
#[case::accepted_bedrock("bedrock/anthropic.claude-sonnet-4-5", None, hi(), json!({}), None)]
|
||||
#[case::unknown_provider(
|
||||
"gpt-4o",
|
||||
Some("openai"),
|
||||
hi(),
|
||||
json!({}),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
)]
|
||||
#[case::unreadable_messages(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("hi"),
|
||||
json!({}),
|
||||
Some("unreadable message list")
|
||||
)]
|
||||
#[case::empty_messages("anthropic/claude-sonnet-4-5", None, json!([]), json!({}), Some("empty message list"))]
|
||||
#[case::streaming(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
hi(),
|
||||
json!({"stream": true}),
|
||||
Some("streaming")
|
||||
)]
|
||||
#[case::unrecognized_param(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
hi(),
|
||||
json!({"not_a_param": 1}),
|
||||
Some("unrecognized request parameter")
|
||||
)]
|
||||
#[case::opens_on_assistant_turn(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "assistant", "content": "hi"}]),
|
||||
json!({}),
|
||||
Some("conversation does not open on a user turn")
|
||||
)]
|
||||
fn decline_reason_names_why_the_core_would_not_serve_the_request(
|
||||
#[case] model: &str,
|
||||
#[case] provider: Option<&str>,
|
||||
#[case] messages: Value,
|
||||
#[case] params: Value,
|
||||
#[case] reason: Option<&str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
chat_completions_decline_reason(model, provider, messages, &object(params)),
|
||||
reason
|
||||
);
|
||||
}
|
||||
|
||||
/// A request the decline check accepts must not be declined by the call itself.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_declined_request_fails_the_call_before_sending(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
optional_params: object(json!({"stream": true})),
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("streaming is declined");
|
||||
|
||||
assert_eq!(error, Error::Unsupported("streaming"));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
|
@ -1,471 +0,0 @@
|
|||
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},
|
||||
};
|
||||
|
||||
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();
|
||||
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 write_response(body: &str) -> String {
|
||||
format!(
|
||||
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
|
||||
body.len(),
|
||||
body
|
||||
)
|
||||
}
|
||||
|
||||
#[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");
|
||||
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":[{"type":"text","text":"hi"}],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":2}}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
|
||||
let response = messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 1024,
|
||||
"messages": [{
|
||||
"role": "user",
|
||||
"content": [{
|
||||
"type": "text",
|
||||
"text": "hi",
|
||||
"cache_control": {"type": "ephemeral", "scope": "global"}
|
||||
}]
|
||||
}]
|
||||
}),
|
||||
api_key: Some("sk-azure"),
|
||||
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");
|
||||
|
||||
assert_eq!(response.content[0]["text"], "hi");
|
||||
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
let (head, body) = request.split_once("\r\n\r\n").expect("has body");
|
||||
assert!(head.starts_with("POST /anthropic/v1/messages "), "{head}");
|
||||
let head_lower = head.to_ascii_lowercase();
|
||||
assert!(head_lower.contains("x-api-key: sk-azure"), "{head}");
|
||||
assert!(
|
||||
head_lower.contains("anthropic-version: 2023-06-01"),
|
||||
"{head}"
|
||||
);
|
||||
assert!(
|
||||
head_lower.contains("content-type: application/json"),
|
||||
"{head}"
|
||||
);
|
||||
|
||||
let sent_body: Value = serde_json::from_str(body).expect("body is json");
|
||||
assert_eq!(
|
||||
sent_body["messages"][0]["content"][0]["cache_control"],
|
||||
json!({"type": "ephemeral"})
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_round_trip_builds_native_anthropic_request() {
|
||||
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":[{"type":"text","text":"hi"}],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":2}}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
|
||||
let response = messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 1024,
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}),
|
||||
api_key: Some("sk-ant"),
|
||||
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");
|
||||
|
||||
assert_eq!(response.content[0]["text"], "hi");
|
||||
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
let (head, _) = request.split_once("\r\n\r\n").expect("has body");
|
||||
assert!(head.starts_with("POST /v1/messages "), "{head}");
|
||||
let head_lower = head.to_ascii_lowercase();
|
||||
assert!(head_lower.contains("x-api-key: sk-ant"), "{head}");
|
||||
assert!(
|
||||
head_lower.contains("anthropic-version: 2023-06-01"),
|
||||
"{head}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() {
|
||||
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_2","type":"message","role":"assistant","content":[],"model":"m"}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
|
||||
let mut headers = Map::new();
|
||||
headers.insert(
|
||||
"x-api-key".to_string(),
|
||||
Value::String("from-python".to_string()),
|
||||
);
|
||||
headers.insert(
|
||||
"anthropic-beta".to_string(),
|
||||
Value::String("token-efficient-tools-2025-02-19".to_string()),
|
||||
);
|
||||
|
||||
messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
||||
api_key: Some("rust-fallback-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("messages request succeeds");
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
let head = request
|
||||
.split_once("\r\n\r\n")
|
||||
.expect("has body")
|
||||
.0
|
||||
.to_ascii_lowercase();
|
||||
let api_key_count = head
|
||||
.lines()
|
||||
.filter(|line| line.starts_with("x-api-key:"))
|
||||
.count();
|
||||
assert_eq!(api_key_count, 1, "{head}");
|
||||
assert!(head.contains("x-api-key: from-python"), "{head}");
|
||||
assert!(
|
||||
head.contains("anthropic-beta: token-efficient-tools-2025-02-19"),
|
||||
"{head}"
|
||||
);
|
||||
assert!(!head.contains("rust-fallback-key"), "{head}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
|
||||
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_3","type":"message","role":"assistant","content":[],"model":"m"}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
|
||||
let mut headers = Map::new();
|
||||
headers.insert(
|
||||
"Authorization".to_string(),
|
||||
Value::String("Bearer entra-token".to_string()),
|
||||
);
|
||||
|
||||
messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
||||
api_key: None,
|
||||
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");
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
let head = request
|
||||
.split_once("\r\n\r\n")
|
||||
.expect("has body")
|
||||
.0
|
||||
.to_ascii_lowercase();
|
||||
assert!(head.contains("authorization: bearer entra-token"), "{head}");
|
||||
assert!(!head.contains("x-api-key"), "{head}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_requires_auth_when_no_key_and_no_header() {
|
||||
let err = messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
||||
api_key: None,
|
||||
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");
|
||||
|
||||
assert!(matches!(err, Error::Auth(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_ignores_malformed_authorization_and_uses_api_key() {
|
||||
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_4","type":"message","role":"assistant","content":[],"model":"m"}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
|
||||
let mut headers = Map::new();
|
||||
headers.insert(
|
||||
"Authorization".to_string(),
|
||||
Value::String("Bearer ".to_string()),
|
||||
);
|
||||
|
||||
messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
||||
api_key: Some("sk-azure"),
|
||||
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");
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
let head = request
|
||||
.split_once("\r\n\r\n")
|
||||
.expect("has body")
|
||||
.0
|
||||
.to_ascii_lowercase();
|
||||
assert!(head.contains("x-api-key: sk-azure"), "{head}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_maps_provider_error_status_to_http_error() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
let addr = listener.local_addr().expect("addr");
|
||||
|
||||
tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
||||
let _ = read_http_request(&mut socket).await;
|
||||
let body = "unauthorized";
|
||||
let response = format!(
|
||||
"HTTP/1.1 401 Unauthorized\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
|
||||
body.len(),
|
||||
body
|
||||
);
|
||||
socket
|
||||
.write_all(response.as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
});
|
||||
|
||||
let err = messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
||||
api_key: Some("sk-azure"),
|
||||
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");
|
||||
|
||||
assert!(matches!(
|
||||
err,
|
||||
Error::Transport(litellm_http::transport::Error::Http { status: 401, .. })
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_rejects_unsupported_provider() {
|
||||
let err = messages(MessagesRequest {
|
||||
model: "claude-3-5-sonnet",
|
||||
body: json!({"model": "claude-3-5-sonnet", "max_tokens": 8, "messages": []}),
|
||||
api_key: Some("sk"),
|
||||
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");
|
||||
|
||||
assert!(matches!(err, Error::InvalidProvider(provider) if provider == "openai"));
|
||||
}
|
||||
203
litellm-rust/crates/core/tests/messages/host.rs
Normal file
203
litellm-rust/crates/core/tests/messages/host.rs
Normal file
|
|
@ -0,0 +1,203 @@
|
|||
use std::{convert::Infallible, sync::Mutex};
|
||||
|
||||
use litellm_core::messages::route::Messages;
|
||||
use litellm_host::{
|
||||
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
|
||||
host::Host,
|
||||
};
|
||||
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
type Rewrite = Box<dyn Fn(WireRequest) -> Result<WireRequest, Error> + Send + Sync>;
|
||||
|
||||
/// Projects like `LocalMessagesHost`, answers `before_send` through `rewrite`, and keeps
|
||||
/// every event the driver emits.
|
||||
struct RecordingHost {
|
||||
call: LocalMessagesHost,
|
||||
rewrite: Rewrite,
|
||||
events: Mutex<Vec<CallEvent>>,
|
||||
optional_params: Mutex<Vec<Value>>,
|
||||
}
|
||||
|
||||
impl RecordingHost {
|
||||
fn new(call: MessagesCall, rewrite: Rewrite) -> Self {
|
||||
Self {
|
||||
call: LocalMessagesHost::new(call),
|
||||
rewrite,
|
||||
events: Mutex::new(Vec::new()),
|
||||
optional_params: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn passthrough(call: MessagesCall) -> Self {
|
||||
Self::new(call, Box::new(Ok))
|
||||
}
|
||||
|
||||
fn raw_responses(&self) -> Vec<String> {
|
||||
self.events
|
||||
.lock()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
CallEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
|
||||
Some(raw.body.clone())
|
||||
}
|
||||
_ => None,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for RecordingHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.project().await
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
}
|
||||
|
||||
async fn before_send(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
self.optional_params
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(context.optional_params.clone());
|
||||
(self.rewrite)(wire)
|
||||
}
|
||||
|
||||
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
|
||||
self.events.lock().unwrap().push(event.clone());
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
}
|
||||
|
||||
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {
|
||||
MessagesCall {
|
||||
api_key: Some("sk-ant".into()),
|
||||
api_base: Some(api_base),
|
||||
..call
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn what_before_send_returns_is_what_the_provider_receives(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let host = RecordingHost::new(
|
||||
authenticated(call, upstream.uri()),
|
||||
Box::new(|wire| {
|
||||
let mut body = wire.body;
|
||||
body["system"] = json!("added by the host");
|
||||
Ok(WireRequest {
|
||||
headers: wire
|
||||
.headers
|
||||
.into_iter()
|
||||
.chain([("x-host".to_string(), "seen".to_string())])
|
||||
.collect(),
|
||||
body,
|
||||
..wire
|
||||
})
|
||||
}),
|
||||
);
|
||||
|
||||
run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(request.json()["system"], "added by the host");
|
||||
assert_eq!(request.header("x-host"), Some("seen"));
|
||||
assert_eq!(request.header("x-api-key"), Some("sk-ant"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_before_send_failure_never_sends(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let host = RecordingHost::new(
|
||||
authenticated(call, upstream.uri()),
|
||||
Box::new(|_| Err(Error::InvalidRequest("vetoed by the host".into()))),
|
||||
);
|
||||
|
||||
let error = run_through(&host)
|
||||
.await
|
||||
.err()
|
||||
.expect("the host failure fails the call");
|
||||
|
||||
assert_eq!(error, Error::InvalidRequest("vetoed by the host".into()));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
assert!(host.raw_responses().is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall) {
|
||||
let raw = message_body();
|
||||
let upstream = upstream([json_response(raw.clone())]).await;
|
||||
let host = RecordingHost::passthrough(authenticated(call, upstream.uri()));
|
||||
|
||||
let output = run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
let [emitted] = <[String; 1]>::try_from(host.raw_responses())
|
||||
.unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len()));
|
||||
assert_eq!(serde_json::from_str::<Value>(&emitted).unwrap(), raw);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::upstream_error(ResponseTemplate::new(500).set_body_string("boom"))]
|
||||
#[case::stream(ResponseTemplate::new(200).set_body_raw("event: message_stop\ndata: {}\n\n", "text/event-stream"))]
|
||||
#[tokio::test]
|
||||
async fn no_raw_response_is_emitted_for_a_stream_or_a_failure(
|
||||
call: MessagesCall,
|
||||
#[case] response: ResponseTemplate,
|
||||
) {
|
||||
let upstream = upstream([response]).await;
|
||||
let host = RecordingHost::passthrough(authenticated(
|
||||
with_fields(call, json!({"stream": true})),
|
||||
upstream.uri(),
|
||||
));
|
||||
|
||||
let _ = run_through(&host).await;
|
||||
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
assert!(host.raw_responses().is_empty());
|
||||
}
|
||||
|
||||
/// Python logs `optional_params` as what it is about to send, so a dropped param must
|
||||
/// not resurface in callbacks.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_request_context_carries_the_shaped_params_without_model_or_messages(
|
||||
call: MessagesCall,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let host = RecordingHost::passthrough(authenticated(
|
||||
MessagesCall {
|
||||
shaping: MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_sampling_params: false,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
drop_params: true,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
..with_fields(call, json!({"temperature": 0.2}))
|
||||
},
|
||||
upstream.uri(),
|
||||
));
|
||||
|
||||
run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
let [optional_params] = <[Value; 1]>::try_from(host.optional_params.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
|
||||
assert_eq!(optional_params, json!({"max_tokens": 16}));
|
||||
}
|
||||
120
litellm-rust/crates/core/tests/messages/main.rs
Normal file
120
litellm-rust/crates/core/tests/messages/main.rs
Normal file
|
|
@ -0,0 +1,120 @@
|
|||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use litellm_core::messages::{
|
||||
Error, MessagesCall, MessagesShaping,
|
||||
route::{LocalMessagesHost, MessagesMachine, MessagesOutput, messages_machine},
|
||||
};
|
||||
use litellm_http::{HttpSettings, Resolution};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
};
|
||||
use rstest::fixture;
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
|
||||
#[path = "../support/mod.rs"]
|
||||
mod support;
|
||||
use support::*;
|
||||
|
||||
mod host;
|
||||
mod request;
|
||||
mod response;
|
||||
mod secrets;
|
||||
mod stream;
|
||||
|
||||
const MODEL: &str = "claude-sonnet-4-5";
|
||||
|
||||
fn object(value: Value) -> Map<String, Value> {
|
||||
let Value::Object(map) = value else {
|
||||
panic!("expected a json object, got {value}");
|
||||
};
|
||||
map
|
||||
}
|
||||
|
||||
fn body(value: Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall {
|
||||
let current = object(serde_json::to_value(&call.body).unwrap());
|
||||
MessagesCall {
|
||||
body: body(Value::Object(
|
||||
current.into_iter().chain(object(fields)).collect(),
|
||||
)),
|
||||
..call
|
||||
}
|
||||
}
|
||||
|
||||
fn with_model(call: MessagesCall, model: &str) -> MessagesCall {
|
||||
with_fields(call, json!({"model": model}))
|
||||
}
|
||||
|
||||
fn message_body() -> Value {
|
||||
json!({
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
"model": MODEL,
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 1, "output_tokens": 2}
|
||||
})
|
||||
}
|
||||
|
||||
fn message_response() -> ResponseTemplate {
|
||||
json_response(message_body())
|
||||
}
|
||||
|
||||
/// A non-streaming call with nothing that would authenticate or route it, so each test
|
||||
/// states the provider, credentials, and base it depends on.
|
||||
#[fixture]
|
||||
fn call() -> MessagesCall {
|
||||
MessagesCall {
|
||||
body: body(json!({
|
||||
"model": MODEL,
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
})),
|
||||
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(),
|
||||
}
|
||||
}
|
||||
|
||||
fn headers<'a>(pairs: impl IntoIterator<Item = (&'a str, &'a str)>) -> Option<Map<String, Value>> {
|
||||
Some(
|
||||
pairs
|
||||
.into_iter()
|
||||
.map(|(name, value)| (name.to_string(), Value::from(value)))
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
fn machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
|
||||
messages_machine(&support::resources(), &http_config(), secrets)
|
||||
.expect("default HTTP settings build a client")
|
||||
}
|
||||
|
||||
async fn run_with(
|
||||
secrets: Arc<RecordingSecrets>,
|
||||
call: MessagesCall,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await
|
||||
}
|
||||
|
||||
/// Runs the route with a secret source that knows nothing, so no environment leaks in.
|
||||
async fn run(call: MessagesCall) -> Result<MessagesOutput, Error> {
|
||||
run_with(Arc::new(RecordingSecrets::empty()), call).await
|
||||
}
|
||||
|
||||
async fn run_message(call: MessagesCall) -> AnthropicMessagesResponse {
|
||||
match run(call).await.expect("messages call succeeds") {
|
||||
MessagesOutput::Message(message) => *message,
|
||||
MessagesOutput::Streamed => panic!("a non-streaming call returned a stream"),
|
||||
}
|
||||
}
|
||||
658
litellm-rust/crates/core/tests/messages/request.rs
Normal file
658
litellm-rust/crates/core/tests/messages/request.rs
Normal file
|
|
@ -0,0 +1,658 @@
|
|||
use litellm_llms::anthropic::common_utils::{AnthropicModelCapabilities, SupportedEffortTiers};
|
||||
use litellm_types::llms::anthropic::{AnthropicBeta, BetaSet};
|
||||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic_key("anthropic", Some("sk-ant"), &[], ("x-api-key", "sk-ant"), &["authorization"])]
|
||||
#[case::azure_key("azure_ai", Some("sk-azure"), &[], ("x-api-key", "sk-azure"), &["authorization"])]
|
||||
#[case::caller_x_api_key_wins(
|
||||
"azure_ai",
|
||||
Some("rust-fallback-key"),
|
||||
&[("x-api-key", "from-python")],
|
||||
("x-api-key", "from-python"),
|
||||
&["authorization"]
|
||||
)]
|
||||
#[case::entra_bearer_without_key(
|
||||
"azure_ai",
|
||||
None,
|
||||
&[("Authorization", "Bearer entra-token")],
|
||||
("authorization", "Bearer entra-token"),
|
||||
&["x-api-key"]
|
||||
)]
|
||||
#[case::empty_bearer_falls_back_to_key(
|
||||
"azure_ai",
|
||||
Some("sk-azure"),
|
||||
&[("Authorization", "Bearer ")],
|
||||
("x-api-key", "sk-azure"),
|
||||
&[]
|
||||
)]
|
||||
#[case::anthropic_forwards_caller_authorization(
|
||||
"anthropic",
|
||||
Some("sk-ant"),
|
||||
&[("Authorization", "Bearer caller")],
|
||||
("authorization", "Bearer caller"),
|
||||
&["x-api-key"]
|
||||
)]
|
||||
#[case::anthropic_oauth_key_becomes_bearer(
|
||||
"anthropic",
|
||||
Some("sk-ant-oat01-token"),
|
||||
&[],
|
||||
("authorization", "Bearer sk-ant-oat01-token"),
|
||||
&["x-api-key"]
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn credentials_become_exactly_one_auth_header(
|
||||
call: MessagesCall,
|
||||
#[case] provider: &str,
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] extra_headers: &[(&str, &str)],
|
||||
#[case] expected: (&str, &str),
|
||||
#[case] absent: &[&str],
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_key: api_key.map(Into::into),
|
||||
api_base: Some(upstream.uri()),
|
||||
extra_headers: headers(extra_headers.iter().copied()),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
let (name, value) = expected;
|
||||
assert_eq!(request.header_values(name), [value]);
|
||||
for name in absent {
|
||||
assert_eq!(request.header(name), None, "{name} must not be sent");
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure_ai("azure_ai")]
|
||||
#[tokio::test]
|
||||
async fn a_call_without_credentials_fails_before_sending(
|
||||
call: MessagesCall,
|
||||
#[case] provider: &str,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
let error = run(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("a call without credentials fails");
|
||||
|
||||
assert!(
|
||||
matches!(
|
||||
error,
|
||||
Error::Auth(litellm_auth::Error::MissingApiKey { .. })
|
||||
),
|
||||
"{error:?}"
|
||||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic(MODEL, Some("anthropic"), "", "/v1/messages")]
|
||||
#[case::anthropic_base_with_trailing_slash(MODEL, Some("anthropic"), "/", "/v1/messages")]
|
||||
#[case::anthropic_base_with_the_messages_path(
|
||||
MODEL,
|
||||
Some("anthropic"),
|
||||
"/v1/messages",
|
||||
"/v1/messages"
|
||||
)]
|
||||
#[case::azure_ai(MODEL, Some("azure_ai"), "", "/anthropic/v1/messages")]
|
||||
#[case::provider_from_model_prefix("anthropic/claude-sonnet-4-5", None, "", "/v1/messages")]
|
||||
#[tokio::test]
|
||||
async fn each_provider_posts_to_its_messages_endpoint(
|
||||
call: MessagesCall,
|
||||
#[case] model: &str,
|
||||
#[case] provider: Option<&str>,
|
||||
#[case] base_suffix: &str,
|
||||
#[case] path: &str,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: provider.map(Into::into),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(format!("{}{base_suffix}", upstream.uri())),
|
||||
..with_model(call, model)
|
||||
})
|
||||
.await;
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(request.method.as_str(), "POST");
|
||||
assert_eq!(request.url.path(), path);
|
||||
assert_eq!(request.json()["model"], MODEL);
|
||||
assert_eq!(request.header_values("anthropic-version"), ["2023-06-01"]);
|
||||
assert_eq!(request.header_values("content-type"), ["application/json"]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::unknown_provider(MODEL, Some("openai"), "openai")]
|
||||
#[case::unresolvable_model(
|
||||
"no-such-model",
|
||||
None,
|
||||
"unable to resolve custom_llm_provider for messages request"
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn unsupported_providers_are_rejected_before_sending(
|
||||
call: MessagesCall,
|
||||
#[case] model: &str,
|
||||
#[case] provider: Option<&str>,
|
||||
#[case] reported: &str,
|
||||
) {
|
||||
let error = run(MessagesCall {
|
||||
custom_llm_provider: provider.map(Into::into),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(UNREACHABLE_BASE.into()),
|
||||
..with_model(call, model)
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("unsupported provider errors");
|
||||
|
||||
assert_eq!(error, Error::InvalidProvider(reported.into()));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn caller_headers_and_provider_scoped_headers_are_forwarded(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let scoped = |provider: &str, value: &str| ProviderSpecificHeader {
|
||||
custom_llm_provider: provider.into(),
|
||||
extra_headers: object(json!({"x-scoped": value})),
|
||||
};
|
||||
|
||||
run_message(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
extra_headers: headers([("anthropic-beta", "token-efficient-tools-2025-02-19")]),
|
||||
provider_specific_header: Some(ProviderSpecificHeaders::Many(vec![
|
||||
scoped("bedrock", "other-provider"),
|
||||
scoped("azure_ai, anthropic", "this-provider"),
|
||||
])),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(
|
||||
request.header("anthropic-beta"),
|
||||
Some("token-efficient-tools-2025-02-19")
|
||||
);
|
||||
assert_eq!(request.header_values("x-scoped"), ["this-provider"]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
api_key: Some("sk-azure".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
body: body(json!({
|
||||
"model": MODEL,
|
||||
"max_tokens": 16,
|
||||
"messages": [{
|
||||
"role": "user",
|
||||
"content": [{
|
||||
"type": "text",
|
||||
"text": "hi",
|
||||
"cache_control": {"type": "ephemeral", "scope": "global"}
|
||||
}]
|
||||
}]
|
||||
})),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.json()["messages"][0]["content"][0]["cache_control"],
|
||||
json!({"type": "ephemeral"})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
additional_drop_params: vec!["temperature".into()],
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
..with_fields(call, json!({"temperature": 0.5, "top_k": 3}))
|
||||
})
|
||||
.await;
|
||||
|
||||
let sent = only_request(&upstream).await.json();
|
||||
assert_eq!(sent.get("temperature"), None);
|
||||
assert_eq!(sent["top_k"], 3);
|
||||
}
|
||||
|
||||
fn sent_betas(request: &wiremock::Request) -> BetaSet {
|
||||
let [header] = <[&str; 1]>::try_from(request.header_values("anthropic-beta"))
|
||||
.unwrap_or_else(|values| panic!("expected one anthropic-beta header, got {values:?}"));
|
||||
header.parse().unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[AnthropicBeta::StructuredOutputs20251113])]
|
||||
#[case::fast_mode(json!({"speed": "fast"}), &[AnthropicBeta::FastMode20260201])]
|
||||
#[case::compaction(json!({"compaction": {"enabled": true}}), &[AnthropicBeta::Compact20260904])]
|
||||
#[case::context_management_edits(
|
||||
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]}}),
|
||||
&[AnthropicBeta::ContextManagement20250627]
|
||||
)]
|
||||
#[case::per_message_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
&[AnthropicBeta::PerTurnControl20260701]
|
||||
)]
|
||||
#[case::advisor_tool(
|
||||
json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": MODEL}]}),
|
||||
&[AnthropicBeta::AdvisorTool20260301]
|
||||
)]
|
||||
#[case::several_features_at_once(
|
||||
json!({"speed": "fast", "output_format": {"type": "json_schema"}}),
|
||||
&[AnthropicBeta::StructuredOutputs20251113, AnthropicBeta::FastMode20260201]
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
|
||||
call: MessagesCall,
|
||||
#[case] fields: Value,
|
||||
#[case] features: &[AnthropicBeta],
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let capabilities = AnthropicModelCapabilities {
|
||||
supports_speed: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
};
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
extra_headers: headers([("Anthropic-Beta", "caller-beta-2025-01-01")]),
|
||||
shaping: MessagesShaping {
|
||||
capabilities,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
..call
|
||||
},
|
||||
fields,
|
||||
))
|
||||
.await;
|
||||
|
||||
let sent = sent_betas(&only_request(&upstream).await);
|
||||
let expected: BetaSet = features
|
||||
.iter()
|
||||
.cloned()
|
||||
.chain([AnthropicBeta::Other("caller-beta-2025-01-01".to_string())])
|
||||
.collect();
|
||||
assert_eq!(sent, expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn an_oauth_key_sends_the_browser_access_header_and_the_oauth_beta(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
api_key: Some("sk-ant-oat01-token".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(
|
||||
request.header("anthropic-dangerous-direct-browser-access"),
|
||||
Some("true")
|
||||
);
|
||||
assert_eq!(
|
||||
sent_betas(&request),
|
||||
BetaSet::from_iter([AnthropicBeta::Oauth20250420])
|
||||
);
|
||||
assert_eq!(request.header("x-api-key"), None);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure_ai("azure_ai")]
|
||||
#[tokio::test]
|
||||
async fn caller_protocol_headers_win_over_the_defaults(call: MessagesCall, #[case] provider: &str) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
extra_headers: headers([
|
||||
("Anthropic-Version", "2024-01-01"),
|
||||
("Content-Type", "application/json; charset=utf-8"),
|
||||
]),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(request.header_values("anthropic-version"), ["2024-01-01"]);
|
||||
assert_eq!(
|
||||
request.header_values("content-type"),
|
||||
["application/json; charset=utf-8"]
|
||||
);
|
||||
}
|
||||
|
||||
fn sampling_removed() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities {
|
||||
supports_sampling_params: false,
|
||||
..AnthropicModelCapabilities::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::sampling_params(sampling_removed(), json!({"temperature": 0.2, "top_p": 0.9, "top_k": 5}), &["temperature", "top_p", "top_k"], "temperature=0.2")]
|
||||
#[case::speed(AnthropicModelCapabilities::default(), json!({"speed": "fast"}), &["speed"], "speed='fast'")]
|
||||
#[tokio::test]
|
||||
async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_it(
|
||||
call: MessagesCall,
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] dropped: &[&str],
|
||||
#[case] rejected_as: &str,
|
||||
) {
|
||||
let upstream = upstream([message_response(), message_response()]).await;
|
||||
let shaped = |drop_params: bool| {
|
||||
with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
capabilities,
|
||||
drop_params,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
body: call.body.clone(),
|
||||
custom_llm_provider: call.custom_llm_provider.clone(),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: call.timeout,
|
||||
},
|
||||
fields.clone(),
|
||||
)
|
||||
};
|
||||
|
||||
let error = run(shaped(false))
|
||||
.await
|
||||
.err()
|
||||
.expect("an unsupported param is rejected without drop_params");
|
||||
assert!(
|
||||
matches!(&error, Error::InvalidRequest(message) if message.contains(rejected_as)),
|
||||
"{error:?}"
|
||||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
|
||||
run_message(shaped(true)).await;
|
||||
let sent = only_request(&upstream).await.json();
|
||||
for name in dropped {
|
||||
assert_eq!(sent.get(*name), None, "{name} must be dropped");
|
||||
}
|
||||
assert_eq!(sent["max_tokens"], 16);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::adaptive_thinking(json!({"type": "adaptive"}), json!({"type": "adaptive", "display": "summarized"}))]
|
||||
#[case::disabled_thinking(json!({"type": "disabled"}), json!({"type": "disabled"}))]
|
||||
#[tokio::test]
|
||||
async fn reasoning_auto_summary_marks_active_thinking_on_the_wire(
|
||||
call: MessagesCall,
|
||||
#[case] thinking: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
reasoning_auto_summary: true,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
..call
|
||||
},
|
||||
json!({"thinking": thinking}),
|
||||
))
|
||||
.await;
|
||||
|
||||
assert_eq!(only_request(&upstream).await.json()["thinking"], expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::reasoning_effort_on_an_adaptive_model(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
supports_output_config: true,
|
||||
effort_tiers: SupportedEffortTiers { high: true, ..SupportedEffortTiers::default() },
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
json!({"reasoning_effort": "high"}),
|
||||
json!({"thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}})
|
||||
)]
|
||||
#[case::reasoning_effort_on_a_legacy_model_caps_the_budget_below_max_tokens(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
json!({"reasoning_effort": "high"}),
|
||||
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
|
||||
)]
|
||||
#[case::adaptive_payload_on_a_legacy_model_becomes_a_capped_budget(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}, "temperature": 0}),
|
||||
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
|
||||
)]
|
||||
#[case::adaptive_payload_on_a_model_without_reasoning_is_dropped(
|
||||
AnthropicModelCapabilities::default(),
|
||||
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}),
|
||||
json!({})
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn reasoning_is_translated_by_the_model_capabilities(
|
||||
call: MessagesCall,
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
capabilities,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
..call
|
||||
},
|
||||
[("max_tokens".to_string(), json!(3000))]
|
||||
.into_iter()
|
||||
.chain(object(fields))
|
||||
.collect(),
|
||||
))
|
||||
.await;
|
||||
|
||||
let sent = only_request(&upstream).await.json();
|
||||
assert_eq!(sent.get("reasoning_effort"), None);
|
||||
assert_eq!(sent.get("temperature"), None);
|
||||
let reasoning: Map<String, Value> = ["thinking", "output_config"]
|
||||
.into_iter()
|
||||
.filter_map(|name| Some((name.to_string(), sent.get(name)?.clone())))
|
||||
.collect();
|
||||
assert_eq!(Value::Object(reasoning), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty_text_blocks(
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": " "}, {"type": "text", "text": "kept"}]}]),
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "kept"}]}])
|
||||
)]
|
||||
#[case::provider_specific_fields(
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "kept", "provider_specific_fields": {"x": 1}}]}]),
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "kept"}]}])
|
||||
)]
|
||||
#[case::unencrypted_web_search_results_become_text(
|
||||
json!([{"role": "assistant", "content": [{
|
||||
"type": "web_search_tool_result",
|
||||
"tool_use_id": "srvtoolu_1",
|
||||
"content": [{"type": "web_search_result", "title": "T", "url": "https://e.x", "page_age": null}]
|
||||
}]}]),
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "Web search results:\n\nTitle: T\nURL: https://e.x"}]}])
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn replayed_history_is_cleaned_before_sending(
|
||||
call: MessagesCall,
|
||||
#[case] history: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
json!({"messages": history}),
|
||||
))
|
||||
.await;
|
||||
|
||||
assert_eq!(only_request(&upstream).await.json()["messages"], expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn metadata_is_reduced_to_the_user_id(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
json!({"metadata": {"user_id": "u-1", "trace_id": "internal", "tags": ["a"]}}),
|
||||
))
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.json()["metadata"],
|
||||
json!({"user_id": "u-1"})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::numeric_user_id(json!({"metadata": {"user_id": 7}}))]
|
||||
#[case::missing_max_tokens(json!({"max_tokens": null}))]
|
||||
#[tokio::test]
|
||||
async fn an_invalid_request_fails_before_sending(call: MessagesCall, #[case] fields: Value) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
let error = run(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
fields,
|
||||
))
|
||||
.await
|
||||
.err()
|
||||
.expect("the request is rejected");
|
||||
|
||||
assert!(error.is_request(), "{error:?}");
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn azure_folds_system_role_messages_into_the_system_prompt(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
api_key: Some("sk-azure".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
json!({
|
||||
"system": "top level",
|
||||
"messages": [
|
||||
{"role": "system", "content": "from a message"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]
|
||||
}),
|
||||
))
|
||||
.await;
|
||||
|
||||
let sent = only_request(&upstream).await.json();
|
||||
assert_eq!(
|
||||
sent["system"],
|
||||
json!([
|
||||
{"type": "text", "text": "top level"},
|
||||
{"type": "text", "text": "from a message"}
|
||||
])
|
||||
);
|
||||
assert_eq!(sent["messages"], json!([{"role": "user", "content": "hi"}]));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::bare_model(MODEL, MODEL)]
|
||||
#[case::one_prefix("anthropic/claude-sonnet-4-5", MODEL)]
|
||||
#[case::doubled_prefix_loses_one_segment(
|
||||
"anthropic/anthropic/claude-sonnet-4-5",
|
||||
"anthropic/claude-sonnet-4-5"
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn the_provider_prefix_is_stripped_exactly_once(
|
||||
call: MessagesCall,
|
||||
#[case] model: &str,
|
||||
#[case] sent_model: &str,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..with_model(call, model)
|
||||
})
|
||||
.await;
|
||||
|
||||
assert_eq!(only_request(&upstream).await.json()["model"], sent_model);
|
||||
}
|
||||
223
litellm-rust/crates/core/tests/messages/response.rs
Normal file
223
litellm-rust/crates/core/tests/messages/response.rs
Normal file
|
|
@ -0,0 +1,223 @@
|
|||
use litellm_core::{
|
||||
Phase,
|
||||
messages::{MessagesResponse, messages, messages_body},
|
||||
};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure_ai("azure_ai")]
|
||||
#[tokio::test]
|
||||
async fn the_provider_message_is_returned(call: MessagesCall, #[case] provider: &str) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
let message = run_message(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
assert_eq!(message.id, "msg_1");
|
||||
assert_eq!(message.content, [json!({"type": "text", "text": "hi"})]);
|
||||
assert_eq!(message.stop_reason.as_deref(), Some("end_turn"));
|
||||
}
|
||||
|
||||
/// A refusal and fields the route does not model come back exactly as the provider sent
|
||||
/// them, since the Python side returns the raw message and the router decides what to do.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_message_passes_through_losslessly(call: MessagesCall) {
|
||||
let upstream_body = json!({
|
||||
"id": "msg_2",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": MODEL,
|
||||
"content": [
|
||||
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "q"}},
|
||||
{"type": "text", "text": "no", "citations": [{"type": "web_search_result_location", "url": "https://e.x"}]}
|
||||
],
|
||||
"stop_reason": "refusal",
|
||||
"stop_sequence": null,
|
||||
"stop_details": {"type": "safeguard", "safeguard_types": ["dangerous_tool_use"]},
|
||||
"container": {"id": "container_1", "expires_at": "2026-01-01T00:00:00Z"},
|
||||
"context_management": {"applied_edits": []},
|
||||
"usage": {"input_tokens": 1, "output_tokens": 2, "server_tool_use": {"web_search_requests": 1}},
|
||||
"unknown_future_field": {"nested": true}
|
||||
});
|
||||
let upstream = upstream([json_response(upstream_body.clone())]).await;
|
||||
|
||||
let message = run_message(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
assert_eq!(message.stop_reason.as_deref(), Some("refusal"));
|
||||
assert_eq!(serde_json::to_value(&message).unwrap(), upstream_body);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_json_error_envelope_is_kept_verbatim(call: MessagesCall) {
|
||||
let envelope =
|
||||
json!({"type": "error", "error": {"type": "invalid_request_error", "message": "bad"}});
|
||||
let upstream = upstream([status_response(400, envelope.clone())]).await;
|
||||
|
||||
let error = run(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
|
||||
let Error::Transport(TransportError::Http { status, body }) = error else {
|
||||
panic!("{error:?}");
|
||||
};
|
||||
assert_eq!(status, 400);
|
||||
assert_eq!(serde_json::from_str::<Value>(&body).unwrap(), envelope);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_long_error_body_is_truncated_at_the_documented_cap(call: MessagesCall) {
|
||||
let long = "x".repeat(600);
|
||||
let upstream = upstream([ResponseTemplate::new(500).set_body_string(long.clone())]).await;
|
||||
|
||||
let error = run(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Transport(TransportError::Http {
|
||||
status: 500,
|
||||
body: format!("{}... (truncated)", &long[..256])
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::bad_request(400)]
|
||||
#[case::unauthorized(401)]
|
||||
#[case::rate_limited(429)]
|
||||
#[case::server_error(500)]
|
||||
#[case::overloaded(529)]
|
||||
#[tokio::test]
|
||||
async fn an_upstream_error_keeps_its_status_and_body(call: MessagesCall, #[case] status: u16) {
|
||||
let upstream =
|
||||
upstream([ResponseTemplate::new(status).set_body_string("upstream said no")]).await;
|
||||
|
||||
let error = run(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: "upstream said no".into()
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::not_json(ResponseTemplate::new(200).set_body_string("not json"))]
|
||||
#[case::not_a_message(json_response(json!({"unexpected": true})))]
|
||||
#[tokio::test]
|
||||
async fn an_unreadable_success_body_is_an_invalid_response(
|
||||
call: MessagesCall,
|
||||
#[case] response: ResponseTemplate,
|
||||
) {
|
||||
let upstream = upstream([response]).await;
|
||||
|
||||
let error = run(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("an unreadable body fails");
|
||||
|
||||
assert_eq!(error.phase(), Phase::AfterSend, "{error:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) {
|
||||
let upstream = upstream([message_response().set_delay(Duration::from_secs(5))]).await;
|
||||
|
||||
let error = run(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
timeout: Some(Duration::from_millis(100)),
|
||||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("the call times out");
|
||||
|
||||
assert!(matches!(error, Error::Transport(_)), "{error:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_facade_sends_through_the_injected_http_pool_configuration(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let base = upstream.uri();
|
||||
let settings = HttpSettings {
|
||||
user_agent: Some("host-owned/1".into()),
|
||||
..HttpSettings::default()
|
||||
};
|
||||
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&Resolution::from(&settings).config,
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
api_key: Some("sk-ant".into()),
|
||||
api_base: Some(base),
|
||||
..call
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
let MessagesResponse::Message(message) = response else {
|
||||
panic!("a non-streaming request returns a message");
|
||||
};
|
||||
assert_eq!(message.id, "msg_1");
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.header("x-api-key"), Some("sk-ant"));
|
||||
assert_eq!(sent.header("user-agent"), Some("host-owned/1"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::mistyped_param(json!({"model": MODEL, "messages": [], "max_tokens": "16"}))]
|
||||
#[case::missing_messages(json!({"model": MODEL, "max_tokens": 16}))]
|
||||
fn a_body_that_does_not_parse_is_an_invalid_request(#[case] raw: Value) {
|
||||
let error = messages_body(object(raw)).expect_err("the body is rejected");
|
||||
|
||||
assert!(
|
||||
matches!(&error, Error::InvalidRequest(message) if message.starts_with("invalid Anthropic messages request: ")),
|
||||
"{error:?}"
|
||||
);
|
||||
}
|
||||
200
litellm-rust/crates/core/tests/messages/secrets.rs
Normal file
200
litellm-rust/crates/core/tests/messages/secrets.rs
Normal file
|
|
@ -0,0 +1,200 @@
|
|||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic(
|
||||
"anthropic",
|
||||
"ANTHROPIC_API_KEY",
|
||||
"ANTHROPIC_BASE_URL",
|
||||
"/v1/messages",
|
||||
&["ANTHROPIC_API_KEY", "ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"]
|
||||
)]
|
||||
#[case::azure_ai(
|
||||
"azure_ai",
|
||||
"AZURE_API_KEY",
|
||||
"AZURE_API_BASE",
|
||||
"/anthropic/v1/messages",
|
||||
&["AZURE_API_KEY", "AZURE_API_BASE"]
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn the_credential_and_base_come_from_the_secret_source(
|
||||
call: MessagesCall,
|
||||
#[case] provider: &str,
|
||||
#[case] key_name: &str,
|
||||
#[case] base_name: &str,
|
||||
#[case] path: &str,
|
||||
#[case] looked_up: &[&str],
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let base = upstream.uri();
|
||||
let secrets = Arc::new(RecordingSecrets::new([
|
||||
(key_name, "sk-from-manager"),
|
||||
(base_name, base.as_str()),
|
||||
]));
|
||||
|
||||
let output = run_with(
|
||||
secrets.clone(),
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
..call
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("messages call succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(request.url.path(), path);
|
||||
assert_eq!(request.header("x-api-key"), Some("sk-from-manager"));
|
||||
assert_eq!(secrets.requested(), looked_up);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn call_arguments_win_over_the_secret_source(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let secrets = Arc::new(RecordingSecrets::new([
|
||||
("ANTHROPIC_API_KEY", "sk-from-manager"),
|
||||
("ANTHROPIC_BASE_URL", UNREACHABLE_BASE),
|
||||
]));
|
||||
|
||||
run_with(
|
||||
secrets,
|
||||
MessagesCall {
|
||||
api_key: Some("sk-from-call".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("messages call succeeds");
|
||||
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header("x-api-key"),
|
||||
Some("sk-from-call")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_secret_manager_failure_fails_the_call_before_sending(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
let error = run_with(
|
||||
Arc::new(RecordingSecrets::failing()),
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
)
|
||||
.await
|
||||
.err()
|
||||
.expect("a secret manager failure fails the call");
|
||||
|
||||
assert!(
|
||||
matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)),
|
||||
"{error:?}"
|
||||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum Base {
|
||||
Upstream,
|
||||
Unreachable,
|
||||
Blank,
|
||||
Absent,
|
||||
}
|
||||
|
||||
fn base_value(base: Base, upstream: &str) -> Option<String> {
|
||||
match base {
|
||||
Base::Upstream => Some(upstream.to_string()),
|
||||
Base::Unreachable => Some(UNREACHABLE_BASE.to_string()),
|
||||
Base::Blank => Some(" ".to_string()),
|
||||
Base::Absent => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_base_beats_base_url(Base::Upstream, Base::Unreachable)]
|
||||
#[case::blank_api_base_falls_through_to_base_url(Base::Blank, Base::Upstream)]
|
||||
#[case::base_url_alone(Base::Absent, Base::Upstream)]
|
||||
#[tokio::test]
|
||||
async fn the_anthropic_base_env_precedence_picks_the_upstream(
|
||||
call: MessagesCall,
|
||||
#[case] api_base: Base,
|
||||
#[case] base_url: Base,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let uri = upstream.uri();
|
||||
let values: Vec<(&str, &str)> = [
|
||||
("ANTHROPIC_API_KEY", Some("sk-env".to_string())),
|
||||
("ANTHROPIC_API_BASE", base_value(api_base, &uri)),
|
||||
("ANTHROPIC_BASE_URL", base_value(base_url, &uri)),
|
||||
]
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((*name, value.as_deref()?)))
|
||||
.map(|(name, value)| (name, Box::leak(value.to_string().into_boxed_str()) as &str))
|
||||
.collect();
|
||||
|
||||
run_with(Arc::new(RecordingSecrets::new(values)), call)
|
||||
.await
|
||||
.expect("messages call reaches the upstream the precedence picks");
|
||||
|
||||
assert_eq!(only_request(&upstream).await.url.path(), "/v1/messages");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::auth_token_alone(
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "tok")],
|
||||
("authorization", "Bearer tok"),
|
||||
"x-api-key"
|
||||
)]
|
||||
#[case::api_key_beats_the_auth_token(
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "tok")],
|
||||
("x-api-key", "sk-env"),
|
||||
"authorization"
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn the_auth_token_env_is_a_bearer_only_without_a_key(
|
||||
call: MessagesCall,
|
||||
#[case] values: &[(&str, &str)],
|
||||
#[case] expected: (&str, &str),
|
||||
#[case] absent: &str,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_with(
|
||||
Arc::new(RecordingSecrets::new(values.iter().copied())),
|
||||
MessagesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("messages call succeeds");
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
let (name, value) = expected;
|
||||
assert_eq!(request.header_values(name), [value]);
|
||||
assert_eq!(request.header(absent), None);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall) {
|
||||
let error = run_with(
|
||||
Arc::new(RecordingSecrets::new([("AZURE_API_KEY", "sk-azure")])),
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
..call
|
||||
},
|
||||
)
|
||||
.await
|
||||
.err()
|
||||
.expect("azure needs a base");
|
||||
|
||||
assert_eq!(error, Error::Auth(litellm_auth::Error::MissingAzureApiBase));
|
||||
}
|
||||
459
litellm-rust/crates/core/tests/messages/stream.rs
Normal file
459
litellm-rust/crates/core/tests/messages/stream.rs
Normal file
|
|
@ -0,0 +1,459 @@
|
|||
use std::{
|
||||
convert::Infallible,
|
||||
sync::{Mutex, mpsc},
|
||||
};
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt};
|
||||
use litellm_core::messages::{
|
||||
MessagesResponse, messages,
|
||||
route::{Messages, MessagesStreamHead},
|
||||
};
|
||||
use litellm_host::host::{Demand, Host};
|
||||
use litellm_tracing::{Logger, Metadata, Record, Sink};
|
||||
use rstest::rstest;
|
||||
use tokio::{
|
||||
io::{AsyncReadExt, AsyncWriteExt},
|
||||
net::TcpListener,
|
||||
task::JoinHandle,
|
||||
};
|
||||
|
||||
use super::*;
|
||||
|
||||
const UPSTREAM_HEADERS: [(&str, &str); 2] = [
|
||||
("request-id", "req_upstream_123"),
|
||||
("anthropic-ratelimit-requests-remaining", "41"),
|
||||
];
|
||||
|
||||
const SSE_BODY: &str = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
|
||||
|
||||
enum Seen {
|
||||
Open(Vec<(String, String)>),
|
||||
Deliver(Bytes),
|
||||
}
|
||||
|
||||
struct TraceSink(mpsc::Sender<(String, Value)>);
|
||||
|
||||
impl Sink for TraceSink {
|
||||
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
|
||||
metadata.target().starts_with("litellm_core::messages")
|
||||
}
|
||||
|
||||
fn emit(&self, record: &Record) {
|
||||
self.0
|
||||
.send((record.message.clone(), Value::Object(record.fields.clone())))
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
/// Projects like `LocalMessagesHost`, records every stream op in the order the route
|
||||
/// performs it, and detaches after `detach_after` ops.
|
||||
struct RecordingStreamHost {
|
||||
call: LocalMessagesHost,
|
||||
detach_after: usize,
|
||||
seen: Mutex<Vec<Seen>>,
|
||||
}
|
||||
|
||||
impl RecordingStreamHost {
|
||||
fn new(call: MessagesCall, detach_after: usize) -> Self {
|
||||
Self {
|
||||
call: LocalMessagesHost::new(call),
|
||||
detach_after,
|
||||
seen: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn record(&self, op: Seen) -> Demand {
|
||||
let mut seen = self.seen.lock().unwrap();
|
||||
seen.push(op);
|
||||
match seen.len() < self.detach_after {
|
||||
true => Demand::More,
|
||||
false => Demand::Detached,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for RecordingStreamHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.project().await
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
}
|
||||
|
||||
async fn open(&self, head: MessagesStreamHead) -> Result<Demand, Error> {
|
||||
Ok(self.record(Seen::Open(head.headers)))
|
||||
}
|
||||
|
||||
async fn deliver(&self, chunk: Bytes) -> Result<Demand, Error> {
|
||||
Ok(self.record(Seen::Deliver(chunk)))
|
||||
}
|
||||
}
|
||||
|
||||
fn streaming(call: MessagesCall, api_base: String) -> MessagesCall {
|
||||
MessagesCall {
|
||||
api_key: Some("sk-ant".into()),
|
||||
api_base: Some(api_base),
|
||||
..with_fields(call, json!({"stream": true}))
|
||||
}
|
||||
}
|
||||
|
||||
fn sse_response() -> ResponseTemplate {
|
||||
UPSTREAM_HEADERS.iter().fold(
|
||||
ResponseTemplate::new(200).set_body_raw(SSE_BODY, "text/event-stream"),
|
||||
|response, (name, value)| response.insert_header(*name, *value),
|
||||
)
|
||||
}
|
||||
|
||||
async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: MessagesCall) {
|
||||
let upstream = upstream([sse_response()]).await;
|
||||
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
|
||||
|
||||
let outcome = stream_through(&host).await.expect("streamed call succeeds");
|
||||
|
||||
assert!(matches!(outcome, MessagesOutput::Streamed));
|
||||
let seen = host.seen.into_inner().unwrap();
|
||||
let [Seen::Open(headers), chunks @ ..] = seen.as_slice() else {
|
||||
panic!("the stream opens before any chunk is delivered");
|
||||
};
|
||||
let surfaced: Vec<(&str, &str)> = headers
|
||||
.iter()
|
||||
.filter(|(name, _)| {
|
||||
UPSTREAM_HEADERS
|
||||
.iter()
|
||||
.any(|(upstream, _)| upstream == name)
|
||||
})
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
.collect();
|
||||
assert_eq!(surfaced, UPSTREAM_HEADERS);
|
||||
let delivered: Vec<u8> = chunks
|
||||
.iter()
|
||||
.flat_map(|step| match step {
|
||||
Seen::Deliver(chunk) => chunk.to_vec(),
|
||||
Seen::Open(_) => panic!("the stream opens exactly once"),
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(delivered, SSE_BODY.as_bytes());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn debug_trace_keeps_provider_input_and_every_stream_chunk(call: MessagesCall) {
|
||||
let upstream = upstream([sse_response()]).await;
|
||||
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
|
||||
let (sender, receiver) = mpsc::channel();
|
||||
|
||||
Logger::new(TraceSink(sender))
|
||||
.instrument(stream_through(&host))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let records: Vec<(String, Value)> = receiver.try_iter().collect();
|
||||
let request = records
|
||||
.iter()
|
||||
.find(|(message, _)| message == "provider request")
|
||||
.unwrap();
|
||||
let body: Value = serde_json::from_str(request.1["body"].as_str().unwrap()).unwrap();
|
||||
assert_eq!(body["messages"][0]["content"], "hi");
|
||||
assert_eq!(request.1["stream"], true);
|
||||
let chunks: String = records
|
||||
.iter()
|
||||
.filter(|(message, fields)| {
|
||||
message == "stream chunk" && fields["stage"] == "provider_response"
|
||||
})
|
||||
.map(|(_, fields)| fields["chunk"].as_str().unwrap())
|
||||
.collect();
|
||||
assert_eq!(chunks, SSE_BODY);
|
||||
assert!(!format!("{records:?}").contains("sk-ant"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::at_open(1)]
|
||||
#[case::after_the_first_chunk(2)]
|
||||
#[tokio::test]
|
||||
async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] detach_after: usize) {
|
||||
let upstream = upstream([sse_response()]).await;
|
||||
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), detach_after);
|
||||
|
||||
let outcome = stream_through(&host)
|
||||
.await
|
||||
.expect("a detached stream still completes");
|
||||
|
||||
assert!(matches!(outcome, MessagesOutput::Streamed));
|
||||
assert_eq!(host.seen.into_inner().unwrap().len(), detach_after);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::text_body(ResponseTemplate::new(429).set_body_string("slow down"), "slow down")]
|
||||
#[case::json_envelope(
|
||||
status_response(429, json!({"type": "error", "error": {"type": "rate_limit_error", "message": "slow down"}})),
|
||||
r#"{"type":"error","error":{"type":"rate_limit_error","message":"slow down"}}"#
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn an_upstream_error_fails_the_call_without_opening_the_stream(
|
||||
call: MessagesCall,
|
||||
#[case] response: ResponseTemplate,
|
||||
#[case] body: &str,
|
||||
) {
|
||||
let upstream = upstream([response]).await;
|
||||
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
|
||||
|
||||
let error = stream_through(&host)
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: 429,
|
||||
body: body.into()
|
||||
})
|
||||
);
|
||||
assert!(host.seen.into_inner().unwrap().is_empty());
|
||||
}
|
||||
|
||||
/// The native route relays bytes as they are. Python's synthetic `api_error` for a stream
|
||||
/// that never reaches `message_stop` lives in its SSE wrapper, above this route.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_stream_that_ends_without_message_stop_is_relayed_as_is(call: MessagesCall) {
|
||||
const INCOMPLETE: &str = "event: message_start\ndata: {\"type\":\"message_start\"}\n\n";
|
||||
let upstream =
|
||||
upstream([ResponseTemplate::new(200).set_body_raw(INCOMPLETE, "text/event-stream")]).await;
|
||||
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
|
||||
|
||||
stream_through(&host).await.expect("streamed call succeeds");
|
||||
|
||||
let delivered: Vec<u8> = host
|
||||
.seen
|
||||
.into_inner()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.flat_map(|step| match step {
|
||||
Seen::Deliver(chunk) => chunk.to_vec(),
|
||||
Seen::Open(_) => Vec::new(),
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(delivered, INCOMPLETE.as_bytes());
|
||||
}
|
||||
|
||||
/// Serves one SSE chunk and then holds the connection open without ever finishing.
|
||||
async fn stalling_upstream() -> (String, JoinHandle<()>) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let base = format!("http://{}", listener.local_addr().unwrap());
|
||||
let connection = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let _ = socket.read(&mut request).await;
|
||||
socket
|
||||
.write_all(
|
||||
b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ntransfer-encoding: chunked\r\n\r\n\
|
||||
1f\r\nevent: message_start\ndata: {}\n\n\r\n",
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let _ = socket.read_to_end(&mut Vec::new()).await;
|
||||
});
|
||||
(base, connection)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
|
||||
let (base, connection) = stalling_upstream().await;
|
||||
let host = RecordingStreamHost::new(
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_millis(300)),
|
||||
..streaming(call, base)
|
||||
},
|
||||
usize::MAX,
|
||||
);
|
||||
|
||||
let error = tokio::time::timeout(Duration::from_secs(5), stream_through(&host))
|
||||
.await
|
||||
.expect("the stalled stream gives up within the timeout")
|
||||
.err()
|
||||
.expect("a stalled body fails the call");
|
||||
|
||||
assert!(matches!(error, Error::Transport(_)), "{error:?}");
|
||||
let seen = host.seen.into_inner().unwrap();
|
||||
assert!(
|
||||
matches!(seen.as_slice(), [Seen::Open(_), Seen::Deliver(chunk)] if chunk.as_ref() == b"event: message_start\ndata: {}\n\n"),
|
||||
"the chunk before the stall reached the caller, saw {} ops",
|
||||
seen.len()
|
||||
);
|
||||
tokio::time::timeout(Duration::from_secs(5), connection)
|
||||
.await
|
||||
.expect("timing out closes the upstream connection")
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure_ai("azure_ai")]
|
||||
#[tokio::test]
|
||||
async fn the_sdk_returns_stream_headers_and_every_sse_byte(
|
||||
call: MessagesCall,
|
||||
#[case] provider: &str,
|
||||
) {
|
||||
let upstream = upstream([sse_response()]).await;
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
..streaming(call, upstream.uri())
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let MessagesResponse::Stream { headers, chunks } = response else {
|
||||
panic!("a streaming request returns a stream");
|
||||
};
|
||||
for (name, value) in UPSTREAM_HEADERS {
|
||||
assert!(headers.contains(&(name.into(), value.into())));
|
||||
}
|
||||
let delivered = chunks.try_collect::<Vec<_>>().await.unwrap().concat();
|
||||
assert_eq!(delivered, SSE_BODY.as_bytes());
|
||||
assert_eq!(only_request(&upstream).await.json()["stream"], true);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_sdk_returns_http_errors_before_opening_a_stream(call: MessagesCall) {
|
||||
let upstream = upstream([ResponseTemplate::new(429).set_body_string("slow down")]).await;
|
||||
let error = messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
streaming(call, upstream.uri()),
|
||||
)
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream failure is returned by messages()");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: 429,
|
||||
body: "slow down".into(),
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::before_reading(false)]
|
||||
#[case::after_reading(true)]
|
||||
#[tokio::test]
|
||||
async fn dropping_the_sdk_stream_closes_the_unfinished_upstream(
|
||||
call: MessagesCall,
|
||||
#[case] read_chunk: bool,
|
||||
) {
|
||||
let (base, connection) = stalling_upstream().await;
|
||||
let response = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_secs(30)),
|
||||
..streaming(call, base)
|
||||
},
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("messages() returns before the upstream finishes")
|
||||
.unwrap();
|
||||
|
||||
let MessagesResponse::Stream { mut chunks, .. } = response else {
|
||||
panic!("a streaming request returns a stream");
|
||||
};
|
||||
if read_chunk {
|
||||
let chunk = tokio::time::timeout(Duration::from_secs(5), chunks.next())
|
||||
.await
|
||||
.expect("the first chunk arrives before the upstream finishes")
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(chunk.as_ref(), b"event: message_start\ndata: {}\n\n");
|
||||
}
|
||||
assert!(!connection.is_finished());
|
||||
drop(chunks);
|
||||
tokio::time::timeout(Duration::from_secs(5), connection)
|
||||
.await
|
||||
.expect("dropping the stream closes the upstream connection")
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_sdk_yields_a_body_error_once_after_delivered_chunks(call: MessagesCall) {
|
||||
let (base, connection) = stalling_upstream().await;
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_millis(300)),
|
||||
..streaming(call, base)
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let MessagesResponse::Stream { mut chunks, .. } = response else {
|
||||
panic!("a streaming request returns a stream");
|
||||
};
|
||||
assert_eq!(
|
||||
chunks.next().await.unwrap().unwrap().as_ref(),
|
||||
b"event: message_start\ndata: {}\n\n"
|
||||
);
|
||||
let error = tokio::time::timeout(Duration::from_secs(5), chunks.next())
|
||||
.await
|
||||
.expect("the stalled body times out")
|
||||
.unwrap()
|
||||
.unwrap_err();
|
||||
assert!(matches!(error, Error::Transport(_)), "{error:?}");
|
||||
assert!(chunks.next().await.is_none());
|
||||
tokio::time::timeout(Duration::from_secs(5), connection)
|
||||
.await
|
||||
.expect("the failed stream closes its upstream connection")
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_host_on_anthropic_sse_is_relayed_byte_for_byte(call: MessagesCall) {
|
||||
let upstream = upstream([sse_response()]).await;
|
||||
let host = RecordingStreamHost::new(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
..streaming(call, upstream.uri())
|
||||
},
|
||||
usize::MAX,
|
||||
);
|
||||
|
||||
let outcome = stream_through(&host).await.expect("azure streams");
|
||||
|
||||
assert!(matches!(outcome, MessagesOutput::Streamed));
|
||||
let seen = host.seen.into_inner().unwrap();
|
||||
let delivered: Vec<u8> = seen
|
||||
.iter()
|
||||
.filter_map(|step| match step {
|
||||
Seen::Deliver(chunk) => Some(chunk.to_vec()),
|
||||
Seen::Open(_) => None,
|
||||
})
|
||||
.flatten()
|
||||
.collect();
|
||||
assert_eq!(delivered, SSE_BODY.as_bytes());
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue