chore: merge main into litellm_batch_jsonl_line_item_callbacks
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-25 05:30:36 +00:00
commit 87d82b9e04
490 changed files with 18674 additions and 70097 deletions

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

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

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

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -167,6 +167,9 @@ MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_M
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600"))
MCP_SSO_ASSERTION_CACHE_TTL_SECONDS: Final = int(os.getenv("MCP_SSO_ASSERTION_CACHE_TTL_SECONDS", "60"))
# mcp_tool_permissions entry that grants every current and future tool on a server
MCP_ALL_TOOLS_WILDCARD: Final = "*"
# Default npm cache directory for STDIO MCP servers.
# npm/npx needs a writable cache dir; in containers the default (~/.npm)
# may not exist or be read-only. /tmp is always writable.
@ -327,6 +330,7 @@ REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float(
WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
DEEPGRAM_DEFAULT_API_BASE: Final = "https://api.deepgram.com/v1"
NADIR_DEFAULT_API_BASE: Final = "https://api.getnadir.com/v1"
DEEPGRAM_LISTEN_DEFAULT_MODEL: Final = "nova-3"
BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY: Final = "litellm.bedrock_realtime.pending_session_update"
@ -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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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:

View file

@ -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]:

View file

@ -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",
]

View file

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

View file

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

View file

@ -1,5 +0,0 @@
"""Azure Document Intelligence OCR module."""
from .transformation import AzureDocumentIntelligenceOCRConfig
__all__ = ["AzureDocumentIntelligenceOCRConfig"]

View file

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

View file

@ -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,
)

View file

@ -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

View file

@ -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",
]

View file

@ -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,
)

View file

@ -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,
)

View file

@ -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:

View file

@ -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={},
)

View file

@ -1,3 +0,0 @@
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
__all__ = ("CohereParseConfig",)

View file

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

View file

@ -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

View file

@ -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,

View file

@ -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
"""

View file

@ -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

View 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

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

View file

@ -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