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

This commit is contained in:
yassin 2026-09-16 16:02:08 +00:00
commit 04aad317c7
1995 changed files with 108672 additions and 28223 deletions

View file

@ -9,7 +9,7 @@ commands:
parameters:
category:
type: enum
enum: ["backend", "client"]
enum: ["backend", "client", "provider-harness"]
default: "backend"
steps:
- run:
@ -147,6 +147,9 @@ commands:
db_name:
type: string
default: circle_test
image:
type: string
default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
steps:
- run:
name: Start PostgreSQL
@ -157,7 +160,7 @@ commands:
-e POSTGRES_PASSWORD=postgres \
-e POSTGRES_DB=<< parameters.db_name >> \
-p 5432:5432 \
postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
<< parameters.image >>
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
@ -1084,9 +1087,7 @@ jobs:
name: Run tests
command: |
mkdir -p test-results
TEST_FILES=$(printf "%s\n%s\n" \
"$(circleci tests glob "tests/ocr_tests/**/test_*.py")" \
"tests/test_litellm/ocr/test_rust_bridge.py")
TEST_FILES=$(circleci tests glob "tests/ocr_tests/**/test_*.py")
echo "$TEST_FILES" | circleci tests run \
--verbose \
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
@ -2914,7 +2915,80 @@ jobs:
exit 1
fi
provider_replay_harness:
docker:
- *python312_image
- image: redis@sha256:e2debfb7956fa12c7ddc79d7e645c8cf26b30c99a6e9161ea9bf4171e1668a5f
working_directory: ~/project
resource_class: medium
environment:
E2E_CACHE_TEST_REDIS_URL: redis://127.0.0.1:6379/0
E2E_PROVIDER_CACHE: "0"
E2E_FIXTURE_MODE: live
steps:
- checkout
- skip_if_unrelated_changes:
category: provider-harness
- setup_litellm_test_deps
- wait_for_service:
url: tcp://localhost:6379
- run:
name: Test provider capture and replay harness
command: |
mkdir -p test-results/provider-replay-harness
uv run --no-sync pytest -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures \
--junitxml=test-results/provider-replay-harness/junit.xml \
tests/e2e/test_provider_edge.py tests/e2e/test_fixture_bundle.py \
tests/e2e/test_fixture_canonical.py tests/e2e/test_fixture_mode.py \
tests/code_coverage_tests/test_provider_replay_harness.py \
tests/code_coverage_tests/test_provider_cache.py
- store_test_results:
path: test-results/provider-replay-harness
integration_contracts:
parameters:
suite:
type: string
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_litellm_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 >>
no_output_timeout: 15m
- run:
name: Stop owned database and Redis
when: always
command: |
mkdir -p test-results/integration-<< parameters.suite >>
docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true
docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
workflows:
integration:
jobs:
- integration_contracts:
name: integration-<< matrix.suite >>
matrix:
parameters:
suite: [management, accounting, providers]
filters:
branches:
only:
- main
- /litellm_.*/
build_and_test:
jobs:
- using_litellm_on_windows:
@ -2923,6 +2997,7 @@ workflows:
only:
- main
- /litellm_.*/
- provider_replay_harness
- base_sdk_install:
filters: *main_branches
- local_testing_part1:

View file

@ -1,13 +1,19 @@
#!/usr/bin/env bash
set -uo pipefail
category="${1:?usage: classify_changes.sh <backend|client|ui>}"
category="${1:?usage: classify_changes.sh <backend|client|ui|provider-harness>}"
has_client=false
has_backend=false
has_ci=false
has_provider_harness=false
while IFS= read -r file || [ -n "$file" ]; do
[ -n "$file" ] || continue
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)
has_provider_harness=true ;;
esac
case "$file" in
ui/* | tests/e2e/ui/*) has_client=true ;;
docs/* | *.md | *.mdx) : ;;
@ -17,6 +23,9 @@ while IFS= read -r file || [ -n "$file" ]; do
done
case "$category" in
provider-harness)
[ "$has_provider_harness" = true ] && echo run || echo skip
;;
backend)
[ "$has_backend" = true ] && echo run || echo skip
;;

View file

@ -1,7 +1,7 @@
#!/usr/bin/env bash
set -uo pipefail
category="${1:?usage: path_filter.sh <backend|client>}"
category="${1:?usage: path_filter.sh <backend|client|provider-harness>}"
here="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
run_full() {
@ -36,5 +36,5 @@ if [ "$decision" = run ]; then
run_full "$category-relevant changes detected"
fi
echo "path-filter[$category]: only unrelated (docs/client) changes detected; halting job as successful"
echo "path-filter[$category]: only unrelated changes detected; halting job as successful"
circleci-agent step halt

View file

@ -0,0 +1,142 @@
#!/usr/bin/env bash
set -euo pipefail
suite="${1:?integration suite required}"
results="test-results/integration-${suite}"
mkdir -p "$results"
integration_identity="$(.venv/bin/python -c 'import uuid; print(uuid.uuid4().hex)')"
upstream_pid=""
proxy_pid=""
peer_pid=""
launched_pid=""
guard_created=false
guard_installed=false
guard6_created=false
guard6_installed=false
cleanup() {
original_status=$?
trap - EXIT INT TERM
sudo .venv/bin/python .circleci/scripts/stop_integration_processes.py \
"$integration_identity" "$(id -u)" "$proxy_pid" "$peer_pid" "$upstream_pid" \
> "$results/process-cleanup.txt" 2>&1 || original_status=1
for owned_pid in "$peer_pid" "$proxy_pid" "$upstream_pid"; do
if [ -n "$owned_pid" ]; then
kill -- "-$owned_pid" 2>/dev/null || true
for _ in {1..50}; do
kill -0 -- "-$owned_pid" 2>/dev/null || break
sleep 0.1
done
if kill -0 -- "-$owned_pid" 2>/dev/null; then
kill -KILL -- "-$owned_pid" 2>/dev/null || true
original_status=1
fi
wait "$owned_pid" 2>/dev/null || true
fi
done
if [ "$guard_installed" = true ]; then
sudo iptables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1
fi
if [ "$guard_created" = true ]; then
sudo iptables -F integration_only || original_status=1
sudo iptables -X integration_only || original_status=1
fi
if [ "$guard6_installed" = true ]; then
sudo ip6tables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1
fi
if [ "$guard6_created" = true ]; then
sudo ip6tables -F integration_only || original_status=1
sudo ip6tables -X integration_only || original_status=1
fi
printf '%s\n' "$original_status" > "$results/exit-status.txt"
exit "$original_status"
}
trap cleanup EXIT
trap 'exit 130' INT
trap 'exit 143' TERM
export PATH="$PWD/.venv/bin:$PATH"
export PYTHONPATH="$PWD:$PWD/tests:$PWD/tests/e2e"
export DATABASE_URL="postgresql://postgres:postgres@127.0.0.1:5432/circle_test"
export REDIS_HOST=127.0.0.1 REDIS_PORT=6379
export LITELLM_MASTER_KEY=sk-integration-master LITELLM_SALT_KEY=sk-integration-salt
export LITELLM_MODE=PRODUCTION LITELLM_LOCAL_MODEL_COST_MAP=True
export STORE_MODEL_IN_DB=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1
export INTEGRATION_PROXY_URL=http://127.0.0.1:4000
export INTEGRATION_PEER_URL=""
export INTEGRATION_UPSTREAM_URL=http://127.0.0.1:8190
export INTEGRATION_MASTER_KEY="$LITELLM_MASTER_KEY"
export INTEGRATION_SEED="$((16#$(git rev-parse --short=8 HEAD)))"
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma > "$results/prisma-generate.log" 2>&1
sudo iptables -N integration_only
guard_created=true
sudo iptables -A integration_only -o lo -j ACCEPT
sudo iptables -A integration_only -m conntrack --ctstate ESTABLISHED,RELATED -j ACCEPT
for service in postgres-db redis-cache; do
address="$(docker inspect --format '{{range .NetworkSettings.Networks}}{{.IPAddress}}{{end}}' "$service")"
sudo iptables -A integration_only -d "$address" -j ACCEPT
done
sudo iptables -A integration_only -j REJECT
sudo iptables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only
guard_installed=true
sudo ip6tables -N integration_only
guard6_created=true
sudo ip6tables -A integration_only -o lo -j ACCEPT
sudo ip6tables -A integration_only -j REJECT
sudo ip6tables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only
guard6_installed=true
if curl --noproxy '*' --connect-timeout 2 -s http://198.51.100.1 >/dev/null 2>&1; then
echo "Unexpected outbound network access" >&2
exit 1
fi
sudo iptables -L integration_only -n -v -x > "$results/egress-guard.txt"
awk '$3 == "REJECT" && $1 > 0 { rejected=1 } END { exit !rejected }' "$results/egress-guard.txt"
setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \
.venv/bin/python -m integration._support.upstream > "$results/upstream.log" 2>&1 &
upstream_pid=$!
start_proxy() {
local port="$1"
local log_name="$2"
setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" \
LITELLM_MODE=PRODUCTION LITELLM_LOCAL_MODEL_COST_MAP=True STORE_MODEL_IN_DB=True \
AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
.venv/bin/python -m integration._support.proxy --config tests/integration/proxy_config.yaml \
--host 127.0.0.1 --port "$port" --num_workers 1 --telemetry False \
--use_prisma_db_push --enforce_prisma_migration_check \
> "$results/$log_name" 2>&1 &
launched_pid=$!
}
start_proxy 4000 proxy.log
proxy_pid="$launched_pid"
.venv/bin/python .circleci/scripts/wait_integration_services.py
if [ "$suite" = management ]; then
export INTEGRATION_PEER_URL=http://127.0.0.1:4001
start_proxy 4001 peer.log
peer_pid="$launched_pid"
.venv/bin/python .circleci/scripts/wait_integration_services.py
fi
if [ "$suite" = providers ]; then
INTEGRATION_RUN_ID="$integration_identity" .venv/bin/python -m pytest --noconftest -o addopts= \
--strict-markers --strict-config -p no:pytest-retry -p no:rerunfailures --timeout=30 \
tests/e2e/test_provider_edge.py::TestReplayMode::test_content_drift_returns_the_miss_status_naming_both_keys \
tests/e2e/test_provider_edge.py::TestReplayMode::test_exhausted_key_returns_the_miss_status \
tests/e2e/test_provider_edge.py::TestReplayLeftover::test_partially_consumed_recording_names_the_leftover \
tests/e2e/test_provider_edge.py::TestStreamingFidelity::test_replay_of_a_stream_makes_no_provider_connection \
--junitxml="$results/replay-controls.xml"
fi
timeout --signal=TERM --kill-after=20s 11m env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_RUN_ID="$integration_identity" \
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
INTEGRATION_PROXY_URL="$INTEGRATION_PROXY_URL" INTEGRATION_PEER_URL="$INTEGRATION_PEER_URL" \
INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \
INTEGRATION_MASTER_KEY="$INTEGRATION_MASTER_KEY" LITELLM_MODE=PRODUCTION \
INTEGRATION_SEED="$INTEGRATION_SEED" \
LITELLM_LOCAL_MODEL_COST_MAP=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
.venv/bin/python tests/integration/run.py "$suite" --results "$results"

View file

@ -0,0 +1,53 @@
import sys
from typing import Final
import psutil
def is_owned(process: psutil.Process, identity: str, owner_uid: int) -> bool:
try:
return process.uids().real == owner_uid and process.environ().get("INTEGRATION_RUN_ID") == identity
except psutil.NoSuchProcess:
return False
def owned_processes(identity: str, owner_uid: int) -> tuple[psutil.Process, ...]:
return tuple(process for process in psutil.process_iter() if is_owned(process, identity, owner_uid))
def main(identity: str, owner_uid: int, root_pids: tuple[int, ...]) -> int:
assert owner_uid > 0, "The integration process owner must be a non-root UID"
owned: Final = owned_processes(identity, owner_uid)
roots: Final = tuple(process for process in owned if process.pid in root_pids)
for process in roots:
try:
process.terminate()
except psutil.NoSuchProcess:
continue
psutil.wait_procs(roots, timeout=30)
residual: Final = owned_processes(identity, owner_uid)
for process in residual:
try:
process.terminate()
except psutil.NoSuchProcess:
continue
psutil.wait_procs(residual, timeout=10)
remaining: Final = owned_processes(identity, owner_uid)
for process in remaining:
try:
process.kill()
except psutil.NoSuchProcess:
continue
psutil.wait_procs(remaining, timeout=2)
survivors: Final = owned_processes(identity, owner_uid)
print(
f"Owned integration processes: {len(owned)}, roots: {len(roots)}, "
f"residual: {len(residual)}, forced: {len(remaining)}, remaining: {len(survivors)}"
)
for process in remaining:
print(f"Forced cleanup was required for PID {process.pid}")
return 1 if remaining or survivors else 0
if __name__ == "__main__":
raise SystemExit(main(sys.argv[1], int(sys.argv[2]), tuple(int(value) for value in sys.argv[3:] if value)))

View file

@ -0,0 +1,43 @@
import os
import time
from typing import Final
import httpx
from redis import Redis
def main() -> None:
primary: Final = os.environ["INTEGRATION_PROXY_URL"]
peer: Final = os.environ.get("INTEGRATION_PEER_URL")
proxies: Final = (primary, peer) if peer else (primary,)
deadline: Final = time.monotonic() + 90
headers: Final = {"Authorization": f"Bearer {os.environ['INTEGRATION_MASTER_KEY']}"}
with httpx.Client(trust_env=False, timeout=2) as client, Redis(
host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), socket_timeout=2
) as cache:
while True:
try:
ready: Final = (
client.get(f"{os.environ['INTEGRATION_UPSTREAM_URL']}/health").status_code == 200
and all(client.get(f"{url}/health/readiness").status_code == 200 for url in proxies)
)
if ready:
for url in proxies:
response: Final = client.get(f"{url}/cache/ping", headers=headers)
response.raise_for_status()
result: Final = response.json()
assert result["status"] == "healthy", result
assert result["cache_type"] == "redis", result
assert result["ping_response"] is True, result
assert result["set_cache_response"] == "success", result
if cache.pubsub_numsub("litellm_proxy.auth_cache_invalidation")[0][1] >= len(proxies):
return
except httpx.TransportError:
pass
if time.monotonic() >= deadline:
raise SystemExit("Integration services or auth-cache subscribers did not become ready")
time.sleep(0.2)
if __name__ == "__main__":
main()

View file

@ -14,6 +14,17 @@ query-filters:
id: py/clear-text-logging-sensitive-data # CWE-312
- exclude:
id: py/polynomial-redos # CWE-730
# Import resolution confuses stdlib types with management_endpoints/types.py.
# The generic cycle query also reports intentional deferred imports.
- exclude:
id: py/cyclic-import
- exclude:
id: py/unsafe-cyclic-import
# Known false positives on live settings and Protocol placeholders.
- exclude:
id: py/unused-global-variable
- exclude:
id: py/ineffectual-statement
paths-ignore:
- tests

View file

@ -3,6 +3,9 @@ import xml.etree.ElementTree as ET
from pathlib import Path
from typing import Final
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "tests/e2e"))
from coverage_registry.management_cases import MANAGEMENT_CASES
def main() -> int:
selected: Final = tuple(sys.argv[2:])
@ -16,6 +19,17 @@ def main() -> int:
case.get("file") for case in cases if all(case.find(tag) is None for tag in ("skipped", "failure", "error"))
)
missing: Final = tuple(path for path in selected if path not in passed)
required_nodes: Final = frozenset(case.node for case in MANAGEMENT_CASES if case.node.split("::", 1)[0] in selected)
passed_nodes: Final = frozenset(
prop.get("value")
for case in cases
if all(case.find(tag) is None for tag in ("skipped", "failure", "error"))
for prop in case.findall("./properties/property")
if prop.get("name") == "management_node"
)
missing_nodes: Final = required_nodes - passed_nodes
for node in sorted(missing_nodes):
_ = sys.stdout.write(f"::error::required management case did not pass: {node}\n")
for path in selected:
collected: Final = sum(case.get("file") == path for case in cases)
skipped: Final = sum(case.get("file") == path and case.find("skipped") is not None for case in cases)
@ -27,6 +41,7 @@ def main() -> int:
if (
selected
and not missing
and not missing_nodes
and not any(case.find(tag) is not None for case in cases for tag in ("failure", "error"))
):
return 0

View file

@ -10,7 +10,7 @@ for pid_file in "${STACK_DIR}"/pids/*.pid; do
rm -f "${pid_file}"
done
for container in e2e-nginx e2e-valkey e2e-jaeger e2e-postgres; do
for container in e2e-nginx e2e-keycloak e2e-valkey e2e-jaeger e2e-postgres; do
docker rm -f "${container}" >/dev/null 2>&1
done

5
.github/e2e-stack/oidc-profile.sh vendored Executable file
View file

@ -0,0 +1,5 @@
#!/usr/bin/env bash
set -euo pipefail
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)"
cd "${REPO_ROOT}"
exec uv run --no-sync python tests/e2e/idp.py "$@"

View file

@ -11,6 +11,9 @@ UNSUPPORTED: Final = re.compile(
)
HARNESS: Final = re.compile(
r"^tests/e2e/[A-Za-z0-9_.-]+\.(py|ini)$"
r"|^tests/e2e/idp_realm\.json$"
r"|^tests/e2e/management/(management_client|jwt_actors|conftest)\.py$"
r"|^tests/e2e/coverage_registry/management_cases\.py$"
r"|^tests/e2e/gateway/"
r"|^\.github/e2e-stack/"
r"|^\.github/workflows/test-e2e-changed\.yml$"

45
.github/e2e-stack/start-idp.sh vendored Normal file
View file

@ -0,0 +1,45 @@
#!/usr/bin/env bash
set -euo pipefail
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)"
KEYCLOAK_IMAGE="${E2E_KEYCLOAK_IMAGE:-quay.io/keycloak/keycloak@sha256:ff4257d0d64efbe99ed1ddfaf07765cc3c36dc7518bf8324d41961327f441c54}"
KEYCLOAK_PORT="${E2E_KEYCLOAK_PORT:-8081}"
POSTGRES_IMAGE="${E2E_POSTGRES_IMAGE:-postgres:16.6}"
: "${DATABASE_HOST:?}" "${DATABASE_PORT:?}" "${DATABASE_USER:?}" "${DATABASE_PASSWORD:?}" "${DATABASE_NAME:?}"
DB_HOST="${DATABASE_HOST}"
DB_NETWORK_ARGS=(--network bridge)
IDP_NETWORK_ARGS=(-p "127.0.0.1:${KEYCLOAK_PORT}:${KEYCLOAK_PORT}")
if [[ "$(uname)" == "Linux" ]]; then
DB_NETWORK_ARGS=(--network host)
IDP_NETWORK_ARGS=(--network host)
elif [[ "${DB_HOST}" == "127.0.0.1" || "${DB_HOST}" == "localhost" ]]; then
DB_HOST=host.docker.internal
fi
docker run --rm "${DB_NETWORK_ARGS[@]}" -e "PGPASSWORD=${DATABASE_PASSWORD}" \
"${POSTGRES_IMAGE}" psql -h "${DB_HOST}" -p "${DATABASE_PORT}" \
-U "${DATABASE_USER}" -d "${DATABASE_NAME}" -v ON_ERROR_STOP=1 \
-c 'CREATE SCHEMA IF NOT EXISTS keycloak' >/dev/null
docker rm -f e2e-keycloak >/dev/null 2>&1 || true
docker run -d --name e2e-keycloak "${IDP_NETWORK_ARGS[@]}" --memory 1536m \
-v "${REPO_ROOT}/tests/e2e/idp_realm.json:/opt/keycloak/data/import/realm.json:ro" \
-e KC_DB=postgres -e "KC_DB_URL_HOST=${DB_HOST}" -e "KC_DB_URL_PORT=${DATABASE_PORT}" \
-e "KC_DB_URL_DATABASE=${DATABASE_NAME}" -e KC_DB_SCHEMA=keycloak \
-e "KC_DB_USERNAME=${DATABASE_USER}" -e "KC_DB_PASSWORD=${DATABASE_PASSWORD}" \
-e KC_DB_POOL_INITIAL_SIZE=2 -e KC_DB_POOL_MIN_SIZE=2 -e KC_DB_POOL_MAX_SIZE=10 \
-e "KC_HTTP_PORT=${KEYCLOAK_PORT}" -e KC_BOOTSTRAP_ADMIN_USERNAME=admin \
-e KC_BOOTSTRAP_ADMIN_PASSWORD=e2e-ephemeral-idp-not-a-secret \
"${KEYCLOAK_IMAGE}" start-dev --import-realm >/dev/null
deadline=$((SECONDS + ${E2E_KEYCLOAK_STARTUP_TIMEOUT:-300}))
until curl -fsS --connect-timeout 2 --max-time 3 \
"http://127.0.0.1:${KEYCLOAK_PORT}/realms/litellm-e2e/.well-known/openid-configuration" >/dev/null 2>&1; do
if ((SECONDS >= deadline)); then
echo 'e2e-stack: timed out waiting for the Keycloak realm' >&2
exit 1
fi
sleep 2
done
echo 'e2e-stack: Keycloak realm is up'

View file

@ -25,6 +25,7 @@ DATABASE_PASSWORD="${E2E_DATABASE_PASSWORD:-dbpassword9090}"
DATABASE_NAME="${E2E_DATABASE_NAME:-litellm}"
JAEGER_OTLP_PORT="${E2E_JAEGER_OTLP_PORT:-4318}"
JAEGER_QUERY_PORT="${E2E_JAEGER_QUERY_PORT:-16686}"
KEYCLOAK_PORT="${E2E_KEYCLOAK_PORT:-8081}"
MASTER_KEY="${LITELLM_MASTER_KEY:-sk-e2e-$(openssl rand -hex 16)}"
@ -124,6 +125,9 @@ SERVER_ENV=(
"OTEL_EXPORTER_OTLP_ENDPOINT=http://127.0.0.1:${JAEGER_OTLP_PORT}"
"SSL_CERT_FILE=${CERTS_DIR}/ca-bundle.pem"
"PYTHONPATH=${REPO_ROOT}"
"JWT_PUBLIC_KEY_URL=http://127.0.0.1:${KEYCLOAK_PORT}/realms/litellm-e2e/protocol/openid-connect/certs"
"JWT_ISSUER=http://127.0.0.1:${KEYCLOAK_PORT}/realms/litellm-e2e"
"JWT_AUDIENCE=litellm-e2e"
)
if [[ -n "${VERTEXAI_CREDENTIALS:-}" ]]; then
printf '%s' "${VERTEXAI_CREDENTIALS}" > "${STACK_DIR}/vertex-adc.json"
@ -132,6 +136,8 @@ fi
cd "${REPO_ROOT}"
env "${SERVER_ENV[@]}" "E2E_KEYCLOAK_PORT=${KEYCLOAK_PORT}" bash .github/e2e-stack/start-idp.sh
log "running migrations"
env "${SERVER_ENV[@]}" uv run --no-sync python migrations/run.py >"${LOGS_DIR}/migrations.log" 2>&1
@ -200,6 +206,9 @@ LITELLM_MASTER_KEY=${MASTER_KEY}
REDIS_HOST=127.0.0.1
REDIS_PORT=${REDIS_PORT}
E2E_OTEL_QUERY_URL=http://127.0.0.1:${JAEGER_QUERY_PORT}
E2E_KEYCLOAK_URL=http://127.0.0.1:${KEYCLOAK_PORT}
E2E_KEYCLOAK_ADMIN_USER=admin
E2E_KEYCLOAK_ADMIN_PASSWORD=e2e-ephemeral-idp-not-a-secret
SSL_CERT_FILE=${CERTS_DIR}/ca-bundle.pem
DATABASE_URL=postgresql://${DATABASE_USER}:${DATABASE_PASSWORD}@${DATABASE_HOST}:${DATABASE_PORT}/${DATABASE_NAME}
EOF

View file

@ -47,6 +47,10 @@ After: the same request comes back with real token counts, so the dashboard show
<!-- e.g., "Fixes #000" -->
## Affected release
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1", and add the `backport-stable` label so the fix is cherry-picked onto the rc line before the stable is tagged. Leave the section blank otherwise -->
## Linear ticket
<!-- if you are an internal contributor, add "Resolves " followed by the Linear ticket e.g., "Resolves LIT-1234" to link the Linear ticket to the GitHub PR. If you don't have one, leave the section blank rather than guessing -->
@ -97,7 +101,8 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac
For bug fixes: Before shows the reproduction, After shows the same steps passing
For new features: Before shows the capability missing, After shows it working end-to-end
If the change applies to all three LLM endpoints (/v1/responses, /v1/chat/completions, /v1/messages), make each endpoint its own case, not just one
For UI changes: before/after screenshots under the same headings -->
For UI changes: before/after screenshots under the same headings
If the main use case runs through a coding tool like Claude Code or Codex, drive that tool interactively the way the user does (never `claude -p`, `codex exec`, or curl on its own) and embed before/after screenshots of its pane under the same headings; curl replays and headless runs can follow as extra cases, never as the only proof -->
## Type
@ -154,3 +159,4 @@ Example checklists:
## Final Attestation
- [ ] The tests check the right things, including the edge cases, and regressions in the respective real-world customer use-cases are not possible after this PR

View file

@ -1,6 +1,7 @@
from __future__ import annotations
import ast
import json
import operator
import pathlib
import re
@ -498,6 +499,69 @@ def _check_shards() -> int:
return 0
def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozenset[str], tuple[Finding, ...]]:
manifest: Final = repo_root / "tests/integration/contracts.json"
if not manifest.exists():
return frozenset(), ()
entries: Final = json.loads(manifest.read_text())
paths: Final = frozenset(node.split("::", 1)[0] for node in entries["tests"])
circle_path: Final = repo_root / ".circleci/config.yml"
circle: Final = yaml.safe_load(circle_path.read_text()) if circle_path.exists() else {}
steps: Final = circle.get("jobs", {}).get("integration_contracts", {}).get("steps", ())
invoked: Final = any(
".circleci/scripts/run_integration.sh" in scalar.value
for scalar in _scalars(steps, "integration_contracts")
if scalar.key == "command"
)
scheduled: Final = frozenset(
suite
for job in circle.get("workflows", {}).get("integration", {}).get("jobs", ())
if isinstance(job, dict) and "integration_contracts" in job
for suite in job["integration_contracts"]
.get("matrix", {})
.get("parameters", {})
.get("suite", (job["integration_contracts"].get("suite"),))
if isinstance(suite, str)
)
required: Final = frozenset(
group
for group, folders in entries["groups"].items()
if any(any(path.startswith(f"tests/integration/{folder}/") for folder in folders) for path in paths)
)
ungrouped: Final = frozenset(
path
for path in paths
if sum(
any(path.startswith(f"tests/integration/{folder}/") for folder in folders)
for folders in entries["groups"].values()
)
!= 1
)
gha_tokens: Final = _invoked_test_tokens(
scalar
for path in (repo_root / ".github/workflows").glob("*.y*ml")
for scalar in _scalars(yaml.safe_load(path.read_text()), path.name)
)
findings: Final = tuple(
Finding(path, "integration contract is also selected by GitHub Actions")
for path in paths
if any(_token_covers(token, path) for token in gha_tokens)
) + tuple(
Finding(path, "canonical integration test file is missing")
for path in paths
if not (repo_root / path).is_file()
)
group_findings: Final = tuple(
Finding(group, "canonical integration group is not scheduled by CircleCI")
for group in sorted(required - scheduled)
) + tuple(Finding(path, "canonical node must have exactly one integration group") for path in sorted(ungrouped))
if not paths or not invoked or not scheduled:
return frozenset(), findings + (
Finding(str(manifest.relative_to(repo_root)), "dedicated CircleCI runner is missing"),
)
return paths, findings + group_findings
def main() -> int:
if "--shards" in sys.argv[1:]:
return _check_shards()
@ -507,7 +571,8 @@ def main() -> int:
allowlist = _load_allowlist()
scalars = _all_scalars()
test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars))
integration_paths, ownership_findings = _integration_ownership()
test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars) | integration_paths) + ownership_findings
dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars))
stale_findings = _stale_allowlist_paths(allowlist, test_files=_test_files(), dockerfiles=_dockerfiles())

View file

@ -1,6 +1,8 @@
import asyncio
import aiohttp
import json
import math
from typing import Any
# Asynchronously fetch data from a given URL
async def fetch_data(url):
@ -21,11 +23,157 @@ async def fetch_data(url):
print("Error fetching data from URL:", e)
return None
FRIENDLI_API_URL = "https://api.friendli.ai/serverless/v1/models"
FRIENDLI_PROVIDER = "friendliai"
INHERITABLE_BASE_KEYS = (
"supports_pdf_input",
"supports_assistant_prefill",
"supports_adaptive_thinking",
"supports_output_config",
)
REASONING_EFFORT_LEVEL_ORDER = ("none", "minimal", "low", "medium", "high", "xhigh", "max")
def _find_base_model_entry(base_model: str, local_data: dict) -> str | None:
if not base_model:
return None
bm_tail = base_model.split("/")[-1].lower()
if base_model in local_data:
return base_model
for key in local_data:
if key.startswith("sample_spec") or key == "fallback_generalizations":
continue
if key.split("/")[-1].lower() == bm_tail:
return key
return None
def _reasoning_effort_levels(reasoning_options: list) -> list:
offered = {
val
for opt in reasoning_options or []
if opt.get("type") == "effort"
for val in opt.get("values", [])
}
return [level for level in REASONING_EFFORT_LEVEL_ORDER if level in offered]
def _valid_token_price(value: object) -> bool:
try:
price = float(value) # pyright: ignore[reportArgumentType] # non-numeric values are rejected via the except
except (TypeError, ValueError):
return False
return math.isfinite(price) and price >= 0
def _has_valid_token_prices(pricing: dict | None) -> bool:
prices = pricing or {}
return _valid_token_price(prices.get("input")) and _valid_token_price(prices.get("output"))
def _pricing(pricing: dict) -> dict:
out: dict[str, Any] = {}
if not pricing:
return out
if "input" in pricing:
out["input_cost_per_token"] = float(pricing["input"])
if "output" in pricing:
out["output_cost_per_token"] = float(pricing["output"])
if "input_cache_read" in pricing and pricing["input_cache_read"] is not None:
out["cache_read_input_token_cost"] = float(pricing["input_cache_read"])
return out
def _modality_flags(input_mods: list) -> dict:
mods = input_mods or []
has_image = "image" in mods
return {
"supports_vision": has_image,
"supports_image_input": has_image,
"supports_video_input": "video" in mods,
}
def transform_friendli_data(data: list, local_data: dict) -> dict:
transformed: dict[str, dict] = {}
if not data:
return transformed
for model in data:
# An unpriced row must never wholesale-replace an already priced local entry:
# missing prices cost-calculate as zero, silently zeroing tracked spend
if not _has_valid_token_prices(model.get("pricing")):
continue
model_id = model["id"]
base_model = model.get("base_model") or ""
entry: dict[str, Any] = {
"litellm_provider": FRIENDLI_PROVIDER,
}
base_key = _find_base_model_entry(base_model, local_data)
if base_key:
base_entry = local_data[base_key]
for k in INHERITABLE_BASE_KEYS:
if k in base_entry:
entry[k] = base_entry[k]
ctx = model.get("context_length")
if ctx is not None:
entry["max_input_tokens"] = int(ctx)
max_out = model.get("max_completion_tokens")
if max_out is not None:
entry["max_output_tokens"] = int(max_out)
entry["max_tokens"] = int(max_out)
pricing = _pricing(model.get("pricing", {}))
entry.update(pricing)
entry["supports_prompt_caching"] = "cache_read_input_token_cost" in pricing
reasoning = model.get("reasoning") is True
entry["supports_reasoning"] = reasoning
if reasoning:
entry["reasoning_effort_levels"] = _reasoning_effort_levels(
model.get("reasoning_options", [])
)
func = model.get("functionality", {})
entry["supports_function_calling"] = func.get("tool_call") is True
entry["supports_parallel_function_calling"] = func.get("parallel_tool_call") is True
is_struct = func.get("structured_output") is True
entry["supports_response_schema"] = is_struct
entry["supports_native_structured_output"] = is_struct
entry["supports_system_messages"] = func.get("system_messages") is True
entry["supports_tool_choice"] = func.get("tool_choice") is True
entry.update(_modality_flags(model.get("input_modalities", [])))
entry["mode"] = model.get("mode", "chat")
desc = model.get("description")
if desc:
entry["comment"] = desc
dep = model.get("deprecation_date")
if dep:
entry["deprecation_date"] = dep.split("T")[0]
entry["source"] = FRIENDLI_API_URL
transformed[f"{FRIENDLI_PROVIDER}/{model_id}"] = entry
return transformed
# Synchronize local data with remote data
def sync_local_data_with_remote(local_data, remote_data):
def sync_local_data_with_remote(local_data, remote_data, replace_keys=frozenset()):
# Update existing keys in local_data with values from remote_data
# (replace_keys entries are swapped wholesale so a field the remote catalog
# dropped, e.g. cache pricing, cannot survive as a stale value)
for key in (set(local_data) & set(remote_data)):
local_data[key].update(remote_data[key])
if key in replace_keys:
local_data[key] = remote_data[key]
else:
local_data[key].update(remote_data[key])
# Add new keys from remote_data to local_data
for key in (set(remote_data) - set(local_data)):
@ -46,6 +194,8 @@ def write_to_file(file_path, data):
# Update the existing models and add the missing models for OpenRouter
def transform_openrouter_data(data):
transformed = {}
if not data:
return transformed
for row in data:
# Add the fields 'max_tokens' and 'input_cost_per_token'
obj = {
@ -84,7 +234,14 @@ def transform_openrouter_data(data):
# Update the existing models and add the missing models for Vercel AI Gateway
def transform_vercel_ai_gateway_data(data):
transformed = {}
if not data:
return transformed
for row in data:
# Rows without token pricing or token limits (video/embedding models) previously KeyError'd the whole sync
if any(row.get(k) is None for k in ("context_window", "max_tokens")) or any(
row.get("pricing", {}).get(k) is None for k in ("input", "output")
):
continue
obj = {
"max_tokens": row["context_window"],
"input_cost_per_token": float(row["pricing"]["input"]),
@ -143,13 +300,16 @@ def main():
vercel_data = asyncio.run(fetch_data(vercel_ai_gateway_url))
# Transform the fetched Vercel AI Gateway data
vercel_data = transform_vercel_ai_gateway_data(vercel_data)
friendli_data = asyncio.run(fetch_data(FRIENDLI_API_URL))
friendli_data = transform_friendli_data(friendli_data, local_data)
# Combine both datasets
all_remote_data = {**openrouter_data, **vercel_data}
all_remote_data = {**openrouter_data, **vercel_data, **friendli_data}
# If both local and openrouter data are available, synchronize and save
if local_data and all_remote_data:
sync_local_data_with_remote(local_data, all_remote_data)
sync_local_data_with_remote(local_data, all_remote_data, replace_keys=frozenset(friendli_data))
write_to_file(local_file_path, local_data)
else:
print("Failed to fetch model data from either local file or URL.")

View file

@ -57,6 +57,7 @@ permissions:
env:
UV_PYTHON: "3.12"
LITELLM_LOCAL_MODEL_COST_MAP: "True"
jobs:
run:
@ -113,6 +114,7 @@ jobs:
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 8
run: |
diff -u model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
uv run --no-sync python -c 'import os, sys; print(sys.version); assert f"{sys.version_info.major}.{sys.version_info.minor}" == os.environ["UV_PYTHON"]'

View file

@ -1,42 +0,0 @@
name: Guard main branch
on:
pull_request:
branches:
- main
merge_group:
permissions: {}
# DO NOT RENAME the job's `name:` — it is referenced by GitHub branch
# protection as a required status check on `main`. Renaming silently
# breaks the gate.
jobs:
guard:
name: Verify PR source branch
runs-on: ubuntu-latest
timeout-minutes: 2
steps:
- name: Reject merge_group events
if: github.event_name == 'merge_group'
run: |
echo "::error::Merge queue is not supported for main. Disable merge queue or update this guard."
exit 1
- name: Check head branch name
env:
HEAD_REF: ${{ github.head_ref }}
HEAD_REPO: ${{ github.event.pull_request.head.repo.full_name }}
BASE_REPO: ${{ github.repository }}
run: |
echo "PR head repo: $HEAD_REPO"
echo "PR head branch: $HEAD_REF"
if [ "$HEAD_REPO" != "$BASE_REPO" ]; then
echo "::error::PRs to main must originate from the canonical repository ($BASE_REPO), not a fork ($HEAD_REPO). External contributors should open PRs against 'litellm_internal_staging' instead."
exit 1
fi
if [ "$HEAD_REF" = "litellm_internal_staging" ] || [[ "$HEAD_REF" == litellm_hotfix_?* ]]; then
echo "Allowed source branch."
exit 0
fi
echo "::error::PRs to main must originate from 'litellm_internal_staging' or a 'litellm_hotfix_*' branch. Got: '$HEAD_REF'. If this is a contribution, retarget the PR against 'litellm_internal_staging' instead."
exit 1

View file

@ -26,6 +26,7 @@ on:
- ui/Dockerfile
- ui/nginx.conf
- .github/workflows/image-scan.yml
- .grype.yaml
schedule:
- cron: "41 6 * * *"
workflow_dispatch:
@ -93,6 +94,7 @@ jobs:
GRYPE_MATCH_PYTHON_USING_CPES: "true"
run: |
"$RUNNER_TEMP/grype" litellm-image-scan:${{ github.sha }} \
--config .grype.yaml \
--only-fixed \
--fail-on high \
--output table

View file

@ -81,7 +81,7 @@ jobs:
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_workflow_job_name_collisions.py
- name: test_e2e_changed_gate
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_changed_gate.py
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_changed_gate.py tests/code_coverage_tests/test_e2e_idp_stack.py
- name: router_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py
@ -178,7 +178,7 @@ jobs:
version: "0.10.9"
- name: Install dependencies
run: uv sync --frozen --extra proxy --python 3.10
run: uv sync --frozen --extra proxy --extra cli --python 3.10
- run: uv run --no-sync python --version
@ -187,3 +187,6 @@ jobs:
- name: Check litellm CLI
run: uv run --no-sync litellm --version
- name: Check lite CLI
run: uv run --no-sync lite version

View file

@ -27,6 +27,8 @@ jobs:
sparse-checkout: |
.github/e2e-stack
tests/e2e/access_control
tests/e2e/management/test_jwt_management_e2e.py
tests/e2e/other/test_jwt_auth_e2e.py
persist-credentials: false
ref: ${{ github.sha }}
@ -45,7 +47,8 @@ jobs:
--jq '.[] | select(.status != "removed") | .filename')"
gh api "repos/${REPO}/pulls/${PR_NUMBER}" --jq '.head.sha' | grep -Fxq "${HEAD_SHA}"
tests="$(printf '%s\n' "${files}" \
| python3 .github/e2e-stack/select_tests.py tests/e2e/access_control/test_*.py)"
| python3 .github/e2e-stack/select_tests.py tests/e2e/access_control/test_*.py \
tests/e2e/management/test_jwt_management_e2e.py tests/e2e/other/test_jwt_auth_e2e.py)"
echo "tests=${tests}" >> "${GITHUB_OUTPUT}"
if [ -n "${tests}" ]; then
echo "any=true" >> "${GITHUB_OUTPUT}"
@ -180,7 +183,7 @@ jobs:
log="${RUNNER_TEMP}/e2e-pass-${pass}.log"
echo "::group::pass ${pass} of 3"
set +e
uv run --no-sync pytest "${test_files[@]}" --rootdir=. -v -p no:cacheprovider \
uv run --no-sync pytest "${test_files[@]}" --rootdir=. -v --reruns 0 -p no:cacheprovider \
-o junit_family=xunit1 --junitxml="${report}" > "${log}" 2>&1
status=$?
uv run --no-sync python .github/e2e-stack/assert_tests_ran.py "${report}" "${test_files[@]}"

3
.gitignore vendored
View file

@ -147,3 +147,6 @@ crash.*.log
ui/litellm-dashboard/out/
litellm.log
.coverage-rust
coverage-rust.xml

13
.grype.yaml Normal file
View file

@ -0,0 +1,13 @@
# Wolfi's security database names zlib 1.3.3-r0 as the fix for CVE-2026-85091,
# but the newest zlib published to the Wolfi apk repo is 1.3.2-r7, so every
# wolfi-base digest reports it and no `apk upgrade` can clear it.
# Drop this once Wolfi ships zlib >= 1.3.3-r0; expected by 2026-10-15.
ignore:
- vulnerability: CVE-2026-85091
package:
name: zlib
type: apk
- vulnerability: GHSA-g5fp-32jq-cfw2
package:
name: zlib
type: apk

View file

@ -25,6 +25,8 @@ Same thing for bug fixes. The tests should make it so that this specific bug can
Never test structure of code only function of it
A test must only fail when litellm code changes. Never pin facts we don't own (a vendor's price, a third party's field, an upstream default, today's date) as literals or as "X must be absent"; assert the invariant our code guarantees instead, e.g. two rows agree, a value is within range, a field is derived from another. If an outside fact is truly load-bearing, cite its source and date next to the assertion so a reader can tell stale from broken
`tests/test_litellm/` mirrors `litellm/` in a parallel path (see `tests/test_litellm/readme.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `CLAUDE.md`

View file

@ -299,6 +299,9 @@ test-rust-extension:
[ "$$#" -eq 1 ] && \
UV_PROJECT_ENVIRONMENT="$$temporary/venv" $(UV) sync --python 3.12 --frozen --no-install-project --all-groups --all-extras && \
$(UV) pip install --python "$$temporary/venv/bin/python" --no-deps "$$1" && \
"$$temporary/venv/bin/python" -I -m mypy.stubtest \
--mypy-config-file tests/test_litellm/rust_bridge/stubtest.ini \
litellm.rust_bridge._native && \
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust

View file

@ -27,10 +27,13 @@ import litellm
from litellm import Router, verbose_logger
from litellm._uuid import uuid
from litellm.caching.caching import DualCache
from litellm.constants import MAX_FILE_LIST_LIMIT
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_metadata,
)
from openai.types.file_deleted import FileDeleted
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.llms.base_llm.managed_resources.isolation import (
build_list_page,
@ -48,7 +51,6 @@ from litellm.proxy._types import (
from litellm.proxy.openai_files_endpoints.common_utils import (
BATCH_CREATE_HIDDEN_PARAM,
FILE_LIST_CONTINUATION_CHUNK_SIZE,
MAX_FILE_LIST_LIMIT,
_is_base64_encoded_unified_file_id,
apply_unified_file_ids,
decode_model_from_file_id,
@ -1787,7 +1789,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
litellm_parent_otel_span: Optional[Span],
llm_router: Router,
**data: Dict,
) -> OpenAIFileObject:
) -> FileDeleted:
# Check if file deletion should be blocked due to batch references
await self._check_file_deletion_allowed(file_id)
@ -1795,7 +1797,6 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# file_id = convert_b64_uid_to_unified_uid(file_id)
model_file_id_mapping = await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span)
delete_response = None
specific_model_file_id_mapping = model_file_id_mapping.get(file_id)
if specific_model_file_id_mapping:
# Remove conflicting keys from data to avoid duplicate keyword arguments
@ -1810,23 +1811,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
else {}
),
}
delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
stored_file_object = await self.delete_unified_file_id(file_id, litellm_parent_otel_span)
await self.delete_unified_file_id(file_id, litellm_parent_otel_span)
# Record successful deletion metric only on actual success
if stored_file_object or delete_response:
prom_logger = self._get_prometheus_logger()
if prom_logger:
prom_logger.record_managed_file_deleted(result="success")
if stored_file_object:
return OpenAIFileObject.model_validate(stored_file_object).model_copy(update={"id": file_id})
elif delete_response:
delete_response.id = file_id
return delete_response
else:
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
prom_logger = self._get_prometheus_logger()
if prom_logger:
prom_logger.record_managed_file_deleted(result="success")
return FileDeleted(id=file_id, object="file", deleted=True)
async def afile_content(
self,

View file

@ -780,7 +780,10 @@ async def update_project(
# Handle budget updates
budget_fields = LiteLLM_BudgetTable.model_fields.keys()
budget_updates = {k: v for k, v in update_data.items() if k in budget_fields}
budget_updates = {
**{k: v for k, v in update_data.items() if k in budget_fields},
**({"max_budget": None} if "max_budget" in data.model_fields_set and data.max_budget is None else {}),
}
if budget_updates and existing_project.budget_id:
# Update existing budget

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.66"
version = "0.1.68"
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.66"
version = "0.1.68"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -3,7 +3,8 @@
The gateway exposes the LLM data-plane surface: chat/completions, embeddings,
audio, batches, files, fine-tuning, rerank, ocr, rag, video, search, image,
responses, vector stores, passthrough providers, realtime websockets, MCP
tool-call endpoints, and operational endpoints (/health, /metrics).
tool-call endpoints, and operational endpoints (/health, /metrics, and the
/debug/memory/summary read of the serving worker's RSS).
Any path not listed here is dropped from the gateway process so management/UI
endpoints don't ride on the same pods.
@ -95,6 +96,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/langfuse/",
"/vllm/",
"/mistral/",
"/nvidia_nim/",
"/groq/",
"/voyage/",
"/cursor/",
@ -121,6 +123,7 @@ GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(
"/docs/oauth2-redirect",
"/redoc",
"/test",
"/debug/memory/summary",
}
)

View file

@ -257,6 +257,14 @@ IAM_TOKEN_DB_AUTH / AZURE_POSTGRESQL_AUTH toggle that only the writer sets.
- name: DATABASE_SCHEMA
value: {{ .schema | quote }}
{{- end }}
{{- if .sslMode }}
- name: DATABASE_SSLMODE
value: {{ .sslMode | quote }}
{{- end }}
{{- if .sslRootCert }}
- name: DATABASE_SSLROOTCERT
value: {{ .sslRootCert | quote }}
{{- end }}
{{- if and .useIAMAuth .useAzureEntraAuth }}
{{- fail "database.writer.useIAMAuth and database.writer.useAzureEntraAuth are mutually exclusive: the database password can only come from one token source" }}
{{- end }}

View file

@ -89,7 +89,7 @@
at "/" Prefix would swallow the whole backend management API) instead of
adding to it.
*/}}
{{- $builtinPathKeys := list "/test|Exact" "/|Prefix" -}}
{{- $builtinPathKeys := list "/test|Exact" "/debug/memory/summary|Exact" "/|Prefix" -}}
apiVersion: networking.k8s.io/v1
kind: Ingress
metadata:
@ -129,6 +129,8 @@ spec:
# --- Gateway data plane ---
# Exact /test only (see the $gatewayPrefixes comment above);
# /test/* MCP management endpoints fall to the backend catch-all.
# Exact /debug/memory/summary reads a serving worker's RSS (the e2e memory
# gate); the rest of /debug/* stays on the backend.
- path: /test
pathType: Exact
backend:
@ -136,6 +138,13 @@ spec:
name: {{ $gatewayName }}
port:
number: {{ $gatewayPort }}
- path: /debug/memory/summary
pathType: Exact
backend:
service:
name: {{ $gatewayName }}
port:
number: {{ $gatewayPort }}
{{- range $gatewayPrefixes }}
{{- $pathType := include "litellm.ingress.pathType" (dict "controller" $controller "path" . "pathType" "Prefix") }}
{{- $builtinPathKeys = append $builtinPathKeys (printf "%s|%s" . $pathType) }}

View file

@ -4,6 +4,7 @@ templates:
- gateway/configmap.yaml
- backend/deployment.yaml
- backend/configmap.yaml
- migrations-job.yaml
values:
- ./values/required.yaml
tests:
@ -67,6 +68,82 @@ tests:
value: "true"
any: true
- it: emits no TLS env by default
template: gateway/deployment.yaml
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLMODE
any: true
- notContains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLROOTCERT
any: true
- it: writer sslMode and sslRootCert reach gateway and backend as DATABASE_SSLMODE and DATABASE_SSLROOTCERT
templates:
- gateway/deployment.yaml
- backend/deployment.yaml
set:
database.writer.useIAMAuth: true
database.writer.sslMode: verify-full
database.writer.sslRootCert: /etc/ssl/certs/ca-certificates.crt
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLMODE
value: verify-full
any: true
- contains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLROOTCERT
value: /etc/ssl/certs/ca-certificates.crt
any: true
- it: writer sslMode and sslRootCert reach the collector sidecar and the migrations job, which dial Postgres themselves
set:
gateway.collector.enabled: true
database.connectionPool.enabled: true
database.writer.sslMode: verify-full
database.writer.sslRootCert: /etc/ssl/certs/ca-certificates.crt
asserts:
- equal:
path: spec.template.spec.containers[1].name
value: collector
template: gateway/deployment.yaml
- contains:
path: spec.template.spec.containers[1].env
content:
name: DATABASE_SSLMODE
value: verify-full
any: true
template: gateway/deployment.yaml
- contains:
path: spec.template.spec.containers[1].env
content:
name: DATABASE_SSLROOTCERT
value: /etc/ssl/certs/ca-certificates.crt
any: true
template: gateway/deployment.yaml
- contains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLMODE
value: verify-full
any: true
template: migrations-job.yaml
- contains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLROOTCERT
value: /etc/ssl/certs/ca-certificates.crt
any: true
template: migrations-job.yaml
- it: writer rejects both token sources at once
template: gateway/deployment.yaml
set:

View file

@ -97,6 +97,16 @@ tests:
name: RELEASE-NAME-litellm-gateway
port:
number: 4000
- contains:
path: spec.rules[0].http.paths
content:
path: /debug/memory/summary
pathType: Exact
backend:
service:
name: RELEASE-NAME-litellm-gateway
port:
number: 4000
- equal:
path: spec.rules[0].http.paths[-1]
value:

View file

@ -288,6 +288,17 @@ tests:
- failedTemplate:
errorMessage: "ingress.extraPaths[0]: path /test with pathType Exact is already routed by this chart, and a duplicate would take it over rather than add to it"
- it: rejects an entry that would take over the exact /debug/memory/summary route
set:
ingress.enabled: true
ingress.extraPaths:
- path: /debug/memory/summary
pathType: Exact
service: backend
asserts:
- failedTemplate:
errorMessage: "ingress.extraPaths[0]: path /debug/memory/summary with pathType Exact is already routed by this chart, and a duplicate would take it over rather than add to it"
- it: allows a built-in path under a different pathType, which is a distinct rule
set:
ingress.enabled: true

View file

@ -208,6 +208,11 @@ database:
name: litellm-writer-secret
usernameKey: username
passwordKey: password
# libpq sslmode / sslrootcert applied to the writer and reader URLs (Prisma and the
# in-container PgBouncer); e.g. verify-full with /etc/ssl/certs/ca-certificates.crt for AWS RDS.
# sslRootCert on its own implies sslMode verify-full
sslMode: ""
sslRootCert: ""
# Optional read-replica routing. When `reader.host` is set, the proxy routes
# reads (find_*, count, group_by, query_raw/_first) to this endpoint while

View file

@ -40,4 +40,4 @@ if not logger.handlers:
logging.Formatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s")
)
logger.addHandler(handler)
logger.setLevel(logging.INFO)
logger.setLevel(os.getenv("LITELLM_LOG", "INFO").upper())

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN IF NOT EXISTS "litellm_call_id" TEXT;

View file

@ -0,0 +1,12 @@
-- CreateIndex (CONCURRENTLY)
--
-- Disclaimer:
-- - CREATE INDEX CONCURRENTLY cannot run inside a transaction. This migration must stay a
-- single statement so Prisma Migrate on PostgreSQL can apply it outside a transaction.
-- - Builds are slower and use more I/O than a blocking CREATE INDEX; if the build is
-- interrupted, Postgres may leave an INVALID index that must be dropped and recreated.
-- - Do not edit this file after it has been applied to any database: Prisma checksums
-- migrations; add a new migration instead.
-- - Requires PostgreSQL that supports CONCURRENTLY with IF NOT EXISTS (use a new migration
-- without IF NOT EXISTS if you must support older versions).
CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_SpendLogs_litellm_call_id_idx" ON "LiteLLM_SpendLogs"("litellm_call_id");

View file

@ -0,0 +1,14 @@
-- AlterTable
ALTER TABLE "LiteLLM_BudgetTable" ADD COLUMN IF NOT EXISTS "tpd_limit" BIGINT;
-- AlterTable
ALTER TABLE "LiteLLM_TeamTable" ADD COLUMN IF NOT EXISTS "tpd_limit" BIGINT;
-- AlterTable
ALTER TABLE "LiteLLM_DeletedTeamTable" ADD COLUMN IF NOT EXISTS "tpd_limit" BIGINT;
-- AlterTable
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN IF NOT EXISTS "tpd_limit" BIGINT;
-- AlterTable
ALTER TABLE "LiteLLM_DeletedVerificationToken" ADD COLUMN IF NOT EXISTS "tpd_limit" BIGINT;

View file

@ -0,0 +1,23 @@
-- AlterTable
ALTER TABLE "LiteLLM_DailyUserSpend" ADD COLUMN IF NOT EXISTS "total_response_time_ms" BIGINT NOT NULL DEFAULT 0;
ALTER TABLE "LiteLLM_DailyUserSpend" ADD COLUMN IF NOT EXISTS "timed_requests" BIGINT NOT NULL DEFAULT 0;
-- AlterTable
ALTER TABLE "LiteLLM_DailyOrganizationSpend" ADD COLUMN IF NOT EXISTS "total_response_time_ms" BIGINT NOT NULL DEFAULT 0;
ALTER TABLE "LiteLLM_DailyOrganizationSpend" ADD COLUMN IF NOT EXISTS "timed_requests" BIGINT NOT NULL DEFAULT 0;
-- AlterTable
ALTER TABLE "LiteLLM_DailyEndUserSpend" ADD COLUMN IF NOT EXISTS "total_response_time_ms" BIGINT NOT NULL DEFAULT 0;
ALTER TABLE "LiteLLM_DailyEndUserSpend" ADD COLUMN IF NOT EXISTS "timed_requests" BIGINT NOT NULL DEFAULT 0;
-- AlterTable
ALTER TABLE "LiteLLM_DailyAgentSpend" ADD COLUMN IF NOT EXISTS "total_response_time_ms" BIGINT NOT NULL DEFAULT 0;
ALTER TABLE "LiteLLM_DailyAgentSpend" ADD COLUMN IF NOT EXISTS "timed_requests" BIGINT NOT NULL DEFAULT 0;
-- AlterTable
ALTER TABLE "LiteLLM_DailyTeamSpend" ADD COLUMN IF NOT EXISTS "total_response_time_ms" BIGINT NOT NULL DEFAULT 0;
ALTER TABLE "LiteLLM_DailyTeamSpend" ADD COLUMN IF NOT EXISTS "timed_requests" BIGINT NOT NULL DEFAULT 0;
-- AlterTable
ALTER TABLE "LiteLLM_DailyTagSpend" ADD COLUMN IF NOT EXISTS "total_response_time_ms" BIGINT NOT NULL DEFAULT 0;
ALTER TABLE "LiteLLM_DailyTagSpend" ADD COLUMN IF NOT EXISTS "timed_requests" BIGINT NOT NULL DEFAULT 0;

View file

@ -0,0 +1,18 @@
-- DropIndex
DROP INDEX IF EXISTS "LiteLLM_JWTKeyMapping_jwt_claim_name_jwt_claim_value_is_act_idx";
-- DropIndex
DROP INDEX IF EXISTS "LiteLLM_JWTKeyMapping_jwt_claim_name_jwt_claim_value_key";
-- AlterTable
-- NOT NULL DEFAULT '' (not nullable): Postgres unique constraints treat every
-- NULL as distinct, so a nullable column would let multiple unscoped mappings
-- collide on the same claim without a constraint violation. The constant
-- default is a fast, metadata-only backfill for existing rows, not a rewrite.
ALTER TABLE "LiteLLM_JWTKeyMapping" ADD COLUMN IF NOT EXISTS "jwt_issuer" TEXT NOT NULL DEFAULT '';
-- CreateIndex
CREATE INDEX IF NOT EXISTS "LiteLLM_JWTKeyMapping_jwt_issuer_jwt_claim_name_jwt_claim_v_idx" ON "LiteLLM_JWTKeyMapping"("jwt_issuer", "jwt_claim_name", "jwt_claim_value", "is_active");
-- CreateIndex
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_JWTKeyMapping_jwt_issuer_jwt_claim_name_jwt_claim_v_key" ON "LiteLLM_JWTKeyMapping"("jwt_issuer", "jwt_claim_name", "jwt_claim_value");

View file

@ -37,11 +37,13 @@ raised it above the deploy default keeps that larger budget for deploy unless
the deploy override says otherwise.
"""
import importlib.util
import math
import os
import shutil
import signal
import subprocess
import sys
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
@ -64,6 +66,7 @@ DEFAULT_PRISMA_BOOTSTRAP_TIMEOUT = 600.0
DEFAULT_PRISMA_MIGRATE_DEPLOY_TIMEOUT = 600.0
BOOTSTRAP_ARG = "--version"
PRISMA_CONSOLE_SCRIPT = "prisma"
@dataclass(frozen=True)
@ -184,6 +187,28 @@ def _kill_process_group(process: "subprocess.Popen[str]") -> None:
return
def prisma_cli_available() -> bool:
"""Whether some way of running the Prisma CLI exists: the console script on PATH or the importable package."""
if shutil.which(PRISMA_CONSOLE_SCRIPT) is not None:
return True
return importlib.util.find_spec(PRISMA_CONSOLE_SCRIPT) is not None
def resolve_prisma_argv(argv: Sequence[str]) -> tuple[str, ...]:
"""Route a bare ``prisma`` command through ``python -m prisma`` when the console script is not on PATH.
The console script and ``python -m prisma`` are the same entry point, but
only the module form survives an interpreter whose ``bin`` directory is
missing from PATH, which is how the proxy gets started under launchers and
init systems. Any other executable name is left untouched.
"""
if not argv or argv[0] != PRISMA_CONSOLE_SCRIPT:
return tuple(argv)
if shutil.which(PRISMA_CONSOLE_SCRIPT) is not None:
return tuple(argv)
return (sys.executable, "-m", PRISMA_CONSOLE_SCRIPT, *argv[1:])
def run_prisma(
argv: Sequence[str],
*,
@ -200,7 +225,7 @@ def run_prisma(
text unless ``stdout``/``stderr`` say otherwise.
"""
with subprocess.Popen(
argv,
resolve_prisma_argv(argv),
env=env,
stdout=stdout,
stderr=stderr,

View file

@ -17,6 +17,7 @@ model LiteLLM_BudgetTable {
max_parallel_requests Int?
tpm_limit BigInt?
rpm_limit BigInt?
tpd_limit BigInt?
model_max_budget Json?
budget_duration String?
budget_reset_at DateTime?
@ -133,6 +134,7 @@ model LiteLLM_TeamTable {
max_parallel_requests Int?
tpm_limit BigInt?
rpm_limit BigInt?
tpd_limit BigInt?
budget_duration String?
budget_reset_at DateTime?
blocked Boolean @default(false)
@ -203,6 +205,7 @@ model LiteLLM_DeletedTeamTable {
max_parallel_requests Int?
tpm_limit BigInt?
rpm_limit BigInt?
tpd_limit BigInt?
budget_duration String?
budget_reset_at DateTime?
blocked Boolean @default(false)
@ -438,6 +441,7 @@ model LiteLLM_VerificationToken {
blocked Boolean?
tpm_limit BigInt?
rpm_limit BigInt?
tpd_limit BigInt?
max_budget Float?
budget_duration String?
budget_reset_at DateTime?
@ -483,6 +487,10 @@ model LiteLLM_VerificationToken {
model LiteLLM_JWTKeyMapping {
id String @id @default(uuid())
jwt_issuer String @default("") // Scopes the mapping to one configured issuer; "" matches any issuer.
// Not nullable: Postgres unique constraints treat every NULL as
// distinct, so a nullable column would let multiple unscoped
// mappings collide on the same claim without a constraint violation.
jwt_claim_name String // e.g. "sub", "email"
jwt_claim_value String // The claim value to match
token String // Hashed virtual key (FK)
@ -495,8 +503,8 @@ model LiteLLM_JWTKeyMapping {
litellm_verification_token LiteLLM_VerificationToken @relation(fields: [token], references: [token], onDelete: Cascade)
@@unique([jwt_claim_name, jwt_claim_value])
@@index([jwt_claim_name, jwt_claim_value, is_active])
@@unique([jwt_issuer, jwt_claim_name, jwt_claim_value])
@@index([jwt_issuer, jwt_claim_name, jwt_claim_value, is_active])
}
// Deprecated keys during grace period - allows old key to work until revoke_at
@ -534,6 +542,7 @@ model LiteLLM_DeletedVerificationToken {
blocked Boolean?
tpm_limit BigInt?
rpm_limit BigInt?
tpd_limit BigInt?
max_budget Float?
budget_duration String?
budget_reset_at DateTime?
@ -659,12 +668,14 @@ model LiteLLM_SpendLogs {
mcp_namespaced_tool_name String?
agent_id String?
proxy_server_request Json? @default("{}")
litellm_call_id String?
created_at DateTime @default(now()) @map("created_at")
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
@@index([startTime])
@@index([startTime, request_id])
@@index([end_user])
@@index([session_id])
@@index([litellm_call_id])
}
model LiteLLM_BudgetWindowSpend {
@ -790,6 +801,8 @@ model LiteLLM_DailyUserSpend {
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@ -826,6 +839,8 @@ model LiteLLM_DailyOrganizationSpend {
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@ -862,6 +877,8 @@ model LiteLLM_DailyEndUserSpend {
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@unique([end_user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@ -897,6 +914,8 @@ model LiteLLM_DailyAgentSpend {
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@unique([agent_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@ -932,6 +951,8 @@ model LiteLLM_DailyTeamSpend {
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
ptu_flat_cost Float @default(0.0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@ -970,6 +991,8 @@ model LiteLLM_DailyTagSpend {
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.96"
version = "0.4.98"
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.96"
version = "0.4.98"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -1,29 +0,0 @@
# Adding a provider / route to litellm-rust
Everything for a route lives in `crates/core/src/<route>/`; `crates/core/src/messages` is the reference. A host (the axum gateway, the Python bridge) only calls the route's entrypoint.
1. **Entrypoint** — `mod.rs`: `pub async fn <route>(request) -> CoreResult<Response>`, the Rust equivalent of `litellm.<route>()`, plus a `<route>_stream` variant when the route streams. It is the only thing a host touches.
2. **Transform contract** — `transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) with types in `types.rs`.
3. **Provider config** — `crates/core/src/providers/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
4. **Prepare + handler** — `prepare.rs` resolves provider/model, credentials, auth headers, and URL, then transforms the request; `handler.rs` performs the provider call through the shared client in `client.rs` and transforms the response.
## Coding standards
Before writing new logic, look for an existing base to extend. When a change is
“the same behavior for one more provider/endpoint/integration”, the codebase
almost always already has a shared abstraction for it (for example, provider
`BaseConfig` transformation classes in `litellm/llms/base_llm/`, shared
helpers in `litellm_core_utils/`, typed request/response models, or factory
functions). Find it first with a search, then add the new variant by inheriting
from or composing that base, overriding only what genuinely differs (model
name, parameter mapping, or auth).
Never copy an existing implementation and edit it in place, and never hand-roll
a parallel version of logic a base already provides. If you catch yourself
writing a second copy of a pattern that exists twice already, stop and extract a
base instead: put the shared shape in one place and make both call sites thin
variants of it. The test for a good abstraction is that adding the next provider
is a few declarative lines, not a new file of duplicated flow. Only diverge from
the base when behavior is genuinely different, and say so explicitly in the PR.
**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run the commands under "Checks" in [CLAUDE.md](CLAUDE.md).

View file

@ -1,45 +0,0 @@
# AGENTS.md
litellm-rust has six crates. A crate is a layer or shared foundation, not a route. Routes (ocr, realtime, chat) and providers (mistral, openai) are modules inside the layers.
## Crates
| Crate | Role |
|-------|------|
| litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. |
| litellm-token-counter | Standalone input token counting shared by host integrations without pulling in the full SDK. |
| litellm-config | Config-loading boundary. Returns resolved core deployment data and optionally delegates loading to Python. |
| litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. |
| litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. |
| litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. Owns API registration, domain wiring, and Python exception mapping. |
Dependency direction is acyclic: `litellm-config` depends on `litellm-core`, the gateway depends on both, and `litellm-python-bridge` depends on the domain layers, `litellm-token-counter`, and `litellm-python-interop`. The token counter and interop foundations depend on no LiteLLM domain crate.
## Where a route lives
A top-level LiteLLM call is a module under `crates/core/src/<route>/`, shaped like `messages`:
```
core/src/messages/
mod.rs # pub async fn messages(..) -> CoreResult<..> (+ messages_stream for SSE)
types.rs # request/response types, MessagesRequest
transformation.rs # the provider template trait
prepare.rs # provider resolution, auth headers, URL
handler.rs # the provider call
client.rs # the shared reqwest client
```
Handlers never live in `ai-gateway`. `ocr`, `audio_transcription`, and `realtime` are still hosted there from before this rule; they move to `core` as they are touched.
Adding a crate: default to a module. A new crate requires a real trigger: separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these.
Adding a crate fails crates/core/tests/workspace_crate_allowlist.rs until you update its allowlist and this file — intentional.
## Style
All Rust in `litellm-rust/` follows the official Rust Style Guide:
https://doc.rust-lang.org/style-guide/
`rustfmt` implements its formatting by default, so run `cargo fmt` before committing; CI gates every PR on `cargo fmt --check`. Do not hand-format against rustfmt or add a `rustfmt.toml` that diverges from the default style.
Beyond formatting, follow the guide's naming and idiom conventions rustfmt cannot auto-apply: `snake_case` items/functions/modules, `UpperCamelCase` types/traits/variants, `SCREAMING_SNAKE_CASE` constants/statics (acronyms as one word, e.g. `HttpClient`), and the import grouping and item ordering it prescribes. See CLAUDE.md for the detailed version.

View file

@ -1,189 +0,0 @@
# CLAUDE.md
This file defines the rules for Rust work in LiteLLM.
## Provider Coding Standards
Before writing new logic, look for an existing base to extend. When a change is
“the same behavior for one more provider/endpoint/integration”, the codebase
almost always already has a shared abstraction for it (for example, provider
`BaseConfig` transformation classes in `litellm/llms/base_llm/`, shared
helpers in `litellm_core_utils/`, typed request/response models, or factory
functions). Find it first with a search, then add the new variant by inheriting
from or composing that base, overriding only what genuinely differs (model
name, parameter mapping, or auth).
Never copy an existing implementation and edit it in place, and never hand-roll
a parallel version of logic a base already provides. If you catch yourself
writing a second copy of a pattern that exists twice already, stop and extract a
base instead: put the shared shape in one place and make both call sites thin
variants of it. The test for a good abstraction is that adding the next provider
is a few declarative lines, not a new file of duplicated flow. Only diverge from
the base when behavior is genuinely different, and say so explicitly in the PR.
## Crates (see AGENTS.md)
`litellm-core` **is** the LiteLLM SDK in Rust: it makes the LLM call.
`litellm-config` is the config-loading boundary and returns resolved core types.
`litellm-ai-gateway` is an HTTP/WebSocket server in front of it, and
`litellm-python-bridge` exposes it to the Python SDK. `litellm-python-interop`
holds domain-neutral PyO3 primitives shared by Python-facing Rust code. A crate
is a layer or shared foundation, not a route; add modules, not crates.
## Core Boundary
`litellm-core` owns the whole call. The Rust equivalent of `litellm.messages()`
is `litellm_core::messages::messages(request).await`: you call it, it does the
provider call, and you get a typed non-streaming response back.
Route-level Rust structure mirrors LiteLLM's Python responsibilities:
- `core/src/<route>/` owns the route end to end: the public entrypoint fn named
after the route in `mod.rs`, the request/response types (`types.rs`), the
provider template trait (`transformation.rs`), the provider/auth/URL
resolution (`prepare.rs`), the HTTP client (`client.rs`), and the handler that
performs the call (`handler.rs`). `core/src/messages` is the reference.
- `core/src/providers/<provider>/<route>/transformation.rs` owns the
provider-specific transform. For Anthropic Messages, this means
`core/src/providers/anthropic/messages/transformation.rs`.
- Handlers live in `core`, never in a host. `ai-gateway` must not contain a
route handler that talks to a provider; its axum route reads the HTTP request,
picks a deployment, and calls the `core` entrypoint. `python-bridge` marshals
Python objects and calls the same entrypoint.
Streaming keeps the same shape: the route entrypoint has a `<route>_stream`
variant in `core` that returns the upstream response so a host can splice it to
its own caller; the host still owns no provider logic.
Call-hook and lifecycle instrumentation, including phase timing, usage
accumulation, and callback payload construction, always lives in `core`.
Hosts feed observed events into core and dispatch the completed payloads through
their I/O logger; hosts must not own callback orchestration.
Allowed in `core`:
- The public entrypoint for a top-level LiteLLM call
- Request/response transforms and stream chunk normalization
- Provider resolution, auth header construction, and URL building
- The provider HTTP call itself, through a shared reused client with connect and
request timeouts
- Shared data types and validation errors
- Deterministic token/cost helper logic
Not allowed in `core`:
- Serving HTTP: axum routes, extractors, and transport concerns stay in the host
- Filesystem access
- Database access
- Config file reading and rollout state
- Logging callbacks, spend writes, or custom callbacks
- Global mutable runtime state
Env reads in `core` are limited to credential fallback inside a route's
`prepare.rs` (the `env_lookup` closure), mirroring what the Python SDK does when
no key is passed. Everything else config-shaped is resolved by the host and
passed in.
Routes still hosted in `ai-gateway` (`ocr`, `audio_transcription`, `realtime`)
predate this rule and are being moved into `core` route modules; do not add new
ones there, and prefer moving one when you touch it.
Python owns rollout state and fallback while Rust is being introduced. Rust
paths must be off by default until parity tests prove equivalence with Python.
A new provider/route may instead be implemented rust-only with no Python
reference; then the Python interface is a thin dispatch that calls Rust with no
fallback, and you state the rust-only choice explicitly in the PR. Either way
the Python side stays minimal (it only marshals inputs and calls the Rust
interface), never add a per-route feature flag, and never push provider
dispatch into `litellm/main.py`; put it in a thin dispatch class under
`litellm/llms/<provider>/<route>/`.
## Production Bar
Rust code in this workspace is held to a strict parity and robustness bar from
the first PR:
- Correctness parity is proven with tests. Do not rely on README claims or
manual inspection for a port that mirrors Python behavior.
- Every provider transform must have unit tests for supported-parameter
filtering, request body shape, response normalization, missing/null fields,
and bad-input errors.
- When Rust is exposed through Python, add Python tests that prove disabled,
enabled, and unavailable-bridge fallback behavior.
- Avoid panics on user/provider input. Return typed errors and let the host map
them to Python exceptions or HTTP responses.
- OCR handles documents that often contain personal data. Do not log document
contents, base64 payloads, provider response bodies, or secrets.
- Error messages must be useful but data-minimized. Truncate or sanitize any
upstream body before it crosses a host boundary.
- Treat empty or whitespace-only credentials, URLs, and config values as absent
at the host/config resolution layer.
- Preserve Python output shape intentionally. If a field is always serialized as
`null` for Python parity, leave a short comment explaining that parity choice.
## Network I/O Rules
These rules apply to every module that executes network I/O, whether it is a
`core` route handler or a host such as `ai-gateway`:
- Set connect and full-request timeouts. No unbounded waits.
- Reuse HTTP clients; do not construct clients per request.
- Prefer rustls TLS for portable Python wheels and Linux images unless there is
a documented reason not to.
- Add request IDs and structured tracing at the host layer, without logging OCR
document contents or secrets.
- Do not echo raw upstream response bodies to callers. Sanitize and bound them.
- Avoid `expect`/`unwrap` in server startup and request paths unless the panic is
impossible by construction and documented.
## Rust Style Guide
All Rust in `litellm-rust/` follows the official Rust Style Guide:
https://doc.rust-lang.org/style-guide/
`rustfmt` implements the guide's formatting rules by default, so the mechanical
side is enforced for you: run `cargo fmt` before committing and CI gates every
PR on `cargo fmt --check` (see Checks). Do not hand-format against rustfmt or add
a `rustfmt.toml` that diverges from the default style; the default style *is* the
guide.
The guide also covers conventions rustfmt cannot auto-apply; follow these too:
- Naming: `snake_case` for items, functions, and modules; `UpperCamelCase` for
types, traits, and enum variants; `SCREAMING_SNAKE_CASE` for constants and
statics; acronyms count as one word (`HttpClient`, not `HTTPClient`).
- Ordering and grouping the guide prescribes: imports grouped std / external /
crate-local, derives before other attributes, and consistent item order.
- Idioms the guide recommends over the formatter fighting you (e.g. prefer
restructuring an over-long expression rather than forcing an awkward wrap).
## Constants
Magic numbers and fixed strings go in a crate-level `constants.rs`, never
hardcoded inline — the Rust mirror of Python's `litellm/constants.py`.
- Each crate that needs them has `src/constants.rs` (declared `mod constants;`);
import from it (`use crate::constants::...`). Don't scatter `const` values at
the top of feature modules.
- An env-overridable tunable still lives in `constants.rs` as its `DEFAULT_*`
value; the env read (with fallback to that default) happens at the host/config
resolution layer, not in `core`/`providers`.
- Exception: a value that is purely local to one function and has no meaning
elsewhere may stay inline, but prefer `constants.rs` when in doubt.
## Checks
Run these before pushing Rust changes. The same checks run in GitHub Actions
for changes under `litellm-rust/`.
```bash
cd litellm-rust
cargo fmt --check
cargo clippy --workspace --all-targets -- -D warnings
cargo clippy -p litellm-core --all-targets --features bedrock-auth -- -D warnings
# the ai-gateway binary + server code is behind the `server` feature
cargo clippy -p litellm-ai-gateway --all-targets --all-features -- -D warnings
cargo test --workspace
cargo test -p litellm-core --features bedrock-auth
# the `auth`, `routes`, `state` and `realtime` tests only exist under `server`
cargo test -p litellm-ai-gateway --features server
```
When a Rust path is exposed through Python, add Python parity tests that compare
the existing Python output with the Rust-backed output.

View file

@ -1948,12 +1948,17 @@ dependencies = [
"azure_core",
"azure_identity",
"base64 0.22.1",
"bytes",
"data-url",
"futures-util",
"gcp_auth",
"mime_guess",
"moka",
"rand 0.8.7",
"reqwest 0.12.28",
"rstest",
"rustls 0.23.42",
"rustls-native-certs",
"serde",
"serde_json",
"serde_path_to_error",
@ -1962,6 +1967,7 @@ dependencies = [
"subtle",
"thiserror 2.0.19",
"tokio",
"tokio-tungstenite",
"tracing",
"tracing-subscriber",
"url",
@ -1974,12 +1980,12 @@ version = "0.1.0"
dependencies = [
"criterion",
"futures-util",
"litellm-ai-gateway",
"litellm-core",
"litellm-python-interop",
"litellm-token-counter",
"pyo3",
"pyo3-async-runtimes",
"rstest",
"serde",
"serde_json",
"tokio",

View file

@ -16,6 +16,7 @@ license = "MIT"
repository = "https://github.com/BerriAI/litellm"
[workspace.dependencies]
bytes = "1"
tracing = "0.1"
tracing-subscriber = { version = "0.3", default-features = false, features = ["registry", "std"] }
litellm-core = { path = "crates/core" }

View file

@ -1,56 +0,0 @@
# LiteLLM Rust
This workspace contains the staged Rust implementation for LiteLLM.
`litellm-core` is the LiteLLM SDK in Rust: one entrypoint per top-level call
that makes the LLM call and hands back a typed response, the same shape as
`litellm.messages()` in Python.
```rust
let response = litellm_core::messages::messages(MessagesRequest {
model: "claude-sonnet-4-5",
body,
api_key: Some(key),
..
})
.await?;
```
Python continues to own configuration, retries, routing policy, logging,
callbacks, spend tracking, and customer plugins until each Rust path has parity
coverage and production evidence.
## Crates
| Crate | Role |
|-------|------|
| litellm-core | The SDK. Per-route entrypoints (`messages::messages()`), types, provider transforms (modules under `providers/`), provider resolution, auth, the provider HTTP call, and the router. |
| litellm-config | Config-loading boundary. Returns resolved deployments and optionally delegates loading to Python. |
| litellm-ai-gateway | The axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
| litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. |
| litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. Owns API registration, domain wiring, and Python exception mapping. |
Dependency direction is acyclic: config depends on core, the gateway depends on config and core, and the Python bridge depends on the domain layers and Python interop.
## Layout
```text
crates/
core/ The SDK: route modules + provider transforms.
src/messages/ mod.rs (entrypoint), types, transformation, prepare, handler, client
src/providers/anthropic/messages/transformation.rs
config/ Config loading and resolved deployments.
ai-gateway/ Axum server + WebSocket hosts; calls core entrypoints.
python-interop/ Domain-neutral PyO3 conversion and GIL primitives.
python-bridge/ PyO3 API adapter for Python LiteLLM.
```
The folder shape follows the Python provider tree:
`core/src/providers/<provider>/<route>/transformation.rs`. The bridge exposes one
function per top-level route, mirroring the core entrypoints.
## Checks
Run the commands under "Checks" in [CLAUDE.md](CLAUDE.md) before pushing Rust
changes. That list is the single source of truth and matches what GitHub Actions
runs for changes under `litellm-rust/`.

View file

@ -1,53 +0,0 @@
# Provider coding standards (litellm-rust)
Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages` (`core/src/messages`, `ANTHROPIC_MESSAGES_CONFIG`) is the reference: a route is a `core` module with a public entrypoint that makes the call and returns a typed response.
## Provider resolution
1. Always resolve the provider/model first with `get_custom_llm_provider` (`core/src/routing_utils/provider.rs`). Nothing downstream may branch on a raw model string.
2. Model/provider is resolved once, in `prepare.rs`, and passed down as typed fields. Don't re-resolve or re-parse it in transforms or handlers.
## Transforms and the base config
3. Every route defines a base config trait with `transform_request` + `transform_response` (+ `complete_url`, `supported_params`), living in `core/src/<route>/transformation.rs` (e.g. `AnthropicMessagesProviderConfig`, mirroring `OcrProviderConfig`).
4. Each provider implements that trait as a `const <PROVIDER>_<ROUTE>_CONFIG` in `core/src/providers/<provider>/<route>/transformation.rs`, mirroring the Python provider tree.
5. Individual configs implement only the request/response transforms. Shared behavior (param filtering, defaults) stays as trait default methods so future providers inherit existing logic instead of reimplementing it.
6. Prefer composition: a provider that extends another reuses the base trait's defaults or wraps another config; don't copy transform bodies between providers.
## Boundaries
7. Layers never cross: `core` = the call itself (entrypoint, types, transforms, provider resolution, auth headers, provider HTTP, lifecycle hooks); `ai-gateway` = serving HTTP/WS (routing, extractors, auth of *our* callers, streaming to the client); `python-bridge` = thin PyO3 adapter. Hosts call the core entrypoint; they never build a provider request.
8. Generic/route files contain zero provider-specific branches. A provider is one module under `core/src/providers/<provider>/<route>/`; a route is a module, never a new crate.
9. Route entry point stays thin: `core::<route>::<route>()` -> `prepare_*` -> handler (or `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing). Axum handlers validate and delegate to a service that calls the entrypoint; no business logic in them.
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Config-shaped env reads happen at the host/config layer with the `DEFAULT_*` fallback defined in `constants.rs`; the only env read in `core` is the credential fallback in a route's `prepare.rs`.
## Types and errors
11. Typed contracts only: no bare `serde_json::Value` / `String` / `Vec<String>` as a transform input or output. Parse wire bytes into typed structs/enums at the host edge; a `type` discriminator is a typed field, not a raw string.
12. Model failures as values: return typed `CoreError`, don't panic. No `unwrap`/`expect`/`panic!` on user or provider input.
13. No mutation: build values in one shot (comprehensions/iterators, `collect`), prefer immutable bindings and owned typed structs over seeding-and-mutating.
14. Early returns over deep nesting; small focused files over god modules.
15. Preserve Python output shape intentionally. If a field is always serialized as `null` for parity, keep it and pin it with a test.
## Safety and data minimization
16. Never log request/response bodies, base64 payloads, document contents, or secrets. Truncate and bound any upstream body before it crosses a host boundary.
17. Treat empty/whitespace credentials, URLs, and config values as absent at the host resolution layer.
18. Network I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
## Tests and rollout
19. Every provider transform ships tests for: supported-param filtering, request body shape, response normalization, missing/null fields, bad input, and `*_match_python` fixture parity.
20. Lifecycle/hook tests cover hook order, success + failure callback payloads, pre-call guardrail blocking before any provider I/O, during-call body mutation, and provider-error mapping.
21. When a route has a Python reference implementation, the Rust path stays off by default and behind Python parity tests (disabled / enabled-equals-Python / bridge-unavailable fallback) until parity is proven. A new provider/route may instead be implemented rust-only with no Python reference; then the Python interface is a thin dispatch to Rust with no fallback, and tests cover the rust-backed path plus the unavailable-bridge error. State the rust-only choice explicitly in the PR.
## Python bridge (SDK side)
22. A Python -> Rust bridge keeps the Python side minimal: the Python interface only marshals inputs and calls the Rust interface, with no transform, handler, or business logic. Aim for well under 100 lines of interface code per route; if the Python grows past that, the logic belongs in Rust.
23. Do not bloat `litellm/main.py`. A route's provider dispatch lives in a thin dispatch class under `litellm/llms/<provider>/<route>/` that calls the Rust bridge; `main.py` only instantiates it and calls its sync/async method.
24. Do not add new feature flags unless explicitly requested. Reuse the existing LiteLLM Rust rollout mechanism (`litellm.rust`); never introduce a per-route env flag such as `LITELLM_USE_RUST_<ROUTE>`.
## Checks before push
25. Run, and keep green, the commands under "Checks" in `litellm-rust/CLAUDE.md`.
That list is the single source of truth and matches what GitHub Actions runs.

View file

@ -1,54 +0,0 @@
# ai-gateway — folder architecture
The Axum server that fronts the Rust gateway. It owns transport + config + auth
only; deployment selection lives in `core::router`, and the LLM call itself
(transforms, auth headers, provider HTTP) lives behind a `core` route entrypoint
such as `litellm_core::messages::messages`. No provider handler lives here.
```
src/
main.rs # entrypoint: build AppState (router + master key), bind, serve
state.rs # AppState — shared Arc<Router> + master_key
auth/ # authentication as an axum extractor — added to handler args
mod.rs # RequireMasterKey: FromRequestParts, single master key (LITELLM_MASTER_KEY)
routes/ # one module per route, all matching the same template
AGENTS.md # ← the route template (read this before adding a route)
mod.rs # app(): merges every module's router()
health.rs # simple route (one file): router() + liveness/readiness
realtime/ # route with logic → axum surface + a no-axum service:
mod.rs # router() + handler + WS<->events adapter (the axum surface)
service.rs # business logic (select deployment, call provider) — no axum, testable
```
## Rules
- **Routes follow one template.** Each route module exposes
`pub fn router() -> Router<AppState>`; `routes/mod.rs` only merges them. Simple
routes are one file; non-trivial routes are a folder (`handler`/`service`/
`transport`). See `routes/AGENTS.md`.
- **Auth is an extractor.** Add `crate::auth::RequireMasterKey` to a handler's
args; it runs during extraction. Never re-implement the check per route.
- **Handlers are thin.** A handler validates and delegates to its `service`. No
business logic, no provider calls, no transforms in handlers.
- **Services call `core`, they don't reimplement it.** A `service` picks the
deployment and calls the `core` route entrypoint. Provider resolution, auth
headers, URL building, and the HTTP call are `core`'s job; a service that
builds a provider request itself is a bug (`routes/messages/service.rs` is
the reference).
- **State is shared and cheap to clone.** Long-lived handles live behind `Arc` in
`state.rs`; read env/config only in `main.rs` when building state.
## Auth (interim)
A single **master key** (`LITELLM_MASTER_KEY`), enforced by the
`auth::RequireMasterKey` extractor: any caller presenting it as
`Authorization: Bearer <key>` may invoke the gateway. Fails closed (500) when
unset; constant-time compare. The server binds `127.0.0.1` by default (`HOST` to
override). Full per-key auth + budgets/rate-limits are delegated to the Python
proxy in a later phase. Health routes don't add the extractor (unauthenticated).
## Python interop
Python-backed loading lives in `litellm-config` and is **load-time only**. The
gateway's `python-config` feature forwards to that crate. The realtime data path
never takes the GIL.

View file

@ -1,14 +0,0 @@
# ai-gateway architecture
The Rust ai-gateway does LLM inference (realtime WebSocket). Spend tracking is an
API callback: it POSTs each finished session to the LiteLLM proxy, which records
spend and runs the usual callbacks.
```mermaid
flowchart LR
C[client] <--> G[Rust ai-gateway<br/>LLM inference]
G <--> O[OpenAI realtime]
G -. spend tracking callback .-> P[litellm proxy]
F[litellm-config<br/>load-time only] --> G
F -. Python backend .-> P
```

View file

@ -13,6 +13,11 @@ name = "litellm-ai-gateway"
path = "src/main.rs"
required-features = ["server"]
[[bin]]
name = "trace-parity-gateway"
path = "src/bin/trace_parity_gateway.rs"
required-features = ["trace-parity"]
[dependencies]
tracing.workspace = true
litellm-core = { workspace = true, features = ["bedrock-auth"] }

View file

@ -1,55 +0,0 @@
# Realtime gateway benchmark — pool on/off
Measures what the gateway adds over talking to OpenAI's realtime WebSocket
directly, and what the pre-warmed connection pool removes. See
`../../src/routes/realtime/README.md` for how the pool works.
## Results
5000 calls / 500 concurrency, gateway at 10 instances, pool ON
(`REALTIME_POOL_SIZE=64`), upstream OpenAI `gpt-realtime`. Each leg run twice.
Times in **ms**. Phases per connection: **dial** = TCP+TLS+WS upgrade,
**session** = upgrade → `session.created` (the phase the pool removes),
**1st-audio** = `response.create` → first audio delta (OpenAI inference),
**total** = full wall-clock.
| metric | Direct OpenAI | Gateway (pool ON) | Overhead (ms) | vs OpenAI |
| ------------------ | ------------- | ----------------- | ------------- | ---------- |
| success rate (%) | 99.8 | 99.8 | — | — |
| dial p50 (ms) | 276 | 158 | −118 | **faster** |
| session p50 (ms) | 7 | 0 | −7 | **faster** |
| 1st-audio p50 (ms) | 440 | 664 | +224 | slower¹ |
| total p50 (ms) | 816 | 1010 | +194 | slower¹ |
| total p95 (ms) | 2152 | 1970 | −182 | **faster** |
| total p99 (ms) | 2692 | 2610 | −82 | **faster** |
The gateway is **faster than direct on 4 of 6 metrics**. The warm pool makes the
**session phase sub-millisecond** at the median — ~76% of connects hit the pool,
~70% had session < 1 ms. ¹ The two "slower" rows are not gateway overhead:
`1st-audio` is OpenAI's own inference time (the gateway only relays it), which ran
slower during the gateway legs and drags `total p50` with it.
**Pool OFF** (control, `REALTIME_POOL_SIZE=0`): session p50 was **367 ms** — the
fresh-dial overhead the pool removes.
## Reproduce
The load generator lives in a separate repo:
**https://github.com/ishaan-berri/litellm-realtime-bench**
```bash
git clone https://github.com/ishaan-berri/litellm-realtime-bench
cd litellm-realtime-bench && go build -o wsbench .
# Direct to OpenAI (baseline)
./wsbench -host api.openai.com -key "$OPENAI_API_KEY" -m gpt-realtime -n 5000 -c 500 -t 60
# Through the gateway — run once with pool ON, once with REALTIME_POOL_SIZE=0
./wsbench -host <gateway-host> -key "$LITELLM_MASTER_KEY" -m gpt-realtime -n 5000 -c 500 -t 60
```
Run the gateway with the env stand-in (`OPENAI_REALTIME_MODEL=gpt-realtime`,
`OPENAI_API_KEY`, `LITELLM_MASTER_KEY`, `REALTIME_POOL_SIZE`, `HOST=0.0.0.0`). At
500 concurrency over N instances, size the pool to `≈ 500 / N` per instance (64 was
used here for 10 instances). The bench repo's README covers running 500-concurrency
legs from a hosted multi-vCPU runner. **Never commit keys — pass them via `-key`.**

View file

@ -277,7 +277,7 @@ fn core_error_kind(error: &Error) -> &'static str {
Error::InvalidProvider(_) => "InvalidProvider",
Error::InvalidRequest(_) => "InvalidRequest",
Error::InvalidType { .. } => "InvalidType",
Error::MissingField(_) => "MissingField",
Error::MissingField(_) | Error::MissingDocumentUrl => "MissingField",
Error::Http { .. } => "HttpError",
Error::InvalidResponse(_) => "InvalidResponse",
Error::Network(_) => "NetworkError",

View file

@ -0,0 +1,42 @@
use std::io::Read;
use serde::Deserialize;
use serde_json::Value;
#[derive(Deserialize)]
struct Input {
path: String,
model_alias: String,
provider_model: String,
api_base: String,
body: Value,
}
#[tokio::main]
async fn main() {
let mut input = String::new();
if let Err(error) = std::io::stdin().read_to_string(&mut input) {
fail(error);
}
let input: Input = match serde_json::from_str(&input) {
Ok(input) => input,
Err(error) => fail(error),
};
let result = litellm_ai_gateway::trace_parity::traced_request(
input.path,
input.model_alias,
input.provider_model,
input.api_base,
input.body,
)
.await;
match serde_json::to_string(&result) {
Ok(result) => println!("{result}"),
Err(error) => fail(error),
}
}
fn fail(error: impl std::fmt::Display) -> ! {
eprintln!("{error}");
std::process::exit(1)
}

View file

@ -1,5 +1,3 @@
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use futures_util::stream::{SplitSink, SplitStream};
@ -10,106 +8,21 @@ use litellm_core::auth::error::MissingCredential;
use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG;
use litellm_core::responses::types::ResponsesWsEvent;
use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig;
use tokio::net::TcpStream;
use tokio::sync::Mutex;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::HeaderValue;
use tokio_tungstenite::tungstenite::http::header::{AUTHORIZATION, HeaderName};
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
use crate::io::tls::connect_upstream;
use litellm_core::responses::websocket::{ResponsesUpstreamWs, connect_upstream};
use crate::constants::{
DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS,
};
const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
pub type ResponsesUpstreamWs = WebSocketStream<MaybeTlsStream<TcpStream>>;
type UpstreamTx = SplitSink<ResponsesUpstreamWs, Message>;
type UpstreamRx = SplitStream<ResponsesUpstreamWs>;
#[derive(Clone)]
pub struct ResponsesWebSocketConnection {
socket: Arc<Mutex<Option<ResponsesUpstreamWs>>>,
}
impl ResponsesWebSocketConnection {
pub async fn connect_url(
url: &str,
headers: &HashMap<String, String>,
timeout: Option<Duration>,
) -> Result<Self, Error> {
let mut request = url
.into_client_request()
.map_err(|error| Error::Network(error.to_string()))?;
for (name, value) in headers {
let header_name = name
.parse::<HeaderName>()
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
let header_value = HeaderValue::from_str(value)
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
request.headers_mut().insert(header_name, header_value);
}
let connect = connect_upstream(request);
let result = match timeout {
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
Error::Network("Responses WebSocket connection timed out".to_string())
})?,
None => connect.await,
};
let (socket, _) = result.map_err(|error| match *error {
tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http {
status: response.status().as_u16(),
body: String::new(),
},
other => Error::Network(other.to_string()),
})?;
Ok(Self {
socket: Arc::new(Mutex::new(Some(socket))),
})
}
pub async fn send_text(&self, text: String) -> Result<(), Error> {
let mut socket = self.socket.lock().await;
let Some(socket) = socket.as_mut() else {
return Err(Error::Network("Responses WebSocket is closed".to_string()));
};
socket
.send(Message::Text(text))
.await
.map_err(|error| Error::Network(error.to_string()))
}
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
let mut socket_guard = self.socket.lock().await;
let Some(socket) = socket_guard.as_mut() else {
return Ok(None);
};
match socket.next().await {
Some(Ok(Message::Text(text))) => Ok(Some(text)),
Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec())
.map(Some)
.map_err(|error| Error::InvalidResponse(error.to_string())),
Some(Ok(Message::Close(_))) | None => Ok(None),
Some(Ok(_)) => Ok(None),
Some(Err(error)) => Err(Error::Network(error.to_string())),
}
}
pub async fn close(&self) -> Result<(), Error> {
let mut socket = self.socket.lock().await;
if let Some(socket) = socket.as_mut() {
socket
.close(None)
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
*socket = None;
Ok(())
}
}
pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result<String, Error> {
api_key
.map(str::trim)

View file

@ -118,7 +118,8 @@ impl IntoResponse for MessagesRouteError {
| Error::Connect(_)
| Error::InvalidResponse(_)
| Error::InvalidType { .. }
| Error::MissingField(_) => (
| Error::MissingField(_)
| Error::MissingDocumentUrl => (
StatusCode::BAD_GATEWAY,
"messages provider request failed".to_string(),
),

View file

@ -10,6 +10,7 @@ use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter};
use serde::Serialize;
use serde_json::Value;
use tower::ServiceExt;
use tracing::instrument::WithSubscriber;
use crate::io::realtime_pool::RealtimePool;
use crate::routes;
@ -21,7 +22,41 @@ pub struct GatewayResponse {
pub body: Value,
}
pub async fn messages_request(
#[derive(Debug, Serialize)]
pub struct TracedGatewayResponse {
pub response: Option<GatewayResponse>,
pub error: Option<String>,
pub trace: Vec<litellm_core::observability::FunctionTraceEvent>,
}
pub async fn traced_request(
path: String,
model_alias: String,
provider_model: String,
api_base: String,
body: Value,
) -> TracedGatewayResponse {
let trace = litellm_core::observability::FunctionTrace::default();
let result = request(path, model_alias, provider_model, api_base, body)
.with_subscriber(trace.dispatcher())
.await;
let events = trace.events();
match result {
Ok(response) => TracedGatewayResponse {
response: Some(response),
error: None,
trace: events,
},
Err(error) => TracedGatewayResponse {
response: None,
error: Some(error.to_string()),
trace: events,
},
}
}
pub async fn request(
path: String,
model_alias: String,
provider_model: String,
api_base: String,
@ -42,7 +77,7 @@ pub async fn messages_request(
};
let request = Request::builder()
.method("POST")
.uri("/v1/messages")
.uri(path)
.header(AUTHORIZATION, "Bearer trace-master-key")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(body.to_string()))

View file

@ -2,10 +2,10 @@
//! API has to resolve its own crypto provider, in a test binary where nothing
//! has installed a process-wide one, and has to leave it uninstalled.
use std::collections::HashMap;
use std::time::Duration;
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection;
use futures_util::{sink, stream};
use litellm_ai_gateway::io::responses_ws::async_responses_websocket;
use tokio::net::TcpListener;
async fn dead_tls_server() -> u16 {
@ -30,10 +30,15 @@ async fn dead_tls_server() -> u16 {
async fn dialing_wss_returns_an_error_instead_of_panicking() {
let port = dead_tls_server().await;
let result = ResponsesWebSocketConnection::connect_url(
&format!("wss://127.0.0.1:{port}/"),
&HashMap::new(),
let result = async_responses_websocket(
"gpt-5",
Some("test-key"),
Some(&format!("wss://127.0.0.1:{port}/")),
None,
Some(Duration::from_secs(10)),
|_| {},
stream::empty(),
sink::drain(),
)
.await;

View file

@ -2,6 +2,6 @@ litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-leve
A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate.
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`.
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or host-specific callback execution. Core owns lifecycle sequencing and callback payload construction; hosts execute the selected integrations. Env reads are limited to credential fallback in a route's `prepare.rs`.
Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates.

View file

@ -1,66 +0,0 @@
# CLAUDE.md
Rules for `litellm-rust/crates/core`.
## Responsibility
`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level
LiteLLM call has a public entrypoint here, named after the route
(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and
calling it returns a typed non-streaming response.
Allowed:
- The public entrypoint for a route, plus its `<route>_stream` variant when the
route supports streaming.
- Provider resolution, auth header construction, URL building, and the provider
HTTP call (shared reused client, connect + request timeouts).
- Shared request/response structs.
- Typed errors with stable, non-sensitive messages.
- Deterministic validation helpers.
- Serialization helpers that intentionally mirror Python output shape.
- Route templates that match Python base config responsibilities, such as
`messages::transformation::AnthropicMessagesProviderConfig`.
Not allowed:
- Serving HTTP: axum routers, extractors, and other transport concerns.
- Filesystem, database, or cache access.
- Config file reading or rollout state; the host resolves those and passes them
in. Env reads are limited to credential fallback in a route's `prepare.rs`.
- Logging callbacks, tracing spans, spend writes, or customer callbacks.
- Provider-specific branching that belongs in `providers`.
- Panics for user/provider-controlled input.
## Typed Contracts (core rule)
Trait and function boundaries MUST be strongly typed. No stringly-typed JSON
(`&str` / `String` / `Vec<String>` / bare `serde_json::Value`) as a transform
input or output. Parse wire bytes into typed structs/enums at the host edge;
`core` and `providers` operate only on those types (e.g. `RealtimeEvent`,
`RealtimeTransformResult`, `OcrRequestData`). A `type`-style discriminator is a
typed field on a struct, not a raw string threaded through the API.
## Structure
Use route names directly under `src/`: `messages`, `ocr`, future
`chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not
invent broad names like `engine` for route contracts.
`src/messages` is the reference shape for a route module:
```
mod.rs pub async fn messages(..) (+ messages_stream)
types.rs request/response types
transformation.rs the provider template trait
prepare.rs provider resolution, auth headers, URL
handler.rs the provider call
client.rs the shared reqwest client
```
## Parity Rules
- Every shared type used by a provider transform needs unit tests for
serialization shape.
- If Python parity requires always emitting a `null` field instead of omitting
it, document that in code and pin it with a test.
- Error enums should preserve enough detail for Python/HTTP hosts to map errors
consistently without exposing document contents or upstream bodies.

View file

@ -6,25 +6,27 @@ license.workspace = true
repository.workspace = true
autotests = false
[[test]]
name = "workspace_crate_allowlist"
path = "tests/workspace_crate_allowlist.rs"
[dependencies]
bytes.workspace = true
futures-util.workspace = true
base64.workspace = true
azure_core.workspace = true
azure_identity.workspace = true
data-url = "0.3.2"
gcp_auth.workspace = true
moka.workspace = true
mime_guess = "2.0.5"
rand.workspace = true
reqwest.workspace = true
rustls.workspace = true
rustls-native-certs.workspace = true
serde.workspace = true
serde_json.workspace = true
serde_path_to_error = "0.1"
strum.workspace = true
subtle.workspace = true
tokio.workspace = true
tokio = { workspace = true, features = ["sync"] }
tokio-tungstenite.workspace = true
thiserror.workspace = true
tracing.workspace = true
tracing-subscriber = { workspace = true, optional = true }

View file

@ -9,6 +9,21 @@ use crate::AuthError;
use super::{ResolvedCredential, SecretValue, TokenProviderHandle};
pub fn credential_index(requested: &str, names: &[String]) -> Option<usize> {
names.iter().position(|name| name == requested)
}
pub fn credential_default_fields<'a>(
supplied: &[String],
credential_fields: &'a [String],
) -> Vec<&'a str> {
credential_fields
.iter()
.filter(|name| !supplied.contains(name))
.map(String::as_str)
.collect()
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum CredentialFileRef {
Path(PathBuf),

View file

@ -49,6 +49,7 @@ impl<T> Sourced<T> {
pub use credential::{
CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialPlan,
CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle,
credential_default_fields, credential_index,
};
pub use http::{CredentialPlacement, RequestAuth};
pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy};

View file

@ -1,167 +0,0 @@
# Call lifecycle
`litellm_core::call_lifecycle` is the shared execution wrapper for LiteLLM call
types migrated to Rust. It owns lifecycle ordering, phase timing, and trace
observer calls. It must not know about OCR, chat, messages, responses,
completions, provider auth, request transforms, or response normalization.
Call-type modules own their domain behavior. For example, OCR owns document
payloads, OCR provider transforms, safe document fetch, guardrail payload shape,
callback payload shape, and provider HTTP execution.
## Runtime order
Every wrapped call runs in this order:
1. `async_pre_call_hook`
2. `async_during_call_hook`
3. provider call
4. `async_log_success_event` or `async_log_failure_event`
`async_pre_call_hook` receives the initial LiteLLM request shape. It is where
pre-call custom guardrails run.
`async_during_call_hook` converts the initial request into the provider-ready
request. It is where provider config selection, parameter mapping, auth/header
resolution, request transforms, and during-call guardrails belong.
The provider call receives only the provider-ready request. It should execute
I/O and call the provider response transform.
Success and failure callbacks receive `CallLifecycleTiming`. Callback failures
must not replace the original provider or guardrail result.
## Trace contract
The lifecycle runner records:
- full call start and end time
- `pre_call` phase timing
- `during_call` phase timing
- `provider_call` phase timing
- `success_callback` phase timing
- `failure_callback` phase timing
`CallLifecycleObserver` receives phase start and end events. The default
observer is a no-op. Future OTEL support should implement this observer instead
of editing OCR, chat, messages, responses, completions, or provider modules.
## Required shape
Each migrated call type should use this folder shape:
```text
litellm-rust/crates/ai-gateway/src/<call_type>/
mod.rs # thin public entrypoint
types.rs # public request, prepared request, provider request, response types
prepare.rs # model/provider/callback/guardrail setup
hooks.rs # CallLifecycleHooks implementation
handler.rs # provider I/O and response normalization
tests.rs # call-type lifecycle and handler tests
```
Provider transforms can live in `litellm-rust/crates/core/src/providers/...`.
Shared call-type helpers can live beside the call type, but generic lifecycle
code stays in this folder.
## Core API
The prepared request implements `CallLifecycleRequest`:
```rust
impl CallLifecycleRequest for PreparedMessagesRequest {
fn lifecycle_context(&self) -> CallLifecycleContext {
CallLifecycleContext::new(
"messages",
self.model.clone(),
self.custom_llm_provider.clone(),
self.litellm_call_id.clone(),
)
}
}
```
The call-type hooks implement `CallLifecycleHooks`:
```rust
impl CallLifecycleHooks<
PreparedMessagesRequest,
ProviderMessagesRequest,
MessagesResponse,
> for MessagesLifecycleHooks {
fn async_pre_call_hook(...) {
// run pre-call custom guardrails against the LiteLLM request shape
}
fn async_during_call_hook(...) {
// map params, validate env, transform request, run during-call guardrails
}
fn async_log_success_event(...) {
// call async_log_success_event on configured custom loggers
}
fn async_log_failure_event(...) {
// call async_log_failure_event without swallowing the original error
}
}
```
The public entrypoint stays thin:
```rust
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<MessagesResponse> {
let PreparedMessagesCall { request, hooks } = prepare_messages_call(request)?;
CallLifecycle::default()
.run_request(request, &hooks, execute_messages_provider_call)
.await
}
```
Use `run_request` for new call types. Keep `run` available only for specialized
tests or existing code that already has a `CallLifecycleContext`.
## Adding a new call type
1. Add `<call_type>/types.rs`
Define the public request accepted by the bridge, the prepared request used by
the lifecycle runner, and the provider request consumed by the handler.
2. Implement `CallLifecycleRequest`
Return `call_type`, `model`, `custom_llm_provider`, and `litellm_call_id`.
Do not put provider-specific logic here.
3. Add `<call_type>/prepare.rs`
Resolve model/provider once, generate or preserve `litellm_call_id`, construct
callback and guardrail runners, and return `Prepared<CallType>Call`.
4. Add `<call_type>/hooks.rs`
Implement `CallLifecycleHooks`. Put pre-call guardrail payload construction,
provider config selection, param mapping, request transform, during-call
guardrail payload construction, and callback payload construction here.
5. Add `<call_type>/handler.rs`
Execute the provider request and normalize the provider response. Do not repeat
provider-specific transforms here; call the provider config.
6. Add tests
Cover hook order, success callback payload, failure callback payload, pre-call
guardrail blocking before provider I/O, during-call body mutation, and provider
error mapping.
## Review checklist
- Core lifecycle has no call-type or provider-specific branches
- Public call-type entrypoint only prepares and calls `run_request`
- Provider behavior lives behind provider config/transformation code
- Hook method names map to the Python custom logger and guardrail concepts
- Phase timing is recorded once in lifecycle, not separately per call type
- Callback failures never hide the original provider or guardrail error
- Tests prove the provider socket is not touched when pre-call guardrails block

View file

@ -0,0 +1,121 @@
use std::future::Future;
use std::pin::Pin;
pub enum HostCallStep<O, C> {
Host(O),
Complete(C),
}
pub type HostCallFuture<'a, O, C> =
Pin<Box<dyn Future<Output = Result<HostCallStep<O, C>, crate::Error>> + Send + 'a>>;
pub trait HostCall: Send + Sync {
type Operation: Send + 'static;
type Result: Send + 'static;
type Complete: Send + 'static;
fn resume(
&mut self,
result: Option<Self::Result>,
) -> HostCallFuture<'_, Self::Operation, Self::Complete>;
fn interrupt(
&mut self,
failure: HostFailure,
) -> HostCallFuture<'_, Self::Operation, Self::Complete>;
}
pub enum HostStep<V, S> {
Ready(V),
Suspend(S),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum HostPhase {
Setup,
DeploymentPreCall,
Prepare,
Execute,
ConstructResponse,
DeploymentPostCall,
Finalize,
Success,
MapFailure,
DeploymentFailure,
Failure,
AsyncFailure,
Complete,
}
#[derive(Clone, Debug)]
pub enum HostFailure {
Error(crate::Error),
Cancelled(crate::Error),
}
pub struct HostLifecycle {
phase: HostPhase,
asynchronous: bool,
}
impl HostLifecycle {
pub fn new(asynchronous: bool) -> Self {
Self {
phase: HostPhase::Setup,
asynchronous,
}
}
pub fn phase(&self) -> HostPhase {
self.phase
}
pub fn accept(&mut self, result: Result<(), HostFailure>) -> Option<crate::Error> {
if let Err(failure) = result {
if self.phase == HostPhase::DeploymentFailure {
self.phase = HostPhase::Failure;
return None;
}
let error = match failure {
HostFailure::Cancelled(error) => {
self.phase = HostPhase::Complete;
return Some(error);
}
HostFailure::Error(error) => error,
};
match self.phase {
HostPhase::Failure | HostPhase::AsyncFailure => {
self.advance();
return None;
}
HostPhase::Success => self.phase = HostPhase::Complete,
HostPhase::Execute | HostPhase::ConstructResponse => {
self.phase = HostPhase::MapFailure;
}
_ => self.phase = HostPhase::Failure,
}
return Some(error);
}
self.advance();
None
}
fn advance(&mut self) {
self.phase = match self.phase {
HostPhase::Setup if self.asynchronous => HostPhase::DeploymentPreCall,
HostPhase::Setup | HostPhase::DeploymentPreCall => HostPhase::Prepare,
HostPhase::Prepare => HostPhase::Execute,
HostPhase::Execute => HostPhase::ConstructResponse,
HostPhase::ConstructResponse if self.asynchronous => HostPhase::DeploymentPostCall,
HostPhase::ConstructResponse | HostPhase::DeploymentPostCall => HostPhase::Finalize,
HostPhase::Finalize => HostPhase::Success,
HostPhase::MapFailure if self.asynchronous => HostPhase::DeploymentFailure,
HostPhase::MapFailure | HostPhase::DeploymentFailure => HostPhase::Failure,
HostPhase::Failure if self.asynchronous => HostPhase::AsyncFailure,
HostPhase::Failure
| HostPhase::AsyncFailure
| HostPhase::Success
| HostPhase::Complete => HostPhase::Complete,
};
}
}

View file

@ -3,6 +3,10 @@ use std::time::{Instant, SystemTime, UNIX_EPOCH};
use crate::Error;
pub mod host;
#[cfg(test)]
#[path = "../../tests/host_lifecycle.rs"]
mod host_tests;
pub mod types;
pub use types::{

View file

@ -46,9 +46,10 @@ pub const FUNCTION_TRACE_TARGET: &str = "litellm::function_trace";
pub(crate) const MEDIA_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const OCR_RESPONSE_MAX_BYTES: usize = 64 * 1024 * 1024;
pub(crate) const OCR_HTTP_TIMEOUT_SECS: u64 = 600;
pub(crate) const OCR_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const OCR_INLINE_MAX_BYTES: usize = 50 * 1024 * 1024;
pub const OCR_INLINE_MAX_BYTES: usize = 50 * 1024 * 1024;
pub(crate) const OCR_DOWNLOAD_MAX_BYTES: u64 = 50 * 1024 * 1024;
pub(crate) const OCR_MAX_FETCH_REDIRECTS: usize = 10;
pub(crate) const OCR_POLL_TIMEOUT_SECS: u64 = 120;
@ -63,3 +64,6 @@ pub(crate) const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY";
pub(crate) const REDUCTO_ID_PREFIX: &str = "reducto://";
pub(crate) const AZURE_AI_OCR_PATH: &str = "/providers/mistral/azure/ocr";
pub(crate) const MISTRAL_OCR_API_BASE: &str = "https://api.mistral.ai/v1";
pub(crate) const COHERE_PARSE_API_BASE: &str = "https://api.cohere.com";
pub(crate) const COHERE_API_KEY_ENV: &str = "COHERE_API_KEY";

View file

@ -1,6 +1,6 @@
use thiserror::Error as ThisError;
#[derive(Debug, ThisError, PartialEq, Eq)]
#[derive(Clone, Debug, ThisError, PartialEq, Eq)]
pub enum Error {
#[error("expected {expected}, got {actual}")]
InvalidType {
@ -9,6 +9,8 @@ pub enum Error {
},
#[error("missing required field: {0}")]
MissingField(&'static str),
#[error("Document URL is required")]
MissingDocumentUrl,
#[error("invalid response: {0}")]
InvalidResponse(String),
#[error("invalid provider: {0}")]
@ -52,6 +54,17 @@ pub enum Error {
Unsupported(&'static str),
}
impl Error {
pub const fn http_status_code(&self) -> Option<u16> {
match self {
Self::InvalidRequest(_) => Some(400),
Self::MissingDocumentUrl => Some(500),
Self::Http { status, .. } => Some(*status),
_ => None,
}
}
}
#[derive(Debug, ThisError)]
pub(crate) enum MediaError {
#[error("media URL rejected by network policy")]
@ -106,6 +119,7 @@ impl From<crate::ocr::error::OcrRequestError> for Error {
fn from(error: crate::ocr::error::OcrRequestError) -> Self {
match error {
crate::ocr::error::OcrRequestError::MissingField(field) => Self::MissingField(field),
crate::ocr::error::OcrRequestError::MissingDocumentUrl => Self::MissingDocumentUrl,
error => Self::InvalidRequest(error.to_string()),
}
}

View file

@ -0,0 +1,131 @@
use super::super::OcrAdapter;
use crate::Error;
use crate::ocr::OcrClient;
use crate::ocr::codecs::cohere::{
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
};
use crate::ocr::document::{inline_remote_document, validate_inline_document};
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use crate::providers::azure_ai::auth::AzureAuthInputs;
use crate::url_utils::ApiUrl;
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
pub(crate) struct AzureCohereAdapter;
impl OcrAdapter for AzureCohereAdapter {
type ProviderResponse = CohereResponse;
const PROVIDER: OcrProvider = OcrProvider::AzureAi;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = super::super::super::wire::decode_request_value::<CohereParams>(
serde_json::Value::Object(request.optional_params.clone()),
"optional_params",
)?;
let mut config = AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
let base = request
.connection
.api_base
.clone()
.or_else(|| credential_env(AZURE_AI_API_BASE_ENV))
.filter(|base| !base.trim().is_empty())
.ok_or_else(|| {
Error::Auth(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(),
)
})?;
let headers =
super::validate_ai_environment(&request.connection, &config, &credential_env).await?;
validate_document(&request.document)?;
let remote = request.document.source().starts_with("http://")
|| request.document.source().starts_with("https://");
let document = inline_remote_document(
client.document_fetcher(),
request.document.clone(),
&request.connection,
)
.await?;
let body = transform_request(&request.model, document, params)?;
transform_request_body(
client,
request,
&complete_url(&base)?,
&headers,
!remote,
body,
|body| {
validate_document(&body.document)?;
validate_inline_document(&body.document)
},
)
.await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_response(&request.model, response)
}
}
fn complete_url(base: &str) -> Result<String, OcrError> {
let mut url = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
if !matches!(url.scheme(), "http" | "https") {
return Err(invalid_api_base().into());
}
let path = url.path().trim_end_matches('/').to_string();
if path.ends_with("/v2/parse") {
url.set_path(&path);
return Ok(url.into());
}
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"]))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base().into())
}
fn invalid_api_base() -> OcrRequestError {
OcrRequestError::RequestField {
path: "api_base".into(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn completes_foundry_urls_without_duplicate_paths_and_preserves_queries() {
for suffix in [
"",
"/models",
"/providers/cohere/v2",
"/providers/cohere/v2/parse",
] {
assert_eq!(
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
"https://example.com/providers/cohere/v2/parse?tenant=a"
);
}
assert_eq!(
complete_url("https://example.com/v2/parse?tenant=a").unwrap(),
"https://example.com/v2/parse?tenant=a"
);
assert!(complete_url("relative/path").is_err());
}
}

View file

@ -10,7 +10,6 @@ use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrResponseFormat};
use crate::ocr::wire::DecodedOcrResponse;
use crate::providers::azure_ai::auth::AzureAuthInputs;
use crate::url_utils::ApiUrl;
@ -32,18 +31,19 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = map_ocr_params(request)?;
let config = AzureAuthInputs::from_sourced_optional_params(
let mut config = AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
let endpoint = nonblank(request.connection.api_base.clone())
.or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV)))
.ok_or_else(|| Error::Auth("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into()))?;
let url = get_complete_url(&endpoint, &request.model, &params)?;
let body = document_intelligence::transform_ocr_request(request.document.clone())?;
transform_request_body(client, request, &url, &headers, body, |_| Ok(())).await
transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await
}
fn transform_ocr_response(
@ -61,7 +61,7 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
url: &str,
headers: &[(String, String)],
request: &LiteLLMOcrRequest,
) -> Result<DecodedOcrResponse<Self::ProviderResponse>, OcrError> {
) -> Result<crate::ocr::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError> {
polling::read_operation_response(
client.polling_http(),
response,
@ -69,6 +69,7 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
headers,
&request.connection,
request.response_format()? == OcrResponseFormat::Native,
&request.hooks,
)
.await
}

View file

@ -1,3 +1,4 @@
use std::sync::Arc;
use std::time::Duration;
use reqwest::Url;
@ -9,6 +10,7 @@ use crate::ocr::codecs::document_intelligence::{
AzureDocumentIntelligenceOperation, OperationStatus,
};
use crate::ocr::error::{OcrError, OcrPollingError, OcrResponseError};
use crate::ocr::hooks::OcrHooks;
use crate::ocr::types::OcrConnection;
use crate::ocr::wire::DecodedOcrResponse;
@ -19,24 +21,33 @@ pub(super) async fn read_operation_response(
headers: &[(String, String)],
connection: &OcrConnection,
native: bool,
hooks: &Arc<dyn OcrHooks>,
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
if response.status() != reqwest::StatusCode::ACCEPTED {
return read_json_response(response, native).await;
let bytes =
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes)
.await?;
crate::ocr::handler::post_call(hooks, &bytes).await?;
return Ok(crate::ocr::wire::decode_response(&bytes, native)?);
}
let location = response
.headers()
.get("operation-location")
.and_then(|value| value.to_str().ok())
.ok_or(OcrPollingError::PollLocation)?;
.ok_or(OcrPollingError::PollLocation)?
.to_string();
let original = Url::parse(original_url).map_err(|_| OcrPollingError::PollOrigin)?;
let operation = Url::parse(location).map_err(|_| OcrPollingError::PollOrigin)?;
let operation = Url::parse(&location).map_err(|_| OcrPollingError::PollOrigin)?;
if original.origin() != operation.origin()
|| !operation.username().is_empty()
|| operation.password().is_some()
{
return Err(OcrPollingError::PollOrigin.into());
}
poll_operation(http_client, operation, headers, connection, native).await
let bytes =
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes).await?;
crate::ocr::handler::post_call(hooks, &bytes).await?;
poll_operation(http_client, operation, headers, connection, native, hooks).await
}
async fn poll_operation(
@ -45,6 +56,7 @@ async fn poll_operation(
headers: &[(String, String)],
connection: &OcrConnection,
native: bool,
hooks: &Arc<dyn OcrHooks>,
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
let deadline = Instant::now()
.checked_add(connection.poll_timeout)
@ -75,12 +87,19 @@ async fn poll_operation(
.max(1);
let decoded = tokio::time::timeout_at(
deadline,
read_json_response::<AzureDocumentIntelligenceOperation>(response, native),
read_json_response::<AzureDocumentIntelligenceOperation>(
response,
native,
connection.max_response_bytes,
),
)
.await
.map_err(|_| OcrPollingError::PollTimeout)??;
match &decoded.data.status {
Some(OperationStatus::Succeeded) => return Ok(decoded),
Some(OperationStatus::Succeeded) => {
crate::ocr::handler::post_call(hooks, decoded.text.as_bytes()).await?;
return Ok(decoded);
}
Some(OperationStatus::Running | OperationStatus::NotStarted) => {
tokio::time::timeout_at(deadline, tokio::time::sleep(Duration::from_secs(retry)))
.await

View file

@ -33,13 +33,16 @@ impl OcrAdapter for AzureMistralAdapter {
known: params,
extra_params: _extra_params,
} = _prepare_ocr_request::<MistralOcrParams>(request)?;
let config = AzureAuthInputs::from_sourced_optional_params(
let mut config = AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
let url = get_complete_url(request.connection.api_base.as_deref(), &credential_env)?;
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
let retains_document = !request.document.source().starts_with("http://")
&& !request.document.source().starts_with("https://");
let document = inline_remote_document(
client.document_fetcher(),
request.document.clone(),
@ -47,9 +50,15 @@ impl OcrAdapter for AzureMistralAdapter {
)
.await?;
let body = mistral::transform_ocr_request(&request.model, document, &params)?;
transform_request_body(client, request, &url, &headers, body, |body| {
validate_inline_document(&body.document)
})
transform_request_body(
client,
request,
&url,
&headers,
retains_document,
body,
|body| validate_inline_document(&body.document),
)
.await
}
@ -83,12 +92,15 @@ fn get_complete_url(
})
}
async fn validate_environment(
pub(in crate::ocr::adapters) async fn validate_environment(
connection: &OcrConnection,
config: &AzureAuthInputs,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
if config.azure_ad_token_provider.is_some() {
super::resolve_entra(config, env_lookup).await?;
}
super::validate_destination(connection, connection.extra_headers_source)?;
return Ok(connection.extra_headers.clone());
}

View file

@ -1,3 +1,4 @@
mod cohere;
mod document_intelligence;
mod mistral;
@ -10,8 +11,10 @@ use crate::ocr::error::OcrError;
use crate::ocr::types::OcrConnection;
use crate::providers::azure_ai::auth::{AzureAuthInputs, AzureAuthService};
pub(crate) use cohere::AzureCohereAdapter;
pub(crate) use document_intelligence::AzureDocumentIntelligenceAdapter;
pub(crate) use mistral::AzureMistralAdapter;
pub(super) use mistral::validate_environment as validate_ai_environment;
async fn resolve_entra(
config: &AzureAuthInputs,
@ -22,6 +25,10 @@ async fn resolve_entra(
.get_or_init(AzureAuthService::default)
.get_azure_ad_token(config, env_lookup)
.await
.or_else(|error| match error {
crate::AuthError::EmptyAzureToken => Ok(None),
other => Err(other),
})
.map(|credential| {
credential.map(|credential| {
let source = credential.source();

View file

@ -0,0 +1,123 @@
use super::OcrAdapter;
use crate::Error;
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
use crate::ocr::OcrClient;
use crate::ocr::codecs::cohere::{
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
};
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
use crate::url_utils::ApiUrl;
pub(crate) struct CohereAdapter;
impl OcrAdapter for CohereAdapter {
type ProviderResponse = CohereResponse;
const PROVIDER: OcrProvider = OcrProvider::Cohere;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = super::super::wire::decode_request_value::<CohereParams>(
serde_json::Value::Object(request.optional_params.clone()),
"optional_params",
)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = complete_url(
request
.connection
.api_base
.as_deref()
.unwrap_or(COHERE_PARSE_API_BASE),
)?;
let body = transform_request(&request.model, request.document.clone(), params)?;
transform_request_body(client, request, &url, &headers, true, body, |body| {
validate_document(&body.document)
})
.await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_response(&request.model, response)
}
}
fn complete_url(base: &str) -> Result<String, OcrError> {
let parsed = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
if !matches!(parsed.scheme(), "http" | "https") {
return Err(invalid_api_base().into());
}
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v2", "parse"]))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base().into())
}
fn invalid_api_base() -> OcrRequestError {
OcrRequestError::RequestField {
path: "api_base".into(),
}
}
fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(COHERE_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.ok_or_else(|| {
Error::Auth("Missing COHERE_API_KEY - set it in the environment or pass api_key".into())
})?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn completes_provider_urls_without_duplicate_paths_and_preserves_queries() {
for suffix in ["", "/v2", "/v2/parse"] {
assert_eq!(
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
"https://example.com/v2/parse?tenant=a"
);
}
}
#[test]
fn rejects_invalid_urls_and_blank_keys() {
assert!(complete_url("relative/path").is_err());
assert!(complete_url("ftp://example.com").is_err());
assert!(matches!(
validate_environment(
&OcrConnection {
api_key: Some(" ".into()),
..Default::default()
},
&|_| None,
),
Err(OcrError::Public(Error::Auth(_)))
));
}
}

View file

@ -33,7 +33,7 @@ impl OcrAdapter for MistralAdapter {
let url = get_complete_url(request.connection.api_base.as_deref())?;
let body =
mistral::transform_ocr_request(&request.model, request.document.clone(), &params)?;
transform_request_body(client, request, &url, &headers, body, |_| Ok(())).await
transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await
}
fn transform_ocr_response(

View file

@ -5,15 +5,16 @@ use serde::de::DeserializeOwned;
use super::OcrClient;
use super::error::{OcrError, OcrResponseError};
use super::registry::OcrProvider;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrResponseFormat};
use super::wire::DecodedOcrResponse;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
mod azure;
mod cohere;
mod mistral;
mod reducto;
mod vertex;
pub(crate) use azure::{AzureDocumentIntelligenceAdapter, AzureMistralAdapter};
pub(crate) use azure::{AzureCohereAdapter, AzureDocumentIntelligenceAdapter, AzureMistralAdapter};
pub(crate) use cohere::CohereAdapter;
pub(crate) use mistral::MistralAdapter;
pub(crate) use reducto::{ReductoLegacyAdapter, ReductoV3Adapter};
pub(crate) use vertex::{VertexDeepSeekAdapter, VertexMistralAdapter};
@ -55,18 +56,27 @@ pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
_url: &str,
_headers: &[(String, String)],
request: &LiteLLMOcrRequest,
) -> impl Future<Output = Result<DecodedOcrResponse<Self::ProviderResponse>, OcrError>> + Send
{
let retain_native = request
.response_format()
.map(|format| format == OcrResponseFormat::Native);
async move { super::client::read_json_response(response, retain_native?).await }
) -> impl Future<
Output = Result<super::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError>,
> + Send {
async move {
let bytes =
super::client::read_response_bytes(response, request.connection.max_response_bytes)
.await?;
super::handler::post_call(&request.hooks, &bytes).await?;
Ok(super::wire::decode_response(
&bytes,
request.response_format()? == super::types::OcrResponseFormat::Native,
)?)
}
}
}
macro_rules! for_each_ocr_adapter {
($callback:ident) => {
$callback! {
Cohere, $crate::ocr::adapters::CohereAdapter, $crate::ocr::adapters::CohereAdapter, Cohere;
AzureCohere, $crate::ocr::adapters::AzureCohereAdapter, $crate::ocr::adapters::AzureCohereAdapter, AzureAi;
Mistral, $crate::ocr::adapters::MistralAdapter, $crate::ocr::adapters::MistralAdapter, Mistral;
AzureMistral, $crate::ocr::adapters::AzureMistralAdapter, $crate::ocr::adapters::AzureMistralAdapter, AzureAi;
AzureDocumentIntelligence, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, AzureAi;

View file

@ -27,7 +27,7 @@ impl OcrAdapter for ReductoLegacyAdapter {
} = _prepare_ocr_request::<ReductoLegacyParams>(request)?;
let headers = super::validate_environment(&request.connection, &credential_env)?;
let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let document = guardrail_document(request, &url).await?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document =
super::prepare_document(client, document, &request.connection, &headers).await?;
let body = reducto::transform_legacy_ocr_request(&request.model, document, &params)?;

View file

@ -93,7 +93,7 @@ pub(super) async fn prepare_document(
.map_err(crate::error::TransportError::from)?;
let uploaded = crate::ocr::client::read_json_response::<
crate::ocr::codecs::reducto::ReductoUploadResponse,
>(response, false)
>(response, false, connection.max_response_bytes)
.await?
.data;
let file_id = uploaded

View file

@ -27,7 +27,7 @@ impl OcrAdapter for ReductoV3Adapter {
} = _prepare_ocr_request::<ReductoV3Params>(request)?;
let headers = super::validate_environment(&request.connection, &credential_env)?;
let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let document = guardrail_document(request, &url).await?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document =
super::prepare_document(client, document, &request.connection, &headers).await?;
let body = reducto::transform_v3_ocr_request(&request.model, document, &params)?;

View file

@ -57,9 +57,15 @@ impl OcrAdapter for VertexDeepSeekAdapter {
let document = request.document.clone();
let body =
deepseek::transform_ocr_request(&provider_model(&request.model), document, &params)?;
transform_request_body(client, request, &url, &authentication.headers, body, |_| {
Ok(())
})
transform_request_body(
client,
request,
&url,
&authentication.headers,
false,
body,
|_| Ok(()),
)
.await
}

View file

@ -54,6 +54,8 @@ impl OcrAdapter for VertexMistralAdapter {
&location,
&request.model,
)?;
let retains_document = !request.document.source().starts_with("http://")
&& !request.document.source().starts_with("https://");
let document = inline_remote_document(
client.document_fetcher(),
request.document.clone(),
@ -66,6 +68,7 @@ impl OcrAdapter for VertexMistralAdapter {
request,
&url,
&authentication.headers,
retains_document,
body,
|body| validate_inline_document(&body.document),
)

View file

@ -1,10 +1,10 @@
use std::sync::OnceLock;
use std::time::Duration;
use bytes::{Bytes, BytesMut};
use serde::de::DeserializeOwned;
use super::error::OcrError;
use super::handler::perform_ocr_request;
use super::error::{OcrError, OcrResponseError};
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use super::wire::{DecodedOcrResponse, decode_response};
use crate::Error;
@ -32,6 +32,10 @@ impl OcrClient {
})
}
pub fn shared() -> Result<Self, Error> {
shared_client()
}
#[tracing::instrument(
name = "ocr",
target = "litellm::function_trace",
@ -39,7 +43,34 @@ impl OcrClient {
skip_all
)]
pub async fn perform(&self, request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
perform_ocr_request(self, request).await
use super::{
NativeOutcome, OcrAdmission, OcrCall, OcrCallStep, OcrHookHost, OcrHost,
OcrHostOperation, OcrHostResult,
};
let host = OcrHookHost::new(request.hooks.clone());
let mut request = Some(request);
let NativeOutcome::Completed(mut call) = OcrCall::admit(self.clone(), OcrAdmission::all())
else {
return Err(Error::InvalidRequest(
"native OCR host admission declined".into(),
));
};
let mut result = None;
loop {
match call.resume(result.take()).await? {
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().ok_or_else(|| {
Error::InvalidRequest("OCR request was already projected".into())
})?),
false,
))))
}
OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await),
OcrCallStep::Complete(response) => return Ok(response),
}
}
}
pub(crate) fn provider_http(&self) -> &reqwest::Client {
@ -77,7 +108,7 @@ fn no_redirect_http() -> Result<reqwest::Client, TransportError> {
.map_err(TransportError::from)
}
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
pub(crate) fn shared_client() -> Result<OcrClient, Error> {
static CLIENT: OnceLock<Result<OcrClient, TransportError>> = OnceLock::new();
let client = CLIENT
.get_or_init(|| {
@ -88,18 +119,50 @@ pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error
.and_then(OcrClient::new)
})
.clone()?;
client.perform(request).await
Ok(client)
}
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
shared_client()?.perform(request).await
}
pub async fn read_json_response<T: DeserializeOwned>(
response: reqwest::Response,
native: bool,
max_response_bytes: usize,
) -> Result<DecodedOcrResponse<T>, OcrError> {
let bytes = read_response_bytes(response, max_response_bytes).await?;
Ok(decode_response(&bytes, native)?)
}
pub(crate) async fn read_response_bytes(
mut response: reqwest::Response,
max_response_bytes: usize,
) -> Result<Bytes, OcrError> {
let status = response.status();
let bytes = response
.bytes()
.await
.map_err(crate::error::TransportError::from)?;
let limit = if status.is_success() {
max_response_bytes
} else {
max_response_bytes.min(4 * (crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS + 1))
};
if status.is_success()
&& response
.content_length()
.is_some_and(|length| length > limit as u64)
{
return Err(OcrResponseError::TooLarge { limit }.into());
}
let mut bytes = BytesMut::new();
while let Some(chunk) = response.chunk().await.map_err(transport_error)? {
let remaining = limit.saturating_sub(bytes.len());
if status.is_success() && chunk.len() > remaining {
return Err(OcrResponseError::TooLarge { limit }.into());
}
bytes.extend_from_slice(&chunk[..chunk.len().min(remaining)]);
if !status.is_success() && bytes.len() == limit {
break;
}
}
if !status.is_success() {
return Err(crate::error::TransportError::Http {
status: status.as_u16(),
@ -107,5 +170,41 @@ pub async fn read_json_response<T: DeserializeOwned>(
}
.into());
}
Ok(decode_response(&bytes, native)?)
Ok(bytes.freeze())
}
pub(crate) fn transport_error(error: reqwest::Error) -> Error {
if error.is_timeout() {
return Error::Http {
status: 408,
body: "OCR request timed out".into(),
};
}
crate::error::TransportError::from(error).into()
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn request_timeout_has_an_http_408_status() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let _connection = listener.accept().await.unwrap();
tokio::time::sleep(Duration::from_secs(1)).await;
});
let error = reqwest::Client::new()
.get(format!("http://{address}"))
.timeout(Duration::from_millis(10))
.send()
.await
.unwrap_err();
assert!(matches!(
transport_error(error),
Error::Http { status: 408, .. }
));
server.abort();
}
}

View file

@ -0,0 +1,254 @@
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value, json};
use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum OutputFormat {
#[default]
Markdown,
Blocks,
}
#[derive(Deserialize)]
pub(crate) struct CohereParams {
#[serde(default)]
pub output_format: OutputFormat,
}
#[derive(Deserialize, Serialize)]
pub(crate) struct CohereRequest {
pub model: String,
pub document: OcrDocument,
pub output_format: OutputFormat,
}
pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), OcrRequestError> {
let OcrDocument::ImageUrl { image_url, .. } = document else {
return Err(OcrRequestError::CohereImageOnly);
};
if image_url.is_empty() {
return Err(OcrRequestError::CohereImageOnly);
}
if let Some(inline) = InlineDocument::parse(image_url)? {
if !inline.mime_type().type_.eq_ignore_ascii_case("image") {
return Err(OcrRequestError::CohereImageOnly);
}
inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
}
Ok(())
}
#[derive(Deserialize)]
pub(crate) struct CohereResponse {
#[serde(default)]
pages: Vec<CoherePage>,
meta: Option<CohereMeta>,
}
#[derive(Deserialize)]
struct CoherePage {
index: Option<i64>,
markdown: Option<CohereMarkdown>,
blocks: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMarkdown {
#[serde(default)]
content: String,
images: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMeta {
billed_units: Option<CohereBilledUnits>,
}
#[derive(Deserialize)]
struct CohereBilledUnits {
pages: Option<i64>,
}
pub(crate) fn transform_response(
model: &str,
response: CohereResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
let pages_processed = response
.meta
.and_then(|meta| meta.billed_units)
.and_then(|units| units.pages)
.map(Ok)
.unwrap_or_else(|| {
i64::try_from(response.pages.len()).map_err(|_| OcrResponseError::NumericRange("pages"))
})?;
let pages = response
.pages
.into_iter()
.enumerate()
.map(|(position, page)| {
let index = page.index.map(Ok).unwrap_or_else(|| {
i64::try_from(position).map_err(|_| OcrResponseError::NumericRange("page index"))
})?;
let (content, images) = page
.markdown
.map(|markdown| {
let images =
markdown
.images
.filter(|images| !images.is_empty())
.map(|images| {
images
.into_iter()
.map(|mut image| {
if let Some(Value::Object(bbox)) =
image.get("bounding_box").cloned()
{
image.insert("bbox".into(), Value::Object(bbox));
}
Value::Object(image)
})
.collect::<Vec<_>>()
});
(markdown.content, images)
})
.unwrap_or_default();
let mut normalized = json!({"index": index, "markdown": content, "images": images});
if let Some(blocks) = page.blocks {
normalized["blocks"] = json!(blocks);
}
Ok(normalized)
})
.collect::<Result<Vec<_>, OcrResponseError>>()?;
Ok(LiteLLMOcrResponse {
pages,
model: model.into(),
document_annotation: None,
usage_info: Some(json!({"pages_processed": pages_processed})),
object: "ocr".into(),
extra_fields: Map::new(),
provider_native_response: None,
})
}
pub(crate) fn transform_request(
model: &str,
document: OcrDocument,
params: CohereParams,
) -> Result<CohereRequest, OcrRequestError> {
validate_document(&document)?;
Ok(CohereRequest {
model: model.into(),
document,
output_format: params.output_format,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn response_normalizes_markdown_images_blocks_and_billed_pages() {
let response = serde_json::from_value(json!({
"pages": [
{
"type":"markdown",
"index":4,
"markdown":{
"content":"receipt",
"images":[{
"id":"image",
"bounding_box":{"top_left_x":1,"bottom_right_x":48},
"bounding_box_normalized":{"top_left_x":0.04,"bottom_right_x":0.15},
"description":"scan",
"category":"logo"
}]
}
},
{"type":"blocks","blocks":[{"type":"text","text":{"content":"total"}}]}
],
"meta":{"api_version":{"version":"2"},"billed_units":{"pages":3}}
}))
.unwrap();
let normalized = transform_response("parse-v5.0", response).unwrap();
assert_eq!(normalized.pages[0]["index"], 4);
assert_eq!(normalized.pages[0]["markdown"], "receipt");
assert_eq!(normalized.pages[0]["images"][0]["bbox"]["top_left_x"], 1);
assert_eq!(
normalized.pages[0]["images"][0]["bounding_box_normalized"]["bottom_right_x"],
0.15
);
assert_eq!(normalized.pages[0]["images"][0]["description"], "scan");
assert_eq!(normalized.pages[0]["images"][0]["category"], "logo");
assert_eq!(normalized.pages[1]["index"], 1);
assert_eq!(normalized.pages[1]["markdown"], "");
assert_eq!(normalized.pages[1]["blocks"][0]["text"]["content"], "total");
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 3);
}
#[test]
fn response_defaults_and_invalid_fields() {
for value in [
json!({}),
json!({"meta":null}),
json!({"pages":[],"meta":{"billed_units":null}}),
] {
let normalized =
transform_response("parse", serde_json::from_value(value).unwrap()).unwrap();
assert!(normalized.pages.is_empty());
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 0);
}
for value in [
json!({"pages":null}),
json!({"pages":[{"markdown":"text"}]}),
json!({"pages":[{"index":"bad"}]}),
] {
assert!(serde_json::from_value::<CohereResponse>(value).is_err());
}
let normalized = transform_response(
"parse",
serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(),
)
.unwrap();
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 1);
assert!(normalized.pages[0]["images"].is_null());
}
#[test]
fn request_requires_image_and_supported_output_format() {
for value in [
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
json!({"type":"image_url","image_url":""}),
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
] {
assert_eq!(
validate_document(&serde_json::from_value(value).unwrap()),
Err(OcrRequestError::CohereImageOnly)
);
}
assert!(serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err());
for format in ["markdown", "blocks"] {
assert!(
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
);
}
let request = transform_request(
"parse-v5.0",
serde_json::from_value(json!({
"type":"image_url",
"image_url":"https://example.com/image.png"
}))
.unwrap(),
serde_json::from_value(json!({})).unwrap(),
)
.unwrap();
assert_eq!(
serde_json::to_value(request).unwrap()["output_format"],
"markdown"
);
}
}

View file

@ -12,13 +12,17 @@ pub(crate) fn transform_ocr_request(
params: &DeepSeekOcrParams,
) -> Result<DeepSeekOcrRequest, OcrRequestError> {
if document.source().is_empty() {
return Err(OcrRequestError::MissingField("document URL"));
return Err(OcrRequestError::MissingDocumentUrl);
}
let content = OcrDocument::ImageUrl {
image_url: document.source().to_string(),
extra_fields: serde_json::Map::new(),
};
Ok(DeepSeekOcrRequest {
model: provider_model.to_string(),
messages: vec![DeepSeekOcrMessage {
role: UserRole::User,
content: vec![document],
content: vec![content],
}],
params: params.clone(),
})

View file

@ -163,6 +163,30 @@ mod tests {
);
}
#[rstest]
#[case(json!([0, 1, 2]), Some("1,2,3"))]
#[case(json!([2, 0, 0, 1]), Some("1,2,3"))]
#[case(json!([]), None)]
#[case(json!("3-9"), Some("3-9"))]
#[case(json!("1-3, 5"), Some("1-3,5"))]
#[case(json!(["1", "3-5"]), Some("1,3-5"))]
fn page_mapping_matches_python(#[case] input: Value, #[case] expected: Option<&str>) {
assert_eq!(
map(json!({"pages": input})).unwrap().pages.as_deref(),
expected
);
}
#[rstest]
#[case(json!("a,b"))]
#[case(json!([-1]))]
#[case(json!([true, false]))]
#[case(json!([1, "2"]))]
#[case(json!(5))]
fn invalid_page_mapping_matches_python(#[case] input: Value) {
assert!(map(json!({"pages": input})).is_err());
}
#[rstest]
#[case(json!(["keyValuePairs"]), "keyValuePairs")]
#[case(json!(["keyValuePairs", "languages"]), "keyValuePairs,languages")]

View file

@ -13,7 +13,7 @@ pub(crate) fn transform_ocr_request(
) -> Result<DocumentIntelligenceRequest, OcrRequestError> {
let source = document.source();
if source.is_empty() {
return Err(OcrRequestError::MissingField("document URL"));
return Err(OcrRequestError::MissingDocumentUrl);
}
Ok(if let Some(document) = InlineDocument::parse(source)? {
DocumentIntelligenceRequest::Base64Source(
@ -46,10 +46,7 @@ pub(crate) fn transform_ocr_response(
let mut extra_fields = Map::new();
extra_fields.insert("content".into(), option_value(result.content));
extra_fields.insert("tables".into(), option_value(result.tables));
extra_fields.insert(
"key_value_pairs".into(),
option_value(result.key_value_pairs),
);
extra_fields.insert("keyValuePairs".into(), option_value(result.key_value_pairs));
Ok(LiteLLMOcrResponse {
pages,
model: model.into(),

View file

@ -114,6 +114,7 @@ mod tests {
#[rstest]
#[case("table_format", json!("html"))]
#[case("confidence_scores_granularity", json!("word"))]
#[case("confidence_scores_granularity", json!("block"))]
#[case("document_annotation_prompt", json!("extract"))]
#[case("include_blocks", json!(true))]
#[case("id", json!("req-123"))]
@ -133,6 +134,7 @@ mod tests {
#[rstest]
#[case("pages", json!([0, 2]))]
#[case("pages", json!("0,2-4"))]
#[case("include_image_base64", json!(true))]
#[case("image_limit", json!(2))]
#[case("image_min_size", json!(100))]
@ -196,8 +198,16 @@ mod tests {
#[rstest]
fn transform_ocr_response_preserves_blocks_and_confidence_scores() {
let response: MistralOcrResponse = serde_json::from_value(json!({
"pages":[{"index":0,"markdown":"hello","blocks":[{"type":"title"}],"confidence_scores":{"mean":0.99}}],
"pages":[{
"index":0,
"markdown":"hello",
"images":[{"id":"img-0","image_base64":"data:image/png;base64,AA=="}],
"dimensions":{"width":612,"height":792,"dpi":72},
"blocks":[{"type":"title","bbox":{"x":1},"confidence_scores":{"mean":0.98}}],
"confidence_scores":{"average_page_confidence_score":0.99,"minimum_page_confidence_score":0.97}
}],
"model":"returned-model",
"document_annotation":"{\"language\":\"en\"}",
"usage_info":{"pages_processed":1}
}))
.unwrap();
@ -205,7 +215,20 @@ mod tests {
.unwrap()
.into_json();
assert_eq!(result["pages"][0]["blocks"][0]["type"], "title");
assert_eq!(result["pages"][0]["confidence_scores"]["mean"], 0.99);
assert_eq!(result["pages"][0]["blocks"][0]["bbox"]["x"], 1);
assert_eq!(
result["pages"][0]["blocks"][0]["confidence_scores"]["mean"],
0.98
);
assert_eq!(
result["pages"][0]["confidence_scores"]["average_page_confidence_score"],
0.99
);
assert_eq!(result["pages"][0]["images"][0]["id"], "img-0");
assert_eq!(result["pages"][0]["dimensions"]["dpi"], 72);
assert_eq!(result["model"], "returned-model");
assert_eq!(result["document_annotation"], "{\"language\":\"en\"}");
assert_eq!(result["usage_info"]["pages_processed"], 1);
}
#[rstest]

View file

@ -3,10 +3,17 @@ use serde_json::{Map, Value};
use crate::ocr::types::OcrDocument;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub(crate) enum MistralOcrPages {
Range(String),
Indices(Vec<i64>),
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub(crate) struct MistralOcrParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub pages: Option<Vec<i64>>,
pub pages: Option<MistralOcrPages>,
#[serde(skip_serializing_if = "Option::is_none")]
pub include_image_base64: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]

View file

@ -1,3 +1,4 @@
pub(crate) mod cohere;
pub(crate) mod deepseek;
pub(crate) mod document_intelligence;
pub(crate) mod mistral;

View file

@ -2,13 +2,90 @@ use base64::{Engine, engine::general_purpose::STANDARD};
use data_url::mime::Mime;
use data_url::{DataUrl, DataUrlError, forgiving_base64::DecodeError};
use reqwest::Url;
use serde_json::Map;
use super::error::{OcrError, OcrRequestError, OcrResponseError};
use super::types::{OcrConnection, OcrDocument};
use crate::constants::OCR_MAX_FETCH_REDIRECTS;
use crate::constants::{OCR_INLINE_MAX_BYTES, OCR_MAX_FETCH_REDIRECTS};
use crate::error::{MediaError, TransportError};
use crate::media::{DownloadPolicy, MediaFetcher};
pub fn encode_file_document(
bytes: &[u8],
file_name: Option<&str>,
mime_type: Option<&str>,
) -> Result<OcrDocument, OcrRequestError> {
if bytes.is_empty() {
return Err(OcrRequestError::EmptyFile);
}
if bytes.len() > OCR_INLINE_MAX_BYTES {
return Err(OcrRequestError::InlineDocumentTooLarge);
}
if let Some(value) = mime_type
&& !valid_mime_type(value)
{
return Err(OcrRequestError::InvalidMimeType(value.into()));
}
let mime_type = mime_type
.map(str::to_string)
.or_else(|| file_name.map(|name| mime_type_for_name(name).to_string()))
.unwrap_or_else(|| "application/octet-stream".into());
let source = format!("data:{mime_type};base64,{}", STANDARD.encode(bytes));
Ok(if mime_type.starts_with("image/") {
OcrDocument::ImageUrl {
image_url: source,
extra_fields: Map::new(),
}
} else {
OcrDocument::DocumentUrl {
document_url: source,
extra_fields: Map::new(),
}
})
}
fn valid_mime_type(value: &str) -> bool {
let Some((kind, subtype)) = value.split_once('/') else {
return false;
};
!kind.is_empty()
&& !subtype.is_empty()
&& kind.chars().chain(subtype.chars()).all(|character| {
character.is_alphanumeric() || matches!(character, '.' | '+' | '-' | '_')
})
}
pub fn mime_type_for_name(name: &str) -> &'static str {
let extension = std::path::Path::new(name)
.extension()
.and_then(|value| value.to_str())
.unwrap_or_default();
match extension.to_ascii_lowercase().as_str() {
"pdf" => "application/pdf",
"png" => "image/png",
"jpg" | "jpeg" => "image/jpeg",
"gif" => "image/gif",
"webp" => "image/webp",
"tiff" | "tif" => "image/tiff",
"bmp" => "image/bmp",
_ => mime_guess::from_path(name)
.first_raw()
.unwrap_or("application/octet-stream"),
}
}
pub fn upload_mime_type<'a>(file_name: Option<&str>, content_type: Option<&'a str>) -> &'a str {
match content_type
.and_then(|value| value.split(';').next())
.map(str::trim)
{
Some(value) if !value.is_empty() && value != "application/octet-stream" => value,
_ => file_name
.map(mime_type_for_name)
.unwrap_or("application/octet-stream"),
}
}
pub(crate) struct InlineDocument<'a>(DataUrl<'a>);
impl<'a> InlineDocument<'a> {
@ -95,9 +172,11 @@ fn map_media_error(error: MediaError) -> OcrError {
body: "OCR document download failed".into(),
}
.into(),
MediaError::Timeout => {
TransportError::Network("OCR document download timed out".into()).into()
MediaError::Timeout => TransportError::Http {
status: 408,
body: "OCR document download timed out".into(),
}
.into(),
MediaError::Transport(error) => error.into(),
}
}
@ -114,6 +193,90 @@ mod tests {
}
}
#[test]
fn file_bytes_are_encoded_with_core_owned_mime_policy() {
assert_eq!(
encode_file_document(b"abc", Some("scan.png"), None).unwrap(),
OcrDocument::ImageUrl {
image_url: "data:image/png;base64,YWJj".into(),
extra_fields: Map::new(),
}
);
assert_eq!(
encode_file_document(b"abc", None, Some("application/pdf")).unwrap(),
document("data:application/pdf;base64,YWJj")
);
}
#[test]
fn file_name_mime_mapping_matches_python() {
for (name, expected) in [
("document.pdf", "application/pdf"),
("image.png", "image/png"),
("photo.jpg", "image/jpeg"),
("photo.jpeg", "image/jpeg"),
("animation.gif", "image/gif"),
("image.webp", "image/webp"),
("scan.tiff", "image/tiff"),
("scan.tif", "image/tiff"),
("bitmap.bmp", "image/bmp"),
("DOCUMENT.PDF", "application/pdf"),
("IMAGE.PNG", "image/png"),
("file.unknown-extension", "application/octet-stream"),
] {
assert_eq!(mime_type_for_name(name), expected);
}
}
#[test]
fn upload_mime_mapping_matches_python() {
assert_eq!(
upload_mime_type(Some("report.pdf"), Some("application/octet-stream")),
"application/pdf"
);
assert_eq!(upload_mime_type(Some("image.png"), None), "image/png");
assert_eq!(upload_mime_type(None, None), "application/octet-stream");
assert_eq!(
upload_mime_type(Some("doc.pdf"), Some("application/pdf; charset=utf-8")),
"application/pdf"
);
assert_eq!(
upload_mime_type(
Some("img.png"),
Some("image/png; charset=utf-8; boundary=something")
),
"image/png"
);
}
#[test]
fn file_encoding_enforces_decoded_size_limit() {
let bytes = vec![b'a'; OCR_INLINE_MAX_BYTES + 1];
assert_eq!(
encode_file_document(&bytes, None, None),
Err(OcrRequestError::InlineDocumentTooLarge)
);
let document = encode_file_document(&bytes[..OCR_INLINE_MAX_BYTES], None, None).unwrap();
let inline = InlineDocument::parse(document.source()).unwrap().unwrap();
assert_eq!(
inline.decode(OCR_INLINE_MAX_BYTES).unwrap(),
bytes[..OCR_INLINE_MAX_BYTES]
);
}
#[test]
fn file_encoding_rejects_empty_bytes_and_invalid_explicit_mime() {
assert!(encode_file_document(b"", None, None).is_err());
for mime in [
"text/plain;bad",
"text/plain/extra",
" text/plain",
"text/plain\n",
] {
assert!(encode_file_document(b"abc", None, Some(mime)).is_err());
}
}
#[test]
fn decodes_data_urls_and_limits_decoded_size() {
for (source, expected) in [

View file

@ -4,15 +4,27 @@ use crate::error::TransportError;
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum OcrRequestError {
#[error("File is empty or could not be read")]
EmptyFile,
#[error("Invalid MIME type: {0}")]
InvalidMimeType(String),
#[error(
"Cohere Parse only accepts `image_url` documents; document_url and PDF inputs are not supported"
)]
CohereImageOnly,
#[error("Invalid `req_format`. Expected 'native' or 'litellm'.")]
RequestFormat,
#[error("invalid OCR request field: {path}")]
RequestField { path: String },
#[error("missing required field: {0}")]
MissingField(&'static str),
#[error("Document URL is required")]
MissingDocumentUrl,
#[error("invalid OCR document data URI")]
InvalidDataUri,
#[error("Reducto requires a reducto:// id or a data URI")]
#[error(
"Reducto requires a reducto:// id or a data URI; plain HTTP URLs are not supported, upload the file first"
)]
ReductoSource,
#[error("inline OCR document exceeds the size limit")]
InlineDocumentTooLarge,
@ -34,6 +46,8 @@ pub enum OcrRequestError {
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum OcrResponseError {
#[error("OCR response exceeds the size limit of {limit} bytes")]
TooLarge { limit: usize },
#[error("invalid OCR response field: {path}")]
ResponseField { path: String },
#[error("OCR response is missing non-empty content")]

View file

@ -1,15 +1,17 @@
use super::OcrClient;
use super::adapters::OcrAdapter;
use super::hooks::OcrLifecycleHooks;
use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest};
use super::registry::OcrAdapterKind;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use crate::Error;
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
use std::sync::Arc;
pub(crate) async fn perform_ocr_request(
client: &OcrClient,
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, Error> {
request.response_format()?;
let context = CallLifecycleContext::new(
"ocr",
request.model.clone(),
@ -23,27 +25,71 @@ pub(crate) async fn perform_ocr_request(
hooks: request.hooks.clone(),
provider_name: context.custom_llm_provider.clone(),
};
CallLifecycle::default().run(context, request, &hooks, |request| async move {
macro_rules! execute_selected_adapter {
CallLifecycle::default()
.run(context, request, &hooks, |request| async move {
PreparedOcrCall::prepare(client.clone(), request)
.await?
.execute()
.await?
.normalize()
})
.await
}
pub(crate) struct PreparedOcrCall {
client: OcrClient,
request: LiteLLMOcrRequest,
http: reqwest::Request,
}
impl PreparedOcrCall {
pub(crate) async fn prepare(
client: OcrClient,
request: LiteLLMOcrRequest,
) -> Result<Self, Error> {
macro_rules! prepare_adapter {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
match request.adapter {
$( OcrAdapterKind::$variant => execute_ocr_provider_call(client, &$instance, request).await, )+
$( OcrAdapterKind::$variant => $instance.prepare_request(&request, &client).await?, )+
}
};
}
super::adapters::for_each_ocr_adapter!(execute_selected_adapter)
}).await
let http = super::adapters::for_each_ocr_adapter!(prepare_adapter);
Ok(Self {
client,
request,
http,
})
}
pub(crate) async fn execute(self) -> Result<OcrProviderResponse, Error> {
let url = self.http.url().to_string();
let headers = request_headers(&self.http)?;
let response = crate::http_utils::http_request(reqwest::RequestBuilder::from_parts(
self.client.provider_http().clone(),
self.http,
))
.await
.map_err(super::client::transport_error)?;
macro_rules! read_adapter {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
match self.request.adapter {
$( OcrAdapterKind::$variant => {
let decoded = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?;
Ok(OcrProviderResponse {
request: self.request,
data: OcrProviderData::$variant(decoded),
})
}, )+
}
};
}
super::adapters::for_each_ocr_adapter!(read_adapter)
}
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn execute_ocr_provider_call<A: OcrAdapter>(
client: &OcrClient,
adapter: &A,
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, Error> {
let provider_request = adapter.prepare_request(&request, client).await?;
let url = provider_request.url().to_string();
let headers = provider_request
fn request_headers(request: &reqwest::Request) -> Result<Vec<(String, String)>, Error> {
request
.headers()
.iter()
.map(|(name, value)| {
@ -53,20 +99,41 @@ async fn execute_ocr_provider_call<A: OcrAdapter>(
.map_err(|_| super::error::OcrRequestError::RequestField {
path: "headers".into(),
})
.map_err(Error::from)
})
.collect::<Result<Vec<_>, _>>()?;
let response = crate::http_utils::http_request(reqwest::RequestBuilder::from_parts(
client.provider_http().clone(),
provider_request,
))
.await
.map_err(crate::error::TransportError::from)?;
let decoded = adapter
.read_response(client, response, &url, &headers, &request)
.await?;
let response = adapter.transform_ocr_response(&request, decoded.data)?;
Ok(LiteLLMOcrResponse {
provider_native_response: decoded.native,
..response
})
.collect()
}
macro_rules! provider_data {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
enum OcrProviderData {
$( $variant(super::wire::DecodedOcrResponse<<$adapter as OcrAdapter>::ProviderResponse>), )+
}
impl OcrProviderResponse {
pub(crate) fn normalize(self) -> Result<LiteLLMOcrResponse, Error> {
match self.data {
$( OcrProviderData::$variant(decoded) => {
let response = $instance.transform_ocr_response(&self.request, decoded.data)?;
Ok(LiteLLMOcrResponse { provider_native_response: decoded.native, ..response })
}, )+
}
}
}
};
}
pub(crate) struct OcrProviderResponse {
request: LiteLLMOcrRequest,
data: OcrProviderData,
}
pub(crate) async fn post_call(hooks: &Arc<dyn OcrHooks>, bytes: &[u8]) -> Result<(), Error> {
let original_response = serde_json::Value::String(String::from_utf8_lossy(bytes).into_owned());
hooks
.post_call(OcrPostCallRequest { original_response })
.await?;
Ok(())
}
super::adapters::for_each_ocr_adapter!(provider_data);

View file

@ -24,11 +24,19 @@ pub struct OcrDuringCallRequest {
pub model: String,
pub custom_llm_provider: String,
pub url: String,
pub headers: Vec<(String, String)>,
pub body: Value,
#[serde(skip)]
pub retained_fields: Vec<String>,
}
#[derive(Clone, Debug, Serialize)]
pub struct OcrPostCallRequest {
pub original_response: Value,
}
pub trait OcrHooks: Send + Sync {
fn has_guardrails(&self) -> bool {
fn intercepts_requests(&self) -> bool {
false
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
@ -40,6 +48,9 @@ pub trait OcrHooks: Send + Sync {
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move { Ok(request) })
}
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move { Ok(request) })
}
fn success<'a>(
&'a self,
_context: &'a CallLifecycleContext,
@ -80,7 +91,7 @@ impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse
request: LiteLLMOcrRequest,
) -> Self::PreCallFuture<'a> {
Box::pin(async move {
if !self.hooks.has_guardrails() {
if !self.hooks.intercepts_requests() {
return Ok(request);
}
let changed = self

View file

@ -0,0 +1,640 @@
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use tokio::sync::{mpsc, oneshot};
use super::handler::perform_ocr_request;
use super::hooks::{
OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest,
OcrPreCallRequest,
};
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient};
use crate::AuthError;
use crate::Error;
use crate::auth::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};
use crate::call_lifecycle::host::{
HostCall, HostCallFuture, HostCallStep, HostFailure, HostLifecycle, HostPhase,
};
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming};
pub type NativeResult<T> = Result<NativeOutcome<T>, Error>;
#[derive(Debug, PartialEq, Eq)]
pub enum NativeOutcome<T> {
Completed(T),
Declined(OcrDecline),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OcrDecline {
ProviderWorkflow,
HostOperations,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct OcrAdmission {
pub provider_workflow: bool,
pub host_operations: bool,
pub asynchronous: bool,
}
impl OcrAdmission {
pub const fn all() -> Self {
Self {
provider_workflow: true,
host_operations: true,
asynchronous: false,
}
}
}
#[derive(Clone, Debug)]
pub enum OcrHostOperation {
ProjectRequest,
Lifecycle(HostPhase),
ConstructResponse(Arc<LiteLLMOcrResponse>),
MapFailure(Error),
Success {
context: CallLifecycleContext,
response: Arc<LiteLLMOcrResponse>,
timing: CallLifecycleTiming,
},
Failure {
context: CallLifecycleContext,
error: Error,
timing: CallLifecycleTiming,
},
AcquireAzureAdToken,
PreCall(OcrPreCallRequest),
DuringCall(OcrDuringCallRequest),
PostCall(OcrPostCallRequest),
}
impl OcrHostOperation {
pub const fn phase(&self) -> Option<HostPhase> {
match self {
Self::Lifecycle(phase) => Some(*phase),
Self::Success { .. } => Some(HostPhase::Success),
Self::Failure { .. } => Some(HostPhase::Failure),
_ => None,
}
}
}
pub enum OcrHostResult {
Request(Result<(Box<LiteLLMOcrRequest>, bool), Error>),
Lifecycle(Result<(), HostFailure>),
AzureAdToken(Result<ResolvedCredential, AuthError>),
PreCall(Result<OcrPreCallRequest, Error>),
DuringCall(Result<OcrDuringCallRequest, Error>),
PostCall(Result<OcrPostCallRequest, Error>),
}
pub type OcrCallStep = HostCallStep<OcrHostOperation, LiteLLMOcrResponse>;
pub struct OcrCall {
lifecycle: HostLifecycle,
execution: OcrExecution,
response: Option<Arc<LiteLLMOcrResponse>>,
error: Option<Error>,
pending: bool,
completed: bool,
projecting: bool,
}
impl OcrCall {
pub fn admit(client: OcrClient, admission: OcrAdmission) -> NativeOutcome<Self> {
if !admission.provider_workflow {
return NativeOutcome::Declined(OcrDecline::ProviderWorkflow);
}
if !admission.host_operations {
return NativeOutcome::Declined(OcrDecline::HostOperations);
}
NativeOutcome::Completed(Self {
lifecycle: HostLifecycle::new(admission.asynchronous),
execution: OcrExecution::new(client),
response: None,
error: None,
pending: false,
completed: false,
projecting: false,
})
}
pub async fn resume(&mut self, result: Option<OcrHostResult>) -> Result<OcrCallStep, Error> {
if self.completed {
return Err(Error::InvalidRequest(
"OCR call cannot be resumed after completion".into(),
));
}
if self.pending != result.is_some() {
return Err(Error::InvalidRequest(
"OCR host operation result does not match pending state".into(),
));
}
match &result {
Some(OcrHostResult::Lifecycle(Ok(())))
if self.lifecycle.phase() == HostPhase::Execute =>
{
return Err(Error::InvalidRequest(
"OCR provider operation requires a typed result".into(),
));
}
Some(result)
if !matches!(result, OcrHostResult::Lifecycle(_))
&& self.lifecycle.phase() != HostPhase::Execute =>
{
return Err(Error::InvalidRequest(
"unexpected OCR provider operation result".into(),
));
}
_ => {}
}
self.pending = false;
let provider_result = match result {
Some(OcrHostResult::Request(result)) if self.projecting => {
self.projecting = false;
match result {
Ok((request, azure_ad_token_provider)) => {
self.execution.request = Some(*request);
self.execution.azure_ad_token_provider = azure_ad_token_provider;
}
Err(error) => self.accept(Err(HostFailure::Error(error))),
}
None
}
Some(OcrHostResult::Request(_)) => {
return Err(Error::InvalidRequest(
"unexpected OCR request projection".into(),
));
}
Some(OcrHostResult::Lifecycle(result)) => {
self.accept(result);
None
}
result => result,
};
if self.lifecycle.phase() == HostPhase::Execute {
if self.execution.request.is_none()
&& self.execution.execution.is_none()
&& !self.execution.completed
{
self.projecting = true;
return Ok(self.host_step(OcrHostOperation::ProjectRequest));
}
match self.execution.resume(provider_result).await {
Ok(OcrCallStep::Host(operation)) => return Ok(self.host_step(operation)),
Ok(OcrCallStep::Complete(response)) => {
self.response = Some(Arc::new(response));
self.accept(Ok(()));
}
Err(error) => self.accept(Err(HostFailure::Error(error))),
}
}
if self.error.is_some() {
self.execution.stop().await;
}
let operation = match self.lifecycle.phase() {
HostPhase::Complete => {
self.completed = true;
return match self.error.take() {
Some(error) => Err(error),
None => self
.response
.take()
.map(Arc::unwrap_or_clone)
.map(OcrCallStep::Complete)
.ok_or_else(|| {
Error::InvalidRequest("OCR completed without a response".into())
}),
};
}
HostPhase::ConstructResponse => OcrHostOperation::ConstructResponse(
self.response
.as_ref()
.ok_or_else(|| Error::InvalidRequest("missing OCR response".into()))?
.clone(),
),
HostPhase::MapFailure => OcrHostOperation::MapFailure(
self.error
.as_ref()
.ok_or_else(|| Error::InvalidRequest("missing OCR failure".into()))?
.clone(),
),
HostPhase::Success | HostPhase::Failure => {
let snapshot = self
.execution
.terminal
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone();
match (self.lifecycle.phase(), snapshot) {
(HostPhase::Success, Some((context, timing))) => OcrHostOperation::Success {
context,
response: self
.response
.as_ref()
.ok_or_else(|| Error::InvalidRequest("missing OCR response".into()))?
.clone(),
timing,
},
(HostPhase::Failure, Some((context, timing))) => OcrHostOperation::Failure {
context,
error: self
.error
.as_ref()
.ok_or_else(|| Error::InvalidRequest("missing OCR failure".into()))?
.clone(),
timing,
},
(phase, _) => OcrHostOperation::Lifecycle(phase),
}
}
phase => OcrHostOperation::Lifecycle(phase),
};
Ok(self.host_step(operation))
}
fn accept(&mut self, result: Result<(), HostFailure>) {
let cancelled = matches!(&result, Err(HostFailure::Cancelled(_)));
if let Some(error) = self.lifecycle.accept(result) {
if cancelled {
self.error = Some(error);
} else {
self.error.get_or_insert(error);
}
self.execution.cancel();
}
}
pub async fn interrupt(&mut self, failure: HostFailure) -> Result<OcrCallStep, Error> {
if self.completed {
return Err(Error::InvalidRequest(
"OCR call cannot be interrupted after completion".into(),
));
}
self.pending = false;
self.accept(Err(failure));
self.resume(None).await
}
fn host_step(&mut self, operation: OcrHostOperation) -> OcrCallStep {
self.pending = true;
OcrCallStep::Host(operation)
}
}
impl HostCall for OcrCall {
type Operation = OcrHostOperation;
type Result = OcrHostResult;
type Complete = LiteLLMOcrResponse;
fn resume(
&mut self,
result: Option<Self::Result>,
) -> HostCallFuture<'_, Self::Operation, Self::Complete> {
Box::pin(OcrCall::resume(self, result))
}
fn interrupt(
&mut self,
failure: HostFailure,
) -> HostCallFuture<'_, Self::Operation, Self::Complete> {
Box::pin(OcrCall::interrupt(self, failure))
}
}
struct PendingOperation {
operation: OcrHostOperation,
result: oneshot::Sender<OcrHostResult>,
}
struct OcrExecution {
client: Option<OcrClient>,
request: Option<LiteLLMOcrRequest>,
operations_tx: mpsc::UnboundedSender<PendingOperation>,
operations_rx: mpsc::UnboundedReceiver<PendingOperation>,
pending_result: Option<oneshot::Sender<OcrHostResult>>,
execution: Option<tokio::task::JoinHandle<Result<LiteLLMOcrResponse, Error>>>,
completed: bool,
azure_ad_token_provider: bool,
terminal: Arc<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
}
impl OcrExecution {
fn new(client: OcrClient) -> Self {
let (operations_tx, operations_rx) = mpsc::unbounded_channel();
Self {
client: Some(client),
request: None,
operations_tx,
operations_rx,
pending_result: None,
execution: None,
completed: false,
azure_ad_token_provider: false,
terminal: Arc::default(),
}
}
pub async fn resume(&mut self, result: Option<OcrHostResult>) -> Result<OcrCallStep, Error> {
if self.completed {
return Err(Error::InvalidRequest(
"OCR call cannot be resumed after completion".into(),
));
}
match (self.pending_result.take(), result) {
(Some(sender), Some(result)) => sender
.send(result)
.map_err(|_| Error::InvalidRequest("OCR host operation was abandoned".into()))?,
(None, None) if self.execution.is_none() => self.start(),
(Some(sender), None) => {
self.pending_result = Some(sender);
return Err(Error::InvalidRequest(
"OCR host operation result is required".into(),
));
}
(None, Some(_)) => {
return Err(Error::InvalidRequest(
"unexpected OCR host operation result".into(),
));
}
(None, None) => {}
}
let execution = self.execution.as_mut().ok_or_else(|| {
Error::InvalidRequest("OCR call cannot be resumed after completion".into())
})?;
tokio::select! {
operation = self.operations_rx.recv() => {
let operation = operation.ok_or_else(|| Error::InvalidRequest("OCR operation channel closed".into()))?;
self.pending_result = Some(operation.result);
Ok(OcrCallStep::Host(operation.operation))
}
result = execution => {
self.execution = None;
self.completed = true;
result
.map_err(|error| Error::Network(format!("OCR execution task failed: {error}")))?
.map(OcrCallStep::Complete)
}
}
}
fn start(&mut self) {
let client = self.client.take().expect("admitted OCR call has a client");
let mut request = self
.request
.take()
.expect("admitted OCR call has a request");
let intercepts_requests = request.hooks.intercepts_requests();
if self.azure_ad_token_provider {
request.azure_ad_token_provider = Some(TokenProviderHandle::new(Arc::new(
OcrAzureAdTokenProvider {
operations: self.operations_tx.clone(),
},
)));
}
request.hooks = Arc::new(ProtocolHooks {
operations: self.operations_tx.clone(),
intercepts_requests,
terminal: self.terminal.clone(),
});
self.execution = Some(tokio::spawn(async move {
perform_ocr_request(&client, request).await
}));
}
fn cancel(&mut self) {
self.pending_result = None;
if let Some(execution) = &self.execution {
execution.abort();
}
}
async fn stop(&mut self) {
self.cancel();
if let Some(execution) = self.execution.as_mut() {
let _ = execution.await;
}
self.execution = None;
}
}
impl Drop for OcrExecution {
fn drop(&mut self) {
if let Some(execution) = &self.execution {
execution.abort();
}
}
}
struct ProtocolHooks {
operations: mpsc::UnboundedSender<PendingOperation>,
intercepts_requests: bool,
terminal: Arc<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
}
#[derive(Debug)]
struct OcrAzureAdTokenProvider {
operations: mpsc::UnboundedSender<PendingOperation>,
}
impl TokenProvider for OcrAzureAdTokenProvider {
fn acquire(&self) -> TokenFuture<'_> {
Box::pin(async move {
let (result, receiver) = oneshot::channel();
self.operations
.send(PendingOperation {
operation: OcrHostOperation::AcquireAzureAdToken,
result,
})
.map_err(|_| {
AuthError::AzureTokenAcquisition("OCR host driver was abandoned".into())
})?;
match receiver.await.map_err(|_| {
AuthError::AzureTokenAcquisition(
"OCR token provider operation was abandoned".into(),
)
})? {
OcrHostResult::AzureAdToken(result) => result,
_ => Err(AuthError::AzureTokenAcquisition(
"invalid OCR token provider host result".into(),
)),
}
})
}
}
impl ProtocolHooks {
async fn invoke(&self, operation: OcrHostOperation) -> Result<OcrHostResult, Error> {
let (result, receiver) = oneshot::channel();
self.operations
.send(PendingOperation { operation, result })
.map_err(|_| Error::InvalidRequest("OCR host driver was abandoned".into()))?;
receiver
.await
.map_err(|_| Error::InvalidRequest("OCR host operation was abandoned".into()))
}
}
impl OcrHooks for ProtocolHooks {
fn intercepts_requests(&self) -> bool {
self.intercepts_requests
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
Box::pin(async move {
match self.invoke(OcrHostOperation::PreCall(request)).await? {
OcrHostResult::PreCall(result) => result,
_ => Err(Error::InvalidRequest(
"invalid OCR pre-call host result".into(),
)),
}
})
}
fn during_call(
&self,
request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
match self.invoke(OcrHostOperation::DuringCall(request)).await? {
OcrHostResult::DuringCall(result) => result,
_ => Err(Error::InvalidRequest(
"invalid OCR during-call host result".into(),
)),
}
})
}
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move {
match self.invoke(OcrHostOperation::PostCall(request)).await? {
OcrHostResult::PostCall(result) => result,
_ => Err(Error::InvalidRequest(
"invalid OCR post-call host result".into(),
)),
}
})
}
fn success<'a>(
&'a self,
context: &'a CallLifecycleContext,
_response: &'a LiteLLMOcrResponse,
timing: &'a CallLifecycleTiming,
) -> OcrLogFuture<'a> {
Box::pin(async move {
*self
.terminal
.lock()
.unwrap_or_else(|error| error.into_inner()) =
Some((context.clone(), timing.clone()));
})
}
fn failure<'a>(
&'a self,
context: &'a CallLifecycleContext,
_error: &'a Error,
timing: &'a CallLifecycleTiming,
) -> OcrLogFuture<'a> {
Box::pin(async move {
*self
.terminal
.lock()
.unwrap_or_else(|error| error.into_inner()) =
Some((context.clone(), timing.clone()));
})
}
}
pub type OcrHostFuture<'a> = Pin<Box<dyn Future<Output = OcrHostResult> + Send + 'a>>;
pub trait OcrHost: Send + Sync {
fn invoke(&self, operation: OcrHostOperation) -> OcrHostFuture<'_>;
}
pub struct NoopOcrHost;
impl OcrHost for NoopOcrHost {
fn invoke(&self, operation: OcrHostOperation) -> OcrHostFuture<'_> {
Box::pin(async move {
match operation {
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err(
Error::InvalidRequest("OCR host has no request projection".into()),
)),
OcrHostOperation::Lifecycle(_)
| OcrHostOperation::ConstructResponse(_)
| OcrHostOperation::MapFailure(_)
| OcrHostOperation::Success { .. }
| OcrHostOperation::Failure { .. } => OcrHostResult::Lifecycle(Ok(())),
OcrHostOperation::AcquireAzureAdToken => {
OcrHostResult::AzureAdToken(Err(AuthError::AzureTokenAcquisition(
"OCR host has no Azure AD token provider".into(),
)))
}
OcrHostOperation::PreCall(request) => OcrHostResult::PreCall(Ok(request)),
OcrHostOperation::DuringCall(request) => OcrHostResult::DuringCall(Ok(request)),
OcrHostOperation::PostCall(request) => OcrHostResult::PostCall(Ok(request)),
}
})
}
}
pub struct OcrHookHost {
hooks: Arc<dyn OcrHooks>,
}
impl OcrHookHost {
pub fn new(hooks: Arc<dyn OcrHooks>) -> Self {
Self { hooks }
}
}
impl OcrHost for OcrHookHost {
fn invoke(&self, operation: OcrHostOperation) -> OcrHostFuture<'_> {
Box::pin(async move {
match operation {
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err(
Error::InvalidRequest("OCR hook host has no request projection".into()),
)),
OcrHostOperation::Success {
context,
response,
timing,
} => {
self.hooks.success(&context, &response, &timing).await;
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::Failure {
context,
error,
timing,
} => {
self.hooks.failure(&context, &error, &timing).await;
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::Lifecycle(_)
| OcrHostOperation::ConstructResponse(_)
| OcrHostOperation::MapFailure(_) => OcrHostResult::Lifecycle(Ok(())),
OcrHostOperation::AcquireAzureAdToken => {
OcrHostResult::AzureAdToken(Err(AuthError::AzureTokenAcquisition(
"OCR hook host has no Azure AD token provider".into(),
)))
}
OcrHostOperation::PreCall(request) => {
OcrHostResult::PreCall(self.hooks.pre_call(request).await)
}
OcrHostOperation::DuringCall(request) => {
OcrHostResult::DuringCall(self.hooks.during_call(request).await)
}
OcrHostOperation::PostCall(request) => {
OcrHostResult::PostCall(self.hooks.post_call(request).await)
}
}
})
}
}

View file

@ -5,12 +5,18 @@ mod document;
pub mod error;
mod handler;
pub mod hooks;
mod lifecycle;
mod prepare;
mod registry;
pub mod types;
pub mod wire;
pub use client::{OcrClient, ocr};
pub use document::{encode_file_document, mime_type_for_name, upload_mime_type};
pub use lifecycle::{
NativeOutcome, NativeResult, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline,
OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult,
};
pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
#[cfg(test)]

View file

@ -62,34 +62,48 @@ pub(crate) async fn transform_request_body<B>(
request: &LiteLLMOcrRequest,
url: &str,
headers: &[(String, String)],
retains_document: bool,
body: B,
validate: impl FnOnce(&B) -> Result<(), OcrRequestError>,
) -> Result<reqwest::Request, OcrError>
where
B: Serialize + DeserializeOwned,
{
let body = if request.hooks.has_guardrails() {
let (body, headers) = if request.hooks.intercepts_requests() {
let body = serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField {
path: "body".into(),
})?;
let retained_fields = request
.optional_params
.keys()
.filter(|name| body.get(*name).is_some())
.cloned()
.chain(retains_document.then(|| "document".to_string()))
.collect();
let changed = request
.hooks
.during_call(OcrDuringCallRequest {
model: request.model.clone(),
custom_llm_provider: request.adapter.provider().as_str().into(),
url: url.into(),
body: serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField {
path: "body".into(),
})?,
headers: headers.to_vec(),
body,
retained_fields,
})
.await?;
let body = OcrWireBody::<B>::decode(changed.body)?;
validate(&body.body)?;
body
(body, changed.headers)
} else {
OcrWireBody {
body,
extra: Map::new(),
}
(
OcrWireBody {
body,
extra: Map::new(),
},
headers.to_vec(),
)
};
build_http_request(client, request, url, headers, &body)
build_http_request(client, request, url, &headers, &body)
}
pub(crate) fn build_http_request<B: Serialize>(
@ -113,9 +127,10 @@ pub(crate) fn build_http_request<B: Serialize>(
pub(crate) async fn guardrail_document(
request: &LiteLLMOcrRequest,
url: &str,
) -> Result<OcrDocument, OcrError> {
if !request.hooks.has_guardrails() {
return Ok(request.document.clone());
headers: &[(String, String)],
) -> Result<(OcrDocument, Vec<(String, String)>), OcrError> {
if !request.hooks.intercepts_requests() {
return Ok((request.document.clone(), headers.to_vec()));
}
let changed = request
.hooks
@ -123,14 +138,17 @@ pub(crate) async fn guardrail_document(
model: request.model.clone(),
custom_llm_provider: request.adapter.provider().as_str().into(),
url: url.into(),
headers: headers.to_vec(),
body: serde_json::to_value(&request.document).map_err(|_| {
OcrRequestError::RequestField {
path: "document".into(),
}
})?,
retained_fields: Vec::new(),
})
.await?;
super::wire::decode_request_value(changed.body, "guardrail.document").map_err(OcrError::from)
let document = super::wire::decode_request_value(changed.body, "guardrail.document")?;
Ok((document, changed.headers))
}
#[derive(Serialize)]

View file

@ -23,6 +23,7 @@ super::adapters::for_each_ocr_adapter!(define_adapter_types);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum OcrProvider {
Cohere,
Mistral,
AzureAi,
Reducto,
@ -32,6 +33,7 @@ pub(crate) enum OcrProvider {
impl OcrProvider {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Cohere => "cohere",
Self::Mistral => "mistral",
Self::AzureAi => "azure_ai",
Self::Reducto => "reducto",
@ -50,6 +52,7 @@ pub(crate) fn resolve_wire_adapter(
custom_llm_provider: OcrProvider::Mistral.as_str(),
});
let typed_provider = match provider.custom_llm_provider {
"cohere" => OcrProvider::Cohere,
"mistral" => OcrProvider::Mistral,
"azure_ai" => OcrProvider::AzureAi,
"reducto" => OcrProvider::Reducto,
@ -57,10 +60,17 @@ pub(crate) fn resolve_wire_adapter(
value => return Err(Error::InvalidProvider(value.to_string())),
};
let adapter = match typed_provider {
OcrProvider::Cohere => OcrAdapterKind::Cohere,
OcrProvider::Mistral => OcrAdapterKind::Mistral,
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
OcrAdapterKind::AzureDocumentIntelligence
}
OcrProvider::AzureAi
if provider.model.to_ascii_lowercase().contains("cohere")
&& provider.model.to_ascii_lowercase().contains("parse") =>
{
OcrAdapterKind::AzureCohere
}
OcrProvider::AzureAi => OcrAdapterKind::AzureMistral,
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
OcrAdapterKind::ReductoLegacy
@ -68,12 +78,7 @@ pub(crate) fn resolve_wire_adapter(
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-v3") => {
OcrAdapterKind::ReductoV3
}
OcrProvider::Reducto => {
return Err(Error::InvalidRequest(format!(
"unsupported Reducto OCR model: {}",
provider.model
)));
}
OcrProvider::Reducto => OcrAdapterKind::ReductoV3,
OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
OcrAdapterKind::VertexDeepSeek
}
@ -107,11 +112,10 @@ mod tests {
}
#[test]
fn unknown_reducto_models_are_rejected() {
assert!(matches!(
resolve_wire_adapter("reducto/future-parse-model", None),
Err(Error::InvalidRequest(_))
));
fn unknown_reducto_models_use_the_current_protocol() {
let (model, adapter) = resolve_wire_adapter("reducto/future-parse-model", None).unwrap();
assert_eq!(model, "future-parse-model");
assert_eq!(adapter, OcrAdapterKind::ReductoV3);
}
#[test]

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