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

This commit is contained in:
yucheng 2026-09-25 18:32:47 +00:00
commit bf803cde01
1126 changed files with 50036 additions and 88141 deletions

View file

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

View file

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

View file

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

View file

@ -5,8 +5,8 @@
"CHAT-TOOL-STREAM": "tests/test_litellm/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",
"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/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",

View file

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

View file

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

44
.github/scripts/read_rc_version.py vendored Normal file
View file

@ -0,0 +1,44 @@
#!/usr/bin/env python3
"""Print `version=X.Y.0` from [project].version in pyproject.toml for $GITHUB_OUTPUT.
Usage
-----
python3 read_rc_version.py [path/to/pyproject.toml] >> "$GITHUB_OUTPUT"
Exit code 1 with a `::error::` line on stderr when the version is not an X.Y.0 release.
"""
from __future__ import annotations
import pathlib
import re
import sys
from typing import Final
if sys.version_info >= (3, 11):
import tomllib
else:
import tomli as tomllib
RELEASE_VERSION: Final = re.compile(r"[0-9]+\.[0-9]+\.0")
def read_version(pyproject: pathlib.Path) -> str:
with pyproject.open("rb") as f:
return tomllib.load(f)["project"]["version"]
def main(argv: list[str]) -> int:
pyproject: Final = pathlib.Path(argv[1]) if len(argv) > 1 else pathlib.Path("pyproject.toml")
version: Final = read_version(pyproject)
if RELEASE_VERSION.fullmatch(version) is None:
print( # noqa: T201 # the ::error:: line to stderr is the workflow's failure signal
f"::error::pyproject.toml version {version} is not an X.Y.0 release version", file=sys.stderr
)
return 1
print(f"version={version}") # noqa: T201 # stdout line is appended to $GITHUB_OUTPUT
return 0
if __name__ == "__main__":
sys.exit(main(sys.argv))

View file

@ -13,6 +13,16 @@ on:
have its path existence-checked like any other token.
required: true
type: string
unit-flag:
description: >-
Codecov flag of the `.circleci/tests.yml` job that now owns part of
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: ""
workers:
description: "Number of pytest-xdist workers"
required: false
@ -92,6 +102,7 @@ jobs:
pull-requests: read
outputs:
decision: ${{ steps.changes.outputs.decision }}
has-coverage: ${{ steps.tests.outputs.has-coverage }}
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
@ -160,10 +171,12 @@ jobs:
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Run tests
id: tests
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: ${{ inputs.timeout-minutes }}
env:
TEST_PATH: ${{ inputs.test-path }}
UNIT_FLAG: ${{ inputs.unit-flag }}
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
@ -171,9 +184,18 @@ jobs:
DIST: ${{ inputs.dist }}
COVERAGE_CORE: sysmon
run: |
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
selection="${TEST_PATH}"
if [ -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; nothing to run"
exit 0
fi
pytest_args=()
existing_paths=0
for token in ${TEST_PATH:?}; do
for token in ${selection}; do
case "${token}" in
-*) pytest_args+=("${token}") ;;
*)
@ -187,7 +209,7 @@ jobs:
esac
done
if [ "${existing_paths}" -eq 0 ]; then
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
echo "No path in the selection exists (${selection}); nothing to run"
exit 0
fi
xdist_args=()
@ -209,8 +231,11 @@ jobs:
--cov-config=pyproject.toml
status=$?
set -e
if [ -f coverage.xml ]; then
echo "has-coverage=true" >> "$GITHUB_OUTPUT"
fi
if [ "$status" -eq 5 ]; then
echo "pytest collected no tests from ${TEST_PATH}; passing"
echo "pytest collected no tests from ${selection}; passing"
exit 0
fi
exit "$status"
@ -226,7 +251,7 @@ jobs:
upload-coverage:
name: Upload coverage to Codecov
needs: run
if: always() && needs.run.outputs.decision != 'skip'
if: always() && needs.run.outputs.decision != 'skip' && needs.run.outputs.has-coverage == 'true'
runs-on: ubuntu-latest
permissions:
contents: read

View file

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

66
.github/workflows/create-rc-branch.yml vendored Normal file
View file

@ -0,0 +1,66 @@
name: Create RC Branch
on:
schedule:
- cron: "0 3 * * 5"
timezone: "America/Los_Angeles"
workflow_dispatch:
permissions: {}
jobs:
create-rc-branch:
name: Create RC Branch
if: github.event_name != 'schedule' || github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Require main
env:
REF: ${{ github.ref }}
run: |
if [ "$REF" != "refs/heads/main" ]; then
echo "::error::rc branches are cut from refs/heads/main only, got $REF"
exit 1
fi
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Read release version
id: version
run: python3 .github/scripts/read_rc_version.py >> "$GITHUB_OUTPUT"
- name: Create rc branch
env:
VERSION: ${{ steps.version.outputs.version }}
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
with:
script: |
const branchName = `rc/${process.env.VERSION}`;
const ref = `heads/${branchName}`;
const existing = await github.rest.git.getRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref,
}).catch((error) => {
if (error.status === 404) {
return null;
}
throw error;
});
if (existing !== null) {
core.setFailed(`Branch ${branchName} already exists at ${existing.data.object.sha}; leaving it untouched`);
return;
}
await github.rest.git.createRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `refs/${ref}`,
sha: context.sha,
});
core.info(`Created branch ${branchName} at ${context.sha}`);

View file

@ -180,6 +180,18 @@ jobs:
echo "No changed tests/e2e Python files; skipping."
fi
- name: Run the claude_code harness unit tests
if: steps.changes.outputs.decision != 'skip'
run: |
if ! git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- ':(glob)tests/e2e/claude_code/**/*.py' ':(glob)tests/e2e/*.py' tests/e2e/claude_code/cron_vm/install_claude_code.sh pyproject.toml uv.lock .github/workflows/test-linting.yml | grep -q .; then
echo "No changed claude_code harness files; skipping."
exit 0
fi
retry() { "$@" || { sleep 15; "$@"; } || { sleep 30; "$@"; }; }
CLAUDE_VERSION="$(retry uv run --no-sync python tests/e2e/claude_code/pr_gate_version_resolver.py)"
tests/e2e/claude_code/cron_vm/install_claude_code.sh "$CLAUDE_VERSION" "$RUNNER_TEMP/claude-cli"
PATH="$RUNNER_TEMP/claude-cli:$PATH" uv run --no-sync pytest -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures tests/e2e/claude_code/_*_unit_tests
- name: Check for circular imports
if: steps.changes.outputs.decision != 'skip'
run: |

View file

@ -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/test_litellm/caching/test_redis_cluster_cache.py"
- "tests/test_litellm/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/unit/test_redis.py \
tests/test_litellm/caching/test_redis_connection_pool.py \
tests/test_litellm/caching/test_redis_cluster_cache.py \
tests/test_litellm/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 \

View file

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

View file

@ -31,10 +31,14 @@ concurrency:
# number, so a partially-specified entry would fail the call rather than fall
# back to the default.
#
# tests/proxy_unit_tests keeps its own caller (test-unit-proxy-db.yml): it is
# already a matrix and carries a shard-coverage guard that reads that file by
# name. Folding it in here is a follow-up, together with generalising that guard
# into assert_ci_coverage.py.
# tests/unit/proxy keeps its own caller (test-unit-proxy-db.yml): it is already
# a matrix and carries a shard-coverage guard that reads that file by name.
# Folding it in here is a follow-up, together with generalising that guard into
# assert_ci_coverage.py.
#
# `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 }}
@ -48,7 +52,8 @@ jobs:
include:
- shard: mcp-integration
artifact-name: mcp-integration
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
test-path: "tests/mcp_tests"
unit-flag: mcp-integration
workers: 2
reruns: 0
timeout-minutes: 20
@ -65,10 +70,9 @@ jobs:
- shard: enterprise-routing
artifact-name: enterprise-routing
test-path: >-
tests/test_litellm/enterprise
tests/test_litellm/google_genai
tests/test_litellm/router_utils
tests/test_litellm/router_strategy
unit-flag: enterprise-routing
workers: 2
reruns: 2
timeout-minutes: 20
@ -101,26 +105,13 @@ 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
@ -200,7 +191,7 @@ jobs:
tests/test_litellm/proxy/types_utils
tests/test_litellm/proxy/logging_endpoints
tests/test_litellm/proxy/test_*.py
tests/test_gateway
unit-flag: proxy-infra
workers: 4
reruns: 2
timeout-minutes: 20
@ -208,11 +199,8 @@ jobs:
- shard: caching-local
artifact-name: caching-local
test-path: >-
tests/local_testing/test_cache_preset_key.py
tests/local_testing/test_caching_handler.py
tests/local_testing/test_responses_stream_cache_keys.py
tests/local_testing/test_unit_test_caching.py
test-path: ""
unit-flag: caching-local
workers: 2
reruns: 2
timeout-minutes: 20
@ -220,7 +208,8 @@ jobs:
- shard: proxy-extras
artifact-name: proxy-extras
test-path: "tests/litellm-proxy-extras"
test-path: ""
unit-flag: proxy-extras
workers: 2
reruns: 2
timeout-minutes: 20
@ -228,7 +217,8 @@ jobs:
- shard: enterprise-package
artifact-name: enterprise-package
test-path: "tests/enterprise"
test-path: ""
unit-flag: enterprise-package
workers: 4
reruns: 2
timeout-minutes: 20
@ -239,7 +229,7 @@ jobs:
test-path: >-
tests/test_litellm/responses
tests/test_litellm/caching
tests/test_litellm/types
unit-flag: responses-caching-types
workers: 2
reruns: 2
timeout-minutes: 20
@ -247,6 +237,7 @@ jobs:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
unit-flag: ${{ matrix.unit-flag || '' }}
workers: ${{ matrix.workers }}
reruns: ${{ matrix.reruns }}
timeout-minutes: ${{ matrix.timeout-minutes }}

View file

@ -96,6 +96,7 @@ Follow these coding conventions for new/updated code (a three-line fix in a lega
- No mutation; don't reassign variables, global or local. Instead of mutable lists and dicts, prefer tuples, frozen dataclasses (with slots=True), `MappingProxyType`, etc.
- Annotate every variable with `: Final` (LIT010). Unpacking and walrus targets cannot carry the annotation, so they are implicitly final. Don't rebind them. Never rebind or mutate function parameters (LIT011); `self`/`cls` attribute stores are the exception. If rebinding or in-place mutation is truly unavoidable, suppress with `# rebind-ok: <reason>`
- Qualify every TypedDict field with `ReadOnly[...]` (LIT012), which nests freely with `Required` / `NotRequired` / `Annotated` in any order. If making the key writable is truly unavoidable, suppress with `# writable-ok: <reason>`
- Comprehensions take at most one `for` clause and one `if` clause (LIT014); split stacked clauses into a helper generator, a named intermediate, or a plain loop. Suppress with `# comprehension-ok: <reason>` only when unavoidable
- Use dependency injection
- Fully typed; no `Any` or coarse types like `dict[str, Any]` or just `dict`. Every function parameter must be strongly typed
- Use tagged unions + match

View file

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

View file

@ -1697,6 +1697,63 @@
"title": "litellm_video_duration_seconds_metric rate",
"type": "timeseries"
},
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"description": "Share of the provider's bill LiteLLM captured as spend over the scheduled capture-rate check's window (needs general_settings.spend_capture_rate_check); NaN while no rate is available",
"fieldConfig": {
"defaults": {
"color": {
"mode": "palette-classic"
},
"custom": {
"drawStyle": "line",
"fillOpacity": 10,
"lineWidth": 1,
"showPoints": "never",
"spanNulls": false
},
"unit": "percentunit"
},
"overrides": []
},
"gridPos": {
"h": 8,
"w": 12,
"x": 12,
"y": 107
},
"id": 111,
"options": {
"legend": {
"calcs": [],
"displayMode": "list",
"placement": "bottom",
"showLegend": true
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "max by (api_provider) (litellm_spend_capture_rate)",
"legendFormat": "{{api_provider}}",
"range": true,
"refId": "A"
}
],
"title": "litellm_spend_capture_rate",
"type": "timeseries"
},
{
"collapsed": false,
"gridPos": {

View file

@ -1,6 +1,6 @@
# LiteLLM All Prometheus Metrics dashboard
Every `litellm_*` metric family the proxy can expose on `/metrics` (134 families across 95 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
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

View file

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

View file

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

View file

@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Final, List, Literal, Optional, Protocol, Tupl
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import (
CLI_SESSION_KEY_PREFIX,
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS,
MAX_OBJECTS_PER_POLL_CYCLE,
)
@ -147,10 +148,12 @@ class CheckBatchCost:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}")
return {}
async def _get_key_alias(self, batch_id: str, api_key: str | None) -> str | None:
async def _get_key_alias(self, batch_id: str, api_key: str | None, created_by: str | None) -> str | None:
"""Resolve the creating virtual key's alias from its hashed token."""
if not api_key:
return None
if created_by and api_key == f"{CLI_SESSION_KEY_PREFIX}-{created_by}":
return api_key
try:
key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table(
self.prisma_client
@ -231,7 +234,7 @@ class CheckBatchCost:
**(await self._get_user_info(batch_id, job.created_by)),
}
key_alias = await self._get_key_alias(batch_id, api_key)
key_alias = await self._get_key_alias(batch_id, api_key, job.created_by)
if key_alias is not None:
metadata["user_api_key_alias"] = key_alias
team_alias = await self._get_team_alias(team_id)

View file

@ -50,6 +50,7 @@ from litellm.proxy._types import (
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.openai_files_endpoints.common_utils import (
BATCH_CREATE_HIDDEN_PARAM,
FILE_LIST_CONTINUATION_CHUNK_SIZE,
@ -359,7 +360,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
from prisma import Json
api_key = user_api_key_dict.api_key or None
api_key = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict) or None
attribution_columns = (
{
**({"api_key": api_key} if api_key is not None else {}),

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.70"
version = "0.1.71"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.70"
version = "0.1.71"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

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

View file

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

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.101"
version = "0.4.102"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.101"
version = "0.4.102"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -8,3 +8,11 @@
- Split a mixed test file along that line instead of widening visibility to move it
- A test for another crate's item belongs in that crate, not in a downstream one
- Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests/<name>/mod.rs` or `tests/<subject>/support.rs` so it is not picked up as a test crate of its own
## 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
- 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

345
litellm-rust/Cargo.lock generated
View file

@ -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"
@ -897,6 +906,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 +929,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"
@ -1458,6 +1479,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 +1572,15 @@ 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 = "equivalent"
version = "1.0.2"
@ -1622,6 +1663,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 +1690,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 +1722,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 +1889,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 +2734,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"
@ -2995,6 +3122,7 @@ dependencies = [
"tokio-tungstenite",
"url",
"veil",
"wiremock",
]
[[package]]
@ -3013,6 +3141,15 @@ dependencies = [
"url",
]
[[package]]
name = "litellm-coroutine"
version = "0.1.0"
dependencies = [
"rstest",
"thiserror 2.0.19",
"tokio",
]
[[package]]
name = "litellm-cost"
version = "0.1.0"
@ -3040,6 +3177,7 @@ name = "litellm-host"
version = "0.1.0"
dependencies = [
"litellm-auth",
"litellm-coroutine",
"rstest",
"serde_json",
"tokio",
@ -3049,6 +3187,7 @@ dependencies = [
name = "litellm-host-python"
version = "0.1.0"
dependencies = [
"bytes",
"futures-util",
"litellm-host",
"pyo3",
@ -3115,14 +3254,14 @@ dependencies = [
name = "litellm-model-catalog"
version = "0.1.0"
dependencies = [
"criterion",
"indexmap 2.14.0",
"litellm-model-catalog",
"jsonschema",
"rstest",
"schemars 1.2.2",
"serde",
"serde_json",
"thiserror 2.0.19",
"time",
]
[[package]]
@ -3344,6 +3483,27 @@ dependencies = [
"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"
@ -3493,6 +3653,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"
@ -3588,6 +3754,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"
@ -3608,6 +3788,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"
@ -3632,6 +3818,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"
@ -4447,6 +4654,23 @@ dependencies = [
"syn 3.0.0",
]
[[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]]
name = "regex"
version = "1.13.1"
@ -5033,6 +5257,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"
@ -5364,6 +5597,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"
@ -5633,6 +5877,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"
@ -5649,9 +5917,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]]
@ -5660,9 +5928,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"
@ -5930,6 +6204,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"
@ -6010,6 +6290,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"
@ -6408,6 +6698,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"
@ -6470,6 +6766,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"
@ -6595,6 +6901,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"
@ -6606,3 +6929,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",
]

View file

@ -12,6 +12,7 @@ repository = "https://github.com/BerriAI/litellm"
litellm-tracing = { path = "crates/tracing" }
tracing = "0.1"
litellm-core = { path = "crates/core" }
litellm-coroutine = { path = "crates/coroutine" }
litellm-host = { path = "crates/host" }
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
litellm-framing = { path = "crates/framer" }
@ -80,6 +81,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"

View file

@ -3,8 +3,8 @@
//! lifetime. No other callback host has that obligation, which is why nothing outside
//! this crate holds them.
use litellm_host::{machine::Machine, route::Route};
use litellm_host_python::{RouteHost, lookup, run_call};
use litellm_host::{machine::Machine, protocol::Protocol};
use litellm_host_python::{ProtocolHost, lookup, run_call};
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,
@ -63,25 +63,25 @@ impl PublicCall {
}
}
/// Runs one native call under the legacy `Logging` contract: the route host projects from
/// 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.
pub fn run_legacy_call<H, M>(
py: Python<'_>,
surface: LegacySurface,
call: PublicCall,
machine: M,
route: H,
host: H,
asynchronous: bool,
) -> PyResult<Py<PyAny>>
where
H: RouteHost + 'static,
M: Machine<Route = H::Route, Complete = <H::Route as Route>::Response> + 'static,
H: ProtocolHost + 'static,
M: Machine<Protocol = H::Protocol, Complete = <H::Protocol as Protocol>::Response> + 'static,
{
let arguments = call.kwargs.clone_ref(py);
run_call(
py,
machine,
route,
host,
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
arguments,
asynchronous,

View file

@ -40,3 +40,4 @@ litellm-auth-gcp.workspace = true
litellm-llms = { workspace = true, features = ["test-support"] }
rstest.workspace = true
rstest_reuse.workspace = true
wiremock = "0.6.5"

View file

@ -85,3 +85,29 @@ pub(super) async fn outbound_request(
other => other,
})
}
#[cfg(test)]
mod tests {
use super::{Error, as_response_error};
#[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);
}
}

View file

@ -736,248 +736,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, .. })
));
}
}
}

View file

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

View file

@ -17,8 +17,10 @@ pub(super) async fn send(
body: &Value,
timeout: Option<Duration>,
) -> Result<reqwest::Response, Error> {
let encoded = serde_json::to_vec(body)
.map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?;
let builder = headers.iter().fold(
http_client().post(url).json(body),
http_client().post(url).body(encoded),
|builder, (key, value)| builder.header(key, value),
);
let builder = match timeout {

View file

@ -1,16 +1,16 @@
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 litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
host::{Demand, Host},
machine::{HostChannel, MachineFault, RouteMachine},
route::Route,
machine::{CallMachine, HostChannel, MachineFault},
protocol::Protocol,
};
use litellm_secrets::source::SecretSource;
use litellm_types::{
@ -21,22 +21,12 @@ use serde_json::{Map, Value};
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;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessagesOp {
ProjectRequest,
}
pub enum MessagesOpResult {
Request(Box<MessagesCall>),
}
/// The caller's request as the host projects it.
pub struct MessagesCall {
pub model: String,
@ -62,15 +52,20 @@ 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 Route for Messages {
impl Protocol for Messages {
type Response = MessagesOutput;
type Error = Error;
type Op = MessagesOp;
type OpResult = MessagesOpResult;
type Projection = MessagesCall;
type Op = Infallible;
type Chunk = Bytes;
type StreamHead = ();
type StreamHead = MessagesStreamHead;
}
impl From<MachineFault> for Error {
@ -78,26 +73,12 @@ impl From<MachineFault> for Error {
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "messages host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("messages {message}"),
MachineFault::Mismatch => "invalid messages host operation result".into(),
})
}
}
pub type MessagesHost = HostChannel<Messages>;
pub type MessagesMachine = RouteMachine<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,
}
}
pub type MessagesMachine = CallMachine<Messages>;
/// The in-process host for a request already in hand. It answers projection once and
/// observes nothing.
@ -114,30 +95,28 @@ impl LocalMessagesHost {
}
impl Host<Messages> for LocalMessagesHost {
async fn route(&self, op: MessagesOp) -> Result<MessagesOpResult, Error> {
match op {
MessagesOp::ProjectRequest => self
.call
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.map(|call| MessagesOpResult::Request(Box::new(call)))
.ok_or_else(|| {
Error::InvalidRequest("messages request was already projected".into())
}),
}
async fn project(&self) -> Result<MessagesCall, Error> {
self.call
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.ok_or_else(|| Error::InvalidRequest("messages request was already projected".into()))
}
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
match op {}
}
}
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
CallMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
}
async fn execute(
host: MessagesHost,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesOutput, Error> {
let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?;
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?;
@ -163,8 +142,11 @@ async fn execute(
model: request.model.clone(),
custom_llm_provider: request.provider.clone(),
optional_params: Value::Object(
call.body
.iter()
request
.body
.as_object()
.into_iter()
.flatten()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages"))
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
@ -204,7 +186,14 @@ async fn relay(
host: &MessagesHost,
mut response: reqwest::Response,
) -> Result<MessagesOutput, Error> {
if host.open(()).await? == Demand::Detached {
let head = MessagesStreamHead {
headers: response
.headers()
.iter()
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string())))
.collect(),
};
if host.open(head).await? == Demand::Detached {
return Ok(MessagesOutput::Streamed);
}
while let Some(chunk) = response.chunk().await.map_err(network)? {

View file

@ -23,9 +23,6 @@ pub fn prepare_document(input: OcrDocumentInput) -> Result<OcrDocument, Error> {
file_name.as_deref(),
mime_type.as_deref(),
)?),
OcrDocumentInput::HostReader { .. } => Err(Error::InvalidRequest(
"OCR file reader was not read by the host".into(),
)),
}
}
@ -207,7 +204,7 @@ mod tests {
}
#[test]
fn byte_documents_are_encoded_and_host_readers_must_be_read_first() {
fn byte_documents_are_encoded() {
assert_eq!(
prepare_document(OcrDocumentInput::Bytes {
bytes: b"abc".as_slice().into(),
@ -217,7 +214,6 @@ mod tests {
.unwrap(),
document("data:application/pdf;base64,YWJj")
);
assert!(prepare_document(OcrDocumentInput::HostReader { mime_type: None }).is_err());
}
#[test]
@ -248,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)
);
}
}

View file

@ -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())
})
}
}

View file

@ -70,20 +70,203 @@ 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(reqwest::Client::new(), reqwest::Client::new())
}
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 {

View file

@ -5,15 +5,15 @@ use litellm_llms::{
transformation::TextractDetectTextConfig,
},
azure_ai::ocr::{
cohere_parse_transformation::AzureAICohereParseConfig,
cohere_parse_transformation::{AZURE_COHERE_PARSE_PATH, AzureAICohereParseConfig},
document_intelligence::transformation::AzureDocumentIntelligenceOcrConfig,
transformation::AzureAiOcrConfig,
transformation::{AZURE_AI_OCR_PATH, AzureAiOcrConfig},
},
base_llm::ocr::{
error::Error,
handler::{self, CallHooks, OcrClient},
transformation::{
BaseOcrConfig, LiteLLMOcrResponse, OcrCredentialInputs, OcrDocument,
BaseOcrConfig, LiteLLMOcrResponse, OcrCredentialInputs, OcrDocument, OcrResponseFormat,
PreparedOcrRequest, ResolvedOcrCredentials,
},
},
@ -157,6 +157,36 @@ pub fn get_health_check_document(
.get_health_check_document())
}
/// Normalize a relayed Azure AI response into the LiteLLM OCR shape when
/// `endpoint` is the OCR route of the model's resolved config.
pub fn passthrough_response(
model: &str,
endpoint: &str,
body: &[u8],
) -> Result<Option<LiteLLMOcrResponse>, Error> {
let (model, config) = resolve_provider_config(model, Some("azure_ai"))?;
let segments: Vec<&str> = endpoint
.split('/')
.filter(|segment| !segment.is_empty())
.collect();
let is_ocr_endpoint = match config {
OcrConfigKind::AzureAi => segments == AZURE_AI_OCR_PATH,
OcrConfigKind::AzureCohere => segments == AZURE_COHERE_PARSE_PATH,
OcrConfigKind::AzureDocumentIntelligence => {
segments == AzureDocumentIntelligenceOcrConfig::analyze_path(&model)?
}
other => {
let provider: &'static str = other.provider().into();
return Err(Error::InvalidProvider(provider.to_owned()));
}
};
if !is_ocr_endpoint {
return Ok(None);
}
with_config!(config, config => config.transform_ocr_response(&model, body, OcrResponseFormat::Litellm))
.map(Some)
}
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
#[strum(serialize_all = "snake_case")]
pub(crate) enum OcrProvider {
@ -520,4 +550,68 @@ mod tests {
assert!(matches!(&error, Error::InvalidProvider(provider) if provider == "not_a_provider"));
assert_eq!(error.http_status_code(), Some(400));
}
#[rstest]
#[case("azure_ai/mistral-document-ai-2512", "providers/mistral/azure/ocr")]
#[case("azure_ai/mistral-document-ai-2512", "/providers/mistral/azure/ocr/")]
#[case("azure_ai/Cohere-parse-v5", "providers/cohere/v2/parse")]
#[case(
"azure_ai/doc-intelligence/prebuilt-layout",
"documentintelligence/documentModels/prebuilt-layout:analyze"
)]
fn passthrough_response_recognizes_the_resolved_config_ocr_route(
#[case] model: &str,
#[case] endpoint: &str,
) {
assert!(passthrough_response(model, endpoint, b"not json").is_err());
}
#[rstest]
#[case("azure_ai/mistral-document-ai-2512", "models/info")]
#[case("azure_ai/mistral-document-ai-2512", "providers/cohere/v2/parse")]
#[case("azure_ai/Cohere-parse-v5", "providers/mistral/azure/ocr")]
#[case(
"azure_ai/doc-intelligence/prebuilt-layout",
"documentintelligence/documentModels/prebuilt-read:analyze"
)]
fn passthrough_response_skips_other_routes(#[case] model: &str, #[case] endpoint: &str) {
assert!(
passthrough_response(model, endpoint, b"not json")
.unwrap()
.is_none()
);
}
#[test]
fn passthrough_response_normalizes_the_mistral_body() {
let body = br#"{
"pages": [{"index": 0, "markdown": "page one"}, {"index": 1, "markdown": "page two"}],
"model": "mistral-document-ai-2512",
"usage_info": {"pages_processed": 2}
}"#;
let json = passthrough_response(
"azure_ai/mistral-document-ai-2512",
"providers/mistral/azure/ocr",
body,
)
.unwrap()
.unwrap()
.into_json();
assert_eq!(json["usage_info"]["pages_processed"], 2);
assert_eq!(json["pages"][0]["markdown"], "page one");
}
#[test]
fn passthrough_response_counts_cohere_billed_pages() {
let body = br#"{"id": "parse-1", "pages": [], "meta": {"billed_units": {"pages": 3}}}"#;
let json = passthrough_response(
"azure_ai/Cohere-parse-v5",
"providers/cohere/v2/parse",
body,
)
.unwrap()
.unwrap()
.into_json();
assert_eq!(json["usage_info"]["pages_processed"], 3);
}
}

File diff suppressed because it is too large Load diff

View file

@ -25,9 +25,6 @@ pub enum OcrDocumentInput {
file_name: Option<String>,
mime_type: Option<String>,
},
HostReader {
mime_type: Option<String>,
},
}
impl From<OcrDocument> for OcrDocumentInput {
@ -45,12 +42,6 @@ impl From<PathBuf> for OcrDocumentInput {
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct OcrFileContent {
pub bytes: Bytes,
pub file_name: Option<String>,
}
/// Caller-supplied connection overrides for a [`LiteLLMOcrRequest`], in the
/// shape hosts receive them: JSON-ish headers, optional timeout, optional
/// credentials, and per-field provenance in `input_sources`.

View file

@ -1,50 +1,250 @@
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([
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 = audio_transcription(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}");
audio_transcription(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();
audio_transcription(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 = audio_transcription(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 = audio_transcription(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 = audio_transcription(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 = audio_transcription(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 = audio_transcription(AudioTranscriptionRequest {
api_base: Some(&base),
..request
})
.await
.expect_err("an unreadable body fails");
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
}

View file

@ -0,0 +1,320 @@
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 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}}"#;
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 = chat_completions(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();
chat_completions(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 = chat_completions(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 = chat_completions(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 = chat_completions(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 = chat_completions(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 = chat_completions(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 = chat_completions(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());
}

View file

@ -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"));
}

View file

@ -0,0 +1,210 @@
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(messages_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 mut body = call.body.clone();
body.insert("stream".into(), json!(true));
let host =
RecordingHost::passthrough(authenticated(MessagesCall { body, ..call }, 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 body: Map<String, Value> = call
.body
.clone()
.into_iter()
.chain([("temperature".to_string(), json!(0.2))])
.collect();
let host = RecordingHost::passthrough(authenticated(
MessagesCall {
body,
shaping: MessagesShaping {
capabilities: AnthropicModelCapabilities {
supports_sampling_params: false,
..AnthropicModelCapabilities::default()
},
drop_params: true,
..MessagesShaping::default()
},
..call
},
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}));
}

View file

@ -0,0 +1,95 @@
use std::{sync::Arc, time::Duration};
use litellm_core::messages::{
Error,
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
types::MessagesShaping,
};
use litellm_types::llms::anthropic_messages::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 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 {
model: MODEL.into(),
body: object(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(),
)
}
async fn run_with(
secrets: Arc<RecordingSecrets>,
call: MessagesCall,
) -> Result<MessagesOutput, Error> {
litellm_host::run::run(messages_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"),
}
}

View file

@ -0,0 +1,675 @@
use litellm_llms::anthropic::common_utils::{
ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_OAUTH_BETA_HEADER, AnthropicModelCapabilities,
SupportedEffortTiers, beta,
};
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 {
model: model.into(),
custom_llm_provider: provider.map(Into::into),
api_key: Some("sk".into()),
api_base: Some(format!("{}{base_suffix}", upstream.uri())),
..call
})
.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 {
model: model.into(),
custom_llm_provider: provider.map(Into::into),
api_key: Some("sk".into()),
api_base: Some(UNREACHABLE_BASE.into()),
..call
})
.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: object(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;
let mut body = call.body.clone();
body.insert("temperature".into(), json!(0.5));
body.insert("top_k".into(), json!(3));
run_message(MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
body,
shaping: MessagesShaping {
additional_drop_params: vec!["temperature".into()],
..MessagesShaping::default()
},
..call
})
.await;
let sent = only_request(&upstream).await.json();
assert_eq!(sent.get("temperature"), None);
assert_eq!(sent["top_k"], 3);
}
fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall {
let body: Map<String, Value> = call.body.into_iter().chain(object(fields)).collect();
MessagesCall { body, ..call }
}
fn sent_betas(request: &wiremock::Request) -> Vec<String> {
let [header] = <[&str; 1]>::try_from(request.header_values("anthropic-beta"))
.unwrap_or_else(|values| panic!("expected one anthropic-beta header, got {values:?}"));
header
.split(',')
.map(str::trim)
.map(str::to_string)
.collect()
}
#[rstest]
#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
#[case::fast_mode(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
#[case::compaction(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
#[case::context_management_edits(
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]}}),
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
)]
#[case::per_message_output_config(
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
&[beta::PER_TURN_CONTROL_2026_07_01]
)]
#[case::advisor_tool(
json!({"tools": [{"type": ANTHROPIC_ADVISOR_TOOL_TYPE, "name": "advisor", "model": MODEL}]}),
&[beta::ADVISOR_TOOL_2026_03_01]
)]
#[case::several_features_at_once(
json!({"speed": "fast", "output_format": {"type": "json_schema"}}),
&[beta::STRUCTURED_OUTPUT, beta::FAST_MODE_2026_02_01]
)]
#[tokio::test]
async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
call: MessagesCall,
#[case] fields: Value,
#[case] features: &[&str],
) {
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 mut expected: Vec<String> = features
.iter()
.map(|feature| feature.to_string())
.chain(["caller-beta-2025-01-01".to_string()])
.collect();
expected.sort();
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), [ANTHROPIC_OAUTH_BETA_HEADER]);
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: capabilities.clone(),
drop_params,
..MessagesShaping::default()
},
body: call.body.clone(),
custom_llm_provider: call.custom_llm_provider.clone(),
extra_headers: None,
provider_specific_header: None,
model: call.model.clone(),
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 {
model: model.into(),
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
})
.await;
assert_eq!(only_request(&upstream).await.json()["model"], sent_model);
}

View file

@ -0,0 +1,221 @@
use litellm_core::messages::{messages, types::MessagesRequest};
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!(error.is_response(), "{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:?}");
}
fn facade_request(body: Value, api_base: &str) -> MessagesRequest<'_> {
MessagesRequest {
model: MODEL,
body,
api_key: Some("sk-ant"),
api_base: Some(api_base),
custom_llm_provider: Some("anthropic"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
}
}
#[tokio::test]
async fn the_facade_runs_the_route_in_process() {
let upstream = upstream([message_response()]).await;
let base = upstream.uri();
let message = messages(facade_request(
json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}),
&base,
))
.await
.expect("messages request succeeds");
assert_eq!(message.id, "msg_1");
assert_eq!(
only_request(&upstream).await.header("x-api-key"),
Some("sk-ant")
);
}
#[tokio::test]
async fn the_facade_rejects_a_body_that_is_not_an_object() {
let error = messages(facade_request(json!([]), UNREACHABLE_BASE))
.await
.expect_err("a non-object body is rejected");
assert_eq!(
error,
Error::InvalidRequest("messages body must be an object".into())
);
}

View 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));
}

View file

@ -0,0 +1,266 @@
use std::{convert::Infallible, sync::Mutex};
use bytes::Bytes;
use litellm_core::messages::route::{Messages, MessagesStreamHead};
use litellm_host::host::{Demand, Host};
use rstest::rstest;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
};
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),
}
/// 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 {
let mut body = call.body.clone();
body.insert("stream".into(), json!(true));
MessagesCall {
api_key: Some("sk-ant".into()),
api_base: Some(api_base),
body,
..call
}
}
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(messages_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]
#[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 {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
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();
std::future::pending::<()>().await;
});
base
}
#[rstest]
#[tokio::test]
async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
let base = 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()
);
}
#[rstest]
#[tokio::test]
async fn streaming_is_refused_for_providers_that_cannot_stream(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 error = stream_through(&host)
.await
.err()
.expect("azure streaming is refused");
assert_eq!(
error,
Error::Unsupported("streaming messages for this provider")
);
assert!(received(&upstream).await.is_empty());
}

View file

@ -0,0 +1,173 @@
use std::{collections::BTreeMap, time::SystemTime};
use litellm_auth_aws::{Credentials, aws_signature_headers, sign_post};
use rstest::rstest;
use time::{PrimitiveDateTime, format_description};
use wiremock::Request;
use super::*;
const ACCESS_KEY_ID: &str = "AKIDEXAMPLE";
const SECRET_ACCESS_KEY: &str = "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY";
const DETECT: &str = "aws_textract/detect-document-text";
const ANALYZE: &str = "aws_textract/analyze-document";
fn textract_request(model: &str, base: &str) -> LiteLLMOcrRequest {
ocr_request_with_document(
model,
&format!("{base}/"),
json!({"type": "image_url", "image_url": "data:image/png;base64,b3JpZ2luYWw="}),
json!({
"aws_access_key_id": ACCESS_KEY_ID,
"aws_secret_access_key": SECRET_ACCESS_KEY,
"aws_region_name": "eu-west-1"
}),
)
}
fn textract_response() -> ResponseTemplate {
json_response(json!({
"DocumentMetadata": {"Pages": 1},
"Blocks": [{"BlockType": "PAGE"}, {"BlockType": "LINE", "Text": "Invoice 12345"}]
}))
}
/// Recomputes SigV4 over the request the upstream received, at the time the client claimed.
fn expected_authorization(url: &str, sent: &Request) -> String {
let format =
format_description::parse_borrowed::<2>("[year][month][day]T[hour][minute][second]Z")
.unwrap();
let signed_at: SystemTime =
PrimitiveDateTime::parse(sent.header("x-amz-date").unwrap(), &format)
.unwrap()
.assume_utc()
.into();
let headers: BTreeMap<String, String> = ["content-type", "x-amz-target"]
.into_iter()
.map(|name| (name.to_string(), sent.header(name).unwrap().to_string()))
.collect();
sign_post(
url,
&sent.body,
&aws_signature_headers(&headers),
"eu-west-1",
"textract",
&Credentials::new(ACCESS_KEY_ID, SECRET_ACCESS_KEY, None, None, "test"),
signed_at,
)
.unwrap()["Authorization"]
.clone()
}
/// The recorded URL names wiremock's host, not the address the client signed for.
fn assert_signed(upstream: &MockServer, sent: &Request) {
let url = format!("{}/", upstream.uri());
assert_eq!(
sent.header("authorization"),
Some(expected_authorization(&url, sent).as_str())
);
}
#[tokio::test]
async fn detect_document_text_is_signed_and_lines_become_the_page() {
let upstream = upstream([textract_response()]).await;
let response = perform_with(LocalOcrHost::new(textract_request(DETECT, &upstream.uri())))
.await
.unwrap();
let sent = only_request(&upstream).await;
assert_eq!(
sent.header("x-amz-target"),
Some("Textract.DetectDocumentText")
);
assert_eq!(
sent.header("content-type"),
Some("application/x-amz-json-1.1")
);
assert_eq!(sent.json(), json!({"Document": {"Bytes": "b3JpZ2luYWw="}}));
assert_signed(&upstream, &sent);
assert_eq!(response.pages[0].markdown, "Invoice 12345");
assert_eq!(response.usage_info.unwrap().pages_processed, Some(1));
}
#[tokio::test]
async fn a_body_rewritten_by_before_send_is_what_gets_signed_and_sent() {
let upstream = upstream([textract_response()]).await;
let host = LocalOcrHost::new(textract_request(DETECT, &upstream.uri())).with_before_send(
|mut wire, _| {
assert!(
!wire
.headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("authorization")),
"the hook ran after signing"
);
wire.body["Document"]["Bytes"] = Value::from("cmVkYWN0ZWQ=");
Ok(wire)
},
);
perform_with(host).await.unwrap();
let sent = only_request(&upstream).await;
assert_eq!(sent.json(), json!({"Document": {"Bytes": "cmVkYWN0ZWQ="}}));
assert_signed(&upstream, &sent);
}
#[tokio::test]
async fn analyze_document_asks_for_layout_and_tables_and_returns_markdown() {
let upstream = upstream([json_response(json!({
"DocumentMetadata": {"Pages": 1},
"Blocks": [
{"Id": "l1", "BlockType": "LINE", "Text": "Quarterly Report"},
{"Id": "t", "BlockType": "LAYOUT_TITLE",
"Relationships": [{"Type": "CHILD", "Ids": ["l1"]}]}
]
}))])
.await;
let response = perform_with(LocalOcrHost::new(textract_request(
ANALYZE,
&upstream.uri(),
)))
.await
.unwrap();
let sent = only_request(&upstream).await;
assert_eq!(
sent.header("x-amz-target"),
Some("Textract.AnalyzeDocument")
);
assert_eq!(sent.json()["FeatureTypes"], json!(["LAYOUT", "TABLES"]));
assert_signed(&upstream, &sent);
assert_eq!(response.pages[0].markdown, "# Quarterly Report");
}
#[rstest]
#[case::detect(DETECT)]
#[case::analyze(ANALYZE)]
#[tokio::test]
async fn a_multi_page_rejection_reaches_the_caller_with_the_single_page_limit(#[case] model: &str) {
let upstream = upstream([status_response(
400,
json!({
"__type": "UnsupportedDocumentException",
"Message": "Request has unsupported document format"
}),
)])
.await;
let error = perform_with(LocalOcrHost::new(textract_request(model, &upstream.uri())))
.await
.unwrap_err();
let Error::Provider { status, body, .. } = error else {
panic!("expected a provider error, got {error:?}");
};
assert_eq!(status, 400);
assert!(
body.contains("multi-page documents are not supported"),
"{body}"
);
}

View file

@ -0,0 +1,270 @@
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use litellm_auth::{
ResolvedCredential, SecretValue, TokenFuture, TokenProvider, TokenProviderHandle,
};
use rstest::rstest;
use super::*;
#[derive(Debug)]
struct CountingToken {
token: fn(usize) -> String,
calls: AtomicUsize,
}
impl CountingToken {
fn new(token: fn(usize) -> String) -> Arc<Self> {
Arc::new(Self {
token,
calls: AtomicUsize::new(0),
})
}
fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
}
impl TokenProvider for CountingToken {
fn acquire(&self) -> TokenFuture<'_> {
let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
let token = SecretValue::new((self.token)(call));
Box::pin(async move {
Ok(ResolvedCredential::AccessToken {
token,
expires_on: None,
})
})
}
}
fn numbered_token(call: usize) -> String {
format!("callback-{call}")
}
fn azure_request(
provider: &Arc<CountingToken>,
api_base: Option<&str>,
api_key: Option<&str>,
extra_headers: Value,
optional_params: Value,
) -> LiteLLMOcrRequest {
let wire = serde_json::from_value(json!({
"model": "azure_ai/mistral-ocr-latest",
"document": {"type": "document_url", "document_url": INLINE_PDF},
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": null,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": 2.0
}))
.unwrap();
let mut request = decode_request(wire).unwrap();
request.azure_ad_token_provider = Some(TokenProviderHandle::new(provider.clone()));
request
}
fn ocr_page() -> ResponseTemplate {
json_response(json!({"pages": [{"index": 0, "markdown": "hello"}]}))
}
#[tokio::test]
async fn mistral_on_azure_sends_the_prepared_bearer_and_the_mistral_body() {
let upstream = upstream([json_response(json!({
"pages": [{"index": 0, "markdown": "hello"}],
"usage_info": {"pages_processed": 1}
}))])
.await;
let request = with_headers(
without_api_key(ocr_request(
"azure_ai/model",
&upstream.uri(),
json!({"include_image_base64": true}),
)),
&[("Authorization", "Bearer python-prepared-token")],
);
let result = perform(request).await.unwrap();
assert_eq!(result.pages[0].markdown, "hello");
let sent = only_request(&upstream).await;
assert_eq!(sent.url.path(), "/providers/mistral/azure/ocr");
assert_eq!(
sent.header("authorization"),
Some("Bearer python-prepared-token")
);
assert_eq!(
sent.json(),
json!({
"model": "model",
"document": {"type": "document_url", "document_url": INLINE_PDF},
"include_image_base64": true
})
);
}
#[tokio::test]
async fn a_static_entra_token_becomes_the_bearer() {
let upstream = upstream([pages_response()]).await;
let request = without_api_key(ocr_request(
"azure_ai/model",
&upstream.uri(),
json!({"azure_ad_token": "rust-owned-token"}),
));
perform(request).await.unwrap();
assert_eq!(
only_request(&upstream).await.header("authorization"),
Some("Bearer rust-owned-token")
);
}
#[tokio::test]
async fn a_guardrail_that_swaps_in_a_remote_document_is_rejected() {
let host = LocalOcrHost::new(ocr_request("azure_ai/model", UNREACHABLE_BASE, json!({})))
.with_before_send(|mut wire, _| {
wire.body["document"] = json!({
"type": "document_url",
"document_url": "https://example.com/not-inline.pdf"
});
Ok(wire)
});
let error = perform_with(host).await.unwrap_err();
assert!(error.to_string().contains("data URI"), "{error}");
}
#[tokio::test]
async fn the_token_provider_is_the_bearer_and_is_acquired_for_each_request() {
let provider = CountingToken::new(numbered_token);
let upstream = upstream([ocr_page(), ocr_page()]).await;
let base = upstream.uri();
for _ in 0..2 {
perform(azure_request(
&provider,
Some(&base),
None,
Value::Null,
json!({}),
))
.await
.unwrap();
}
assert_eq!(provider.calls(), 2);
let authorizations: Vec<String> = received(&upstream)
.await
.iter()
.map(|request| {
request
.header("authorization")
.unwrap_or_default()
.to_string()
})
.collect();
assert_eq!(authorizations, ["Bearer callback-1", "Bearer callback-2"]);
}
#[rstest]
#[case::api_key_skips_provider(Some("resource-key"), Value::Null, json!({}), "Bearer resource-key", 0)]
#[case::provider_beats_static_token(
None,
Value::Null,
json!({"azure_ad_token": "static-token"}),
"Bearer callback-1",
1
)]
#[case::header_wins_on_the_wire_but_provider_still_runs(
None,
json!({"Authorization": "Bearer override"}),
json!({}),
"Bearer override",
1
)]
#[tokio::test]
async fn credential_precedence(
#[case] api_key: Option<&str>,
#[case] extra_headers: Value,
#[case] optional_params: Value,
#[case] expected_authorization: &str,
#[case] expected_calls: usize,
) {
let provider = CountingToken::new(numbered_token);
let upstream = upstream([ocr_page()]).await;
perform(azure_request(
&provider,
Some(&upstream.uri()),
api_key,
extra_headers,
optional_params,
))
.await
.unwrap();
assert_eq!(provider.calls(), expected_calls);
assert_eq!(
only_request(&upstream).await.header_values("authorization"),
[expected_authorization]
);
}
#[rstest]
#[case::missing_api_base(
false,
json!({}),
numbered_token,
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::MissingApiBase {
provider: "Azure AI",
environment_variable: "AZURE_AI_API_BASE",
})),
0
)]
#[case::unsupported_oidc_reference(
true,
json!({"azure_ad_token": "oidc/assertion", "client_id": "client", "tenant_id": "tenant"}),
numbered_token,
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::UnsupportedOidcReference)),
0
)]
#[case::empty_provider_token_ignores_static_token(
true,
json!({"azure_ad_token": "static-token"}),
|_| String::new(),
|error: &Error| matches!(error, Error::MissingAzureAiCredentials),
1
)]
#[tokio::test]
async fn credential_failures_send_no_provider_request(
#[case] with_api_base: bool,
#[case] optional_params: Value,
#[case] token: fn(usize) -> String,
#[case] expected: fn(&Error) -> bool,
#[case] expected_calls: usize,
) {
let provider = CountingToken::new(token);
let upstream = upstream([ocr_page()]).await;
let base = upstream.uri();
let error = perform(azure_request(
&provider,
with_api_base.then_some(base.as_str()),
None,
Value::Null,
optional_params,
))
.await
.unwrap_err();
assert!(expected(&error), "unexpected error: {error:?}");
assert_eq!(provider.calls(), expected_calls);
assert!(received(&upstream).await.is_empty());
}

View file

@ -0,0 +1,441 @@
use std::{
sync::{Arc, Mutex},
time::Duration,
};
use litellm_host::event::{CallEvent, MachineEvent};
use litellm_llms::base_llm::ocr::settings::OcrSettings;
use rstest::rstest;
use super::*;
const MODEL: &str = "azure_ai/doc-intelligence/prebuilt-read";
fn read_request(base: &str, options: Value) -> LiteLLMOcrRequest {
ocr_request(MODEL, base, options)
}
#[tokio::test]
async fn pages_features_and_extra_options_map_to_the_analyze_call() {
let upstream = upstream([json_response(json!({
"status": "succeeded",
"analyzeResult": {"pages": []}
}))])
.await;
let request = read_request(
&upstream.uri(),
json!({
"pages": [2, 0, 0, 1],
"features": ["keyValuePairs", "languages"],
"future_option": {"nested": null},
"extra_body": {"provider_option": false}
}),
)
.with_document(
document(
json!({"type": "document_url", "document_url": "https://example.com/document.pdf"}),
)
.into(),
);
perform(request).await.unwrap();
let sent = only_request(&upstream).await;
assert!(
sent.url.path().ends_with("/prebuilt-read:analyze"),
"{}",
sent.url
);
assert_eq!(sent.query("pages").as_deref(), Some("1,2,3"));
assert_eq!(
sent.query("features").as_deref(),
Some("keyValuePairs,languages")
);
assert_eq!(
sent.json(),
json!({
"urlSource": "https://example.com/document.pdf",
"future_option": {"nested": null},
"provider_option": false
})
);
}
#[rstest]
#[case(json!({"pages": [true]}), Error::Pages("expected only integers or only strings".into()))]
#[case(json!({"pages": [1, "2"]}), Error::Pages("expected only integers or only strings".into()))]
#[case(json!({"pages": [-1]}), Error::Pages("negative page index".into()))]
#[case(json!({"pages": "1&&features=bad"}), Error::Pages("invalid native page range".into()))]
#[case(json!({"features": "languages&pages=1"}), Error::Features)]
#[case(json!({"req_format": "azure"}), Error::RequestFormat)]
#[tokio::test]
async fn invalid_pages_features_and_format_are_rejected_before_sending(
#[case] options: Value,
#[case] expected: Error,
) {
let upstream = upstream([json_response(json!({}))]).await;
let result = match decode_request(wire(
MODEL,
&upstream.uri(),
json!({"type": "document_url", "document_url": "https://example.com/a.pdf"}),
options.clone(),
)) {
Ok(request) => perform(request).await,
Err(error) => Err(error),
};
assert!(
received(&upstream).await.is_empty(),
"sent invalid options: {options}"
);
let error = result.unwrap_err();
assert_eq!(
std::mem::discriminant(&error),
std::mem::discriminant(&expected)
);
assert_eq!(error.http_status_code(), Some(400));
assert_eq!(error.to_string(), expected.to_string());
}
#[rstest]
#[case::no_options(json!({}))]
#[case::litellm_format(json!({"req_format": "litellm"}))]
#[tokio::test]
async fn an_inline_document_is_sent_as_base64_and_only_page_text_is_kept(#[case] options: Value) {
let upstream = upstream([json_response(json!({
"status": "succeeded",
"analyzeResult": {"pages": [{"pageNumber": 1, "lines": [{"content": "hello"}]}]}
}))])
.await;
let response = perform(read_request(&upstream.uri(), options))
.await
.unwrap();
assert_eq!(response.pages.len(), 1);
assert_eq!(response.pages[0].index, 0);
assert_eq!(response.pages[0].markdown, "hello");
assert_eq!(response.provider_native_response, None);
let serialized = response.into_json();
for field in ["content", "tables", "keyValuePairs"] {
assert_eq!(serialized.get(field), Some(&Value::Null), "{field}");
}
let sent = only_request(&upstream).await;
for field in ["pages", "features", "req_format"] {
assert_eq!(sent.query(field), None, "{field}");
}
assert_eq!(sent.json(), json!({"base64Source": "YWJj"}));
}
#[tokio::test]
async fn native_format_normalizes_pages_and_keeps_the_provider_response() {
let operation = json!({
"status": "succeeded",
"operationExtension": 42,
"analyzeResult": {
"content": "A\n\nB",
"tables": [{"cells": []}],
"keyValuePairs": [{"key": {"content": "A"}}],
"pages": [{
"pageNumber": "2",
"width": "8.5",
"height": 11,
"unit": "inch",
"lines": [{"content": "A"}, {"content": null}, {"content": "B"}]
}]
}
});
let upstream = upstream([json_response(operation.clone())]).await;
let result = perform(read_request(
&upstream.uri(),
json!({"req_format": "native"}),
))
.await
.unwrap();
assert_eq!(result.pages[0].index, 1);
assert_eq!(result.pages[0].markdown, "A\n\nB");
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width": 816, "height": 1056, "dpi": 96})
);
assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1));
let serialized = result.clone().into_json();
assert_eq!(serialized["content"], "A\n\nB");
assert_eq!(serialized["tables"], json!([{"cells": []}]));
assert_eq!(
serialized["keyValuePairs"],
json!([{"key": {"content": "A"}}])
);
assert!(serialized.get("key_value_pairs").is_none());
assert_eq!(
result.provider_native_response.map(Value::Object),
Some(operation)
);
}
#[tokio::test]
async fn client_settings_choose_the_api_version_and_the_inch_to_pixel_dpi() {
let upstream = upstream([json_response(json!({
"status": "succeeded",
"analyzeResult": {"pages": [{"pageNumber": 1, "width": 8.5, "height": 11, "unit": "inch"}]}
}))])
.await;
let client = ocr_client().with_settings(OcrSettings {
document_intelligence_api_version: "2099-01-01".into(),
document_intelligence_dpi: 72,
..OcrSettings::default()
});
let result =
litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({})))
.await
.unwrap();
assert_eq!(
only_request(&upstream)
.await
.query("api-version")
.as_deref(),
Some("2099-01-01")
);
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width": 612, "height": 792, "dpi": 72})
);
}
#[tokio::test]
async fn an_accepted_response_polls_to_success_with_only_credentials() {
let operation = json!({"status": "succeeded", "analyzeResult": {"pages": []}});
let upstream = MockServer::start().await;
respond_in_order(
&upstream,
[
accepted(&upstream, json!({})),
json_response(json!({"status": "running"})).insert_header("Retry-After", "0"),
json_response(operation.clone()),
],
)
.await;
let request = with_headers(
read_request(&upstream.uri(), json!({"req_format": "native"})),
&[("X-Trace", "initial-only")],
);
let result = perform(request).await.unwrap();
assert_eq!(
result.provider_native_response.map(Value::Object),
Some(operation)
);
let requests = received(&upstream).await;
assert_eq!(requests.len(), 3);
assert_eq!(requests[0].header("x-trace"), Some("initial-only"));
for poll in &requests[1..] {
assert_eq!(poll.method.as_str(), "GET");
assert_eq!(poll.url.path(), "/operation");
assert_eq!(poll.header("x-trace"), None);
assert_eq!(poll.header("ocp-apim-subscription-key"), Some("test-key"));
}
}
#[tokio::test]
async fn polling_forwards_bearer_credentials() {
let upstream = MockServer::start().await;
respond_in_order(
&upstream,
[
accepted(&upstream, json!({})),
json_response(json!({"status": "succeeded"})),
],
)
.await;
let request = with_headers(
without_api_key(read_request(&upstream.uri(), json!({}))),
&[("Authorization", "Bearer token")],
);
perform(request).await.unwrap();
assert_eq!(
received(&upstream).await[1].header("authorization"),
Some("Bearer token")
);
}
#[tokio::test]
async fn response_received_fires_for_the_submission_and_the_completed_poll() {
let upstream = MockServer::start().await;
respond_in_order(
&upstream,
[
accepted(&upstream, json!({"submitted": true})),
json_response(json!({"status": "succeeded"})),
],
)
.await;
let observed = Arc::new(Mutex::new(Vec::new()));
let recorder = observed.clone();
let host =
LocalOcrHost::new(read_request(&upstream.uri(), json!({}))).with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
recorder.lock().unwrap().push(raw.body.clone());
}
});
perform_with(host).await.unwrap();
assert_eq!(received(&upstream).await.len(), 2);
assert_eq!(
*observed.lock().unwrap(),
[r#"{"submitted":true}"#, r#"{"status":"succeeded"}"#]
);
}
#[tokio::test]
async fn polling_does_not_follow_redirects() {
let upstream = MockServer::start().await;
respond_in_order(
&upstream,
[
accepted(&upstream, json!({})),
ResponseTemplate::new(302)
.insert_header("Location", format!("{}/redirected", upstream.uri())),
json_response(json!({"status": "succeeded"})),
],
)
.await;
let error = perform(read_request(&upstream.uri(), json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains("status 302"), "{error}");
assert_eq!(received(&upstream).await.len(), 2);
}
#[tokio::test]
async fn a_failed_operation_is_an_error() {
let upstream = MockServer::start().await;
respond_in_order(
&upstream,
[
accepted(&upstream, json!({})),
json_response(json!({"status": "failed"})),
],
)
.await;
let error = perform(read_request(&upstream.uri(), json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains("status failed"), "{error}");
}
#[tokio::test]
async fn the_polling_deadline_bounds_the_retry_delay() {
let upstream = MockServer::start().await;
respond_in_order(
&upstream,
[
accepted(&upstream, json!({})),
json_response(json!({"status": "notStarted"})).insert_header("Retry-After", "9999"),
],
)
.await;
let client = ocr_client().with_settings(OcrSettings {
poll_timeout: Duration::from_millis(100),
..OcrSettings::default()
});
let error = tokio::time::timeout(
Duration::from_secs(1),
litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({}))),
)
.await
.expect("the deadline cuts the retry delay short")
.unwrap_err();
assert!(error.to_string().contains("timed out"), "{error}");
}
#[rstest]
#[case::null_pages(json!({"pages": null}), "pages")]
#[case::null_page(json!({"pages": [null]}), "pages[0]")]
#[case::null_lines(json!({"pages": [{"lines": null}]}), "lines")]
#[case::bad_width(json!({"pages": [{"width": "bad"}]}), "width")]
#[tokio::test]
async fn malformed_provider_pages_report_the_response_path(
#[case] analysis: Value,
#[case] path: &str,
) {
let upstream = upstream([json_response(json!({
"status": "succeeded",
"analyzeResult": analysis
}))])
.await;
let error = perform(read_request(&upstream.uri(), json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains(path), "{error}");
}
#[rstest]
#[case::missing(None)]
#[case::relative(Some("/relative"))]
#[case::cross_origin(Some("http://example.com/operation"))]
#[case::with_userinfo(Some("http://user:password@127.0.0.1/operation"))]
#[tokio::test]
async fn an_unusable_operation_location_is_rejected(#[case] location: Option<&str>) {
let response = location
.into_iter()
.fold(ResponseTemplate::new(202), |response, location| {
response.insert_header("Operation-Location", location)
});
let upstream = upstream([response]).await;
let error = perform(read_request(&upstream.uri(), json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains("operation-location"), "{error}");
assert_eq!(received(&upstream).await.len(), 1);
}
#[tokio::test]
async fn the_model_id_is_percent_encoded() {
let upstream = upstream([json_response(json!({"status": "succeeded"}))]).await;
perform(ocr_request(
"azure_ai/doc-intelligence/a ?#é",
&upstream.uri(),
json!({}),
))
.await
.unwrap();
let sent = only_request(&upstream).await;
assert!(
sent.url.path().ends_with("/a%20%3F%23%C3%A9:analyze"),
"{}",
sent.url
);
}
#[rstest]
#[case::dot("azure_ai/doc-intelligence/.")]
#[case::dot_dot("azure_ai/doc-intelligence/..")]
#[tokio::test]
async fn dot_segment_model_ids_are_rejected(#[case] model: &str) {
let error = perform(ocr_request(model, UNREACHABLE_BASE, json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains("dot segment"), "{error}");
}

View file

@ -0,0 +1,42 @@
use rstest::rstest;
use super::*;
#[rstest]
#[case::cohere("cohere/parse-v5.0", "/v2/parse")]
#[case::azure_ai("azure_ai/Cohere-parse-v5.0", "/providers/cohere/v2/parse")]
#[tokio::test]
async fn an_image_goes_to_the_parse_endpoint_with_the_bearer_key(
#[case] model: &str,
#[case] path: &str,
) {
let upstream = upstream([pages_response()]).await;
let request = ocr_request_with_document(
model,
&upstream.uri(),
json!({"type": "image_url", "image_url": "data:image/png;base64,YWJj"}),
json!({}),
);
perform(request).await.unwrap();
let sent = only_request(&upstream).await;
assert_eq!(sent.method.as_str(), "POST");
assert_eq!(sent.url.path(), path);
assert_eq!(sent.header("authorization"), Some("Bearer test-key"));
}
#[rstest]
#[tokio::test]
async fn a_non_image_document_is_rejected_before_sending(
#[values("cohere/parse-v5.0", "azure_ai/Cohere-parse-v5.0")] model: &str,
) {
let upstream = upstream([pages_response()]).await;
let error = perform(ocr_request(model, &upstream.uri(), json!({})))
.await
.unwrap_err();
assert!(matches!(error, Error::CohereImageOnly), "{error:?}");
assert!(received(&upstream).await.is_empty());
}

View file

@ -0,0 +1,182 @@
use base64::Engine;
use litellm_core::ocr::types::OcrDocumentInput;
use litellm_host::event::WireRequest;
use rstest::rstest;
use wiremock::{Mock, matchers::any};
use super::*;
const SERVED_DOCUMENT: &[u8] = b"\x89PNG served document";
const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ=";
#[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 Guardrail {
Detached,
ReplacesDocument,
}
impl Guardrail {
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::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::Detached | Self::ReplacesDocument => (name, value),
})
.collect();
WireRequest {
body: Value::Object(body),
..wire
}
}
}
/// Serves [`SERVED_DOCUMENT`] as `image/png` to every request.
async fn document_server() -> MockServer {
let server = MockServer::start().await;
Mock::given(any())
.respond_with(ResponseTemplate::new(200).set_body_raw(SERVED_DOCUMENT, "image/png"))
.mount(&server)
.await;
server
}
/// Sends a remote document through `route` and returns the document the provider saw.
async fn provider_document(route: Route, guardrail: Guardrail) -> Value {
let documents = document_server().await;
let upstream = upstream([pages_response()]).await;
let document_type = route.document_type();
let request = ocr_request_with_document(
route.model(),
&upstream.uri(),
json!({"type": document_type, document_type: format!("{}/scan.png", documents.uri())}),
route.options(),
);
let host =
LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(guardrail.before_send(wire)));
perform_with(host).await.unwrap();
only_request(&upstream).await.json()["document"][document_type].clone()
}
#[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 expected = format!(
"data:image/png;base64,{}",
base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT)
);
assert_eq!(
provider_document(route, Guardrail::Detached).await,
expected
);
}
#[rstest]
#[tokio::test]
async fn a_document_replaced_by_the_host_reaches_the_provider(
#[values(
Route::Mistral,
Route::AzureAi,
Route::VertexMistral,
Route::AzureCohereParse,
Route::Cohere
)]
route: Route,
) {
assert_eq!(
provider_document(route, Guardrail::ReplacesDocument).await,
REPLACED_DOCUMENT
);
}
#[tokio::test]
async fn an_empty_byte_document_fails_before_sending() {
let upstream = upstream([pages_response()]).await;
let request = ocr_request("mistral/model", &upstream.uri(), json!({})).with_document(
OcrDocumentInput::Bytes {
bytes: Default::default(),
file_name: None,
mime_type: None,
},
);
let error = perform(request).await.unwrap_err();
assert!(matches!(error, Error::EmptyFile), "{error:?}");
assert!(received(&upstream).await.is_empty());
}
#[tokio::test]
async fn a_missing_path_document_fails_before_sending() {
let upstream = upstream([pages_response()]).await;
let path =
std::env::temp_dir().join(format!("litellm-ocr-missing-{}.png", rand::random::<u64>()));
let request = ocr_request("mistral/model", &upstream.uri(), json!({})).with_document(
OcrDocumentInput::Path {
path: path.clone(),
mime_type: None,
},
);
let error = perform(request).await.unwrap_err();
assert!(
matches!(
&error,
Error::FileRead { path: failed, source }
if *failed == path && source.kind() == std::io::ErrorKind::NotFound
),
"{error:?}"
);
assert!(received(&upstream).await.is_empty());
}

View file

@ -0,0 +1,269 @@
use std::sync::{Arc, Mutex};
use litellm_core::ocr::{
route::{Ocr, OcrOp, OcrProjection, ocr_machine},
types::OcrDocumentInput,
};
use litellm_host::{
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
host::Host,
};
use rstest::rstest;
use super::*;
pub(crate) fn event_name(event: &CallEvent) -> &'static str {
match event {
CallEvent::Started { .. } => "started",
CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response",
CallEvent::Succeeded { .. } => "success",
CallEvent::Failed { .. } => "failure",
}
}
fn recording_host(
request: LiteLLMOcrRequest,
events: Arc<Mutex<Vec<&'static str>>>,
block: bool,
) -> LocalOcrHost {
let before_send_events = events.clone();
LocalOcrHost::new(request)
.with_before_send(move |wire, _| {
before_send_events.lock().unwrap().push("before_send");
match block {
true => Err(Error::InvalidRequest("blocked".into())),
false => Ok(wire),
}
})
.with_observer(move |event| events.lock().unwrap().push(event_name(event)))
}
#[tokio::test]
async fn hooks_run_in_order_and_one_success_is_emitted() {
let upstream = upstream([pages_response()]).await;
let events = Arc::new(Mutex::new(Vec::new()));
perform_with(recording_host(
ocr_request("mistral/model", &upstream.uri(), json!({})),
events.clone(),
false,
))
.await
.unwrap();
assert_eq!(
*events.lock().unwrap(),
["started", "before_send", "response", "success"]
);
assert_eq!(received(&upstream).await.len(), 1);
}
#[tokio::test]
async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() {
let upstream = upstream([pages_response()]).await;
let events = Arc::new(Mutex::new(Vec::new()));
let error = perform_with(recording_host(
ocr_request("mistral/model", &upstream.uri(), json!({})),
events.clone(),
true,
))
.await
.unwrap_err();
assert!(
matches!(&error, Error::InvalidRequest(message) if message == "blocked"),
"{error:?}"
);
assert_eq!(
*events.lock().unwrap(),
["started", "before_send", "failure"]
);
assert!(received(&upstream).await.is_empty());
}
#[tokio::test]
async fn an_upstream_failure_emits_one_terminal_failure() {
let upstream = upstream([status_response(500, json!({"error": "failed"}))]).await;
let events = Arc::new(Mutex::new(Vec::new()));
let result = perform_with(recording_host(
ocr_request("mistral/model", &upstream.uri(), json!({})),
events.clone(),
false,
))
.await;
assert!(result.is_err());
assert_eq!(
*events.lock().unwrap(),
["started", "before_send", "failure"]
);
assert_eq!(received(&upstream).await.len(), 1);
}
#[tokio::test]
async fn an_invalid_provider_response_is_observed_before_normalization_fails() {
let upstream = upstream([json_response(json!({"pages": "invalid"}))]).await;
let observed = Arc::new(Mutex::new(Vec::new()));
let recorder = observed.clone();
let host = LocalOcrHost::new(ocr_request("mistral/model", &upstream.uri(), json!({})))
.with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
recorder.lock().unwrap().push(raw.body.clone());
}
});
let error = perform_with(host).await.unwrap_err();
assert!(matches!(error, Error::ResponseField { .. }), "{error:?}");
assert_eq!(*observed.lock().unwrap(), [r#"{"pages":"invalid"}"#]);
}
#[tokio::test]
async fn headers_returned_by_before_send_are_sent() {
let upstream = upstream([pages_response()]).await;
let host = LocalOcrHost::new(ocr_request("mistral/model", &upstream.uri(), json!({})))
.with_before_send(|mut wire, _| {
wire.headers
.push(("x-core-callback".into(), "edited".into()));
Ok(wire)
});
perform_with(host).await.unwrap();
assert_eq!(
only_request(&upstream).await.header("x-core-callback"),
Some("edited")
);
}
async fn before_send_context(request: LiteLLMOcrRequest) -> (WireRequest, RequestContext) {
let observed = Arc::new(Mutex::new(None));
let captured = observed.clone();
let host = LocalOcrHost::new(request).with_before_send(move |wire, context| {
*captured.lock().unwrap() = Some((wire.clone(), context.clone()));
Ok(wire)
});
perform_with(host).await.unwrap();
let context = observed.lock().unwrap().take();
context.expect("before_send ran")
}
#[tokio::test]
async fn before_send_sees_the_route_its_params_and_the_body() {
let upstream = upstream([pages_response()]).await;
let (wire, context) = before_send_context(ocr_request(
"mistral/model",
&upstream.uri(),
json!({"pages": [0], "req_format": "native"}),
))
.await;
assert_eq!(context.custom_llm_provider, "mistral");
assert_eq!(context.model, "model");
assert_eq!(context.optional_params["req_format"], "native");
assert!(context.secret_fields.is_empty());
assert_eq!(wire.body["pages"], json!([0]));
}
#[rstest]
#[case::client_secret(json!({"client_secret": "shh", "tenant_id": "t"}), &["client_secret"])]
#[case::no_secrets(json!({"tenant_id": "t"}), &[])]
#[tokio::test]
async fn before_send_names_the_secret_params(#[case] options: Value, #[case] secrets: &[&str]) {
let upstream = upstream([pages_response()]).await;
let request = ocr_request("azure_ai/model", &upstream.uri(), options).with_document(
OcrDocumentInput::Bytes {
bytes: b"abc".as_slice().into(),
file_name: None,
mime_type: Some("application/pdf".into()),
},
);
let (_, context) = before_send_context(request).await;
assert_eq!(context.secret_fields, secrets);
}
/// Hands the route a caller-owned Azure token and rewrites the bearer in `before_send`.
struct CallerTokenHost {
request: Mutex<Option<LiteLLMOcrRequest>>,
trace: Mutex<Vec<String>>,
}
impl Host<Ocr> for CallerTokenHost {
async fn project(&self) -> Result<OcrProjection, Error> {
self.trace.lock().unwrap().push("project".into());
Ok(OcrProjection {
request: self.request.lock().unwrap().take().unwrap(),
caller_token: true,
})
}
async fn custom_op(&self, op: OcrOp) -> Result<(), Error> {
match op {
OcrOp::AcquireAzureAdToken(reply) => {
self.trace.lock().unwrap().push("token".into());
reply.send(litellm_auth::ResolvedCredential::Static(
litellm_auth::SecretValue::new("caller-token"),
));
Ok(())
}
}
}
async fn before_send(
&self,
wire: WireRequest,
_: &RequestContext,
) -> Result<WireRequest, Error> {
let is_authorization = |name: &str| name.eq_ignore_ascii_case("authorization");
let authorization = wire
.headers
.iter()
.find(|(name, _)| is_authorization(name))
.map(|(_, value)| value.clone())
.unwrap_or_default();
self.trace
.lock()
.unwrap()
.push(format!("before_send:{authorization}"));
let headers = wire
.headers
.into_iter()
.map(|(name, value)| match is_authorization(&name) {
true => (name, "Bearer edited".to_string()),
false => (name, value),
})
.collect();
Ok(WireRequest { headers, ..wire })
}
}
#[tokio::test]
async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_replace_it() {
let upstream = upstream([pages_response()]).await;
let host = CallerTokenHost {
request: Mutex::new(Some(without_api_key(ocr_request(
"azure_ai/model",
&upstream.uri(),
json!({}),
)))),
trace: Mutex::new(Vec::new()),
};
litellm_host::run::run(ocr_machine(ocr_client()), &host)
.await
.unwrap();
assert_eq!(
*host.trace.lock().unwrap(),
["project", "token", "before_send:Bearer caller-token"]
);
assert_eq!(
only_request(&upstream).await.header_values("authorization"),
["Bearer edited"]
);
}

View file

@ -0,0 +1,284 @@
use std::{
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
time::Duration,
};
use litellm_core::ocr::{
route::{OcrMachine, OcrOp, OcrProjection},
types::OcrDocumentInput,
};
use litellm_host::{
event::{CallEvent, WireRequest},
host::{Host, HostOp},
machine::{HostFailure, Machine, MachineStep},
};
use litellm_llms::base_llm::ocr::transformation::OcrTransportConfig;
use rstest::rstest;
use tokio::{io::AsyncReadExt, net::TcpListener, sync::Notify};
use super::{lifecycle::event_name, *};
/// Drives the machine by hand, answering every op through `host` except `before_send`,
/// which `intercept` answers so a test can fail or cancel exactly there.
async fn drive_until(
host: &LocalOcrHost,
mut intercept: impl FnMut(WireRequest) -> Result<WireRequest, HostFailure<Error>>,
) -> (
Result<LiteLLMOcrResponse, Error>,
Vec<&'static str>,
OcrMachine,
) {
let mut machine = ocr_machine(ocr_client());
let mut ops = Vec::new();
let outcome = loop {
let op = match machine.resume().await {
Ok(MachineStep::Host(op)) => op,
Ok(MachineStep::Complete(response)) => break Ok(response),
Err(error) => break Err(error),
};
let answer = match op {
HostOp::Project(reply) => {
ops.push("Project");
host.project()
.await
.map(|projection| reply.send(projection))
.map_err(HostFailure::Error)
}
HostOp::Custom(op) => {
ops.push(match op {
OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken",
});
host.custom_op(op).await.map_err(HostFailure::Error)
}
HostOp::BeforeSend { wire, reply, .. } => {
ops.push("BeforeSend");
intercept(*wire).map(|wire| reply.send(wire))
}
HostOp::Emit(event, reply) => {
let event = CallEvent::Machine(event);
ops.push(event_name(&event));
host.emit(&event)
.await
.map(|()| reply.send(()))
.map_err(HostFailure::Error)
}
};
if let Err(failure) = answer {
break machine.interrupt(failure).await;
}
};
(outcome, ops, machine)
}
/// Answers every op until `stop` fires, leaving the machine suspended mid-call.
async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, stop: &Notify) {
tokio::time::timeout(Duration::from_secs(2), async {
loop {
tokio::select! {
_ = stop.notified() => break,
step = machine.resume() => {
match step.unwrap() {
MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()),
MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(),
MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire),
MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()),
MachineStep::Complete(_) => panic!("the stalled call completed"),
}
}
}
}
})
.await
.expect("the call reached the stall point");
}
#[tokio::test]
async fn a_hand_driven_machine_performs_the_same_call() {
let upstream = upstream([json_response(json!({
"pages": [{"index": 0, "markdown": "native"}]
}))])
.await;
let host = LocalOcrHost::new(ocr_request("mistral/model", &upstream.uri(), json!({})));
let (outcome, ops, mut machine) = drive_until(&host, Ok).await;
assert_eq!(outcome.unwrap().pages[0].markdown, "native");
assert_eq!(received(&upstream).await.len(), 1);
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
assert!(matches!(
machine.resume().await,
Err(Error::InvalidRequest(_))
));
}
#[tokio::test]
async fn a_path_document_is_read_by_core_without_a_host_operation() {
let upstream = upstream([json_response(json!({
"pages": [{"index": 0, "markdown": "path"}]
}))])
.await;
let dir = std::env::temp_dir().join(format!("litellm-ocr-{}", rand::random::<u64>()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("scan.png");
std::fs::write(&path, b"abc").unwrap();
let request = ocr_request("mistral/model", &upstream.uri(), json!({})).with_document(
OcrDocumentInput::Path {
path,
mime_type: None,
},
);
let (response, ops, _) = drive_until(&LocalOcrHost::new(request), Ok).await;
std::fs::remove_dir_all(&dir).unwrap();
assert_eq!(response.unwrap().pages[0].markdown, "path");
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
assert_eq!(
only_request(&upstream).await.json()["document"]["image_url"],
"data:image/png;base64,YWJj"
);
}
#[rstest]
#[case::failed(HostFailure::Error(Error::InvalidRequest("before_send failed".into())), "before_send failed")]
#[case::cancelled(HostFailure::Cancelled(Error::InvalidRequest("cancelled".into())), "cancelled")]
#[tokio::test]
async fn a_before_send_failure_ends_the_call_without_reaching_transport(
#[case] failure: HostFailure<Error>,
#[case] message: &str,
) {
let upstream = upstream([pages_response()]).await;
let host = LocalOcrHost::new(ocr_request("mistral/model", &upstream.uri(), json!({})));
let failure = Arc::new(std::sync::Mutex::new(Some(failure)));
let (outcome, ops, mut machine) = drive_until(&host, |_| {
Err(failure
.lock()
.unwrap()
.take()
.expect("before_send is asked once"))
})
.await;
assert!(
matches!(&outcome, Err(Error::InvalidRequest(actual)) if actual == message),
"{outcome:?}"
);
assert_eq!(ops, ["Project", "BeforeSend"]);
assert!(machine.resume().await.is_err());
assert!(received(&upstream).await.is_empty());
}
#[tokio::test]
async fn resuming_before_answering_keeps_the_pending_operation() {
let request = ocr_request("mistral/model", UNREACHABLE_BASE, json!({}));
let mut machine = ocr_machine(ocr_client());
let Ok(MachineStep::Host(HostOp::Project(reply))) = machine.resume().await else {
panic!("expected the projection op first");
};
assert!(machine.resume().await.is_err());
reply.send(OcrProjection {
request,
caller_token: false,
});
assert!(matches!(
machine.resume().await,
Ok(MachineStep::Host(HostOp::BeforeSend { .. }))
));
}
#[derive(Debug)]
struct PendingToken {
entered: Arc<Notify>,
dropped: Arc<AtomicBool>,
}
struct TokenFutureDrop(Arc<AtomicBool>);
impl Drop for TokenFutureDrop {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
impl litellm_auth::TokenProvider for PendingToken {
fn acquire(&self) -> litellm_auth::TokenFuture<'_> {
Box::pin(async move {
let _guard = TokenFutureDrop(self.dropped.clone());
self.entered.notify_one();
std::future::pending().await
})
}
}
#[tokio::test]
async fn interrupt_drops_provider_captures_before_returning() {
let entered = Arc::new(Notify::new());
let dropped = Arc::new(AtomicBool::new(false));
let mut request = ocr_request("azure_ai/mistral-ocr", "https://example.invalid", json!({}));
request.transport = OcrTransportConfig {
extra_headers: vec![("authorization".into(), "Bearer test-key".into())],
..request.transport
};
request.azure_ad_token_provider = Some(litellm_auth::TokenProviderHandle::new(Arc::new(
PendingToken {
entered: entered.clone(),
dropped: dropped.clone(),
},
)));
let host = LocalOcrHost::new(request);
let mut machine = ocr_machine(ocr_client());
drive_until_notified(&mut machine, &host, &entered).await;
assert!(!dropped.load(Ordering::SeqCst));
let acknowledgement = machine.interrupt(HostFailure::Cancelled(Error::InvalidRequest(
"cancelled".into(),
)));
assert!(
dropped.load(Ordering::SeqCst),
"interrupt returned while provider captures were still alive"
);
assert!(
matches!(acknowledgement.await, Err(Error::InvalidRequest(message)) if message == "cancelled")
);
}
#[tokio::test]
async fn interrupting_an_in_flight_provider_request_closes_its_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let received = Arc::new(Notify::new());
let server_received = received.clone();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = Vec::new();
let mut buffer = [0u8; 4096];
while !request.windows(4).any(|window| window == b"\r\n\r\n") {
let read = socket.read(&mut buffer).await.unwrap();
request.extend_from_slice(&buffer[..read]);
}
server_received.notify_one();
while socket.read(&mut buffer).await.unwrap() != 0 {}
});
let host = LocalOcrHost::new(ocr_request("mistral/model", &base, json!({})));
let mut machine = ocr_machine(ocr_client());
drive_until_notified(&mut machine, &host, &received).await;
let cancelled = Error::InvalidRequest("cancelled".into());
assert!(
machine
.interrupt(HostFailure::Cancelled(cancelled))
.await
.is_err()
);
tokio::time::timeout(Duration::from_secs(1), server)
.await
.expect("the provider connection stayed open after the interrupt")
.unwrap();
}

View file

@ -0,0 +1,125 @@
use litellm_core::ocr::{
document::prepare_document,
route::{LocalOcrHost, ocr_machine},
types::LiteLLMOcrRequest,
wire::{OcrWireRequest, decode_request},
};
use litellm_llms::base_llm::ocr::{
error::Error,
handler::OcrClient,
transformation::{LiteLLMOcrResponse, OcrDocument},
};
use serde_json::{Map, Value, json};
use wiremock::{MockServer, ResponseTemplate};
#[path = "../support/mod.rs"]
mod support;
use support::*;
mod aws_textract;
mod azure_ai;
mod azure_document_intelligence;
mod cohere;
mod documents;
mod lifecycle;
mod machine;
mod mistral;
mod reducto;
mod vertex_ai;
const INLINE_PDF: &str = "data:application/pdf;base64,YWJj";
fn object(value: Value) -> Map<String, Value> {
let Value::Object(map) = value else {
panic!("expected a json object, got {value}");
};
map
}
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)
}
async fn perform(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
litellm_core::ocr::client::perform(&ocr_client(), request).await
}
async fn perform_with(host: LocalOcrHost) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::run::run(ocr_machine(ocr_client()), &host).await
}
fn wire(model: &str, base: &str, document: Value, options: Value) -> OcrWireRequest {
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: object(options),
input_sources: Default::default(),
timeout_seconds: Some(2.0),
}
}
/// A request for an inline PDF, authenticated with `test-key`.
fn ocr_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
ocr_request_with_document(
model,
base,
json!({"type": "document_url", "document_url": INLINE_PDF}),
options,
)
}
fn ocr_request_with_document(
model: &str,
base: &str,
document: Value,
options: Value,
) -> LiteLLMOcrRequest {
decode_request(wire(model, base, document, options)).expect("request decodes")
}
fn document(value: Value) -> OcrDocument {
serde_json::from_value(value).expect("document parses")
}
/// Points the request's resolved document at `source`, keeping its type.
fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest {
let resolved = request
.map_document(prepare_document)
.expect("document resolves");
let document = resolved.document.clone().with_source(source.into());
resolved.with_document(document.into())
}
fn with_headers(request: LiteLLMOcrRequest, headers: &[(&str, &str)]) -> LiteLLMOcrRequest {
let mut request = request;
request.transport.extra_headers = headers
.iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect();
request
}
fn without_api_key(request: LiteLLMOcrRequest) -> LiteLLMOcrRequest {
let mut request = request;
request.credentials.api_key = None;
request
}
fn pages_response() -> ResponseTemplate {
json_response(json!({"pages": []}))
}
/// An Azure Document Intelligence 202 whose operation lives on `server`.
fn accepted(server: &MockServer, body: Value) -> ResponseTemplate {
ResponseTemplate::new(202)
.insert_header("Operation-Location", format!("{}/operation", server.uri()))
.set_body_json(body)
}

View file

@ -0,0 +1,248 @@
use std::sync::Arc;
use litellm_auth_gcp::VertexAuth;
use litellm_http::{
HttpClientPool, HttpSettings, Resolution,
media::{PublicDnsResolver, UrlPolicy},
};
use litellm_llms::{
base_llm::ocr::{
settings::OcrSettings,
transformation::{BaseOcrConfig, OCR_RESPONSE_MAX_BYTES},
},
mistral::ocr::transformation::MistralOcrConfig,
};
use rstest::rstest;
use super::*;
#[tokio::test]
async fn direct_mistral_sends_one_request_with_every_option() {
let upstream = upstream([json_response(json!({
"pages": [{"index": 0, "markdown": "hello", "custom": "preserved"}],
"usage_info": {"pages_processed": 1}
}))])
.await;
let result = perform(ocr_request(
"mistral/model",
&upstream.uri(),
json!({"pages": "0,2-4", "extract_header": true, "unknown": "ignored"}),
))
.await
.unwrap();
assert_eq!(result.pages[0].markdown, "hello");
assert_eq!(result.pages[0].extra_fields["custom"], "preserved");
let sent = only_request(&upstream).await;
assert_eq!(sent.url.path(), "/v1/ocr");
assert_eq!(sent.header("authorization"), Some("Bearer test-key"));
assert_eq!(
sent.json(),
json!({
"model": "model",
"document": {"type": "document_url", "document_url": INLINE_PDF},
"pages": "0,2-4",
"extract_header": true,
"unknown": "ignored"
})
);
}
#[rstest]
#[case::litellm_format(json!({}), false)]
#[case::native_format(json!({"req_format": "native"}), true)]
#[tokio::test]
async fn the_native_response_is_kept_only_when_requested(
#[case] options: Value,
#[case] kept: bool,
) {
let provider_response = json!({
"pages": [{"index": 0, "markdown": "hello"}],
"usage_info": {"pages_processed": 1},
"provider_only": "preserved"
});
let upstream = upstream([json_response(provider_response.clone())]).await;
let response = perform(ocr_request("mistral/model", &upstream.uri(), options))
.await
.unwrap();
assert_eq!(
response.provider_native_response.map(Value::Object),
kept.then_some(provider_response)
);
}
#[rstest]
#[case::mistral("mistral/model", json!({}))]
#[case::vertex(
"vertex_ai/mistral-ocr-latest",
json!({"vertex_project": "test-project", "vertex_location": "us-central1"})
)]
#[tokio::test]
async fn an_upstream_error_keeps_its_status_whole_body_and_headers(
#[case] model: &str,
#[case] options: Value,
) {
let payload = json!({"message": format!("{} END-OF-PROVIDER-BODY", "x".repeat(4096))});
let expected_body = serde_json::to_string(&payload).unwrap();
let upstream = upstream([status_response(422, payload)
.insert_header("Retry-After", "17")
.insert_header("X-Request-ID", "request-123")
.insert_header("X-Future-Header", "retained")])
.await;
let error = perform(ocr_request(model, &upstream.uri(), options))
.await
.unwrap_err();
assert_eq!(received(&upstream).await.len(), 1);
let Error::Provider {
status,
body,
headers,
} = error
else {
panic!("expected provider error, got {error:?}");
};
assert_eq!(status, 422);
for (name, value) in [
("retry-after", "17"),
("x-request-id", "request-123"),
("x-future-header", "retained"),
] {
assert!(
headers
.iter()
.any(|(key, actual)| key.eq_ignore_ascii_case(name) && actual == value),
"{name} missing from {headers:?}"
);
}
assert_eq!(body, expected_body);
}
#[rstest]
#[case::mistral_prefix("mistral/model", None, true)]
#[case::unknown_provider("model", Some("unknown"), false)]
fn decoding_accepts_known_providers_and_rejects_unknown_ones(
#[case] model: &str,
#[case] provider: Option<&str>,
#[case] accepted: bool,
) {
let request = OcrWireRequest {
custom_llm_provider: provider.map(Into::into),
..wire(
model,
"https://example.com",
json!({"type": "document_url", "document_url": "https://example.com/doc.pdf"}),
json!({"extract_header": true, "unknown": 42}),
)
};
assert_eq!(decode_request(request).is_ok(), accepted);
}
#[rstest]
#[case::plain_key(&[("MISTRAL_API_KEY", "plain")], "plain")]
#[case::azure_key_wins(&[("MISTRAL_AZURE_API_KEY", "azure"), ("MISTRAL_API_KEY", "plain")], "azure")]
#[case::empty_azure_key_falls_through(&[("MISTRAL_AZURE_API_KEY", ""), ("MISTRAL_API_KEY", "plain")], "plain")]
#[tokio::test]
async fn missing_credentials_come_from_the_injected_secret_source(
#[case] secrets: &[(&str, &str)],
#[case] expected_key: &str,
) {
let upstream = upstream([pages_response()]).await;
let base = upstream.uri();
let source = Arc::new(RecordingSecrets::new(
secrets
.iter()
.copied()
.chain([("MISTRAL_AZURE_API_BASE", base.as_str())]),
));
let client = ocr_client().with_secrets(source.clone());
let request = decode_request(OcrWireRequest {
api_key: None,
api_base: None,
..wire(
"mistral/model",
&base,
json!({"type": "document_url", "document_url": INLINE_PDF}),
json!({}),
)
})
.unwrap();
litellm_core::ocr::client::perform(&client, request)
.await
.unwrap();
assert_eq!(source.requested(), MistralOcrConfig.secret_names());
assert_eq!(
only_request(&upstream).await.header("authorization"),
Some(format!("Bearer {expected_key}").as_str())
);
}
#[tokio::test]
async fn the_client_uses_the_injected_http_pool_configuration() {
let upstream = upstream([pages_response()]).await;
let settings = HttpSettings {
user_agent: Some("host-owned/1".into()),
..HttpSettings::default()
};
let client = OcrClient::new(
&HttpClientPool::new(Arc::new(PublicDnsResolver)),
&Resolution::from(&settings).config,
UrlPolicy::default(),
VertexAuth::default(),
OcrSettings::default(),
Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
)
.unwrap();
litellm_core::ocr::client::perform(
&client,
ocr_request("mistral/model", &upstream.uri(), json!({})),
)
.await
.unwrap();
assert_eq!(
only_request(&upstream).await.header("user-agent"),
Some("host-owned/1")
);
}
#[test]
fn a_valid_response_limit_is_consumed_and_not_forwarded() {
let request = ocr_request(
"mistral/model",
UNREACHABLE_BASE,
json!({"max_response_bytes": 123}),
);
assert_eq!(request.transport.max_response_bytes, 123);
assert!(!request.optional_params.contains_key("max_response_bytes"));
}
#[rstest]
#[case::zero(json!(0))]
#[case::negative(json!(-1))]
#[case::boolean(json!(true))]
#[case::string(json!("123"))]
#[case::fraction(json!(1.5))]
#[case::above_the_cap(json!(OCR_RESPONSE_MAX_BYTES + 1))]
#[case::null(Value::Null)]
fn an_invalid_response_limit_is_rejected(#[case] limit: Value) {
let Err(error) = decode_request(wire(
"mistral/model",
UNREACHABLE_BASE,
json!({"type": "document_url", "document_url": INLINE_PDF}),
json!({"max_response_bytes": limit}),
)) else {
panic!("invalid response limit {limit} accepted");
};
assert!(error.to_string().contains("max_response_bytes"), "{error}");
}

View file

@ -0,0 +1,321 @@
use std::sync::{Arc, Mutex};
use litellm_host::event::{CallEvent, MachineEvent, WireRequest};
use rstest::rstest;
use super::*;
fn upload_response() -> ResponseTemplate {
json_response(json!({"file_id": "reducto://uploaded.pdf"}))
}
fn chunks_response(chunks: Value) -> ResponseTemplate {
json_response(json!({"result": {"chunks": chunks}}))
}
fn source_field(model: &str) -> &'static str {
match model.ends_with("parse-legacy") {
true => "document_url",
false => "input",
}
}
#[rstest]
#[case::v3(
"reducto/parse-v3",
json!({
"formatting": {"table_output_format": "html"},
"retrieval": {"chunk_mode": "section"},
"settings": {"ocr_system": "standard"},
"future_ocr_option": true,
"extra_body": {"provider_option": "value"}
}),
"reducto://already.pdf",
json!({
"input": "reducto://already.pdf",
"formatting": {"table_output_format": "html"},
"retrieval": {"chunk_mode": "section"},
"settings": {"ocr_system": "standard"},
"future_ocr_option": true,
"provider_option": "value"
})
)]
#[case::legacy(
"reducto/parse-legacy",
json!({
"enhance": {"agentic": [{"type": "table"}]},
"future_ocr_option": true,
"extra_body": {"provider_option": "value"}
}),
"reducto://legacy.pdf",
json!({
"document_url": "reducto://legacy.pdf",
"options": {"enhance": {"agentic": [{"type": "table"}]}},
"future_ocr_option": true,
"provider_option": "value"
})
)]
#[tokio::test]
async fn an_uploaded_document_is_parsed_with_mapped_options(
#[case] model: &str,
#[case] options: Value,
#[case] source: &str,
#[case] expected: Value,
) {
let upstream = upstream([chunks_response(json!([]))]).await;
perform(with_source(
ocr_request(model, &upstream.uri(), options),
source,
))
.await
.unwrap();
let sent = only_request(&upstream).await;
assert_eq!(sent.url.path(), "/parse");
assert_eq!(sent.json(), expected);
}
#[rstest]
#[tokio::test]
async fn an_inline_document_is_uploaded_as_multipart_then_parsed(
#[values("parse-v3", "parse-legacy")] model: &str,
#[values("application/pdf", "image/png")] mime_type: &str,
) {
let upstream = upstream([
upload_response(),
chunks_response(json!([{"content": "hello"}])),
])
.await;
let data_uri = format!("data:{mime_type};base64,YWJj");
let document = match mime_type.starts_with("image/") {
true => json!({"type": "image_url", "image_url": data_uri}),
false => json!({"type": "document_url", "document_url": data_uri}),
};
let request = with_headers(
ocr_request_with_document(
&format!("reducto/{model}"),
&upstream.uri(),
document,
json!({}),
),
&[
("Content-Type", "application/json"),
("X-Trace", "upload-test"),
],
);
let response = perform(request).await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = received(&upstream).await;
let [upload, parse] = requests.as_slice() else {
panic!(
"expected an upload and a parse, got {} requests",
requests.len()
);
};
assert_eq!(upload.url.path(), "/upload");
assert!(
upload
.header("content-type")
.is_some_and(|value| value.starts_with("multipart/form-data; boundary=")),
"{:?}",
upload.header("content-type")
);
assert_eq!(upload.header("x-trace"), Some("upload-test"));
let multipart = upload.body_text();
assert!(
multipart.contains(&format!("Content-Type: {mime_type}\r\n")),
"{multipart}"
);
assert!(multipart.contains("\r\n\r\nabc\r\n--"), "{multipart}");
assert_eq!(parse.url.path(), "/parse");
assert_eq!(
parse.json(),
json!({source_field(model): "reducto://uploaded.pdf"})
);
for request in &requests {
assert_eq!(request.header("authorization"), Some("Bearer test-key"));
}
}
#[tokio::test]
async fn response_received_fires_once_for_the_parse_response() {
let upstream = upstream([upload_response(), chunks_response(json!([]))]).await;
let observed = Arc::new(Mutex::new(Vec::new()));
let recorder = observed.clone();
let host = LocalOcrHost::new(ocr_request("reducto/parse-v3", &upstream.uri(), json!({})))
.with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
recorder.lock().unwrap().push(raw.body.clone());
}
});
perform_with(host).await.unwrap();
assert_eq!(received(&upstream).await.len(), 2);
assert_eq!(*observed.lock().unwrap(), [r#"{"result":{"chunks":[]}}"#]);
}
#[rstest]
#[case::empty_id(json_response(json!({"file_id": ""})))]
#[case::missing_id(json_response(json!({})))]
#[case::null_id(json_response(json!({"file_id": null})))]
#[case::upload_failure(status_response(503, json!({"error": "unavailable"})))]
#[tokio::test]
async fn a_failed_upload_stops_before_parse(#[case] upload: ResponseTemplate) {
let upstream = upstream([upload]).await;
let result = perform(ocr_request("reducto/parse-v3", &upstream.uri(), json!({}))).await;
assert!(result.is_err());
assert_eq!(received(&upstream).await.len(), 1);
}
#[rstest]
#[case::remote_url("https://example.com/a.pdf", Error::ReductoSource)]
#[case::empty_file_id("reducto://", Error::RequestField { path: "document file id".into() })]
#[case::data_uri_without_payload("data:application/pdf;base64", Error::InvalidDataUri)]
#[case::invalid_base64("data:application/pdf;base64,INVALID!", Error::InvalidDataUri)]
#[tokio::test]
async fn invalid_document_sources_are_rejected_before_sending(
#[case] source: &str,
#[case] expected: Error,
) {
let upstream = upstream([json_response(json!({}))]).await;
let result = perform(with_source(
ocr_request("reducto/parse-v3", &upstream.uri(), json!({})),
source,
))
.await;
assert!(
received(&upstream).await.is_empty(),
"sent invalid source: {source}"
);
let error = result.unwrap_err();
assert_eq!(
std::mem::discriminant(&error),
std::mem::discriminant(&expected)
);
assert_eq!(error.http_status_code(), Some(400));
assert_eq!(error.to_string(), expected.to_string());
}
#[tokio::test]
async fn a_forwarded_authorization_wins_and_the_native_response_is_omitted_by_default() {
let upstream = upstream([json_response(
json!({"job_id": "job-1", "result": {"chunks": []}}),
)])
.await;
let request = with_headers(
with_source(
ocr_request("reducto/parse-v3", &upstream.uri(), json!({})),
"reducto://ready.pdf",
),
&[("authorization", "Bearer existing")],
);
let response = perform(request).await.unwrap();
assert_eq!(response.provider_native_response, None);
assert_eq!(
only_request(&upstream).await.header_values("authorization"),
["Bearer existing"]
);
}
#[tokio::test]
async fn native_format_retains_the_provider_response() {
let raw = json!({
"result": {"chunks": [{"content": "native OCR response"}]},
"usage": {"num_pages": 1}
});
let upstream = upstream([json_response(raw.clone())]).await;
let response = perform(with_source(
ocr_request(
"reducto/parse-v3",
&upstream.uri(),
json!({"req_format": "native"}),
),
"reducto://ready.pdf",
))
.await
.unwrap();
assert_eq!(response.pages[0].markdown, "native OCR response");
assert_eq!(
response.provider_native_response.map(Value::Object),
Some(raw)
);
}
#[tokio::test]
async fn an_unknown_model_reaches_parse_and_keeps_its_name() {
let upstream = upstream([chunks_response(
json!([{"content": "future model response"}]),
)])
.await;
let response = perform(with_source(
ocr_request("reducto/future-parse-model", &upstream.uri(), json!({})),
"reducto://ready.pdf",
))
.await
.unwrap();
assert_eq!(response.model, "future-parse-model");
assert_eq!(response.pages[0].markdown, "future model response");
let sent = only_request(&upstream).await;
assert_eq!(sent.url.path(), "/parse");
assert_eq!(sent.json(), json!({"input": "reducto://ready.pdf"}));
}
#[tokio::test]
async fn a_guardrail_can_replace_the_document_before_upload() {
let upstream = upstream([chunks_response(json!([]))]).await;
let host = LocalOcrHost::new(ocr_request("reducto/parse-v3", &upstream.uri(), json!({})))
.with_before_send(|wire, _| {
assert_eq!(wire.body["document_url"], INLINE_PDF);
Ok(WireRequest {
body: json!({"type": "document_url", "document_url": "reducto://guarded.pdf"}),
..wire
})
});
perform_with(host).await.unwrap();
let sent = only_request(&upstream).await;
assert_eq!(sent.url.path(), "/parse");
assert_eq!(sent.json(), json!({"input": "reducto://guarded.pdf"}));
}
#[rstest]
#[tokio::test]
async fn guardrail_headers_reach_both_upload_and_parse(
#[values("reducto/parse-v3", "reducto/parse-legacy")] model: &str,
) {
let upstream = upstream([upload_response(), chunks_response(json!([]))]).await;
let request = with_headers(
ocr_request(model, &upstream.uri(), json!({})),
&[("authorization", "Bearer original")],
);
let host = LocalOcrHost::new(request).with_before_send(|wire, _| {
Ok(WireRequest {
headers: vec![("authorization".into(), "Bearer guarded".into())],
..wire
})
});
perform_with(host).await.unwrap();
let requests = received(&upstream).await;
let paths: Vec<&str> = requests.iter().map(|request| request.url.path()).collect();
assert_eq!(paths, ["/upload", "/parse"]);
for request in &requests {
assert_eq!(request.header_values("authorization"), ["Bearer guarded"]);
}
}

View file

@ -0,0 +1,184 @@
use litellm_auth::{InputSource, Sourced};
use litellm_core::ocr::arguments::is_supported_request;
use litellm_llms::base_llm::ocr::settings::OcrSettings;
use rstest::rstest;
use super::*;
#[tokio::test]
async fn mistral_is_served_at_the_resolved_project_and_location() {
let upstream = upstream([json_response(json!({
"pages": [{"index": 0, "markdown": "hello"}],
"usage_info": {"pages_processed": 1}
}))])
.await;
let response = perform(ocr_request(
"vertex_ai/mistral-ocr-maas",
&upstream.uri(),
json!({
"vertex_project": "project-1",
"vertex_location": "europe-west4",
"extract_footer": true
}),
))
.await
.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let sent = only_request(&upstream).await;
assert_eq!(
sent.url.path(),
"/v1/projects/project-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
assert_eq!(sent.header("authorization"), Some("Bearer test-key"));
assert_eq!(
sent.json(),
json!({
"model": "mistral-ocr-maas",
"document": {"type": "document_url", "document_url": INLINE_PDF},
"extract_footer": true
})
);
}
#[tokio::test]
async fn configured_project_and_location_apply_when_the_call_sets_neither() {
let upstream = upstream([pages_response()]).await;
let client = ocr_client().with_settings(OcrSettings {
vertex_project: Some("configured-project".into()),
vertex_location: Some("europe-west4".into()),
..OcrSettings::default()
});
litellm_core::ocr::client::perform(
&client,
ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})),
)
.await
.unwrap();
assert_eq!(
only_request(&upstream).await.url.path(),
"/v1/projects/configured-project/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
}
#[tokio::test]
async fn a_supplied_authorization_is_forwarded_without_a_static_token() {
let upstream = upstream([pages_response()]).await;
let request = with_headers(
without_api_key(ocr_request(
"vertex_ai/model",
&upstream.uri(),
json!({"vertex_project": "project-1"}),
)),
&[("authorization", "Bearer supplied")],
);
perform(request).await.unwrap();
assert_eq!(
only_request(&upstream).await.header_values("authorization"),
["Bearer supplied"]
);
}
#[tokio::test]
async fn invalid_credentials_fail_before_sending() {
let error = perform(ocr_request(
"vertex_ai/model",
UNREACHABLE_BASE,
json!({"vertex_credentials": true}),
))
.await
.unwrap_err();
assert!(error.to_string().contains("vertex_credentials"), "{error}");
}
#[rstest]
#[tokio::test]
async fn a_request_controlled_api_base_is_rejected_before_vertex_auth(
#[values("vertex_ai/mistral-ocr-maas", "vertex_ai/deepseek-ocr-maas")] model: &str,
) {
let mut request = ocr_request(
model,
"https://caller.example",
json!({"vertex_project": "project-1"}),
);
request.credentials.api_base = Some(Sourced::new(
"https://caller.example".into(),
InputSource::Request,
));
let error = perform(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint"),
"{error}"
);
}
#[tokio::test]
async fn deepseek_is_served_at_the_openai_compatible_endpoint() {
let upstream = upstream([json_response(json!({
"choices": [{"message": {"content": "recognized"}}],
"usage": {"prompt_tokens": 1}
}))])
.await;
let request = with_source(
ocr_request(
"vertex_ai/deepseek-ocr-maas",
&upstream.uri(),
json!({
"vertex_project": "project-1",
"vertex_location": "europe-west4",
"temperature": 0.1,
"future_ocr_option": true,
"extra_body": {"provider_option": "value"}
}),
),
"gs://bucket/document.pdf",
);
let response = perform(request).await.unwrap();
assert_eq!(response.pages[0].markdown, "recognized");
assert_eq!(
response.usage_info.unwrap().extra_fields["prompt_tokens"],
1
);
let sent = only_request(&upstream).await;
assert_eq!(
sent.url.path(),
"/v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions"
);
assert_eq!(sent.header("authorization"), Some("Bearer test-key"));
let body = sent.json();
assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(body["temperature"], 0.1);
assert_eq!(body["future_ocr_option"], true);
assert_eq!(body["provider_option"], "value");
assert!(body.get("vertex_project").is_none());
assert!(body.get("extra_body").is_none());
assert_eq!(
body["messages"][0]["content"][0],
json!({"type": "image_url", "image_url": "gs://bucket/document.pdf"})
);
}
#[rstest]
#[case::deepseek("deepseek-ocr-maas", Some("vertex_ai"), true)]
#[case::mistral("mistral-ocr-maas", Some("vertex_ai"), true)]
#[case::prefixed("vertex_ai/mistral-ocr-maas", None, true)]
#[case::unknown_provider("model", Some("unknown"), false)]
fn supported_requests_follow_the_registered_configs(
#[case] model: &str,
#[case] provider: Option<&str>,
#[case] supported: bool,
) {
assert_eq!(is_supported_request(model, provider), supported);
}

View file

@ -0,0 +1,155 @@
//! Shared fixtures for route integration tests: a scripted upstream and a recording
//! secret source.
#![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset
use std::sync::Mutex;
use futures_util::future::BoxFuture;
use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::Value;
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
/// A port nothing listens on, for calls that must fail before any request is sent.
pub const UNREACHABLE_BASE: &str = "http://127.0.0.1:1";
/// Starts an upstream that answers its n-th request with the n-th response and 404s after.
pub async fn upstream(responses: impl IntoIterator<Item = ResponseTemplate>) -> MockServer {
let server = MockServer::start().await;
respond_in_order(&server, responses).await;
server
}
/// Scripts responses on a started server, for responses that need its address.
pub async fn respond_in_order(
server: &MockServer,
responses: impl IntoIterator<Item = ResponseTemplate>,
) {
for response in responses {
Mock::given(any())
.respond_with(response)
.up_to_n_times(1)
.mount(server)
.await;
}
}
pub async fn received(server: &MockServer) -> Vec<Request> {
server
.received_requests()
.await
.expect("request recording is on")
}
pub async fn only_request(server: &MockServer) -> Request {
let [request] = <[Request; 1]>::try_from(received(server).await)
.unwrap_or_else(|requests| panic!("expected one request, got {}", requests.len()));
request
}
pub fn json_response(body: Value) -> ResponseTemplate {
ResponseTemplate::new(200).set_body_json(body)
}
pub fn status_response(status: u16, body: Value) -> ResponseTemplate {
ResponseTemplate::new(status).set_body_json(body)
}
pub trait ReceivedRequest {
fn header(&self, name: &str) -> Option<&str>;
fn header_values(&self, name: &str) -> Vec<&str>;
fn json(&self) -> Value;
fn body_text(&self) -> String;
/// The path and query, as the request line carried them.
fn target(&self) -> String;
fn query(&self, name: &str) -> Option<String>;
}
impl ReceivedRequest for Request {
fn header(&self, name: &str) -> Option<&str> {
self.headers.get(name).and_then(|value| value.to_str().ok())
}
fn header_values(&self, name: &str) -> Vec<&str> {
self.headers
.get_all(name)
.iter()
.filter_map(|value| value.to_str().ok())
.collect()
}
fn json(&self) -> Value {
serde_json::from_slice(&self.body).expect("request body is json")
}
fn body_text(&self) -> String {
String::from_utf8_lossy(&self.body).into_owned()
}
fn target(&self) -> String {
match self.url.query() {
Some(query) => format!("{}?{query}", self.url.path()),
None => self.url.path().to_string(),
}
}
fn query(&self, name: &str) -> Option<String> {
self.url
.query_pairs()
.find_map(|(key, value)| (key == name).then(|| value.into_owned()))
}
}
/// A secret source that answers from a fixed table and records every name it was asked for.
pub struct RecordingSecrets {
values: Vec<(String, String)>,
fails: bool,
requested: Mutex<Vec<String>>,
}
impl RecordingSecrets {
pub fn new<'a>(values: impl IntoIterator<Item = (&'a str, &'a str)>) -> Self {
Self {
values: values
.into_iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect(),
fails: false,
requested: Mutex::new(Vec::new()),
}
}
pub fn empty() -> Self {
Self::new([])
}
pub fn failing() -> Self {
Self {
fails: true,
..Self::empty()
}
}
pub fn requested(&self) -> Vec<String> {
self.requested.lock().unwrap().clone()
}
}
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())))
})
}
}

View file

@ -0,0 +1,31 @@
# Requirements
Core must pause mid-call to ask the host for things it cannot do itself (Python callbacks, secret and token reads, `before_send` rewrites, stream demand), then continue where it stopped. Any change to this crate must keep every requirement below; the alternatives section says which one each rejected design breaks
- R1 Core never calls the host: it names an op and waits for the answer, so it stays free of PyO3 and of any other host runtime
- R2 Async host work is awaited by the host's own driver in the caller's asyncio task (`litellm/rust_bridge/lifecycle.py`), so `contextvars` writes reach the caller; a Rust-side `into_future` would run it in a copied context
- R3 The body awaits real I/O (HTTP, `spawn_blocking`, timers) between yields, so `resume` is itself a future driven by the caller's runtime
- R4 Route code stays straight-line async (`host.route(OcrOp::ReadDocument).await?`) instead of hand-written states
- R5 Each op fixes its answer type at compile time: a host cannot answer `ReadDocument` with a token, and core never matches a result variant it did not ask for
- R6 A yield the body makes while being resumed is returned by that same poll, so the host driver's inline first poll needs no extra event-loop turn per op
- R7 No task is spawned: `cancel`, or dropping the coroutine, drops the body, and nothing waits forever on an answer that cannot come
- R8 Several yields can be pending at once, since route code hands clones of its `Co` to token providers and hooks
- R9 Stable Rust
# Other implementations and why they do not fit
- Nightly `std::ops::Coroutine`: breaks R9, and its body cannot await futures between yields (R3)
- `genawaiter`: resumes async bodies only with a noop waker, so the body cannot await real I/O (R3)
- `simple_coro`: typestate `Coro` makes answering before resuming a compile-time rule, but its body cannot await arbitrary futures (R3) and its reply type `R` is fixed per coroutine (R5)
- `corosensei` and other stackful coroutines: sync bodies on their own stack, no async I/O inside (R3)
- A hand-written phase enum with an `advance` match (the old `HostPhase`): every await point becomes a state (R4)
- An injected host trait with `async fn`s: core would call the host itself (R1, R2)
- Sans-IO, where core does no I/O and HTTP becomes one more host op: keeps every requirement and makes `resume` a pure step function, but HTTP, streaming, retries and timeouts would move out of core into every bridge; the one real alternative, not taken
- Temporal's Rust workflow SDK (`WorkflowFuture`, `WfContext`) is the closest precedent: an `async fn` polled in place, commands sent over a channel with a oneshot to unblock them. Roles are inverted there (the language SDK owns the program, core answers), and its workflow body may not do real I/O
# Tradeoffs accepted
- A tokio `mpsc` channel plus a `oneshot` per yield instead of compiler-generated states
- Protocol mistakes (resuming before answering, resuming after the end) are runtime `ResumeError`s, not compile errors
- Pending yields come out one per `resume`, in the order they were made, and each reply goes back to the yield that made it (R8)
- An answer sent after its yield stopped waiting (for example the body timed out on it) is discarded, since the body already moved on

View file

@ -0,0 +1,15 @@
[package]
name = "litellm-coroutine"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
description = "Async coroutines on stable Rust whose every yield carries its own typed reply"
[dependencies]
thiserror.workspace = true
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
rstest.workspace = true
tokio = { workspace = true, features = ["rt", "macros", "time"] }

View file

@ -0,0 +1,42 @@
use std::sync::Weak;
use tokio::sync::mpsc;
use crate::{Abandoned, Reply, reply};
pub(crate) struct Request<Y> {
pub(crate) value: Y,
pub(crate) outstanding: Weak<()>,
}
/// The body's handle for yielding, `genawaiter`'s `Co`.
pub struct Co<Y> {
yields: mpsc::UnboundedSender<Request<Y>>,
}
impl<Y> Clone for Co<Y> {
fn clone(&self) -> Self {
Self {
yields: self.yields.clone(),
}
}
}
impl<Y> Co<Y> {
pub(crate) fn new(yields: mpsc::UnboundedSender<Request<Y>>) -> Self {
Self { yields }
}
/// Yields the value `ask` builds around a fresh [`Reply`] and waits for its answer.
pub async fn yield_<A>(&self, ask: impl FnOnce(Reply<A>) -> Y) -> Result<A, Abandoned> {
let (reply, answer) = reply();
let outstanding = reply.outstanding();
self.yields
.send(Request {
value: ask(reply),
outstanding,
})
.map_err(|_| Abandoned)?;
answer.await
}
}

View file

@ -0,0 +1,94 @@
use std::{
future::{Future, poll_fn},
pin::Pin,
sync::Weak,
task::{Context, Poll},
};
use tokio::sync::mpsc;
use crate::{Co, ResumeError, co::Request};
/// What one `resume` produced, as in [`std::ops::CoroutineState`].
#[derive(Debug, PartialEq, Eq)]
pub enum CoroutineState<Y, C> {
Yielded(Y),
Complete(C),
}
type Body<C> = Pin<Box<dyn Future<Output = C> + Send>>;
enum Step<Y, C> {
Yielded(Request<Y>),
Complete(C),
}
fn queued<Y>(
yields: &mut mpsc::UnboundedReceiver<Request<Y>>,
context: &mut Context<'_>,
) -> Option<Request<Y>> {
match yields.poll_recv(context) {
Poll::Ready(request) => request,
Poll::Pending => None,
}
}
pub struct Coroutine<Y, C> {
body: Option<Body<C>>,
yields: mpsc::UnboundedReceiver<Request<Y>>,
outstanding: Weak<()>,
}
impl<Y, C> Coroutine<Y, C> {
/// Builds the body from `producer`. Nothing runs until the first `resume`.
pub fn new<F>(producer: impl FnOnce(Co<Y>) -> F) -> Self
where
F: Future<Output = C> + Send + 'static,
{
let (sender, yields) = mpsc::unbounded_channel();
Self {
body: Some(Box::pin(producer(Co::new(sender)))),
yields,
outstanding: Weak::new(),
}
}
pub async fn resume(&mut self) -> Result<CoroutineState<Y, C>, ResumeError> {
let Some(body) = self.body.as_mut() else {
return Err(ResumeError::Finished);
};
if self.outstanding.strong_count() > 0 {
return Err(ResumeError::Unanswered);
}
let yields = &mut self.yields;
let step = poll_fn(|context| {
if let Some(request) = queued(yields, context) {
return Poll::Ready(Step::Yielded(request));
}
if let Poll::Ready(output) = body.as_mut().poll(context) {
return Poll::Ready(Step::Complete(output));
}
queued(yields, context)
.map_or(Poll::Pending, |request| Poll::Ready(Step::Yielded(request)))
})
.await;
match step {
Step::Yielded(Request { value, outstanding }) => {
self.outstanding = outstanding;
Ok(CoroutineState::Yielded(value))
}
Step::Complete(output) => {
self.cancel();
Ok(CoroutineState::Complete(output))
}
}
}
/// Drops the body and fails every yield still waiting, or yet to be made, with
/// [`Abandoned`](crate::Abandoned).
pub fn cancel(&mut self) {
self.body = None;
self.yields.close();
while self.yields.try_recv().is_ok() {}
}
}

View file

@ -0,0 +1,14 @@
/// A `resume` the coroutine refused, leaving it as it was.
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
pub enum ResumeError {
#[error("coroutine resumed after it finished")]
Finished,
#[error("coroutine resumed before the reply to its last yield was sent or dropped")]
Unanswered,
}
/// No answer will come to a yield: its [`Reply`](crate::Reply) was dropped unsent, or the
/// coroutine it was sent to is gone.
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
#[error("the yield was abandoned before it was answered")]
pub struct Abandoned;

View file

@ -0,0 +1,12 @@
//! Async coroutines on stable Rust whose every yield carries its own typed [`Reply`].
//! See `AGENTS.md` for the requirement, the alternatives and the contracts.
mod co;
mod coroutine;
mod error;
mod reply;
pub use co::Co;
pub use coroutine::{Coroutine, CoroutineState};
pub use error::{Abandoned, ResumeError};
pub use reply::{Answer, Reply, reply};

View file

@ -0,0 +1,60 @@
use std::{
fmt,
future::Future,
pin::Pin,
sync::{Arc, Weak},
task::{Context, Poll},
};
use tokio::sync::oneshot;
use crate::Abandoned;
/// The one way to answer a yield. Sending or dropping it settles the yield.
pub struct Reply<A> {
slot: oneshot::Sender<A>,
outstanding: Arc<()>,
}
impl<A> Reply<A> {
/// An answer the yield no longer awaits is discarded.
pub fn send(self, answer: A) {
let _ = self.slot.send(answer);
}
/// Alive until this reply is sent or dropped.
pub(crate) fn outstanding(&self) -> Weak<()> {
Arc::downgrade(&self.outstanding)
}
}
impl<A> fmt::Debug for Reply<A> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("Reply")
}
}
/// The waiting end of a [`Reply`].
pub struct Answer<A> {
slot: oneshot::Receiver<A>,
}
impl<A> Future for Answer<A> {
type Output = Result<A, Abandoned>;
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
Pin::new(&mut self.slot)
.poll(context)
.map(|answer| answer.map_err(|_| Abandoned))
}
}
/// A reply outside any coroutine, for answering a host operation directly.
pub fn reply<A>() -> (Reply<A>, Answer<A>) {
let (slot, answer) = oneshot::channel();
let reply = Reply {
slot,
outstanding: Arc::new(()),
};
(reply, Answer { slot: answer })
}

View file

@ -0,0 +1,256 @@
use std::{
future::Future,
sync::{Arc, Mutex},
time::Duration,
};
use litellm_coroutine::{Abandoned, Co, Coroutine, CoroutineState, Reply, ResumeError, reply};
use rstest::rstest;
use tokio::time::timeout;
#[derive(Debug)]
enum Ask {
Name(Reply<&'static str>),
Count(Reply<u32>),
}
type Test<C> = Coroutine<Ask, C>;
fn yielded<C>(state: Result<CoroutineState<Ask, C>, ResumeError>) -> Ask {
match state {
Ok(CoroutineState::Yielded(ask)) => ask,
Ok(CoroutineState::Complete(_)) => panic!("expected a yield, the body returned"),
Err(error) => panic!("expected a yield, resume failed: {error}"),
}
}
fn complete<C>(state: Result<CoroutineState<Ask, C>, ResumeError>) -> C {
match state {
Ok(CoroutineState::Complete(output)) => output,
Ok(CoroutineState::Yielded(ask)) => panic!("expected completion, got {ask:?}"),
Err(error) => panic!("expected completion, resume failed: {error}"),
}
}
fn name(ask: Ask) -> Reply<&'static str> {
match ask {
Ask::Name(reply) => reply,
other => panic!("expected a name ask, got {other:?}"),
}
}
fn count(ask: Ask) -> Reply<u32> {
match ask {
Ask::Count(reply) => reply,
other => panic!("expected a count ask, got {other:?}"),
}
}
/// A body parked at one name ask, with nothing else going on.
fn suspended_once() -> Test<Result<&'static str, Abandoned>> {
Coroutine::new(|co| async move { co.yield_(Ask::Name).await })
}
#[tokio::test]
async fn each_typed_answer_resumes_the_yield_that_asked_for_it() {
let mut coroutine: Test<String> = Coroutine::new(|co| async move {
let first = co.yield_(Ask::Name).await.unwrap();
let second = co.yield_(Ask::Count).await.unwrap();
format!("{first}+{second}")
});
name(yielded(coroutine.resume().await)).send("a");
count(yielded(coroutine.resume().await)).send(2);
assert_eq!(complete(coroutine.resume().await), "a+2");
}
/// A driver that polls `resume` once, inline, sees every yield the body makes during
/// that poll instead of being sent back to its event loop.
#[test]
fn a_yield_made_while_resuming_is_returned_by_that_same_poll() {
let mut coroutine: Test<u32> = Coroutine::new(|co| async move {
let first = co.yield_(Ask::Count).await.unwrap();
let second = co.yield_(Ask::Count).await.unwrap();
first + second
});
let mut context = std::task::Context::from_waker(std::task::Waker::noop());
let mut poll_once =
|coroutine: &mut Test<u32>| match std::pin::pin!(coroutine.resume()).poll(&mut context) {
std::task::Poll::Ready(state) => state,
std::task::Poll::Pending => panic!("resume needed a second poll"),
};
count(yielded(poll_once(&mut coroutine))).send(1);
count(yielded(poll_once(&mut coroutine))).send(2);
assert_eq!(complete(poll_once(&mut coroutine)), 3);
}
#[tokio::test]
async fn the_body_awaits_real_futures_between_yields() {
let mut coroutine: Test<u32> = Coroutine::new(|co| async move {
tokio::time::sleep(Duration::from_millis(5)).await;
co.yield_(Ask::Count).await.unwrap()
});
count(yielded(coroutine.resume().await)).send(7);
assert_eq!(complete(coroutine.resume().await), 7);
}
#[tokio::test]
async fn concurrent_yields_come_out_in_order_and_are_answered_separately() {
let mut coroutine: Test<(&str, u32)> = Coroutine::new(|co| async move {
let (first, second) = tokio::join!(co.yield_(Ask::Name), co.yield_(Ask::Count));
(first.unwrap(), second.unwrap())
});
name(yielded(coroutine.resume().await)).send("one");
count(yielded(coroutine.resume().await)).send(2);
assert_eq!(complete(coroutine.resume().await), ("one", 2));
}
#[tokio::test]
async fn resuming_before_the_reply_is_settled_is_refused_and_keeps_the_yield_waiting() {
let mut coroutine = suspended_once();
let reply = name(yielded(coroutine.resume().await));
assert_eq!(
coroutine.resume().await.unwrap_err(),
ResumeError::Unanswered
);
reply.send("real");
assert_eq!(complete(coroutine.resume().await), Ok("real"));
}
#[tokio::test]
async fn a_dropped_reply_abandons_its_yield() {
let mut coroutine = suspended_once();
drop(yielded(coroutine.resume().await));
assert_eq!(complete(coroutine.resume().await), Err(Abandoned));
}
#[tokio::test]
async fn an_answer_the_yield_no_longer_awaits_is_discarded() {
let mut coroutine: Test<&str> = Coroutine::new(|co| async move {
tokio::select! {
biased;
_ = co.yield_(Ask::Name) => unreachable!("the answer comes after the body moved on"),
() = std::future::ready(()) => {}
}
co.yield_(Ask::Name).await.unwrap()
});
let stale = name(yielded(coroutine.resume().await));
stale.send("stale");
name(yielded(coroutine.resume().await)).send("fresh");
assert_eq!(complete(coroutine.resume().await), "fresh");
}
#[rstest]
#[case::returned(false)]
#[case::cancelled(true)]
#[tokio::test]
async fn a_finished_coroutine_refuses_to_resume(#[case] cancel: bool) {
let mut coroutine = suspended_once();
let reply = name(yielded(coroutine.resume().await));
if cancel {
coroutine.cancel();
} else {
reply.send("done");
complete(coroutine.resume().await).unwrap();
}
assert_eq!(coroutine.resume().await.unwrap_err(), ResumeError::Finished);
}
#[tokio::test]
async fn a_dropped_resume_leaves_the_coroutine_resumable() {
let mut coroutine: Test<u32> = Coroutine::new(|co| async move {
tokio::time::sleep(Duration::from_millis(20)).await;
co.yield_(Ask::Count).await.unwrap()
});
assert!(
timeout(Duration::from_millis(1), coroutine.resume())
.await
.is_err()
);
count(yielded(coroutine.resume().await)).send(3);
assert_eq!(complete(coroutine.resume().await), 3);
}
struct Dropped(Arc<Mutex<bool>>);
impl Drop for Dropped {
fn drop(&mut self) {
*self.0.lock().unwrap() = true;
}
}
#[tokio::test]
async fn cancel_drops_the_body() {
let dropped = Arc::new(Mutex::new(false));
let guard = Dropped(Arc::clone(&dropped));
let mut coroutine: Test<()> = Coroutine::new(|co| async move {
let _guard = guard;
co.yield_(Ask::Count).await.unwrap();
});
let _reply = yielded(coroutine.resume().await);
coroutine.cancel();
assert!(*dropped.lock().unwrap());
}
#[rstest]
#[case::cancelled(true)]
#[case::dropped(false)]
#[tokio::test]
async fn a_co_that_escaped_the_body_is_abandoned_once_the_coroutine_ends(#[case] cancel: bool) {
let escaped: Arc<Mutex<Option<Co<Ask>>>> = Arc::default();
let slot = Arc::clone(&escaped);
let mut coroutine: Test<()> = Coroutine::new(move |co| {
*slot.lock().unwrap() = Some(co.clone());
async move {
co.yield_(Ask::Count).await.unwrap();
}
});
let _reply = yielded(coroutine.resume().await);
let co = escaped.lock().unwrap().take().unwrap();
let waiting = tokio::spawn(async move { co.yield_(Ask::Name).await });
tokio::task::yield_now().await;
if cancel {
coroutine.cancel();
} else {
drop(coroutine);
}
let outcome = timeout(Duration::from_secs(1), waiting)
.await
.expect("an escaped yield waits forever")
.unwrap();
assert_eq!(outcome, Err(Abandoned));
}
#[rstest]
#[case::sent(true)]
#[case::dropped(false)]
#[tokio::test]
async fn a_detached_reply_settles_its_answer(#[case] send: bool) {
let (reply, answer) = reply::<u32>();
if send {
reply.send(5);
} else {
drop(reply);
}
assert_eq!(answer.await, if send { Ok(5) } else { Err(Abandoned) });
}

View file

@ -1,8 +1,8 @@
- Target invariants; implementation and runtime validation may lag these rules
- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`RouteHost` traits
- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`ProtocolHost` traits
- No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features
- The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business
- `RouteHost::invoke` receives the keyword view the adapter's `begin` returned, not the caller's dict; a route host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance)
- `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, not the caller's dict; a protocol host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance)
- A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is
- A failing `classify` is raised with the native error's text as its `__context__`, never swallowed
- Use standard PyO3 ownership and conversion APIs

View file

@ -6,13 +6,14 @@ license.workspace = true
repository.workspace = true
[dependencies]
bytes.workspace = true
futures-util.workspace = true
litellm-host.workspace = true
pyo3.workspace = true
pyo3-async-runtimes.workspace = true
pythonize.workspace = true
serde.workspace = true
tokio = { workspace = true, features = ["sync"] }
tokio = { workspace = true, features = ["rt", "sync"] }
[dev-dependencies]
rstest.workspace = true

View file

@ -1,5 +1,5 @@
use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest};
use litellm_host::route::Route;
use litellm_host::protocol::Protocol;
use pyo3::exceptions::PyRuntimeError;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
@ -85,7 +85,7 @@ pub trait PythonLifecycle: Send + Sync {
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
}
/// Why a route operation the host answered did not produce a result: the route's own code
/// Why a custom operation the host answered did not produce a result: the route's own code
/// rejected it, which the route classifies like any other native failure, or Python code
/// raised, which reaches the caller as it was raised.
#[derive(Debug)]
@ -100,45 +100,61 @@ impl<E> From<PyErr> for InvokeError<E> {
}
}
/// The Python side of one route: answers the route's own operations, builds the public
/// The Python side of one protocol: answers its custom operations, builds the public
/// response and classifies native failures into public exceptions.
pub trait RouteHost: Send + Sync {
type Route: Route<Error: std::fmt::Display>;
pub trait ProtocolHost: Send + Sync {
type Protocol: Protocol<Error: std::fmt::Display>;
/// The public exception a native failure maps to, kept as a value until the driver
/// raises it.
type Failure: Into<PyErr>;
/// `arguments` is the keyword view the lifecycle's `begin` produced, not the
/// caller's own dict. A route host that projects from it inherits whatever that
/// adapter rewrote.
fn invoke(
/// Projects the call's request. `arguments` is the keyword view the lifecycle's
/// `begin` produced, not the caller's own dict, so the projection inherits whatever
/// that adapter rewrote.
fn project(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: <Self::Route as Route>::Op,
) -> Result<<Self::Route as Route>::OpResult, InvokeError<<Self::Route as Route>::Error>>;
) -> Result<
<Self::Protocol as Protocol>::Projection,
InvokeError<<Self::Protocol as Protocol>::Error>,
>;
/// Answers `op` through its reply.
fn invoke(
&mut self,
py: Python<'_>,
op: <Self::Protocol as Protocol>::Op,
) -> Result<(), InvokeError<<Self::Protocol as Protocol>::Error>>;
fn complete(
&mut self,
py: Python<'_>,
response: <Self::Route as Route>::Response,
response: <Self::Protocol as Protocol>::Response,
) -> PyResult<Py<PyAny>>;
/// What the stream carries at hand-off, as the caller's stream receives it.
fn head(
&mut self,
py: Python<'_>,
head: <Self::Protocol as Protocol>::StreamHead,
) -> PyResult<Py<PyAny>>;
/// One streamed chunk as the caller receives it.
fn chunk(
&mut self,
py: Python<'_>,
chunk: <Self::Route as Route>::Chunk,
chunk: <Self::Protocol as Protocol>::Chunk,
) -> PyResult<Py<PyAny>>;
fn classify(
&self,
py: Python<'_>,
error: <Self::Route as Route>::Error,
error: <Self::Protocol as Protocol>::Error,
) -> PyResult<Self::Failure>;
fn host_error(error: &PyErr) -> <Self::Route as Route>::Error;
fn host_error(error: &PyErr) -> <Self::Protocol as Protocol>::Error;
fn close(&mut self, py: Python<'_>);

View file

@ -2,10 +2,11 @@ use std::sync::Arc;
use std::task::Poll;
use futures_util::future::{AbortHandle, Abortable};
use litellm_host::event::WireRequest;
use litellm_host::event::{FailureOrigin, Timing, epoch_seconds};
use litellm_host::host::{Demand, HostOp, HostResult, HostStep};
use litellm_host::host::{Demand, HostOp, HostStep, Reply};
use litellm_host::machine::{HostFailure, Machine, MachineStep};
use litellm_host::route::Route;
use litellm_host::protocol::Protocol;
use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError};
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
@ -13,21 +14,21 @@ use pyo3::types::PyDict;
use tokio::sync::Mutex;
use crate::adapter::{
InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state,
InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state,
};
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
use crate::handle::{Execution, ExecutionBody, ExecutionStep};
type RouteOf<H> = <H as RouteHost>::Route;
type ErrorOf<H> = <RouteOf<H> as Route>::Error;
type ResponseOf<H> = <RouteOf<H> as Route>::Response;
type NativeStep<H> = MachineStep<RouteOf<H>, ResponseOf<H>>;
type ProtocolOf<H> = <H as ProtocolHost>::Protocol;
type ErrorOf<H> = <ProtocolOf<H> as Protocol>::Error;
type ResponseOf<H> = <ProtocolOf<H> as Protocol>::Response;
type NativeStep<H> = MachineStep<ProtocolOf<H>, ResponseOf<H>>;
type NativeResult<H> = Result<NativeStep<H>, ErrorOf<H>>;
type NativeResume<H> = Option<Result<HostResult<RouteOf<H>>, HostFailure<ErrorOf<H>>>>;
type Interruption<H> = Option<HostFailure<ErrorOf<H>>>;
type MachineResult<M> = Result<
MachineStep<<M as Machine>::Route, <M as Machine>::Complete>,
<<M as Machine>::Route as Route>::Error,
MachineStep<<M as Machine>::Protocol, <M as Machine>::Complete>,
<<M as Machine>::Protocol as Protocol>::Error,
>;
struct MachineState<M: Machine> {
@ -44,12 +45,11 @@ enum Stage {
Failed(Py<PyBaseException>),
}
#[derive(Clone, Copy)]
enum Expect {
Started,
Arguments,
Wire,
Emitted,
Wire(Reply<WireRequest>),
Emitted(Reply<()>),
Response,
Terminal,
}
@ -58,20 +58,30 @@ enum Pending {
Native,
Adapter(Expect),
/// The stream handed to the caller waits for its next read or its close.
Consumer,
Consumer(Reply<Demand>),
}
enum Next<H: RouteHost> {
/// A route answer as the driver resumes on it: a Python exception interrupts the call as
/// raised, a native rejection resumes the machine with it.
fn answered<E>(answer: Result<(), InvokeError<E>>) -> PyResult<Result<(), E>> {
match answer {
Ok(()) => Ok(Ok(())),
Err(InvokeError::Native(error)) => Ok(Err(error)),
Err(InvokeError::Python(error)) => Err(error),
}
}
enum Next<H: ProtocolHost> {
Return(ExecutionStep),
Continue(HostStep<NativeResult<H>, Py<PyAny>>),
}
struct PythonDriver<H, M>
where
H: RouteHost,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
route: H,
host: H,
adapter: Box<dyn PythonLifecycle>,
machine: Option<Arc<Mutex<MachineState<M>>>>,
arguments: Option<Py<PyDict>>,
@ -89,17 +99,17 @@ where
pub fn run_call<H, M>(
py: Python<'_>,
machine: M,
route: H,
host: H,
adapter: Box<dyn PythonLifecycle>,
arguments: Py<PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>>
where
H: RouteHost + 'static,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost + 'static,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
let mut driver = PythonDriver {
route,
host,
adapter,
machine: Some(Arc::new(Mutex::new(MachineState {
machine,
@ -124,10 +134,10 @@ where
}
match driver.resume(None)? {
ExecutionStep::Return(value) => Ok(value),
ExecutionStep::Open => py
ExecutionStep::Open(head) => py
.import("litellm.rust_bridge.lifecycle")?
.getattr("SyncStream")?
.call1((Py::new(py, Execution::suspended(driver))?,))
.call1((Py::new(py, Execution::suspended(driver))?, head))
.map(Bound::unbind),
ExecutionStep::Await(_) | ExecutionStep::Yield(_) => {
Err(PyRuntimeError::new_err("sync call suspended"))
@ -141,8 +151,8 @@ fn is_cancellation(py: Python<'_>, error: &PyErr) -> bool {
impl<H, M> PythonDriver<H, M>
where
H: RouteHost,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
fn timing(&self) -> Timing {
Timing {
@ -172,13 +182,13 @@ where
self.run_steps(py, HostStep::Ready(result))
}
(Some(Pending::Native), Some(Err(error))) => self.interrupt(py, error),
(Some(Pending::Consumer), Some(read)) => {
let demand = if read.is_ok() {
(Some(Pending::Consumer(reply)), Some(read)) => {
reply.send(if read.is_ok() {
Demand::More
} else {
Demand::Detached
};
self.resume_machine(py, Some(Ok(HostResult::Demand(demand))))
});
self.resume_machine(py, None)
}
(Some(Pending::Adapter(expect)), Some(result)) => {
match self.adapter.resume(py, result) {
@ -196,22 +206,24 @@ where
step: LifecycleStep,
expect: Expect,
) -> PyResult<ExecutionStep> {
if let LifecycleStep::Await(awaitable) = step {
self.pending = Some(Pending::Adapter(expect));
return Ok(ExecutionStep::Await(awaitable));
}
match (expect, step) {
(_, LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(expect));
Ok(ExecutionStep::Await(awaitable))
}
(Expect::Started, LifecycleStep::Done) => self.begin(py),
(Expect::Arguments, LifecycleStep::Arguments(arguments)) => {
self.arguments = Some(arguments);
self.stage = Stage::Call;
self.resume_machine(py, None)
}
(Expect::Wire, LifecycleStep::Wire(wire)) => {
self.resume_machine(py, Some(Ok(HostResult::BeforeSend(wire))))
(Expect::Wire(reply), LifecycleStep::Wire(wire)) => {
reply.send(*wire);
self.resume_machine(py, None)
}
(Expect::Emitted, LifecycleStep::Done) => {
self.resume_machine(py, Some(Ok(HostResult::Emitted)))
(Expect::Emitted(reply), LifecycleStep::Done) => {
reply.send(());
self.resume_machine(py, None)
}
(Expect::Response, LifecycleStep::Response(response)) => self.succeeded(py, response),
(Expect::Terminal, LifecycleStep::Done) => match &self.stage {
@ -242,9 +254,9 @@ where
fn resume_machine(
&mut self,
py: Python<'_>,
result: NativeResume<H>,
interruption: Interruption<H>,
) -> PyResult<ExecutionStep> {
let step = self.resume_core(py, result)?;
let step = self.resume_core(py, interruption)?;
self.run_steps(py, step)
}
@ -277,54 +289,72 @@ where
}
Err(error) => return self.machine_failed(py, error).map(Next::Return),
};
let answer = match op {
HostOp::Route(op) => {
let answered = match op {
HostOp::Project(reply) => {
let arguments = self.arguments.as_ref().ok_or_else(missing_state)?;
match self.route.invoke(py, arguments.bind(py), op) {
Ok(result) => Ok(HostResult::Route(result)),
Err(InvokeError::Native(error)) => {
return self
.resume_core(py, Some(Err(HostFailure::Error(error))))
.map(Next::Continue);
}
Err(InvokeError::Python(error)) => Err(error),
}
let projected = self.host.project(py, arguments.bind(py));
answered(projected.map(|projection| reply.send(projection)))
}
HostOp::BeforeSend { wire, context } => {
match self.adapter.before_send(py, wire, &context) {
Ok(LifecycleStep::Wire(wire)) => Ok(HostResult::BeforeSend(wire)),
HostOp::Custom(op) => answered(self.host.invoke(py, op)),
HostOp::BeforeSend {
wire,
context,
reply,
} => match self.adapter.before_send(py, wire, &context) {
Ok(LifecycleStep::Wire(wire)) => {
reply.send(*wire);
Ok(Ok(()))
}
Ok(LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(Expect::Wire(reply)));
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
}
Ok(_) => return Err(missing_state()),
Err(error) => Err(error),
},
HostOp::Open(head, reply) => return self.opened(py, head, reply).map(Next::Return),
HostOp::Deliver(chunk, reply) => {
return self.delivered(py, chunk, reply).map(Next::Return);
}
HostOp::Emit(event, reply) => {
match self.adapter.emit(py, LifecycleEvent::Machine(&event)) {
Ok(LifecycleStep::Done) => {
reply.send(());
Ok(Ok(()))
}
Ok(LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(Expect::Wire));
self.pending = Some(Pending::Adapter(Expect::Emitted(reply)));
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
}
Ok(_) => return Err(missing_state()),
Err(error) => Err(error),
}
}
HostOp::Open(_) => return self.opened(py).map(Next::Return),
HostOp::Deliver(chunk) => return self.delivered(py, chunk).map(Next::Return),
HostOp::Emit(event) => match self.adapter.emit(py, LifecycleEvent::Machine(&event)) {
Ok(LifecycleStep::Done) => Ok(HostResult::Emitted),
Ok(LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(Expect::Emitted));
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
}
Ok(_) => return Err(missing_state()),
Err(error) => Err(error),
},
};
match answer {
Ok(answer) => self.resume_core(py, Some(Ok(answer))).map(Next::Continue),
match answered {
Ok(Ok(())) => self.resume_core(py, None).map(Next::Continue),
Ok(Err(native)) => self
.resume_core(py, Some(HostFailure::Error(native)))
.map(Next::Continue),
Err(error) => self.interrupt(py, error).map(Next::Return),
}
}
fn opened(&mut self, py: Python<'_>) -> PyResult<ExecutionStep> {
fn opened(
&mut self,
py: Python<'_>,
head: <ProtocolOf<H> as Protocol>::StreamHead,
reply: Reply<Demand>,
) -> PyResult<ExecutionStep> {
self.stage = Stage::Streaming;
let head = match self.host.head(py, head) {
Ok(head) => head,
Err(error) => return self.interrupt(py, error),
};
match self.adapter.opened(py) {
Ok(()) => {
self.pending = Some(Pending::Consumer);
Ok(ExecutionStep::Open)
self.pending = Some(Pending::Consumer(reply));
Ok(ExecutionStep::Open(head))
}
Err(error) => self.interrupt(py, error),
}
@ -333,15 +363,16 @@ where
fn delivered(
&mut self,
py: Python<'_>,
chunk: <RouteOf<H> as Route>::Chunk,
chunk: <ProtocolOf<H> as Protocol>::Chunk,
reply: Reply<Demand>,
) -> PyResult<ExecutionStep> {
let chunk = match self.route.chunk(py, chunk) {
let chunk = match self.host.chunk(py, chunk) {
Ok(chunk) => chunk,
Err(error) => return self.interrupt(py, error),
};
match self.adapter.delivered(py, &chunk) {
Ok(()) => {
self.pending = Some(Pending::Consumer);
self.pending = Some(Pending::Consumer(reply));
Ok(ExecutionStep::Yield(chunk))
}
Err(error) => self.interrupt(py, error),
@ -357,25 +388,24 @@ where
} else {
HostFailure::Error(native)
};
self.resume_machine(py, Some(Err(failure)))
self.resume_machine(py, Some(failure))
}
fn resume_core(
&mut self,
py: Python<'_>,
result: NativeResume<H>,
interruption: Interruption<H>,
) -> PyResult<HostStep<NativeResult<H>, Py<PyAny>>> {
let state = Arc::clone(self.machine.as_ref().ok_or_else(missing_state)?);
let future = async move {
let mut state = state.lock().await;
let result = match result {
Some(Err(failure)) => state
let result = match interruption {
Some(failure) => state
.machine
.interrupt(failure)
.await
.map(MachineStep::Complete),
Some(Ok(result)) => state.machine.resume(Some(result)).await,
None => state.machine.resume(None).await,
None => state.machine.resume().await,
};
state.result = Some(result);
Ok(())
@ -414,7 +444,7 @@ where
fn completed(&mut self, py: Python<'_>, response: ResponseOf<H>) -> PyResult<ExecutionStep> {
self.ended_at = Some(epoch_seconds());
let public = match self.route.complete(py, response) {
let public = match self.host.complete(py, response) {
Ok(public) => public,
Err(error) => return self.failure(py, error, FailureOrigin::Call),
};
@ -441,7 +471,7 @@ where
/// fails, that failure is raised with the native error's text as its `__context__`.
fn classified(&self, py: Python<'_>, error: ErrorOf<H>) -> PyErr {
let native = error.to_string();
let classifier_error = match self.route.classify(py, error) {
let classifier_error = match self.host.classify(py, error) {
Ok(failure) => return failure.into(),
Err(classifier_error) => classifier_error,
};
@ -486,7 +516,7 @@ where
if self.machine.take().is_some() {
Python::attach(|py| {
self.adapter.close(py);
self.route.close(py);
self.host.close(py);
});
}
}
@ -494,15 +524,15 @@ where
impl<H, M> ExecutionBody for PythonDriver<H, M>
where
H: RouteHost,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
fn resume(&mut self, result: Option<PyResult<Py<PyAny>>>) -> PyResult<ExecutionStep> {
Python::attach(|py| self.drive(py, result))
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
self.route.traverse(visit)?;
self.host.traverse(visit)?;
self.adapter.traverse(visit)?;
visit.call(&self.arguments)?;
visit.call(&self.interrupted)?;
@ -516,8 +546,8 @@ where
impl<H, M> Drop for PythonDriver<H, M>
where
H: RouteHost,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
fn drop(&mut self) {
self.clear();
@ -528,8 +558,8 @@ where
mod tests {
use std::sync::{Arc, Mutex};
use litellm_host::event::{MachineEvent, RequestContext, WireRequest};
use litellm_host::machine::{Interrupted, Step};
use litellm_host::event::{MachineEvent, RawResponse, RequestContext};
use litellm_host::machine::{CallMachine, MachineFault};
use pyo3::exceptions::{PyBaseException, PyValueError};
use pyo3::types::PyDict;
@ -573,22 +603,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
struct Synthetic;
impl Route for Synthetic {
type Response = String;
type Error = Error;
type Op = &'static str;
type OpResult = String;
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
impl From<MachineFault> for Error {
fn from(fault: MachineFault) -> Self {
Self(format!("{fault:?}"))
}
}
/// Yields the scripted ops in order, then completes or fails as scripted.
struct ScriptedMachine {
ops: Vec<HostOp<Synthetic>>,
outcome: Option<Result<String, Error>>,
answers: Vec<String>,
struct Synthetic;
impl Protocol for Synthetic {
type Response = String;
type Error = Error;
type Projection = String;
type Op = (&'static str, Reply<String>);
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
fn wire() -> WireRequest {
@ -609,37 +638,6 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
impl Machine for ScriptedMachine {
type Route = Synthetic;
type Complete = String;
fn resume(&mut self, result: Option<HostResult<Synthetic>>) -> Step<'_, Self> {
Box::pin(async move {
if let Some(result) = result {
self.answers.push(match result {
HostResult::Route(value) => value,
HostResult::BeforeSend(wire) => wire.url,
HostResult::Emitted => "emitted".into(),
HostResult::Demand(demand) => format!("{demand:?}"),
});
}
if !self.ops.is_empty() {
return Ok(MachineStep::Host(self.ops.remove(0)));
}
self.outcome
.take()
.ok_or_else(|| Error("resumed after completion".into()))?
.map(MachineStep::Complete)
})
}
fn interrupt(&mut self, failure: HostFailure<Error>) -> Interrupted<'_, Self> {
self.ops.clear();
self.outcome = None;
Box::pin(async move { Err(failure.into_error()) })
}
}
#[derive(Default)]
struct Log(Arc<Mutex<Vec<String>>>);
@ -677,22 +675,41 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
impl RouteHost for SyntheticHost {
type Route = Synthetic;
impl SyntheticHost {
fn answer(&self, value: impl FnOnce() -> String) -> Result<String, InvokeError<Error>> {
match self.op {
OpScript::Answer => Ok(value()),
OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()),
OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))),
}
}
}
impl ProtocolHost for SyntheticHost {
type Protocol = Synthetic;
type Failure = Classified;
fn project(
&mut self,
_: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> Result<String, InvokeError<Error>> {
self.log.push("project");
self.answer(|| format!("project:{}", arguments.len()))
}
fn invoke(
&mut self,
_: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: &'static str,
) -> Result<String, InvokeError<Error>> {
self.log.push(format!("route:{op}"));
match self.op {
OpScript::Answer => Ok(format!("{op}:{}", arguments.len())),
OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()),
OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))),
}
(op, reply): (&'static str, Reply<String>),
) -> Result<(), InvokeError<Error>> {
self.log.push(format!("op:{op}"));
self.answer(|| op.to_string())
.map(|answer| reply.send(answer))
}
fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult<Py<PyAny>> {
match head {}
}
fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult<Py<PyAny>> {
@ -719,7 +736,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
fn close(&mut self, _: Python<'_>) {
self.log.push("route.close");
self.log.push("host.close");
}
fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> {
@ -828,7 +845,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
fn run_scripted(
py: Python<'_>,
machine: ScriptedMachine,
machine: CallMachine<Synthetic>,
op: OpScript,
script: AdapterScript,
asynchronous: bool,
@ -848,12 +865,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
fn run_hosted(
py: Python<'_>,
machine: ScriptedMachine,
route: SyntheticHost,
machine: CallMachine<Synthetic>,
host: SyntheticHost,
script: AdapterScript,
asynchronous: bool,
) -> (PyResult<Py<PyAny>>, Vec<String>) {
let log = Log(route.log.0.clone());
let log = Log(host.log.0.clone());
let adapter = SyntheticAdapter {
log: Log(log.0.clone()),
script,
@ -863,7 +880,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
let result = run_call(
py,
machine,
route,
host,
Box::new(adapter),
arguments.unbind(),
asynchronous,
@ -884,21 +901,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
(result, log.entries())
}
fn success_machine() -> ScriptedMachine {
ScriptedMachine {
ops: vec![
HostOp::Route("project"),
HostOp::BeforeSend {
wire: Box::new(wire()),
context: Box::new(context()),
},
HostOp::Emit(MachineEvent::ResponseReceived {
raw: litellm_host::event::RawResponse { body: "raw".into() },
}),
],
outcome: Some(Ok("done".into())),
answers: Vec::new(),
}
/// Answers to projection, to the route op and to `before_send` all reach the
/// response, so a driver that misroutes a reply changes what the call returns.
fn success_machine() -> CallMachine<Synthetic> {
CallMachine::new(|host| {
Box::pin(async move {
let projected = host.project().await?;
let signed = host.custom_op(|reply| ("sign", reply)).await?;
let wire = host.before_send(wire(), context()).await?;
host.emit(MachineEvent::ResponseReceived {
raw: RawResponse { body: "raw".into() },
})
.await?;
Ok(format!("{projected}|{signed}|{}", wire.url))
})
})
}
#[test]
@ -917,32 +934,194 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
AdapterScript::Plain,
asynchronous,
);
assert_eq!(result.unwrap().extract::<String>(py).unwrap(), "done");
assert_eq!(
result.unwrap().extract::<String>(py).unwrap(),
"project:1|sign|rewritten"
);
assert_eq!(
log,
[
"started",
"begin",
"route:project",
"project",
"op:sign",
"before_send",
"response:raw",
"complete",
"after_success",
"succeeded:done",
"succeeded:project:1|sign|rewritten",
"adapter.close",
"route.close",
"host.close",
]
);
}
});
}
fn failing_machine() -> ScriptedMachine {
ScriptedMachine {
ops: vec![HostOp::Route("project")],
outcome: Some(Err(Error("provider exploded".into()))),
answers: Vec::new(),
struct Streaming;
impl Protocol for Streaming {
type Response = ();
type Error = Error;
type Projection = ();
type Op = std::convert::Infallible;
type Chunk = &'static str;
type StreamHead = Vec<(&'static str, &'static str)>;
}
struct StreamingHost;
impl ProtocolHost for StreamingHost {
type Protocol = Streaming;
type Failure = Classified;
fn project(
&mut self,
_: Python<'_>,
_: &Bound<'_, PyDict>,
) -> Result<(), InvokeError<Error>> {
Ok(())
}
fn invoke(
&mut self,
_: Python<'_>,
op: std::convert::Infallible,
) -> Result<(), InvokeError<Error>> {
match op {}
}
fn head(
&mut self,
py: Python<'_>,
head: Vec<(&'static str, &'static str)>,
) -> PyResult<Py<PyAny>> {
let headers = PyDict::new(py);
for (name, value) in head {
headers.set_item(name, value)?;
}
let hidden = PyDict::new(py);
hidden.set_item("additional_headers", headers)?;
Ok(hidden.into_any().unbind())
}
fn chunk(&mut self, py: Python<'_>, chunk: &'static str) -> PyResult<Py<PyAny>> {
Ok(pyo3::types::PyString::new(py, chunk).into_any().unbind())
}
fn complete(&mut self, py: Python<'_>, (): ()) -> PyResult<Py<PyAny>> {
Ok(py.None())
}
fn classify(&self, _: Python<'_>, error: Error) -> PyResult<Classified> {
Ok(Classified(error.0))
}
fn host_error(error: &PyErr) -> Error {
Error(error.to_string())
}
fn close(&mut self, _: Python<'_>) {}
fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> {
Ok(())
}
}
fn streaming_machine() -> CallMachine<Streaming> {
CallMachine::new(|host| {
Box::pin(async move {
host.project().await?;
if host.open(vec![("request-id", "req_1")]).await? == Demand::Detached {
return Ok(());
}
for chunk in ["first", "second"] {
if host.deliver(chunk).await? == Demand::Detached {
break;
}
}
Ok(())
})
})
}
/// Drives a `Stream` (async) or `SyncStream` to completion from a sync test.
fn read_all(py: Python<'_>, stream: &Bound<'_, PyAny>, asynchronous: bool) -> Vec<String> {
if !asynchronous {
return stream
.try_iter()
.unwrap()
.map(|chunk| chunk.unwrap().extract().unwrap())
.collect();
}
std::iter::from_fn(|| {
let stop = stream
.call_method0("__anext__")
.unwrap()
.call_method1("send", (py.None(),))
.unwrap_err();
if stop.is_instance_of::<pyo3::exceptions::PyStopAsyncIteration>(py) {
return None;
}
assert!(stop.is_instance_of::<pyo3::exceptions::PyStopIteration>(py));
Some(stop.value(py).getattr("value").unwrap().extract().unwrap())
})
.collect()
}
#[test]
fn a_stream_carries_its_head_as_hidden_params_before_the_first_chunk() {
let _guard = PYTHON_GLOBALS
.lock()
.unwrap_or_else(|error| error.into_inner());
crate::initialize_python();
Python::attach(|py| {
install_lifecycle_module(py);
for asynchronous in [false, true] {
let log = Log::default();
let adapter = SyntheticAdapter {
log: Log(log.0.clone()),
script: AdapterScript::Plain,
};
let handed = run_call(
py,
streaming_machine(),
StreamingHost,
Box::new(adapter),
PyDict::new(py).unbind(),
asynchronous,
)
.unwrap();
let stream = if asynchronous {
let stop = handed.call_method1(py, "send", (py.None(),)).unwrap_err();
stop.value(py).getattr("value").unwrap()
} else {
handed.into_bound(py)
};
let hidden: std::collections::HashMap<
String,
std::collections::HashMap<String, String>,
> = stream.getattr("_hidden_params").unwrap().extract().unwrap();
assert_eq!(
hidden["additional_headers"],
std::collections::HashMap::from([(
"request-id".to_string(),
"req_1".to_string()
)])
);
assert_eq!(log.entries(), ["started", "begin", "opened"]);
assert_eq!(read_all(py, &stream, asynchronous), ["first", "second"]);
}
});
}
fn failing_machine() -> CallMachine<Synthetic> {
CallMachine::new(|host| {
Box::pin(async move {
host.project().await?;
Err(Error("provider exploded".into()))
})
})
}
#[test]
@ -969,11 +1148,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
[
"started",
"begin",
"route:project",
"project",
"classify:provider exploded",
"failed:Call:classified: provider exploded",
"adapter.close",
"route.close",
"host.close",
]
);
}
@ -1003,11 +1182,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
[
"started",
"begin",
"route:project",
"project",
"classify:op rejected",
"failed:Call:classified: op rejected",
"adapter.close",
"route.close",
"host.close",
]
);
});
@ -1035,10 +1214,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
[
"started",
"begin",
"route:project",
"project",
"failed:Call:op failed",
"adapter.close",
"route.close",
"host.close",
]
);
});
@ -1073,11 +1252,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
[
"started",
"begin",
"route:project",
"project",
"classify:provider exploded",
"failed:Call:classifier failed",
"adapter.close",
"route.close",
"host.close",
]
);
});
@ -1106,7 +1285,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
"begin",
"failed:Host:begin failed",
"adapter.close",
"route.close"
"host.close"
]
);
});
@ -1130,7 +1309,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
);
assert_eq!(result.unwrap().extract::<String>(py).unwrap(), "replaced");
assert!(log.contains(&"succeeded:replaced".to_string()));
assert!(!log.contains(&"succeeded:done".to_string()));
assert!(!log.contains(&"succeeded:project:1|rewritten".to_string()));
}
});
}
@ -1159,7 +1338,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
"after_success",
"failed:Host:after_success failed",
"adapter.close",
"route.close"
"host.close"
]
);
assert!(!log.iter().any(|entry| entry.starts_with("succeeded")));
@ -1175,18 +1354,31 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
crate::initialize_python();
Python::attach(|py| {
struct Cancelling(Log);
impl RouteHost for Cancelling {
type Route = Synthetic;
impl ProtocolHost for Cancelling {
type Protocol = Synthetic;
type Failure = Classified;
fn invoke(
fn project(
&mut self,
_: Python<'_>,
_: &Bound<'_, PyDict>,
_: &'static str,
) -> Result<String, InvokeError<Error>> {
self.0.push("route");
self.0.push("project");
Err(pyo3::exceptions::asyncio::CancelledError::new_err(()).into())
}
fn invoke(
&mut self,
_: Python<'_>,
_: (&'static str, Reply<String>),
) -> Result<(), InvokeError<Error>> {
Err(missing_state().into())
}
fn head(
&mut self,
_: Python<'_>,
head: std::convert::Infallible,
) -> PyResult<Py<PyAny>> {
match head {}
}
fn chunk(
&mut self,
_: Python<'_>,
@ -1210,7 +1402,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
let log = Log::default();
let route = Cancelling(Log(log.0.clone()));
let host = Cancelling(Log(log.0.clone()));
let adapter = SyntheticAdapter {
log: Log(log.0.clone()),
script: AdapterScript::Plain,
@ -1218,7 +1410,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
let error = run_call(
py,
success_machine(),
route,
host,
Box::new(adapter),
PyDict::new(py).unbind(),
false,
@ -1227,7 +1419,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
assert!(!error.is_instance_of::<pyo3::exceptions::PyException>(py));
assert_eq!(
log.entries(),
["started", "begin", "route", "adapter.close"]
["started", "begin", "project", "adapter.close"]
);
});
}

View file

@ -225,29 +225,12 @@ mod tests {
use pyo3::exceptions::PyLookupError;
use pyo3::panic::PanicException;
use pyo3::types::{PyDict, PyModule};
use rstest::{fixture, rstest};
use rstest::rstest;
use serde::Serializer;
use tokio::runtime::Builder;
use super::*;
struct InitializedPython;
impl InitializedPython {
fn attach<F, R>(&self, f: F) -> R
where
F: for<'py> FnOnce(Python<'py>) -> R,
{
Python::attach(f)
}
}
#[fixture]
#[once]
fn initialized_python() -> InitializedPython {
crate::initialize_python();
InitializedPython
}
use crate::{InitializedPython, initialized_python};
#[derive(Debug)]
struct Error(String);

View file

@ -0,0 +1,241 @@
//! A caller's file-like object: anything with a callable `read`, kept as a handle and read
//! once, on the host's thread, into bytes Rust owns.
use bytes::Bytes;
use pyo3::{
exceptions::PyTypeError,
gc::{PyTraverseError, PyVisit},
prelude::*,
pybacked::PyBackedBytes,
types::{PyBytes, PyString},
};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct FileContent {
pub bytes: Bytes,
pub file_name: Option<String>,
}
#[derive(Debug)]
pub struct PythonFileReader {
reader: Py<PyAny>,
name: Option<String>,
}
impl PythonFileReader {
/// `None` when `file` has no callable `read`. The object's `name` is read now, its
/// contents only on [`read`](Self::read).
pub fn from_file_like(file: &Bound<'_, PyAny>) -> PyResult<Option<Self>> {
let reader = file
.getattr_opt("read")?
.filter(|value| value.is_callable());
let Some(reader) = reader else {
return Ok(None);
};
let name = file
.getattr_opt("name")?
.filter(|value| !value.is_none())
.map(|value| value.extract::<String>())
.transpose()?;
Ok(Some(Self {
reader: reader.unbind(),
name,
}))
}
pub fn read(&self, py: Python<'_>) -> PyResult<FileContent> {
let value = self.reader.bind(py).call0()?;
let bytes = if value.is_instance_of::<PyString>() {
Bytes::from(value.extract::<String>()?)
} else if value.is_instance_of::<PyBytes>() {
py_bytes(&value)?
} else {
return Err(PyTypeError::new_err(format!(
"file read must return bytes or str, got {}",
value.get_type(),
)));
};
Ok(FileContent {
bytes,
file_name: self.name.clone(),
})
}
pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.reader)
}
}
/// An exact `bytes` object is shared without copying and keeps the Python object alive;
/// a `bytes` subclass is copied.
pub fn py_bytes(value: &Bound<'_, PyAny>) -> PyResult<Bytes> {
if value.is_exact_instance_of::<PyBytes>() {
return Ok(Bytes::from_owner(value.extract::<PyBackedBytes>()?));
}
Ok(Bytes::copy_from_slice(
value.extract::<PyBackedBytes>()?.as_ref(),
))
}
#[cfg(test)]
mod tests {
use pyo3::{exceptions::PyTypeError, types::PyDict};
use super::*;
fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> {
let locals = PyDict::new(py);
py.run(source, Some(&locals), Some(&locals)).unwrap();
locals
}
fn reader<'py>(locals: &Bound<'py, PyDict>, name: &str) -> PythonFileReader {
PythonFileReader::from_file_like(&locals.get_item(name).unwrap().unwrap())
.unwrap()
.unwrap()
}
#[test]
fn objects_without_a_callable_read_are_not_readers() {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"
class Attribute:
read = 'not callable'
plain = object()
attribute = Attribute()
",
);
for name in ["plain", "attribute"] {
let file = locals.get_item(name).unwrap().unwrap();
assert!(PythonFileReader::from_file_like(&file).unwrap().is_none());
}
});
}
#[test]
fn the_name_is_taken_up_front_and_the_contents_only_on_read() {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"
class Reader:
name = 'scan.png'
def __init__(self):
self.reads = 0
def read(self):
self.reads += 1
return b'abc'
file = Reader()
",
);
let reads = || {
locals
.get_item("file")
.unwrap()
.unwrap()
.getattr("reads")
.unwrap()
.extract::<usize>()
.unwrap()
};
let file = reader(&locals, "file");
assert_eq!(reads(), 0);
let content = file.read(py).unwrap();
assert_eq!(reads(), 1);
assert_eq!(
content,
FileContent {
bytes: b"abc".as_slice().into(),
file_name: Some("scan.png".into()),
}
);
});
}
#[test]
fn read_results_are_normalized_and_exceptions_keep_their_identity() {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"
failure = KeyError('reader failed')
class Raising:
def read(self):
raise failure
class Text:
def read(self):
return 'héllo'
class Wrong:
def read(self):
return 7
raising = Raising()
text = Text()
wrong = Wrong()
",
);
let error = reader(&locals, "raising").read(py).unwrap_err();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
assert_eq!(
reader(&locals, "text").read(py).unwrap().bytes.as_ref(),
"héllo".as_bytes()
);
let error = reader(&locals, "wrong").read(py).unwrap_err();
assert!(error.is_instance_of::<PyTypeError>(py));
assert!(error.to_string().contains("bytes or str"));
});
}
#[rstest::rstest]
#[case::read("read")]
#[case::name("name")]
fn attribute_failures_keep_their_identity(#[case] attribute: &str) {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"
failure = LookupError('file property failed')
class File:
def __getattribute__(self, name):
if name == attribute:
raise failure
return super().__getattribute__(name)
name = 'scan.pdf'
def read(self):
return b'abc'
file = File()
",
);
locals.set_item("attribute", attribute).unwrap();
let error =
PythonFileReader::from_file_like(&locals.get_item("file").unwrap().unwrap())
.unwrap_err();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
});
}
#[test]
fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() {
Python::initialize();
let (bytes, pointer) = Python::attach(|py| {
let value = PyBytes::new(py, b"document bytes");
let pointer = value.as_bytes().as_ptr() as usize;
(py_bytes(value.as_any()).unwrap(), pointer)
});
assert_eq!(bytes.as_ptr() as usize, pointer);
assert_eq!(bytes.as_ref(), b"document bytes");
}
}

View file

@ -1,6 +1,57 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{
Arc, Mutex,
atomic::{AtomicU64, Ordering},
};
use pyo3::prelude::*;
use pyo3::{exceptions::PyRuntimeError, prelude::*};
/// The caller's `contextvars` context, captured at the Python entry point so blocking Python
/// work started from Rust sees the same request-local values as the Python caller.
#[derive(Clone)]
pub struct PythonContext(Arc<Py<PyAny>>);
impl PythonContext {
pub fn capture(py: Python<'_>) -> PyResult<Self> {
Ok(Self(Arc::new(
py.import("contextvars")?
.call_method0("copy_context")?
.unbind(),
)))
}
/// Runs `f` inside a fresh copy of the captured context. The copy is what lets two blocking
/// calls run concurrently: a `contextvars.Context` cannot be entered twice at once.
pub fn enter<T, F>(&self, py: Python<'_>, f: F) -> PyResult<T>
where
F: FnOnce(Python<'_>) -> T + Send + Sync + 'static,
T: Send + 'static,
{
let copy = self.0.bind(py).call_method0("copy")?;
let body = Arc::new(Mutex::new(Some(f)));
let value = Arc::new(Mutex::new(None::<T>));
let callback = {
let body = Arc::clone(&body);
let value = Arc::clone(&value);
pyo3::types::PyCFunction::new_closure(py, None, None, move |_, _| {
Python::attach(|py| -> PyResult<()> {
let body = body
.lock()
.expect("context body slot poisoned")
.take()
.expect("the context body ran more than once");
*value.lock().expect("context value slot poisoned") = Some(body(py));
Ok(())
})
})?
};
copy.call_method1("run", (callback,))?;
value
.lock()
.expect("context value slot poisoned")
.take()
.ok_or_else(|| PyRuntimeError::new_err("the context body produced no value"))
}
}
static GIL_RELEASES: AtomicU64 = AtomicU64::new(0);
@ -19,3 +70,270 @@ where
pub fn release_count() -> u64 {
GIL_RELEASES.load(Ordering::Relaxed)
}
/// Runs Python work that may block, such as a secret manager read or a callback that does
/// I/O, on the runtime's blocking pool so the async workers stay free to poll other calls.
/// The work runs inside a copy of `context` so request-local `contextvars` survive the hop.
pub async fn attach_blocking<T, F>(context: PythonContext, f: F) -> PyResult<T>
where
F: for<'py> FnOnce(Python<'py>) -> T + Send + Sync + 'static,
T: Send + 'static,
{
match tokio::task::spawn_blocking(move || Python::attach(|py| context.enter(py, f))).await {
Ok(value) => value,
Err(error) if error.is_panic() => std::panic::resume_unwind(error.into_panic()),
Err(error) => panic!("the blocking pool dropped a Python call: {error}"),
}
}
#[cfg(test)]
mod tests {
use std::time::{Duration, Instant};
use pyo3::{exceptions::PyRuntimeError, prelude::*, types::PyDict};
use rstest::{fixture, rstest};
use super::{PythonContext, attach_blocking};
use crate::{InitializedPython, initialized_python, run_sync_value};
#[fixture]
fn namespace(#[from(initialized_python)] python: &InitializedPython) -> Py<PyDict> {
python.attach(|py| {
let namespace = PyDict::new(py);
py.run(
c"
import threading, time
finished = False
def work(seconds):
global finished
time.sleep(seconds)
finished = True
def observe(expression):
return eval(expression)
",
Some(&namespace),
None,
)
.unwrap();
namespace.unbind()
})
}
fn observe<T: for<'a, 'py> FromPyObject<'a, 'py, Error: std::fmt::Debug>>(
namespace: &Py<PyDict>,
py: Python<'_>,
expression: &str,
) -> T {
namespace
.bind(py)
.get_item("observe")
.unwrap()
.unwrap()
.call1((expression,))
.unwrap()
.extract()
.unwrap()
}
fn work(namespace: &Py<PyDict>, py: Python<'_>, seconds: f64) {
namespace
.bind(py)
.get_item("work")
.unwrap()
.unwrap()
.call1((seconds,))
.unwrap();
}
fn shared(namespace: &Py<PyDict>) -> Py<PyDict> {
Python::attach(|py| namespace.clone_ref(py))
}
fn context() -> PythonContext {
Python::attach(|py| PythonContext::capture(py).unwrap())
}
#[fixture]
fn request_context() -> (PythonContext, Py<PyAny>) {
Python::initialize();
Python::attach(|py| {
let namespace = PyDict::new(py);
py.run(
c"import contextvars\nrequest_var = contextvars.ContextVar('request_var', default='unset')",
Some(&namespace),
None,
)
.unwrap();
let var = namespace.get_item("request_var").unwrap().unwrap();
var.call_method1("set", ("request-value",)).unwrap();
(PythonContext::capture(py).unwrap(), var.unbind())
})
}
#[rstest]
#[tokio::test]
async fn blocking_python_work_leaves_the_runtime_free_to_run_other_tasks(
namespace: Py<PyDict>,
) {
let (python_done, timer_done) = tokio::join!(
attach_blocking(context(), move |py| {
work(&namespace, py, 0.3);
Instant::now()
}),
async {
tokio::time::sleep(Duration::from_millis(30)).await;
Instant::now()
}
);
assert!(
timer_done < python_done.unwrap(),
"the timer only finished after the Python call: the call ran inline on the worker"
);
}
#[rstest]
#[tokio::test]
async fn python_work_runs_off_the_thread_polling_the_future(namespace: Py<PyDict>) {
let polling: u64 = Python::attach(|py| observe(&namespace, py, "threading.get_ident()"));
let worker: u64 = attach_blocking(context(), move |py| {
observe(&namespace, py, "threading.get_ident()")
})
.await
.unwrap();
assert_ne!(worker, polling);
}
#[rstest]
#[tokio::test]
async fn a_dropped_await_never_interrupts_the_python_call(namespace: Py<PyDict>) {
let handle = shared(&namespace);
let started = tokio::time::timeout(
Duration::from_millis(10),
attach_blocking(context(), move |py| work(&handle, py, 0.1)),
)
.await;
assert!(
started.is_err(),
"the await was dropped before the call returned"
);
tokio::time::sleep(Duration::from_millis(300)).await;
let finished: bool = Python::attach(|py| observe(&namespace, py, "finished"));
assert!(finished);
}
#[rstest]
#[tokio::test]
async fn a_panic_in_python_work_reaches_the_awaiting_task(
#[from(initialized_python)] _python: &InitializedPython,
) {
let joined = tokio::spawn(attach_blocking(context(), |_| -> () {
panic!("python work failed")
}))
.await;
let error = joined.expect_err("the panic propagates instead of being swallowed");
assert!(error.is_panic());
}
#[rstest]
fn a_sync_route_can_await_python_work_without_deadlocking_on_the_gil(
#[from(initialized_python)] python: &InitializedPython,
) {
let value = python
.attach(|py| {
let context = PythonContext::capture(py).unwrap();
run_sync_value(py, async move {
tokio::time::timeout(
Duration::from_secs(5),
attach_blocking(context, |_| Python::version_str().len()),
)
.await
.map_err(|_| {
PyRuntimeError::new_err("the blocking call never re-acquired the GIL")
})
})
})
.unwrap()
.unwrap();
assert!(value > 0);
}
#[rstest]
#[tokio::test]
async fn blocking_work_sees_the_callers_contextvars(
request_context: (PythonContext, Py<PyAny>),
) {
let (context, var) = request_context;
let seen: String = attach_blocking(context, move |py| {
var.bind(py)
.call_method0("get")
.unwrap()
.extract::<String>()
.unwrap()
})
.await
.unwrap();
assert_eq!(seen, "request-value");
}
#[rstest]
#[tokio::test]
async fn concurrent_blocking_calls_each_enter_a_context_copy(
request_context: (PythonContext, Py<PyAny>),
) {
let (context, var) = request_context;
let first = Python::attach(|py| var.clone_ref(py));
let second = var;
let (first_seen, second_seen) = tokio::join!(
attach_blocking(context.clone(), move |py| {
first
.bind(py)
.call_method0("get")
.unwrap()
.extract::<String>()
.unwrap()
}),
attach_blocking(context, move |py| {
second
.bind(py)
.call_method0("get")
.unwrap()
.extract::<String>()
.unwrap()
}),
);
assert_eq!(first_seen.unwrap(), "request-value");
assert_eq!(second_seen.unwrap(), "request-value");
}
#[rstest]
#[tokio::test]
async fn writes_inside_blocking_work_do_not_leak_back_to_the_caller(
request_context: (PythonContext, Py<PyAny>),
) {
let (context, var) = request_context;
let leaked = Python::attach(|py| var.clone_ref(py));
attach_blocking(context, move |py| {
var.bind(py).call_method1("set", ("worker-value",)).unwrap();
})
.await
.unwrap();
let caller_value: String = Python::attach(|py| {
leaked
.bind(py)
.call_method0("get")
.unwrap()
.extract()
.unwrap()
});
assert_ne!(caller_value, "worker-value");
}
}

View file

@ -8,9 +8,9 @@ use pyo3::prelude::*;
pub enum ExecutionStep {
Return(Py<PyAny>),
Await(Py<PyAny>),
/// The call streams: the caller gets a stream over this execution, which stays
/// suspended until the stream asks for a chunk.
Open,
/// The call streams: the caller gets a stream over this execution carrying this head,
/// and the execution stays suspended until the stream asks for a chunk.
Open(Py<PyAny>),
Yield(Py<PyAny>),
}
@ -75,7 +75,7 @@ impl Execution {
let step = body.resume(result)?;
let (tag, value, suspended) = match step {
ExecutionStep::Await(value) => ("Await", value, true),
ExecutionStep::Open => ("Open", py.None(), true),
ExecutionStep::Open(head) => ("Open", head, true),
ExecutionStep::Yield(value) => ("Yield", value, true),
ExecutionStep::Return(value) => ("Complete", value, false),
};

View file

@ -1,6 +1,6 @@
//! The CPython runtime adapter: value marshalling, interpreter detachment, the tokio and
//! asyncio glue, and the driver that runs a native [`Machine`](litellm_host::machine::Machine)
//! against a Python route host and a Python lifecycle. Everything here is Python-specific by
//! against a Python protocol host and a Python lifecycle. Everything here is Python-specific by
//! construction; another host language gets its own crate of the same shape.
mod adapter;
@ -8,13 +8,14 @@ mod argument;
mod callable;
mod driver;
mod execution;
mod file_reader;
mod fork_gate;
mod gil;
mod handle;
mod marshal;
pub use adapter::{
InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state,
InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state,
};
pub use argument::lookup;
pub use callable::wrap_failure;
@ -24,8 +25,9 @@ pub use execution::{
reserve_process_for_forking, run_async, run_async_value, run_sync, run_sync_value,
runtime_started,
};
pub use file_reader::{FileContent, PythonFileReader, py_bytes};
pub use fork_gate::RuntimeAlreadyStarted;
pub use gil::{release_count, release_gil};
pub use gil::{PythonContext, attach_blocking, release_count, release_gil};
pub use handle::{Execution, ExecutionBody, ExecutionStep};
pub use marshal::{
Pythonized, from_py, from_py_argument, json_loads, json_object_field, panic_to_pyerr, to_py,
@ -43,3 +45,24 @@ pub(crate) fn initialize_python() {
});
});
}
#[cfg(test)]
pub(crate) struct InitializedPython;
#[cfg(test)]
impl InitializedPython {
pub(crate) fn attach<F, R>(&self, f: F) -> R
where
F: for<'py> FnOnce(pyo3::Python<'py>) -> R,
{
pyo3::Python::attach(f)
}
}
#[cfg(test)]
#[rstest::fixture]
#[once]
pub(crate) fn initialized_python() -> InitializedPython {
initialize_python();
InitializedPython
}

View file

@ -7,6 +7,7 @@ repository.workspace = true
[dependencies]
litellm-auth.workspace = true
litellm-coroutine.workspace = true
serde_json.workspace = true
tokio = { workspace = true, features = ["sync"] }

View file

@ -1,28 +1,27 @@
use std::future::Future;
use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
use crate::route::Route;
pub use litellm_coroutine::{Abandoned, Answer, Reply, reply};
/// One suspension point of a native call, performed by the host.
pub enum HostOp<R: Route> {
Route(R::Op),
use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
use crate::protocol::Protocol;
/// One suspension point of a native call, performed by the host and answered through the
/// [`Reply`] it carries.
pub enum HostOp<R: Protocol> {
/// The first op of every call: the caller's request as the host projects it.
Project(Reply<R::Projection>),
Custom(R::Op),
BeforeSend {
wire: Box<WireRequest>,
context: Box<RequestContext>,
reply: Reply<WireRequest>,
},
Emit(MachineEvent),
Emit(MachineEvent, Reply<()>),
/// The response streams: the host hands the caller a stream and answers once the
/// caller asks for the first chunk or goes away.
Open(R::StreamHead),
Open(R::StreamHead, Reply<Demand>),
/// The next chunk of an open stream, answered once the caller asks for the one after.
Deliver(R::Chunk),
}
pub enum HostResult<R: Route> {
Route(R::OpResult),
BeforeSend(Box<WireRequest>),
Emitted,
Demand(Demand),
Deliver(R::Chunk, Reply<Demand>),
}
/// Whether the caller of a streamed call still reads it.
@ -39,10 +38,13 @@ pub enum HostStep<V, S> {
Suspend(S),
}
/// An in-process host: answers route operations and observes the call without leaving
/// An in-process host: answers custom operations and observes the call without leaving
/// the Rust runtime. Language hosts implement their own driver instead.
pub trait Host<R: Route>: Send + Sync {
fn route(&self, op: R::Op) -> impl Future<Output = Result<R::OpResult, R::Error>> + Send;
pub trait Host<R: Protocol>: Send + Sync {
fn project(&self) -> impl Future<Output = Result<R::Projection, R::Error>> + Send;
/// Answers `op` through its reply, or fails the call.
fn custom_op(&self, op: R::Op) -> impl Future<Output = Result<(), R::Error>> + Send;
fn before_send(
&self,

View file

@ -1,12 +1,13 @@
//! The contract between a native call and the host runtime that drives it.
//!
//! A host is whatever sits on the far side of the language boundary: CPython today,
//! another runtime later. Core runs each route on a [`machine::RouteMachine`] and never learns
//! another runtime later. Core runs each route on a [`machine::CallMachine`] and never learns
//! which host is on the other end. The machine yields [`host::HostOp`]s; a driver answers
//! them, observes [`event::CallEvent`]s and may rewrite the wire request before it is sent.
//! each through the typed [`host::Reply`] it carries, observes [`event::CallEvent`]s and
//! may rewrite the wire request before it is sent.
pub mod event;
pub mod host;
pub mod machine;
pub mod route;
pub mod protocol;
pub mod run;

View file

@ -1,22 +1,21 @@
use std::sync::Arc;
use super::{HostChannel, MachineFault};
use crate::route::Route;
use crate::{host::Reply, protocol::Protocol};
use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};
/// A route whose host can mint credentials on the call's behalf.
pub trait TokenRoute: Route {
fn acquire_token_op() -> Self::Op;
fn token_credential(result: Self::OpResult) -> Option<ResolvedCredential>;
/// A protocol whose host can mint credentials on the call's behalf.
pub trait TokenProtocol: Protocol {
fn acquire_token_op(reply: Reply<ResolvedCredential>) -> Self::Op;
}
/// A [`TokenProvider`] that asks the host for each credential through the call's own
/// operation channel, so the host answers it on the caller's thread and context.
pub struct HostTokenProvider<R: Route> {
pub struct HostTokenProvider<R: Protocol> {
channel: HostChannel<R>,
}
impl<R: Route> std::fmt::Debug for HostTokenProvider<R> {
impl<R: Protocol> std::fmt::Debug for HostTokenProvider<R> {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("HostTokenProvider")
}
@ -24,7 +23,7 @@ impl<R: Route> std::fmt::Debug for HostTokenProvider<R> {
impl<R> HostTokenProvider<R>
where
R: TokenRoute,
R: TokenProtocol,
R::Error: From<MachineFault> + std::fmt::Display,
{
pub fn handle(channel: HostChannel<R>) -> TokenProviderHandle {
@ -34,19 +33,15 @@ where
impl<R> TokenProvider for HostTokenProvider<R>
where
R: TokenRoute,
R: TokenProtocol,
R::Error: From<MachineFault> + std::fmt::Display,
{
fn acquire(&self) -> TokenFuture<'_> {
Box::pin(async move {
let result = self
.channel
.route(R::acquire_token_op())
self.channel
.custom_op(R::acquire_token_op)
.await
.map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?;
R::token_credential(result).ok_or_else(|| {
Error::AzureTokenAcquisition("invalid token provider host result".into())
})
.map_err(|error| Error::AzureTokenAcquisition(error.to_string()))
})
}
}

View file

@ -0,0 +1,137 @@
//! The one machine every route runs on: the route's provider future as a
//! [`Coroutine`] that yields [`HostOp`]s, each answered through its own typed reply. No
//! task is spawned; dropping the machine drops the in-flight call.
use std::{future::Future, pin::Pin};
use litellm_coroutine::{Co, Coroutine, CoroutineState, ResumeError};
use super::{HostFailure, Interrupted, Machine, MachineStep, Step};
use crate::{
event::{MachineEvent, RequestContext, WireRequest},
host::{Demand, HostOp, Reply},
protocol::Protocol,
};
/// The machine's own failures, distinct from anything the provider call reports.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MachineFault {
/// The host dropped an op's reply unanswered, or went away while the call waited.
Abandoned,
/// The host resumed the call out of turn.
Protocol(ResumeError),
}
pub type ExecuteFuture<R> =
Pin<Box<dyn Future<Output = Result<<R as Protocol>::Response, <R as Protocol>::Error>> + Send>>;
/// The provider side of the machine: how the in-flight call reaches its host.
pub struct HostChannel<R: Protocol> {
co: Co<HostOp<R>>,
}
impl<R: Protocol> Clone for HostChannel<R> {
fn clone(&self) -> Self {
Self {
co: self.co.clone(),
}
}
}
impl<R: Protocol> HostChannel<R>
where
R::Error: From<MachineFault>,
{
async fn yield_<A: Send>(
&self,
ask: impl FnOnce(Reply<A>) -> HostOp<R> + Send,
) -> Result<A, R::Error> {
self.co
.yield_(ask)
.await
.map_err(|_| MachineFault::Abandoned.into())
}
pub async fn project(&self) -> Result<R::Projection, R::Error> {
self.yield_(HostOp::Project).await
}
/// Asks the host to perform the custom operation `ask` builds around its reply, as in
/// `host.custom_op(OcrOp::AcquireAzureAdToken)`.
pub async fn custom_op<A: Send>(
&self,
ask: impl FnOnce(Reply<A>) -> R::Op + Send,
) -> Result<A, R::Error> {
self.yield_(|reply| HostOp::Custom(ask(reply))).await
}
pub async fn before_send(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, R::Error> {
self.yield_(|reply| HostOp::BeforeSend {
wire: Box::new(wire),
context: Box::new(context),
reply,
})
.await
}
pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> {
self.yield_(|reply| HostOp::Emit(event, reply)).await
}
pub async fn open(&self, head: R::StreamHead) -> Result<Demand, R::Error> {
self.yield_(|reply| HostOp::Open(head, reply)).await
}
pub async fn deliver(&self, chunk: R::Chunk) -> Result<Demand, R::Error> {
self.yield_(|reply| HostOp::Deliver(chunk, reply)).await
}
}
type CallCoroutine<R> =
Coroutine<HostOp<R>, Result<<R as Protocol>::Response, <R as Protocol>::Error>>;
pub struct CallMachine<R: Protocol> {
coroutine: CallCoroutine<R>,
}
impl<R: Protocol> CallMachine<R>
where
R::Error: From<MachineFault>,
{
pub fn new(execute: impl FnOnce(HostChannel<R>) -> ExecuteFuture<R> + Send + 'static) -> Self {
Self {
coroutine: Coroutine::new(|co| execute(HostChannel { co })),
}
}
}
impl<R: Protocol> Machine for CallMachine<R>
where
R::Error: From<MachineFault>,
{
type Protocol = R;
type Complete = R::Response;
fn resume(&mut self) -> Step<'_, Self> {
Box::pin(async move {
match self
.coroutine
.resume()
.await
.map_err(MachineFault::Protocol)?
{
CoroutineState::Yielded(op) => Ok(MachineStep::Host(op)),
CoroutineState::Complete(outcome) => outcome.map(MachineStep::Complete),
}
})
}
fn interrupt(&mut self, failure: HostFailure<R::Error>) -> Interrupted<'_, Self> {
self.coroutine.cancel();
Box::pin(async move { Err(failure.into_error()) })
}
}

View file

@ -1,16 +1,16 @@
mod auth;
mod route_machine;
mod call_machine;
use std::future::Future;
use std::pin::Pin;
pub use auth::{HostTokenProvider, TokenRoute};
pub use route_machine::{ExecuteFuture, HostChannel, MachineFault, RouteMachine};
pub use auth::{HostTokenProvider, TokenProtocol};
pub use call_machine::{CallMachine, ExecuteFuture, HostChannel, MachineFault};
use crate::host::{HostOp, HostResult};
use crate::route::Route;
use crate::host::HostOp;
use crate::protocol::Protocol;
pub enum MachineStep<R: Route, C> {
pub enum MachineStep<R: Protocol, C> {
Host(HostOp<R>),
Complete(C),
}
@ -19,8 +19,8 @@ pub type Step<'a, M> = Pin<
Box<
dyn Future<
Output = Result<
MachineStep<<M as Machine>::Route, <M as Machine>::Complete>,
<<M as Machine>::Route as Route>::Error,
MachineStep<<M as Machine>::Protocol, <M as Machine>::Complete>,
<<M as Machine>::Protocol as Protocol>::Error,
>,
> + Send
+ 'a,
@ -30,7 +30,10 @@ pub type Step<'a, M> = Pin<
pub type Interrupted<'a, M> = Pin<
Box<
dyn Future<
Output = Result<<M as Machine>::Complete, <<M as Machine>::Route as Route>::Error>,
Output = Result<
<M as Machine>::Complete,
<<M as Machine>::Protocol as Protocol>::Error,
>,
> + Send
+ 'a,
>,
@ -51,19 +54,18 @@ impl<E> HostFailure<E> {
}
/// A resumable call. Core implements it per route; a host drives it. Every suspension
/// point is an op the host performs and answers with a result.
/// point is an op the host performs and answers through the op's own reply before it
/// resumes the call again.
pub trait Machine: Send {
type Route: Route;
type Protocol: Protocol;
type Complete: Send + 'static;
/// `None` on the first call and whenever the previous step completed without
/// yielding an op; otherwise the result of the op last yielded.
fn resume(&mut self, result: Option<HostResult<Self::Route>>) -> Step<'_, Self>;
fn resume(&mut self) -> Step<'_, Self>;
/// The host failed to perform the pending op, or the caller cancelled. The call
/// yields no further ops.
fn interrupt(
&mut self,
failure: HostFailure<<Self::Route as Route>::Error>,
failure: HostFailure<<Self::Protocol as Protocol>::Error>,
) -> Interrupted<'_, Self>;
}

View file

@ -1,199 +0,0 @@
//! The one machine every route runs on: it owns the route's provider future, polls it in
//! place, and turns the host operations that future requests into [`Machine`] steps. No
//! task is spawned; dropping the machine drops the in-flight call.
use std::{future::Future, pin::Pin};
use tokio::sync::{mpsc, oneshot};
use super::{HostFailure, Interrupted, Machine, MachineStep, Step};
use crate::{
event::{MachineEvent, RequestContext, WireRequest},
host::{Demand, HostOp, HostResult},
route::Route,
};
/// The machine's own failures, distinct from anything the provider call reports.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MachineFault {
/// The host driver went away while the call was waiting on it.
Abandoned,
/// The host answered out of turn: a result with nothing pending, or nothing when a
/// result was pending.
Protocol(&'static str),
/// The host answered a route operation with the wrong result variant.
Mismatch,
}
pub type ExecuteFuture<R> =
Pin<Box<dyn Future<Output = Result<<R as Route>::Response, <R as Route>::Error>> + Send>>;
struct PendingOp<R: Route> {
op: HostOp<R>,
reply: oneshot::Sender<HostResult<R>>,
}
/// The provider side of the machine: how the in-flight call reaches its host.
pub struct HostChannel<R: Route> {
ops: mpsc::UnboundedSender<PendingOp<R>>,
}
impl<R: Route> Clone for HostChannel<R> {
fn clone(&self) -> Self {
Self {
ops: self.ops.clone(),
}
}
}
impl<R: Route> HostChannel<R>
where
R::Error: From<MachineFault>,
{
async fn invoke(&self, op: HostOp<R>) -> Result<HostResult<R>, R::Error> {
let (reply, answer) = oneshot::channel();
self.ops
.send(PendingOp { op, reply })
.map_err(|_| MachineFault::Abandoned)?;
answer.await.map_err(|_| MachineFault::Abandoned.into())
}
pub async fn route(&self, op: R::Op) -> Result<R::OpResult, R::Error> {
match self.invoke(HostOp::Route(op)).await? {
HostResult::Route(result) => Ok(result),
_ => Err(MachineFault::Mismatch.into()),
}
}
pub async fn before_send(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, R::Error> {
let op = HostOp::BeforeSend {
wire: Box::new(wire),
context: Box::new(context),
};
match self.invoke(op).await? {
HostResult::BeforeSend(wire) => Ok(*wire),
_ => Err(MachineFault::Mismatch.into()),
}
}
pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> {
match self.invoke(HostOp::Emit(event)).await? {
HostResult::Emitted => Ok(()),
_ => Err(MachineFault::Mismatch.into()),
}
}
pub async fn open(&self, head: R::StreamHead) -> Result<Demand, R::Error> {
self.demand(HostOp::Open(head)).await
}
pub async fn deliver(&self, chunk: R::Chunk) -> Result<Demand, R::Error> {
self.demand(HostOp::Deliver(chunk)).await
}
async fn demand(&self, op: HostOp<R>) -> Result<Demand, R::Error> {
match self.invoke(op).await? {
HostResult::Demand(demand) => Ok(demand),
_ => Err(MachineFault::Mismatch.into()),
}
}
}
enum Execution<R: Route> {
Unstarted(Box<dyn FnOnce(HostChannel<R>) -> ExecuteFuture<R> + Send>),
Running(ExecuteFuture<R>),
Done,
}
pub struct RouteMachine<R: Route> {
execution: Execution<R>,
ops: mpsc::UnboundedReceiver<PendingOp<R>>,
channel: HostChannel<R>,
reply: Option<oneshot::Sender<HostResult<R>>>,
}
impl<R: Route> RouteMachine<R>
where
R::Error: From<MachineFault>,
{
pub fn new(execute: impl FnOnce(HostChannel<R>) -> ExecuteFuture<R> + Send + 'static) -> Self {
let (ops_tx, ops) = mpsc::unbounded_channel();
Self {
execution: Execution::Unstarted(Box::new(execute)),
ops,
channel: HostChannel { ops: ops_tx },
reply: None,
}
}
async fn step(
&mut self,
result: Option<HostResult<R>>,
) -> Result<MachineStep<R, R::Response>, R::Error> {
match (self.reply.take(), result) {
(Some(reply), Some(result)) => {
reply
.send(result)
.map_err(|_| MachineFault::Protocol("the call stopped waiting on the host"))?;
}
(None, None) if matches!(self.execution, Execution::Unstarted(_)) => {}
(Some(reply), None) => {
self.reply = Some(reply);
return Err(MachineFault::Protocol("host operation result is required").into());
}
(None, Some(_)) => {
return Err(MachineFault::Protocol("unexpected host operation result").into());
}
(None, None) => {
return Err(
MachineFault::Protocol("call cannot be resumed after completion").into(),
);
}
}
if let Execution::Unstarted(_) = self.execution {
let Execution::Unstarted(start) =
std::mem::replace(&mut self.execution, Execution::Done)
else {
unreachable!()
};
self.execution = Execution::Running(start(self.channel.clone()));
}
let Execution::Running(future) = &mut self.execution else {
return Err(MachineFault::Protocol("call cannot be resumed after completion").into());
};
tokio::select! {
biased;
pending = self.ops.recv() => {
let pending = pending.ok_or(MachineFault::Abandoned)?;
self.reply = Some(pending.reply);
Ok(MachineStep::Host(pending.op))
}
outcome = future => {
self.execution = Execution::Done;
outcome.map(MachineStep::Complete)
}
}
}
}
impl<R: Route> Machine for RouteMachine<R>
where
R::Error: From<MachineFault>,
{
type Route = R;
type Complete = R::Response;
fn resume(&mut self, result: Option<HostResult<R>>) -> Step<'_, Self> {
Box::pin(self.step(result))
}
fn interrupt(&mut self, failure: HostFailure<R::Error>) -> Interrupted<'_, Self> {
self.reply = None;
self.execution = Execution::Done;
Box::pin(async move { Err(failure.into_error()) })
}
}

View file

@ -0,0 +1,17 @@
/// One public call surface: what a completed call produces, how it fails, what the host
/// projects the caller's request into, and the protocol-specific operations only its host
/// can perform mid-call (token acquisition, for one).
pub trait Protocol: Send + Sync + 'static {
type Response: Send + 'static;
type Error: Clone + Send + Sync + 'static;
/// The caller's request as the host projects it, answered once before anything else.
type Projection: Send + 'static;
/// Each operation carries the [`Reply`](crate::host::Reply) its answer goes through.
/// A protocol with no operations of its own uses `Infallible`.
type Op: Send + 'static;
/// One piece of a streamed response, handed to the caller as it arrives. A protocol
/// that never streams uses `Infallible`.
type Chunk: Send + 'static;
/// What the call knows once a streamed response starts, before its first chunk.
type StreamHead: Send + 'static;
}

View file

@ -1,14 +0,0 @@
/// One public call surface: what a completed call produces, how it fails, and the
/// route-specific operations only its host can perform (request projection, file reads,
/// token acquisition).
pub trait Route: Send + Sync + 'static {
type Response: Send + 'static;
type Error: Clone + Send + Sync + 'static;
type Op: Send + 'static;
type OpResult: Send + 'static;
/// One piece of a streamed response, handed to the caller as it arrives. A route
/// that never streams uses `Infallible`.
type Chunk: Send + 'static;
/// What the route knows once a streamed response starts, before its first chunk.
type StreamHead: Send + 'static;
}

View file

@ -1,40 +1,28 @@
use crate::event::{CallEvent, FailureOrigin, Timing, epoch_seconds};
use crate::host::{Host, HostOp, HostResult};
use crate::host::{Host, HostOp};
use crate::machine::{HostFailure, Machine, MachineStep};
use crate::route::Route;
use crate::protocol::Protocol;
/// Drives a machine to completion against an in-process host and emits exactly one
/// terminal event.
pub async fn run<M, H>(mut machine: M, host: &H) -> Result<M::Complete, <M::Route as Route>::Error>
pub async fn run<M, H>(
mut machine: M,
host: &H,
) -> Result<M::Complete, <M::Protocol as Protocol>::Error>
where
M: Machine,
H: Host<M::Route>,
H: Host<M::Protocol>,
{
let start_time = epoch_seconds();
let _ = host.emit(&CallEvent::Started { start_time }).await;
let mut result = None;
let outcome = loop {
let step = match machine.resume(result.take()).await {
let op = match machine.resume().await {
Ok(MachineStep::Complete(complete)) => break Ok(complete),
Ok(MachineStep::Host(op)) => op,
Err(error) => break Err(error),
};
let answer = match step {
HostOp::Route(op) => host.route(op).await.map(HostResult::Route),
HostOp::BeforeSend { wire, context } => host
.before_send(*wire, &context)
.await
.map(|wire| HostResult::BeforeSend(Box::new(wire))),
HostOp::Emit(event) => host
.emit(&CallEvent::Machine(event))
.await
.map(|()| HostResult::Emitted),
HostOp::Open(head) => host.open(head).await.map(HostResult::Demand),
HostOp::Deliver(chunk) => host.deliver(chunk).await.map(HostResult::Demand),
};
match answer {
Ok(answer) => result = Some(answer),
Err(error) => break machine.interrupt(HostFailure::Error(error)).await,
if let Err(error) = perform(host, op).await {
break machine.interrupt(HostFailure::Error(error)).await;
}
};
let timing = Timing {
@ -52,44 +40,52 @@ where
outcome
}
async fn perform<R: Protocol, H: Host<R>>(host: &H, op: HostOp<R>) -> Result<(), R::Error> {
match op {
HostOp::Project(reply) => host
.project()
.await
.map(|projection| reply.send(projection)),
HostOp::Custom(op) => host.custom_op(op).await,
HostOp::BeforeSend {
wire,
context,
reply,
} => host
.before_send(*wire, &context)
.await
.map(|wire| reply.send(wire)),
HostOp::Emit(event, reply) => host
.emit(&CallEvent::Machine(event))
.await
.map(|()| reply.send(())),
HostOp::Open(head, reply) => host.open(head).await.map(|demand| reply.send(demand)),
HostOp::Deliver(chunk, reply) => host.deliver(chunk).await.map(|demand| reply.send(demand)),
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use super::*;
use crate::machine::{Interrupted, Step};
use crate::host::Reply;
use crate::machine::{CallMachine, MachineFault};
struct Unit;
impl Route for Unit {
impl Protocol for Unit {
type Response = ();
type Error = &'static str;
type Op = &'static str;
type OpResult = ();
type Projection = ();
type Op = (&'static str, Reply<()>);
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
struct Scripted {
ops: Vec<&'static str>,
outcome: Result<(), &'static str>,
}
impl Machine for Scripted {
type Route = Unit;
type Complete = ();
fn resume(&mut self, _: Option<HostResult<Unit>>) -> Step<'_, Self> {
Box::pin(async move {
if !self.ops.is_empty() {
return Ok(MachineStep::Host(HostOp::Route(self.ops.remove(0))));
}
self.outcome.map(MachineStep::Complete)
})
}
fn interrupt(&mut self, failure: HostFailure<&'static str>) -> Interrupted<'_, Self> {
Box::pin(async move { Err(failure.into_error()) })
impl From<MachineFault> for &'static str {
fn from(_: MachineFault) -> Self {
"machine fault"
}
}
@ -100,12 +96,21 @@ mod tests {
}
impl Host<Unit> for Recording {
async fn route(&self, op: &'static str) -> Result<(), &'static str> {
self.seen.lock().unwrap().push(format!("route:{op}"));
match self.fail {
Some(failing) if failing == op => Err("host failed"),
_ => Ok(()),
async fn project(&self) -> Result<(), &'static str> {
self.seen.lock().unwrap().push("project".into());
Ok(())
}
async fn custom_op(
&self,
(op, reply): (&'static str, Reply<()>),
) -> Result<(), &'static str> {
self.seen.lock().unwrap().push(format!("op:{op}"));
if self.fail == Some(op) {
return Err("host failed");
}
reply.send(());
Ok(())
}
async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> {
@ -119,21 +124,29 @@ mod tests {
}
}
fn scripted(ops: &[&'static str], outcome: Result<(), &'static str>) -> Scripted {
Scripted {
ops: ops.to_vec(),
outcome,
}
fn scripted(
ops: &'static [&'static str],
outcome: Result<(), &'static str>,
) -> CallMachine<Unit> {
CallMachine::new(move |host| {
Box::pin(async move {
host.project().await?;
for op in ops {
host.custom_op(|reply| (*op, reply)).await?;
}
outcome
})
})
}
#[tokio::test]
async fn forwards_every_op_then_emits_one_succeeded() {
let host = Recording::default();
let outcome = run(scripted(&["project", "send"], Ok(())), &host).await;
let outcome = run(scripted(&["sign", "send"], Ok(())), &host).await;
assert_eq!(outcome, Ok(()));
assert_eq!(
*host.seen.lock().unwrap(),
["started", "route:project", "route:send", "succeeded"]
["started", "project", "op:sign", "op:send", "succeeded"]
);
}
@ -142,24 +155,32 @@ mod tests {
let host = Recording::default();
let outcome = run(scripted(&[], Err("boom")), &host).await;
assert_eq!(outcome, Err("boom"));
assert_eq!(*host.seen.lock().unwrap(), ["started", "failed"]);
assert_eq!(*host.seen.lock().unwrap(), ["started", "project", "failed"]);
let host = Recording {
fail: Some("send"),
..Recording::default()
};
let outcome = run(scripted(&["project", "send", "never"], Ok(())), &host).await;
let outcome = run(scripted(&["sign", "send", "never"], Ok(())), &host).await;
assert_eq!(outcome, Err("host failed"));
assert_eq!(
*host.seen.lock().unwrap(),
["started", "route:project", "route:send", "failed"]
["started", "project", "op:sign", "op:send", "failed"]
);
}
struct StartTimes(Mutex<Vec<f64>>);
impl Host<Unit> for StartTimes {
async fn route(&self, _: &'static str) -> Result<(), &'static str> {
async fn project(&self) -> Result<(), &'static str> {
Ok(())
}
async fn custom_op(
&self,
(_, reply): (&'static str, Reply<()>),
) -> Result<(), &'static str> {
reply.send(());
Ok(())
}
@ -178,7 +199,7 @@ mod tests {
#[tokio::test]
async fn started_opens_the_call_at_the_terminal_start_time_and_cannot_fail_it() {
let host = StartTimes(Mutex::default());
assert_eq!(run(scripted(&["project"], Ok(())), &host).await, Ok(()));
assert_eq!(run(scripted(&["send"], Ok(())), &host).await, Ok(()));
let times = host.0.lock().unwrap();
assert_eq!(times.len(), 2);
assert_eq!(times[0], times[1]);

View file

@ -16,6 +16,8 @@ use crate::{
},
};
pub const AZURE_COHERE_PARSE_PATH: [&str; 4] = ["providers", "cohere", "v2", "parse"];
#[derive(Default)]
pub struct AzureAICohereParseConfig;
@ -131,7 +133,7 @@ impl AzureAICohereParseConfig {
}
url.set_path(path.strip_suffix("/models").unwrap_or(&path));
ApiUrl::parse(url.as_str())
.and_then(|url| url.complete_path(&["providers", "cohere", "v2", "parse"]))
.and_then(|url| url.complete_path(&AZURE_COHERE_PARSE_PATH))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base())
}

View file

@ -561,6 +561,14 @@ async fn poll_operation(
}
impl AzureDocumentIntelligenceOcrConfig {
pub fn analyze_path(model: &str) -> Result<[String; 3], Error> {
Ok([
"documentintelligence".into(),
"documentModels".into(),
format!("{}:analyze", model_id(model)?),
])
}
fn build_ocr_url(
&self,
endpoint: &str,
@ -568,9 +576,9 @@ impl AzureDocumentIntelligenceOcrConfig {
params: &DocumentIntelligenceParams,
api_version: &str,
) -> Result<String, Error> {
let model = format!("{}:analyze", model_id(model)?);
let path = Self::analyze_path(model)?;
ApiUrl::parse(endpoint)
.and_then(|url| url.complete_path(&["documentintelligence", "documentModels", &model]))
.and_then(|url| url.complete_path(&path.each_ref().map(String::as_str)))
.map(|url| {
url.append_query_pairs(
[("api-version", api_version)]

View file

@ -17,7 +17,7 @@ use crate::{
mistral::ocr::transformation::{MistralOcrConfig, MistralOcrRequest},
};
const AZURE_AI_OCR_PATH: &str = "/providers/mistral/azure/ocr";
pub const AZURE_AI_OCR_PATH: [&str; 4] = ["providers", "mistral", "azure", "ocr"];
const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY";
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
@ -179,9 +179,8 @@ impl AzureAiOcrConfig {
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
let base = Self::resolve_api_base(api_base, env_lookup)?;
let path: Vec<&str> = AZURE_AI_OCR_PATH.trim_matches('/').split('/').collect();
ApiUrl::parse(&base)
.and_then(|url| url.complete_path(&path))
.and_then(|url| url.complete_path(&AZURE_AI_OCR_PATH))
.map(|url| url.into_string())
.map_err(|_| Error::RequestField {
path: "api_base".into(),

View file

@ -118,7 +118,6 @@ impl From<litellm_host::machine::MachineFault> for Error {
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "OCR host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("OCR {message}"),
MachineFault::Mismatch => "invalid OCR host operation result".into(),
})
}
}

View file

@ -554,6 +554,50 @@ async fn upload_bytes_async(
mod tests {
use super::*;
#[tokio::test]
async fn v3_body_keeps_explicit_null_options_and_drops_unknown_ones() {
use crate::base_llm::ocr::{handler::OcrClient, transformation::OcrRequestContext};
let overrides =
serde_json::from_value(json!({"formatting":null,"settings":{},"unknown":true}))
.unwrap();
let params = ReductoParseV3Config
.map_ocr_params(&overrides, "parse-v3")
.unwrap();
let client = OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new());
let connection = OcrConnection::default();
let document = serde_json::from_value(
json!({"type":"document_url","document_url":"reducto://ready.pdf"}),
)
.unwrap();
let body = ReductoParseV3Config
.async_transform_ocr_request(
"parse-v3",
document,
&params,
&[],
OcrRequestContext {
client: &client,
connection: &connection,
},
)
.await
.unwrap();
assert_eq!(
serde_json::to_value(body).unwrap(),
json!({"input":"reducto://ready.pdf", "formatting":null, "settings":{}})
);
let absent = ReductoParseV3Config
.map_ocr_params(
&litellm_core_utils::call_arguments::CallArguments::default(),
"parse-v3",
)
.unwrap();
assert_eq!(serde_json::to_value(absent).unwrap(), json!({}));
}
#[test]
fn options_preserve_null_and_select_the_provider_fields() {
let overrides = serde_json::from_value(json!({

View file

@ -0,0 +1,79 @@
use std::time::Duration;
use litellm_llms::base_llm::ocr::{error::Error, handler::read_response_bytes};
use rstest::rstest;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
};
/// Answers one request with raw `response` bytes and then holds the connection open, so a
/// read that waits for the rest of an oversized body hangs instead of passing.
async fn read_bounded(response: String, limit: usize) -> Result<bytes::Bytes, Error> {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0; 4096];
assert!(socket.read(&mut request).await.unwrap() > 0);
socket.write_all(response.as_bytes()).await.unwrap();
std::future::pending::<()>().await;
});
let response = reqwest::Client::new()
.get(format!("http://{address}"))
.send()
.await
.unwrap();
let result =
tokio::time::timeout(Duration::from_secs(2), read_response_bytes(response, limit)).await;
server.abort();
result.expect("bounded reads must finish without waiting for the rest of an oversized body")
}
#[rstest]
#[case::declared("HTTP/1.1 200 OK\r\nContent-Length: 8\r\n\r\nabcdefgh")]
#[case::chunked(
"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n4\r\nefgh\r\n0\r\n\r\n"
)]
#[tokio::test]
async fn a_body_of_exactly_the_limit_is_read(#[case] response: &str) {
assert_eq!(read_bounded(response.into(), 8).await.unwrap(), "abcdefgh");
}
#[rstest]
#[case::declared("HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\n")]
#[case::chunked("HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n5\r\nefghi\r\n")]
#[tokio::test]
async fn a_body_over_the_limit_is_rejected(#[case] response: &str) {
assert!(matches!(
read_bounded(response.into(), 8).await,
Err(Error::TooLarge { limit: 8 })
));
}
#[rstest]
#[case::declared("Content-Length: 1000000")]
#[case::chunked("Transfer-Encoding: chunked")]
#[tokio::test]
async fn an_oversized_error_keeps_its_status_and_a_bounded_body_without_draining(
#[case] headers: &str,
) {
let prefix = "x".repeat(4096);
let body = match headers.starts_with("Transfer") {
true => format!("{:x}\r\n{prefix}\r\n", prefix.len()),
false => prefix.clone(),
};
let error = read_bounded(
format!("HTTP/1.1 429 Too Many Requests\r\n{headers}\r\n\r\n{body}"),
prefix.len(),
)
.await
.unwrap_err();
let Error::Transport(litellm_http::transport::Error::Http { status, body }) = error else {
panic!("unexpected error: {error}");
};
assert_eq!(status, 429);
assert_eq!(body, prefix);
}

View file

@ -0,0 +1,6 @@
## Validation
For `model_prices_and_context_window.json` validation, we should eventually:
- Remove any schema file like `model_prices_and_context_window.schema.json`
- Stop skipping this crate's tests

View file

@ -14,12 +14,8 @@ schemars = { version = "1.0", optional = true }
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
time.workspace = true
[dev-dependencies]
criterion.workspace = true
jsonschema = { version = "0.55.1", default-features = false }
rstest.workspace = true
litellm-model-catalog = { path = ".", features = ["schema"] }
[[bench]]
name = "catalog"
harness = false

View file

@ -1,25 +0,0 @@
# Model catalog
`litellm-model-catalog` builds an immutable snapshot from caller supplied JSON bytes. It has no network, Python, registration, or refresh behavior. The caller supplies optional source, revision, and ETag provenance. Parse and validation are separate so small synthetic catalogs can use explicit integrity limits
The parser treats `sample_spec` and `fallback_generalizations` as reserved top level metadata. `fallback_rules()` exposes the typed rule array when present; this crate does not execute regex generalizations. Model entries retain all JSON fields except `aliases`, including unknown fields. `field()` returns `None` for an absent key and a JSON null, false, or zero value for a present key. The returned values are borrowed, so callers cannot mutate the snapshot
Each entry also deserializes into `ModelInfo`, a typed mirror of `model_prices_and_context_window.schema.json`'s `modelEntry` definition, reachable via `ModelEntry::info()`. All schema fields are optional on `ModelInfo`, including `litellm_provider` which the schema marks required, so small synthetic catalogs still parse. Unknown fields are not part of `ModelInfo`; they remain on `fields()`. Building with the `schema` feature adds `schemars` derives and exposes `model_entry_json_schema()` for emitting the entry's JSON Schema. Parse and validation failures are reported by the `Error` enum in `error.rs`, while catalog logic lives in `catalog.rs`
The integration tests read the repository's catalog and schema files at test time, assert every entry round-trips through `ModelInfo`, and verify that the generated schema's properties match the repository schema
Aliases point to their canonical entries. An alias that exactly matches any canonical key is skipped; the first canonical entry claiming an alias wins. Invalid alias lists and nonstring names are skipped and reported by `alias_issues()`. Exact lookup wins. For a case insensitive miss, the last key with the same lowercase spelling wins, following Python's lowercase map built after aliases are appended. This uses Rust Unicode lowercasing, which can differ from Python for unusual Unicode model IDs
`validate()` counts canonical entries before alias expansion and excludes both reserved keys. It enforces an explicit minimum and backup shrink ratio, with Python defaults of 50 models and 0.5. Parsing rejects nonobject model entries and known fields with the wrong JSON type, but ignores unknown fields. It does not enforce every constraint in the JSON schema, calculate prices, resolve providers, or check provenance authenticity. The caller decides how to handle validation failures
This snapshot does not represent Python's live mutable `litellm.model_cost`, nested dict and list mutation, or mutation of dicts previously returned by Python APIs. It has no bridge or runtime integration
## Benchmarks
`cargo bench -p litellm-model-catalog --bench catalog` measures parsing plus alias indexing and exact lookup. For a local Python baseline on the same fixture, use:
```sh
python3 -m timeit -s 'import json, pathlib; body = pathlib.Path("../model_prices_and_context_window.json").read_bytes()' 'json.loads(body)'
```
Run these commands from `litellm-rust`. Python's command measures JSON loading only, without alias expansion or snapshot construction. The Rust benchmark does not include future Python object materialization, so these numbers are not an end to end runtime comparison

View file

@ -1,21 +0,0 @@
use criterion::{Criterion, criterion_group, criterion_main};
use litellm_model_catalog::{Catalog, Provenance};
use std::hint::black_box;
fn benchmarks(c: &mut Criterion) {
let body = include_bytes!("../../../../model_prices_and_context_window.json");
c.bench_function("parse_current_catalog", |b| {
b.iter(|| Catalog::parse(black_box(body), Provenance::default()).unwrap())
});
let catalog = Catalog::parse(body, Provenance::default()).unwrap();
let key = catalog
.model_names()
.next()
.expect("catalog must have a benchmark key");
c.bench_function("lookup_catalog_key", |b| {
b.iter(|| black_box(&catalog).lookup(black_box(key)))
});
}
criterion_group!(benches, benchmarks);
criterion_main!(benches);

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