mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
chore: merge main into litellm_batch_jsonl_line_item_callbacks
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
87d82b9e04
490 changed files with 18674 additions and 70097 deletions
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))
|
||||
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: |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -658,6 +658,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()
|
||||
|
|
@ -894,6 +895,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":
|
||||
|
|
@ -1084,6 +1087,7 @@ model_list = list(
|
|||
| azure_anthropic_models
|
||||
| anyscale_models
|
||||
| cerebras_models
|
||||
| nadir_models
|
||||
| galadriel_models
|
||||
| nvidia_nim_models
|
||||
| nvidia_riva_models
|
||||
|
|
@ -1192,6 +1196,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,
|
||||
|
|
@ -1995,6 +2000,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"
|
||||
|
|
@ -708,6 +712,7 @@ LITELLM_CHAT_PROVIDERS: Final = [
|
|||
"gigachat",
|
||||
"nvidia_nim",
|
||||
"cerebras",
|
||||
"nadir",
|
||||
"baseten",
|
||||
"ai21_chat",
|
||||
"volcengine",
|
||||
|
|
@ -901,6 +906,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",
|
||||
|
|
@ -2122,6 +2128,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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -31,13 +31,17 @@ from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
|||
async_inline_remote_media,
|
||||
inline_remote_image_urls,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import (
|
||||
place_mid_conversation_system,
|
||||
split_leading_system_run,
|
||||
)
|
||||
from litellm.llms.base_llm.base_utils import type_to_response_format_param
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.types.llms.anthropic import (
|
||||
ANTHROPIC_ADVISOR_TOOL_TYPE,
|
||||
ANTHROPIC_BETA_HEADER_VALUES,
|
||||
ANTHROPIC_HOSTED_TOOLS,
|
||||
AllAnthropicMessageValues,
|
||||
AllAnthropicPassThroughMessageValues,
|
||||
AllAnthropicToolsValues,
|
||||
AnthropicCodeExecutionTool,
|
||||
AnthropicComputerTool,
|
||||
|
|
@ -88,6 +92,7 @@ from litellm.utils import (
|
|||
get_max_tokens,
|
||||
has_tool_call_blocks,
|
||||
last_assistant_with_tool_calls_has_no_thinking_blocks,
|
||||
supports_mid_conversation_system,
|
||||
supports_reasoning,
|
||||
token_counter,
|
||||
)
|
||||
|
|
@ -1744,10 +1749,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def add_code_execution_tool(
|
||||
self,
|
||||
messages: list[AllAnthropicMessageValues],
|
||||
messages: list[AllAnthropicPassThroughMessageValues],
|
||||
tools: list[AllAnthropicToolsValues | dict],
|
||||
) -> list[AllAnthropicToolsValues | dict]:
|
||||
"""if 'container_upload' in messages, add code_execution tool"""
|
||||
"""if 'container_upload' in messages, add code_execution tool
|
||||
|
||||
Takes the pass-through union because the translator emits ``role: "system"``
|
||||
in ``messages`` for models that accept it; only ``content`` is read here."""
|
||||
add_code_execution_tool = False
|
||||
for message in messages:
|
||||
message_content = message.get("content", None)
|
||||
|
|
@ -1967,16 +1975,27 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if _name_reverse_map and isinstance(litellm_params, dict):
|
||||
litellm_params[ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY] = _name_reverse_map
|
||||
|
||||
# Separate system prompt from rest of message
|
||||
anthropic_system_message_list: Final = self.translate_system_message(messages=messages)
|
||||
# Only the leading system run becomes the top-level system prompt. A later
|
||||
# system message stays in the conversation: hoisting it rewrites the cached
|
||||
# prefix and re-bills the whole history at cache-write pricing (#36559).
|
||||
leading_system_run, later_messages = split_leading_system_run(messages)
|
||||
anthropic_system_message_list: Final = self.translate_system_message(
|
||||
messages=list(leading_system_run) # mutable-ok: translate_system_message pops from the list it is given
|
||||
)
|
||||
# Handling anthropic API Prompt Caching
|
||||
if len(anthropic_system_message_list) > 0:
|
||||
optional_params["system"] = anthropic_system_message_list
|
||||
conversation: Final = place_mid_conversation_system(
|
||||
later_messages,
|
||||
supports_mid_conversation_system=supports_mid_conversation_system(
|
||||
model=model, custom_llm_provider=self.custom_llm_provider
|
||||
),
|
||||
)
|
||||
# Format rest of message according to anthropic guidelines
|
||||
try:
|
||||
anthropic_messages = anthropic_messages_pt(
|
||||
model=model,
|
||||
messages=messages,
|
||||
messages=list(conversation), # mutable-ok: anthropic_messages_pt rewrites entries in place
|
||||
llm_provider=self._resolved_provider,
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -2,9 +2,7 @@ from collections.abc import Mapping, Sequence
|
|||
from itertools import groupby
|
||||
from typing import Final
|
||||
|
||||
CONVERTED_SYSTEM_NOTE: Final = (
|
||||
"Operator note (not from the user): the following was originally a mid-conversation system-role reminder."
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import CONVERTED_SYSTEM_NOTE
|
||||
|
||||
|
||||
def as_system_content_blocks(value: object) -> list[object]:
|
||||
|
|
|
|||
|
|
@ -1,15 +0,0 @@
|
|||
"""Azure AI OCR module."""
|
||||
|
||||
from .cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from .common_utils import get_azure_ai_ocr_config
|
||||
from .document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
)
|
||||
from .transformation import AzureAIOCRConfig
|
||||
|
||||
__all__ = [
|
||||
"AzureAICohereParseConfig",
|
||||
"AzureAIOCRConfig",
|
||||
"AzureDocumentIntelligenceOCRConfig",
|
||||
"get_azure_ai_ocr_config",
|
||||
]
|
||||
|
|
@ -1,91 +0,0 @@
|
|||
"""Cohere Parse served from Azure AI Foundry (`/providers/cohere/v2/parse`)."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
||||
async_convert_url_to_base64,
|
||||
convert_url_to_base64,
|
||||
)
|
||||
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
|
||||
from litellm.llms.cohere.ocr.transformation import COHERE_PARSE_PATH, CohereParseConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
AZURE_AI_API_KEY_ENV_VAR: Final = "AZURE_AI_API_KEY"
|
||||
AZURE_AI_API_BASE_ENV_VAR: Final = "AZURE_AI_API_BASE"
|
||||
AZURE_AI_COHERE_PROVIDER_PATH: Final = "/providers/cohere"
|
||||
AZURE_AI_MODELS_PATH_SUFFIX: Final = "/models"
|
||||
|
||||
|
||||
class AzureAICohereParseConfig(CohereParseConfig):
|
||||
"""Same request and response shape as Cohere Parse, behind Azure AI auth and URL layout.
|
||||
|
||||
Foundry cannot fetch external URLs, so remote images are inlined as base64 data URIs.
|
||||
"""
|
||||
|
||||
def get_api_key_env_var(self) -> str | None:
|
||||
return AZURE_AI_API_KEY_ENV_VAR
|
||||
|
||||
def _llm_provider(self) -> str:
|
||||
return "azure_ai"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.validate_environment signature
|
||||
) -> dict[str, str]: # mutable-ok: BaseOCRConfig signature
|
||||
resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR)
|
||||
if resolved_base is None:
|
||||
raise ValueError(
|
||||
f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable "
|
||||
"or pass api_base parameter"
|
||||
)
|
||||
resolved_key: Final = api_key or get_secret_str(AZURE_AI_API_KEY_ENV_VAR)
|
||||
return { # mutable-ok: BaseOCRConfig signature
|
||||
**get_azure_ai_auth_headers(api_key=resolved_key, litellm_params=litellm_params),
|
||||
"Content-Type": "application/json",
|
||||
**headers,
|
||||
}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.get_complete_url signature
|
||||
) -> str:
|
||||
resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR)
|
||||
if resolved_base is None:
|
||||
raise ValueError(
|
||||
f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable "
|
||||
"or pass api_base parameter"
|
||||
)
|
||||
url: Final = httpx.URL(resolved_base)
|
||||
if not url.is_absolute_url:
|
||||
raise ValueError(
|
||||
"Azure AI API Base must be an absolute URL including scheme (e.g. "
|
||||
f"'https://<resource>.services.ai.azure.com'). Got api_base={resolved_base!r}."
|
||||
)
|
||||
path: Final = url.path.rstrip("/")
|
||||
if path.endswith(COHERE_PARSE_PATH):
|
||||
return str(url.copy_with(path=path))
|
||||
if path.endswith(f"{AZURE_AI_COHERE_PROVIDER_PATH}/v2"):
|
||||
return str(url.copy_with(path=f"{path}/parse"))
|
||||
return str(
|
||||
url.copy_with(
|
||||
path=f"{path.removesuffix(AZURE_AI_MODELS_PATH_SUFFIX)}{AZURE_AI_COHERE_PROVIDER_PATH}{COHERE_PARSE_PATH}"
|
||||
)
|
||||
)
|
||||
|
||||
def _resolve_image_url_sync(self, image_url: str) -> str:
|
||||
return convert_url_to_base64(image_url)
|
||||
|
||||
async def _resolve_image_url_async(self, image_url: str) -> str:
|
||||
return await async_convert_url_to_base64(image_url)
|
||||
|
|
@ -1,71 +0,0 @@
|
|||
"""
|
||||
Common utilities for Azure AI OCR providers.
|
||||
|
||||
This module provides routing logic to determine which OCR configuration to use
|
||||
based on the model name.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Final, Optional
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
|
||||
|
||||
|
||||
def is_azure_document_intelligence_model(model: str) -> bool:
|
||||
"""Whether an azure_ai OCR model routes to Azure Document Intelligence.
|
||||
|
||||
Azure AI exposes two OCR services on the same provider; the sub-route in the
|
||||
model name (`azure_ai/doc-intelligence/<model>`) selects Document Intelligence
|
||||
over Mistral OCR. This is the single source of truth for that routing decision.
|
||||
"""
|
||||
lowered: Final = model.lower()
|
||||
return "doc-intelligence" in lowered or "documentintelligence" in lowered
|
||||
|
||||
|
||||
def is_azure_cohere_parse_model(model: str) -> bool:
|
||||
lowered: Final = model.lower()
|
||||
return "cohere" in lowered and "parse" in lowered
|
||||
|
||||
|
||||
def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
|
||||
"""
|
||||
Determine which Azure AI OCR configuration to use based on the model name.
|
||||
|
||||
Azure AI supports multiple OCR services:
|
||||
- Azure Document Intelligence: azure_ai/doc-intelligence/<model>
|
||||
- Mistral OCR (via Azure AI): azure_ai/<model>
|
||||
|
||||
Args:
|
||||
model: The model name (e.g., "azure_ai/doc-intelligence/prebuilt-read",
|
||||
"azure_ai/pixtral-12b-2409")
|
||||
|
||||
Returns:
|
||||
OCR configuration instance for the specified model
|
||||
|
||||
Examples:
|
||||
>>> get_azure_ai_ocr_config("azure_ai/doc-intelligence/prebuilt-read")
|
||||
<AzureDocumentIntelligenceOCRConfig object>
|
||||
|
||||
>>> get_azure_ai_ocr_config("azure_ai/pixtral-12b-2409")
|
||||
<AzureAIOCRConfig object>
|
||||
"""
|
||||
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
)
|
||||
from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
|
||||
|
||||
# Check for Azure Document Intelligence models
|
||||
if is_azure_document_intelligence_model(model):
|
||||
verbose_logger.debug("Routing %s to Azure Document Intelligence OCR config", model)
|
||||
return AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
if is_azure_cohere_parse_model(model):
|
||||
verbose_logger.debug("Routing %s to Azure AI Cohere Parse config", model)
|
||||
return AzureAICohereParseConfig()
|
||||
|
||||
# Default to Mistral-based OCR for other azure_ai models
|
||||
verbose_logger.debug("Routing %s to Azure AI (Mistral) OCR config", model)
|
||||
return AzureAIOCRConfig()
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
"""Azure Document Intelligence OCR module."""
|
||||
|
||||
from .transformation import AzureDocumentIntelligenceOCRConfig
|
||||
|
||||
__all__ = ["AzureDocumentIntelligenceOCRConfig"]
|
||||
|
|
@ -1,806 +0,0 @@
|
|||
"""
|
||||
Azure Document Intelligence OCR transformation implementation.
|
||||
|
||||
Azure Document Intelligence (formerly Form Recognizer) provides advanced document analysis capabilities.
|
||||
This implementation transforms between Mistral OCR format and Azure Document Intelligence API v4.0.
|
||||
|
||||
Note: Azure Document Intelligence API is async - POST returns 202 Accepted with Operation-Location header.
|
||||
The operation location must be polled until the analysis completes.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import (
|
||||
AZURE_DOCUMENT_INTELLIGENCE_API_VERSION,
|
||||
AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI,
|
||||
AZURE_OPERATION_POLLING_TIMEOUT,
|
||||
)
|
||||
from litellm.exceptions import UnsupportedParamsError
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin, encode_url_path_segment
|
||||
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
|
||||
from litellm.llms.base_llm.ocr.transformation import (
|
||||
OCR_REQUEST_FORMAT_PARAM,
|
||||
BaseOCRConfig,
|
||||
DocumentType,
|
||||
OCRPage,
|
||||
OCRPageDimensions,
|
||||
OCRRequestData,
|
||||
OCRRequestFormat,
|
||||
OCRResponse,
|
||||
OCRUsageInfo,
|
||||
parse_ocr_request_format,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR: Final = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"
|
||||
|
||||
|
||||
class AzureDocumentIntelligenceLine(BaseModel):
|
||||
content: str | None = None
|
||||
|
||||
|
||||
class AzureDocumentIntelligencePage(BaseModel):
|
||||
pageNumber: int | None = None
|
||||
width: float | None = None
|
||||
height: float | None = None
|
||||
unit: str | None = None
|
||||
lines: tuple[AzureDocumentIntelligenceLine, ...] = ()
|
||||
|
||||
|
||||
class AzureDocumentIntelligenceAnalyzeResult(BaseModel):
|
||||
content: str | None = None
|
||||
pages: tuple[AzureDocumentIntelligencePage, ...] = ()
|
||||
tables: list[dict[str, object]] | None = None
|
||||
keyValuePairs: list[dict[str, object]] | None = None
|
||||
|
||||
|
||||
class AzureDocumentIntelligenceOperation(BaseModel):
|
||||
status: str | None = None
|
||||
analyzeResult: AzureDocumentIntelligenceAnalyzeResult | None = None
|
||||
|
||||
|
||||
class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
||||
"""
|
||||
Azure Document Intelligence OCR transformation configuration.
|
||||
|
||||
Supports Azure Document Intelligence v4.0 (2024-11-30) API.
|
||||
Model route: azure_ai/doc-intelligence/<model>
|
||||
|
||||
Supported models:
|
||||
- prebuilt-layout: Extracts text with markdown, tables, and structure (closest to Mistral OCR)
|
||||
- prebuilt-read: Basic text extraction optimized for reading
|
||||
- prebuilt-document: General document analysis
|
||||
|
||||
Reference: https://learn.microsoft.com/en-us/azure/ai-services/document-intelligence/
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def get_api_key_env_var(self) -> str | None:
|
||||
return AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR
|
||||
|
||||
def resolve_connection_params(
|
||||
self,
|
||||
*,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
dynamic_api_key: str | None,
|
||||
dynamic_api_base: str | None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
explicit_api_key: Final = None if api_key is None else dynamic_api_key or api_key
|
||||
explicit_api_base: Final = None if api_base is None else dynamic_api_base or api_base
|
||||
return explicit_api_key, explicit_api_base
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list:
|
||||
"""
|
||||
Get supported OCR parameters for Azure Document Intelligence.
|
||||
|
||||
Azure DI exposes a `pages` query parameter on the analyze endpoint
|
||||
(1-based, e.g. "1-3,5,7-9"). To keep the public request shape
|
||||
aligned with Mistral OCR, callers pass `pages` using Mistral
|
||||
semantics — a list of 0-based integers — or a pre-formatted
|
||||
Azure-style string. Azure DI also exposes a `features` query
|
||||
parameter enabling add-on capabilities (e.g. "keyValuePairs",
|
||||
"languages"), passed as a list of feature names or a
|
||||
comma-separated string. Other Mistral-specific params (e.g.
|
||||
`include_image_base64`) are not supported by Azure DI and are
|
||||
ignored during transformation.
|
||||
|
||||
`req_format` selects the response shape: "litellm" (default) returns
|
||||
the normalized OCR schema, "native" returns Azure DI's own analyze
|
||||
operation payload as-is.
|
||||
"""
|
||||
return ["pages", "features", OCR_REQUEST_FORMAT_PARAM]
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OCR params to Azure DI format.
|
||||
|
||||
Translates Mistral-style `pages` (list[int], 0-based) into Azure's
|
||||
`pages` query string (1-based, e.g. "1,2,3" or "1-3,5"). A raw
|
||||
string that already matches Azure's format is passed through
|
||||
unchanged. `features` (list[str] or comma-separated string) is
|
||||
normalized into Azure's comma-joined `features` query string.
|
||||
"""
|
||||
pages: Final = non_default_params.get("pages")
|
||||
features: Final = non_default_params.get("features")
|
||||
request_format: Final = non_default_params.get(OCR_REQUEST_FORMAT_PARAM)
|
||||
normalized_pages: Final = self._normalize_pages_param(pages) if pages is not None else ""
|
||||
normalized_features: Final = self._normalize_features_param(features) if features is not None else ""
|
||||
return {
|
||||
**optional_params,
|
||||
**({"pages": normalized_pages} if normalized_pages else {}),
|
||||
**({"features": normalized_features} if normalized_features else {}),
|
||||
**(
|
||||
{OCR_REQUEST_FORMAT_PARAM: self._parse_request_format(request_format, model)}
|
||||
if request_format is not None
|
||||
else {}
|
||||
),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _parse_request_format(request_format: object, model: str) -> OCRRequestFormat:
|
||||
try:
|
||||
return parse_ocr_request_format(request_format)
|
||||
except ValueError as e:
|
||||
raise UnsupportedParamsError(message=f"{e}", model=model, llm_provider="azure_ai") from e
|
||||
|
||||
@staticmethod
|
||||
def _normalize_pages_param(pages: object) -> str:
|
||||
"""
|
||||
Convert a caller-provided `pages` value to Azure DI's query-string
|
||||
form. Azure expects 1-based page numbers, grammar: `^(\\d+(-\\d+)?)(,\\s*(\\d+(-\\d+)?))*$`.
|
||||
|
||||
Accepted inputs:
|
||||
- list[int]: Mistral-style 0-based indices. Converted to 1-based
|
||||
and joined (e.g. [0,1,2] -> "1,2,3").
|
||||
- list[str]: tokens like "1" or "3-5". Validated, joined as-is
|
||||
(treated as Azure-native, i.e. 1-based).
|
||||
- str: already in Azure format. Validated and whitespace-stripped.
|
||||
"""
|
||||
pages_pattern: Final = re.compile(r"^\s*\d+(-\d+)?(\s*,\s*\d+(-\d+)?)*\s*$")
|
||||
|
||||
if isinstance(pages, str):
|
||||
if not pages_pattern.match(pages):
|
||||
raise ValueError(
|
||||
f"Invalid `pages` string for Azure Document Intelligence: "
|
||||
f"{pages!r}. Expected format like '1-3,5,7-9'."
|
||||
)
|
||||
return pages.replace(" ", "")
|
||||
|
||||
if isinstance(pages, list):
|
||||
if len(pages) == 0:
|
||||
return ""
|
||||
if any(isinstance(p, bool) for p in pages):
|
||||
raise ValueError("`pages` must be integers, not booleans")
|
||||
if all(isinstance(p, int) for p in pages):
|
||||
if any(p < 0 for p in pages):
|
||||
raise ValueError("`pages` integers must be >= 0 (Mistral 0-based indices)")
|
||||
# Mistral 0-based -> Azure 1-based.
|
||||
return ",".join(str(p + 1) for p in sorted(set(pages)))
|
||||
if all(isinstance(p, str) for p in pages):
|
||||
joined: Final = ",".join(p.strip() for p in pages)
|
||||
if not pages_pattern.match(joined):
|
||||
raise ValueError(
|
||||
f"Invalid `pages` list for Azure Document Intelligence: "
|
||||
f"{pages!r}. Expected tokens like '1' or '3-5'."
|
||||
)
|
||||
return joined
|
||||
|
||||
raise ValueError("`pages` must be a list[int] (0-based, Mistral-style) or a string like '1-3,5,7-9'.")
|
||||
|
||||
@staticmethod
|
||||
def _normalize_features_param(features: object) -> str:
|
||||
"""
|
||||
Convert a caller-provided `features` value to Azure DI's query-string
|
||||
form (comma-joined feature names, e.g. "keyValuePairs,languages").
|
||||
|
||||
Accepted inputs:
|
||||
- list[str]: feature names like ["keyValuePairs", "languages"].
|
||||
- str: a single feature name or comma-separated names.
|
||||
"""
|
||||
invalid_features_error: Final = ValueError(
|
||||
f"Invalid `features` for Azure Document Intelligence: {features!r}. "
|
||||
f"Expected a list of feature names or a comma-separated string like "
|
||||
f"'keyValuePairs' or 'keyValuePairs,languages'."
|
||||
)
|
||||
|
||||
if isinstance(features, str):
|
||||
raw_tokens = features.split(",")
|
||||
elif isinstance(features, list):
|
||||
if len(features) == 0:
|
||||
return ""
|
||||
raw_tokens = [feature for feature in features if isinstance(feature, str)]
|
||||
if len(raw_tokens) != len(features):
|
||||
raise invalid_features_error
|
||||
else:
|
||||
raise invalid_features_error
|
||||
|
||||
tokens: Final = tuple(token.strip() for token in raw_tokens)
|
||||
feature_pattern: Final = re.compile(r"^[A-Za-z][A-Za-z0-9]*$")
|
||||
if not all(feature_pattern.match(token) for token in tokens):
|
||||
raise invalid_features_error
|
||||
return ",".join(tokens)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
**kwargs,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate environment and return headers for Azure Document Intelligence.
|
||||
|
||||
Authentication uses the Ocp-Apim-Subscription-Key header, or an Entra ID / OAuth bearer
|
||||
token when no subscription key is set.
|
||||
"""
|
||||
# Get API key from environment if not provided
|
||||
if api_key is None:
|
||||
api_key = get_secret_str(AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR)
|
||||
|
||||
# Validate API base/endpoint is provided
|
||||
if api_base is None:
|
||||
api_base = get_secret_str("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"Missing Azure Document Intelligence Endpoint - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT environment variable or pass api_base parameter"
|
||||
)
|
||||
|
||||
headers = {
|
||||
**get_azure_ai_auth_headers(
|
||||
api_key=api_key,
|
||||
litellm_params=litellm_params,
|
||||
api_key_header="Ocp-Apim-Subscription-Key",
|
||||
api_key_env_var=AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR,
|
||||
),
|
||||
"Content-Type": "application/json",
|
||||
**headers,
|
||||
}
|
||||
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict | None = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Get complete URL for Azure Document Intelligence endpoint.
|
||||
|
||||
Format: {endpoint}/documentintelligence/documentModels/{modelId}:analyze?api-version=2024-11-30
|
||||
|
||||
Note: API version 2024-11-30 uses /documentintelligence/ path (not /formrecognizer/)
|
||||
|
||||
Args:
|
||||
api_base: Azure Document Intelligence endpoint (e.g., https://your-resource.cognitiveservices.azure.com)
|
||||
model: Model ID (e.g., "prebuilt-layout", "prebuilt-read")
|
||||
optional_params: Optional parameters
|
||||
|
||||
Returns: Complete URL for Azure DI analyze endpoint
|
||||
"""
|
||||
if api_base is None:
|
||||
api_base = get_secret_str("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"Missing Azure Document Intelligence Endpoint - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT environment variable or pass api_base parameter"
|
||||
)
|
||||
|
||||
# Ensure no trailing slash
|
||||
api_base = api_base.rstrip("/")
|
||||
|
||||
# Extract model ID from full model path if needed
|
||||
# Model can be "prebuilt-layout" or "azure_ai/doc-intelligence/prebuilt-layout"
|
||||
model_id = model
|
||||
if "/" in model:
|
||||
# Extract the last part after the last slash
|
||||
model_id = model.split("/")[-1]
|
||||
encoded_model_id: Final = encode_url_path_segment(model_id, field_name="model_id")
|
||||
|
||||
# Azure Document Intelligence analyze endpoint
|
||||
# Note: API version 2024-11-30+ uses /documentintelligence/ (not /formrecognizer/)
|
||||
url: Final = (
|
||||
f"{api_base}/documentintelligence/documentModels/{encoded_model_id}:analyze"
|
||||
f"?api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}"
|
||||
)
|
||||
|
||||
# Azure DI accepts `pages` (1-based, e.g. "1-3,5") and `features`
|
||||
# (comma-joined names, e.g. "keyValuePairs") as query params.
|
||||
# `optional_params` has already been normalized in `map_ocr_params`.
|
||||
pages: Final = optional_params.get("pages") if optional_params else None
|
||||
features: Final = optional_params.get("features") if optional_params else None
|
||||
pages_query: Final = f"&pages={quote(str(pages), safe=',-')}" if pages else ""
|
||||
features_query: Final = f"&features={quote(str(features), safe=',')}" if features else ""
|
||||
|
||||
return f"{url}{pages_query}{features_query}"
|
||||
|
||||
def _extract_base64_from_data_uri(self, data_uri: str) -> str:
|
||||
"""
|
||||
Extract base64 content from a data URI.
|
||||
|
||||
Args:
|
||||
data_uri: Data URI like "data:application/pdf;base64,..."
|
||||
|
||||
Returns:
|
||||
Base64 string without the data URI prefix
|
||||
"""
|
||||
# Match pattern: data:[<mediatype>][;base64],<data>
|
||||
match: Final = re.match(r"data:([^;]+)(?:;base64)?,(.+)", data_uri)
|
||||
if match:
|
||||
return match.group(2)
|
||||
return data_uri
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Transform OCR request to Azure Document Intelligence format.
|
||||
|
||||
Mistral OCR format:
|
||||
{
|
||||
"document": {
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/doc.pdf"
|
||||
}
|
||||
}
|
||||
|
||||
Azure DI format:
|
||||
{
|
||||
"urlSource": "https://example.com/doc.pdf"
|
||||
}
|
||||
OR
|
||||
{
|
||||
"base64Source": "base64_encoded_content"
|
||||
}
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
document: Document dict from user (Mistral format)
|
||||
optional_params: Already mapped optional parameters
|
||||
headers: Request headers
|
||||
|
||||
Returns:
|
||||
OCRRequestData with JSON data
|
||||
"""
|
||||
verbose_logger.debug("Azure Document Intelligence transform_ocr_request - model: %s", model)
|
||||
|
||||
if not isinstance(document, dict):
|
||||
raise ValueError(f"Expected document dict, got {type(document)}")
|
||||
|
||||
# Extract document URL from Mistral format
|
||||
doc_type: Final = document.get("type")
|
||||
document_url = None
|
||||
|
||||
if doc_type == "document_url":
|
||||
document_url = document.get("document_url", "")
|
||||
elif doc_type == "image_url":
|
||||
document_url = document.get("image_url", "")
|
||||
else:
|
||||
raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url' or 'image_url'")
|
||||
|
||||
if not document_url:
|
||||
raise ValueError("Document URL is required")
|
||||
|
||||
# Build Azure DI request
|
||||
data: Final[dict[str, str]] = {}
|
||||
|
||||
# Check if it's a data URI (base64)
|
||||
if document_url.startswith("data:"):
|
||||
# Extract base64 content
|
||||
base64_content: Final = self._extract_base64_from_data_uri(document_url)
|
||||
data["base64Source"] = base64_content
|
||||
verbose_logger.debug("Using base64Source for Azure Document Intelligence")
|
||||
else:
|
||||
# Regular URL
|
||||
data["urlSource"] = document_url
|
||||
verbose_logger.debug("Using urlSource for Azure Document Intelligence")
|
||||
|
||||
# Azure DI: `pages` is a query param (wired in get_complete_url),
|
||||
# not a body field. Other Mistral-specific params (e.g.
|
||||
# include_image_base64, image_limit) are unsupported and ignored.
|
||||
|
||||
return OCRRequestData(data=data, files=None)
|
||||
|
||||
def _transform_azure_page(self, azure_page: AzureDocumentIntelligencePage) -> OCRPage:
|
||||
page_number: Final = azure_page.pageNumber if azure_page.pageNumber is not None else 1
|
||||
markdown: Final = "\n".join(line.content or "" for line in azure_page.lines)
|
||||
dimensions: Final = self._convert_dimensions(
|
||||
width=azure_page.width if azure_page.width is not None else 8.5,
|
||||
height=azure_page.height if azure_page.height is not None else 11,
|
||||
unit=azure_page.unit if azure_page.unit is not None else "inch",
|
||||
)
|
||||
return OCRPage(index=page_number - 1, markdown=markdown, dimensions=dimensions)
|
||||
|
||||
def _convert_dimensions(self, width: float, height: float, unit: str) -> OCRPageDimensions:
|
||||
"""
|
||||
Convert Azure DI dimensions to pixels.
|
||||
|
||||
Azure DI provides dimensions in inches. We convert to pixels using configured DPI.
|
||||
|
||||
Args:
|
||||
width: Width in specified unit
|
||||
height: Height in specified unit
|
||||
unit: Unit of measurement (e.g., "inch")
|
||||
|
||||
Returns:
|
||||
OCRPageDimensions with pixel values
|
||||
"""
|
||||
# Convert to pixels using configured DPI
|
||||
dpi: Final = AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI
|
||||
if unit == "inch":
|
||||
width_px = int(width * dpi)
|
||||
height_px = int(height * dpi)
|
||||
else:
|
||||
# If unit is not inches, assume it's already in pixels
|
||||
width_px = int(width)
|
||||
height_px = int(height)
|
||||
|
||||
return OCRPageDimensions(width=width_px, height=height_px, dpi=dpi)
|
||||
|
||||
@staticmethod
|
||||
def _check_timeout(start_time: float, timeout_secs: int) -> None:
|
||||
"""
|
||||
Check if operation has timed out.
|
||||
|
||||
Args:
|
||||
start_time: Start time of the operation
|
||||
timeout_secs: Timeout duration in seconds
|
||||
|
||||
Raises:
|
||||
TimeoutError: If operation has exceeded timeout
|
||||
"""
|
||||
if time.time() - start_time > timeout_secs:
|
||||
raise TimeoutError(f"Azure Document Intelligence operation polling timed out after {timeout_secs} seconds")
|
||||
|
||||
@staticmethod
|
||||
def _get_retry_after(response: httpx.Response) -> int:
|
||||
"""
|
||||
Get retry-after duration from response headers.
|
||||
|
||||
Args:
|
||||
response: HTTP response
|
||||
|
||||
Returns:
|
||||
Retry-after duration in seconds (default: 2)
|
||||
"""
|
||||
retry_after: Final = int(response.headers.get("retry-after", "2"))
|
||||
verbose_logger.debug("Retry polling after: %s seconds", retry_after)
|
||||
return retry_after
|
||||
|
||||
@staticmethod
|
||||
def _check_operation_status(response: httpx.Response) -> str:
|
||||
"""
|
||||
Check Azure DI operation status from response.
|
||||
|
||||
Args:
|
||||
response: HTTP response from operation endpoint
|
||||
|
||||
Returns:
|
||||
Operation status string
|
||||
|
||||
Raises:
|
||||
ValueError: If operation failed or status is unknown
|
||||
"""
|
||||
try:
|
||||
result: Final = response.json()
|
||||
status: Final = result.get("status")
|
||||
|
||||
verbose_logger.debug("Azure DI operation status: %s", status)
|
||||
|
||||
if status == "succeeded":
|
||||
return "succeeded"
|
||||
elif status == "failed":
|
||||
error_msg: Final = result.get("error", {}).get("message", "Unknown error")
|
||||
raise ValueError(f"Azure Document Intelligence analysis failed: {error_msg}")
|
||||
elif status in ["running", "notStarted"]:
|
||||
return "running"
|
||||
else:
|
||||
raise ValueError(f"Unknown operation status: {status}")
|
||||
|
||||
except Exception as e:
|
||||
if "succeeded" in str(e) or "failed" in str(e):
|
||||
raise
|
||||
# If we can't parse JSON, something went wrong
|
||||
raise ValueError(f"Failed to parse Azure DI operation response: {e}")
|
||||
|
||||
def _poll_operation_sync(
|
||||
self,
|
||||
operation_url: str,
|
||||
headers: dict[str, str],
|
||||
timeout_secs: int,
|
||||
) -> httpx.Response:
|
||||
"""
|
||||
Poll Azure Document Intelligence operation until completion (sync).
|
||||
|
||||
Azure DI POST returns 202 with Operation-Location header.
|
||||
We need to poll that URL until status is "succeeded" or "failed".
|
||||
|
||||
Args:
|
||||
operation_url: The Operation-Location URL to poll
|
||||
headers: Request headers (including auth)
|
||||
timeout_secs: Total timeout in seconds
|
||||
|
||||
Returns:
|
||||
Final response with completed analysis
|
||||
"""
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
|
||||
client: Final = _get_httpx_client()
|
||||
start_time: Final = time.time()
|
||||
|
||||
verbose_logger.debug("Polling Azure DI operation: %s", operation_url)
|
||||
|
||||
while True:
|
||||
self._check_timeout(start_time=start_time, timeout_secs=timeout_secs)
|
||||
|
||||
# Poll the operation status
|
||||
response = client.get(url=operation_url, headers=headers)
|
||||
|
||||
# Check operation status
|
||||
status = self._check_operation_status(response=response)
|
||||
|
||||
if status == "succeeded":
|
||||
return response
|
||||
elif status == "running":
|
||||
# Wait before polling again
|
||||
retry_after = self._get_retry_after(response=response)
|
||||
time.sleep(retry_after)
|
||||
|
||||
async def _poll_operation_async(
|
||||
self,
|
||||
operation_url: str,
|
||||
headers: dict[str, str],
|
||||
timeout_secs: int,
|
||||
) -> httpx.Response:
|
||||
"""
|
||||
Poll Azure Document Intelligence operation until completion (async).
|
||||
|
||||
Args:
|
||||
operation_url: The Operation-Location URL to poll
|
||||
headers: Request headers (including auth)
|
||||
timeout_secs: Total timeout in seconds
|
||||
|
||||
Returns:
|
||||
Final response with completed analysis
|
||||
"""
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
|
||||
client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.AZURE_AI)
|
||||
start_time: Final = time.time()
|
||||
|
||||
verbose_logger.debug("Polling Azure DI operation (async): %s", operation_url)
|
||||
|
||||
while True:
|
||||
self._check_timeout(start_time=start_time, timeout_secs=timeout_secs)
|
||||
|
||||
# Poll the operation status
|
||||
response = await client.get(url=operation_url, headers=headers)
|
||||
|
||||
# Check operation status
|
||||
status = self._check_operation_status(response=response)
|
||||
|
||||
if status == "succeeded":
|
||||
return response
|
||||
elif status == "running":
|
||||
# Wait before polling again
|
||||
retry_after = self._get_retry_after(response=response)
|
||||
await asyncio.sleep(retry_after)
|
||||
|
||||
def _get_polling_target(self, raw_response: httpx.Response) -> tuple[str, dict[str, str]]:
|
||||
operation_url: Final = raw_response.headers.get("Operation-Location")
|
||||
if not operation_url:
|
||||
raise ValueError("Azure Document Intelligence returned 202 but no Operation-Location header found")
|
||||
|
||||
# Reject cross-origin polling URLs — the auth headers
|
||||
# below would otherwise leak to whatever URL the upstream
|
||||
# (or an attacker-controlled upstream) returns. VERIA-51.
|
||||
try:
|
||||
assert_same_origin(operation_url, str(raw_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise ValueError(f"Azure Document Intelligence: rejected polling URL ({ssrf_err})")
|
||||
|
||||
poll_headers: Final = {
|
||||
header: raw_response.request.headers[header]
|
||||
for header in ("Ocp-Apim-Subscription-Key", "Authorization")
|
||||
if header in raw_response.request.headers
|
||||
}
|
||||
return operation_url, poll_headers
|
||||
|
||||
@staticmethod
|
||||
def _get_request_format(optional_params: object) -> OCRRequestFormat:
|
||||
if not isinstance(optional_params, dict):
|
||||
return "litellm"
|
||||
request_format: Final = optional_params.get(OCR_REQUEST_FORMAT_PARAM)
|
||||
if request_format is None:
|
||||
return "litellm"
|
||||
return parse_ocr_request_format(request_format)
|
||||
|
||||
def _transform_completed_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
request_format: OCRRequestFormat,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Transform a completed Azure Document Intelligence analyze operation
|
||||
into the Mistral OCR response shape, preserving Azure-native
|
||||
`analyzeResult` fields (`content`, `tables`, `keyValuePairs`) as
|
||||
top-level response fields.
|
||||
|
||||
When `request_format` is "native", the untouched Azure operation
|
||||
payload is attached to the response's hidden params so the proxy can
|
||||
return it verbatim while cost tracking still reads `usage_info`.
|
||||
"""
|
||||
raw_operation: Final[Mapping[str, object]] = raw_response.json()
|
||||
operation: Final = AzureDocumentIntelligenceOperation.model_validate(raw_operation)
|
||||
|
||||
verbose_logger.debug("Azure Document Intelligence response status: %s", operation.status)
|
||||
|
||||
if operation.status != "succeeded":
|
||||
raise ValueError(f"Azure Document Intelligence analysis failed with status: {operation.status}")
|
||||
|
||||
analyze_result: Final = (
|
||||
operation.analyzeResult if operation.analyzeResult is not None else AzureDocumentIntelligenceAnalyzeResult()
|
||||
)
|
||||
mistral_pages: Final = [self._transform_azure_page(azure_page) for azure_page in analyze_result.pages]
|
||||
usage_info: Final = OCRUsageInfo(pages_processed=len(mistral_pages), doc_size_bytes=None)
|
||||
|
||||
response: Final = OCRResponse(
|
||||
pages=mistral_pages,
|
||||
model=model,
|
||||
usage_info=usage_info,
|
||||
object="ocr",
|
||||
content=analyze_result.content,
|
||||
tables=analyze_result.tables,
|
||||
keyValuePairs=analyze_result.keyValuePairs,
|
||||
)
|
||||
|
||||
if request_format == "native":
|
||||
response.set_provider_native_response(raw_operation)
|
||||
|
||||
return response
|
||||
|
||||
def transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Transform Azure Document Intelligence response to Mistral OCR format.
|
||||
|
||||
Handles async operation polling: If response is 202 Accepted, polls Operation-Location
|
||||
until analysis completes.
|
||||
|
||||
Azure DI response (after polling):
|
||||
{
|
||||
"status": "succeeded",
|
||||
"analyzeResult": {
|
||||
"content": "Full document text...",
|
||||
"pages": [
|
||||
{
|
||||
"pageNumber": 1,
|
||||
"width": 8.5,
|
||||
"height": 11,
|
||||
"unit": "inch",
|
||||
"lines": [{"content": "text", "boundingBox": [...]}]
|
||||
}
|
||||
],
|
||||
"tables": [...],
|
||||
"keyValuePairs": [...]
|
||||
}
|
||||
}
|
||||
|
||||
Mistral OCR format (with Azure-native fields preserved):
|
||||
{
|
||||
"pages": [
|
||||
{
|
||||
"index": 0,
|
||||
"markdown": "extracted text",
|
||||
"dimensions": {"width": 816, "height": 1056, "dpi": 96}
|
||||
}
|
||||
],
|
||||
"model": "azure_ai/doc-intelligence/prebuilt-layout",
|
||||
"usage_info": {"pages_processed": 1},
|
||||
"object": "ocr",
|
||||
"content": "Full document text...",
|
||||
"tables": [...],
|
||||
"keyValuePairs": [...]
|
||||
}
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
raw_response: Raw HTTP response from Azure DI (may be 202 Accepted)
|
||||
logging_obj: Logging object
|
||||
|
||||
Returns:
|
||||
OCRResponse in Mistral format
|
||||
"""
|
||||
request_format: Final = self._get_request_format(kwargs.get("optional_params"))
|
||||
|
||||
if raw_response.status_code != 202:
|
||||
return self._transform_completed_response(
|
||||
model=model, raw_response=raw_response, request_format=request_format
|
||||
)
|
||||
|
||||
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation...")
|
||||
operation_url, poll_headers = self._get_polling_target(raw_response)
|
||||
completed_response: Final = self._poll_operation_sync(
|
||||
operation_url=operation_url,
|
||||
headers=poll_headers,
|
||||
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
|
||||
)
|
||||
return self._transform_completed_response(
|
||||
model=model, raw_response=completed_response, request_format=request_format
|
||||
)
|
||||
|
||||
async def async_transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Async transform Azure Document Intelligence response to Mistral OCR format.
|
||||
|
||||
Handles async operation polling: If response is 202 Accepted, polls Operation-Location
|
||||
until analysis completes using async polling.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
raw_response: Raw HTTP response from Azure DI (may be 202 Accepted)
|
||||
logging_obj: Logging object
|
||||
|
||||
Returns:
|
||||
OCRResponse in Mistral format
|
||||
"""
|
||||
request_format: Final = self._get_request_format(kwargs.get("optional_params"))
|
||||
|
||||
if raw_response.status_code != 202:
|
||||
return self._transform_completed_response(
|
||||
model=model, raw_response=raw_response, request_format=request_format
|
||||
)
|
||||
|
||||
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation (async)...")
|
||||
operation_url, poll_headers = self._get_polling_target(raw_response)
|
||||
completed_response: Final = await self._poll_operation_async(
|
||||
operation_url=operation_url,
|
||||
headers=poll_headers,
|
||||
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
|
||||
)
|
||||
return self._transform_completed_response(
|
||||
model=model, raw_response=completed_response, request_format=request_format
|
||||
)
|
||||
|
|
@ -1,263 +0,0 @@
|
|||
"""
|
||||
Azure AI OCR transformation implementation.
|
||||
"""
|
||||
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
||||
async_convert_url_to_base64,
|
||||
convert_url_to_base64,
|
||||
)
|
||||
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
|
||||
from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRRequestData
|
||||
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
AZURE_AI_OCR_API_KEY_ENV_VAR: Final = "AZURE_AI_API_KEY"
|
||||
|
||||
|
||||
class AzureAIOCRConfig(MistralOCRConfig):
|
||||
"""
|
||||
Azure AI OCR transformation configuration.
|
||||
|
||||
Azure AI uses Mistral's OCR API but with a different endpoint format.
|
||||
Inherits transformation logic from MistralOCRConfig since they use the same format.
|
||||
|
||||
Reference: Azure AI Foundry OCR documentation
|
||||
|
||||
Important: Azure AI only supports base64 data URIs (data:image/..., data:application/pdf;base64,...).
|
||||
Regular URLs are not supported.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def get_api_key_env_var(self) -> str | None:
|
||||
return AZURE_AI_OCR_API_KEY_ENV_VAR
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
**kwargs,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate environment and return headers for Azure AI OCR.
|
||||
|
||||
Authenticates with AZURE_AI_API_KEY, or with an Entra ID / OAuth token when no key is set.
|
||||
"""
|
||||
# Get API key from environment if not provided
|
||||
if api_key is None:
|
||||
api_key = get_secret_str(AZURE_AI_OCR_API_KEY_ENV_VAR)
|
||||
|
||||
# Validate API base is provided
|
||||
if api_base is None:
|
||||
api_base = get_secret_str("AZURE_AI_API_BASE")
|
||||
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter"
|
||||
)
|
||||
|
||||
headers = {
|
||||
**get_azure_ai_auth_headers(api_key=api_key, litellm_params=litellm_params),
|
||||
"Content-Type": "application/json",
|
||||
**headers,
|
||||
}
|
||||
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict | None = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Get complete URL for Azure AI OCR endpoint.
|
||||
|
||||
Azure AI endpoint format: https://<api_base>/providers/mistral/azure/ocr
|
||||
|
||||
Args:
|
||||
api_base: Azure AI API base URL
|
||||
model: Model name (not used in URL construction)
|
||||
optional_params: Optional parameters
|
||||
|
||||
Returns: Complete URL for Azure AI OCR endpoint
|
||||
"""
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter"
|
||||
)
|
||||
|
||||
# Ensure no trailing slash
|
||||
api_base = api_base.rstrip("/")
|
||||
|
||||
# Azure AI OCR endpoint format
|
||||
return f"{api_base}/providers/mistral/azure/ocr"
|
||||
|
||||
def _convert_url_to_data_uri_sync(self, url: str) -> str:
|
||||
"""
|
||||
Synchronously convert a URL to a base64 data URI.
|
||||
|
||||
Azure AI OCR doesn't have internet access, so we need to fetch URLs
|
||||
and convert them to base64 data URIs.
|
||||
|
||||
Args:
|
||||
url: The URL to convert
|
||||
|
||||
Returns:
|
||||
Base64 data URI string
|
||||
"""
|
||||
verbose_logger.debug("Azure AI OCR: Converting URL to base64 data URI (sync): %s", url)
|
||||
|
||||
# Fetch and convert to base64 data URI
|
||||
# convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
|
||||
data_uri: Final = convert_url_to_base64(url=url)
|
||||
|
||||
verbose_logger.debug("Azure AI OCR: Converted URL to data URI (length: %s)", len(data_uri))
|
||||
|
||||
return data_uri
|
||||
|
||||
async def _convert_url_to_data_uri_async(self, url: str) -> str:
|
||||
"""
|
||||
Asynchronously convert a URL to a base64 data URI.
|
||||
|
||||
Azure AI OCR doesn't have internet access, so we need to fetch URLs
|
||||
and convert them to base64 data URIs.
|
||||
|
||||
Args:
|
||||
url: The URL to convert
|
||||
|
||||
Returns:
|
||||
Base64 data URI string
|
||||
"""
|
||||
verbose_logger.debug("Azure AI OCR: Converting URL to base64 data URI (async): %s", url)
|
||||
|
||||
# Fetch and convert to base64 data URI asynchronously
|
||||
# async_convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
|
||||
data_uri: Final = await async_convert_url_to_base64(url=url)
|
||||
|
||||
verbose_logger.debug("Azure AI OCR: Converted URL to data URI (length: %s)", len(data_uri))
|
||||
|
||||
return data_uri
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Transform OCR request for Azure AI, converting URLs to base64 data URIs (sync).
|
||||
|
||||
Azure AI OCR doesn't have internet access, so we automatically fetch
|
||||
any URLs and convert them to base64 data URIs synchronously.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
document: Document dict from user
|
||||
optional_params: Already mapped optional parameters
|
||||
headers: Request headers
|
||||
**kwargs: Additional arguments
|
||||
|
||||
Returns:
|
||||
OCRRequestData with JSON data
|
||||
"""
|
||||
verbose_logger.debug("Azure AI OCR transform_ocr_request (sync) - model: %s", model)
|
||||
|
||||
if not isinstance(document, dict):
|
||||
raise ValueError(f"Expected document dict, got {type(document)}")
|
||||
|
||||
# Check if we need to convert URL to base64
|
||||
doc_type: Final = document.get("type")
|
||||
transformed_document: Final = document.copy()
|
||||
|
||||
if doc_type == "document_url":
|
||||
document_url: Final = document.get("document_url", "")
|
||||
# If it's not already a data URI, convert it
|
||||
if document_url and not document_url.startswith("data:"):
|
||||
verbose_logger.debug("Azure AI OCR: Converting document URL to base64 data URI (sync)")
|
||||
data_uri = self._convert_url_to_data_uri_sync(url=document_url)
|
||||
transformed_document["document_url"] = data_uri
|
||||
elif doc_type == "image_url":
|
||||
image_url: Final = document.get("image_url", "")
|
||||
# If it's not already a data URI, convert it
|
||||
if image_url and not image_url.startswith("data:"):
|
||||
verbose_logger.debug("Azure AI OCR: Converting image URL to base64 data URI (sync)")
|
||||
data_uri = self._convert_url_to_data_uri_sync(url=image_url)
|
||||
transformed_document["image_url"] = data_uri
|
||||
|
||||
# Call parent's transform to build the request
|
||||
return super().transform_ocr_request(
|
||||
model=model,
|
||||
document=transformed_document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def async_transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Transform OCR request for Azure AI, converting URLs to base64 data URIs (async).
|
||||
|
||||
Azure AI OCR doesn't have internet access, so we automatically fetch
|
||||
any URLs and convert them to base64 data URIs asynchronously.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
document: Document dict from user
|
||||
optional_params: Already mapped optional parameters
|
||||
headers: Request headers
|
||||
**kwargs: Additional arguments
|
||||
|
||||
Returns:
|
||||
OCRRequestData with JSON data
|
||||
"""
|
||||
verbose_logger.debug("Azure AI OCR async_transform_ocr_request - model: %s", model)
|
||||
|
||||
if not isinstance(document, dict):
|
||||
raise ValueError(f"Expected document dict, got {type(document)}")
|
||||
|
||||
# Check if we need to convert URL to base64
|
||||
doc_type: Final = document.get("type")
|
||||
transformed_document: Final = document.copy()
|
||||
|
||||
if doc_type == "document_url":
|
||||
document_url: Final = document.get("document_url", "")
|
||||
# If it's not already a data URI, convert it
|
||||
if document_url and not document_url.startswith("data:"):
|
||||
verbose_logger.debug("Azure AI OCR: Converting document URL to base64 data URI (async)")
|
||||
data_uri = await self._convert_url_to_data_uri_async(url=document_url)
|
||||
transformed_document["document_url"] = data_uri
|
||||
elif doc_type == "image_url":
|
||||
image_url: Final = document.get("image_url", "")
|
||||
# If it's not already a data URI, convert it
|
||||
if image_url and not image_url.startswith("data:"):
|
||||
verbose_logger.debug("Azure AI OCR: Converting image URL to base64 data URI (async)")
|
||||
data_uri = await self._convert_url_to_data_uri_async(url=image_url)
|
||||
transformed_document["image_url"] = data_uri
|
||||
|
||||
# Call parent's transform to build the request
|
||||
return super().transform_ocr_request(
|
||||
model=model,
|
||||
document=transformed_document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
|
|
@ -13,7 +13,7 @@ from litellm.llms.azure_ai.common_utils import (
|
|||
api_key_header_for_base,
|
||||
get_azure_ai_auth_headers,
|
||||
)
|
||||
from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.llms.base_llm.passthrough.transformation import (
|
||||
BasePassthroughConfig,
|
||||
RelayShape,
|
||||
|
|
@ -22,6 +22,7 @@ from litellm.llms.base_llm.passthrough.transformation import (
|
|||
relayed_body,
|
||||
strip_leading_model_segment,
|
||||
)
|
||||
from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR_PASSTHROUGH_RESPONSE, NativeOcrPassthroughResponse
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.rerank import RerankResponse
|
||||
from litellm.types.utils import CallTypes, ImageResponse, StandardPassThroughResponseObject
|
||||
|
|
@ -30,7 +31,6 @@ if TYPE_CHECKING:
|
|||
from httpx import URL, Response
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
|
||||
from litellm.llms.base_llm.passthrough.transformation import LoggedRelayResponse
|
||||
|
||||
|
||||
|
|
@ -92,9 +92,9 @@ FOUNDRY_RELAY_SHAPES: Final = (
|
|||
|
||||
|
||||
class AzureAIPassthroughConfig(AzureFoundryModelInfo, BasePassthroughConfig):
|
||||
def __init__(self, ocr_config_for: Callable[[str], BaseOCRConfig | None] = get_azure_ai_ocr_config) -> None:
|
||||
def __init__(self, passthrough_ocr: NativeOcrPassthroughResponse | None = None) -> None:
|
||||
super().__init__()
|
||||
self.ocr_config_for: Final = ocr_config_for
|
||||
self._passthrough_ocr: Final = passthrough_ocr
|
||||
|
||||
def is_streaming_request(self, endpoint: str, request_data: Mapping[str, object]) -> bool:
|
||||
return bool(request_data.get("stream"))
|
||||
|
|
@ -168,28 +168,22 @@ class AzureAIPassthroughConfig(AzureFoundryModelInfo, BasePassthroughConfig):
|
|||
def logged_ocr_response(
|
||||
self, model: str, httpx_response: Response, logging_obj: Logging, endpoint: str
|
||||
) -> OCRResponse | None:
|
||||
ocr_config: Final = self.ocr_config_for(model)
|
||||
if ocr_config is None or httpx_response.status_code != 200:
|
||||
return None
|
||||
relayed_url: Final = httpx_response.request.url
|
||||
relayed_origin: Final = str(relayed_url.copy_with(path="/", query=None, fragment=None)).rstrip("/")
|
||||
ocr_url: Final = httpx.URL(
|
||||
ocr_config.get_complete_url(
|
||||
api_base=relayed_origin,
|
||||
model=model,
|
||||
optional_params={}, # mutable-ok: BaseOCRConfig wants a dict
|
||||
)
|
||||
passthrough_ocr: Final = (
|
||||
self._passthrough_ocr if self._passthrough_ocr is not None else NATIVE_OCR_PASSTHROUGH_RESPONSE.load()
|
||||
)
|
||||
if passthrough_ocr is None or httpx_response.status_code != 200:
|
||||
return None
|
||||
known_prefixes: Final = (model, model_group_from(logging_obj.litellm_params))
|
||||
native_endpoint: Final = strip_leading_model_segment(endpoint, known_prefixes)
|
||||
if f"/{native_endpoint.strip('/')}" != ocr_url.path:
|
||||
return None
|
||||
try:
|
||||
ocr_response: Final = ocr_config.transform_ocr_response(
|
||||
model=model, raw_response=httpx_response, logging_obj=logging_obj
|
||||
result: Final = passthrough_ocr(model, native_endpoint, httpx_response.content)
|
||||
ocr_response: Final = OCRResponse.model_validate(result) if result is not None else None
|
||||
except (ValueError, RuntimeError) as error:
|
||||
verbose_logger.warning(
|
||||
"azure_ai passthrough: OCR body from %s is not costable: %s", httpx_response.request.url, error
|
||||
)
|
||||
except (ValueError, AttributeError) as error:
|
||||
verbose_logger.warning("azure_ai passthrough: OCR body from %s is not costable: %s", ocr_url, error)
|
||||
return None
|
||||
if ocr_response is None:
|
||||
return None
|
||||
logging_obj.call_type = CallTypes.aocr.value # rebind-ok: routes cost calculation to the per-page OCR path
|
||||
return ocr_response
|
||||
|
|
|
|||
|
|
@ -1,23 +1,19 @@
|
|||
"""Base OCR transformation module."""
|
||||
|
||||
from .transformation import (
|
||||
BaseOCRConfig,
|
||||
DocumentType,
|
||||
OCRPage,
|
||||
OCRPageDimensions,
|
||||
OCRPageImage,
|
||||
OCRRequestData,
|
||||
OCRResponse,
|
||||
OCRUsageInfo,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BaseOCRConfig",
|
||||
"DocumentType",
|
||||
"OCRPage",
|
||||
"OCRPageDimensions",
|
||||
"OCRPageImage",
|
||||
"OCRRequestData",
|
||||
"OCRResponse",
|
||||
"OCRUsageInfo",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,27 +1,16 @@
|
|||
"""
|
||||
Base OCR transformation configuration.
|
||||
Base OCR types shared by the Rust OCR route and Python consumers.
|
||||
"""
|
||||
|
||||
import builtins
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
import httpx
|
||||
from pydantic import PrivateAttr
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.llms.base import LiteLLMPydanticObjectBase
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
# DocumentType for OCR - providers always receive a dict with
|
||||
# type="document_url" or type="image_url" (str values only).
|
||||
# File-type inputs are preprocessed to this format in litellm/ocr/main.py.
|
||||
DocumentType = dict[str, str]
|
||||
DocumentType = Mapping[str, object]
|
||||
|
||||
OCRRequestFormat = Literal["litellm", "native"]
|
||||
|
||||
|
|
@ -33,8 +22,6 @@ OCR_REQUEST_FORMAT_HEADER: Final = "x-req-format"
|
|||
|
||||
PROVIDER_NATIVE_RESPONSE_KEY: Final = "provider_native_response"
|
||||
|
||||
HEALTH_CHECK_PDF_DATA_URI: Final = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y="
|
||||
|
||||
|
||||
def parse_ocr_request_format(value: object) -> OCRRequestFormat:
|
||||
if value == "litellm":
|
||||
|
|
@ -102,7 +89,6 @@ class OCRResponse(LiteLLMPydanticObjectBase):
|
|||
|
||||
model_config = {"extra": "allow"}
|
||||
|
||||
# Define private attributes using PrivateAttr
|
||||
_hidden_params: dict = PrivateAttr(default_factory=dict)
|
||||
|
||||
def set_provider_native_response(self, native_response: Mapping[str, builtins.object]) -> None:
|
||||
|
|
@ -113,203 +99,3 @@ class OCRResponse(LiteLLMPydanticObjectBase):
|
|||
"""The provider's own response payload, when `req_format=native` was requested."""
|
||||
native_response: Final = self._hidden_params.get(PROVIDER_NATIVE_RESPONSE_KEY)
|
||||
return native_response if isinstance(native_response, dict) else None
|
||||
|
||||
|
||||
class OCRRequestData(LiteLLMPydanticObjectBase):
|
||||
"""OCR request data structure."""
|
||||
|
||||
data: dict | bytes | None = None
|
||||
files: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class BaseOCRConfig:
|
||||
"""
|
||||
Base configuration for OCR transformations.
|
||||
Handles provider-agnostic OCR operations.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list:
|
||||
"""
|
||||
Get supported OCR parameters for this provider.
|
||||
Override this method in provider-specific implementations.
|
||||
"""
|
||||
return []
|
||||
|
||||
def get_api_key_env_var(self) -> str | None:
|
||||
"""
|
||||
Return the provider-specific API key environment variable name, if any.
|
||||
"""
|
||||
return None
|
||||
|
||||
def resolve_connection_params(
|
||||
self,
|
||||
*,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
dynamic_api_key: str | None,
|
||||
dynamic_api_base: str | None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
return dynamic_api_key or api_key, dynamic_api_base or api_base
|
||||
|
||||
def get_health_check_document(self) -> DocumentType:
|
||||
return { # mutable-ok: litellm.aocr rejects any document that is not a dict
|
||||
"type": "document_url",
|
||||
"document_url": HEALTH_CHECK_PDF_DATA_URI,
|
||||
}
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
) -> dict:
|
||||
"""Map OCR parameters to provider-specific parameters."""
|
||||
return optional_params
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
**kwargs,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate environment and return headers.
|
||||
Override in provider-specific implementations.
|
||||
"""
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict | None = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Get complete URL for OCR endpoint.
|
||||
Override in provider-specific implementations.
|
||||
"""
|
||||
raise NotImplementedError("get_complete_url must be implemented by provider")
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Transform OCR request to provider-specific format.
|
||||
Override in provider-specific implementations.
|
||||
|
||||
Note: By the time this method is called, any file-type documents have already
|
||||
been converted to document_url/image_url format with base64 data URIs by
|
||||
the preprocessing in litellm/ocr/main.py.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
document: Document to process - always a dict with type="document_url" or type="image_url"
|
||||
optional_params: Optional parameters for the request
|
||||
headers: Request headers
|
||||
|
||||
Returns:
|
||||
OCRRequestData with data and files fields
|
||||
"""
|
||||
raise NotImplementedError("transform_ocr_request must be implemented by provider")
|
||||
|
||||
async def async_transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Async transform OCR request to provider-specific format.
|
||||
Optional method - providers can override if they need async transformations
|
||||
(e.g., Azure AI for URL-to-base64 conversion).
|
||||
|
||||
Default implementation falls back to sync transform_ocr_request.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
document: Document to process (Mistral format dict, or file path, bytes, etc.)
|
||||
optional_params: Optional parameters for the request
|
||||
headers: Request headers
|
||||
|
||||
Returns:
|
||||
OCRRequestData with data and files fields
|
||||
"""
|
||||
# Default implementation: call sync version
|
||||
return self.transform_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Transform provider-specific OCR response to standard format.
|
||||
Override in provider-specific implementations.
|
||||
"""
|
||||
raise NotImplementedError("transform_ocr_response must be implemented by provider")
|
||||
|
||||
async def async_transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Async transform provider-specific OCR response to standard format.
|
||||
Optional method - providers can override if they need async transformations
|
||||
(e.g., Azure Document Intelligence for async operation polling).
|
||||
|
||||
Default implementation falls back to sync transform_ocr_response.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
raw_response: Raw HTTP response
|
||||
logging_obj: Logging object
|
||||
|
||||
Returns:
|
||||
OCRResponse in standard format
|
||||
"""
|
||||
# Default implementation: call sync version
|
||||
return self.transform_ocr_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
logging_obj=logging_obj,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: dict,
|
||||
) -> Exception:
|
||||
"""Get appropriate error class for the provider."""
|
||||
return BaseLLMException(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import base64
|
||||
from typing import Final, NoReturn
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -16,14 +16,6 @@ from litellm.rust_bridge.transcription.native import (
|
|||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
|
||||
def _no_python_implementation() -> NoReturn:
|
||||
raise NotImplementedError("Bedrock audio transcription is implemented in Rust only")
|
||||
|
||||
|
||||
async def _no_async_python_implementation() -> NoReturn:
|
||||
_no_python_implementation()
|
||||
|
||||
|
||||
class BedrockAudioTranscriptionRustDispatch:
|
||||
@staticmethod
|
||||
def _audio_payload(audio_file: FileTypes) -> dict[str, object]:
|
||||
|
|
@ -77,7 +69,7 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
RouteContext(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model),
|
||||
binding=NATIVE_TRANSCRIPTION,
|
||||
native=native,
|
||||
python=_no_python_implementation,
|
||||
python=runtime.NO_PYTHON,
|
||||
)
|
||||
|
||||
async def async_audio_transcriptions(
|
||||
|
|
@ -110,5 +102,5 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
RouteContext(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model),
|
||||
binding=NATIVE_ATRANSCRIPTION,
|
||||
native=native,
|
||||
python=_no_async_python_implementation,
|
||||
python=runtime.NO_PYTHON,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -7,7 +7,8 @@ import json
|
|||
import re
|
||||
import time
|
||||
import types
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from itertools import chain
|
||||
from typing import TYPE_CHECKING, Final, Literal, cast, overload
|
||||
|
||||
import httpx
|
||||
|
|
@ -34,6 +35,12 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
_bedrock_tools_pt,
|
||||
make_valid_bedrock_tool_name,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import (
|
||||
CONVERTED_SYSTEM_NOTE,
|
||||
is_system_message,
|
||||
message_field,
|
||||
parts_of,
|
||||
)
|
||||
from litellm.llms.anthropic.chat.transformation import (
|
||||
DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING,
|
||||
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
|
||||
|
|
@ -55,9 +62,11 @@ from litellm.types.llms.openai import (
|
|||
ChatCompletionAnnotation,
|
||||
ChatCompletionAssistantMessage,
|
||||
ChatCompletionAssistantToolCall,
|
||||
ChatCompletionCachedContent,
|
||||
ChatCompletionRedactedThinkingBlock,
|
||||
ChatCompletionResponseMessage,
|
||||
ChatCompletionSystemMessage,
|
||||
ChatCompletionTextObject,
|
||||
ChatCompletionThinkingBlock,
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolCallFunctionChunk,
|
||||
|
|
@ -1343,30 +1352,157 @@ class AmazonConverseConfig(BaseConfig):
|
|||
cache_point["ttl"] = ttl
|
||||
return cache_point
|
||||
|
||||
@staticmethod
|
||||
def _assistant_has_tool_calls(message: object) -> bool:
|
||||
return message_field(message, "role") == "assistant" and bool(message_field(message, "tool_calls"))
|
||||
|
||||
@staticmethod
|
||||
def _opens_with_tool_result(message: object) -> bool:
|
||||
"""Whether the message starts a tool-result turn on Converse.
|
||||
|
||||
``_bedrock_converse_messages_pt`` builds ``toolResult`` blocks from ``tool``
|
||||
messages only, so a ``function`` message never opens one."""
|
||||
role: Final = message_field(message, "role")
|
||||
if role == "tool":
|
||||
return True
|
||||
if role != "user":
|
||||
return False
|
||||
first_part: Final = next(iter(parts_of(message_field(message, "content"))), None)
|
||||
return message_field(first_part, "type") == "tool_result"
|
||||
|
||||
def _system_run_before(self, messages: Sequence[AllMessageValues], index: int) -> Sequence[AllMessageValues]:
|
||||
start: Final = next(
|
||||
(j + 1 for j in range(index - 1, -1, -1) if not is_system_message(messages[j])),
|
||||
0,
|
||||
)
|
||||
return messages[start:index]
|
||||
|
||||
def _system_run_end(self, messages: Sequence[AllMessageValues], index: int) -> int:
|
||||
return next(
|
||||
(j for j in range(index, len(messages)) if not is_system_message(messages[j])),
|
||||
len(messages),
|
||||
)
|
||||
|
||||
def _reordered_around_tool_results(
|
||||
self, messages: Sequence[AllMessageValues], index: int
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
"""Move a system run wedged between an assistant tool-call turn and its
|
||||
tool-result turn(s) to after the tool results.
|
||||
|
||||
A converted system entry becomes a user turn, and a user turn between
|
||||
a tool call and its result would split them. Everything else stays in
|
||||
place so the cached prefix stays byte-identical."""
|
||||
message: Final = messages[index]
|
||||
if self._opens_with_tool_result(message):
|
||||
if index + 1 < len(messages) and self._opens_with_tool_result(messages[index + 1]):
|
||||
return (message,)
|
||||
tool_run_start: Final = next(
|
||||
(j + 1 for j in range(index, -1, -1) if not self._opens_with_tool_result(messages[j])),
|
||||
0,
|
||||
)
|
||||
run: Final = self._system_run_before(messages, tool_run_start)
|
||||
prev_idx: Final = tool_run_start - len(run) - 1
|
||||
if run and prev_idx >= 0 and self._assistant_has_tool_calls(messages[prev_idx]):
|
||||
return (message, *run)
|
||||
return (message,)
|
||||
if not is_system_message(message):
|
||||
return (message,)
|
||||
run_start: Final = next(
|
||||
(j + 1 for j in range(index - 1, -1, -1) if not is_system_message(messages[j])),
|
||||
0,
|
||||
)
|
||||
run_end: Final = self._system_run_end(messages, index)
|
||||
follower: Final = messages[run_end] if run_end < len(messages) else None
|
||||
if (
|
||||
follower is not None
|
||||
and self._opens_with_tool_result(follower)
|
||||
and run_start > 0
|
||||
and self._assistant_has_tool_calls(messages[run_start - 1])
|
||||
):
|
||||
return ()
|
||||
return (message,)
|
||||
|
||||
def _system_role_message_as_user(self, message: ChatCompletionSystemMessage) -> ChatCompletionUserMessage | None:
|
||||
"""Convert a mid-conversation system entry to a user turn, in place.
|
||||
|
||||
The Converse API only accepts user/assistant roles in ``messages``,
|
||||
so keeping the role is not an option. Hoisting it to the top-level
|
||||
``system`` block would mutate the system prefix and collapse implicit
|
||||
prompt caching; converting in place keeps everything before the entry
|
||||
byte-identical. An entry that carries no text becomes ``None``."""
|
||||
text_blocks: Final = self._converted_text_blocks(message)
|
||||
if not text_blocks:
|
||||
return None
|
||||
note: Final = ChatCompletionTextObject(type="text", text=CONVERTED_SYSTEM_NOTE)
|
||||
body: Final = [ # mutable-ok: _bedrock_converse_messages_pt narrows content with isinstance(list)
|
||||
note,
|
||||
*text_blocks,
|
||||
]
|
||||
return ChatCompletionUserMessage(role="user", content=body)
|
||||
|
||||
def _converted_or_kept(self, message: AllMessageValues) -> AllMessageValues | None:
|
||||
if not is_system_message(message):
|
||||
return message
|
||||
return self._system_role_message_as_user(
|
||||
cast(ChatCompletionSystemMessage, message) # cast-ok: the role is checked on the line above
|
||||
)
|
||||
|
||||
def _converted_text_blocks(self, message: ChatCompletionSystemMessage) -> tuple[ChatCompletionTextObject, ...]:
|
||||
content: Final = message["content"]
|
||||
if isinstance(content, str):
|
||||
return (self._converted_text_block(content, message.get("cache_control")),) if content else ()
|
||||
parts: Final[Sequence[object]] = content or ()
|
||||
return tuple(
|
||||
self._converted_text_block(part["text"], part.get("cache_control"))
|
||||
for part in map(self._text_part, parts)
|
||||
if part is not None
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _text_part(part: object) -> ChatCompletionTextObject | None:
|
||||
if not isinstance(part, dict) or part.get("type") != "text" or not part.get("text"):
|
||||
return None
|
||||
return cast(ChatCompletionTextObject, part) # cast-ok: the shape is checked on the line above
|
||||
|
||||
@staticmethod
|
||||
def _converted_text_block(text: str, cache_control: ChatCompletionCachedContent | None) -> ChatCompletionTextObject:
|
||||
if cache_control is None:
|
||||
return ChatCompletionTextObject(type="text", text=text)
|
||||
return ChatCompletionTextObject(type="text", text=text, cache_control=cache_control)
|
||||
|
||||
def _transform_system_message(
|
||||
self, messages: list[AllMessageValues], model: str | None = None
|
||||
) -> tuple[list[AllMessageValues], list[SystemContentBlock]]:
|
||||
system_prompt_indices: Final = []
|
||||
leading_count: Final = next(
|
||||
(i for i, m in enumerate(messages) if not is_system_message(m)),
|
||||
len(messages),
|
||||
)
|
||||
hoisted: Final = messages[:leading_count]
|
||||
remaining: Final = messages[leading_count:]
|
||||
system_content_blocks: Final[list[SystemContentBlock]] = []
|
||||
for idx, message in enumerate(messages):
|
||||
if message["role"] == "system":
|
||||
system_prompt_indices.append(idx)
|
||||
if isinstance(message["content"], str) and message["content"]:
|
||||
system_content_blocks.append(SystemContentBlock(text=message["content"]))
|
||||
cache_block = self.get_cache_point_block(message, block_type="system", model=model)
|
||||
if cache_block:
|
||||
system_content_blocks.append(cache_block)
|
||||
elif isinstance(message["content"], list):
|
||||
for m in message["content"]:
|
||||
if m.get("type") == "text" and m.get("text"):
|
||||
system_content_blocks.append(SystemContentBlock(text=m["text"]))
|
||||
cache_block = self.get_cache_point_block(m, block_type="system", model=model)
|
||||
if cache_block:
|
||||
system_content_blocks.append(cache_block)
|
||||
if len(system_prompt_indices) > 0:
|
||||
for idx in reversed(system_prompt_indices):
|
||||
messages.pop(idx)
|
||||
return messages, system_content_blocks
|
||||
for message in hoisted:
|
||||
if message["role"] != "system":
|
||||
continue
|
||||
if isinstance(message["content"], str) and message["content"]:
|
||||
system_content_blocks.append(SystemContentBlock(text=message["content"]))
|
||||
cache_block = self.get_cache_point_block(message, block_type="system", model=model)
|
||||
if cache_block:
|
||||
system_content_blocks.append(cache_block)
|
||||
elif isinstance(message["content"], list):
|
||||
for m in message["content"]:
|
||||
if m.get("type") == "text" and m.get("text"):
|
||||
system_content_blocks.append(SystemContentBlock(text=m["text"]))
|
||||
cache_block = self.get_cache_point_block(m, block_type="system", model=model)
|
||||
if cache_block:
|
||||
system_content_blocks.append(cache_block)
|
||||
reordered: Final = tuple(
|
||||
chain.from_iterable(
|
||||
self._reordered_around_tool_results(remaining, index) for index in range(len(remaining))
|
||||
)
|
||||
)
|
||||
converted: Final = tuple(self._converted_or_kept(message) for message in reordered)
|
||||
kept: Final = [message for message in converted if message is not None] # mutable-ok: converse pt takes a list
|
||||
return kept, system_content_blocks
|
||||
|
||||
def _transform_inference_params(self, inference_params: dict) -> InferenceConfig:
|
||||
if "top_k" in inference_params:
|
||||
|
|
|
|||
|
|
@ -58,8 +58,8 @@ from litellm.types.llms.openai import (
|
|||
OpenAIFileObject,
|
||||
PathLike,
|
||||
)
|
||||
from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums
|
||||
from litellm.utils import get_llm_provider
|
||||
from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums, all_litellm_params
|
||||
from litellm.utils import get_llm_provider, get_optional_params
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import (
|
||||
|
|
@ -88,6 +88,14 @@ def _frozen_mapping(items: Iterable[tuple[str, object]]) -> Mapping[str, object]
|
|||
return MappingProxyType(dict(items))
|
||||
|
||||
|
||||
_LITELLM_PARAMS_THE_MAPPER_TAKES: Final = frozenset({"allowed_openai_params"})
|
||||
_MAPPED_PARAMS_THE_REQUEST_HANDLER_STRIPS: Final = frozenset({"json_mode"})
|
||||
|
||||
|
||||
def _invoke_route_model(model: str) -> str:
|
||||
return f"invoke/{_strip_llm_routing_prefix(model).removeprefix('invoke/')}"
|
||||
|
||||
|
||||
def _strip_llm_routing_prefix(model: str) -> str:
|
||||
try:
|
||||
stripped_model, _, _, _ = get_llm_provider(model=model, custom_llm_provider=None)
|
||||
|
|
@ -891,16 +899,24 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
)
|
||||
|
||||
config: Final = AmazonAnthropicClaudeConfig()
|
||||
mapped_params = config.map_openai_params(
|
||||
non_default_params={},
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=False,
|
||||
mapped_params = get_optional_params(
|
||||
model=_invoke_route_model(model),
|
||||
custom_llm_provider="bedrock",
|
||||
messages=messages,
|
||||
**MappingProxyType(
|
||||
{
|
||||
k: v
|
||||
for k, v in optional_params.items()
|
||||
if k not in all_litellm_params or k in _LITELLM_PARAMS_THE_MAPPER_TAKES
|
||||
}
|
||||
),
|
||||
)
|
||||
return config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=mapped_params,
|
||||
optional_params={
|
||||
k: v for k, v in mapped_params.items() if k not in _MAPPED_PARAMS_THE_REQUEST_HANDLER_STRIPS
|
||||
},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,3 +0,0 @@
|
|||
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
|
||||
|
||||
__all__ = ("CohereParseConfig",)
|
||||
|
|
@ -1,298 +0,0 @@
|
|||
"""Cohere Parse (`POST /v2/parse`) exposed through LiteLLM's OCR interface."""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.exceptions import BadRequestError, UnsupportedParamsError
|
||||
from litellm.llms.base_llm.ocr.transformation import (
|
||||
OCR_REQUEST_FORMAT_PARAM,
|
||||
BaseOCRConfig,
|
||||
DocumentType,
|
||||
OCRPage,
|
||||
OCRPageImage,
|
||||
OCRRequestData,
|
||||
OCRRequestFormat,
|
||||
OCRResponse,
|
||||
OCRUsageInfo,
|
||||
parse_ocr_request_format,
|
||||
)
|
||||
from litellm.llms.cohere.common_utils import CohereError
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
COHERE_API_KEY_ENV_VAR: Final = "COHERE_API_KEY"
|
||||
COHERE_PARSE_API_BASE: Final = "https://api.cohere.com"
|
||||
COHERE_PARSE_PATH: Final = "/v2/parse"
|
||||
COHERE_PARSE_OUTPUT_FORMAT_PARAM: Final = "output_format"
|
||||
COHERE_PARSE_OUTPUT_FORMATS: Final = ("markdown", "blocks")
|
||||
COHERE_PARSE_DEFAULT_OUTPUT_FORMAT: Final = "markdown"
|
||||
COHERE_PARSE_SUPPORTED_PARAMS: Final = (COHERE_PARSE_OUTPUT_FORMAT_PARAM, OCR_REQUEST_FORMAT_PARAM)
|
||||
COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI: Final = (
|
||||
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
|
||||
)
|
||||
COHERE_PARSE_IMAGE_ONLY_MESSAGE: Final = (
|
||||
"Cohere Parse only accepts `image_url` documents (an image URL or a base64 image data URI); "
|
||||
"`document_url` and PDF inputs are not supported."
|
||||
)
|
||||
|
||||
_NATIVE_RESPONSE_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_BOUNDING_BOX_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
class _CohereParseDocument(TypedDict):
|
||||
type: ReadOnly[Literal["image_url"]]
|
||||
image_url: ReadOnly[str]
|
||||
|
||||
|
||||
class _CohereParseRequestBody(TypedDict):
|
||||
model: ReadOnly[str]
|
||||
document: ReadOnly[_CohereParseDocument]
|
||||
output_format: ReadOnly[str]
|
||||
|
||||
|
||||
class _MarkdownPage(TypedDict):
|
||||
index: ReadOnly[int]
|
||||
markdown: ReadOnly[str]
|
||||
images: ReadOnly[Sequence[OCRPageImage] | None]
|
||||
|
||||
|
||||
class _BlocksPage(_MarkdownPage):
|
||||
blocks: ReadOnly[Sequence[Mapping[str, object]]]
|
||||
|
||||
|
||||
class _CohereParseMarkdown(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
content: str = ""
|
||||
images: Sequence[Mapping[str, object]] | None = None
|
||||
|
||||
|
||||
class _CohereParsePage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
index: int | None = None
|
||||
markdown: _CohereParseMarkdown | None = None
|
||||
blocks: Sequence[Mapping[str, object]] | None = None
|
||||
|
||||
|
||||
class _CohereParseBilledUnits(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
pages: int | None = None
|
||||
|
||||
|
||||
class _CohereParseMeta(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
billed_units: _CohereParseBilledUnits | None = None
|
||||
|
||||
|
||||
class _CohereParseResponse(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
pages: Sequence[_CohereParsePage] = ()
|
||||
meta: _CohereParseMeta | None = None
|
||||
|
||||
|
||||
def _requested_format(optional_params: Mapping[str, object] | None) -> OCRRequestFormat:
|
||||
if optional_params is None:
|
||||
return "litellm"
|
||||
return "native" if optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" else "litellm"
|
||||
|
||||
|
||||
def _page_image(image: Mapping[str, object]) -> OCRPageImage:
|
||||
bounding_box: Final = image.get("bounding_box")
|
||||
if not isinstance(bounding_box, Mapping):
|
||||
return OCRPageImage.model_validate(image)
|
||||
bbox: Final = _BOUNDING_BOX_ADAPTER.validate_python(bounding_box)
|
||||
return OCRPageImage.model_validate(MappingProxyType({**image, "bbox": bbox}))
|
||||
|
||||
|
||||
def _normalize_page(page: _CohereParsePage, position: int) -> OCRPage:
|
||||
markdown: Final = page.markdown
|
||||
images: Final = tuple(_page_image(image) for image in markdown.images) if markdown and markdown.images else None
|
||||
normalized: Final[_MarkdownPage] = {
|
||||
"index": page.index if page.index is not None else position,
|
||||
"markdown": markdown.content if markdown else "",
|
||||
"images": images,
|
||||
}
|
||||
if page.blocks is None:
|
||||
return OCRPage.model_validate(normalized)
|
||||
with_blocks: Final[_BlocksPage] = {**normalized, "blocks": page.blocks}
|
||||
return OCRPage.model_validate(with_blocks)
|
||||
|
||||
|
||||
def _billed_pages(parsed: _CohereParseResponse) -> int | None:
|
||||
if parsed.meta is None or parsed.meta.billed_units is None:
|
||||
return None
|
||||
return parsed.meta.billed_units.pages
|
||||
|
||||
|
||||
class CohereParseConfig(BaseOCRConfig):
|
||||
"""Cohere Parse, an image-only document understanding endpoint returning markdown or blocks."""
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list[str]: # mutable-ok: BaseOCRConfig signature
|
||||
return list(COHERE_PARSE_SUPPORTED_PARAMS) # mutable-ok: BaseOCRConfig signature
|
||||
|
||||
def get_api_key_env_var(self) -> str | None:
|
||||
return COHERE_API_KEY_ENV_VAR
|
||||
|
||||
def get_health_check_document(self) -> DocumentType:
|
||||
return { # mutable-ok: litellm.aocr rejects any document that is not a dict
|
||||
"type": "image_url",
|
||||
"image_url": COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI,
|
||||
}
|
||||
|
||||
def _llm_provider(self) -> str:
|
||||
return "cohere"
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
optional_params: Mapping[str, object],
|
||||
model: str,
|
||||
) -> dict[str, object]: # mutable-ok: BaseOCRConfig signature
|
||||
output_format: Final = non_default_params.get(COHERE_PARSE_OUTPUT_FORMAT_PARAM)
|
||||
if output_format is not None and output_format not in COHERE_PARSE_OUTPUT_FORMATS:
|
||||
raise UnsupportedParamsError(
|
||||
message=(
|
||||
f"Invalid `{COHERE_PARSE_OUTPUT_FORMAT_PARAM}`: {output_format!r}. "
|
||||
f"Expected one of {', '.join(COHERE_PARSE_OUTPUT_FORMATS)}."
|
||||
),
|
||||
model=model,
|
||||
llm_provider=self._llm_provider(),
|
||||
)
|
||||
requested_format: Final = non_default_params.get(OCR_REQUEST_FORMAT_PARAM)
|
||||
request_format: Final = parse_ocr_request_format(requested_format) if requested_format is not None else None
|
||||
overrides: Final = tuple(
|
||||
(key, value)
|
||||
for key, value in (
|
||||
(COHERE_PARSE_OUTPUT_FORMAT_PARAM, output_format),
|
||||
(OCR_REQUEST_FORMAT_PARAM, request_format),
|
||||
)
|
||||
if value is not None
|
||||
)
|
||||
return {**optional_params, **dict(overrides)} # mutable-ok: BaseOCRConfig signature
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.validate_environment signature
|
||||
) -> dict[str, str]: # mutable-ok: BaseOCRConfig signature
|
||||
resolved_key: Final = api_key or get_secret_str(COHERE_API_KEY_ENV_VAR)
|
||||
if resolved_key is None:
|
||||
raise ValueError(
|
||||
f"Missing {COHERE_API_KEY_ENV_VAR} - set it in the environment or pass api_key to "
|
||||
"litellm.ocr()/litellm.aocr()"
|
||||
)
|
||||
return { # mutable-ok: BaseOCRConfig signature
|
||||
"Authorization": f"Bearer {resolved_key}",
|
||||
"Content-Type": "application/json",
|
||||
**headers,
|
||||
}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.get_complete_url signature
|
||||
) -> str:
|
||||
url: Final = httpx.URL(api_base or COHERE_PARSE_API_BASE)
|
||||
path: Final = url.path.rstrip("/")
|
||||
if path.endswith(COHERE_PARSE_PATH):
|
||||
return str(url.copy_with(path=path))
|
||||
if path.endswith("/v2"):
|
||||
return str(url.copy_with(path=f"{path}/parse"))
|
||||
return str(url.copy_with(path=f"{path}{COHERE_PARSE_PATH}"))
|
||||
|
||||
def _image_url(self, document: DocumentType, model: str) -> str:
|
||||
image_url: Final = document.get("image_url", "")
|
||||
if document.get("type") != "image_url" or not image_url or image_url.startswith("data:application/pdf"):
|
||||
raise BadRequestError(
|
||||
message=COHERE_PARSE_IMAGE_ONLY_MESSAGE,
|
||||
model=model,
|
||||
llm_provider=self._llm_provider(),
|
||||
)
|
||||
return image_url
|
||||
|
||||
def _resolve_image_url_sync(self, image_url: str) -> str:
|
||||
return image_url
|
||||
|
||||
async def _resolve_image_url_async(self, image_url: str) -> str:
|
||||
return image_url
|
||||
|
||||
def _build_request(self, model: str, image_url: str, optional_params: Mapping[str, object]) -> OCRRequestData:
|
||||
body: Final[_CohereParseRequestBody] = {
|
||||
"model": model,
|
||||
"document": {"type": "image_url", "image_url": image_url},
|
||||
"output_format": str(
|
||||
optional_params.get(COHERE_PARSE_OUTPUT_FORMAT_PARAM, COHERE_PARSE_DEFAULT_OUTPUT_FORMAT)
|
||||
),
|
||||
}
|
||||
return OCRRequestData(data=dict(body), files=None) # mutable-ok: OCRRequestData.data is a dict
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: Mapping[str, object],
|
||||
headers: Mapping[str, str],
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.transform_ocr_request signature
|
||||
) -> OCRRequestData:
|
||||
image_url: Final = self._resolve_image_url_sync(self._image_url(document, model))
|
||||
return self._build_request(model=model, image_url=image_url, optional_params=optional_params)
|
||||
|
||||
async def async_transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: Mapping[str, object],
|
||||
headers: Mapping[str, str],
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.async_transform_ocr_request signature
|
||||
) -> OCRRequestData:
|
||||
image_url: Final = await self._resolve_image_url_async(self._image_url(document, model))
|
||||
return self._build_request(model=model, image_url=image_url, optional_params=optional_params)
|
||||
|
||||
def transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
optional_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.transform_ocr_response signature
|
||||
) -> OCRResponse:
|
||||
native: Final = _NATIVE_RESPONSE_ADAPTER.validate_python(raw_response.json())
|
||||
parsed: Final = _CohereParseResponse.model_validate(native)
|
||||
pages: Final = [ # mutable-ok: OCRResponse.pages is a list
|
||||
_normalize_page(page, position) for position, page in enumerate(parsed.pages)
|
||||
]
|
||||
billed_pages: Final = _billed_pages(parsed)
|
||||
response: Final = OCRResponse(
|
||||
pages=pages,
|
||||
model=model,
|
||||
usage_info=OCRUsageInfo(pages_processed=billed_pages if billed_pages is not None else len(pages)),
|
||||
)
|
||||
if _requested_format(optional_params) == "native":
|
||||
response.set_provider_native_response(native)
|
||||
return response
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: Mapping[str, str],
|
||||
) -> Exception:
|
||||
return CohereError(status_code=status_code, message=error_message)
|
||||
|
|
@ -81,7 +81,6 @@ from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
|||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse
|
||||
from litellm.llms.base_llm.realtime.http_transformation import BaseRealtimeHTTPConfig
|
||||
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
|
||||
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
|
||||
|
|
@ -1636,320 +1635,6 @@ class BaseLLMHTTPHandler:
|
|||
api_key=api_key,
|
||||
)
|
||||
|
||||
def _prepare_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, str],
|
||||
optional_params: dict,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
headers: dict[str, object] | None,
|
||||
provider_config: BaseOCRConfig,
|
||||
litellm_params: dict,
|
||||
) -> tuple[dict[str, object], str, dict[str, object], None]:
|
||||
"""
|
||||
Shared logic for preparing OCR requests.
|
||||
Returns: (headers, complete_url, data, files)
|
||||
"""
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRRequestData
|
||||
|
||||
headers = provider_config.validate_environment(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers or {},
|
||||
model=model,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
complete_url: Final = provider_config.get_complete_url(
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
# Transform the request to get data and files
|
||||
transformed_result: Final = provider_config.transform_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params={key: value for key, value in optional_params.items() if key != OCR_REQUEST_FORMAT_PARAM},
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
if not isinstance(transformed_result, OCRRequestData):
|
||||
raise ValueError(f"Provider {provider_config.__class__.__name__} must return OCRRequestData")
|
||||
|
||||
# Data is always a dict for Mistral OCR format
|
||||
if not isinstance(transformed_result.data, dict):
|
||||
raise ValueError(f"Expected dict data for OCR request, got {type(transformed_result.data)}")
|
||||
|
||||
data: Final = transformed_result.data
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": complete_url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
return headers, complete_url, data, None
|
||||
|
||||
async def _async_prepare_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, str],
|
||||
optional_params: dict,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
headers: dict[str, object] | None,
|
||||
provider_config: BaseOCRConfig,
|
||||
litellm_params: dict,
|
||||
) -> tuple[dict[str, object], str, dict[str, object], None]:
|
||||
"""
|
||||
Async version of _prepare_ocr_request for providers that need async transforms.
|
||||
Returns: (headers, complete_url, data, files)
|
||||
"""
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRRequestData
|
||||
|
||||
headers = provider_config.validate_environment(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers or {},
|
||||
model=model,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
complete_url: Final = provider_config.get_complete_url(
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
# Use async transform (providers can override this method if they need async operations)
|
||||
transformed_result: Final = await provider_config.async_transform_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params={key: value for key, value in optional_params.items() if key != OCR_REQUEST_FORMAT_PARAM},
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
if not isinstance(transformed_result, OCRRequestData):
|
||||
raise ValueError(f"Provider {provider_config.__class__.__name__} must return OCRRequestData")
|
||||
|
||||
# Data is always a dict for Mistral OCR format
|
||||
if not isinstance(transformed_result.data, dict):
|
||||
raise ValueError(f"Expected dict data for OCR request, got {type(transformed_result.data)}")
|
||||
|
||||
data: Final = transformed_result.data
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": complete_url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
return headers, complete_url, data, None
|
||||
|
||||
def _transform_ocr_response(
|
||||
self,
|
||||
provider_config: BaseOCRConfig,
|
||||
model: str,
|
||||
response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
optional_params: Mapping[str, object],
|
||||
) -> OCRResponse:
|
||||
"""Shared logic for transforming OCR responses."""
|
||||
normalized: Final = provider_config.transform_ocr_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
return self._finalize_ocr_response(normalized, response, optional_params)
|
||||
|
||||
@staticmethod
|
||||
def _finalize_ocr_response(
|
||||
normalized: OCRResponse,
|
||||
response: httpx.Response,
|
||||
optional_params: Mapping[str, object],
|
||||
) -> OCRResponse:
|
||||
if (
|
||||
optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"
|
||||
and normalized.get_provider_native_response() is None
|
||||
):
|
||||
normalized.set_provider_native_response(response.json())
|
||||
return normalized
|
||||
|
||||
def ocr(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, str],
|
||||
optional_params: dict,
|
||||
timeout: float | httpx.Timeout,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str,
|
||||
client: HTTPHandler | AsyncHTTPHandler | None = None,
|
||||
aocr: bool = False,
|
||||
headers: dict[str, object] | None = None,
|
||||
provider_config: BaseOCRConfig | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
"""
|
||||
Sync OCR handler.
|
||||
"""
|
||||
if provider_config is None:
|
||||
raise ValueError(f"No provider config found for model: {model} and provider: {custom_llm_provider}")
|
||||
|
||||
if litellm_params is None:
|
||||
litellm_params = {}
|
||||
|
||||
if aocr is True:
|
||||
return self.async_ocr(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
client=client,
|
||||
headers=headers,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
# Prepare the request
|
||||
headers, complete_url, data, files = self._prepare_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
client = _get_httpx_client()
|
||||
|
||||
try:
|
||||
# Make the POST request with JSON data (Mistral format)
|
||||
response: Final = client.post(
|
||||
url=complete_url,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
logging_obj.post_call(
|
||||
api_key=api_key,
|
||||
original_response=response.text,
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
|
||||
return self._transform_ocr_response(
|
||||
provider_config=provider_config,
|
||||
model=model,
|
||||
response=response,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
async def async_ocr(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, str],
|
||||
optional_params: dict,
|
||||
timeout: float | httpx.Timeout,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str,
|
||||
client: HTTPHandler | AsyncHTTPHandler | None = None,
|
||||
headers: dict[str, object] | None = None,
|
||||
provider_config: BaseOCRConfig | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Async OCR handler.
|
||||
"""
|
||||
if provider_config is None:
|
||||
raise ValueError(f"No provider config found for model: {model} and provider: {custom_llm_provider}")
|
||||
|
||||
if litellm_params is None:
|
||||
litellm_params = {}
|
||||
|
||||
# Prepare the request using async prepare method
|
||||
headers, complete_url, data, files = await self._async_prepare_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
try:
|
||||
# Make the async POST request with JSON data (Mistral format)
|
||||
response: Final = await async_httpx_client.post(
|
||||
url=complete_url,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
logging_obj.post_call(
|
||||
api_key=api_key,
|
||||
original_response=response.text,
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
|
||||
# Use async response transform for async operations
|
||||
normalized: Final = await provider_config.async_transform_ocr_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
return self._finalize_ocr_response(normalized, response, optional_params)
|
||||
|
||||
def search(
|
||||
self,
|
||||
query: str | list[str],
|
||||
|
|
@ -6211,7 +5896,6 @@ class BaseLLMHTTPHandler:
|
|||
BaseGoogleGenAIGenerateContentConfig,
|
||||
BaseAnthropicMessagesConfig,
|
||||
BaseBatchesConfig,
|
||||
BaseOCRConfig,
|
||||
BaseVideoConfig,
|
||||
BaseSearchConfig,
|
||||
BaseTextToSpeechConfig,
|
||||
|
|
@ -6254,12 +5938,6 @@ class BaseLLMHTTPHandler:
|
|||
status_code=status_code,
|
||||
headers=error_headers,
|
||||
)
|
||||
if (
|
||||
isinstance(provider_config, BaseOCRConfig)
|
||||
and isinstance(provider_error, BaseLLMException)
|
||||
and isinstance(error_response, httpx.Response)
|
||||
):
|
||||
provider_error.response = error_response
|
||||
if not isinstance(received_status_code, int):
|
||||
provider_error.status_code_is_synthesized = True
|
||||
raise provider_error
|
||||
|
|
|
|||
|
|
@ -142,14 +142,35 @@ def _resolution_key(resolution: object) -> str | None:
|
|||
return str(resolution)
|
||||
|
||||
|
||||
def _resolution_cost_per_image(entry: Mapping[str, object] | None, resolution: object) -> float | None:
|
||||
resolution_key: Final = _resolution_key(resolution)
|
||||
if entry is None or resolution_key is None:
|
||||
return None
|
||||
cost: Final = entry.get(f"output_cost_per_image_{resolution_key}")
|
||||
return float(cost) if isinstance(cost, (int, float)) else None
|
||||
|
||||
|
||||
def _requested_image_count(request_body: Mapping[str, object]) -> int:
|
||||
num_images: Final = request_body.get("num_images")
|
||||
return num_images if type(num_images) is int and num_images > 0 else 1
|
||||
|
||||
|
||||
def _passthrough_cost_per_image(entry: Mapping[str, object], request_body: Mapping[str, object]) -> float | None:
|
||||
resolution_cost: Final = _resolution_cost_per_image(entry, request_body.get("resolution"))
|
||||
if resolution_cost is not None:
|
||||
return resolution_cost
|
||||
cost: Final = entry.get("output_cost_per_image")
|
||||
return float(cost) if isinstance(cost, (int, float)) else None
|
||||
|
||||
|
||||
def fal_ai_passthrough_cost(model: str, request_body: Mapping[str, object]) -> float | None:
|
||||
entry: Final = _entry(f"{litellm.LlmProviders.FAL_AI.value}/{model}")
|
||||
if entry is None:
|
||||
return None
|
||||
resolution: Final = _resolution_key(request_body.get("resolution"))
|
||||
keyed_cost: Final = entry.get(f"output_cost_per_image_{resolution}") if resolution is not None else None
|
||||
cost: Final = keyed_cost if isinstance(keyed_cost, (int, float)) else entry.get("output_cost_per_image")
|
||||
return float(cost) if isinstance(cost, (int, float)) else None
|
||||
cost_per_image: Final = _passthrough_cost_per_image(entry, request_body)
|
||||
if cost_per_image is None:
|
||||
return None
|
||||
return cost_per_image * _requested_image_count(request_body)
|
||||
|
||||
|
||||
def cost_calculator(
|
||||
|
|
@ -172,6 +193,11 @@ def cost_calculator(
|
|||
if deployment_cost_per_image is not None:
|
||||
return deployment_cost_per_image * len(images)
|
||||
params: Final[Mapping[str, object]] = optional_params or MappingProxyType({})
|
||||
resolution_cost_per_image: Final = _resolution_cost_per_image(
|
||||
_entry(f"{litellm.LlmProviders.FAL_AI.value}/{normalized_model}"), params.get("resolution")
|
||||
)
|
||||
if resolution_cost_per_image is not None:
|
||||
return resolution_cost_per_image * len(images)
|
||||
keyed_costs: Final = tuple(
|
||||
_keyed_cost_per_image(
|
||||
model=normalized_model,
|
||||
|
|
|
|||
|
|
@ -8,12 +8,17 @@ from .transformation import FalAIBaseConfig
|
|||
|
||||
class FalAINanoBananaConfig(FalAIBaseConfig):
|
||||
"""
|
||||
Configuration for Fal AI's Nano Banana / Gemini 2.5 Flash Image models.
|
||||
Configuration for Fal AI's Nano Banana family (Gemini Flash / Pro Image models).
|
||||
|
||||
Serves the imagen4 deprecation migration path. The same underlying model is
|
||||
exposed under two endpoints that share an identical schema:
|
||||
Serves the imagen4 deprecation migration path. Every endpoint shares the same
|
||||
request schema, so one config covers all of them:
|
||||
- fal-ai/nano-banana
|
||||
- fal-ai/gemini-25-flash-image
|
||||
- fal-ai/nano-banana-2
|
||||
- fal-ai/nano-banana-pro
|
||||
|
||||
Provider-specific params such as ``resolution`` ("0.5K", "1K", "2K", "4K") are
|
||||
forwarded as-is and drive the per-resolution price in the cost map.
|
||||
|
||||
Documentation: https://fal.ai/models/fal-ai/nano-banana
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,244 +0,0 @@
|
|||
"""
|
||||
Mistral OCR transformation implementation.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.ocr.transformation import (
|
||||
BaseOCRConfig,
|
||||
DocumentType,
|
||||
OCRRequestData,
|
||||
OCRResponse,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
MISTRAL_OCR_API_KEY_ENV_VAR: Final = "MISTRAL_API_KEY"
|
||||
|
||||
|
||||
class MistralOCRConfig(BaseOCRConfig):
|
||||
"""
|
||||
Mistral OCR transformation configuration.
|
||||
|
||||
Reference: https://docs.mistral.ai/api/#tag/ocr
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list:
|
||||
"""
|
||||
Get supported OCR parameters for Mistral OCR.
|
||||
|
||||
Mistral OCR supports:
|
||||
- pages: List of page numbers to process
|
||||
- include_image_base64: Whether to include base64 encoded images
|
||||
- image_limit: Maximum number of images to return
|
||||
- image_min_size: Minimum size of images to include
|
||||
- bbox_annotation_format: Format for bounding box annotations
|
||||
- document_annotation_format: Format for document annotations
|
||||
- document_annotation_prompt: Prompt for document annotation extraction
|
||||
- extract_header: Whether to extract document header
|
||||
- extract_footer: Whether to extract document footer
|
||||
- table_format: Table output format ("markdown" or "html")
|
||||
- confidence_scores_granularity: Confidence score level ("word" or "page")
|
||||
- include_blocks: Whether to return paragraph-level bounding boxes and typed content blocks (OCR 4)
|
||||
- id: Request identifier
|
||||
"""
|
||||
return [
|
||||
"pages",
|
||||
"include_image_base64",
|
||||
"image_limit",
|
||||
"image_min_size",
|
||||
"bbox_annotation_format",
|
||||
"document_annotation_format",
|
||||
"document_annotation_prompt",
|
||||
"extract_header",
|
||||
"extract_footer",
|
||||
"table_format",
|
||||
"confidence_scores_granularity",
|
||||
"include_blocks",
|
||||
"id",
|
||||
]
|
||||
|
||||
def get_api_key_env_var(self) -> str | None:
|
||||
return MISTRAL_OCR_API_KEY_ENV_VAR
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OCR parameters to Mistral-specific format.
|
||||
|
||||
Mistral accepts these parameters directly, so no transformation needed.
|
||||
Just filter out unsupported params.
|
||||
"""
|
||||
supported_params: Final = self.get_supported_ocr_params(model=model)
|
||||
|
||||
# Only include params that are in the supported list
|
||||
mapped_params: Final = {}
|
||||
for param, value in non_default_params.items():
|
||||
if param in supported_params:
|
||||
mapped_params[param] = value
|
||||
|
||||
return mapped_params
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
**kwargs,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate environment and return headers for Mistral OCR.
|
||||
"""
|
||||
# Get API key from environment if not provided
|
||||
if api_key is None:
|
||||
api_key = get_secret_str(MISTRAL_OCR_API_KEY_ENV_VAR)
|
||||
|
||||
if api_key is None:
|
||||
raise ValueError(
|
||||
"Missing Mistral API Key - A call is being made to Mistral but no key is set either in the environment variables or via params"
|
||||
)
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
**headers,
|
||||
}
|
||||
|
||||
# Don't set Content-Type for multipart/form-data - httpx will handle it
|
||||
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict | None = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Get complete URL for Mistral OCR endpoint.
|
||||
|
||||
Returns: https://api.mistral.ai/v1/ocr
|
||||
"""
|
||||
if api_base is None:
|
||||
api_base = "https://api.mistral.ai/v1"
|
||||
|
||||
# Ensure no trailing slash
|
||||
api_base = api_base.rstrip("/")
|
||||
|
||||
# Remove /v1 if it's already in the base to avoid duplication
|
||||
if api_base.endswith("/v1"):
|
||||
return f"{api_base}/ocr"
|
||||
|
||||
return f"{api_base}/v1/ocr"
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
"""
|
||||
Transform OCR request to Mistral-specific format.
|
||||
|
||||
Mistral OCR API accepts:
|
||||
{
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": {
|
||||
"type": "document_url",
|
||||
"document_url": "<https-url or data-uri>"
|
||||
},
|
||||
"pages": [0], # optional
|
||||
"include_image_base64": false, # optional
|
||||
...
|
||||
}
|
||||
|
||||
Args:
|
||||
model: Model name (e.g., "mistral-ocr-latest")
|
||||
document: Document dict from user (Mistral format) - already validated in main.py
|
||||
optional_params: Already mapped optional parameters
|
||||
headers: Request headers
|
||||
|
||||
Returns:
|
||||
OCRRequestData with JSON data
|
||||
"""
|
||||
verbose_logger.debug("Mistral OCR transform_ocr_request - model: %s", model)
|
||||
|
||||
# Document parameter is the Mistral-format dict from the user
|
||||
# Just pass it through as-is to the Mistral API
|
||||
if not isinstance(document, dict):
|
||||
raise ValueError(f"Expected document dict, got {type(document)}")
|
||||
|
||||
# Build request data - use document dict directly
|
||||
data: Final = {
|
||||
"model": model,
|
||||
"document": document, # Pass through the Mistral-format document dict
|
||||
}
|
||||
|
||||
# Add all optional parameters from the already-mapped optional_params
|
||||
data.update(optional_params)
|
||||
|
||||
# No multipart files - using JSON
|
||||
return OCRRequestData(data=data, files=None)
|
||||
|
||||
def transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Return Mistral OCR response in native format.
|
||||
|
||||
Mistral OCR is the standard format for LiteLLM OCR responses.
|
||||
No transformation needed - return native response.
|
||||
|
||||
Mistral OCR returns:
|
||||
{
|
||||
"pages": [
|
||||
{
|
||||
"index": 0,
|
||||
"markdown": "extracted text content",
|
||||
"images": [...],
|
||||
"dimensions": {...}
|
||||
},
|
||||
...
|
||||
],
|
||||
"model": "mistral-ocr-2505-completion",
|
||||
"document_annotation": null,
|
||||
"usage_info": {...}
|
||||
}
|
||||
"""
|
||||
try:
|
||||
response_json: Final = raw_response.json()
|
||||
|
||||
verbose_logger.debug("Mistral OCR response keys: %s", response_json.keys())
|
||||
|
||||
# Return native Mistral format - no transformation
|
||||
return OCRResponse(
|
||||
pages=response_json.get("pages", []),
|
||||
model=response_json.get("model", model),
|
||||
document_annotation=response_json.get("document_annotation"),
|
||||
usage_info=response_json.get("usage_info"),
|
||||
object="ocr",
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.error("Error parsing Mistral OCR response: %s", e)
|
||||
raise e
|
||||
68
litellm/llms/nadir/chat/transformation.py
Normal file
68
litellm/llms/nadir/chat/transformation.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
import math
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.core_helpers import set_response_cost_in_hidden_params
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
_SUPPORTED_OPENAI_PARAMS: Final = (
|
||||
"extra_headers",
|
||||
"frequency_penalty",
|
||||
"max_retries",
|
||||
"max_tokens",
|
||||
"presence_penalty",
|
||||
"response_format",
|
||||
"stream",
|
||||
"temperature",
|
||||
"top_p",
|
||||
)
|
||||
|
||||
|
||||
def _reported_cost_usd(raw_response: httpx.Response) -> float | None:
|
||||
try:
|
||||
cost: Final = raw_response.json()["nadir_metadata"]["cost"]["total_cost_usd"]
|
||||
except (ValueError, KeyError, TypeError):
|
||||
return None
|
||||
if isinstance(cost, bool) or not isinstance(cost, (int, float)):
|
||||
return None
|
||||
if not math.isfinite(cost) or cost < 0:
|
||||
return None
|
||||
return float(cost)
|
||||
|
||||
|
||||
class NadirConfig(OpenAIGPTConfig):
|
||||
def get_supported_openai_params(self, model: str) -> list: # mutable-ok: return type fixed by the base interface
|
||||
return list(_SUPPORTED_OPENAI_PARAMS) # mutable-ok: the base interface returns a list
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: object,
|
||||
request_data: dict, # mutable-ok: signature fixed by the base interface
|
||||
messages: list[AllMessageValues], # mutable-ok: signature fixed by the base interface
|
||||
optional_params: dict, # mutable-ok: signature fixed by the base interface
|
||||
litellm_params: dict, # mutable-ok: signature fixed by the base interface
|
||||
encoding: object,
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
transformed: Final = super().transform_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
model_response=model_response,
|
||||
logging_obj=logging_obj,
|
||||
request_data=request_data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
encoding=encoding,
|
||||
api_key=api_key,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
set_response_cost_in_hidden_params(transformed, _reported_cost_usd(raw_response))
|
||||
return transformed
|
||||
133
litellm/llms/openai/organization_costs.py
Normal file
133
litellm/llms/openai/organization_costs.py
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
"""OpenAI's organization costs endpoint: the USD the organization was billed per UTC day, read with an admin key."""
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime, timedelta, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
from litellm.constants import (
|
||||
OPENAI_ORGANIZATION_COSTS_PAGE_LIMIT,
|
||||
OPENAI_ORGANIZATION_COSTS_URL,
|
||||
PROVIDER_BILLING_TIMEOUT_SECONDS,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
OPENAI_ADMIN_KEY_ENV_VAR: Final = "OPENAI_ADMIN_KEY"
|
||||
|
||||
BillingHttpGet: TypeAlias = Callable[
|
||||
[str, Mapping[str, object], Mapping[str, str]], # mutable-ok: Callable parameter list is type syntax
|
||||
Awaitable[httpx.Response],
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OpenAICostsRequestFailed:
|
||||
detail: str
|
||||
|
||||
|
||||
class _OpenAICostAmount(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
value: float
|
||||
currency: Literal["usd"]
|
||||
|
||||
|
||||
class _OpenAICostResult(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
amount: _OpenAICostAmount
|
||||
|
||||
|
||||
class _OpenAICostBucket(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
start_time: int
|
||||
results: tuple[_OpenAICostResult, ...] = ()
|
||||
|
||||
|
||||
class _OpenAICostsPage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
data: tuple[_OpenAICostBucket, ...]
|
||||
has_more: bool = False
|
||||
next_page: str | None = None
|
||||
|
||||
|
||||
async def provider_billing_get(url: str, params: Mapping[str, object], headers: Mapping[str, str]) -> httpx.Response:
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.ProviderBilling)
|
||||
return await client.get(
|
||||
url,
|
||||
params=dict(params), # mutable-ok: AsyncHTTPHandler.get takes dict params
|
||||
headers=dict(headers), # mutable-ok: AsyncHTTPHandler.get takes dict headers
|
||||
timeout=PROVIDER_BILLING_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
|
||||
def _utc_midnight(day: date) -> int:
|
||||
return int(datetime(day.year, day.month, day.day, tzinfo=timezone.utc).timestamp())
|
||||
|
||||
|
||||
def _bucket_day(bucket: _OpenAICostBucket) -> str:
|
||||
return datetime.fromtimestamp(bucket.start_time, tz=timezone.utc).date().isoformat()
|
||||
|
||||
|
||||
async def fetch_openai_daily_costs(
|
||||
start_date: date,
|
||||
end_date: date,
|
||||
*,
|
||||
admin_key: str,
|
||||
project_ids: Sequence[str] = (),
|
||||
http_get: BillingHttpGet = provider_billing_get,
|
||||
) -> Mapping[str, float] | OpenAICostsRequestFailed:
|
||||
"""USD billed by OpenAI per UTC day (ISO date) over the closed range, following pagination to the end."""
|
||||
scope: Final = (("project_ids[]", tuple(project_ids)),) if project_ids else ()
|
||||
window: Final[Mapping[str, object]] = MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in (
|
||||
("start_time", _utc_midnight(start_date)),
|
||||
("end_time", _utc_midnight(end_date + timedelta(days=1))),
|
||||
("bucket_width", "1d"),
|
||||
("limit", OPENAI_ORGANIZATION_COSTS_PAGE_LIMIT),
|
||||
*scope,
|
||||
)
|
||||
}
|
||||
)
|
||||
headers: Final[Mapping[str, str]] = MappingProxyType({"Authorization": f"Bearer {admin_key}"})
|
||||
|
||||
async def fetch_from(page: str | None) -> tuple[_OpenAICostBucket, ...] | OpenAICostsRequestFailed:
|
||||
params: Final[Mapping[str, object]] = MappingProxyType(
|
||||
{key: value for key, value in (*window.items(), ("page", page)) if value is not None}
|
||||
)
|
||||
try:
|
||||
response: Final = await http_get(OPENAI_ORGANIZATION_COSTS_URL, params, headers)
|
||||
except httpx.HTTPError as exc:
|
||||
return OpenAICostsRequestFailed(f"request failed: {exc}")
|
||||
if response.status_code != 200:
|
||||
return OpenAICostsRequestFailed(f"HTTP {response.status_code}: {response.text[:300]}")
|
||||
try:
|
||||
parsed: Final = _OpenAICostsPage.model_validate(response.json())
|
||||
except (ValueError, ValidationError) as exc:
|
||||
return OpenAICostsRequestFailed(f"unexpected response shape: {exc}")
|
||||
if not parsed.has_more or parsed.next_page is None:
|
||||
return parsed.data
|
||||
rest: Final = await fetch_from(parsed.next_page)
|
||||
return rest if isinstance(rest, OpenAICostsRequestFailed) else parsed.data + rest
|
||||
|
||||
buckets: Final = await fetch_from(None)
|
||||
if isinstance(buckets, OpenAICostsRequestFailed):
|
||||
return buckets
|
||||
days: Final = frozenset(_bucket_day(bucket) for bucket in buckets)
|
||||
return MappingProxyType(
|
||||
{
|
||||
day: sum(
|
||||
result.amount.value for bucket in buckets if _bucket_day(bucket) == day for result in bucket.results
|
||||
)
|
||||
for day in days
|
||||
}
|
||||
)
|
||||
|
|
@ -1,149 +0,0 @@
|
|||
import base64
|
||||
import binascii
|
||||
from collections import defaultdict
|
||||
from typing import TYPE_CHECKING, Any, Final, NoReturn
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.constants import request_timeout
|
||||
|
||||
REDUCTO_API_BASE: Final = "https://platform.reducto.ai"
|
||||
REDUCTO_ID_PREFIX: Final = "reducto://"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRPage
|
||||
|
||||
|
||||
def _normalize_api_base(api_base: str | None) -> str:
|
||||
return (api_base or REDUCTO_API_BASE).rstrip("/")
|
||||
|
||||
|
||||
def _raise_bad_request(message: str, model: str) -> NoReturn:
|
||||
import litellm
|
||||
|
||||
raise litellm.BadRequestError(
|
||||
message=message,
|
||||
model=model,
|
||||
llm_provider="reducto",
|
||||
)
|
||||
|
||||
|
||||
def extract_file_id_or_bytes(
|
||||
source_url: str,
|
||||
model: str,
|
||||
) -> tuple[str | None, bytes | None, str | None]:
|
||||
if source_url.startswith(REDUCTO_ID_PREFIX):
|
||||
return source_url, None, None
|
||||
|
||||
if source_url.startswith("http://") or source_url.startswith("https://"):
|
||||
_raise_bad_request(
|
||||
"Reducto requires type='file' (auto-uploaded) or a reducto:// id. Plain http(s) URLs are not supported; upload the file first.",
|
||||
model=model,
|
||||
)
|
||||
|
||||
if not source_url.startswith("data:"):
|
||||
_raise_bad_request(
|
||||
"Reducto requires a reducto:// id or a base64 data URI after OCR preprocessing.",
|
||||
model=model,
|
||||
)
|
||||
|
||||
try:
|
||||
header, encoded = source_url.split(",", 1)
|
||||
except ValueError:
|
||||
_raise_bad_request("Invalid Reducto data URI provided.", model=model)
|
||||
|
||||
if ";base64" not in header:
|
||||
_raise_bad_request("Reducto only supports base64-encoded data URIs.", model=model)
|
||||
|
||||
mime: Final = header.removeprefix("data:").split(";")[0] or "application/octet-stream"
|
||||
try:
|
||||
raw_bytes: Final = base64.b64decode(encoded, validate=True)
|
||||
except (binascii.Error, ValueError):
|
||||
_raise_bad_request("Invalid Reducto base64 payload provided.", model=model)
|
||||
|
||||
return None, raw_bytes, mime
|
||||
|
||||
|
||||
def _extract_file_id_from_upload_response(response: httpx.Response) -> str:
|
||||
try:
|
||||
payload: Final = response.json()
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"Reducto /upload returned a non-JSON 200 response: {response.text}") from exc
|
||||
file_id: Final = (payload or {}).get("file_id") if isinstance(payload, dict) else None
|
||||
if not isinstance(file_id, str) or not file_id:
|
||||
raise ValueError(f"Reducto /upload returned 200 without a file_id; got payload={payload}")
|
||||
return file_id
|
||||
|
||||
|
||||
def upload_bytes_sync(
|
||||
raw_bytes: bytes,
|
||||
mime: str | None,
|
||||
api_key: str,
|
||||
api_base: str | None,
|
||||
) -> str:
|
||||
import litellm
|
||||
|
||||
response: Final = litellm.module_level_client.post(
|
||||
url="{}{}".format(_normalize_api_base(api_base), "/upload"),
|
||||
headers={"Authorization": f"Bearer {api_key}"},
|
||||
files={"file": ("document", raw_bytes, mime or "application/octet-stream")},
|
||||
timeout=request_timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return _extract_file_id_from_upload_response(response)
|
||||
|
||||
|
||||
async def upload_bytes_async(
|
||||
raw_bytes: bytes,
|
||||
mime: str | None,
|
||||
api_key: str,
|
||||
api_base: str | None,
|
||||
) -> str:
|
||||
import litellm
|
||||
|
||||
response: Final = await litellm.module_level_aclient.post(
|
||||
url="{}{}".format(_normalize_api_base(api_base), "/upload"),
|
||||
headers={"Authorization": f"Bearer {api_key}"},
|
||||
files={"file": ("document", raw_bytes, mime or "application/octet-stream")},
|
||||
timeout=request_timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return _extract_file_id_from_upload_response(response)
|
||||
|
||||
|
||||
def build_pages_from_reducto(result: dict[str, Any]) -> list["OCRPage"]:
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRPage
|
||||
|
||||
chunks: Final = result.get("chunks", []) or []
|
||||
blocks_by_page: Final[dict[int, list[dict[str, Any]]]] = defaultdict(list)
|
||||
|
||||
for chunk in chunks:
|
||||
for block in chunk.get("blocks", []) or []:
|
||||
page_no = (block.get("bbox") or {}).get("page")
|
||||
if page_no is None:
|
||||
continue
|
||||
try:
|
||||
normalized_page = int(page_no)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
blocks_by_page[normalized_page].append(block)
|
||||
|
||||
if not blocks_by_page:
|
||||
fallback_markdown: Final = "\n\n".join(chunk.get("content", "") for chunk in chunks if chunk.get("content"))
|
||||
if fallback_markdown == "":
|
||||
return []
|
||||
return [OCRPage(index=0, markdown=fallback_markdown)]
|
||||
|
||||
pages: Final[list[OCRPage]] = []
|
||||
for page_no, blocks in sorted(blocks_by_page.items()):
|
||||
markdown = "\n\n".join(block.get("content", "") for block in blocks if block.get("content"))
|
||||
page_index = max(page_no - 1, 0)
|
||||
page = OCRPage(
|
||||
index=page_index,
|
||||
markdown=markdown,
|
||||
)
|
||||
# OCRPage accepts extra keys at runtime; assign blocks after construction
|
||||
# so static typing does not reject provider-specific metadata.
|
||||
setattr(page, "blocks", blocks)
|
||||
pages.append(page)
|
||||
return pages
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue