Merge remote-tracking branch 'upstream/main' into litellm_scheduler_remove_admitted_requests

# Conflicts:
#	tests/test_litellm/test_router.py
This commit is contained in:
RachelHuangZW 2026-09-25 09:03:23 -04:00
commit 471cc693e5
805 changed files with 33501 additions and 74401 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

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

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

View file

@ -13,6 +13,15 @@ on:
have its path existence-checked like any other token.
required: true
type: string
fork-flag:
description: >-
Codecov flag of the `.circleci/tests.yml` job that now owns part of
this shard. CircleCI does not run on pull requests from forks, so on
those events this shard also runs the files
`.circleci/scripts/unit_selection.sh` lists for the flag.
required: false
type: string
default: ""
workers:
description: "Number of pytest-xdist workers"
required: false
@ -92,6 +101,7 @@ jobs:
pull-requests: read
outputs:
decision: ${{ steps.changes.outputs.decision }}
has-coverage: ${{ steps.tests.outputs.has-coverage }}
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
@ -160,10 +170,13 @@ jobs:
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Run tests
id: tests
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: ${{ inputs.timeout-minutes }}
env:
TEST_PATH: ${{ inputs.test-path }}
FORK_FLAG: ${{ inputs.fork-flag }}
IS_FORK: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository }}
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
@ -171,9 +184,18 @@ jobs:
DIST: ${{ inputs.dist }}
COVERAGE_CORE: sysmon
run: |
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
selection="${TEST_PATH}"
if [ "${IS_FORK}" = "true" ] && [ -n "${FORK_FLAG}" ]; then
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${FORK_FLAG}" | tr '\n' ' ')"
fi
if [ -z "${selection// /}" ]; then
echo "shard selection is empty on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run"
exit 0
fi
pytest_args=()
existing_paths=0
for token in ${TEST_PATH:?}; do
for token in ${selection}; do
case "${token}" in
-*) pytest_args+=("${token}") ;;
*)
@ -187,7 +209,7 @@ jobs:
esac
done
if [ "${existing_paths}" -eq 0 ]; then
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
echo "No path in the selection exists (${selection}); nothing to run"
exit 0
fi
xdist_args=()
@ -209,8 +231,11 @@ jobs:
--cov-config=pyproject.toml
status=$?
set -e
if [ -f coverage.xml ]; then
echo "has-coverage=true" >> "$GITHUB_OUTPUT"
fi
if [ "$status" -eq 5 ]; then
echo "pytest collected no tests from ${TEST_PATH}; passing"
echo "pytest collected no tests from ${selection}; passing"
exit 0
fi
exit "$status"
@ -226,7 +251,7 @@ jobs:
upload-coverage:
name: Upload coverage to Codecov
needs: run
if: always() && needs.run.outputs.decision != 'skip'
if: always() && needs.run.outputs.decision != 'skip' && needs.run.outputs.has-coverage == 'true'
runs-on: ubuntu-latest
permissions:
contents: read

View file

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

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

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

View file

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

View file

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

View file

@ -31,10 +31,14 @@ concurrency:
# number, so a partially-specified entry would fail the call rather than fall
# back to the default.
#
# tests/proxy_unit_tests keeps its own caller (test-unit-proxy-db.yml): it is
# already a matrix and carries a shard-coverage guard that reads that file by
# name. Folding it in here is a follow-up, together with generalising that guard
# into assert_ci_coverage.py.
# tests/unit/proxy keeps its own caller (test-unit-proxy-db.yml): it is already
# a matrix and carries a shard-coverage guard that reads that file by name.
# Folding it in here is a follow-up, together with generalising that guard into
# assert_ci_coverage.py.
#
# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the
# shard under the same Codecov flag. CircleCI does not build pull requests from
# forks, so the shard still runs those files there and skips them elsewhere.
jobs:
unit:
name: ${{ matrix.shard }}
@ -49,6 +53,7 @@ jobs:
- shard: mcp-integration
artifact-name: mcp-integration
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
fork-flag: mcp-integration
workers: 2
reruns: 0
timeout-minutes: 20
@ -65,10 +70,10 @@ jobs:
- shard: enterprise-routing
artifact-name: enterprise-routing
test-path: >-
tests/test_litellm/enterprise
tests/test_litellm/google_genai
tests/test_litellm/router_utils
tests/test_litellm/router_strategy
fork-flag: enterprise-routing
workers: 2
reruns: 2
timeout-minutes: 20
@ -200,7 +205,7 @@ jobs:
tests/test_litellm/proxy/types_utils
tests/test_litellm/proxy/logging_endpoints
tests/test_litellm/proxy/test_*.py
tests/test_gateway
fork-flag: proxy-infra
workers: 4
reruns: 2
timeout-minutes: 20
@ -208,11 +213,8 @@ jobs:
- shard: caching-local
artifact-name: caching-local
test-path: >-
tests/local_testing/test_cache_preset_key.py
tests/local_testing/test_caching_handler.py
tests/local_testing/test_responses_stream_cache_keys.py
tests/local_testing/test_unit_test_caching.py
test-path: ""
fork-flag: caching-local
workers: 2
reruns: 2
timeout-minutes: 20
@ -220,7 +222,8 @@ jobs:
- shard: proxy-extras
artifact-name: proxy-extras
test-path: "tests/litellm-proxy-extras"
test-path: ""
fork-flag: proxy-extras
workers: 2
reruns: 2
timeout-minutes: 20
@ -228,7 +231,8 @@ jobs:
- shard: enterprise-package
artifact-name: enterprise-package
test-path: "tests/enterprise"
test-path: ""
fork-flag: enterprise-package
workers: 4
reruns: 2
timeout-minutes: 20
@ -247,6 +251,7 @@ jobs:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
fork-flag: ${{ matrix.fork-flag || '' }}
workers: ${{ matrix.workers }}
reruns: ${{ matrix.reruns }}
timeout-minutes: ${{ matrix.timeout-minutes }}

View file

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

View file

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

View file

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

View file

@ -1,6 +1,6 @@
# LiteLLM All Prometheus Metrics dashboard
Every `litellm_*` metric family the proxy can expose on `/metrics` (134 families across 95 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -8,3 +8,11 @@
- Split a mixed test file along that line instead of widening visibility to move it
- A test for another crate's item belongs in that crate, not in a downstream one
- Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests/<name>/mod.rs` or `tests/<subject>/support.rs` so it is not picked up as a test crate of its own
## Error definitions
- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs`
- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each
- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string
- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return
- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it

View file

@ -3013,6 +3013,15 @@ dependencies = [
"url",
]
[[package]]
name = "litellm-coroutine"
version = "0.1.0"
dependencies = [
"rstest",
"thiserror 2.0.19",
"tokio",
]
[[package]]
name = "litellm-cost"
version = "0.1.0"
@ -3040,6 +3049,7 @@ name = "litellm-host"
version = "0.1.0"
dependencies = [
"litellm-auth",
"litellm-coroutine",
"rstest",
"serde_json",
"tokio",
@ -3049,6 +3059,7 @@ dependencies = [
name = "litellm-host-python"
version = "0.1.0"
dependencies = [
"bytes",
"futures-util",
"litellm-host",
"pyo3",

View file

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

View file

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

View file

@ -1,4 +1,5 @@
use std::{
convert::Infallible,
sync::{Arc, Mutex},
time::Duration,
};
@ -9,8 +10,8 @@ use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
host::{Demand, Host},
machine::{HostChannel, MachineFault, RouteMachine},
route::Route,
machine::{CallMachine, HostChannel, MachineFault},
protocol::Protocol,
};
use litellm_secrets::source::SecretSource;
use litellm_types::{
@ -28,15 +29,6 @@ use super::{
};
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessagesOp {
ProjectRequest,
}
pub enum MessagesOpResult {
Request(Box<MessagesCall>),
}
/// The caller's request as the host projects it.
pub struct MessagesCall {
pub model: String,
@ -64,11 +56,11 @@ pub enum MessagesOutput {
pub struct Messages;
impl Route for Messages {
impl Protocol for Messages {
type Response = MessagesOutput;
type Error = Error;
type Op = MessagesOp;
type OpResult = MessagesOpResult;
type Projection = MessagesCall;
type Op = Infallible;
type Chunk = Bytes;
type StreamHead = ();
}
@ -78,13 +70,12 @@ impl From<MachineFault> for Error {
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "messages host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("messages {message}"),
MachineFault::Mismatch => "invalid messages host operation result".into(),
})
}
}
pub type MessagesHost = HostChannel<Messages>;
pub type MessagesMachine = RouteMachine<Messages>;
pub type MessagesMachine = CallMachine<Messages>;
/// Whether this route serves the request, decided before any callback runs so a host
/// can still run its own path.
@ -114,30 +105,28 @@ impl LocalMessagesHost {
}
impl Host<Messages> for LocalMessagesHost {
async fn route(&self, op: MessagesOp) -> Result<MessagesOpResult, Error> {
match op {
MessagesOp::ProjectRequest => self
.call
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.map(|call| MessagesOpResult::Request(Box::new(call)))
.ok_or_else(|| {
Error::InvalidRequest("messages request was already projected".into())
}),
}
async fn project(&self) -> Result<MessagesCall, Error> {
self.call
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.ok_or_else(|| Error::InvalidRequest("messages request was already projected".into()))
}
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
match op {}
}
}
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
CallMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
}
async fn execute(
host: MessagesHost,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesOutput, Error> {
let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?;
let call = host.project().await?;
let stream = call.streams();
let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
let secrets = secrets.resolve(resolved.config.secret_names()).await?;

View file

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

View file

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

View file

@ -3,102 +3,73 @@ use std::sync::{Arc, Mutex};
use litellm_auth::ResolvedCredential;
use litellm_host::{
event::{CallEvent, RequestContext, WireRequest},
machine::{HostChannel, HostTokenProvider, MachineFault, RouteMachine, TokenRoute},
route::Route,
host::Reply,
machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol},
protocol::Protocol,
};
use litellm_llms::base_llm::ocr::{
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
};
use super::handler::perform_ocr_request;
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, OcrFileContent, ResolvedOcrRequest};
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OcrOp {
ProjectRequest,
ReadDocument,
AcquireAzureAdToken,
AcquireAzureAdToken(Reply<ResolvedCredential>),
}
pub enum OcrOpResult {
Request {
request: Box<LiteLLMOcrRequest<OcrDocumentInput>>,
caller_token: bool,
},
Document(OcrFileContent),
AzureAdToken(ResolvedCredential),
/// The caller's request as the host projects it.
pub struct OcrProjection {
pub request: LiteLLMOcrRequest<OcrDocumentInput>,
/// The caller passed its own Azure AD token provider, which the host keeps.
pub caller_token: bool,
}
pub struct Ocr;
impl Route for Ocr {
impl Protocol for Ocr {
type Response = LiteLLMOcrResponse;
type Error = Error;
type Projection = OcrProjection;
type Op = OcrOp;
type OpResult = OcrOpResult;
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
impl TokenRoute for Ocr {
fn acquire_token_op() -> OcrOp {
OcrOp::AcquireAzureAdToken
}
fn token_credential(result: OcrOpResult) -> Option<ResolvedCredential> {
match result {
OcrOpResult::AzureAdToken(credential) => Some(credential),
_ => None,
}
impl TokenProtocol for Ocr {
fn acquire_token_op(reply: Reply<ResolvedCredential>) -> OcrOp {
OcrOp::AcquireAzureAdToken(reply)
}
}
pub type OcrHost = HostChannel<Ocr>;
pub type OcrMachine = RouteMachine<Ocr>;
pub type OcrMachine = CallMachine<Ocr>;
/// The OCR call as a machine: projection, document reading and token acquisition are
/// host operations; everything else runs in Rust.
/// The OCR call as a machine: projection and token acquisition are host operations;
/// everything else runs in Rust.
pub fn ocr_machine(client: OcrClient) -> OcrMachine {
RouteMachine::new(move |host| Box::pin(execute(client, host)))
CallMachine::new(move |host| Box::pin(execute(client, host)))
}
async fn execute(client: OcrClient, host: OcrHost) -> Result<LiteLLMOcrResponse, Error> {
let OcrOpResult::Request {
let OcrProjection {
request,
caller_token,
} = host.route(OcrOp::ProjectRequest).await?
else {
return Err(MachineFault::Mismatch.into());
};
} = host.project().await?;
let request = LiteLLMOcrRequest {
azure_ad_token_provider: caller_token
.then(|| HostTokenProvider::handle(host.clone()))
.or(request.azure_ad_token_provider),
..*request
..request
};
let caller_document = matches!(request.document, OcrDocumentInput::Document(_));
let request = prepare_request_document(request, &host).await?;
let request = prepare_request_document(request).await?;
perform_ocr_request(&client, request, &host, caller_document).await
}
async fn prepare_request_document(
request: LiteLLMOcrRequest<OcrDocumentInput>,
host: &OcrHost,
) -> Result<ResolvedOcrRequest, Error> {
let request = match &request.document {
OcrDocumentInput::HostReader { mime_type } => {
let mime_type = mime_type.clone();
let OcrOpResult::Document(content) = host.route(OcrOp::ReadDocument).await? else {
return Err(MachineFault::Mismatch.into());
};
request.with_document(OcrDocumentInput::Bytes {
bytes: content.bytes,
file_name: content.file_name,
mime_type,
})
}
_ => request,
};
if let OcrDocumentInput::Document(_) = &request.document {
return request.map_document(super::document::prepare_document);
}
@ -107,7 +78,6 @@ async fn prepare_request_document(
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
}
type Reader = Box<dyn Fn() -> Result<OcrFileContent, Error> + Send + Sync>;
type BeforeSend =
Box<dyn Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error> + Send + Sync>;
type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
@ -116,7 +86,6 @@ type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
/// projection, and the optional observer sees and may rewrite the wire request.
pub struct LocalOcrHost {
request: Mutex<Option<LiteLLMOcrRequest<OcrDocumentInput>>>,
reader: Option<Reader>,
before_send: Option<BeforeSend>,
observer: Option<Observer>,
}
@ -125,22 +94,11 @@ impl LocalOcrHost {
pub fn new(request: LiteLLMOcrRequest<OcrDocumentInput>) -> Self {
Self {
request: Mutex::new(Some(request)),
reader: None,
before_send: None,
observer: None,
}
}
pub fn with_reader(
self,
reader: impl Fn() -> Result<OcrFileContent, Error> + Send + Sync + 'static,
) -> Self {
Self {
reader: Some(Box::new(reader)),
..self
}
}
pub fn with_before_send(
self,
before_send: impl Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error>
@ -163,25 +121,21 @@ impl LocalOcrHost {
}
impl litellm_host::host::Host<Ocr> for LocalOcrHost {
async fn route(&self, op: OcrOp) -> Result<OcrOpResult, Error> {
async fn project(&self) -> Result<OcrProjection, Error> {
self.request
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.map(|request| OcrProjection {
request,
caller_token: false,
})
.ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into()))
}
async fn custom_op(&self, op: OcrOp) -> Result<(), Error> {
match op {
OcrOp::ProjectRequest => self
.request
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.map(|request| OcrOpResult::Request {
request: Box::new(request),
caller_token: false,
})
.ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())),
OcrOp::ReadDocument => self
.reader
.as_ref()
.ok_or_else(|| Error::InvalidRequest("OCR host has no document reader".into()))
.and_then(|reader| reader())
.map(OcrOpResult::Document),
OcrOp::AcquireAzureAdToken => {
OcrOp::AcquireAzureAdToken(_) => {
Err(Error::Auth(litellm_auth::Error::AzureTokenAcquisition(
"OCR host has no Azure AD token provider".into(),
)))
@ -2757,7 +2711,7 @@ pub(crate) mod tests {
use litellm_auth_gcp::VertexAuth;
use litellm_host::{
event::{CallEvent, MachineEvent, WireRequest},
host::{Host, HostOp, HostResult},
host::{Host, HostOp},
machine::{HostFailure, Machine, MachineStep},
};
use litellm_http::{
@ -2776,7 +2730,7 @@ pub(crate) mod tests {
use rstest::rstest;
use serde_json::{Value, json};
use crate::ocr::route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine};
use crate::ocr::route::{LocalOcrHost, OcrOp, OcrProjection, ocr_machine};
use crate::ocr::{
test_support::{
MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request,
@ -3212,42 +3166,42 @@ pub(crate) mod tests {
crate::ocr::route::OcrMachine,
) {
let mut machine = ocr_machine(client);
let mut result = None;
let mut ops = Vec::new();
let outcome = loop {
let op = match machine.resume(result.take()).await {
let op = match machine.resume().await {
Ok(MachineStep::Host(op)) => op,
Ok(MachineStep::Complete(response)) => break Ok(response),
Err(error) => break Err(error),
};
let answer = match op {
HostOp::Route(op) => {
ops.push(match op {
OcrOp::ProjectRequest => "ProjectRequest",
OcrOp::ReadDocument => "ReadDocument",
OcrOp::AcquireAzureAdToken => "AcquireAzureAdToken",
});
host.route(op)
HostOp::Project(reply) => {
ops.push("Project");
host.project()
.await
.map(HostResult::Route)
.map(|projection| reply.send(projection))
.map_err(HostFailure::Error)
}
HostOp::BeforeSend { wire, .. } => {
ops.push("BeforeSend");
intercept(*wire).map(|wire| HostResult::BeforeSend(Box::new(wire)))
HostOp::Custom(op) => {
ops.push(match op {
OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken",
});
host.custom_op(op).await.map_err(HostFailure::Error)
}
HostOp::Emit(event) => {
HostOp::BeforeSend { wire, reply, .. } => {
ops.push("BeforeSend");
intercept(*wire).map(|wire| reply.send(wire))
}
HostOp::Emit(event, reply) => {
let event = CallEvent::Machine(event);
ops.push(event_name(&event));
host.emit(&event)
.await
.map(|()| HostResult::Emitted)
.map(|()| reply.send(()))
.map_err(HostFailure::Error)
}
};
match answer {
Ok(answer) => result = Some(answer),
Err(failure) => break machine.interrupt(failure).await,
if let Err(failure) = answer {
break machine.interrupt(failure).await;
}
};
(outcome, ops, machine)
@ -3269,8 +3223,8 @@ pub(crate) mod tests {
assert!(
matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "before_send failed")
);
assert_eq!(ops, ["ProjectRequest", "BeforeSend"]);
assert!(machine.resume(None).await.is_err());
assert_eq!(ops, ["Project", "BeforeSend"]);
assert!(machine.resume().await.is_err());
}
#[tokio::test]
@ -3306,80 +3260,24 @@ pub(crate) mod tests {
server.await.unwrap();
assert_eq!(outcome.unwrap().pages[0].markdown, "native");
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(ops, ["ProjectRequest", "BeforeSend", "response"]);
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
assert!(matches!(
machine.resume(None).await,
machine.resume().await,
Err(OcrError::InvalidRequest(_))
));
}
async fn drive_native_file_call(
request: crate::ocr::types::LiteLLMOcrRequest<crate::ocr::types::OcrDocumentInput>,
content: Result<crate::ocr::types::OcrFileContent, OcrError>,
) -> (Result<LiteLLMOcrResponse, OcrError>, usize) {
let reads = Arc::new(Mutex::new(0));
let counted = reads.clone();
let content = Mutex::new(Some(content));
let host = LocalOcrHost::new(request).with_reader(move || {
*counted.lock().unwrap() += 1;
content.lock().unwrap().take().unwrap()
});
let outcome = perform_ocr_with(host).await;
let reads = *reads.lock().unwrap();
(outcome, reads)
}
#[tokio::test]
async fn host_reader_documents_are_read_once_at_the_core_selected_point_and_encoded() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"file"}]
}))])
.await;
let request = wire_request("mistral/model", &base, json!({})).with_document(
crate::ocr::types::OcrDocumentInput::HostReader {
mime_type: Some("application/pdf".into()),
},
);
let (response, reads) = drive_native_file_call(
request,
Ok(crate::ocr::types::OcrFileContent {
bytes: b"abc".as_slice().into(),
file_name: Some("scan.png".into()),
}),
)
.await;
server.await.unwrap();
assert_eq!(response.unwrap().pages[0].markdown, "file");
assert_eq!(reads, 1);
assert!(seen.lock().unwrap()[0].contains("data:application/pdf;base64,YWJj"));
}
#[tokio::test]
async fn host_reader_failures_and_empty_files_fail_before_the_provider_is_called() {
async fn empty_byte_documents_fail_before_the_provider_is_called() {
let (base, seen, _server) = mock_server(vec![]).await;
let request = wire_request("mistral/model", &base, json!({}));
let failure = OcrError::InvalidRequest("reader exploded".into());
let (response, reads) = drive_native_file_call(
request
.with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }),
Err(failure.clone()),
)
.await;
assert!(
matches!(response.unwrap_err(), OcrError::InvalidRequest(message) if message == "reader exploded")
);
assert_eq!(reads, 1);
let request = wire_request("mistral/model", &base, json!({}));
let (response, _) = drive_native_file_call(
request
.with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }),
Ok(crate::ocr::types::OcrFileContent {
let request = wire_request("mistral/model", &base, json!({})).with_document(
crate::ocr::types::OcrDocumentInput::Bytes {
bytes: Default::default(),
file_name: None,
}),
)
.await;
mime_type: None,
},
);
let response = perform_ocr_with(LocalOcrHost::new(request)).await;
assert!(matches!(response.unwrap_err(), OcrError::EmptyFile));
assert!(seen.lock().unwrap().is_empty());
}
@ -3400,24 +3298,21 @@ pub(crate) mod tests {
mime_type: None,
},
);
let (response, reads) =
drive_native_file_call(request, Err(OcrError::InvalidRequest("unused".into()))).await;
let (response, ops, _) = drive_until(ocr_client(), &LocalOcrHost::new(request), Ok).await;
server.await.unwrap();
std::fs::remove_dir_all(&dir).unwrap();
assert_eq!(response.unwrap().pages[0].markdown, "path");
assert_eq!(reads, 0);
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
assert!(seen.lock().unwrap()[0].contains("data:image/png;base64,YWJj"));
let (base, seen, _server) = mock_server(vec![]).await;
let request = wire_request("mistral/model", &base, json!({}));
let (response, _) = drive_native_file_call(
request.with_document(crate::ocr::types::OcrDocumentInput::Path {
let request = wire_request("mistral/model", &base, json!({})).with_document(
crate::ocr::types::OcrDocumentInput::Path {
path: path.clone(),
mime_type: None,
}),
Err(OcrError::InvalidRequest("unused".into())),
)
.await;
},
);
let response = perform_ocr_with(LocalOcrHost::new(request)).await;
assert!(matches!(
response.unwrap_err(),
OcrError::FileRead { path: failed, source } if failed == path && source.kind() == std::io::ErrorKind::NotFound
@ -3441,28 +3336,25 @@ pub(crate) mod tests {
assert!(
matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "cancelled")
);
assert_eq!(ops, ["ProjectRequest", "BeforeSend"]);
assert!(machine.resume(Some(HostResult::Emitted)).await.is_err());
assert_eq!(ops, ["Project", "BeforeSend"]);
assert!(machine.resume().await.is_err());
}
#[tokio::test]
async fn missing_host_result_preserves_pending_operation() {
async fn resuming_before_answering_preserves_pending_operation() {
let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({}));
let mut machine = ocr_machine(ocr_client());
let Ok(MachineStep::Host(HostOp::Project(reply))) = machine.resume().await else {
panic!("expected the projection op first");
};
assert!(machine.resume().await.is_err());
reply.send(OcrProjection {
request,
caller_token: false,
});
assert!(matches!(
machine.resume(None).await.unwrap(),
MachineStep::Host(HostOp::Route(OcrOp::ProjectRequest))
));
assert!(machine.resume(None).await.is_err());
assert!(matches!(
machine
.resume(Some(HostResult::Route(OcrOpResult::Request {
request: Box::new(request),
caller_token: false,
})))
.await
.unwrap(),
MachineStep::Host(HostOp::BeforeSend { .. })
machine.resume().await,
Ok(MachineStep::Host(HostOp::BeforeSend { .. }))
));
}
@ -3623,20 +3515,18 @@ pub(crate) mod tests {
};
let host = LocalOcrHost::new(request);
let mut machine = ocr_machine(ocr_client());
let mut result = None;
tokio::time::timeout(std::time::Duration::from_secs(2), async {
loop {
tokio::select! {
_ = entered.notified() => break,
step = machine.resume(result.take()) => {
result = Some(match step.unwrap() {
MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()),
MachineStep::Host(HostOp::BeforeSend { wire, .. }) => {
HostResult::BeforeSend(wire)
}
MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted,
step = machine.resume() => {
match step.unwrap() {
MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()),
MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(),
MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire),
MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()),
MachineStep::Complete(_) => panic!("pending provider completed"),
});
}
}
}
}
@ -3661,24 +3551,23 @@ pub(crate) mod tests {
}
impl Host<crate::ocr::route::Ocr> for CallerTokenHost {
async fn route(&self, op: OcrOp) -> Result<OcrOpResult, OcrError> {
async fn project(&self) -> Result<OcrProjection, OcrError> {
self.trace.lock().unwrap().push("project".into());
Ok(OcrProjection {
request: self.request.lock().unwrap().take().unwrap(),
caller_token: true,
})
}
async fn custom_op(&self, op: OcrOp) -> Result<(), OcrError> {
match op {
OcrOp::ProjectRequest => {
self.trace.lock().unwrap().push("project".into());
Ok(OcrOpResult::Request {
request: Box::new(self.request.lock().unwrap().take().unwrap()),
caller_token: true,
})
}
OcrOp::AcquireAzureAdToken => {
OcrOp::AcquireAzureAdToken(reply) => {
self.trace.lock().unwrap().push("token".into());
Ok(OcrOpResult::AzureAdToken(
litellm_auth::ResolvedCredential::Static(litellm_auth::SecretValue::new(
"caller-token",
)),
))
reply.send(litellm_auth::ResolvedCredential::Static(
litellm_auth::SecretValue::new("caller-token"),
));
Ok(())
}
OcrOp::ReadDocument => Err(OcrError::InvalidRequest("no reader".into())),
}
}
@ -3761,18 +3650,18 @@ pub(crate) mod tests {
});
let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({})));
let mut machine = ocr_machine(ocr_client());
let mut result = None;
tokio::time::timeout(std::time::Duration::from_secs(2), async {
loop {
tokio::select! {
_ = received.notified() => break,
step = machine.resume(result.take()) => {
result = Some(match step.unwrap() {
MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()),
MachineStep::Host(HostOp::BeforeSend { wire, .. }) => HostResult::BeforeSend(wire),
MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted,
step = machine.resume() => {
match step.unwrap() {
MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()),
MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(),
MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire),
MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()),
MachineStep::Complete(_) => panic!("the stalled provider completed"),
});
}
}
}
}

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -7,6 +7,26 @@
- Python, Rust SDK and gateway use one lifecycle-bearing core route entrypoint; provider helpers stay private, never bridge-accessible transport drivers
- Built-in provider/config/secret/auth/document preparation stays in Rust; caller-authored callbacks and focused Python-file reads run only at core-selected points
- Target GIL-enabled CPython explicitly with `#[pymodule(gil_used = true)]`; detach Rust-only work
- GIL and tokio invariants, each pinned by a test in `host-python` (`execution.rs`,
`gil.rs`) so a regression fails there before it deadlocks a proxy:
- Never hold the GIL while waiting on the runtime. A sync entrypoint releases it with
`release_gil` around `block_on`, because every task that attaches would otherwise wait
on the thread that is waiting on them ([pyo3 parallelism](https://pyo3.rs/v0.29.2/parallelism.html))
- Never `block_on` from a tokio worker; the sync entrypoints refuse with "cannot run from
a Tokio context" instead of panicking inside the runtime ([tokio `Runtime::block_on`](https://docs.rs/tokio/latest/tokio/runtime/struct.Runtime.html#method.block_on))
- Inside a future, `Python::attach` only for GIL-cheap work: cloning a `Py<T>`, building
a small value, reading a settings snapshot. Anything that can block (a secret manager
read, a callback that does I/O, an import, a network call) goes through
`litellm_host_python::attach_blocking`, which runs it on the blocking pool so the async
workers keep polling other calls ([tokio `spawn_blocking`](https://docs.rs/tokio/latest/tokio/task/fn.spawn_blocking.html)).
`block_in_place` is not an alternative: it needs a multi-thread worker and still steals it
- `attach_blocking` work runs on a thread the interpreter did not create (pinned by the
`threading.get_ident()` test). Like any foreign-thread attach it therefore has no running
asyncio loop and a fresh `contextvars` context: do not hand it a coroutine or anything
bound to the caller's loop
- Dropping the await (an asyncio cancel) does not interrupt the Python call; it runs to
completion and its result is discarded. A panic in it reaches the awaiting task as a panic
- Add a case to `gil.rs` when a new seam changes any of these; the tests are the spec
- Free-threading requires separate runtime/concurrency validation; omitting the attribute does not opt out on PyO3 0.28+
- Preserve public argument binding and Python object provenance
- Project only consumed fields at reference read points; no eager whole-graph serialization or equality-based alias reconstruction
@ -56,6 +76,14 @@ GIL handling to `litellm-host-python`.
decides whether to raise or fall back. For a rust-only provider/route (no
Python reference), the Python side is a thin dispatch that calls Rust and
raises when the bridge is unavailable, with no fallback.
- Declare it by passing `python=NO_PYTHON` (`litellm.rust_bridge.runtime`)
to `PublicDispatch.run`/`arun` or `runtime.run`/`arun`, never a stand-in
callable that raises, and give every context of it a `RUST_REQUIRED`
catalog rule
- Any other decision, an unprojectable call, or a bypass raises
`NoPythonImplementationError` before native runs, so a misdeclared route
fails in tests instead of reaching deleted code. When deleting a route's
Python implementation, switch its dispatch to `NO_PYTHON` in the same change
- Keep the Python interface minimal (well under 100 lines per route): it only
marshals inputs and calls Rust. Do not add per-route feature flags, and do
not put provider dispatch in `litellm/main.py`; it lives in a thin dispatch

View file

@ -1,10 +0,0 @@
Native OCR uses `litellm_secrets::source::SecretSource`. Built-in secret managers resolve to retained Rust backends. Custom Python managers and overrides keep the callback path. Readable managers still require the Rust secret-manager binding to be enabled
The shared proxy initializer captures native configuration without loading the extension or doing native I/O. `_SecretManagerRuntime.from_client` constructs a backend on first use and keeps its handle on the Python client. The secret-manager dispatcher selects Python or Rust through `catalog.py`. Native reads call that handle; Rust routes extract the backend directly. Configuration changes replace the handle, while calls already bound to the previous backend keep using it. Handles cannot be reused after fork. Directly constructed LiteLLM managers are adapted on first native use. Manually supplied SDK clients keep their Python behavior because their credentials cannot be inferred safely. Provider implementations contain no bridge registration
Retention describes ownership and lifetime. `callbacks-legacy-python::PublicCall` owns Python references for one call to preserve identity. A native cache or secret-manager handle owns shared Rust state across calls to preserve connection pools and caches. Both use existing `Py<T>` and shared Rust ownership, with execution and GIL transitions handled by `litellm-host-python`
Cache and secret-manager catalog entries remain Python-only, including when `LITELLM_RUST=1`. This wiring does not change rollout policy
OCR provider requests use the shared `litellm-http` pool. AWS and Google secret-manager SDK clients keep their SDK transports, which do not yet inherit the pool's proxy, TLS, certificate, timeout, or observability configuration. Preserve those SDK transports and configure them equivalently instead of forcing them through reqwest

View file

@ -34,7 +34,7 @@ mod _native {
#[pymodule_export]
use crate::routes::messages::{amessages, messages};
#[pymodule_export]
use crate::routes::ocr::{aocr, ocr};
use crate::routes::ocr::{aocr, ocr, ocr_health_check_document, ocr_passthrough_response};
#[pymodule_export]
use crate::routes::responses::{ResponsesWebSocketConnection, aresponses, responses};
#[pymodule_export]
@ -85,6 +85,8 @@ mod tests {
"ProcessReservedForForking",
"ocr",
"aocr",
"ocr_health_check_document",
"ocr_passthrough_response",
"embedding",
"aembedding",
"transcription",

View file

@ -1,9 +1,8 @@
use std::sync::OnceLock;
use litellm_host::{
host::HostResult,
machine::{HostFailure, Interrupted, Machine, Step},
route::Route,
protocol::Protocol,
};
use litellm_tracing::Logger;
use pyo3::Python;
@ -23,17 +22,17 @@ impl<M> LoggedMachine<M> {
}
impl<M: Machine> Machine for LoggedMachine<M> {
type Route = M::Route;
type Protocol = M::Protocol;
type Complete = M::Complete;
fn resume(&mut self, result: Option<HostResult<Self::Route>>) -> Step<'_, Self> {
fn resume(&mut self) -> Step<'_, Self> {
let logger = self.logger.get_or_init(|| Python::attach(super::capture));
Box::pin(logger.instrument(logger.scope(|| self.machine.resume(result))))
Box::pin(logger.instrument(logger.scope(|| self.machine.resume())))
}
fn interrupt(
&mut self,
failure: HostFailure<<Self::Route as Route>::Error>,
failure: HostFailure<<Self::Protocol as Protocol>::Error>,
) -> Interrupted<'_, Self> {
let logger = self.logger.get_or_init(|| Python::attach(super::capture));
Box::pin(logger.instrument(logger.scope(|| self.machine.interrupt(failure))))

View file

@ -1,29 +1,28 @@
use std::{process::Command, task::Poll};
use litellm_host::{
host::HostResult,
machine::{HostFailure, Interrupted, Machine, MachineStep, Step},
route::Route,
protocol::Protocol,
};
use pyo3::{prelude::*, types::PyDict};
struct DiagnosticMachine;
impl Route for DiagnosticMachine {
impl Protocol for DiagnosticMachine {
type Response = ();
type Error = String;
type Projection = ();
type Op = ();
type OpResult = ();
type Chunk = ();
type StreamHead = ();
}
impl Machine for DiagnosticMachine {
type Route = Self;
type Protocol = Self;
type Complete = ();
fn resume(&mut self, _: Option<HostResult<Self>>) -> Step<'_, Self> {
fn resume(&mut self) -> Step<'_, Self> {
litellm_tracing::warn!("machine started");
Box::pin(async {
tokio::task::yield_now().await;
@ -45,7 +44,7 @@ fn machine_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
let mut machine = super::LoggedMachine::new(DiagnosticMachine);
let mut future = Box::pin(async move {
machine
.resume(None)
.resume()
.await
.map_err(pyo3::exceptions::PyValueError::new_err)?;
machine

View file

@ -1,10 +1,12 @@
use std::convert::Infallible;
use bytes::Bytes;
use litellm_core::messages::{
Error,
route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput},
route::{Messages, MessagesCall, MessagesOutput},
types::MessagesShaping,
};
use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py};
use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py};
use litellm_http::transport::Error as TransportError;
use litellm_types::utils::ProviderSpecificHeaders;
use pyo3::{
@ -80,16 +82,16 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
/// The Python side of the Messages route: projects the prepared arguments and builds the
/// public response, chunks and exceptions.
pub(super) struct MessagesRouteHost {
pub(super) struct MessagesPythonHost {
request: Py<PyAny>,
}
impl MessagesRouteHost {
impl MessagesPythonHost {
pub(super) fn new(request: Py<PyAny>) -> Self {
Self { request }
}
fn project(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<MessagesCall> {
fn projection(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<MessagesCall> {
let request = self.request.bind(py);
let argument = |name: &str| -> PyResult<Option<Bound<'_, PyAny>>> {
Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none()))
@ -208,22 +210,21 @@ impl MessagesRouteHost {
}
}
impl RouteHost for MessagesRouteHost {
type Route = Messages;
impl ProtocolHost for MessagesPythonHost {
type Protocol = Messages;
type Failure = PyErr;
fn invoke(
fn project(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: MessagesOp,
) -> Result<MessagesOpResult, InvokeError<Error>> {
match op {
MessagesOp::ProjectRequest => self
.project(py, arguments)
.map(|call| MessagesOpResult::Request(Box::new(call)))
.map_err(|error| InvokeError::Python(self.map_failure(py, error))),
}
) -> Result<MessagesCall, InvokeError<Error>> {
self.projection(py, arguments)
.map_err(|error| InvokeError::Python(self.map_failure(py, error)))
}
fn invoke(&mut self, _: Python<'_>, op: Infallible) -> Result<(), InvokeError<Error>> {
match op {}
}
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {

View file

@ -1,6 +1,6 @@
mod host;
use host::MessagesRouteHost;
use host::MessagesPythonHost;
use litellm_callbacks_legacy_python::{
LegacySurface, PassThroughStream, PublicCall, run_legacy_call,
};
@ -45,7 +45,7 @@ fn run_messages(
SURFACE,
PublicCall::capture(&request, &args, &kwargs)?,
crate::logger::LoggedMachine::new(messages_machine(secrets)),
MessagesRouteHost::new(request.unbind()),
MessagesPythonHost::new(request.unbind()),
asynchronous,
)
}

View file

@ -1,58 +1,38 @@
use std::path::PathBuf;
use bytes::Bytes;
use litellm_core::ocr::types::{OcrDocumentInput, OcrFileContent};
use litellm_core::ocr::types::OcrDocumentInput;
use litellm_host_python::{PythonFileReader, py_bytes};
use pyo3::{
exceptions::{PyTypeError, PyValueError},
gc::{PyTraverseError, PyVisit},
exceptions::PyValueError,
prelude::*,
pybacked::PyBackedBytes,
sync::PyOnceLock,
types::{PyBytes, PyString, PyType},
};
#[derive(Debug)]
pub(super) struct PythonFileReader {
reader: Py<PyAny>,
name: Option<String>,
/// A `type='file'` document as projected: paths and bytes are typed inputs already; a
/// file-like object is a reader the projection consumes once every other field is read.
pub(super) enum FileDocumentInput {
Ready(OcrDocumentInput),
Deferred {
reader: PythonFileReader,
mime_type: Option<String>,
},
}
impl PythonFileReader {
pub(super) fn read(&self, py: Python<'_>) -> PyResult<OcrFileContent> {
let value = self.reader.bind(py).call0()?;
let bytes = if value.is_instance_of::<PyString>() {
Bytes::from(value.extract::<String>()?)
} else if value.is_instance_of::<PyBytes>() {
extract_bytes(&value)?
} else {
return Err(PyTypeError::new_err(format!(
"OCR file read must return bytes or str, got {}",
value.get_type(),
)));
};
Ok(OcrFileContent {
bytes,
file_name: self.name.clone(),
})
impl FileDocumentInput {
pub(super) fn resolve(self, py: Python<'_>) -> PyResult<OcrDocumentInput> {
match self {
Self::Ready(input) => Ok(input),
Self::Deferred { reader, mime_type } => {
let content = reader.read(py)?;
Ok(OcrDocumentInput::Bytes {
bytes: content.bytes,
file_name: content.file_name,
mime_type,
})
}
}
}
pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.reader)
}
}
fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult<Bytes> {
if value.is_exact_instance_of::<PyBytes>() {
return Ok(Bytes::from_owner(value.extract::<PyBackedBytes>()?));
}
Ok(Bytes::copy_from_slice(
value.extract::<PyBackedBytes>()?.as_ref(),
))
}
pub(super) struct FileDocumentInput {
pub input: OcrDocumentInput,
pub reader: Option<PythonFileReader>,
}
impl FromPyObject<'_, '_> for FileDocumentInput {
@ -87,51 +67,31 @@ impl FromPyObject<'_, '_> for FileDocumentInput {
}
static PATH_LIKE: PyOnceLock<Py<PyType>> = PyOnceLock::new();
if file.is_instance(PATH_LIKE.import(py, "os", "PathLike")?)? {
return Ok(Self {
input: OcrDocumentInput::Path {
path: file.extract::<PathBuf>()?,
mime_type,
},
reader: None,
});
return Ok(Self::Ready(OcrDocumentInput::Path {
path: file.extract::<PathBuf>()?,
mime_type,
}));
}
if file.is_instance_of::<PyBytes>() {
return Ok(Self {
input: OcrDocumentInput::Bytes {
bytes: extract_bytes(&file)?,
file_name: None,
mime_type,
},
reader: None,
});
return Ok(Self::Ready(OcrDocumentInput::Bytes {
bytes: py_bytes(&file)?,
file_name: None,
mime_type,
}));
}
let reader = file
.getattr_opt("read")?
.filter(|value| value.is_callable());
let Some(reader) = reader else {
return Err(PyValueError::new_err(format!(
match PythonFileReader::from_file_like(&file)? {
Some(reader) => Ok(Self::Deferred { reader, mime_type }),
None => Err(PyValueError::new_err(format!(
"Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.",
file.get_type(),
)));
};
let name = file
.getattr_opt("name")?
.filter(|value| !value.is_none())
.map(|value| value.extract::<String>())
.transpose()?;
Ok(Self {
input: OcrDocumentInput::HostReader { mime_type },
reader: Some(PythonFileReader {
reader: reader.unbind(),
name,
}),
})
))),
}
}
}
#[cfg(test)]
mod tests {
use pyo3::types::PyDict;
use pyo3::{exceptions::PyTypeError, types::PyDict};
use super::*;
@ -141,6 +101,13 @@ mod tests {
locals
}
fn ready(input: FileDocumentInput) -> OcrDocumentInput {
match input {
FileDocumentInput::Ready(input) => input,
FileDocumentInput::Deferred { .. } => panic!("expected a ready document"),
}
}
#[test]
fn extraction_validates_required_file_and_optional_mime_type() {
Python::initialize();
@ -167,13 +134,19 @@ mod tests {
.unwrap();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("bare str"));
let error = py
.eval(c"{'file': object()}", None, None)
.unwrap()
.extract::<FileDocumentInput>()
.err()
.unwrap();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("Unsupported file input type"));
let document = py
.eval(c"{'file': b'abc', 'mime_type': 'image/png'}", None, None)
.unwrap();
let input: FileDocumentInput = document.extract().unwrap();
assert!(input.reader.is_none());
assert_eq!(
input.input,
ready(document.extract().unwrap()),
OcrDocumentInput::Bytes {
bytes: b"abc".as_slice().into(),
file_name: None,
@ -199,7 +172,7 @@ class Reader:
return b'abc'
reader = Reader()
document = {'file': reader, 'mime_type': 7}
reader_document = {'file': reader}
reader_document = {'file': reader, 'mime_type': 'application/pdf'}
path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_type': 'image/png'}",
);
let document = locals.get_item("document").unwrap().unwrap();
@ -208,10 +181,6 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ
let document = locals.get_item("reader_document").unwrap().unwrap();
let input: FileDocumentInput = document.extract().unwrap();
assert_eq!(
input.input,
OcrDocumentInput::HostReader { mime_type: None }
);
let reads = || {
locals
.get_item("reader")
@ -223,21 +192,20 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ
.unwrap()
};
assert_eq!(reads(), 0);
let content = input.reader.unwrap().read(py).unwrap();
let resolved = input.resolve(py).unwrap();
assert_eq!(reads(), 1);
assert_eq!(
content,
OcrFileContent {
resolved,
OcrDocumentInput::Bytes {
bytes: b"abc".as_slice().into(),
file_name: Some("scan.png".into()),
mime_type: Some("application/pdf".into()),
}
);
let document = locals.get_item("path_document").unwrap().unwrap();
let input: FileDocumentInput = document.extract().unwrap();
assert!(input.reader.is_none());
assert_eq!(
input.input,
ready(document.extract().unwrap()),
OcrDocumentInput::Path {
path: PathBuf::from("/nonexistent/ocr-projection-test.pdf"),
mime_type: Some("image/png".into()),
@ -245,97 +213,4 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ
);
});
}
#[test]
fn reader_results_are_normalized_and_exceptions_keep_their_identity() {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"failure = KeyError('reader failed')
class Raising:
def read(self):
raise failure
class Text:
def read(self):
return 'héllo'
class Wrong:
def read(self):
return 7
raising = {'file': Raising()}
text = {'file': Text()}
wrong = {'file': Wrong()}",
);
let reader = |name: &str| {
locals
.get_item(name)
.unwrap()
.unwrap()
.extract::<FileDocumentInput>()
.unwrap()
.reader
.unwrap()
};
let error = reader("raising").read(py).unwrap_err();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
assert_eq!(
reader("text").read(py).unwrap().bytes.as_ref(),
"héllo".as_bytes()
);
let error = reader("wrong").read(py).unwrap_err();
assert!(error.is_instance_of::<PyTypeError>(py));
assert!(error.to_string().contains("bytes or str"));
});
}
#[rstest::rstest]
#[case::read("read")]
#[case::name("name")]
fn reader_attribute_failures_keep_their_identity(#[case] attribute: &str) {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"failure = LookupError('file property failed')
class File:
def __getattribute__(self, name):
if name == attribute:
raise failure
return super().__getattribute__(name)
name = 'scan.pdf'
def read(self):
return b'abc'
document = {'file': File()}",
);
locals.set_item("attribute", attribute).unwrap();
let error = locals
.get_item("document")
.unwrap()
.unwrap()
.extract::<FileDocumentInput>()
.err()
.unwrap();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
});
}
#[test]
fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() {
Python::initialize();
let (bytes, pointer) = Python::attach(|py| {
let value = PyBytes::new(py, b"document bytes");
let pointer = value.as_bytes().as_ptr() as usize;
(extract_bytes(value.as_any()).unwrap(), pointer)
});
assert_eq!(bytes.as_ptr() as usize, pointer);
assert_eq!(bytes.as_ref(), b"document bytes");
}
}

View file

@ -1,6 +1,6 @@
use litellm_auth::ResolvedCredential;
use litellm_core::ocr::route::{Ocr, OcrOp, OcrOpResult};
use litellm_host_python::{InvokeError, RouteHost, missing_state, to_py};
use litellm_core::ocr::route::{Ocr, OcrOp, OcrProjection};
use litellm_host_python::{InvokeError, ProtocolHost, missing_state, to_py};
use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse};
use pyo3::{
exceptions::{PyBaseException, PyException},
@ -20,14 +20,15 @@ enum OcrHostData {
Released,
}
/// The Python side of the OCR route: projects the prepared arguments, reads file-like
/// documents, acquires Azure AD tokens, and builds the public response and exception.
pub(super) struct OcrRouteHost {
/// The Python side of the OCR route: projects the prepared arguments (reading a file-like
/// document as it goes), acquires Azure AD tokens, and builds the public response and
/// exception.
pub(super) struct OcrPythonHost {
request: Py<PyAny>,
data: OcrHostData,
}
impl OcrRouteHost {
impl OcrPythonHost {
pub(super) fn new(request: Py<PyAny>) -> Self {
Self {
request,
@ -42,14 +43,6 @@ impl OcrRouteHost {
}
}
fn read_document(&self, py: Python<'_>) -> PyResult<litellm_core::ocr::types::OcrFileContent> {
self.handles()?
.reader
.as_ref()
.ok_or_else(missing_state)?
.read(py)
}
fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult<ResolvedCredential> {
self.handles()?
.azure_ad_token_provider
@ -58,30 +51,21 @@ impl OcrRouteHost {
.acquire(py)
}
fn answer(
fn projection(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: OcrOp,
) -> PyResult<OcrOpResult> {
match op {
OcrOp::ProjectRequest => {
let OcrHostData::Unprojected = self.data else {
return Err(missing_state());
};
let (request, handles) = project_request(self.request.bind(py), arguments)?;
let caller_token = handles.azure_ad_token_provider.is_some();
self.data = OcrHostData::Projected(Box::new(handles));
Ok(OcrOpResult::Request {
request: Box::new(request),
caller_token,
})
}
OcrOp::ReadDocument => self.read_document(py).map(OcrOpResult::Document),
OcrOp::AcquireAzureAdToken => self
.acquire_azure_ad_token(py)
.map(OcrOpResult::AzureAdToken),
}
) -> PyResult<OcrProjection> {
let OcrHostData::Unprojected = self.data else {
return Err(missing_state());
};
let (request, handles) = project_request(self.request.bind(py), arguments)?;
let caller_token = handles.azure_ad_token_provider.is_some();
self.data = OcrHostData::Projected(Box::new(handles));
Ok(OcrProjection {
request,
caller_token,
})
}
fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr {
@ -104,20 +88,28 @@ impl OcrRouteHost {
}
}
impl RouteHost for OcrRouteHost {
type Route = Ocr;
impl ProtocolHost for OcrPythonHost {
type Protocol = Ocr;
type Failure = PyErr;
fn invoke(
fn project(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: OcrOp,
) -> Result<OcrOpResult, InvokeError<Error>> {
self.answer(py, arguments, op)
) -> Result<OcrProjection, InvokeError<Error>> {
self.projection(py, arguments)
.map_err(|error| InvokeError::Python(self.map_failure(py, error)))
}
fn invoke(&mut self, py: Python<'_>, op: OcrOp) -> Result<(), InvokeError<Error>> {
match op {
OcrOp::AcquireAzureAdToken(reply) => self
.acquire_azure_ad_token(py)
.map(|token| reply.send(token))
.map_err(|error| InvokeError::Python(self.map_failure(py, error))),
}
}
fn complete(&mut self, py: Python<'_>, response: LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {
py.import("litellm.rust_bridge.ocr.route_host")?
.getattr("response")?
@ -148,13 +140,10 @@ impl RouteHost for OcrRouteHost {
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.request)?;
if let OcrHostData::Projected(handles) = &self.data {
if let Some(reader) = &handles.reader {
reader.traverse(visit)?;
}
if let Some(provider) = &handles.azure_ad_token_provider {
provider.traverse(visit)?;
}
if let OcrHostData::Projected(handles) = &self.data
&& let Some(provider) = &handles.azure_ad_token_provider
{
provider.traverse(visit)?;
}
Ok(())
}
@ -205,20 +194,13 @@ del provider
.unwrap()
.cast_into::<PyDict>()
.unwrap();
let mut host = OcrRouteHost::new(py.None());
let projected = host.invoke(py, &kwargs, OcrOp::ProjectRequest).unwrap();
assert!(matches!(
projected,
OcrOpResult::Request {
caller_token: true,
..
}
));
let mut host = OcrPythonHost::new(py.None());
assert!(host.project(py, &kwargs).unwrap().caller_token);
locals.del_item("kwargs").unwrap();
drop(kwargs);
let (reply, _) = litellm_host::host::reply();
assert_eq!(
host.invoke(py, &PyDict::new(py), OcrOp::AcquireAzureAdToken)
.is_ok(),
host.invoke(py, OcrOp::AcquireAzureAdToken(reply)).is_ok(),
succeeds
);
let alive = || {

View file

@ -5,11 +5,12 @@ mod project;
use std::sync::LazyLock;
use host::OcrRouteHost;
use host::OcrPythonHost;
use litellm_auth_gcp::VertexAuth;
use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call};
use litellm_core::ocr::route::ocr_machine;
use litellm_core::ocr::{provider_config, route::ocr_machine};
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_host_python::to_py;
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
use pyo3::{
prelude::*,
@ -68,7 +69,7 @@ fn run_ocr(
if asynchronous { ASYNC_SURFACE } else { SURFACE },
PublicCall::capture(&request, &args, &kwargs)?,
crate::logger::LoggedMachine::new(ocr_machine(client)),
OcrRouteHost::new(request.unbind()),
OcrPythonHost::new(request.unbind()),
asynchronous,
)
}
@ -106,6 +107,30 @@ pub(crate) fn aocr(
run_ocr(py, request, args, kwargs, true)
}
#[pyfunction]
pub(crate) fn ocr_health_check_document(
py: Python<'_>,
model: &str,
custom_llm_provider: Option<&str>,
) -> PyResult<Py<PyAny>> {
let document = provider_config::get_health_check_document(model, custom_llm_provider)
.map_err(errors::to_pyerr)?;
to_py(py, &document)
}
#[pyfunction]
pub(crate) fn ocr_passthrough_response(
py: Python<'_>,
model: &str,
endpoint: &str,
body: &[u8],
) -> PyResult<Option<Py<PyAny>>> {
provider_config::passthrough_response(model, endpoint, body)
.map_err(errors::to_pyerr)?
.map(|response| to_py(py, &response.into_json()))
.transpose()
}
#[cfg(test)]
mod tests {
use pyo3::prelude::*;

View file

@ -8,19 +8,15 @@ use litellm_llms::base_llm::ocr::error::Error;
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
use serde_json::{Map, Value};
use super::{
document::{FileDocumentInput, PythonFileReader},
errors::to_pyerr as ocr_error_to_pyerr,
};
use super::{document::FileDocumentInput, errors::to_pyerr as ocr_error_to_pyerr};
use crate::{
credentials::{self, CallerTokenProvider},
marshal::{project_optional_fields, python_timeout_seconds, request_input_sources},
};
/// What the host keeps after projection: the caller's callables that answer the document
/// read and token operations, and the provider name the failure mapping reports.
/// What the host keeps after projection: the caller's token callable that answers the
/// token operation, and the provider name the failure mapping reports.
pub(super) struct OcrHostHandles {
pub reader: Option<PythonFileReader>,
pub azure_ad_token_provider: Option<CallerTokenProvider>,
pub provider: &'static str,
}
@ -104,13 +100,11 @@ impl ProjectedDocument {
Ok(Self::File(document.extract()?))
}
fn into_parts(self) -> PyResult<(OcrDocumentInput, Option<PythonFileReader>)> {
/// Reads a file-like document now, so it runs after every other argument was read.
fn resolve(self, py: Python<'_>) -> PyResult<OcrDocumentInput> {
match self {
Self::File(FileDocumentInput { input, reader }) => Ok((input, reader)),
Self::Other(wire) => Ok((
decode_document(wire).map_err(ocr_error_to_pyerr)?.into(),
None,
)),
Self::File(file) => file.resolve(py),
Self::Other(wire) => Ok(decode_document(wire).map_err(ocr_error_to_pyerr)?.into()),
}
}
}
@ -136,24 +130,25 @@ pub(super) fn project_request(
.chain(["api_key", "api_base", "extra_headers"]),
)?;
let azure_ad_token_provider = credentials::azure_ad_token_provider(kwargs)?;
let (document, reader) = document.into_parts()?;
let api_base = arguments.api_base()?;
let extra_headers = arguments.extra_headers()?;
let timeout_seconds = arguments.timeout_seconds()?;
let wire = OcrWireRequest {
model,
document,
document: document.resolve(request.py())?,
api_key,
api_base: arguments.api_base()?,
api_base,
custom_llm_provider,
extra_headers: arguments.extra_headers()?,
extra_headers,
optional_params,
input_sources,
timeout_seconds: arguments.timeout_seconds()?,
timeout_seconds,
};
let request = decode_request_input(wire).map_err(ocr_error_to_pyerr)?;
let provider = request.provider_name();
Ok((
request,
OcrHostHandles {
reader,
azure_ad_token_provider,
provider,
},
@ -180,10 +175,8 @@ mod tests {
OcrArguments { request, kwargs }
}
fn project_document(
document: &Bound<'_, PyAny>,
) -> PyResult<(OcrDocumentInput, Option<PythonFileReader>)> {
ProjectedDocument::project(document)?.into_parts()
fn project_document(document: &Bound<'_, PyAny>) -> PyResult<OcrDocumentInput> {
ProjectedDocument::project(document)?.resolve(document.py())
}
fn url_document(url: &str) -> OcrDocumentInput {
@ -342,8 +335,11 @@ kwargs = {}
});
}
/// A reader that rewrites the request while it runs shows which arguments projection
/// read before it and which after: every other argument is read first, and the read
/// happens exactly once.
#[test]
fn document_readers_are_not_consumed_during_projection() {
fn document_readers_are_read_once_after_every_other_argument() {
Python::initialize();
Python::attach(|py| {
stub_timeout_conversion(py);
@ -351,17 +347,24 @@ kwargs = {}
py,
c"
class Request:
api_base = 'original'
model = 'mistral/mistral-ocr-latest'
custom_llm_provider = None
api_key = None
api_base = 'https://original.example.com'
extra_headers = {'x-source': 'original'}
timeout = 1
@property
def document(self):
return document
class Reader:
reads = 0
def read(self):
Request.api_base = 'mutated'
Reader.reads += 1
Request.api_base = 'https://mutated.example.com'
Request.extra_headers = {'x-source': 'mutated'}
Request.timeout = 9
return b'abc'
document = {'type': 'file', 'file': Reader()}
document = {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'}
request = Request()
kwargs = {}
",
@ -373,15 +376,38 @@ kwargs = {}
.unwrap()
.cast_into::<PyDict>()
.unwrap();
let arguments = arguments(&request, &kwargs);
let document = arguments.document().unwrap();
let (input, reader) = project_document(&document).unwrap();
assert_eq!(input, OcrDocumentInput::HostReader { mime_type: None });
assert_eq!(arguments.api_base().unwrap().as_deref(), Some("original"));
assert_eq!(arguments.timeout_seconds().unwrap(), Some(1.0));
reader.unwrap().read(py).unwrap();
assert_eq!(arguments.api_base().unwrap().as_deref(), Some("mutated"));
assert_eq!(arguments.timeout_seconds().unwrap(), Some(9.0));
let (projected, _) = project_request(&request, &kwargs).unwrap();
assert_eq!(
py.eval(c"Reader.reads", Some(&locals), Some(&locals))
.unwrap()
.extract::<usize>()
.unwrap(),
1
);
assert_eq!(
projected.document,
OcrDocumentInput::Bytes {
bytes: b"abc".as_slice().into(),
file_name: None,
mime_type: Some("application/pdf".into()),
}
);
assert_eq!(
projected
.credentials
.api_base
.as_ref()
.map(|base| base.value().as_str()),
Some("https://original.example.com")
);
assert_eq!(
projected.transport.extra_headers,
[("x-source".to_string(), "original".to_string())]
);
assert_eq!(
projected.transport.timeout,
Some(std::time::Duration::from_secs(1))
);
});
}
@ -396,16 +422,14 @@ kwargs = {}
None,
)
.unwrap();
let (input, reader) = project_document(&file).unwrap();
assert_eq!(
input,
project_document(&file).unwrap(),
OcrDocumentInput::Bytes {
bytes: b"%PDF-1.4".as_slice().into(),
file_name: None,
mime_type: Some("application/pdf".into()),
}
);
assert!(reader.is_none());
let original = py
.eval(
@ -414,8 +438,10 @@ kwargs = {}
None,
)
.unwrap();
let (input, _) = project_document(&original).unwrap();
assert_eq!(input, url_document("https://example.com/a.pdf"));
assert_eq!(
project_document(&original).unwrap(),
url_document("https://example.com/a.pdf")
);
});
}
@ -617,7 +643,7 @@ document = Document()
",
);
let document = locals.get_item("document").unwrap().unwrap();
let (input, _) = project_document(&document).unwrap();
let input = project_document(&document).unwrap();
assert!(matches!(input, OcrDocumentInput::Bytes { .. }));
let reads: Vec<String> = document.getattr("reads").unwrap().extract().unwrap();
assert_eq!(reads, ["type", "mime_type", "file"]);

View file

@ -1,6 +1,7 @@
use std::{future::Future, pin::Pin};
use std::{future::Future, pin::Pin, sync::Arc};
use litellm_core_utils::settings::Lookup;
use litellm_host_python::{PythonContext, attach_blocking};
use litellm_secrets::{
Error, ExternalSecretManager, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue,
};
@ -19,6 +20,11 @@ const ENVIRONMENT_FALLBACK_LOG: &str =
/// A secret manager whose reads execute in Python: a custom manager, a legacy compatible
/// client, or a manually assigned SDK client.
pub(crate) struct PythonSecretManager {
client: Arc<PythonClient>,
context: PythonContext,
}
struct PythonClient {
client: Py<PyAny>,
system: Option<KeyManagementSystem>,
settings: Option<Py<PyAny>>,
@ -29,14 +35,20 @@ impl PythonSecretManager {
client: Py<PyAny>,
system: Option<KeyManagementSystem>,
settings: Option<Py<PyAny>>,
context: PythonContext,
) -> Self {
Self {
client,
system,
settings,
client: Arc::new(PythonClient {
client,
system,
settings,
}),
context,
}
}
}
impl PythonClient {
fn read(&self, py: Python<'_>, name: &str) -> PyResult<Option<String>> {
let client = self.client.bind(py);
let kwargs = PyDict::new(py);
@ -76,7 +88,7 @@ fn python_name(system: KeyManagementSystem) -> &'static str {
impl ExternalSecretManager for PythonSecretManager {
fn system(&self) -> KeyManagementSystem {
self.system.unwrap_or(KeyManagementSystem::Custom)
self.client.system.unwrap_or(KeyManagementSystem::Custom)
}
fn read_secret<'a>(
@ -85,18 +97,26 @@ impl ExternalSecretManager for PythonSecretManager {
_settings: &'a KeyManagementSettings,
_environment: &'a (dyn Lookup + Send + Sync),
) -> Pin<Box<dyn Future<Output = Result<Option<Secret>, Error>> + Send + 'a>> {
let client = Arc::clone(&self.client);
let context = self.context.clone();
let name = name.to_owned();
Box::pin(async move {
Python::attach(|py| match self.read(py, name) {
match attach_blocking(context, move |py| match client.read(py, &name) {
Ok(value) => Ok(value.map(SecretValue::new).map(Secret::String)),
// `get_secret` answers a failed manager read from the process environment, but
// only for `Exception`: cancellation and other `BaseException`s propagate.
Err(error) if error.is_instance_of::<PyException>(py) => {
log_environment_fallback(py, name, &error)
log_environment_fallback(py, &name, &error)
.map_err(|error| external_error(py, error))?;
Err(read_error(py, error))
}
Err(error) => Err(external_error(py, error)),
})
.await
{
Ok(result) => result,
Err(error) => Python::attach(|py| Err(external_error(py, error))),
}
})
}
}
@ -116,8 +136,9 @@ fn log_environment_fallback(py: Python<'_>, name: &str, error: &PyErr) -> PyResu
}
#[cfg(test)]
#[allow(clippy::await_holding_lock)]
mod tests {
use std::sync::Arc;
use std::sync::{Arc, Mutex, MutexGuard};
use litellm_secrets::{
FailurePolicy, KeyManagementSettings, KeyManagementSystem, OidcResolver, SecretManager,
@ -126,15 +147,27 @@ mod tests {
use pyo3::{prelude::*, types::PyDict};
use rstest::rstest;
use litellm_host_python::PythonContext;
use super::{HANDLER_MODULE, PythonSecretManager, python_name};
use crate::secrets::python_error;
/// `sys.modules` is interpreter-global, so tests that install or rely on the handler module
/// cannot overlap with any other test on this list.
static HANDLER_LOCK: Mutex<()> = Mutex::new(());
fn handler_guard() -> MutexGuard<'static, ()> {
HANDLER_LOCK.lock().expect("handler lock poisoned")
}
/// A resolver over a Python manager whose reads raise `failure_type`, with the chained
/// exceptions Python attaches, and `fallback` as the process environment.
/// exceptions Python attaches, and `fallback` as the process environment. The returned guard
/// keeps other module-mutating tests out for the lifetime of the returned resolver.
fn failing_resolver(
failure_type: &str,
fallback: Option<&'static str>,
) -> (SecretResolver, Py<PyDict>) {
) -> (SecretResolver, Py<PyDict>, MutexGuard<'static, ()>) {
let handler = handler_guard();
Python::initialize();
let (reader, locals) = Python::attach(|py| {
let locals = PyDict::new(py);
@ -167,6 +200,7 @@ handler.get_secret_from_manager = get_secret_from_manager
locals.get_item("manager").unwrap().unwrap().unbind(),
None,
None,
PythonContext::capture(py).unwrap(),
);
(reader, locals.unbind())
});
@ -179,7 +213,7 @@ handler.get_secret_from_manager = get_secret_from_manager
OidcResolver::default(),
)
.with_failure_policy(FailurePolicy::EnvironmentFallback);
(resolver, locals)
(resolver, locals, handler)
}
#[rstest]
@ -191,7 +225,7 @@ handler.get_secret_from_manager = get_secret_from_manager
#[case] failure_type: &str,
#[case] fallback: Option<&'static str>,
) {
let (resolver, locals) = failing_resolver(failure_type, fallback);
let (resolver, locals, _handler) = failing_resolver(failure_type, fallback);
let error = resolver.get_secret("API_KEY", None).await.unwrap_err();
Python::attach(|py| {
let original = python_error(py, &error).unwrap();
@ -264,7 +298,7 @@ sys.modules.setdefault('litellm._logging', logging)
#[case] fallback: Option<&'static str>,
#[case] name: &str,
) {
let (resolver, _locals) = failing_resolver(failure_type, fallback);
let (resolver, _locals, _handler) = failing_resolver(failure_type, fallback);
Python::attach(|py| assert!(logged_errors(py, name).is_empty()));
let secret = resolver.get_secret(name, None).await.unwrap();
assert_eq!(
@ -290,6 +324,7 @@ sys.modules.setdefault('litellm._logging', logging)
/// Installs a fake `get_secret_from_manager` that records its kwargs, runs `body`, and
/// removes the fake handler again; parent package stubs persist for concurrent tests.
/// Callers hold `handler_guard` before attaching so the GIL is never held while waiting on it.
fn with_fake_handler<'py>(py: Python<'py>, body: impl FnOnce(&Bound<'py, PyDict>)) {
let locals = PyDict::new(py);
py.run(
@ -330,6 +365,7 @@ else:
#[case("123")]
#[case("{'key': 'value'}")]
fn nonstring_results_are_absent_without_a_read_failure(#[case] expression: &str) {
let _handler = handler_guard();
Python::initialize();
Python::attach(|py| {
with_fake_handler(py, |locals| {
@ -340,8 +376,13 @@ else:
Some(locals),
)
.unwrap();
let reader = PythonSecretManager::new(py.None(), None, None);
assert_eq!(reader.read(py, "KEY").unwrap(), None);
let reader = PythonSecretManager::new(
py.None(),
None,
None,
PythonContext::capture(py).unwrap(),
);
assert_eq!(reader.client.read(py, "KEY").unwrap(), None);
});
});
}
@ -365,6 +406,7 @@ else:
#[test]
fn configured_systems_dispatch_through_the_python_handler_with_the_original_settings() {
let _handler = handler_guard();
Python::initialize();
Python::attach(|py| {
with_fake_handler(py, |locals| {
@ -374,9 +416,10 @@ else:
client.clone().unbind(),
Some(KeyManagementSystem::AzureKeyVault),
Some(settings.clone().unbind()),
PythonContext::capture(py).unwrap(),
);
assert_eq!(
reader.read(py, "API_KEY").unwrap().as_deref(),
reader.client.read(py, "API_KEY").unwrap().as_deref(),
Some("handled-API_KEY")
);
assert!(py.import(HANDLER_MODULE).is_ok());
@ -416,6 +459,7 @@ else:
#[case] system: Option<KeyManagementSystem>,
#[case] key_manager: &str,
) {
let _handler = handler_guard();
Python::initialize();
Python::attach(|py| {
with_fake_handler(py, |locals| {
@ -434,9 +478,14 @@ manager = Manager()
)
.unwrap();
let manager = locals.get_item("manager").unwrap().unwrap();
let reader = PythonSecretManager::new(manager.clone().unbind(), system, None);
let reader = PythonSecretManager::new(
manager.clone().unbind(),
system,
None,
PythonContext::capture(py).unwrap(),
);
assert_eq!(
reader.read(py, "API_KEY").unwrap().as_deref(),
reader.client.read(py, "API_KEY").unwrap().as_deref(),
Some("handled-API_KEY")
);
assert_eq!(

View file

@ -5,6 +5,8 @@ use litellm_secrets_types::{AccessMode, KeyManagementSettings, KeyManagementSyst
use pyo3::prelude::*;
use serde_json::Value;
use litellm_host_python::PythonContext;
use super::callback::PythonSecretManager;
use crate::{
coercion::{Field, FieldSpec, ProjectionError},
@ -87,7 +89,7 @@ pub(crate) struct SecretManagerSnapshot {
}
impl SecretManagerSnapshot {
pub(crate) fn into_state(self) -> Arc<SecretManagerState> {
pub(crate) fn into_state(self, context: PythonContext) -> Arc<SecretManagerState> {
match self.client {
SecretManagerClient::Native(backend) => {
Arc::new(SecretManagerState::new(*backend, self.settings))
@ -98,6 +100,7 @@ impl SecretManagerSnapshot {
client,
self.system,
self.settings_object,
context,
))),
self.settings,
)),

View file

@ -4,102 +4,30 @@ mod error;
mod mutation;
mod operations;
mod provider;
mod python;
pub(crate) mod resolved;
pub(crate) mod runtime;
mod vault;
use std::sync::Arc;
use litellm_secrets::source::{EnvironmentSecrets, SecretSource};
use pyo3::prelude::*;
pub(crate) use error::python_error;
use litellm_secrets::source::SecretSource;
use pyo3::prelude::*;
use python::PythonSecrets;
use resolved::ResolvedSecrets;
use crate::{
coercion::FieldSpec,
errors::RustBridgeDeclined,
python_settings::{PythonSettings, Snapshot},
};
use crate::{coercion::FieldSpec, python_settings::PythonSettings};
const READABLE: FieldSpec<bool> = FieldSpec::new("readable", |field| field.schema_bool());
const NATIVE: FieldSpec<bool> = FieldSpec::new("native", |field| field.schema_bool());
/// Where a Rust route reads provider secrets from, as `litellm.get_secret` would.
/// Where a Rust route reads provider secrets from. Python's `get_secret_str` until a
/// `SecretManagerRule` in `catalog.py` moves the configured system off `PYTHON_ONLY`, then the
/// native secret manager.
pub(crate) fn source(py: Python<'_>) -> PyResult<Arc<dyn SecretSource>> {
select(&PythonSettings::SecretManager.read(py)?, || {
Ok(Arc::new(ResolvedSecrets::new(config::read(py)?)))
})
}
fn select(
manager: &Snapshot<'_>,
resolved: impl FnOnce() -> PyResult<Arc<dyn SecretSource>>,
) -> PyResult<Arc<dyn SecretSource>> {
if !manager.read(&READABLE)? {
return Ok(Arc::new(EnvironmentSecrets::python_compatible()));
}
if !manager.read(&NATIVE)? {
return Err(RustBridgeDeclined::new_err(
"the configured secret manager is not enabled for the Rust bridge",
));
}
resolved()
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use litellm_secrets::source::{EnvironmentSecrets, SecretSource};
use pyo3::{prelude::*, types::PyDict};
use rstest::rstest;
use super::select;
use crate::{errors::RustBridgeDeclined, python_settings::PythonSettings};
enum Selected {
Environment,
Declined,
Resolved,
}
#[rstest]
#[case::unreadable(false, false, Selected::Environment)]
#[case::unreadable_even_if_native(false, true, Selected::Environment)]
#[case::readable_python_only(true, false, Selected::Declined)]
#[case::readable_native(true, true, Selected::Resolved)]
fn readable_and_native_select_the_secret_source(
#[case] readable: bool,
#[case] native: bool,
#[case] expected: Selected,
) {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
locals.set_item("readable", readable).unwrap();
locals.set_item("native", native).unwrap();
let manager = py
.eval(
c"__import__('types').SimpleNamespace(readable=readable, native=native)",
None,
Some(&locals),
)
.unwrap();
let mut resolved_called = false;
let selected = select(&PythonSettings::SecretManager.snapshot(manager), || {
resolved_called = true;
Ok(Arc::new(EnvironmentSecrets::python_compatible()) as Arc<dyn SecretSource>)
});
match expected {
Selected::Environment => assert!(selected.is_ok() && !resolved_called),
Selected::Resolved => assert!(selected.is_ok() && resolved_called),
Selected::Declined => {
let error = selected.err().expect("the Rust route declines");
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
assert!(!resolved_called);
}
}
});
}
if PythonSettings::SecretManager.read(py)?.read(&NATIVE)? {
let context = litellm_host_python::PythonContext::capture(py)?;
return Ok(Arc::new(ResolvedSecrets::new(config::read(py)?, context)));
}
Ok(Arc::new(PythonSecrets::new(py)?))
}

View file

@ -0,0 +1,190 @@
use std::sync::Arc;
use futures_util::future::BoxFuture;
use litellm_host_python::{PythonContext, attach_blocking};
use litellm_secrets::{Error, SecretValue, source::SecretSource};
use pyo3::prelude::*;
use super::error::external_error;
/// Reads each secret through Python's `get_secret_str`, so the configured manager, the key
/// management settings and the environment fallback behave exactly as they do in Python.
pub(super) struct PythonSecrets {
get_secret_str: Arc<Py<PyAny>>,
context: PythonContext,
}
impl PythonSecrets {
pub(super) fn new(py: Python<'_>) -> PyResult<Self> {
Ok(Self::reading_with(
py.import("litellm.secret_managers.main")?
.getattr("get_secret_str")?
.unbind(),
PythonContext::capture(py)?,
))
}
fn reading_with(get_secret_str: Py<PyAny>, context: PythonContext) -> Self {
Self {
get_secret_str: Arc::new(get_secret_str),
context,
}
}
}
impl SecretSource for PythonSecrets {
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> BoxFuture<'a, Result<Option<SecretValue>, Error>> {
let get_secret_str = Arc::clone(&self.get_secret_str);
let context = self.context.clone();
let name = name.to_owned();
Box::pin(async move {
match attach_blocking(context, move |py| {
get_secret_str
.bind(py)
.call1((name,))
.and_then(|value| value.extract::<Option<String>>())
.map(|value| value.map(SecretValue::new))
.map_err(|error| external_error(py, error))
})
.await
{
Ok(result) => result,
Err(error) => Python::attach(|py| Err(external_error(py, error))),
}
})
}
}
#[cfg(test)]
mod tests {
use litellm_secrets::source::SecretSource;
use pyo3::{prelude::*, types::PyDict};
use rstest::{fixture, rstest};
use super::PythonSecrets;
use crate::secrets::python_error;
use litellm_host_python::PythonContext;
#[fixture]
fn namespace() -> Py<PyDict> {
Python::initialize();
Python::attach(|py| {
let namespace = PyDict::new(py);
py.run(
c"
import contextvars
import threading
read_on = None
request_var = contextvars.ContextVar('request_var', default=None)
seen_context_values = []
raised = KeyboardInterrupt('secret manager stopped')
def get_secret_str(name):
global read_on
read_on = threading.get_ident()
seen_context_values.append(request_var.get())
if name == 'RAISING':
raise raised
return {'MISTRAL_API_KEY': 'vault-key'}.get(name)
",
Some(&namespace),
None,
)
.unwrap();
namespace.unbind()
})
}
#[fixture]
fn secrets(namespace: Py<PyDict>) -> (PythonSecrets, Py<PyDict>) {
let (reader, context) = Python::attach(|py| {
let namespace = namespace.bind(py);
namespace
.get_item("request_var")
.unwrap()
.unwrap()
.call_method1("set", ("request-value",))
.unwrap();
(
namespace
.get_item("get_secret_str")
.unwrap()
.unwrap()
.unbind(),
PythonContext::capture(py).unwrap(),
)
});
(PythonSecrets::reading_with(reader, context), namespace)
}
fn global<T: for<'a, 'py> FromPyObject<'a, 'py, Error: std::fmt::Debug>>(
namespace: &Py<PyDict>,
py: Python<'_>,
name: &str,
) -> T {
namespace
.bind(py)
.get_item(name)
.unwrap()
.unwrap()
.extract()
.unwrap()
}
#[rstest]
#[case::found("MISTRAL_API_KEY", Some("vault-key"))]
#[case::missing("OTHER", None)]
#[tokio::test]
async fn returns_what_get_secret_str_returns(
secrets: (PythonSecrets, Py<PyDict>),
#[case] name: &str,
#[case] expected: Option<&str>,
) {
let value = secrets.0.get_secret_str(name).await.unwrap();
assert_eq!(value.as_ref().map(|value| value.expose()), expected);
}
#[rstest]
#[tokio::test]
async fn exceptions_surface_as_the_original_python_object(
secrets: (PythonSecrets, Py<PyDict>),
) {
let error = secrets.0.get_secret_str("RAISING").await.unwrap_err();
Python::attach(|py| {
let surfaced = python_error(py, &error).expect("the Python exception is preserved");
let raised: Py<PyAny> = global(&secrets.1, py, "raised");
assert!(surfaced.value(py).is(raised.bind(py)));
});
}
#[rstest]
#[tokio::test]
async fn reads_run_off_the_thread_polling_the_route(secrets: (PythonSecrets, Py<PyDict>)) {
let polling: u64 = Python::attach(|py| {
py.import("threading")
.unwrap()
.call_method0("get_ident")
.unwrap()
.extract()
.unwrap()
});
secrets.0.get_secret_str("MISTRAL_API_KEY").await.unwrap();
let read_on: u64 = Python::attach(|py| global(&secrets.1, py, "read_on"));
assert_ne!(read_on, polling);
}
#[rstest]
#[tokio::test]
async fn reads_see_the_callers_contextvars(secrets: (PythonSecrets, Py<PyDict>)) {
secrets.0.get_secret_str("MISTRAL_API_KEY").await.unwrap();
let seen: Vec<String> = Python::attach(|py| global(&secrets.1, py, "seen_context_values"));
assert_eq!(seen, vec!["request-value".to_owned()]);
}
}

View file

@ -2,6 +2,7 @@ use std::sync::Arc;
use futures_util::future::BoxFuture;
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_host_python::PythonContext;
use litellm_secrets::source::SecretSource;
use litellm_secrets::{
Error, FailurePolicy, OidcResolver, SecretManagerState, SecretResolver, SecretValue,
@ -14,8 +15,8 @@ pub(crate) struct ResolvedSecrets {
}
impl ResolvedSecrets {
pub(crate) fn new(snapshot: SecretManagerSnapshot) -> Self {
Self::from_state(snapshot.into_state())
pub(crate) fn new(snapshot: SecretManagerSnapshot, context: PythonContext) -> Self {
Self::from_state(snapshot.into_state(context))
}
fn from_state(state: Arc<SecretManagerState>) -> Self {

View file

@ -657,6 +657,7 @@ azure_anthropic_models: Set = set()
azure_text_models: Set = set()
anyscale_models: Set = set()
cerebras_models: Set = set()
nadir_models: Set = set() # mutable-ok: provider registry, filled from model_cost at import like every sibling provider
galadriel_models: Set = set()
nvidia_nim_models: Set = set()
nvidia_riva_models: Set = set()
@ -893,6 +894,8 @@ def _populate_provider_model_sets(model_cost_map: Dict) -> None:
anyscale_models.add(key)
elif value.get("litellm_provider") == "cerebras":
cerebras_models.add(key)
elif value.get("litellm_provider") == "nadir":
nadir_models.add(key)
elif value.get("litellm_provider") == "galadriel":
galadriel_models.add(key)
elif value.get("litellm_provider") == "nvidia_nim":
@ -1083,6 +1086,7 @@ model_list = list(
| azure_anthropic_models
| anyscale_models
| cerebras_models
| nadir_models
| galadriel_models
| nvidia_nim_models
| nvidia_riva_models
@ -1191,6 +1195,7 @@ def _build_models_by_provider() -> dict:
"azure_text": azure_text_models,
"anyscale": anyscale_models,
"cerebras": cerebras_models,
"nadir": nadir_models,
"galadriel": galadriel_models,
"nvidia_nim": nvidia_nim_models,
"nvidia_riva": nvidia_riva_models,
@ -1994,6 +1999,7 @@ if TYPE_CHECKING:
FeatherlessAIConfig as FeatherlessAIConfig,
)
from .llms.cerebras.chat import CerebrasConfig as CerebrasConfig
from .llms.nadir.chat.transformation import NadirConfig as NadirConfig
from .llms.baseten.chat import BasetenConfig as BasetenConfig
from .llms.sambanova.chat import SambanovaConfig as SambanovaConfig
from .llms.sambanova.embedding.transformation import (

View file

@ -264,6 +264,7 @@ LLM_CONFIG_NAMES: Final = (
"NvidiaNimEmbeddingConfig",
"FeatherlessAIConfig",
"CerebrasConfig",
"NadirConfig",
"BasetenConfig",
"SambanovaConfig",
"SambaNovaEmbeddingConfig",
@ -394,12 +395,10 @@ UTILS_MODULE_NAMES: Final = (
"redact_message_input_output_from_logging",
"CustomStreamWrapper",
"BaseGoogleGenAIGenerateContentConfig",
"BaseOCRConfig",
"BaseSearchConfig",
"BaseTextToSpeechConfig",
"BedrockModelInfo",
"CohereModelInfo",
"MistralOCRConfig",
"Rules",
"AsyncHTTPHandler",
"HTTPHandler",
@ -1063,6 +1062,7 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
"FeatherlessAIConfig",
),
"CerebrasConfig": (".llms.cerebras.chat", "CerebrasConfig"),
"NadirConfig": (".llms.nadir.chat.transformation", "NadirConfig"),
"BasetenConfig": (".llms.baseten.chat", "BasetenConfig"),
"SambanovaConfig": (".llms.sambanova.chat", "SambanovaConfig"),
"SambaNovaEmbeddingConfig": (
@ -1367,7 +1367,6 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
"litellm.llms.base_llm.google_genai.transformation",
"BaseGoogleGenAIGenerateContentConfig",
),
"BaseOCRConfig": ("litellm.llms.base_llm.ocr.transformation", "BaseOCRConfig"),
"BaseSearchConfig": (
"litellm.llms.base_llm.search.transformation",
"BaseSearchConfig",
@ -1378,7 +1377,6 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
),
"BedrockModelInfo": ("litellm.llms.bedrock.common_utils", "BedrockModelInfo"),
"CohereModelInfo": ("litellm.llms.cohere.common_utils", "CohereModelInfo"),
"MistralOCRConfig": ("litellm.llms.mistral.ocr.transformation", "MistralOCRConfig"),
"Rules": ("litellm.litellm_core_utils.rules", "Rules"),
"AsyncHTTPHandler": ("litellm.llms.custom_httpx.http_handler", "AsyncHTTPHandler"),
"HTTPHandler": ("litellm.llms.custom_httpx.http_handler", "HTTPHandler"),

View file

@ -167,6 +167,9 @@ MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_M
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600"))
MCP_SSO_ASSERTION_CACHE_TTL_SECONDS: Final = int(os.getenv("MCP_SSO_ASSERTION_CACHE_TTL_SECONDS", "60"))
# mcp_tool_permissions entry that grants every current and future tool on a server
MCP_ALL_TOOLS_WILDCARD: Final = "*"
# Default npm cache directory for STDIO MCP servers.
# npm/npx needs a writable cache dir; in containers the default (~/.npm)
# may not exist or be read-only. /tmp is always writable.
@ -327,6 +330,7 @@ REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float(
WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
DEEPGRAM_DEFAULT_API_BASE: Final = "https://api.deepgram.com/v1"
NADIR_DEFAULT_API_BASE: Final = "https://api.getnadir.com/v1"
DEEPGRAM_LISTEN_DEFAULT_MODEL: Final = "nova-3"
BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY: Final = "litellm.bedrock_realtime.pending_session_update"
@ -557,6 +561,8 @@ SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS: Final[float] = float(
request_timeout: float = float(os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS))))
request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ
DEFAULT_A2A_AGENT_TIMEOUT: Final[float] = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes
AGENT_KILL_SWITCH_TIMEOUT_SECONDS: Final = 10.0
AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS: Final = 2000
# Patterns that indicate a localhost/internal URL in A2A agent cards that should be
# replaced with the original base_url. This is a common misconfiguration where
# developers deploy agents with development URLs in their agent cards.
@ -592,6 +598,7 @@ FIREWORKS_AI_DEFAULT_CACHE_READ_RATE_RATIO: Final = 0.5
#### Logging callback constants ####
REDACTED_BY_LITELM_STRING: Final = "REDACTED_BY_LITELM"
MAX_LANGFUSE_INITIALIZED_CLIENTS: Final = int(os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50))
LANGFUSE_SHUTDOWN_FLUSH_TIMEOUT_MILLIS: Final = 10_000
# Backpressure + lifetime bounds for the /v1/messages streaming relay (see
# BaseAnthropicMessagesStreamingIterator.async_sse_wrapper). The relay queue is
# bounded so a slow client throttles the upstream pump instead of letting it
@ -706,6 +713,7 @@ LITELLM_CHAT_PROVIDERS: Final = [
"gigachat",
"nvidia_nim",
"cerebras",
"nadir",
"baseten",
"ai21_chat",
"volcengine",
@ -899,6 +907,7 @@ openai_compatible_endpoints: Final[list] = [
"codestral.mistral.ai/v1/fim/completions",
"api.groq.com/openai/v1",
"https://integrate.api.nvidia.com/v1",
NADIR_DEFAULT_API_BASE,
"api.deepseek.com/v1",
"api.together.ai/v1",
"api.together.xyz/v1",
@ -2120,6 +2129,14 @@ PTU_LAPSED_ALERT_LIMIT: Final[int] = 10
DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID: Final[str] = "daily_global_spend_reconcile_job"
DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS: Final[int] = 3600
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM: Final[str] = "daily_global_spend_reconciled_through"
SPEND_CAPTURE_RATE_CHECK_JOB_ID: Final[str] = "spend_capture_rate_check_job"
SPEND_CAPTURE_RATE_CHECK_LOCK_TTL_SECONDS: Final[int] = 900
SPEND_CAPTURE_RATE_MAX_RANGE_DAYS: Final[int] = 180
SPEND_CAPTURE_RATE_DOCS_URL: Final[str] = "https://docs.litellm.ai/docs/proxy/spend_capture_rate"
OPENAI_ORGANIZATION_COSTS_URL: Final[str] = "https://api.openai.com/v1/organization/costs"
# Buckets per page the OpenAI costs endpoint allows (1 to 180, default 7), 2026-09-24
OPENAI_ORGANIZATION_COSTS_PAGE_LIMIT: Final[int] = 180
PROVIDER_BILLING_TIMEOUT_SECONDS: Final[float] = 30.0
# Slack allowed when deciding a sentinel row is stale. The row's updated_at and the
# run's cutoff are stamped by different hosts, so clock skew between them must not let
# one run delete a charge another just wrote. A stale row is hours old and a concurrent

View file

@ -790,6 +790,16 @@ def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None
return value if isinstance(value, str) and value else None
_NON_TOKEN_RATE_FIELDS: Final = frozenset({"input_cost_per_second", "input_cost_per_query", "tiered_pricing"})
def _cost_map_entry_prices_anything(entry: Mapping[str, object]) -> bool:
return any(
value is not None and (field in _NON_TOKEN_RATE_FIELDS or ("cost_per" in field and "token" in field))
for field, value in entry.items()
)
def _select_model_name_for_cost_calc(
model: str | None,
completion_response: object | None,
@ -828,12 +838,7 @@ def _select_model_name_for_cost_calc(
if custom_pricing is True:
if router_model_id is not None and router_model_id in litellm.model_cost:
entry: Final = litellm.model_cost[router_model_id]
if (
entry.get("input_cost_per_token") is not None
or entry.get("input_cost_per_second") is not None
or entry.get("input_cost_per_query") is not None
or entry.get("tiered_pricing") is not None
):
if _cost_map_entry_prices_anything(entry):
return_model = router_model_id
else:
return_model = model
@ -1699,6 +1704,8 @@ def completion_cost(
litellm_model_name=model,
data_residency=data_residency,
litellm_logging_obj=litellm_logging_obj,
custom_pricing_model=selected_model if custom_pricing else None,
base_pricing_model=(selected_model if base_model is not None and not custom_pricing else None),
)
elif call_type == _MCP_CALL_TYPE:
from litellm.proxy._experimental.mcp_server.cost_calculator import (
@ -2870,14 +2877,20 @@ def _candidate_realtime_token_costs(
def _cost_map_entry_declares_pricing(model_name: str, custom_llm_provider: str) -> bool:
"""Whether the entry behind ``model_name`` sets any rate of its own, even a zero one.
The name is resolved the way ``get_model_info`` resolves it before the raw entry is read,
because a deployment-scoped name arrives here already carrying its provider prefix. Two raw
lookups cannot strip that prefix, so a zero-rated override read as declaring nothing, and a
session that should bill nothing fell through to the public rates instead.
"""
resolved: Final = _get_model_info_or_none(model_name, custom_llm_provider)
entries: Final = (
litellm.model_cost.get(resolved.get("key")) if resolved is not None else None,
litellm.model_cost.get(model_name),
litellm.model_cost.get(f"{custom_llm_provider}/{model_name}"),
)
return any(
entry is not None and any("cost_per" in field and value is not None for field, value in entry.items())
for entry in entries
)
return any(entry is not None and _cost_map_entry_prices_anything(entry) for entry in entries)
def _first_priced_realtime_token_costs(
@ -2917,6 +2930,8 @@ def handle_realtime_stream_cost_calculation(
litellm_model_name: str,
data_residency: str | None = None,
litellm_logging_obj: LitellmLoggingObject | None = None,
custom_pricing_model: str | None = None,
base_pricing_model: str | None = None,
) -> float:
"""
Handles the cost calculation for realtime stream responses.
@ -2925,9 +2940,13 @@ def handle_realtime_stream_cost_calculation(
Args:
results: A list of OpenAIRealtimeStreamBaseObject objects
custom_pricing_model: deployment-scoped pricing key from the deployment's
custom rates, tried ahead of the session-reported model
base_pricing_model: the deployment's resolved base_model, tried ahead of the
session-reported model but after custom rates
"""
received_model = None
potential_model_names: Final = []
potential_model_names: Final = [custom_pricing_model, base_pricing_model]
for result in results:
if result["type"] == "session.created":
received_model = cast(OpenAIRealtimeStreamSessionEvents, result)["session"].get("model", None)
@ -2945,6 +2964,7 @@ def handle_realtime_stream_cost_calculation(
results=results,
custom_llm_provider=custom_llm_provider,
litellm_model_name=litellm_model_name,
custom_pricing_model=custom_pricing_model,
)
if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results)
else 0.0
@ -2968,6 +2988,7 @@ def handle_realtime_transcription_cost_calculation(
results: OpenAIRealtimeStreamList,
custom_llm_provider: str,
litellm_model_name: str,
custom_pricing_model: str | None = None,
) -> float:
"""
Cost for realtime transcription sessions (e.g. gpt-realtime-whisper).
@ -2985,15 +3006,15 @@ def handle_realtime_transcription_cost_calculation(
return 0.0
model_name: Final = _get_transcription_model_name_from_results(results) or litellm_model_name
try:
model_info = litellm.get_model_info(model=model_name, custom_llm_provider=custom_llm_provider)
except Exception:
model_info = None
model_info: Final = _get_model_info_or_none(model_name, custom_llm_provider)
override_info: Final = (
_get_model_info_or_none(custom_pricing_model, custom_llm_provider) if custom_pricing_model is not None else None
)
total_cost = 0.0
for event in completed_events:
usage = event.get("usage") or {}
total_cost += _transcription_usage_cost(usage, model_info)
total_cost += _transcription_usage_cost(usage, model_info, override_info)
return total_cost
@ -3018,23 +3039,57 @@ def _get_transcription_model_name_from_results(
return None
def _transcription_usage_cost(usage: dict, model_info: ModelInfo | None) -> float:
if model_info is None:
def _get_model_info_or_none(model: str, custom_llm_provider: str) -> ModelInfo | None:
try:
return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception:
return None
def _declared_transcription_rate(info: ModelInfo | None, keys: tuple[str, ...]) -> float | None:
"""First of ``keys`` this entry prices, read off the raw ``litellm.model_cost`` entry
because ``get_model_info`` synthesizes zero token rates for entries that omit them."""
if info is None:
return None
declared: Final = litellm.model_cost.get(info.get("key"))
if declared is None:
return None
return next(
(float(value) for key in keys if declared.get(key) is not None and (value := info.get(key)) is not None),
None,
)
def _transcription_rate(keys: tuple[str, ...], override: ModelInfo | None, base: ModelInfo | None) -> float:
rates: Final = (_declared_transcription_rate(info, keys) for info in (override, base))
return next((rate for rate in rates if rate is not None), 0.0)
def _transcription_usage_cost(
usage: dict,
model_info: ModelInfo | None,
override_info: ModelInfo | None = None,
) -> float:
if model_info is None and override_info is None:
return 0.0
usage_type: Final = usage.get("type")
if usage_type == "duration":
seconds: Final = usage.get("seconds") or 0.0
per_second: Final = model_info.get("input_cost_per_second") or 0.0
return float(seconds) * float(per_second)
return float(seconds) * _transcription_rate(("input_cost_per_second",), override_info, model_info)
if usage_type == "tokens":
input_token_details: Final = usage.get("input_token_details") or {}
audio_tokens: Final = input_token_details.get("audio_tokens") or 0
text_tokens: Final = input_token_details.get("text_tokens") or 0
output_tokens: Final = usage.get("output_tokens") or 0
audio_cost: Final = float(audio_tokens) * float(
model_info.get("input_cost_per_audio_token") or model_info.get("input_cost_per_token") or 0.0
audio_cost: Final = float(audio_tokens) * _transcription_rate(
("input_cost_per_audio_token", "input_cost_per_token"), override_info, model_info
)
text_cost: Final = float(text_tokens) * _transcription_rate(
("input_cost_per_token",), override_info, model_info
)
output_cost: Final = float(output_tokens) * _transcription_rate(
("output_cost_per_token",), override_info, model_info
)
text_cost: Final = float(text_tokens) * float(model_info.get("input_cost_per_token") or 0.0)
output_cost: Final = float(output_tokens) * float(model_info.get("output_cost_per_token") or 0.0)
return audio_cost + text_cost + output_cost
return 0.0

View file

@ -118,6 +118,7 @@ class SlackAlerting(CustomBatchLogger):
self.default_webhook_url = default_webhook_url
self.flush_lock = asyncio.Lock()
self.periodic_started = False
self._periodic_flush_task: asyncio.Task[None] | None = None
self.hanging_request_check = AlertingHangingRequestCheck(
slack_alerting_object=self,
)
@ -129,6 +130,12 @@ class SlackAlerting(CustomBatchLogger):
self.digest_lock = asyncio.Lock()
super().__init__(**kwargs, flush_lock=self.flush_lock)
def _ensure_periodic_flush_task(self) -> None:
if self.periodic_started and (self._periodic_flush_task is None or not self._periodic_flush_task.done()):
return
self._periodic_flush_task = asyncio.create_task(self.periodic_flush())
self.periodic_started = True
def update_values(
self,
alerting: list | None = None,
@ -141,17 +148,14 @@ class SlackAlerting(CustomBatchLogger):
):
if alerting is not None:
self.alerting = alerting
asyncio.create_task(self.periodic_flush())
self.periodic_started = True
self._ensure_periodic_flush_task()
if alerting_threshold is not None:
self.alerting_threshold = alerting_threshold
if alert_types is not None:
self.alert_types = alert_types
if alerting_args is not None:
self.alerting_args = SlackAlertingArgs(**alerting_args)
if not self.periodic_started:
asyncio.create_task(self.periodic_flush())
self.periodic_started = True
self._ensure_periodic_flush_task()
if alert_type_config is not None:
for key, val in alert_type_config.items():
self.alert_type_config[key] = AlertTypeConfig(**val) if isinstance(val, dict) else val
@ -1446,9 +1450,8 @@ Model Info:
return
# Start periodic flush if not already started
if not self.periodic_started and self.alerting is not None and len(self.alerting) > 0:
asyncio.create_task(self.periodic_flush())
self.periodic_started = True
if self.alerting is not None and len(self.alerting) > 0:
self._ensure_periodic_flush_task()
if "webhook" in self.alerting and alert_type == "budget_alerts" and user_info is not None:
await self.send_webhook_alert(webhook_event=user_info)

View file

@ -3,9 +3,11 @@ Utils used for slack alerting
"""
import asyncio
from collections.abc import Callable
from typing import TYPE_CHECKING, Any, Final
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import AlertType
from litellm.secret_managers.main import get_secret
@ -66,25 +68,27 @@ async def _add_langfuse_trace_id_to_alert(
-> trace_id
-> litellm_call_id
"""
if "langfuse" not in litellm.logging_callback_manager._get_all_callbacks():
from litellm.integrations.langfuse.langfuse import LangFuseLogger, resolve_langfuse_host
callbacks: Final[list[CustomLogger | Callable[..., object] | str]] = (
litellm.logging_callback_manager._get_all_callbacks()
)
if not any(callback == "langfuse" or isinstance(callback, LangFuseLogger) for callback in callbacks):
return None
#########################################################
# Only run if langfuse is added as a callback
#########################################################
if request_data is not None and request_data.get("litellm_logging_obj", None) is not None:
trace_id: str | None = None
litellm_logging_obj: Final[Logging] = request_data["litellm_logging_obj"]
if request_data is None or request_data.get("litellm_logging_obj", None) is None:
return None
for _ in range(3):
trace_id = litellm_logging_obj._get_trace_id(service_name="langfuse")
if trace_id is not None:
break
await asyncio.sleep(3) # wait 3s before retrying for trace id
#########################################################
langfuse_object: Final = litellm_logging_obj._get_callback_object(service_name="langfuse")
if langfuse_object is not None:
base_url: Final = langfuse_object.Langfuse.base_url
return f"{base_url}/trace/{trace_id}"
litellm_logging_obj: Final[Logging] = request_data["litellm_logging_obj"]
instance_host: Final = next(
(callback.langfuse_host for callback in callbacks if isinstance(callback, LangFuseLogger)), None
)
host: Final = resolve_langfuse_host(
litellm_logging_obj.standard_callback_dynamic_params.get("langfuse_host") or instance_host
)
for _ in range(3):
if (trace_id := litellm_logging_obj._get_trace_id(service_name="langfuse")) is not None:
return f"{host}/trace/{trace_id}"
await asyncio.sleep(3) # wait 3s before retrying for trace id
return None

View file

@ -421,6 +421,24 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm
pass
async def async_filter_listed_models(
self,
user_api_key_dict: UserAPIKeyAuth,
model_names: Sequence[str],
) -> Sequence[str]:
"""Runs on the model listing routes (`/v1/models`, `/v1/models/{id}`, `/model/info`,
`/model_group/info`) with the public model names the route would otherwise return, so a
lookup of one model may offer just that name: decide per name, never by position in the
sequence. Return the names to keep as a sequence of strings; a name left out disappears
from every listing, any alias of it offered in the same call goes with it, and
`/v1/models/{id}` answers 404 for it, exactly as for a model that does not exist. Names
outside `model_names` are ignored, so a callback can only narrow the listing, never widen
it. Under `use_team_public_model_name: false`, `/v1/models` and `/model_group/info` list a
team model by its internal routing name while `/model/info` keeps its public name, so hide
both names to hide it on every route.
"""
return model_names
async def async_post_call_response_headers_hook(
self,
data: dict,

View file

@ -1,14 +1,14 @@
#### What this does ####
# On success, logs events to Langfuse
import inspect
import os
import re
import traceback
from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import datetime
from functools import lru_cache
from importlib.metadata import PackageNotFoundError, version
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, runtime_checkable
from packaging.version import Version
@ -45,13 +45,13 @@ from litellm.types.utils import (
)
if TYPE_CHECKING:
from langfuse.client import Langfuse, StatefulTraceClient
from litellm.integrations.langfuse.langfuse_sdk import LangfuseApiClient, LangfuseObservation, LangfuseTracing
from litellm.litellm_core_utils.litellm_logging import DynamicLoggingCache
else:
DynamicLoggingCache = Any
StatefulTraceClient = Any
Langfuse = Any
LangfuseApiClient = Any
LangfuseObservation = Any
LangfuseTracing = Any
_DENIED_STEERING_KEYS: Final = frozenset({"headers", "endpoint", "caching_groups", "previous_models"})
@ -142,6 +142,20 @@ def _logging_id(start_time: datetime | None, response_obj: object) -> str | None
return litellm.utils.get_logging_id(start_time, response_obj)
@runtime_checkable
class _ResponseWithId(Protocol):
"""Response payloads (ModelResponse and friends, or a plain dict) expose their provider id via ``get``."""
def get(self, key: Literal["id"], default: None = None, /) -> object: ...
def _lookup_ids(litellm_call_id: str | None, response_obj: object) -> Mapping[str, str]:
"""v2 carried the response id inside the generation id; v4 hashes ids to 16 hex chars, so they ride in metadata."""
response_id: Final[object] = response_obj.get("id") if isinstance(response_obj, _ResponseWithId) else None
ids: Final[tuple[tuple[str, object], ...]] = (("litellm_call_id", litellm_call_id), ("response_id", response_id))
return MappingProxyType({key: str(value) for key, value in ids if value is not None})
def _as_steering_flag(value: object) -> bool:
"""A string ``str_to_bool`` does not recognise falls back to its truthiness."""
if isinstance(value, str):
@ -158,6 +172,68 @@ def _as_steering_key_sequence(value: object) -> tuple[str, ...]:
return ()
MINIMUM_LANGFUSE_VERSION: Final = "4.7"
UNSUPPORTED_LANGFUSE_VERSION: Final = "5"
PROMPT_CACHE_TTL_ENV: Final = "LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS"
def installed_langfuse_version() -> str:
"""Only ``importlib.metadata`` reads correctly on every major.
``langfuse.version`` was removed in v4, ``langfuse.__version__`` does not
exist in v3, and in v2 it reports a different value from the distribution
that is actually installed.
"""
return version("langfuse")
def raise_if_unsupported_langfuse_version(installed_version: str) -> None:
"""Fail at logger construction rather than dropping every event at request time.
v4 moved the callback onto OpenTelemetry, so on an older SDK the import of
`LangfuseOtelSpanAttributes` raises inside the per-request handler and the
broad except there turns it into silent total data loss.
"""
installed: Final = Version(installed_version)
# compare majors, not versions: "5.0.0rc1" sorts below "5" but is just as unsupported
if Version(MINIMUM_LANGFUSE_VERSION) <= installed and installed.major < Version(UNSUPPORTED_LANGFUSE_VERSION).major:
return
raise ImportError(
f"\033[91mlitellm requires langfuse>={MINIMUM_LANGFUSE_VERSION},<{UNSUPPORTED_LANGFUSE_VERSION} for the "
f"'langfuse' callback, but {installed_version} is installed. Run "
f"'pip install \"langfuse>={MINIMUM_LANGFUSE_VERSION},<{UNSUPPORTED_LANGFUSE_VERSION}\"' to upgrade, or use "
f"the 'langfuse_otel' callback, which does not depend on the langfuse SDK\033[0m"
)
def whole_number(raw: str) -> int | None:
try:
return int(raw)
except ValueError:
return None
def raise_if_unusable_prompt_cache_ttl() -> None:
"""The v4 SDK runs ``int()`` on this variable while it is being imported, so a value that is not a whole
number has to be named here, before that import fails with a bare ``ValueError`` on every request."""
raw: Final = os.environ.get(PROMPT_CACHE_TTL_ENV)
if raw is None or whole_number(raw) is not None:
return
raise ValueError(f"\033[91m{PROMPT_CACHE_TTL_ENV}={raw!r} must be a whole number of seconds\033[0m")
def _optional_str(value: object) -> str | None:
"""v4 sets attribute values raw; a non-string version would be dropped by the server."""
return str(value) if value is not None else None
def _trace_public_flag(value: object) -> bool | None:
"""``trace_public`` reaches here as a bool from metadata or a string from a ``langfuse_*`` header."""
if value is None:
return None
return _as_steering_flag(value)
def resolve_langfuse_credentials(
langfuse_public_key=None,
langfuse_secret=None,
@ -172,9 +248,29 @@ def resolve_langfuse_credentials(
secret_key = langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY")
public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY")
resolved_host: Final = langfuse_host or os.getenv("LANGFUSE_HOST", "https://cloud.langfuse.com")
return public_key, secret_key, resolve_langfuse_host(langfuse_host)
return public_key, secret_key, resolved_host
def resolve_langfuse_host(langfuse_host: object = None) -> str:
"""The Langfuse base URL for ``langfuse_host`` with the env fallbacks, always carrying a scheme."""
resolved: Final = str(
langfuse_host or os.getenv("LANGFUSE_HOST") or os.getenv("LANGFUSE_BASE_URL") or "https://cloud.langfuse.com"
)
return resolved if resolved.startswith(("http://", "https://")) else f"http://{resolved}"
def warn_if_upstream_langfuse_configured() -> None:
if os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") is None:
return
verbose_logger.warning(
"UPSTREAM_LANGFUSE_* is no longer supported: the langfuse callback moved to SDK v4, "
"which has no second ingestion client. The values are ignored."
)
def parse_langfuse_debug(raw_value: str | None) -> bool:
"""Parse the LANGFUSE_DEBUG value into the boolean flag the langfuse client expects."""
return raw_value is not None and raw_value.strip().lower() in ("true", "1")
@lru_cache(maxsize=8)
@ -199,29 +295,29 @@ class LangFuseLogger:
allow_env_credentials: bool = True,
):
try:
import langfuse
from langfuse import Langfuse
except Exception as e:
self.langfuse_sdk_version: str = installed_langfuse_version()
except PackageNotFoundError as e:
raise Exception(
f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n{traceback.format_exc()}\033[0m"
)
f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\033[0m"
) from e
raise_if_unsupported_langfuse_version(self.langfuse_sdk_version)
raise_if_unusable_prompt_cache_ttl()
from litellm.integrations.langfuse.langfuse_sdk import configured_release
self.public_key, self.secret_key, self.langfuse_host = resolve_langfuse_credentials(
langfuse_public_key=langfuse_public_key,
langfuse_secret=langfuse_secret,
langfuse_host=langfuse_host,
allow_env_credentials=allow_env_credentials,
)
if not (self.langfuse_host.startswith("http://") or self.langfuse_host.startswith("https://")):
# add http:// if unset, assume communicating over private network - e.g. render
self.langfuse_host = "http://" + self.langfuse_host
_env_override: Final = str(langfuse_environment).strip() if langfuse_environment is not None else None
if _env_override:
validate_langfuse_environment_value(_env_override)
self.langfuse_environment: str | None = _env_override
else:
self.langfuse_environment = self.resolve_deployment_environment()
self.langfuse_release = os.getenv("LANGFUSE_RELEASE")
self.langfuse_debug = os.getenv("LANGFUSE_DEBUG")
self.langfuse_release = configured_release()
self.langfuse_debug = parse_langfuse_debug(os.getenv("LANGFUSE_DEBUG"))
self.langfuse_flush_interval = LangFuseLogger._get_langfuse_flush_interval(flush_interval)
if should_use_langfuse_mock():
@ -232,22 +328,9 @@ class LangFuseLogger:
self.langfuse_client = self._http_handler.client
self.is_mock_mode = False
parameters: Final = {
"public_key": self.public_key,
"secret_key": self.secret_key,
"host": self.langfuse_host,
"release": self.langfuse_release,
"debug": self.langfuse_debug,
"flush_interval": self.langfuse_flush_interval, # flush interval in seconds
"httpx_client": self.langfuse_client,
}
self.langfuse_sdk_version: str = langfuse.version.__version__
if "environment" in inspect.signature(Langfuse.__init__).parameters:
parameters["environment"] = self.langfuse_environment
if Version(self.langfuse_sdk_version) >= Version("2.6.0"):
parameters["sdk_integration"] = "litellm"
self.Langfuse: Langfuse = self.safe_init_langfuse_client(parameters)
self.api_client: LangfuseApiClient
self.tracing: LangfuseTracing
self.api_client, self.tracing = self.safe_init_langfuse_client()
# set the current langfuse project id in the environ
# this is used by Alerting to link to the correct project
@ -256,49 +339,62 @@ class LangFuseLogger:
verbose_logger.debug("Langfuse Mock: Using mock project ID")
else:
try:
project_id = self.Langfuse.client.projects.get().data[0].id
os.environ["LANGFUSE_PROJECT_ID"] = project_id
project_id: Final = self.api_client.project_id()
if project_id is not None:
os.environ["LANGFUSE_PROJECT_ID"] = project_id
except Exception:
project_id = None
verbose_logger.debug("Langfuse project id unavailable, alerting links will omit it")
if os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") is not None:
upstream_langfuse_debug_env: Final = os.getenv("UPSTREAM_LANGFUSE_DEBUG")
upstream_langfuse_debug: Final = (
str_to_bool(upstream_langfuse_debug_env) if upstream_langfuse_debug_env is not None else None
)
self.upstream_langfuse_secret_key = os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY")
self.upstream_langfuse_public_key = os.getenv("UPSTREAM_LANGFUSE_PUBLIC_KEY")
self.upstream_langfuse_host = os.getenv("UPSTREAM_LANGFUSE_HOST")
self.upstream_langfuse_release = os.getenv("UPSTREAM_LANGFUSE_RELEASE")
self.upstream_langfuse_debug = upstream_langfuse_debug_env
self.upstream_langfuse = Langfuse(
public_key=self.upstream_langfuse_public_key,
secret_key=self.upstream_langfuse_secret_key,
host=self.upstream_langfuse_host,
release=self.upstream_langfuse_release,
debug=(upstream_langfuse_debug if upstream_langfuse_debug is not None else False),
)
else:
self.upstream_langfuse = None
warn_if_upstream_langfuse_configured()
def safe_init_langfuse_client(self, parameters: dict) -> Langfuse:
def safe_init_langfuse_client(self) -> "tuple[LangfuseApiClient, LangfuseTracing]":
"""Build the REST client and export channel while the process is under its logger budget.
The budget dates from the SDK client, which started a consumer thread per instance and once
pinned a CPU at 100% when many were built; it still bounds the number of per-key loggers.
"""
Safely init a langfuse client if the number of initialized clients is less than the max
Note:
- Langfuse initializes 1 thread everytime a client is initialized.
- We've had an incident in the past where we reached 100% cpu utilization because Langfuse was initialized several times.
"""
from langfuse import Langfuse
if litellm.initialized_langfuse_clients >= MAX_LANGFUSE_INITIALIZED_CLIENTS:
raise Exception(
f"Max langfuse clients reached: {litellm.initialized_langfuse_clients} is greater than {MAX_LANGFUSE_INITIALIZED_CLIENTS}"
)
langfuse_client: Final = Langfuse(**parameters)
from litellm.integrations.langfuse.langfuse_sdk import (
acquire_langfuse_tracing,
build_langfuse_client,
release_langfuse_tracing,
)
tracing: Final = acquire_langfuse_tracing(
public_key=str(self.public_key),
secret_key=str(self.secret_key),
base_url=self.langfuse_host,
environment=self.langfuse_environment,
release=self.langfuse_release,
flush_interval=self.langfuse_flush_interval,
mock_mode=self.is_mock_mode,
)
try:
api_client: Final = build_langfuse_client(
public_key=self.public_key,
secret_key=self.secret_key,
base_url=self.langfuse_host,
httpx_client=self.langfuse_client,
)
except Exception:
release_langfuse_tracing(tracing, grace_seconds=0.0)
raise
litellm.initialized_langfuse_clients += 1
verbose_logger.debug("Created langfuse client number %s", litellm.initialized_langfuse_clients)
return langfuse_client
return api_client, tracing
def flush(self) -> None:
"""Push every queued observation to Langfuse before the process goes away."""
self.tracing.flush()
def stop(self) -> None:
"""Give the export channel back; ``DynamicLoggingCache`` calls this when a per-key logger expires."""
from litellm.integrations.langfuse.langfuse_sdk import release_langfuse_tracing
release_langfuse_tracing(self.tracing)
@staticmethod
def add_metadata_from_header(litellm_params: dict, metadata: dict) -> dict[str, object]:
@ -349,7 +445,7 @@ class LangFuseLogger:
user_id: str | None = None,
level: str = "DEFAULT",
status_message: str | None = None,
) -> dict:
) -> LangfuseLoggedEvent:
"""
Logs a success or error event on Langfuse
"""
@ -411,10 +507,10 @@ class LangFuseLogger:
verbose_logger.debug("Langfuse Layer Logging - final response object: %s", response_obj)
verbose_logger.info("Langfuse Layer Logging - logging success")
return {"trace_id": trace_id, "generation_id": generation_id}
return LangfuseLoggedEvent(trace_id=trace_id, generation_id=generation_id)
except Exception as e:
verbose_logger.exception("Langfuse Layer Error(): Exception occured - %s", e)
return {"trace_id": None, "generation_id": None}
return LangfuseLoggedEvent(trace_id=None, generation_id=None)
def _get_langfuse_input_output_content(
self,
@ -518,18 +614,14 @@ class LangFuseLogger:
level: str,
litellm_call_id: str | None,
) -> tuple:
verbose_logger.debug("Langfuse Layer Logging - logging to langfuse v2")
verbose_logger.debug("Langfuse Layer Logging - logging to langfuse via sdk v%s", self.langfuse_sdk_version)
try:
standard_logging_object: Final[StandardLoggingPayload | None] = cast(
StandardLoggingPayload | None,
kwargs.get("standard_logging_object", None),
)
tags = (
self._get_langfuse_tags(standard_logging_object=standard_logging_object)
if self._supports_tags()
else []
)
tags = self._get_langfuse_tags(standard_logging_object=standard_logging_object)
allowlisted_metadata: Final[StandardLoggingMetadata | Mapping[str, object]] = (
standard_logging_object["metadata"] if standard_logging_object is not None else _NO_METADATA
@ -581,17 +673,17 @@ class LangFuseLogger:
# This allows continuing an existing trace while still returning the correct trace_id
if existing_trace_id is not None:
trace_id = existing_trace_id
resolved_trace_id: Final = (
call_trace_id: Final = (
litellm_call_id or trace_id
if existing_trace_id is None
and _is_session_header_trace(trace_id, session_id, litellm_params.get("proxy_server_request"))
else trace_id
)
if resolved_trace_id != trace_id:
if call_trace_id != trace_id:
verbose_logger.debug(
"Langfuse: trace_id %s came from a session header; using call id %s so each call gets its own trace",
trace_id,
resolved_trace_id,
call_trace_id,
)
requested_trace_keys: Final = _as_steering_key_sequence(clean_metadata.pop("update_trace_keys", ()))
update_trace_keys: Final = (
@ -647,7 +739,7 @@ class LangFuseLogger:
trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm"
else: # don't overwrite an existing trace
trace_params = {
"id": resolved_trace_id,
"id": call_trace_id,
"name": trace_name,
"session_id": session_id,
"input": masked_input if not mask_input else "redacted-by-litellm",
@ -659,10 +751,7 @@ class LangFuseLogger:
for key in list(filter(lambda key: key.startswith("trace_"), clean_metadata.keys())):
trace_params[key.replace("trace_", "")] = clean_metadata.pop(key, None)
if level == "ERROR":
trace_params["status_message"] = masked_output
else:
trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm"
trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm"
if debug is True or (isinstance(debug, str) and debug.lower() == "true"):
debug_metadata: Final = {
@ -697,17 +786,16 @@ class LangFuseLogger:
("api_base", api_base, bool(api_base)),
("vertex_location", vertex_location, bool(vertex_location)),
("aws_region_name", aws_region_name, bool(aws_region_name)),
("cache_hit", kwargs.get("cache_hit") or False, self._supports_tags() and "cache_hit" in kwargs),
("cache_hit", kwargs.get("cache_hit") or False, "cache_hit" in kwargs),
)
enrichments: Final[Mapping[str, object]] = {
key: value for key, value, include in candidate_enrichments if include
}
if self._supports_tags():
if "cache_hit" in kwargs and kwargs["cache_hit"] is None:
kwargs["cache_hit"] = False # rebind-ok: pre-existing normalization other integrations rely on
if existing_trace_id is None:
trace_params.update({"tags": tags})
if "cache_hit" in kwargs and kwargs["cache_hit"] is None:
kwargs["cache_hit"] = False # rebind-ok: pre-existing normalization other integrations rely on
if existing_trace_id is None:
trace_params.update({"tags": tags})
proxy_server_request: Final = litellm_params.get("proxy_server_request", None)
if proxy_server_request:
@ -721,17 +809,6 @@ class LangFuseLogger:
if key.lower() not in _REDACTED_PROXY_HEADERS:
clean_headers[key] = value
trace: Final[StatefulTraceClient] = self.Langfuse.trace(**trace_params)
# Log provider specific information as a span
log_provider_specific_information_as_span(trace, enrichments)
# Log guardrail information as a span
self._log_guardrail_information_as_span(
trace=trace,
standard_logging_object=standard_logging_object,
)
generation_id = None
usage = None
usage_details = None
@ -753,7 +830,7 @@ class LangFuseLogger:
usage = {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_cost": cost if self._supports_costs() else None,
"total_cost": cost,
}
# According to langfuse documentation: "the input value must be reduced by the number of cache_read_input_tokens"
input_tokens: Final = prompt_tokens - cache_read_input_tokens
@ -765,15 +842,15 @@ class LangFuseLogger:
cache_read_input_tokens=cache_read_input_tokens,
)
generation_name = clean_metadata.pop("generation_name", None)
if generation_name is None:
# if `generation_name` is None, use sensible default values
# If using litellm proxy user `key_alias` if not None
# If `key_alias` is None, just log `litellm-{call_type}` as the generation name
_user_api_key_alias: Final = cast(str | None, clean_metadata.get("user_api_key_alias", None))
generation_name = f"litellm-{cast(str, kwargs.get('call_type', 'completion'))}"
if _user_api_key_alias is not None:
generation_name = f"litellm:{_user_api_key_alias}"
requested_generation_name: Final = clean_metadata.pop("generation_name", None)
_user_api_key_alias: Final = cast(str | None, clean_metadata.get("user_api_key_alias", None))
generation_name: Final = (
str(requested_generation_name)
if requested_generation_name is not None
else f"litellm:{_user_api_key_alias}"
if _user_api_key_alias is not None
else f"litellm-{cast(str, kwargs.get('call_type', 'completion'))}"
)
if response_obj is not None:
system_fingerprint = getattr(response_obj, "system_fingerprint", None)
@ -789,53 +866,97 @@ class LangFuseLogger:
generation_params = {
"name": generation_name,
"id": clean_metadata.pop("generation_id", generation_id),
"start_time": start_time,
"end_time": end_time,
"model": model_name,
"model_parameters": optional_params,
"input": masked_input if not mask_input else "redacted-by-litellm",
"output": masked_output if not mask_output else "redacted-by-litellm",
"usage": usage,
"usage_details": usage_details,
"metadata": {
**log_requester_metadata(redact_user_api_key_info(metadata=allowlisted_metadata)),
"cost_details": {"total": cost} # mutable-ok: langfuse serializes this payload
if usage is not None and isinstance(cost, (int, float))
else None,
"metadata": { # mutable-ok: langfuse serializes this payload, a proxy is not json-encodable
**log_requester_metadata(redact_user_api_key_info(metadata=allowlisted_metadata)), # pyright: ignore[reportArgumentType] # TypedDict in, plain metadata dict out
**enrichments,
**_lookup_ids(litellm_call_id, response_obj),
},
"level": level,
"version": clean_metadata.pop("version", None),
"version": _optional_str(clean_metadata.pop("version", None)),
}
parent_observation_id: Final = metadata.get("parent_observation_id", None)
if parent_observation_id is not None:
generation_params["parent_observation_id"] = parent_observation_id
if self._supports_prompt():
generation_params = _add_prompt_to_generation_params(
generation_params=generation_params,
clean_metadata=clean_metadata,
prompt_management_metadata=prompt_management_metadata,
langfuse_client=self.Langfuse,
)
generation_params = _add_prompt_to_generation_params(
generation_params=generation_params,
clean_metadata=clean_metadata,
prompt_management_metadata=prompt_management_metadata,
langfuse_client=self.api_client,
)
if masked_output is not None and isinstance(masked_output, str) and level == "ERROR":
generation_params["status_message"] = masked_output
if self._supports_completion_start_time():
generation_params["completion_start_time"] = kwargs.get("completion_start_time", None)
# langfuse ships in the proxy-runtime extra, so this module must import cleanly without it
from litellm.integrations.langfuse.langfuse_sdk import (
observation_attributes,
resolve_observation_id,
resolve_trace_id,
start_generation,
trace_attributes,
)
generation_client: Final = trace.generation(**generation_params)
resolved_trace_id: Final = resolve_trace_id(call_trace_id) # pyright: ignore[reportArgumentType] # metadata value, str or None at runtime
continued_trace: Final = existing_trace_id is not None
generation_is_trace_root: Final = not continued_trace and parent_observation_id is None
trace_public: Final = _trace_public_flag(trace_params.get("public"))
trace_input: Final = trace_params.get("input")
trace_output: Final = trace_params.get("output")
trace_level_attributes: Final = trace_attributes(
name=trace_params.get("name"),
user_id=trace_params.get("user_id"),
session_id=trace_params.get("session_id"),
version=trace_params.get("version"),
release=trace_params.get("release"),
tags=trace_params.get("tags"),
metadata=trace_params.get("metadata"),
public=trace_public,
input=None if generation_is_trace_root and trace_input == generation_params["input"] else trace_input,
output=None
if generation_is_trace_root and trace_output == generation_params["output"]
else trace_output,
)
generation_attributes: Final = observation_attributes(
observation_type="generation",
input=generation_params["input"],
output=generation_params["output"],
metadata=generation_params["metadata"],
level=level,
status_message=generation_params.get("status_message"),
version=generation_params["version"],
model=model_name,
model_parameters=optional_params,
usage_details=usage_details,
cost_details=generation_params["cost_details"],
completion_start_time=kwargs.get("completion_start_time", None),
prompt=generation_params.get("prompt"),
)
generation: Final = start_generation(
tracing=self.tracing,
trace_id=resolved_trace_id,
parent_observation_id=resolve_observation_id(parent_observation_id), # pyright: ignore[reportArgumentType] # metadata value, str or None at runtime
existing_trace=continued_trace,
observation_id=resolve_observation_id(generation_params["id"]),
name=generation_params["name"], # pyright: ignore[reportArgumentType] # always the str set a few lines up
start_time=start_time,
public=trace_public,
attributes=MappingProxyType({**generation_attributes, **trace_level_attributes}),
)
try:
log_provider_specific_information_as_span(
tracing=self.tracing, parent=generation, enrichments=enrichments
)
self._log_guardrail_information_as_span(
tracing=self.tracing, parent=generation, standard_logging_object=standard_logging_object
)
finally:
generation.end(end_time)
# Return the trace_id we set (which should be litellm_call_id when no explicit trace_id provided)
# We explicitly set trace_id in trace_params["id"], so langfuse should use it
# Verify langfuse accepted our trace_id; if it differs, log a warning but still return our intended value
# to match expected test behavior
if hasattr(generation_client, "trace_id") and generation_client.trace_id:
if generation_client.trace_id != resolved_trace_id:
verbose_logger.warning(
"Langfuse trace_id mismatch: set %s, but langfuse returned %s. Using our intended trace_id for consistency.",
resolved_trace_id,
generation_client.trace_id,
)
return resolved_trace_id, generation_id
# log_event_on_langfuse tuple-unpacks this and re-wraps it in the dict callers cache.
# The observation id is the requested generation_id after resolve_observation_id.
return resolved_trace_id, generation.id
except Exception:
verbose_logger.error("Langfuse Layer Error - %s", traceback.format_exc())
return None, None
@ -904,27 +1025,11 @@ class LangFuseLogger:
_cache_key = _hidden_params.get("cache_key", None)
if _cache_key is None and litellm.cache is not None:
# fallback to using "preset_cache_key"
_preset_cache_key: Final = litellm.cache._get_preset_cache_key_from_kwargs(**kwargs)
_preset_cache_key: Final = litellm.cache._get_preset_cache_key_from_kwargs(**kwargs) # pyright: ignore[reportPrivateUsage] # kwargs-ok: no public preset-cache-key accessor
_cache_key = _preset_cache_key
tags.append(f"cache_key:{_cache_key}")
return tags
def _supports_tags(self):
"""Check if current langfuse version supports tags"""
return Version(self.langfuse_sdk_version) >= Version("2.6.3")
def _supports_prompt(self):
"""Check if current langfuse version supports prompt"""
return Version(self.langfuse_sdk_version) >= Version("2.7.3")
def _supports_costs(self):
"""Check if current langfuse version supports costs"""
return Version(self.langfuse_sdk_version) >= Version("2.7.3")
def _supports_completion_start_time(self):
"""Check if current langfuse version supports completion start time"""
return Version(self.langfuse_sdk_version) >= Version("2.7.3")
@staticmethod
def _apply_masking_function(data: object, masking_function: Callable[[object], object]) -> object:
"""
@ -973,23 +1078,24 @@ class LangFuseLogger:
@staticmethod
def _get_langfuse_flush_interval(flush_interval: int) -> int:
"""
Get the langfuse flush interval to initialize the Langfuse client
Reads `LANGFUSE_FLUSH_INTERVAL` from the environment variable.
If not set, uses the flush interval passed in as an argument.
Args:
flush_interval: The flush interval to use if LANGFUSE_FLUSH_INTERVAL is not set
Returns:
[int] The flush interval to use to initialize the Langfuse client
"""
return int(os.getenv("LANGFUSE_FLUSH_INTERVAL") or flush_interval)
"""``LANGFUSE_FLUSH_INTERVAL`` in whole seconds above 0 (the export scheduler's delay), else ``flush_interval``."""
raw: Final = os.getenv("LANGFUSE_FLUSH_INTERVAL")
if not raw:
return flush_interval
parsed: Final = int(raw) if raw.strip().isdigit() else None
if parsed is None or parsed <= 0:
verbose_logger.warning(
"LANGFUSE_FLUSH_INTERVAL=%r is not a whole number of seconds above 0; flushing every %d s",
raw,
flush_interval,
)
return flush_interval
return parsed
def _log_guardrail_information_as_span(
self,
trace: StatefulTraceClient,
tracing: "LangfuseTracing",
parent: "LangfuseObservation",
standard_logging_object: StandardLoggingPayload | None,
):
"""
@ -1011,6 +1117,8 @@ class LangFuseLogger:
)
return
from litellm.integrations.langfuse.langfuse_sdk import observation_attributes, start_child_span
for guardrail_entry in guardrail_information:
if not isinstance(guardrail_entry, dict):
verbose_logger.debug(
@ -1019,30 +1127,35 @@ class LangFuseLogger:
)
continue
span = trace.span(
span = start_child_span(
tracing=tracing,
parent=parent,
name="guardrail",
input=guardrail_entry.get("guardrail_request", None),
output=guardrail_entry.get("guardrail_response", None),
metadata={
"guardrail_name": guardrail_entry.get("guardrail_name", None),
"guardrail_mode": guardrail_entry.get("guardrail_mode", None),
"guardrail_masked_entity_count": guardrail_entry.get("masked_entity_count", None),
},
start_time=guardrail_entry.get("start_time", None),
end_time=guardrail_entry.get("end_time", None),
attributes=observation_attributes(
observation_type="span",
input=guardrail_entry.get("guardrail_request", None),
output=guardrail_entry.get("guardrail_response", None),
metadata=MappingProxyType(
{
"guardrail_name": guardrail_entry.get("guardrail_name", None),
"guardrail_mode": guardrail_entry.get("guardrail_mode", None),
"guardrail_masked_entity_count": guardrail_entry.get("masked_entity_count", None),
}
),
),
)
verbose_logger.debug("Logged guardrail information as span: %s", span)
span.end()
span.end(guardrail_entry.get("end_time", None))
def _add_prompt_to_generation_params(
generation_params: dict,
clean_metadata: dict,
prompt_management_metadata: StandardLoggingPromptManagementMetadata | None,
langfuse_client: object,
langfuse_client: "LangfuseApiClient",
) -> dict:
from langfuse import Langfuse
from langfuse.model import (
ChatPromptClient,
Prompt_Chat,
@ -1050,8 +1163,6 @@ def _add_prompt_to_generation_params(
TextPromptClient,
)
langfuse_client = cast(Langfuse, langfuse_client)
user_prompt: Final = clean_metadata.pop("prompt", None)
if user_prompt is None and prompt_management_metadata is None:
pass
@ -1075,7 +1186,7 @@ def _add_prompt_to_generation_params(
if "labels" in prompt_text_params and "tags" in prompt_text_params:
_data["labels"] = user_prompt.get("labels", []) or []
_data["tags"] = user_prompt.get("tags", []) or []
_prompt_obj = Prompt_Text(**_data)
_prompt_obj = Prompt_Text(**_data) # pyright: ignore[reportArgumentType] # kwargs-ok: shape mirrors the pydantic model, values from the user's prompt dict
generation_params["prompt"] = TextPromptClient(prompt=_prompt_obj)
elif isinstance(user_prompt["prompt"], list):
@ -1090,7 +1201,7 @@ def _add_prompt_to_generation_params(
_data["labels"] = user_prompt.get("labels", []) or []
_data["tags"] = user_prompt.get("tags", []) or []
_prompt_obj = Prompt_Chat(**_data)
_prompt_obj = Prompt_Chat(**_data) # pyright: ignore[reportArgumentType] # kwargs-ok: shape mirrors the pydantic model, values from the user's prompt dict
generation_params["prompt"] = ChatPromptClient(prompt=_prompt_obj)
else:
@ -1110,21 +1221,14 @@ def _add_prompt_to_generation_params(
def log_provider_specific_information_as_span(
trace,
clean_metadata: Mapping[str, Any],
*,
tracing: "LangfuseTracing",
parent: "LangfuseObservation",
enrichments: Mapping[str, Any],
):
"""
Logs provider-specific information as spans.
"""Logs provider-specific information as spans under the generation."""
Parameters:
trace: The tracing object used to log spans.
clean_metadata: A dictionary containing metadata to be logged.
Returns:
None
"""
_hidden_params: Final[Mapping[str, object] | None] = clean_metadata.get("hidden_params", None)
_hidden_params: Final[Mapping[str, object] | None] = enrichments.get("hidden_params", None)
if _hidden_params is None:
return
@ -1135,22 +1239,27 @@ def log_provider_specific_information_as_span(
for elem in vertex_ai_grounding_metadata:
if isinstance(elem, dict):
for key, value in elem.items():
trace.span(
name=key,
input=value,
)
_end_grounding_span(tracing=tracing, parent=parent, name=key, value=value)
else:
trace.span(
name="vertex_ai_grounding_metadata",
input=elem,
)
_end_grounding_span(tracing=tracing, parent=parent, name="vertex_ai_grounding_metadata", value=elem)
else:
trace.span(
name="vertex_ai_grounding_metadata",
input=vertex_ai_grounding_metadata,
_end_grounding_span(
tracing=tracing, parent=parent, name="vertex_ai_grounding_metadata", value=vertex_ai_grounding_metadata
)
def _end_grounding_span(*, tracing: "LangfuseTracing", parent: "LangfuseObservation", name: str, value: object) -> None:
from litellm.integrations.langfuse.langfuse_sdk import observation_attributes, start_child_span
start_child_span(
tracing=tracing,
parent=parent,
name=name,
start_time=None,
attributes=observation_attributes(observation_type="span", input=value),
).end()
def log_requester_metadata(clean_metadata: Mapping[str, Any]):
returned_metadata: Final = {}
requester_metadata: Final = clean_metadata.get("requester_metadata") or {}

View file

@ -2,16 +2,14 @@
Call Hook for LiteLLM Proxy which allows Langfuse prompt management.
"""
import inspect
import os
from functools import lru_cache
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, cast
from packaging.version import Version
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prompt_management_base import PromptManagementClient
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.types.integrations.langfuse import LangfuseLoggedEvent
from litellm.types.llms.openai import AllMessageValues, ChatCompletionSystemMessage
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.types.utils import StandardCallbackDynamicParams, StandardLoggingPayload
@ -19,17 +17,27 @@ from litellm.types.utils import StandardCallbackDynamicParams, StandardLoggingPa
from ...litellm_core_utils.specialty_caches.dynamic_logging_cache import (
DynamicLoggingCache,
)
from ...litellm_core_utils.specialty_caches.service_trace_id_cache import in_memory_trace_id_cache
from ..prompt_management_base import PromptManagementBase
from .langfuse import LangFuseLogger, resolve_langfuse_credentials
from .langfuse import (
LangFuseLogger,
installed_langfuse_version,
raise_if_unsupported_langfuse_version,
raise_if_unusable_prompt_cache_ttl,
resolve_langfuse_credentials,
warn_if_upstream_langfuse_configured,
)
from .langfuse_handler import LangFuseHandler
from .langfuse_mock_client import create_mock_langfuse_client, should_use_langfuse_mock
if TYPE_CHECKING:
from langfuse import Langfuse
from langfuse.client import ChatPromptClient, TextPromptClient
from langfuse.model import ChatPromptClient, TextPromptClient
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
LangfuseClass: TypeAlias = Langfuse
from .langfuse_sdk import LangfuseApiClient
LangfuseClass: TypeAlias = LangfuseApiClient
PROMPT_CLIENT = TextPromptClient | ChatPromptClient
else:
@ -49,23 +57,24 @@ def langfuse_client_init(
allow_env_credentials: bool = True,
) -> LangfuseClass:
"""
Initialize Langfuse client with caching to prevent multiple initializations.
Initialize the Langfuse REST client with caching to prevent multiple initializations.
Args:
langfuse_public_key (str, optional): Public key for Langfuse. Defaults to None.
langfuse_secret (str, optional): Secret key for Langfuse. Defaults to None.
langfuse_host (str, optional): Host URL for Langfuse. Defaults to None.
flush_interval (int, optional): Flush interval in seconds. Defaults to 1.
flush_interval (int, optional): Kept in the signature so cached callers keep their cache key.
Returns:
Langfuse: Initialized Langfuse client instance
LangfuseApiClient: prompt, auth and project lookups for one credential set
Raises:
Exception: If langfuse package is not installed
"""
raise_if_unsupported_langfuse_version(installed_langfuse_version())
raise_if_unusable_prompt_cache_ttl()
try:
import langfuse
from langfuse import Langfuse
from .langfuse_sdk import build_langfuse_client
except Exception as e:
raise Exception(
f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n\033[0m"
@ -83,39 +92,22 @@ def langfuse_client_init(
# add http:// if unset, assume communicating over private network - e.g. render
langfuse_host = "http://" + langfuse_host
langfuse_release: Final = os.getenv("LANGFUSE_RELEASE")
langfuse_debug: Final = os.getenv("LANGFUSE_DEBUG")
warn_if_upstream_langfuse_configured()
parameters: Final = {
"public_key": public_key,
"secret_key": secret_key,
"host": langfuse_host,
"release": langfuse_release,
"debug": langfuse_debug,
"flush_interval": LangFuseLogger._get_langfuse_flush_interval(flush_interval), # flush interval in seconds
}
httpx_client: Final = create_mock_langfuse_client() if should_use_langfuse_mock() else HTTPHandler().client
return build_langfuse_client(
public_key=public_key,
secret_key=secret_key,
base_url=langfuse_host,
httpx_client=httpx_client,
)
if Version(langfuse.version.__version__) >= Version("2.6.0"):
parameters["sdk_integration"] = "litellm"
if Version(langfuse.version.__version__) >= Version("2.7.3"):
import httpx
import litellm
from ...llms.custom_httpx.http_handler import get_ssl_configuration
parameters["httpx_client"] = httpx.Client(
verify=get_ssl_configuration(),
cert=os.getenv("SSL_CERTIFICATE", litellm.ssl_certificate),
)
if "environment" in inspect.signature(Langfuse.__init__).parameters:
parameters["environment"] = LangFuseLogger.resolve_deployment_environment()
client: Final = Langfuse(**parameters)
return client
def _remember_trace_id(litellm_call_id: object, logged: LangfuseLoggedEvent) -> None:
trace_id: Final = logged["trace_id"]
if not isinstance(litellm_call_id, str) or trace_id is None:
return
in_memory_trace_id_cache.set_cache(litellm_call_id=litellm_call_id, service_name="langfuse", trace_id=trace_id)
class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogger):
@ -126,15 +118,33 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge
langfuse_host=None,
flush_interval=1,
):
import langfuse
self.langfuse_sdk_version = langfuse.version.__version__
self.Langfuse = langfuse_client_init(
self.langfuse_sdk_version = installed_langfuse_version()
raise_if_unsupported_langfuse_version(self.langfuse_sdk_version)
raise_if_unusable_prompt_cache_ttl()
from .langfuse_sdk import acquire_langfuse_tracing, configured_release
self.api_client = langfuse_client_init(
langfuse_public_key=langfuse_public_key,
langfuse_secret=langfuse_secret,
langfuse_host=langfuse_host,
flush_interval=flush_interval,
)
self.public_key, self.secret_key, self.langfuse_host = resolve_langfuse_credentials(
langfuse_public_key=langfuse_public_key,
langfuse_secret=langfuse_secret,
langfuse_host=langfuse_host,
)
self.tracing = acquire_langfuse_tracing(
public_key=str(self.public_key),
secret_key=str(self.secret_key),
base_url=self.langfuse_host,
environment=LangFuseLogger.resolve_deployment_environment(),
release=configured_release(),
flush_interval=LangFuseLogger._get_langfuse_flush_interval(flush_interval), # pyright: ignore[reportPrivateUsage] # shared env-fallback helper, not part of the logger's API
mock_mode=should_use_langfuse_mock(),
)
@property
def integration_name(self):
@ -228,11 +238,8 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge
langfuse_host=dynamic_callback_params.get("langfuse_host"),
allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None,
)
langfuse_prompt_client: Final = self._get_prompt_from_id(
langfuse_prompt_id=prompt_id,
langfuse_client=langfuse_client,
)
return langfuse_prompt_client is not None
self._get_prompt_from_id(langfuse_prompt_id=prompt_id, langfuse_client=langfuse_client)
return True
def _compile_prompt_helper(
self,
@ -311,13 +318,14 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge
standard_callback_dynamic_params=standard_callback_dynamic_params,
in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache,
)
langfuse_logger_to_use.log_event_on_langfuse(
logged: Final = langfuse_logger_to_use.log_event_on_langfuse(
kwargs=kwargs,
response_obj=response_obj,
start_time=start_time,
end_time=end_time,
user_id=kwargs.get("user", None),
)
_remember_trace_id(litellm_call_id=kwargs.get("litellm_call_id"), logged=logged)
except Exception as e:
from litellm._logging import verbose_logger
@ -339,7 +347,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge
status_message = str(kwargs.get("exception", "Unknown error"))
if standard_logging_object is not None:
status_message = standard_logging_object.get("error_str", None) or status_message
langfuse_logger_to_use.log_event_on_langfuse(
logged: Final = langfuse_logger_to_use.log_event_on_langfuse(
start_time=start_time,
end_time=end_time,
response_obj=None,
@ -348,6 +356,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge
level="ERROR",
kwargs=kwargs,
)
_remember_trace_id(litellm_call_id=kwargs.get("litellm_call_id"), logged=logged)
except Exception as e:
from litellm._logging import verbose_logger

File diff suppressed because it is too large Load diff

View file

@ -729,6 +729,15 @@ class PrometheusLogger(CustomLogger):
labelnames=self.get_labels_for_metric("litellm_zero_cost_requests_total"),
)
self.litellm_spend_capture_rate = self._gauge_factory(
"litellm_spend_capture_rate",
(
"Share of the provider's bill LiteLLM captured as spend over the scheduled check's window "
"(captured spend / provider bill), by api_provider; NaN when the last check produced no rate"
),
labelnames=self.get_labels_for_metric("litellm_spend_capture_rate"),
)
# Cache metrics
self.litellm_cache_hits_metric = self._counter_factory(
name="litellm_cache_hits_metric",
@ -2028,6 +2037,15 @@ class PrometheusLogger(CustomLogger):
)
self.litellm_zero_cost_requests_total.labels(**labels).inc()
def set_spend_capture_rate(self, api_provider: str, capture_rate: float | None) -> None:
labels: Final = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric("litellm_spend_capture_rate"),
enum_values=UserAPIKeyLabelValues(api_provider=api_provider),
)
gauge: Final = self.litellm_spend_capture_rate
series: Final = gauge.labels(**labels) if labels else gauge
series.set(math.nan if capture_rate is None else capture_rate)
@staticmethod
def _get_remaining_from_v3_rate_limit_headers(
standard_logging_payload: StandardLoggingPayload | None,
@ -2605,6 +2623,7 @@ class PrometheusLogger(CustomLogger):
from litellm.litellm_core_utils.litellm_logging import (
StandardLoggingPayloadSetup,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
status_code: Final = self._extract_status_code(exception=original_exception)
@ -2623,7 +2642,9 @@ class PrometheusLogger(CustomLogger):
end_user=user_api_key_dict.end_user_id,
user=user_api_key_dict.user_id,
user_email=user_api_key_dict.user_email,
hashed_api_key=None if status_code == 401 else user_api_key_dict.api_key,
hashed_api_key=None
if status_code == 401
else LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
api_key_alias=user_api_key_dict.key_alias,
team=user_api_key_dict.team_id,
team_alias=user_api_key_dict.team_alias,

View file

@ -1687,7 +1687,7 @@ class WebSearchInterceptionLogger(CustomLogger):
**user_api_key_metadata,
**parent_correlation.as_search_metadata(),
"model_group": search_tool_name,
"user_api_key": user_api_key_auth.api_key,
"user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_auth),
"user_api_key_auth": user_api_key_auth,
}

View file

@ -2,7 +2,11 @@ from typing import Final, cast
from urllib.parse import urlparse
import litellm
from litellm.constants import PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO, REPLICATE_MODEL_NAME_WITH_ID_LENGTH
from litellm.constants import (
NADIR_DEFAULT_API_BASE,
PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO,
REPLICATE_MODEL_NAME_WITH_ID_LENGTH,
)
from litellm.litellm_core_utils.fallback_generalizations import (
match_routing_generalization,
)
@ -139,6 +143,18 @@ def declared_authenticating_provider(model: str | None, custom_llm_provider: str
return declared if declared in PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO else None
def inferred_provider(model: str | None) -> str | None:
if not model:
return None
declared: Final = declared_authenticating_provider(model)
if declared is not None:
return declared
try:
return get_llm_provider(model=model)[1]
except Exception: # noqa: BLE001 # get_llm_provider raises for an unknown name, which then has no provider
return None
def get_llm_provider(
model: str,
custom_llm_provider: str | None = None,
@ -265,6 +281,11 @@ def get_llm_provider(
elif endpoint == "https://api.cerebras.ai/v1":
custom_llm_provider = "cerebras"
dynamic_api_key = get_secret_str("CEREBRAS_API_KEY")
elif endpoint == NADIR_DEFAULT_API_BASE:
custom_llm_provider = "nadir" # rebind-ok: mirrors sibling endpoint branches
dynamic_api_key = (
get_secret_str("NADIR_API_KEY") if api_base.lower().startswith("https://") else None
)
elif endpoint == "https://inference.baseten.co/v1":
custom_llm_provider = "baseten"
dynamic_api_key = get_secret_str("BASETEN_API_KEY")
@ -637,6 +658,13 @@ def _get_openai_compatible_provider_info(
elif custom_llm_provider == "cerebras":
api_base = api_base or get_secret("CEREBRAS_API_BASE") or "https://api.cerebras.ai/v1"
dynamic_api_key = api_key or get_secret_str("CEREBRAS_API_KEY")
elif custom_llm_provider == "nadir":
default_nadir_base: Final = get_secret_str("NADIR_API_BASE") or NADIR_DEFAULT_API_BASE
caller_base: Final = api_base
api_base = api_base or default_nadir_base # rebind-ok: mirrors sibling provider branches
trusted_base: Final = caller_base is None or caller_base.rstrip("/") == default_nadir_base.rstrip("/")
env_key: Final = get_secret_str("NADIR_API_KEY") if trusted_base else None
dynamic_api_key = api_key or env_key # rebind-ok: mirrors sibling provider branches
elif custom_llm_provider == "baseten":
# Use BasetenConfig to determine the appropriate API base URL
if api_base is None:

View file

@ -91,6 +91,8 @@ def get_supported_openai_params(
return litellm.nvidiaNimEmbeddingConfig.get_supported_openai_params()
elif custom_llm_provider == "cerebras":
return litellm.CerebrasConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "nadir":
return litellm.NadirConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "baseten":
return litellm.BasetenConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "xai":

View file

@ -6,8 +6,10 @@ import base64
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Final, Literal
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, DocumentType
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS, LlmProviders
from litellm.llms.base_llm.ocr.transformation import DocumentType
from litellm.rust_bridge import runtime
from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR_HEALTH_CHECK_DOCUMENT
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
@ -29,11 +31,12 @@ def get_image_file_for_health_check() -> bytes:
def _ocr_health_check_document(model: str, custom_llm_provider: str) -> DocumentType:
from litellm.utils import ProviderConfigManager
provider: Final = next((known for known in LlmProviders if known.value == custom_llm_provider), None)
config: Final = ProviderConfigManager.get_provider_ocr_config(model=model, provider=provider) if provider else None
return (config or BaseOCRConfig()).get_health_check_document()
native: Final = NATIVE_OCR_HEALTH_CHECK_DOCUMENT.load()
if native is None:
raise runtime.NoPythonImplementationError(
"ocr health check documents are resolved by the Rust extension, which is not available"
)
return native(model, custom_llm_provider)
class HealthCheckHelpers:

View file

@ -36,7 +36,7 @@ from litellm._logging import (
)
from litellm._uuid import uuid
from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final
from litellm.caching.caching import DualCache, InMemoryCache
from litellm.caching.caching import DualCache
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.constants import (
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
@ -221,6 +221,7 @@ from .initialize_dynamic_callback_params import (
initialize_standard_callback_dynamic_params as _initialize_standard_callback_dynamic_params,
)
from .specialty_caches.dynamic_logging_cache import DynamicLoggingCache
from .specialty_caches.service_trace_id_cache import in_memory_trace_id_cache
if TYPE_CHECKING:
from mcp.types import CallToolResult, EmbeddedResource, ImageContent, TextContent
@ -349,21 +350,6 @@ last_fetched_at_keys: Final = None
####
class ServiceTraceIDCache:
def __init__(self) -> None:
self.cache = InMemoryCache()
def get_cache(self, litellm_call_id: str, service_name: str) -> str | None:
key_name: Final = f"{service_name}:{litellm_call_id}"
response: Final = self.cache.get_cache(key=key_name)
return response
def set_cache(self, litellm_call_id: str, service_name: str, trace_id: str) -> None:
key_name: Final = f"{service_name}:{litellm_call_id}"
self.cache.set_cache(key=key_name, value=trace_id)
in_memory_trace_id_cache: Final = ServiceTraceIDCache()
in_memory_dynamic_logger_cache: Final = DynamicLoggingCache()
# Cached lazy import for PrometheusLogger
@ -385,6 +371,10 @@ def _get_cached_prometheus_logger():
return _PrometheusLogger
class RawRequestCaptured(Exception):
pass
_DEPLOYMENT_PRICING_KEYS: Final = (
"input_cost_per_token",
"output_cost_per_token",
@ -591,6 +581,7 @@ class Logging(LiteLLMLoggingBaseClass):
kwargs: dict | None = None,
log_raw_request_response: bool = False,
supports_correlation_logging: bool = True,
raw_request_only: bool = False,
):
_input: Final[str | None] = messages # save original value of messages
if messages is not None:
@ -650,6 +641,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.streaming_chunks: list[Any] = [] # for generating complete stream response
self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response
self.log_raw_request_response = log_raw_request_response
self.raw_request_only = raw_request_only
# Initialize dynamic callbacks
self.dynamic_input_callbacks: list[str | Callable | CustomLogger] | None = dynamic_input_callbacks
@ -1476,6 +1468,9 @@ class Logging(LiteLLMLoggingBaseClass):
if capture_exception: # log this error to sentry for debugging
capture_exception(e)
if self.raw_request_only:
raise RawRequestCaptured()
def _print_llm_call_debugging_log(
self,
api_base: str,
@ -3970,40 +3965,6 @@ class Logging(LiteLLMLoggingBaseClass):
return trace_id
def _get_callback_object(self, service_name: Literal["langfuse"]) -> Any | None:
"""
Return dynamic callback object.
Meant to solve issue when doing key-based/team-based logging
"""
global langFuseLogger
if service_name == "langfuse":
if langFuseLogger is None or (
(
self.standard_callback_dynamic_params.get("langfuse_public_key") is not None
and self.standard_callback_dynamic_params.get("langfuse_public_key") != langFuseLogger.public_key
)
or (
self.standard_callback_dynamic_params.get("langfuse_public_key") is not None
and self.standard_callback_dynamic_params.get("langfuse_public_key") != langFuseLogger.public_key
)
or (
self.standard_callback_dynamic_params.get("langfuse_host") is not None
and self.standard_callback_dynamic_params.get("langfuse_host") != langFuseLogger.langfuse_host
)
):
return LangFuseLogger(
langfuse_public_key=self.standard_callback_dynamic_params.get("langfuse_public_key"),
langfuse_secret=self.standard_callback_dynamic_params.get("langfuse_secret")
or self.standard_callback_dynamic_params.get("langfuse_secret_key"),
langfuse_host=self.standard_callback_dynamic_params.get("langfuse_host"),
allow_env_credentials=self.standard_callback_dynamic_params.get("langfuse_host") is None,
)
return langFuseLogger
return None
def handle_sync_success_callbacks_for_async_calls(
self,
result: Any,

View file

@ -1989,6 +1989,26 @@ def is_encrypted_reasoning_block(block: object) -> bool:
return _carries_encrypted_reasoning(_encrypted_reasoning_field(mapping))
def is_unsignable_thinking_block(block: object) -> bool:
"""A thinking block Anthropic cannot accept on input.
Anthropic verifies the thinking signature cryptographically, so a block whose
signature is null, empty, or missing (e.g. from an open-source reasoning model)
is rejected with a 400 and must be dropped rather than blanked or repaired, and
so is a block whose signature or data carries another provider's encrypted
reasoning. A `redacted_thinking` block Anthropic minted is always kept.
"""
if is_encrypted_reasoning_block(block):
return True
if not isinstance(block, Mapping):
return False
mapping: Final = cast(Mapping[str, object], block) # cast-ok: narrowed by isinstance
if mapping.get("type") != "thinking":
return False
signature: Final = mapping.get("signature")
return not (isinstance(signature, str) and len(signature) > 0)
def strip_encrypted_reasoning_from_messages(messages: object) -> None:
"""Drop the bridge-tagged reasoning blocks a routed deployment cannot decrypt from
Anthropic-shaped history.

View file

@ -7,7 +7,7 @@ import re
import xml.etree.ElementTree as ET
from collections.abc import Iterator, Mapping, Sequence
from enum import Enum
from typing import Any, Final, TypedDict, cast, overload
from typing import Any, Final, TypeAlias, TypedDict, cast, overload
from jinja2.sandbox import ImmutableSandboxedEnvironment
@ -17,6 +17,7 @@ import litellm.types.llms
from litellm import verbose_logger
from litellm._uuid import uuid
from litellm.constants import REDACTED_BY_LITELLM
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import anthropic_system_messages
from litellm.litellm_core_utils.url_utils import async_safe_get, safe_get
from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client
from litellm.types.files import get_file_extension_from_mime_type
@ -48,8 +49,8 @@ from litellm.types.utils import GenericImageParsingChunk
from .common_utils import (
convert_content_list_to_str,
infer_content_type_from_url_and_content,
is_encrypted_reasoning_block,
is_non_content_values_set,
is_unsignable_thinking_block,
parse_tool_call_arguments,
)
from .image_handling import convert_url_to_base64
@ -2329,37 +2330,25 @@ def sanitize_messages_for_tool_calling(
return sanitized_messages
def _is_unsignable_thinking_block(block: object) -> bool:
"""A thinking block that Anthropic cannot accept on input.
Anthropic verifies the thinking signature cryptographically, so a block whose
signature is null, empty, or missing (e.g. from an open-source reasoning model)
is rejected with a 400 and must be dropped rather than blanked or repaired, and
so is a block whose signature or data carries another provider's encrypted
reasoning. A `redacted_thinking` block Anthropic minted is always kept.
"""
if is_encrypted_reasoning_block(block):
return True
if not isinstance(block, dict) or block.get("type") != "thinking":
return False
signature: Final = block.get("signature")
return not (isinstance(signature, str) and len(signature) > 0)
def _drop_unsignable_thinking_blocks(
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock],
) -> list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock]:
return [block for block in thinking_blocks if not _is_unsignable_thinking_block(block)]
return [block for block in thinking_blocks if not is_unsignable_thinking_block(block)]
_AnthropicMessageList: TypeAlias = list[AllAnthropicPassThroughMessageValues]
def anthropic_messages_pt(
messages: list[AllMessageValues],
model: str,
llm_provider: str,
) -> list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]:
) -> _AnthropicMessageList:
"""
format messages for anthropic
1. Anthropic supports roles like "user" and "assistant" (system prompt sent separately)
1. Anthropic supports roles like "user" and "assistant" (system prompt sent separately).
Models flagged ``supports_mid_conversation_system`` also accept "system" inside
messages after a user turn; the caller decides placement, this keeps such messages.
2. The first message always needs to be of role "user"
3. Each message must alternate between "user" and "assistant" (this is not addressed as now by litellm)
4. final assistant content cannot end with trailing whitespace (anthropic raises an error otherwise)
@ -2384,7 +2373,7 @@ def anthropic_messages_pt(
# add role=tool support to allow function call result/error submission
user_message_types: Final = {"user", "tool", "function"}
# reformat messages to ensure user/assistant are alternating, if there's either 2 consecutive 'user' messages or 2 consecutive 'assistant' message, merge them.
new_messages: Final[list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]] = []
new_messages: Final[_AnthropicMessageList] = [] # mutable-ok: accumulator behind the mutable return contract
if len(messages) == 0:
if not litellm.modify_params:
@ -2697,7 +2686,7 @@ def anthropic_messages_pt(
if (
m.get("type", "") == "thinking"
and len(thinking_block) > 0
and not _is_unsignable_thinking_block(m)
and not is_unsignable_thinking_block(m)
): # don't pass empty text blocks. anthropic api raises errors.
anthropic_message: ChatCompletionThinkingBlock | AnthropicMessagesTextParam = cast(
ChatCompletionThinkingBlock, m
@ -2777,6 +2766,11 @@ def anthropic_messages_pt(
if assistant_content:
new_messages.append({"role": "assistant", "content": assistant_content})
## MID-CONVERSATION SYSTEM MESSAGES (placement is the caller's job) ##
while msg_i < len(messages) and messages[msg_i]["role"] == "system":
new_messages.extend(anthropic_system_messages(messages[msg_i]))
msg_i += 1
if msg_i == init_msg_i: # prevent infinite loops
raise litellm.BadRequestError(
message=BAD_MESSAGE_ERROR_STR + f"passed in {messages[msg_i]}",

View file

@ -0,0 +1,418 @@
"""Placement policy for ``role: "system"`` messages that appear after the first turn
of an Anthropic-shaped chat completions request.
Only the leading run of system messages belongs in the top-level ``system``
parameter. Hoisting a later one there rewrites the cached prefix, so the provider
re-bills the whole conversation at cache-write pricing on every reminder (#36559).
Models flagged ``supports_mid_conversation_system`` in the cost map accept the role
inside ``messages`` under Anthropic's placement rules: the message must directly
follow a user turn, must be the last entry or be followed by an assistant turn, and
must not sit next to another system message. OpenAI-shaped clients put system
messages anywhere, so this module places each run by its neighbours alone: a run
after a user turn stays with that turn, a run after an assistant turn slides
behind the user turn that immediately follows it, and a run that ends the array
or precedes an assistant turn becomes a user turn in place. Runs that land on the
same slot merge into one system message. No later message can move an earlier
run, so a client that replays the conversation with more turns appended sends a
byte-identical prefix and preserved thinking blocks keep their binding.
Models without the flag reject the role inside ``messages``. Their system messages
become user turns in place, prefixed with an operator note so the model can tell
the instruction apart from the user's own words. A run caught between a tool call
and its result moves to just after the result so the ``tool_result`` block stays
first in the merged user turn.
Every transformation here is a pure function of the message sequence: turn N's
output stays a prefix of turn N+1's output, which is what keeps the provider-side
prompt cache readable across turns. Messages are handled in OpenAI format; the
Anthropic wire shape is built later by ``anthropic_messages_pt``.
"""
from collections.abc import Iterator, Mapping, Sequence
from itertools import chain, groupby
from typing import Final, Literal, TypeAlias
from litellm.types.llms.anthropic import AnthropicMessagesSystemMessageParam, AnthropicSystemMessageContent
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionCachedContent,
ChatCompletionSystemMessage,
ChatCompletionTextObject,
ChatCompletionUserMessage,
)
from .common_utils import is_unsignable_thinking_block
CONVERTED_SYSTEM_NOTE: Final = (
"Operator note (not from the user): the following was originally a mid-conversation system-role reminder."
)
_USER_TYPE_ROLES: Final = frozenset({"user", "tool", "function"})
_TOOL_ROLES: Final = frozenset({"tool", "function"})
_RENDERED_PART_TYPES: Final = frozenset({"text", "image_url", "document", "file"})
_RENDERED_ASSISTANT_PART_TYPES: Final = frozenset({"text", "server_tool_use"})
_THINKING_BLOCK_TYPES: Final = frozenset({"thinking", "redacted_thinking"})
_MessageKind: TypeAlias = Literal["system", "tool", "user", "other"]
_TextPart: TypeAlias = tuple[str, ChatCompletionCachedContent | None]
def _as_mapping(value: object) -> Mapping[str, object] | None:
return value if isinstance(value, Mapping) else None
def parts_of(value: object) -> tuple[object, ...]:
return tuple(value) if isinstance(value, Sequence) and not isinstance(value, str) else ()
def message_field(message: object, key: str) -> object:
"""A message field, whether the message is a dict or a pydantic ``Message``.
Clients replay assistant turns straight from a response, so a history mixes
plain dicts with ``litellm.Message`` objects; every predicate reads through here.
"""
mapping: Final = _as_mapping(message)
return mapping.get(key) if mapping is not None else getattr(message, key, None)
def is_system_message(message: object) -> bool:
return message_field(message, "role") == "system"
def _is_user_type(message: object) -> bool:
return message_field(message, "role") in _USER_TYPE_ROLES
def _kind(message: object) -> _MessageKind:
role: Final = message_field(message, "role")
if role == "system":
return "system"
if role in _TOOL_ROLES:
return "tool"
if role == "user":
return "user"
return "other"
def split_leading_system_run(
messages: Sequence[AllMessageValues],
) -> tuple[tuple[AllMessageValues, ...], tuple[AllMessageValues, ...]]:
"""Split ``messages`` into the leading run of system messages and everything after it."""
leading_count: Final = next(
(index for index, message in enumerate(messages) if not is_system_message(message)),
len(messages),
)
return tuple(messages[:leading_count]), tuple(messages[leading_count:])
def _cache_control(holder: object) -> ChatCompletionCachedContent | None:
"""The client's ``cache_control`` rebuilt in the only shape Anthropic accepts."""
value: Final = _as_mapping(message_field(holder, "cache_control"))
if value is None or value.get("type") != "ephemeral":
return None
ttl: Final = value.get("ttl")
if ttl == "1h":
one_hour: Final[ChatCompletionCachedContent] = {"type": "ephemeral", "ttl": "1h"}
return one_hour
if ttl == "5m":
five_minutes: Final[ChatCompletionCachedContent] = {"type": "ephemeral", "ttl": "5m"}
return five_minutes
ephemeral: Final[ChatCompletionCachedContent] = {"type": "ephemeral"}
return ephemeral
def _text_parts(message: object) -> tuple[_TextPart, ...]:
"""``(text, cache_control)`` for each non-empty text part of a system message.
Anthropic rejects empty text blocks and only accepts text in system content. A
``cache_control`` on the message itself belongs to the block built from string
content; block-level ``cache_control`` stays with its block.
"""
content: Final = message_field(message, "content")
if isinstance(content, str):
return ((content, _cache_control(message)),) if content else ()
return tuple(part for part in map(_text_part, parts_of(content)) if part is not None)
def _text_part(part: object) -> _TextPart | None:
if message_field(part, "type") != "text":
return None
text: Final = message_field(part, "text")
return (text, _cache_control(part)) if isinstance(text, str) and text else None
def _openai_text_block(part: _TextPart) -> ChatCompletionTextObject:
text, cache_control = part
if cache_control is None:
plain: Final[ChatCompletionTextObject] = {"type": "text", "text": text}
return plain
cached: Final[ChatCompletionTextObject] = {"type": "text", "text": text, "cache_control": cache_control}
return cached
def _anthropic_text_block(part: _TextPart) -> AnthropicSystemMessageContent:
text, cache_control = part
if cache_control is None:
plain: Final[AnthropicSystemMessageContent] = {"type": "text", "text": text}
return plain
cached: Final[AnthropicSystemMessageContent] = {"type": "text", "text": text, "cache_control": cache_control}
return cached
def anthropic_system_messages(message: object) -> tuple[AnthropicMessagesSystemMessageParam, ...]:
"""The Anthropic wire message for a system message, or nothing when it carries no text."""
blocks: Final = tuple(_anthropic_text_block(part) for part in _text_parts(message))
if not blocks:
return ()
wire: Final[AnthropicMessagesSystemMessageParam] = {
"role": "system",
"content": list(blocks), # mutable-ok: wire payload; cache_control hooks edit content blocks in place
}
return (wire,)
def system_message_as_user(message: object) -> ChatCompletionUserMessage:
"""A system message re-rolled as a user turn, prefixed with the operator note."""
note: Final[ChatCompletionTextObject] = {"type": "text", "text": CONVERTED_SYSTEM_NOTE}
content: Final[list[ChatCompletionTextObject]] = [ # mutable-ok: anthropic_messages_pt only recognises list content
note,
*(_openai_text_block(part) for part in _text_parts(message)),
]
turn: Final[ChatCompletionUserMessage] = {"role": "user", "content": content}
return turn
def _merged_system_message(run: Sequence[object]) -> tuple[ChatCompletionSystemMessage, ...]:
parts: Final = tuple(chain.from_iterable(_text_parts(message) for message in run))
if not parts:
return ()
content: Final[list[ChatCompletionTextObject]] = [ # mutable-ok: anthropic_messages_pt only recognises list content
_openai_text_block(part) for part in parts
]
merged: Final[ChatCompletionSystemMessage] = {"role": "system", "content": content}
return (merged,)
def _converted_user_turns(run: Sequence[object]) -> tuple[ChatCompletionUserMessage, ...]:
return tuple(system_message_as_user(message) for message in run if _text_parts(message))
def _runs(messages: Sequence[AllMessageValues]) -> tuple[tuple[_MessageKind, tuple[AllMessageValues, ...]], ...]:
return tuple((kind, tuple(group)) for kind, group in groupby(messages, key=_kind))
def _converted_for_unflagged_model(messages: Sequence[AllMessageValues]) -> tuple[AllMessageValues, ...]:
"""Convert every system message to a user turn in place.
A system run whose follower is a tool message is emitted after that tool run:
``tool_result`` blocks have to open the merged user turn.
"""
runs: Final = _runs(messages)
def emit(index: int) -> tuple[AllMessageValues, ...]:
kind, run = runs[index]
follower: Final = runs[index + 1][0] if index + 1 < len(runs) else None
if kind == "system":
return () if follower == "tool" else _converted_user_turns(run)
if kind == "tool" and index > 0 and runs[index - 1][0] == "system":
return (*run, *_converted_user_turns(runs[index - 1][1]))
return run
return tuple(chain.from_iterable(emit(index) for index in range(len(runs))))
def _user_type_blocks(messages: Sequence[AllMessageValues]) -> tuple[tuple[bool, tuple[int, ...]], ...]:
"""Maximal groups of consecutive non-system messages, keyed by whether they are user-type.
Consecutive user-type messages become one user turn on the wire, so a group is
the unit a system message can validly follow.
"""
indexed: Final = tuple((index, message) for index, message in enumerate(messages) if not is_system_message(message))
return tuple(
(is_user, tuple(index for index, _ in group))
for is_user, group in groupby(indexed, key=lambda pair: _is_user_type(pair[1]))
)
def _system_runs(messages: Sequence[AllMessageValues]) -> tuple[tuple[int, ...], ...]:
"""Index runs of consecutive system messages."""
system_indices: Final = tuple(index for index, message in enumerate(messages) if is_system_message(message))
return tuple(
tuple(index for _, index in group)
for _, group in groupby(enumerate(system_indices), key=lambda pair: pair[1] - pair[0])
)
def _block_containing(message_index: int, blocks: Sequence[tuple[bool, tuple[int, ...]]]) -> int:
return next(index for index, (_, indices) in enumerate(blocks) if message_index in indices)
def _thinking_block_renders(block: object) -> bool:
"""A thinking block the converter keeps: one Anthropic can verify, so never bridged encrypted reasoning."""
return message_field(block, "type") in _THINKING_BLOCK_TYPES and not is_unsignable_thinking_block(block)
def _assistant_part_renders(part: object) -> bool:
"""A text part always renders: the converter pads empty text with a placeholder."""
part_type: Final = message_field(part, "type")
if part_type == "thinking":
thinking: Final = message_field(part, "thinking")
return isinstance(thinking, str) and bool(thinking) and _thinking_block_renders(part)
return part_type in _RENDERED_ASSISTANT_PART_TYPES or (
isinstance(part_type, str) and part_type.endswith("_tool_result")
)
def _separate_thinking_blocks_render(message: object, parts: Sequence[object]) -> bool:
"""``thinking_blocks`` reach the wire only when no inline thinking part claims the slot.
The converter skips the separate blocks as soon as the content list carries a
``thinking`` or ``redacted_thinking`` part, whether or not that part itself renders.
"""
if any(message_field(part, "type") in _THINKING_BLOCK_TYPES for part in parts):
return False
return any(_thinking_block_renders(block) for block in parts_of(message_field(message, "thinking_blocks")))
def _assistant_renders(message: object) -> bool:
"""Whether ``anthropic_messages_pt`` puts a block on the wire for this assistant message.
String content (the converter pads an empty one with a placeholder), a text part,
a signed thinking part, a server tool part, tool calls, a function call, a kept
thinking block and compaction blocks each render. An assistant message with none
of them, such as ``content: None`` or an empty list, vanishes from the wire.
"""
content: Final = message_field(message, "content")
if isinstance(content, str):
return True
parts: Final = parts_of(content)
return (
any(_assistant_part_renders(part) for part in parts)
or _separate_thinking_blocks_render(message, parts)
or bool(message_field(message, "tool_calls"))
or bool(message_field(message, "function_call"))
or bool(message_field(message_field(message, "provider_specific_fields"), "compaction_blocks"))
)
def _renders(message: object) -> bool:
"""Whether ``anthropic_messages_pt`` puts a block on the wire for this message.
A tool message always becomes a ``tool_result`` and a user message with string
content always becomes a text block (empty text gets a placeholder). A user list
renders only through parts of a type the converter emits; ``None``, an empty list,
and a list of other parts vanish. Assistant messages follow ``_assistant_renders``.
"""
role: Final = message_field(message, "role")
if role in _TOOL_ROLES:
return True
if role == "assistant":
return _assistant_renders(message)
content: Final = message_field(message, "content")
return isinstance(content, str) or any(
message_field(part, "type") in _RENDERED_PART_TYPES for part in parts_of(content)
)
def _rendered_block(
message_index: int,
messages: Sequence[AllMessageValues],
blocks: Sequence[tuple[bool, tuple[int, ...]]],
) -> int | None:
block_index: Final = _block_containing(message_index, blocks)
_, indices = blocks[block_index]
return block_index if any(_renders(messages[index]) for index in indices) else None
def _system_may_follow(
block_index: int,
messages: Sequence[AllMessageValues],
blocks: Sequence[tuple[bool, tuple[int, ...]]],
) -> bool:
"""Whether a system message behind this block precedes an assistant turn or ends the array on the wire.
Blocks alternate between user-type and assistant, so the check is whether the
first later block that puts anything on the wire is an assistant block.
"""
return next(
(
not is_user
for is_user, indices in blocks[block_index + 1 :]
if any(_renders(messages[index]) for index in indices)
),
True,
)
def _anchor_block(
run: Sequence[int],
messages: Sequence[AllMessageValues],
blocks: Sequence[tuple[bool, tuple[int, ...]]],
) -> int | None:
"""The user-type block a system run must follow, or ``None`` when it converts in place.
The run never starts at 0: the leading system run was split off before this
policy runs, so the message before a run is always a non-system message. Only
the run's neighbours decide, so a request that replays these messages with more
turns appended places the run identically. A block that puts nothing on the wire
cannot anchor a run: the system message would land first or behind an assistant
turn, so the run converts in place instead. The same happens when the assistant
turn after the anchor puts nothing on the wire and a user turn follows it: the
system message would sit directly before that user turn, which Anthropic rejects.
"""
previous: Final = run[0] - 1
neighbour: Final = previous if _is_user_type(messages[previous]) else run[-1] + 1
if neighbour >= len(messages) or not _is_user_type(messages[neighbour]):
return None
block_index: Final = _rendered_block(neighbour, messages, blocks)
if block_index is None or not _system_may_follow(block_index, messages, blocks):
return None
return block_index
def _placed_for_flagged_model(messages: Sequence[AllMessageValues]) -> tuple[AllMessageValues, ...]:
"""Keep system messages as ``role: "system"`` at a placement Anthropic accepts.
A run already sitting after a user-type message stays with that user turn. A
run after an assistant turn moves behind the user turn that immediately follows
it. A run that ends the array or is followed by an assistant turn becomes user
turns in place, so replaying the same messages with more turns appended cannot
move it. Runs that share a user turn merge into one system message.
"""
blocks: Final = _user_type_blocks(messages)
anchors: Final = tuple((run, _anchor_block(run, messages, blocks)) for run in _system_runs(messages))
def messages_of(run: tuple[int, ...]) -> tuple[AllMessageValues, ...]:
return tuple(messages[index] for index in run)
def anchored_to(block_index: int) -> tuple[AllMessageValues, ...]:
anchored_runs: Final = tuple(run for run, anchor in anchors if anchor == block_index)
return tuple(chain.from_iterable(map(messages_of, anchored_runs)))
def converted_after(message_index: int) -> tuple[ChatCompletionUserMessage, ...]:
following_runs: Final = tuple(run for run, anchor in anchors if anchor is None and run[0] == message_index + 1)
return tuple(chain.from_iterable(_converted_user_turns(messages_of(run)) for run in following_runs))
def emit(block_index: int) -> Iterator[AllMessageValues]:
is_user, indices = blocks[block_index]
for index in indices:
yield messages[index]
yield from converted_after(index)
if is_user:
yield from _merged_system_message(anchored_to(block_index))
return tuple(chain.from_iterable(emit(block_index) for block_index in range(len(blocks))))
def place_mid_conversation_system(
messages: Sequence[AllMessageValues],
*,
supports_mid_conversation_system: bool,
) -> tuple[AllMessageValues, ...]:
"""Apply the placement policy to the messages after the leading system run."""
if not any(is_system_message(message) for message in messages):
return tuple(messages)
if supports_mid_conversation_system:
return _placed_for_flagged_model(messages)
return _converted_for_unflagged_model(messages)

View file

@ -1,10 +1,8 @@
"""
This is a cache for LangfuseLoggers.
Langfuse Python SDK initializes a thread for each client.
This ensures we do
1. Proper cleanup of Langfuse initialized clients.
1. Release the initialized-client slot a LangfuseLogger holds when it expires.
2. Re-use created langfuse clients.
"""
@ -21,45 +19,34 @@ from ...caching import InMemoryCache
class LangfuseInMemoryCache(InMemoryCache):
"""
Ensures we do proper cleanup of Langfuse initialized clients.
Decrements ``litellm.initialized_langfuse_clients`` when a LangFuseLogger entry expires.
Langfuse Python SDK initializes a thread for each client, we need to call Langfuse.shutdown() to properly cleanup.
This ensures we do proper cleanup of Langfuse initialized clients.
The counter is a soft budget: loggers built concurrently for one credential set before the
first lands in the cache each take a slot, and only the cached one gives it back on expiry.
The logger's ``stop()`` below hands its shared export channel back
(https://github.com/BerriAI/litellm/issues/11169).
"""
def _remove_key(self, key: str) -> None:
"""
Override _remove_key in InMemoryCache to ensure we do proper cleanup of Langfuse initialized clients.
LangfuseLoggers consume threads when initalized, this shuts them down when they are expired
Relevant Issue: https://github.com/BerriAI/litellm/issues/11169
"""
from litellm.integrations.langfuse.langfuse import LangFuseLogger
if isinstance(self.cache_dict[key], LangFuseLogger):
_created_langfuse_logger: Final[LangFuseLogger] = self.cache_dict[key]
#########################################################
# Clean up Langfuse initialized clients
#########################################################
evicted: Final = self.cache_dict.pop(key, None)
self.ttl_dict.pop(key, None)
if evicted is None:
return
if isinstance(evicted, LangFuseLogger):
litellm.initialized_langfuse_clients -= 1
_created_langfuse_logger.Langfuse.flush()
_created_langfuse_logger.Langfuse.shutdown()
# Loggers with a periodic flush task (e.g. NewRelicMetricsLogger) expose
# stop() so eviction actually ends the task instead of leaking it.
_evicted_stop: Final = getattr(self.cache_dict[key], "stop", None)
if callable(_evicted_stop):
try:
_evicted_stop()
except Exception: # noqa: BLE001 # a failing stop() must not block eviction
verbose_logger.debug("DynamicLoggingCache: stop() raised during eviction", exc_info=True)
#########################################################
# Call parent class to remove key from cache
#########################################################
return super()._remove_key(key)
_evicted_stop: Final = getattr(evicted, "stop", None)
if not callable(_evicted_stop):
return
try:
_evicted_stop()
except Exception: # noqa: BLE001 # a failing stop() must not block eviction
verbose_logger.debug("DynamicLoggingCache: stop() raised during eviction", exc_info=True)
class DynamicLoggingCache:

View file

@ -0,0 +1,20 @@
from typing import Final
from ...caching import InMemoryCache
class ServiceTraceIDCache:
def __init__(self) -> None:
self.cache = InMemoryCache()
def get_cache(self, litellm_call_id: str, service_name: str) -> str | None:
key_name: Final = f"{service_name}:{litellm_call_id}"
response: Final = self.cache.get_cache(key=key_name)
return response
def set_cache(self, litellm_call_id: str, service_name: str, trace_id: str) -> None:
key_name: Final = f"{service_name}:{litellm_call_id}"
self.cache.set_cache(key=key_name, value=trace_id)
in_memory_trace_id_cache: Final = ServiceTraceIDCache()

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