mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'upstream/main' into litellm_scheduler_remove_admitted_requests
# Conflicts: # tests/test_litellm/test_router.py
This commit is contained in:
commit
471cc693e5
805 changed files with 33501 additions and 74401 deletions
|
|
@ -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
|
||||
|
|
|
|||
140
.circleci/scripts/unit_selection.sh
Executable file
140
.circleci/scripts/unit_selection.sh
Executable file
|
|
@ -0,0 +1,140 @@
|
|||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
flag="${1:?usage: unit_selection.sh <codecov flag>}"
|
||||
|
||||
legacy_flags=(
|
||||
caching-local
|
||||
enterprise-package
|
||||
enterprise-routing
|
||||
mcp-integration
|
||||
proxy-db-auth-checks
|
||||
proxy-db-budgets
|
||||
proxy-db-custom-logging
|
||||
proxy-db-db-and-spend
|
||||
proxy-db-endpoints-and-responses
|
||||
proxy-db-guardrails-hooks
|
||||
proxy-db-jwt-and-keys
|
||||
proxy-db-key-generation
|
||||
proxy-db-logging-misc
|
||||
proxy-db-proxy-runtime
|
||||
proxy-db-proxy-server-core
|
||||
proxy-db-proxy-utils
|
||||
proxy-extras
|
||||
proxy-infra
|
||||
)
|
||||
|
||||
legacy_paths() {
|
||||
case "$1" in
|
||||
caching-local) echo tests/unit/caching ;;
|
||||
enterprise-package)
|
||||
echo tests/unit/enterprise/integrations
|
||||
echo tests/unit/enterprise/proxy/auth
|
||||
echo tests/unit/enterprise/proxy/guardrails
|
||||
echo tests/unit/enterprise/proxy/hooks
|
||||
echo tests/unit/enterprise/proxy/management_endpoints
|
||||
echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
|
||||
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
|
||||
enterprise-routing)
|
||||
echo tests/unit/enterprise/enterprise_callbacks/send_emails
|
||||
echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py
|
||||
echo tests/unit/enterprise/proxy/test_enterprise_routes.py
|
||||
echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_access_check.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;;
|
||||
mcp-integration)
|
||||
echo tests/unit/proxy/_experimental/mcp_server
|
||||
echo tests/unit/responses/mcp
|
||||
echo tests/mcp_tests/test_proxy_mcp_e2e.py ;;
|
||||
proxy-db-auth-checks)
|
||||
echo tests/unit/proxy/auth/test_auth_checks.py
|
||||
echo tests/unit/proxy/auth/test_user_api_key_auth.py
|
||||
echo tests/unit/proxy/test_deprecated_key_grace_period.py ;;
|
||||
proxy-db-budgets)
|
||||
echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py
|
||||
echo tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py
|
||||
echo tests/unit/proxy/test_zero_cost_model_budget_bypass.py ;;
|
||||
proxy-db-custom-logging)
|
||||
echo tests/unit/proxy/test_custom_callback_input.py
|
||||
echo tests/unit/proxy/test_custom_logger_s3_gcs.py ;;
|
||||
proxy-db-db-and-spend)
|
||||
echo tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py
|
||||
echo tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py
|
||||
echo tests/unit/proxy/db/test_update_daily_tag_spend.py
|
||||
echo tests/unit/proxy/test_db_schema_changes.py
|
||||
echo tests/unit/proxy/test_prisma_client_backoff_retry.py
|
||||
echo tests/unit/proxy/test_update_spend.py
|
||||
echo tests/unit/skills/test_skills_db.py ;;
|
||||
proxy-db-endpoints-and-responses)
|
||||
echo tests/unit/proxy/auth/test_models_fallback_endpoint.py
|
||||
echo tests/unit/proxy/common_utils/test_check_batch_cost.py
|
||||
echo tests/unit/proxy/common_utils/test_check_responses_cost.py
|
||||
echo tests/unit/proxy/common_utils/test_realtime_cache.py
|
||||
echo tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py
|
||||
echo tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py
|
||||
echo tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py
|
||||
echo tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py
|
||||
echo tests/unit/proxy/response_polling/test_response_polling_handler.py
|
||||
echo tests/unit/proxy/test_custom_tokenizer_bug.py
|
||||
echo tests/unit/proxy/test_get_favicon.py
|
||||
echo tests/unit/proxy/test_get_image.py
|
||||
echo tests/unit/proxy/test_prompt_test_endpoint.py
|
||||
echo tests/unit/proxy/test_reducto_ocr_route.py
|
||||
echo tests/unit/proxy/test_response_polling_pre_call_checks.py
|
||||
echo tests/unit/proxy/test_ui_path_detection.py ;;
|
||||
proxy-db-guardrails-hooks)
|
||||
echo tests/unit/proxy/hooks/test_banned_keyword_list.py
|
||||
echo tests/unit/proxy/test_proxy_setting_guardrails.py
|
||||
echo tests/unit/proxy/test_unit_test_proxy_hooks.py ;;
|
||||
proxy-db-jwt-and-keys)
|
||||
echo tests/unit/proxy/auth/test_jwt.py
|
||||
echo tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py
|
||||
echo tests/unit/proxy/test_proxy_custom_auth.py ;;
|
||||
proxy-db-key-generation) echo tests/unit/proxy/management_endpoints/test_key_generate_prisma.py ;;
|
||||
proxy-db-logging-misc)
|
||||
echo tests/unit/proxy/management_helpers/test_audit_logs_proxy.py
|
||||
echo tests/unit/proxy/spend_tracking/test_search_api_logging.py
|
||||
echo tests/unit/proxy/test_proxy_reject_logging.py ;;
|
||||
proxy-db-proxy-runtime)
|
||||
echo tests/unit/proxy/auth/test_multipart_bypass_repro.py
|
||||
echo tests/unit/proxy/auth/test_proxy_routes.py
|
||||
echo tests/unit/proxy/middleware/test_request_size_limit_middleware.py
|
||||
echo tests/unit/proxy/test_proxy_config_unit_test.py
|
||||
echo tests/unit/proxy/test_proxy_token_counter.py
|
||||
echo tests/unit/proxy/test_server_root_path.py ;;
|
||||
proxy-db-proxy-server-core)
|
||||
echo tests/unit/proxy/test_aproxy_startup.py
|
||||
echo tests/unit/proxy/test_proxy_server.py ;;
|
||||
proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;;
|
||||
proxy-extras) echo tests/unit/litellm_proxy_extras ;;
|
||||
proxy-infra) echo tests/unit/gateway ;;
|
||||
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
|
||||
esac
|
||||
}
|
||||
|
||||
expand() {
|
||||
while read -r path; do
|
||||
if [ -d "$path" ]; then
|
||||
find "$path" -name 'test_*.py'
|
||||
elif [ -f "$path" ]; then
|
||||
echo "$path"
|
||||
else
|
||||
echo "unit_selection.sh: $path does not exist" >&2
|
||||
exit 1
|
||||
fi
|
||||
done
|
||||
}
|
||||
|
||||
if [ "$flag" = unit ]; then
|
||||
comm -23 \
|
||||
<(find tests/unit -name 'test_*.py' | sort) \
|
||||
<(for legacy in "${legacy_flags[@]}"; do legacy_paths "$legacy"; done | expand | sort)
|
||||
exit 0
|
||||
fi
|
||||
|
||||
legacy_paths "$flag" | expand | sort
|
||||
|
|
@ -74,6 +74,7 @@ commands:
|
|||
steps:
|
||||
- run:
|
||||
name: Install Codecov CLI (pinned v11.3.1)
|
||||
when: always
|
||||
command: |
|
||||
curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov
|
||||
curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM
|
||||
|
|
@ -90,7 +91,6 @@ commands:
|
|||
uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)"
|
||||
setup_test_deps:
|
||||
steps:
|
||||
- checkout
|
||||
- install_uv
|
||||
- install_rust
|
||||
- restore_cache:
|
||||
|
|
@ -165,42 +165,72 @@ commands:
|
|||
jobs:
|
||||
unit:
|
||||
parameters:
|
||||
tests_path:
|
||||
type: string
|
||||
default: tests/unit
|
||||
flag:
|
||||
type: string
|
||||
default: unit
|
||||
shards:
|
||||
type: integer
|
||||
default: 6
|
||||
workers:
|
||||
type: integer
|
||||
default: 4
|
||||
dist:
|
||||
type: string
|
||||
default: loadscope
|
||||
base_ref:
|
||||
type: string
|
||||
default: ""
|
||||
pull_request_url:
|
||||
type: string
|
||||
default: ""
|
||||
legacy_mcp_peer:
|
||||
type: boolean
|
||||
default: false
|
||||
reruns:
|
||||
type: integer
|
||||
default: 0
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
parallelism: << parameters.shards >>
|
||||
environment:
|
||||
COVERAGE_CORE: sysmon
|
||||
LITELLM_LOCAL_MODEL_COST_MAP: "True"
|
||||
steps:
|
||||
- setup_test_deps
|
||||
- checkout
|
||||
- skip_unless_relevant:
|
||||
base_ref: << parameters.base_ref >>
|
||||
pull_request_url: << parameters.pull_request_url >>
|
||||
- setup_test_deps
|
||||
- when:
|
||||
condition: << parameters.legacy_mcp_peer >>
|
||||
steps:
|
||||
- run:
|
||||
name: Install MCP SDK1 peer
|
||||
command: |
|
||||
uv venv --python 3.12 .venv-mcp-peer
|
||||
uv pip install --python .venv-mcp-peer 'mcp==1.28.1' 'langchain-mcp-adapters==0.2.1'
|
||||
echo "export MCP_TEST_PEER_PYTHON=$PWD/.venv-mcp-peer/bin/python" >> "$BASH_ENV"
|
||||
- run:
|
||||
name: "Run << parameters.tests_path >> shard"
|
||||
name: "Run << parameters.flag >> shard"
|
||||
no_output_timeout: 20m
|
||||
command: |
|
||||
mkdir -p test-results/<< parameters.flag >>
|
||||
mapfile -t files < <(find << parameters.tests_path >> -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)
|
||||
if [ "${#files[@]}" -eq 0 ]; then echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.tests_path >> files; nothing to run"; exit 0; fi
|
||||
selection="$(bash .circleci/scripts/unit_selection.sh << parameters.flag >>)" || { echo "unit_selection.sh failed for << parameters.flag >>"; exit 1; }
|
||||
[ -n "${selection}" ] || { echo "unit_selection.sh produced no files for << parameters.flag >>"; exit 1; }
|
||||
shard="$(printf '%s\n' "${selection}" | circleci tests split --split-by=timings --timings-type=filename)" || { echo "circleci tests split failed for << parameters.flag >>"; exit 1; }
|
||||
[ -n "${shard}" ] || { echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.flag >> files; nothing to run"; exit 0; }
|
||||
mapfile -t files < <(printf '%s\n' "${shard}")
|
||||
xdist_args=()
|
||||
if [ "<< parameters.workers >>" -gt 0 ]; then xdist_args=(-n << parameters.workers >> --dist=<< parameters.dist >>); fi
|
||||
rerun_args=(-p no:rerunfailures)
|
||||
if [ "<< parameters.reruns >>" -gt 0 ]; then rerun_args=(--reruns << parameters.reruns >> --reruns-delay 1 --rerun-except "from pytest-timeout"); fi
|
||||
test_env=(PATH="$PATH" HOME="$HOME" CI=true COVERAGE_CORE="$COVERAGE_CORE" LITELLM_LOCAL_MODEL_COST_MAP="$LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
if [ -n "${MCP_TEST_PEER_PYTHON:-}" ]; then test_env+=(MCP_TEST_PEER_PYTHON="$MCP_TEST_PEER_PYTHON"); fi
|
||||
set +e
|
||||
uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
|
||||
env -i "${test_env[@]}" \
|
||||
uv run --no-sync pytest "${files[@]}" "${rerun_args[@]}" -p no:pytest-retry --timeout=90 "${xdist_args[@]}" --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
|
||||
status=$?
|
||||
set -e
|
||||
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from the shard; passing"; exit 0; fi
|
||||
|
|
@ -224,6 +254,7 @@ jobs:
|
|||
resource_class: large
|
||||
working_directory: ~/project
|
||||
steps:
|
||||
- checkout
|
||||
- setup_test_deps
|
||||
- run:
|
||||
name: Checkout litellm-docs
|
||||
|
|
@ -250,16 +281,17 @@ jobs:
|
|||
resource_class: large
|
||||
working_directory: ~/project
|
||||
steps:
|
||||
- setup_test_deps
|
||||
- checkout
|
||||
- skip_unless_relevant:
|
||||
base_ref: << parameters.base_ref >>
|
||||
pull_request_url: << parameters.pull_request_url >>
|
||||
- setup_test_deps
|
||||
- start_postgres:
|
||||
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
|
||||
- start_redis
|
||||
- run:
|
||||
name: Run owned integration contracts
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
|
||||
command: env -i PATH="$PATH" HOME="$HOME" CIRCLE_SHA1="$CIRCLE_SHA1" CIRCLE_WORKFLOW_ID="$CIRCLE_WORKFLOW_ID" bash .circleci/scripts/run_integration.sh << parameters.suite >>
|
||||
no_output_timeout: 15m
|
||||
- run:
|
||||
name: Stop owned database and Redis
|
||||
|
|
@ -282,6 +314,61 @@ workflows:
|
|||
- unit:
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-<< matrix.flag >>
|
||||
shards: 1
|
||||
workers: 2
|
||||
reruns: 2
|
||||
matrix:
|
||||
parameters:
|
||||
flag: [caching-local, proxy-extras, enterprise-routing]
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-mcp-integration
|
||||
flag: mcp-integration
|
||||
shards: 1
|
||||
workers: 2
|
||||
legacy_mcp_peer: true
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-<< matrix.flag >>
|
||||
shards: 1
|
||||
reruns: 2
|
||||
matrix:
|
||||
parameters:
|
||||
flag:
|
||||
- enterprise-package
|
||||
- proxy-infra
|
||||
- proxy-db-auth-checks
|
||||
- proxy-db-jwt-and-keys
|
||||
- proxy-db-proxy-server-core
|
||||
- proxy-db-proxy-runtime
|
||||
- proxy-db-custom-logging
|
||||
- proxy-db-logging-misc
|
||||
- proxy-db-db-and-spend
|
||||
- proxy-db-guardrails-hooks
|
||||
- proxy-db-budgets
|
||||
- proxy-db-endpoints-and-responses
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-proxy-db-proxy-utils
|
||||
flag: proxy-db-proxy-utils
|
||||
shards: 1
|
||||
reruns: 2
|
||||
dist: worksteal
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-proxy-db-key-generation
|
||||
flag: proxy-db-key-generation
|
||||
shards: 1
|
||||
workers: 0
|
||||
reruns: 2
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- documentation
|
||||
- integration:
|
||||
name: integration-<< matrix.suite >>
|
||||
|
|
|
|||
6
.github/pull_request_template.md
vendored
6
.github/pull_request_template.md
vendored
|
|
@ -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:
|
||||
|
||||
|
|
|
|||
13
.github/scripts/assert_ci_coverage.py
vendored
13
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -34,7 +34,6 @@ GLOB_CHARS = frozenset("*?")
|
|||
# tests has to be named by some shard or it runs nowhere. A child listed here is
|
||||
# itself decomposed one level deeper and is checked through its own entry.
|
||||
SHARDED_ROOTS: tuple[str, ...] = (
|
||||
"tests/proxy_unit_tests",
|
||||
"tests/test_litellm",
|
||||
"tests/test_litellm/proxy",
|
||||
)
|
||||
|
|
@ -120,6 +119,13 @@ def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
|||
)
|
||||
|
||||
|
||||
def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]:
|
||||
script: Final = repo_root / ".circleci/scripts/unit_selection.sh"
|
||||
if not script.is_file():
|
||||
return frozenset()
|
||||
return frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(_uncommented(script.read_text())))
|
||||
|
||||
|
||||
def _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
||||
return frozenset(
|
||||
match.group(0)
|
||||
|
|
@ -611,7 +617,10 @@ def main() -> int:
|
|||
scalars = _all_scalars()
|
||||
|
||||
integration_paths, ownership_findings = _integration_ownership()
|
||||
test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars) | integration_paths) + ownership_findings
|
||||
test_findings = (
|
||||
_uncovered_tests(allowlist, _invoked_test_tokens(scalars) | _unit_selection_tokens() | integration_paths)
|
||||
+ ownership_findings
|
||||
)
|
||||
dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars))
|
||||
stale_findings = _stale_allowlist_paths(allowlist, test_files=_test_files(), dockerfiles=_dockerfiles())
|
||||
|
||||
|
|
|
|||
44
.github/scripts/read_rc_version.py
vendored
Normal file
44
.github/scripts/read_rc_version.py
vendored
Normal 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))
|
||||
33
.github/workflows/_test-unit-base.yml
vendored
33
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -13,6 +13,15 @@ on:
|
|||
have its path existence-checked like any other token.
|
||||
required: true
|
||||
type: string
|
||||
fork-flag:
|
||||
description: >-
|
||||
Codecov flag of the `.circleci/tests.yml` job that now owns part of
|
||||
this shard. CircleCI does not run on pull requests from forks, so on
|
||||
those events this shard also runs the files
|
||||
`.circleci/scripts/unit_selection.sh` lists for the flag.
|
||||
required: false
|
||||
type: string
|
||||
default: ""
|
||||
workers:
|
||||
description: "Number of pytest-xdist workers"
|
||||
required: false
|
||||
|
|
@ -92,6 +101,7 @@ jobs:
|
|||
pull-requests: read
|
||||
outputs:
|
||||
decision: ${{ steps.changes.outputs.decision }}
|
||||
has-coverage: ${{ steps.tests.outputs.has-coverage }}
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
|
|
@ -160,10 +170,13 @@ jobs:
|
|||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Run tests
|
||||
id: tests
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
timeout-minutes: ${{ inputs.timeout-minutes }}
|
||||
env:
|
||||
TEST_PATH: ${{ inputs.test-path }}
|
||||
FORK_FLAG: ${{ inputs.fork-flag }}
|
||||
IS_FORK: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository }}
|
||||
MAX_FAILURES: ${{ inputs.max-failures }}
|
||||
WORKERS: ${{ inputs.workers }}
|
||||
RERUNS: ${{ inputs.reruns }}
|
||||
|
|
@ -171,9 +184,18 @@ jobs:
|
|||
DIST: ${{ inputs.dist }}
|
||||
COVERAGE_CORE: sysmon
|
||||
run: |
|
||||
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
|
||||
selection="${TEST_PATH}"
|
||||
if [ "${IS_FORK}" = "true" ] && [ -n "${FORK_FLAG}" ]; then
|
||||
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${FORK_FLAG}" | tr '\n' ' ')"
|
||||
fi
|
||||
if [ -z "${selection// /}" ]; then
|
||||
echo "shard selection is empty on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
pytest_args=()
|
||||
existing_paths=0
|
||||
for token in ${TEST_PATH:?}; do
|
||||
for token in ${selection}; do
|
||||
case "${token}" in
|
||||
-*) pytest_args+=("${token}") ;;
|
||||
*)
|
||||
|
|
@ -187,7 +209,7 @@ jobs:
|
|||
esac
|
||||
done
|
||||
if [ "${existing_paths}" -eq 0 ]; then
|
||||
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
|
||||
echo "No path in the selection exists (${selection}); nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
xdist_args=()
|
||||
|
|
@ -209,8 +231,11 @@ jobs:
|
|||
--cov-config=pyproject.toml
|
||||
status=$?
|
||||
set -e
|
||||
if [ -f coverage.xml ]; then
|
||||
echo "has-coverage=true" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
if [ "$status" -eq 5 ]; then
|
||||
echo "pytest collected no tests from ${TEST_PATH}; passing"
|
||||
echo "pytest collected no tests from ${selection}; passing"
|
||||
exit 0
|
||||
fi
|
||||
exit "$status"
|
||||
|
|
@ -226,7 +251,7 @@ jobs:
|
|||
upload-coverage:
|
||||
name: Upload coverage to Codecov
|
||||
needs: run
|
||||
if: always() && needs.run.outputs.decision != 'skip'
|
||||
if: always() && needs.run.outputs.decision != 'skip' && needs.run.outputs.has-coverage == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
13
.github/workflows/compat-matrix-image.yml
vendored
13
.github/workflows/compat-matrix-image.yml
vendored
|
|
@ -4,6 +4,7 @@ on:
|
|||
pull_request:
|
||||
paths:
|
||||
- tests/e2e/claude_code/cron_vm/**
|
||||
- tests/e2e/claude_code/pr_gate_version_resolver.py
|
||||
- .github/workflows/compat-matrix-image.yml
|
||||
workflow_dispatch:
|
||||
|
||||
|
|
@ -28,6 +29,14 @@ jobs:
|
|||
- name: Build the Render cron image
|
||||
run: docker build -f tests/e2e/claude_code/cron_vm/Dockerfile -t compat-matrix:${{ github.sha }} tests/e2e
|
||||
|
||||
- name: Run the pinned binaries as the cron user
|
||||
- name: Resolve and install the Claude Code CLI as the cron user
|
||||
run: |
|
||||
docker run --rm compat-matrix:${{ github.sha }} bash -c 'set -e; whoami; claude --version; gh --version; uv --version'
|
||||
docker run --rm compat-matrix:${{ github.sha }} bash -c '
|
||||
set -euo pipefail
|
||||
whoami
|
||||
gh --version
|
||||
uv --version
|
||||
version="$(uv run --no-project --python 3.12 python /opt/litellm/tests/e2e/claude_code/pr_gate_version_resolver.py)"
|
||||
/opt/litellm/tests/e2e/claude_code/cron_vm/install_claude_code.sh "${version}" /tmp/claude-cli
|
||||
/tmp/claude-cli/claude --version
|
||||
'
|
||||
|
|
|
|||
66
.github/workflows/create-rc-branch.yml
vendored
Normal file
66
.github/workflows/create-rc-branch.yml
vendored
Normal 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}`);
|
||||
12
.github/workflows/test-linting.yml
vendored
12
.github/workflows/test-linting.yml
vendored
|
|
@ -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: |
|
||||
|
|
|
|||
97
.github/workflows/test-unit-proxy-db.yml
vendored
97
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -20,6 +20,12 @@ concurrency:
|
|||
# rather than alphabetical letter ranges. Adding a new test file means adding it
|
||||
# to whichever group it belongs to, not reshuffling slices.
|
||||
#
|
||||
# `.circleci/tests.yml` runs each group's files on same-repo events under the
|
||||
# `proxy-db-<group>` Codecov flag; `.circleci/scripts/unit_selection.sh` holds
|
||||
# the file lists. CircleCI does not build pull requests from forks, so `fork-flag`
|
||||
# makes the shard run that list there. `test-path` keeps the files that still
|
||||
# reach real providers and never left tests/proxy_unit_tests.
|
||||
#
|
||||
# Design targets:
|
||||
# * Every shard runs in <= 7 minutes of wall-clock on the default runner.
|
||||
# Most of a shard's time is pytest plugin load + xdist worker imports +
|
||||
|
|
@ -58,7 +64,7 @@ jobs:
|
|||
proxy-db:
|
||||
needs: assert-shard-coverage
|
||||
# Display only the semantic shard name in the checks UI instead of GHA's
|
||||
# default "proxy-db (key-generation, tests/proxy_unit_tests/…, 0, loadscope, 20)"
|
||||
# default "proxy-db (key-generation, tests/unit/proxy/…, 0, loadscope, 20)"
|
||||
# which includes every matrix field and gets truncated past the test-path.
|
||||
name: ${{ matrix.test-group }}
|
||||
permissions:
|
||||
|
|
@ -71,132 +77,93 @@ jobs:
|
|||
include:
|
||||
# Must run serially — event-loop conflict with the logging worker.
|
||||
- test-group: key-generation
|
||||
test-path: "tests/proxy_unit_tests/test_key_generate_prisma.py"
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-key-generation
|
||||
workers: 0
|
||||
dist: loadscope
|
||||
timeout: 20
|
||||
|
||||
# ---- auth: split into 2 shards ----
|
||||
- test-group: auth-checks
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_auth_checks.py
|
||||
tests/proxy_unit_tests/test_user_api_key_auth.py
|
||||
tests/proxy_unit_tests/test_deprecated_key_grace_period.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-auth-checks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: jwt-and-keys
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_jwt.py
|
||||
tests/proxy_unit_tests/test_jwt_key_mapping.py
|
||||
tests/proxy_unit_tests/test_proxy_custom_auth.py
|
||||
tests/proxy_unit_tests/test_key_generate_dynamodb.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-jwt-and-keys
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
|
||||
- test-group: proxy-utils
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_utils.py"
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-utils
|
||||
workers: 4
|
||||
dist: worksteal
|
||||
timeout: 15
|
||||
|
||||
# ---- proxy server: split into 2 shards ----
|
||||
- test-group: proxy-server-core
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_proxy_server.py
|
||||
tests/proxy_unit_tests/test_aproxy_startup.py
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py"
|
||||
fork-flag: proxy-db-proxy-server-core
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: proxy-runtime
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_proxy_config_unit_test.py
|
||||
tests/proxy_unit_tests/test_proxy_routes.py
|
||||
tests/proxy_unit_tests/test_server_root_path.py
|
||||
tests/proxy_unit_tests/test_proxy_pass_user_config.py
|
||||
tests/proxy_unit_tests/test_proxy_token_counter.py
|
||||
tests/proxy_unit_tests/test_request_size_limit_middleware.py
|
||||
tests/proxy_unit_tests/test_multipart_bypass_repro.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-runtime
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
# ---- logging: split into 2 shards ----
|
||||
- test-group: custom-logging
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_custom_callback_input.py
|
||||
tests/proxy_unit_tests/test_custom_logger_s3_gcs.py
|
||||
tests/proxy_unit_tests/test_proxy_custom_logger.py
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py"
|
||||
fork-flag: proxy-db-custom-logging
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: logging-misc
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_proxy_reject_logging.py
|
||||
tests/proxy_unit_tests/test_audit_logs_proxy.py
|
||||
tests/proxy_unit_tests/test_search_api_logging.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-logging-misc
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: db-and-spend
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_prisma_client_backoff_retry.py
|
||||
tests/proxy_unit_tests/test_db_schema_changes.py
|
||||
tests/proxy_unit_tests/test_e2e_pod_lock_manager.py
|
||||
tests/proxy_unit_tests/test_skills_db.py
|
||||
tests/proxy_unit_tests/test_update_daily_tag_spend.py
|
||||
tests/proxy_unit_tests/test_update_spend.py
|
||||
tests/proxy_unit_tests/test_proxy_encrypt_decrypt.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-db-and-spend
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
# ---- guardrails + budget + hooks: split into 2 ----
|
||||
- test-group: guardrails-hooks
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_proxy_setting_guardrails.py
|
||||
tests/proxy_unit_tests/test_banned_keyword_list.py
|
||||
tests/proxy_unit_tests/test_unit_test_proxy_hooks.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-guardrails-hooks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: budgets
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_default_end_user_budget_simple.py
|
||||
tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py
|
||||
tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-budgets
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: endpoints-and-responses
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_blog_posts_endpoint.py
|
||||
tests/proxy_unit_tests/test_models_fallback_endpoint.py
|
||||
tests/proxy_unit_tests/test_google_endpoint_routing.py
|
||||
tests/proxy_unit_tests/test_google_gemini_proxy_request.py
|
||||
tests/proxy_unit_tests/test_gemini_agents_endpoints.py
|
||||
tests/proxy_unit_tests/test_get_favicon.py
|
||||
tests/proxy_unit_tests/test_get_image.py
|
||||
tests/proxy_unit_tests/test_reducto_ocr_route.py
|
||||
tests/proxy_unit_tests/test_ui_path_detection.py
|
||||
tests/proxy_unit_tests/test_prompt_test_endpoint.py
|
||||
tests/proxy_unit_tests/test_check_batch_cost.py
|
||||
tests/proxy_unit_tests/test_check_responses_cost.py
|
||||
tests/proxy_unit_tests/test_response_polling_handler.py
|
||||
tests/proxy_unit_tests/test_response_polling_pre_call_checks.py
|
||||
tests/proxy_unit_tests/test_realtime_cache.py
|
||||
tests/proxy_unit_tests/test_proxy_exception_mapping.py
|
||||
tests/proxy_unit_tests/test_custom_tokenizer_bug.py
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py"
|
||||
fork-flag: proxy-db-endpoints-and-responses
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: ${{ matrix.test-path }}
|
||||
fork-flag: ${{ matrix.fork-flag }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: 2
|
||||
timeout-minutes: ${{ matrix.timeout }}
|
||||
|
|
|
|||
31
.github/workflows/test-unit.yml
vendored
31
.github/workflows/test-unit.yml
vendored
|
|
@ -31,10 +31,14 @@ concurrency:
|
|||
# number, so a partially-specified entry would fail the call rather than fall
|
||||
# back to the default.
|
||||
#
|
||||
# tests/proxy_unit_tests keeps its own caller (test-unit-proxy-db.yml): it is
|
||||
# already a matrix and carries a shard-coverage guard that reads that file by
|
||||
# name. Folding it in here is a follow-up, together with generalising that guard
|
||||
# into assert_ci_coverage.py.
|
||||
# tests/unit/proxy keeps its own caller (test-unit-proxy-db.yml): it is already
|
||||
# a matrix and carries a shard-coverage guard that reads that file by name.
|
||||
# Folding it in here is a follow-up, together with generalising that guard into
|
||||
# assert_ci_coverage.py.
|
||||
#
|
||||
# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the
|
||||
# shard under the same Codecov flag. CircleCI does not build pull requests from
|
||||
# forks, so the shard still runs those files there and skips them elsewhere.
|
||||
jobs:
|
||||
unit:
|
||||
name: ${{ matrix.shard }}
|
||||
|
|
@ -49,6 +53,7 @@ jobs:
|
|||
- shard: mcp-integration
|
||||
artifact-name: mcp-integration
|
||||
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
|
||||
fork-flag: mcp-integration
|
||||
workers: 2
|
||||
reruns: 0
|
||||
timeout-minutes: 20
|
||||
|
|
@ -65,10 +70,10 @@ jobs:
|
|||
- shard: enterprise-routing
|
||||
artifact-name: enterprise-routing
|
||||
test-path: >-
|
||||
tests/test_litellm/enterprise
|
||||
tests/test_litellm/google_genai
|
||||
tests/test_litellm/router_utils
|
||||
tests/test_litellm/router_strategy
|
||||
fork-flag: enterprise-routing
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
12
Makefile
12
Makefile
|
|
@ -51,8 +51,8 @@ help:
|
|||
@echo " make test-unit-core-utils - Run core utils tests (~32 files)"
|
||||
@echo " make test-unit-other - Run other tests (caching, responses, etc., ~69 files)"
|
||||
@echo " make test-unit-root - Run root-level tests (~34 files)"
|
||||
@echo " make test-proxy-unit-a - Run proxy_unit_tests (a-o, ~20 files)"
|
||||
@echo " make test-proxy-unit-b - Run proxy_unit_tests (p-z, ~28 files)"
|
||||
@echo " make test-proxy-unit-a - Run tests/unit/proxy (a-o)"
|
||||
@echo " make test-proxy-unit-b - Run tests/unit/proxy (p-z)"
|
||||
@echo " make test-integration - Run integration tests"
|
||||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
|
||||
|
|
@ -332,17 +332,17 @@ test-unit-core-utils: install-test-deps
|
|||
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-unit-other: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/test_litellm/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/unit/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-root: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
# Proxy unit tests (tests/proxy_unit_tests split alphabetically)
|
||||
# Proxy unit tests (tests/unit/proxy split alphabetically)
|
||||
test-proxy-unit-a: install-test-deps
|
||||
$(UV_RUN) pytest tests/proxy_unit_tests/test_[a-o]*.py --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/proxy --ignore-glob='tests/unit/proxy/test_[p-z]*.py' --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-proxy-unit-b: install-test-deps
|
||||
$(UV_RUN) pytest tests/proxy_unit_tests/test_[p-z]*.py --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/proxy/test_[p-z]*.py tests/unit/skills --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-integration: install-test-deps
|
||||
$(UV_RUN) pytest tests/ -k "not test_litellm"
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
43
db_scripts/backfill_key_total_spend.sql
Normal file
43
db_scripts/backfill_key_total_spend.sql
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
-- One-shot backfill of LiteLLM_VerificationToken.total_spend (lifetime spend)
|
||||
-- for keys created before the column was introduced in LiteLLM v1.103.0.
|
||||
--
|
||||
-- The column was added with DEFAULT 0 and no backfill, so keys that predate
|
||||
-- the upgrade report lifetime spend below their current period spend. New
|
||||
-- deployments do not need this script: total_spend is updated at request
|
||||
-- time from the moment the release is deployed. Run it only if you want
|
||||
-- pre-upgrade keys to show their historical lifetime spend. It sets lifetime
|
||||
-- spend to at least the current spend on every key, active and archived,
|
||||
-- because current period spend is a valid lower bound on lifetime spend.
|
||||
-- For keys with no budget reset that is already the exact lifetime value;
|
||||
-- for resetting keys it only recovers the current period. It is idempotent:
|
||||
-- it only touches rows where total_spend is below spend, so re-running is a
|
||||
-- no-op. It touches no spend logs and runs in seconds.
|
||||
--
|
||||
-- IMPORTANT caveats before running:
|
||||
--
|
||||
-- 1. Take a backup of the affected tables first:
|
||||
-- pg_dump "$DATABASE_URL" -t '"LiteLLM_VerificationToken"' -t '"LiteLLM_DeletedVerificationToken"' > key_total_spend_backup.sql
|
||||
--
|
||||
-- 2. A key "resets" when its own budget_duration IS NOT NULL, or when its
|
||||
-- budget_id links to a LiteLLM_BudgetTable row whose budget_duration IS
|
||||
-- NOT NULL (a linked budget resets the key's spend each period too). For
|
||||
-- those keys this script only recovers the current period;
|
||||
-- db_scripts/backfill_key_total_spend_from_spend_logs.sql is an optional
|
||||
-- follow-up that rebuilds the earlier periods from LiteLLM_SpendLogs.
|
||||
--
|
||||
-- 3. No proxy restart is needed. The proxy picks up the corrected values on
|
||||
-- its next read of each key.
|
||||
--
|
||||
-- Usage:
|
||||
-- psql "$DATABASE_URL" -f db_scripts/backfill_key_total_spend.sql
|
||||
|
||||
UPDATE "LiteLLM_VerificationToken"
|
||||
SET total_spend = spend
|
||||
WHERE total_spend < spend;
|
||||
|
||||
UPDATE "LiteLLM_DeletedVerificationToken"
|
||||
SET total_spend = spend
|
||||
WHERE total_spend < spend;
|
||||
|
||||
-- Verify: this should return 0.
|
||||
-- SELECT count(*) FROM "LiteLLM_VerificationToken" WHERE total_spend < spend;
|
||||
89
db_scripts/backfill_key_total_spend_from_spend_logs.sql
Normal file
89
db_scripts/backfill_key_total_spend_from_spend_logs.sql
Normal file
|
|
@ -0,0 +1,89 @@
|
|||
-- Optional follow-up to db_scripts/backfill_key_total_spend.sql. Run that
|
||||
-- script first; this one rebuilds earlier budget periods for the keys it
|
||||
-- can only partially fix: keys whose spend resets each period, because their own
|
||||
-- budget_duration IS NOT NULL or because their budget_id links to a
|
||||
-- LiteLLM_BudgetTable row whose budget_duration IS NOT NULL.
|
||||
--
|
||||
-- For those keys the "spend" column only covers the current period, so
|
||||
-- lifetime spend is reconstructed from LiteLLM_SpendLogs. The join matches
|
||||
-- l.api_key against both the stored token and its second sha256
|
||||
-- (encode(sha256(convert_to(token, 'UTF8')), 'hex')), because spend logs
|
||||
-- written by older paths recorded the re-hashed digest instead of the
|
||||
-- token. It is idempotent and never lowers a value: every statement only
|
||||
-- touches rows where total_spend is below the rebuilt sum, so re-running is
|
||||
-- a no-op, and a key whose log history is shorter than its current period
|
||||
-- keeps the value backfill_key_total_spend.sql already gave it.
|
||||
--
|
||||
-- IMPORTANT caveats before running:
|
||||
--
|
||||
-- 1. Take a backup of the affected tables first:
|
||||
-- pg_dump "$DATABASE_URL" -t '"LiteLLM_VerificationToken"' -t '"LiteLLM_DeletedVerificationToken"' > key_total_spend_backup.sql
|
||||
--
|
||||
-- 2. It requires spend logs to have been enabled, and coverage is bounded
|
||||
-- by maximum_spend_logs_retention_period: spend older than the retention
|
||||
-- window is already gone and cannot be recovered.
|
||||
--
|
||||
-- 3. On a large SpendLogs table the join scan is slow, so run it off peak.
|
||||
--
|
||||
-- 4. Run it while the proxy is idle (or with traffic paused). The proxy
|
||||
-- flushes spend logs in batches, so a request that already raised
|
||||
-- total_spend but whose log is still queued is missing from the sum, and
|
||||
-- the rebuilt value would be short by that in-flight amount.
|
||||
--
|
||||
-- 5. A custom token can be deleted and recreated, so the archived table can
|
||||
-- hold several lifetimes of one token. The update only rewrites archived
|
||||
-- rows that reset, and the log sum covers every lifetime of that token.
|
||||
--
|
||||
-- 6. No proxy restart is needed. The proxy picks up the corrected values on
|
||||
-- its next read of each key.
|
||||
--
|
||||
-- Usage:
|
||||
-- psql "$DATABASE_URL" -f db_scripts/backfill_key_total_spend_from_spend_logs.sql
|
||||
|
||||
-- Active keys whose spend resets (own budget_duration, or a linked
|
||||
-- LiteLLM_BudgetTable row with one). Rebuild from LiteLLM_SpendLogs,
|
||||
-- matching api_key against the stored token and its second sha256 digest.
|
||||
UPDATE "LiteLLM_VerificationToken" k
|
||||
SET total_spend = s.sum_spend
|
||||
FROM (
|
||||
SELECT k2.token, SUM(l.spend) AS sum_spend
|
||||
FROM "LiteLLM_VerificationToken" k2
|
||||
JOIN "LiteLLM_SpendLogs" l
|
||||
ON l.api_key IN (k2.token, encode(sha256(convert_to(k2.token, 'UTF8')), 'hex'))
|
||||
WHERE k2.budget_duration IS NOT NULL
|
||||
OR k2.budget_id IN (
|
||||
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
|
||||
)
|
||||
GROUP BY k2.token
|
||||
) s
|
||||
WHERE k.token = s.token
|
||||
AND k.total_spend < s.sum_spend;
|
||||
|
||||
-- Archived tokens are not unique, so collapse them to one row per token
|
||||
-- before joining spend logs; the update then hits every resetting archived
|
||||
-- row.
|
||||
UPDATE "LiteLLM_DeletedVerificationToken" k
|
||||
SET total_spend = s.sum_spend
|
||||
FROM (
|
||||
SELECT k2.token, SUM(l.spend) AS sum_spend
|
||||
FROM (
|
||||
SELECT DISTINCT token
|
||||
FROM "LiteLLM_DeletedVerificationToken"
|
||||
WHERE budget_duration IS NOT NULL
|
||||
OR budget_id IN (
|
||||
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
|
||||
)
|
||||
) k2
|
||||
JOIN "LiteLLM_SpendLogs" l
|
||||
ON l.api_key IN (k2.token, encode(sha256(convert_to(k2.token, 'UTF8')), 'hex'))
|
||||
GROUP BY k2.token
|
||||
) s
|
||||
WHERE k.token = s.token
|
||||
AND k.total_spend < s.sum_spend
|
||||
AND (k.budget_duration IS NOT NULL
|
||||
OR k.budget_id IN (
|
||||
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
|
||||
));
|
||||
|
||||
-- Verify: this should return 0.
|
||||
-- SELECT count(*) FROM "LiteLLM_VerificationToken" WHERE total_spend < spend;
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 {}),
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "kill_switch" JSONB;
|
||||
|
|
@ -72,6 +72,7 @@ model LiteLLM_AgentsTable {
|
|||
agent_card_params Json
|
||||
static_headers Json? @default("{}")
|
||||
extra_headers String[] @default([])
|
||||
kill_switch Json?
|
||||
agent_access_groups String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
object_permission_id String?
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
11
litellm-rust/Cargo.lock
generated
11
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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" }
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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?;
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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`.
|
||||
|
|
|
|||
31
litellm-rust/crates/coroutine/AGENTS.md
Normal file
31
litellm-rust/crates/coroutine/AGENTS.md
Normal 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
|
||||
15
litellm-rust/crates/coroutine/Cargo.toml
Normal file
15
litellm-rust/crates/coroutine/Cargo.toml
Normal 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"] }
|
||||
42
litellm-rust/crates/coroutine/src/co.rs
Normal file
42
litellm-rust/crates/coroutine/src/co.rs
Normal 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
|
||||
}
|
||||
}
|
||||
94
litellm-rust/crates/coroutine/src/coroutine.rs
Normal file
94
litellm-rust/crates/coroutine/src/coroutine.rs
Normal 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() {}
|
||||
}
|
||||
}
|
||||
14
litellm-rust/crates/coroutine/src/error.rs
Normal file
14
litellm-rust/crates/coroutine/src/error.rs
Normal 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;
|
||||
12
litellm-rust/crates/coroutine/src/lib.rs
Normal file
12
litellm-rust/crates/coroutine/src/lib.rs
Normal 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};
|
||||
60
litellm-rust/crates/coroutine/src/reply.rs
Normal file
60
litellm-rust/crates/coroutine/src/reply.rs
Normal 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 })
|
||||
}
|
||||
256
litellm-rust/crates/coroutine/tests/coroutine.rs
Normal file
256
litellm-rust/crates/coroutine/tests/coroutine.rs
Normal 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) });
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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<'_>);
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
);
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
241
litellm-rust/crates/host-python/src/file_reader.rs
Normal file
241
litellm-rust/crates/host-python/src/file_reader.rs
Normal 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");
|
||||
}
|
||||
}
|
||||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
137
litellm-rust/crates/host/src/machine/call_machine.rs
Normal file
137
litellm-rust/crates/host/src/machine/call_machine.rs
Normal 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()) })
|
||||
}
|
||||
}
|
||||
|
|
@ -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>;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()) })
|
||||
}
|
||||
}
|
||||
17
litellm-rust/crates/host/src/protocol.rs
Normal file
17
litellm-rust/crates/host/src/protocol.rs
Normal 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;
|
||||
}
|
||||
|
|
@ -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;
|
||||
}
|
||||
|
|
@ -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]);
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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))))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>> {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 = || {
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
|
|
|||
|
|
@ -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"]);
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)),
|
||||
|
|
|
|||
|
|
@ -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)?))
|
||||
}
|
||||
|
|
|
|||
190
litellm-rust/crates/python-bridge/src/secrets/python.rs
Normal file
190
litellm-rust/crates/python-bridge/src/secrets/python.rs
Normal 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()]);
|
||||
}
|
||||
}
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
1213
litellm/integrations/langfuse/langfuse_sdk.py
Normal file
1213
litellm/integrations/langfuse/langfuse_sdk.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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]}",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue