refactor(rust): share call lifecycle across route-owned inference (#43461)

* feat(rust): expand gateway configuration parsing

* refactor(rust): unify core calls and host lifecycle

* fix(config): accept environment references for model rate limits

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(rust): unify core calls and host lifecycle

* style(rust): apply rustfmt

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(rust): read environment secrets when litellm is not importable

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(core): drop the duplicate rstest attribute

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(rust): make shared route dispatch route-owned

* fix(rust): satisfy Clippy in messages regression test

---------

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-27 16:05:30 -07:00 • committed by GitHub
parent b94f5bdbed
commit 36784e3b79
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
95 changed files with 3806 additions and 2447 deletions

View file

@ -3279,6 +3279,7 @@ dependencies = [
name = "litellm-host"
version = "0.1.0"
dependencies = [
"futures-util",
"litellm-auth",
"litellm-coroutine",
"rstest",

View file

@ -2,7 +2,7 @@
- This crate is the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on)
- Smell test: if a future callback host (`callbacks-v1-python`, WASM, in-process Rust) could share a piece of this crate, it does not belong here
- SDK request policy (credential inheritance, the budget and retry-count limits) is the driver's preflight, supplied by `python-bridge`; this crate only adopts the keyword view it produces
- The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call
- The driver in `litellm-host-python`, the routes and core see one `PythonCallHooks`; they never learn which Python objects consume a call
- Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json`
- The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it
- Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython`

View file

@ -2,12 +2,12 @@
//! raises is answered with the same `Logging` calls, in the same order, as the Python
//! `@client` path makes them.
use litellm_host_python::PythonOwned;
use litellm_host::event::{
FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest, epoch_seconds,
};
use litellm_host_python::{
LifecycleEvent, LifecycleStep, PythonLifecycle, from_py, missing_state, to_py,
};
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, from_py, missing_state, to_py};
use pyo3::{
exceptions::{PyBaseException, PyException},
gc::{PyTraverseError, PyVisit},
@ -48,11 +48,10 @@ struct DeliveredStream {
first_chunk: Option<Py<PyAny>>,
}
enum Pending {
DeploymentPreCall,
DeploymentPostCall,
DeploymentFailure,
AsyncFailure,
struct LoggedRequest {
body: Py<PyDict>,
headers: Py<PyDict>,
context: RequestContext,
}
pub struct LegacyLogging {
@ -63,13 +62,10 @@ pub struct LegacyLogging {
end: Option<Py<PyAny>>,
response: Option<Py<PyAny>>,
error: Option<Py<PyBaseException>>,
body: Option<Py<PyDict>>,
headers: Option<Py<PyDict>>,
context: Option<RequestContext>,
request: Option<LoggedRequest>,
stream: Option<DeliveredStream>,
asynchronous: bool,
internal: bool,
pending: Option<Pending>,
}
fn datetime(py: Python<'_>, epoch_seconds: f64) -> PyResult<Py<PyAny>> {
@ -95,13 +91,10 @@ impl LegacyLogging {
end: None,
response: None,
error: None,
body: None,
headers: None,
context: None,
request: None,
stream: None,
asynchronous,
internal: false,
pending: None,
}
}
@ -120,14 +113,14 @@ impl LegacyLogging {
/// The keyword view the rest of the call reads: a copy, so the deployment hook's own
/// dict is left as the hook returned it, carrying the logger as `@client` injects it.
/// The driver's preflight rewrites this same dict before the host projects from it.
fn prepare(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
fn prepare(&mut self, py: Python<'_>) -> PyResult<HookStep<Self, Py<PyDict>>> {
let prepared = self.call.kwargs().bind(py).copy()?;
prepared.set_item("litellm_logging_obj", self.logger()?.object(py))?;
self.call.set_kwargs(prepared.unbind());
Ok(LifecycleStep::Arguments(self.call.kwargs().clone_ref(py)))
Ok(HookStep::Ready(self.call.kwargs().clone_ref(py)))
}
fn finalize(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
fn finalize(&mut self, py: Python<'_>) -> PyResult<HookStep<Self, Py<PyAny>>> {
finalize(
py,
&self.response,
@ -138,7 +131,7 @@ impl LegacyLogging {
)?;
self.response
.as_ref()
.map(|response| LifecycleStep::Response(response.clone_ref(py)))
.map(|response| HookStep::Ready(response.clone_ref(py)))
.ok_or_else(missing_state)
}
@ -195,7 +188,10 @@ impl LegacyLogging {
logger.object(py),
billing.url_route,
billing.endpoint_type,
&self.body,
&self
.request
.as_ref()
.map(|request| request.body.clone_ref(py)),
&stream.chunks,
&self.start,
&self.end,
@ -214,11 +210,11 @@ impl LegacyLogging {
/// A failure after the stream reached the caller bills the delivered chunks as
/// partial usage. The sync path has no loop to schedule that on, so it falls back to
/// the plain failure handler.
fn stream_failure(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
fn stream_failure(&mut self, py: Python<'_>) -> PyResult<HookStep<Self, ()>> {
let (Some(logger), Some(error), Some(stream), Some(billing)) =
(&self.logger, &self.error, &self.stream, self.surface.stream)
else {
return Ok(LifecycleStep::Done);
return Ok(HookStep::Ready(()));
};
if !self.asynchronous {
return self.dispatch_failure(py);
@ -228,30 +224,33 @@ impl LegacyLogging {
(
logger.object(py),
billing.endpoint_type,
&self.body,
&self
.request
.as_ref()
.map(|request| request.body.clone_ref(py)),
&stream.chunks,
error,
),
);
match scheduled {
Ok(awaitable) => {
self.pending = Some(Pending::AsyncFailure);
Ok(LifecycleStep::Await(awaitable.unbind()))
}
Ok(awaitable) => Ok(HookStep::Await(
awaitable.unbind(),
Self::resume_async_failure,
)),
Err(failure) if is_cancellation(py, &failure) => Err(failure),
Err(_) => Ok(LifecycleStep::Done),
Err(_) => Ok(HookStep::Ready(())),
}
}
/// The sync failure handler, then the async one for async calls. Ordinary handler
/// errors never replace the selected failure or suppress the other family; a
/// cancellation does end the call.
fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult<HookStep<Self, ()>> {
let (Some(logger), Some(error)) = (&self.logger, &self.error) else {
return Ok(LifecycleStep::Done);
return Ok(HookStep::Ready(()));
};
if self.asynchronous && self.internal {
return Ok(LifecycleStep::Done);
return Ok(HookStep::Ready(()));
}
if let Err(failure) = logger.failure(py, error, &self.start, &self.end, false)
&& is_cancellation(py, &failure)
@ -259,27 +258,64 @@ impl LegacyLogging {
return Err(failure);
}
if !self.asynchronous {
return Ok(LifecycleStep::Done);
return Ok(HookStep::Ready(()));
}
match logger.failure(py, error, &self.start, &self.end, true) {
Ok(Some(awaitable)) => {
self.pending = Some(Pending::AsyncFailure);
Ok(LifecycleStep::Await(awaitable))
}
Ok(None) => Ok(LifecycleStep::Done),
Ok(Some(awaitable)) => Ok(HookStep::Await(awaitable, Self::resume_async_failure)),
Ok(None) => Ok(HookStep::Ready(())),
Err(failure) if is_cancellation(py, &failure) => Err(failure),
Err(_) => Ok(LifecycleStep::Done),
Err(_) => Ok(HookStep::Ready(())),
}
}
}
impl PythonLifecycle for LegacyLogging {
fn begin(
impl LegacyLogging {
fn resume_begin(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, Py<PyDict>>> {
self.call
.set_kwargs(result?.into_bound(py).cast_into::<PyDict>()?.unbind());
self.prepare(py)
}
fn resume_after_success(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, Py<PyAny>>> {
self.response = Some(result?);
self.finalize(py)
}
fn resume_deployment_failure(
&mut self,
py: Python<'_>,
_: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, ()>> {
self.dispatch_failure(py)
}
fn resume_async_failure(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, ()>> {
match result {
Err(error) if is_cancellation(py, &error) => Err(error),
_ => Ok(HookStep::Ready(())),
}
}
}
impl PythonCallHooks for LegacyLogging {
fn prepare_arguments(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
started_at: f64,
) -> PyResult<LifecycleStep> {
) -> PyResult<HookStep<Self, Py<PyDict>>> {
self.call.set_kwargs(arguments);
self.start = datetime(py, started_at)?;
self.internal = is_internal_call(py)?;
@ -294,22 +330,20 @@ impl PythonLifecycle for LegacyLogging {
self.logger = Some(result.logger()?);
self.call.set_kwargs(result.kwargs()?);
if self.runs_deployment_hooks() {
self.pending = Some(Pending::DeploymentPreCall);
return Ok(LifecycleStep::Await(DeploymentHooks::before_call(
py,
self.call.kwargs(),
self.surface.call_type,
)?));
return Ok(HookStep::Await(
DeploymentHooks::before_call(py, self.call.kwargs(), self.surface.call_type)?,
Self::resume_begin,
));
}
self.prepare(py)
}
fn before_send(
fn before_provider_request(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<LifecycleStep> {
) -> PyResult<HookStep<Self, Box<WireRequest>>> {
let logger = self.logger()?;
logger.update_from_kwargs(py, self.call.kwargs(), &wire, context)?;
let body = to_py(py, &wire.body)?
@ -326,9 +360,11 @@ impl PythonLifecycle for LegacyLogging {
for (name, value) in &wire.headers {
headers.set_item(name, value)?;
}
self.body = Some(body.clone().unbind());
self.headers = Some(headers.clone().unbind());
self.context = Some(context.clone());
self.request = Some(LoggedRequest {
body: body.clone().unbind(),
headers: headers.clone().unbind(),
context: context.clone(),
});
self.logger()?.pre_call(
py,
self.surface.input_description,
@ -341,61 +377,63 @@ impl PythonLifecycle for LegacyLogging {
.iter()
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
.collect::<PyResult<Vec<_>>>()?;
Ok(LifecycleStep::Wire(Box::new(WireRequest {
Ok(HookStep::Ready(Box::new(WireRequest {
body: from_py(&body)?,
headers,
..*wire
})))
}
fn after_success(
fn transform_response(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
timing: Timing,
) -> PyResult<LifecycleStep> {
) -> PyResult<HookStep<Self, Py<PyAny>>> {
self.end = Some(datetime(py, timing.end_time)?);
self.response = Some(response);
if self.runs_deployment_hooks() {
self.pending = Some(Pending::DeploymentPostCall);
return Ok(LifecycleStep::Await(DeploymentHooks::after_success(
py,
self.call.kwargs(),
&self.response,
self.surface.call_type,
)?));
return Ok(HookStep::Await(
DeploymentHooks::after_success(
py,
self.call.kwargs(),
&self.response,
self.surface.call_type,
)?,
Self::resume_after_success,
));
}
self.finalize(py)
}
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep> {
fn on_event(&mut self, py: Python<'_>, event: HookEvent<'_>) -> PyResult<HookStep<Self, ()>> {
match event {
LifecycleEvent::Started { .. } => Ok(LifecycleStep::Done),
LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
HookEvent::Started { .. } => Ok(HookStep::Ready(())),
HookEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
let api_key = self
.context
.request
.as_ref()
.and_then(|context| context.api_key.as_ref())
.and_then(|request| request.context.api_key.as_ref())
.map(|api_key| api_key.expose());
self.logger()?.post_call(
py,
&raw.body,
api_key,
self.body.as_ref(),
self.headers.as_ref(),
self.request.as_ref().map(|request| &request.body),
self.request.as_ref().map(|request| &request.headers),
)?;
Ok(LifecycleStep::Done)
Ok(HookStep::Ready(()))
}
LifecycleEvent::Succeeded { timing, response } => {
HookEvent::Succeeded { timing, response } => {
self.end = Some(datetime(py, timing.end_time)?);
self.response = Some(response.clone_ref(py));
match &self.stream {
Some(stream) => self.stream_success(py, stream)?,
None => self.dispatch_success(py)?,
}
Ok(LifecycleStep::Done)
Ok(HookStep::Ready(()))
}
LifecycleEvent::Failed {
HookEvent::Failed {
timing,
origin,
error,
@ -410,20 +448,22 @@ impl PythonLifecycle for LegacyLogging {
&& self.runs_deployment_hooks()
{
let error = self.error.as_ref().ok_or_else(missing_state)?;
self.pending = Some(Pending::DeploymentFailure);
return Ok(LifecycleStep::Await(DeploymentHooks::after_failure(
py,
self.call.kwargs(),
error,
self.surface.call_type,
)?));
return Ok(HookStep::Await(
DeploymentHooks::after_failure(
py,
self.call.kwargs(),
error,
self.surface.call_type,
)?,
Self::resume_deployment_failure,
));
}
self.dispatch_failure(py)
}
}
}
fn opened(&mut self, py: Python<'_>) -> PyResult<()> {
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
if self.surface.stream.is_none() {
return Err(missing_state());
}
@ -435,45 +475,25 @@ impl PythonLifecycle for LegacyLogging {
Ok(())
}
fn delivered(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
let stream = self.stream.as_mut().ok_or_else(missing_state)?;
if stream.first_chunk.is_none() {
stream.first_chunk = Some(datetime(py, epoch_seconds())?);
}
stream.chunks.bind(py).append(chunk)
}
}
fn resume(&mut self, py: Python<'_>, result: PyResult<Py<PyAny>>) -> PyResult<LifecycleStep> {
match self.pending.take().ok_or_else(missing_state)? {
Pending::DeploymentPreCall => {
self.call
.set_kwargs(result?.into_bound(py).cast_into::<PyDict>()?.unbind());
self.prepare(py)
}
Pending::DeploymentPostCall => {
self.response = Some(result?);
self.finalize(py)
}
Pending::DeploymentFailure => self.dispatch_failure(py),
Pending::AsyncFailure => match result {
Err(failure) if is_cancellation(py, &failure) => Err(failure),
_ => Ok(LifecycleStep::Done),
},
}
}
impl PythonOwned for LegacyLogging {
fn close(&mut self, py: Python<'_>) {
if let Some(logger) = self.logger.take()
&& let Err(error) = logger.restore_context(py)
{
error.write_unraisable(py, None);
}
self.body = None;
self.headers = None;
self.context = None;
self.request = None;
self.stream = None;
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
self.call.traverse(visit)?;
if let Some(logger) = &self.logger {
@ -487,8 +507,11 @@ impl PythonLifecycle for LegacyLogging {
visit.call(&stream.chunks)?;
visit.call(&stream.first_chunk)?;
}
visit.call(&self.body)?;
visit.call(&self.headers)
if let Some(request) = &self.request {
visit.call(&request.body)?;
visit.call(&request.headers)?;
}
Ok(())
}
}
@ -497,7 +520,7 @@ mod deployment_hooks_tests {
use std::ffi::CStr;
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks};
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
use pyo3::types::PyDict;
@ -506,6 +529,18 @@ mod deployment_hooks_tests {
use super::LegacyLogging;
use crate::test_support::{legacy_call, local, namespace, run};
fn resume<T>(
logging: &mut LegacyLogging,
step: HookStep<LegacyLogging, T>,
py: Python<'_>,
value: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<LegacyLogging, T>> {
let HookStep::Await(_, continuation) = step else {
panic!("expected suspension")
};
continuation(logging, py, value)
}
const CALL: &CStr = c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
kwargs = {'logger': logger, 'document': document}
@ -520,25 +555,28 @@ kwargs = {'logger': logger, 'document': document}
py: Python<'py>,
locals: &Bound<'py, PyDict>,
asynchronous: bool,
) -> (LegacyLogging, LifecycleStep) {
) -> (LegacyLogging, HookStep<LegacyLogging, Py<PyDict>>) {
let mut logging = legacy_call(py, locals, asynchronous);
let kwargs = local(locals, "kwargs")
.cast_into::<PyDict>()
.unwrap()
.unbind();
let step = logging.begin(py, kwargs, 0.0).unwrap();
let step = logging.prepare_arguments(py, kwargs, 0.0).unwrap();
(logging, step)
}
fn arguments<'py>(py: Python<'py>, step: LifecycleStep) -> Bound<'py, PyDict> {
let LifecycleStep::Arguments(arguments) = step else {
fn arguments<'py>(
py: Python<'py>,
step: HookStep<LegacyLogging, Py<PyDict>>,
) -> Bound<'py, PyDict> {
let HookStep::Ready(arguments) = step else {
panic!("expected the prepared arguments");
};
arguments.into_bound(py)
}
fn awaits_deployment_hook(step: &LifecycleStep) -> bool {
matches!(step, LifecycleStep::Await(_))
fn awaits_deployment_hook<T>(step: &HookStep<LegacyLogging, T>) -> bool {
matches!(step, HookStep::Await(_, _))
}
#[rstest]
@ -559,7 +597,7 @@ kwargs = {'logger': logger, 'document': document}
});
}
#[test]
#[rstest::rstest]
fn kwargs_returned_by_the_pre_call_hook_are_what_the_call_prepares() {
Python::initialize();
Python::attach(|py| {
@ -574,9 +612,13 @@ replaced_kwargs = {'logger': logger, 'document': replacement, 'pages': [0]}
);
let (mut logging, step) = begin(py, &locals, true);
assert!(awaits_deployment_hook(&step));
let step = logging
.resume(py, Ok(local(&locals, "replaced_kwargs").unbind()))
.unwrap();
let step = resume(
&mut logging,
step,
py,
Ok(local(&locals, "replaced_kwargs").unbind()),
)
.unwrap();
locals.set_item("prepared", arguments(py, step)).unwrap();
run(
py,
@ -610,7 +652,9 @@ kwargs = {'logger': logger, 'vendor_extension': opaque}
);
let (mut logging, step) = begin(py, &locals, asynchronous);
let step = match step {
LifecycleStep::Await(hook_result) => logging.resume(py, Ok(hook_result)).unwrap(),
HookStep::Await(hook_result, resume) => {
resume(&mut logging, py, Ok(hook_result)).unwrap()
}
step => step,
};
locals.set_item("prepared", arguments(py, step)).unwrap();
@ -626,7 +670,7 @@ assert hooked == ([opaque] if asynchronous else []), hooked
});
}
#[test]
#[rstest::rstest]
fn response_returned_by_the_post_call_hook_is_finalized_and_returned() {
Python::initialize();
Python::attach(|py| {
@ -639,18 +683,26 @@ replacement = object()
logger.hooks = {'pre': lambda kwargs: kwargs}
",
);
let (mut logging, _) = begin(py, &locals, true);
logging
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
let (mut logging, step) = begin(py, &locals, true);
resume(
&mut logging,
step,
py,
Ok(local(&locals, "kwargs").unbind()),
)
.unwrap();
let step = logging
.after_success(py, local(&locals, "response").unbind(), TIMING)
.transform_response(py, local(&locals, "response").unbind(), TIMING)
.unwrap();
assert!(awaits_deployment_hook(&step));
let step = logging
.resume(py, Ok(local(&locals, "replacement").unbind()))
.unwrap();
let LifecycleStep::Response(returned) = step else {
let step = resume(
&mut logging,
step,
py,
Ok(local(&locals, "replacement").unbind()),
)
.unwrap();
let HookStep::Ready(returned) = step else {
panic!("expected the finalized response");
};
assert!(returned.bind(py).is(local(&locals, "replacement")));
@ -672,18 +724,28 @@ assert finalized is replacement
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"kwargs = {'logger': logger}\nresponse = object()");
let (mut logging, _) = begin(py, &locals, true);
if post_call {
logging
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
logging
.after_success(py, local(&locals, "response").unbind(), TIMING)
.unwrap();
}
let (mut logging, step) = begin(py, &locals, true);
let cancellation = CancelledError::new_err("cancelled");
let cancelled = cancellation.value(py).clone();
let error = logging.resume(py, Err(cancellation)).err().unwrap();
let error = if post_call {
resume(
&mut logging,
step,
py,
Ok(local(&locals, "kwargs").unbind()),
)
.unwrap();
let step = logging
.transform_response(py, local(&locals, "response").unbind(), TIMING)
.unwrap();
resume(&mut logging, step, py, Err(cancellation))
.err()
.unwrap()
} else {
resume(&mut logging, step, py, Err(cancellation))
.err()
.unwrap()
};
assert!(error.value(py).is(&cancelled));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
@ -704,17 +766,21 @@ assert finalized is replacement
py,
c"kwargs = {'logger': logger}\nfailure = ValueError('provider')",
);
let (mut logging, _) = begin(py, &locals, true);
logging
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
let (mut logging, step) = begin(py, &locals, true);
resume(
&mut logging,
step,
py,
Ok(local(&locals, "kwargs").unbind()),
)
.unwrap();
let failure = PyErr::from_value(local(&locals, "failure"));
let failed = LifecycleEvent::Failed {
let failed = HookEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Call,
error: &failure,
};
let step = logging.emit(py, failed).unwrap();
let step = logging.on_event(py, failed).unwrap();
assert!(awaits_deployment_hook(&step));
let hook_result = if cancelled {
Err(CancelledError::new_err("cancelled"))
@ -722,8 +788,8 @@ assert finalized is replacement
Ok(py.None())
};
assert!(matches!(
logging.resume(py, hook_result).unwrap(),
LifecycleStep::Await(_)
resume(&mut logging, step, py, hook_result).unwrap(),
HookStep::Await(_, _)
));
run(
py,
@ -743,7 +809,7 @@ mod payload_tests {
use litellm_auth::SecretValue;
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py};
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, PythonOwned, to_py};
use proptest::prelude::*;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
@ -788,11 +854,11 @@ check = lambda: None
json!({"type": "document_url", "document_url": source})
}
fn before_send(script: &CStr, body: Value) -> WireRequest {
fn before_provider_request(script: &CStr, body: Value) -> WireRequest {
before_send_with_secrets(script, json!({}), body, &[])
}
/// Runs `before_send` over `body` for a route whose parameters are `optional_params`, with
/// Runs `before_provider_request` over `body` for a route whose parameters are `optional_params`, with
/// the Python objects `script` binds, then delivers the provider's raw response the way the
/// driver does and runs the script's `check()`.
fn before_send_with_secrets(
@ -837,30 +903,35 @@ check = lambda: None
};
let (_, step) = send_and_receive(py, &mut logging, wire, &context);
run(py, &locals, c"check()");
let LifecycleStep::Wire(wire) = step else {
panic!("before_send did not hand back the wire request");
let HookStep::Ready(wire) = step else {
panic!("before_provider_request did not hand back the wire request");
};
*wire
})
}
/// `before_send` over `wire`, then the provider's raw response the way the driver
/// `before_provider_request` over `wire`, then the provider's raw response the way the driver
/// delivers it, so `pre_call` and `post_call` have both seen the retained payload.
fn send_and_receive<'a>(
py: Python<'_>,
logging: &'a mut LegacyLogging,
wire: WireRequest,
context: &RequestContext,
) -> (&'a mut LegacyLogging, LifecycleStep) {
let step = logging.before_send(py, Box::new(wire), context).unwrap();
) -> (
&'a mut LegacyLogging,
HookStep<LegacyLogging, Box<WireRequest>>,
) {
let step = logging
.before_provider_request(py, Box::new(wire), context)
.unwrap();
let raw = MachineEvent::ResponseReceived {
raw: RawResponse {
body: "raw response".into(),
},
};
assert!(matches!(
logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(),
LifecycleStep::Done
logging.on_event(py, HookEvent::Machine(&raw)).unwrap(),
HookStep::Ready(())
));
(logging, step)
}
@ -904,7 +975,7 @@ check = lambda: None
}
}
#[test]
#[rstest::rstest]
fn a_cycle_through_the_retained_headers_is_collected() {
Python::initialize();
Python::attach(|py| {
@ -940,7 +1011,7 @@ assert reference() is None
});
}
#[test]
#[rstest::rstest]
fn close_releases_the_retained_headers() {
Python::initialize();
Python::attach(|py| {
@ -998,13 +1069,13 @@ def check():
")]
fn passthrough_keys_reach_pre_call_as_the_callers_own_objects(#[case] script: &CStr) {
let body = json!({"model": "model", "document": document(DOCUMENT), "pages": [0]});
let wire = before_send(script, body.clone());
let wire = before_provider_request(script, body.clone());
assert_eq!(wire.body, body);
}
#[test]
#[rstest::rstest]
fn pre_call_edit_of_a_passthrough_object_reaches_the_caller_and_the_wire() {
let wire = before_send(
let wire = before_provider_request(
c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
kwargs = {'document': document}
@ -1018,9 +1089,9 @@ def check():
assert_eq!(wire.body["document"], document(EDITED));
}
#[test]
#[rstest::rstest]
fn a_body_key_the_route_rewrote_is_not_the_callers_object() {
let wire = before_send(
let wire = before_provider_request(
c"
document = {'type': 'document_url', 'document_url': 'https://example.invalid/scan.pdf'}
kwargs = {'document': document}
@ -1040,10 +1111,10 @@ def check():
);
}
#[test]
#[rstest::rstest]
fn a_caller_value_with_no_json_form_is_left_out_of_realiasing() {
let body = json!({"pages": [0]});
let wire = before_send(
let wire = before_provider_request(
c"
opaque = object()
kwargs = {'pages': opaque}
@ -1072,14 +1143,14 @@ def on_pre_call(args):
)]
fn rebinding_the_payload_envelope_does_not_reach_the_wire(#[case] script: &CStr) {
let body = json!({"document": document(DOCUMENT)});
let wire = before_send(script, body.clone());
let wire = before_provider_request(script, body.clone());
assert_eq!(wire.body, body);
assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]);
}
#[test]
#[rstest::rstest]
fn pre_call_header_edit_reaches_the_wire() {
let wire = before_send(
let wire = before_provider_request(
c"
def on_pre_call(args):
args['headers']['x-callback'] = 'edited'
@ -1095,7 +1166,7 @@ def on_pre_call(args):
);
}
#[test]
#[rstest::rstest]
fn pre_call_receives_the_wire_request_and_the_logger_its_redacted_request() {
let body = json!({"model": "model", "document": document(DOCUMENT)});
before_send_with_secrets(
@ -1167,13 +1238,13 @@ def on_pre_call(args):
)]
fn pre_call_body_edits_reach_the_wire(#[case] script: &CStr, #[case] expected: Value) {
let body = json!({"document": document(DOCUMENT)});
let wire = before_send(script, body);
let wire = before_provider_request(script, body);
assert_eq!(wire.body, expected);
}
#[test]
#[rstest::rstest]
fn retained_headers_edited_after_rebinding_reach_the_wire() {
let wire = before_send(
let wire = before_provider_request(
c"
def on_pre_call(args):
retained = args['headers']
@ -1191,9 +1262,9 @@ def on_pre_call(args):
);
}
#[test]
#[rstest::rstest]
fn post_call_receives_the_raw_response_the_route_key_and_the_body_and_headers_pre_call_saw() {
before_send(
before_provider_request(
c"
def check():
original_response, api_key, additional_args = logger.post
@ -1210,9 +1281,9 @@ def check():
);
}
#[test]
#[rstest::rstest]
fn every_request_runs_the_full_pre_call_and_post_call() {
let wire = before_send(
let wire = before_provider_request(
c"
def on_pre_call(args):
args['complete_input_dict']['include_image_base64'] = True
@ -1343,7 +1414,7 @@ def check():
/// For any body, any caller keywords and any callback edit: every keyword the route
/// sends unchanged reaches `pre_call` as the caller's own object, and the provider is
/// sent exactly what the model says, so a callback that edits nothing changes nothing.
#[test]
#[rstest::rstest]
fn the_wire_is_the_body_pre_call_received_as_the_callback_left_it(
fields in prop::collection::btree_map(key(), (json_value(), caller()), 0..5),
edit in edit(),
@ -1389,7 +1460,7 @@ mod terminal_tests {
use std::ffi::CStr;
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, PythonOwned};
use pyo3::exceptions::PyRuntimeError;
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
@ -1416,12 +1487,12 @@ mod terminal_tests {
py: Python<'_>,
locals: &Bound<'_, PyDict>,
logging: &mut LegacyLogging,
) -> LifecycleStep {
) -> HookStep<LegacyLogging, ()> {
let response = local(locals, "response").unbind();
logging
.emit(
.on_event(
py,
LifecycleEvent::Succeeded {
HookEvent::Succeeded {
timing: TIMING,
response: &response,
},
@ -1433,12 +1504,12 @@ mod terminal_tests {
py: Python<'_>,
locals: &Bound<'_, PyDict>,
logging: &mut LegacyLogging,
) -> LifecycleStep {
) -> HookStep<LegacyLogging, ()> {
let failure = PyErr::from_value(local(locals, "failure"));
logging
.emit(
.on_event(
py,
LifecycleEvent::Failed {
HookEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Host,
error: &failure,
@ -1468,7 +1539,7 @@ mod terminal_tests {
let mut logging = logged(py, &locals, asynchronous);
assert!(matches!(
succeed(py, &locals, &mut logging),
LifecycleStep::Done
HookStep::Ready(())
));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
@ -1503,7 +1574,7 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy
};
assert!(matches!(
fail(py, &locals, &mut logging),
LifecycleStep::Done
HookStep::Ready(())
));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
@ -1514,7 +1585,7 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy
});
}
#[test]
#[rstest::rstest]
fn internal_async_calls_skip_the_async_success_fan_out() {
Python::initialize();
Python::attach(|py| {
@ -1532,7 +1603,7 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy
});
}
#[test]
#[rstest::rstest]
fn a_failing_success_callback_is_reported_without_replacing_the_response() {
Python::initialize();
Python::attach(|py| {
@ -1552,7 +1623,7 @@ logger = FailingLogger()
let mut logging = logged(py, &locals, true);
assert!(matches!(
succeed(py, &locals, &mut logging),
LifecycleStep::Done
HookStep::Ready(())
));
assert!(
logging
@ -1581,10 +1652,7 @@ logger = FailingLogger()
let mut logging = logged(py, &locals, asynchronous);
let step = fail(py, &locals, &mut logging);
let awaits_async_handler = expected.contains(&"async_failure_handler");
assert_eq!(
matches!(step, LifecycleStep::Await(_)),
awaits_async_handler
);
assert_eq!(matches!(step, HookStep::Await(_, _)), awaits_async_handler);
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
@ -1599,7 +1667,7 @@ logger = FailingLogger()
});
}
#[test]
#[rstest::rstest]
fn a_failing_sync_failure_callback_keeps_the_error_and_still_runs_the_async_family() {
Python::initialize();
Python::attach(|py| {
@ -1619,7 +1687,7 @@ logger = FailingLogger()
let mut logging = logged(py, &locals, true);
assert!(matches!(
fail(py, &locals, &mut logging),
LifecycleStep::Await(_)
HookStep::Await(_, _)
));
assert!(
logging
@ -1649,15 +1717,17 @@ logger = FailingLogger()
Python::attach(|py| {
let locals = namespace(py, c"failure = ValueError('provider')");
let mut logging = logged(py, &locals, true);
fail(py, &locals, &mut logging);
let HookStep::Await(_, resume) = fail(py, &locals, &mut logging) else {
panic!("expected async failure handler")
};
let result = match error {
None => Ok(py.None()),
Some(false) => Err(PyRuntimeError::new_err("handler failed")),
Some(true) => Err(CancelledError::new_err("cancelled")),
};
let expected = result.as_ref().err().map(|error| error.value(py).clone());
match logging.resume(py, result) {
Ok(step) => assert!(done && matches!(step, LifecycleStep::Done)),
match resume(&mut logging, py, result) {
Ok(step) => assert!(done && matches!(step, HookStep::Ready(()))),
Err(propagated) => {
assert!(!done);
assert!(propagated.value(py).is(expected.unwrap()));
@ -1666,7 +1736,7 @@ logger = FailingLogger()
});
}
#[test]
#[rstest::rstest]
fn closing_restores_the_correlation_context_once() {
Python::initialize();
Python::attach(|py| {

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, protocol::Protocol};
use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call};
use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol};
use litellm_host_python::{Preflight, PythonBinding, PythonHostCalls, lookup, run_call};
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,
@ -71,21 +71,22 @@ pub fn run_legacy_call<H, M>(
py: Python<'_>,
surface: LegacySurface,
call: PublicCall,
machine: M,
start: impl FnOnce(<H::Protocol as Protocol>::Request) -> M + Send + Sync + 'static,
host: H,
preflight: Preflight,
asynchronous: bool,
) -> PyResult<Py<PyAny>>
where
H: ProtocolHost + 'static,
M: Machine<Protocol = H::Protocol, Complete = <H::Protocol as Protocol>::Response> + 'static,
H: PythonBinding + PythonHostCalls<H::Protocol> + 'static,
M: Machine<Protocol = H::Protocol> + 'static,
M::Complete: Into<HostedCompletion<<H::Protocol as Protocol>::Response>>,
{
let arguments = call.kwargs.clone_ref(py);
run_call(
py,
machine,
start,
host,
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
LegacyLogging::new(py, surface, call, asynchronous),
preflight,
arguments,
asynchronous,
@ -110,7 +111,7 @@ mod tests {
(call, locals)
}
#[test]
#[rstest::rstest]
fn capture_copies_the_keyword_dict_without_copying_its_values() {
Python::initialize();
Python::attach(|py| {

View file

@ -1,7 +1,7 @@
//! The legacy `@client` wrapper as the native call sees it: litellm's `Logging` object, the
//! sync and async callback registries it fans out to, the deployment hooks and the deferred
//! proxy release. All of it sits behind one
//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and
//! [`PythonCallHooks`](litellm_host_python::PythonCallHooks), so the driver, the routes and
//! core never learn which Python object is on the other end. The SDK's own request policy
//! (credential inheritance, the budget and retry limits) is the driver's preflight, not this
//! crate's.

View file

@ -1,9 +1,26 @@
use strum::{EnumString, IntoStaticStr};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CustomLlmProvider<'a> {
pub model: &'a str,
pub custom_llm_provider: &'a str,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)]
#[strum(serialize_all = "snake_case")]
pub enum LlmProviders {
Anthropic,
AwsTextract,
AzureAi,
Bedrock,
Cohere,
Mistral,
Openai,
OpenaiLike,
Reducto,
VertexAi,
}
pub fn get_custom_llm_provider<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,

View file

@ -1,6 +1,12 @@
litellm-core is the LiteLLM SDK in Rust. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route. `messages::messages()` returns `MessagesResponse::Message` for a completed response or `MessagesResponse::Stream { headers, chunks }` when the request sets `stream: true`. The chunks are Anthropic SSE bytes in a `Stream<Item = Result<Bytes, Error>>`. Dropping the stream cancels the call. The Python bridge drives `messages::route::messages_machine()` instead, because Python has to answer the call's operations on its own thread; the gateway and the Rust SDK call the plain entrypoint
litellm-core owns route orchestration. Messages and HTTP Responses return `litellm_host::call::CallOutput`, containing either a completed response or a stream head and chunks. OCR and currently non-streaming Chat Completions return their completed response directly
A route module has the same five pieces, in the order Python runs them. `types.rs` holds the call, the provider request, and the response. `prepare.rs` resolves the provider and credentials and shapes the request (Python's `validate_environment`, `get_complete_url`, `transform_request`). `handler.rs` resolves auth, offers the wire request to `litellm_host::hooks::RouteHooks::before_send`, sends it, reports the raw response through `emit`, and normalizes the response or stream (`pre_call`, `post`, `post_call`, `transform_response`). `mod.rs` exposes the entrypoint that runs prepare then handler with no hooks (`()`). `route.rs`, where a host needs it, wraps the same two calls in a `CallMachine` whose `HostChannel` is the hooks, and pumps a stream through `open` and `deliver`. A handler takes `&impl RouteHooks<Error>` and never a `HostChannel` directly, so it runs without a coroutine. Keep provider transport and transformation details out of the machine driver
Hosts assemble route objects from shared `CoreResources`, HTTP settings, and secret sources. Each route owns its provider client and authentication dependencies. Gateway routes live for the gateway lifetime; Python assembles routes per call from its settings snapshot
Chat Completions, Messages, and OCR execute through their route objects. Calls pass `RouteHooks` directly; use `&()` when no hooks are needed. Construction does no work; preparation and lifecycle observation begin when the future is polled. Handlers accept `RouteHooks`, never a concrete `ChannelHooks`. Native observers receive start and terminal events through the shared call runner; a stream retains its lifecycle until exhaustion, error, or drop. A host channel has no native observer because its driver owns terminal dispatch
`route.rs` declares the concrete `Protocol` and implements a route method that accepts a typed request and constructs a `litellm_host::call::HostedMachine` with `hosted_call`. The shared call plumbing owns stream opening, delivery, backpressure, and detachment. Request decoding belongs to the boundary before the machine starts. Route closures only supply execution dependencies and route-specific host capabilities such as an OCR token provider. Use `run_hosted` for a native host so detachment is reported as cancellation. Python uses its own shared driver and preserves caller-task callback execution
Responses WebSocket sessions remain separate from the HTTP call driver because a connection can accept multiple requests while receiving events
## Crate layering

View file

@ -17,7 +17,7 @@ pub async fn execute_audio_transcription_provider_call(
) -> Result<Value, Error> {
let env_lookup = |key: &str| request.secrets.get(key);
let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?;
let response = crate::outbound::outbound_request(
let outbound = crate::outbound::outbound_request(
authenticated,
request.url.clone(),
&request.body,
@ -26,12 +26,12 @@ pub async fn execute_audio_transcription_provider_call(
.timeout
.unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)),
),
)?
.send(http)
.await
.map_err(|error| {
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
})?;
)?;
let response = crate::outbound::send(outbound, http)
.await
.map_err(|error| {
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
})?;
let status = response.status();
let text = response.text().await.map_err(|error| {
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))

View file

@ -1,22 +1,39 @@
use litellm_secrets::source::SecretSource;
pub mod types;
pub use crate::error::RouteError as Error;
mod handler;
mod prepare;
pub use handler::execute_audio_transcription_provider_call;
use litellm_http::{ClientVariant, HttpClientConfig};
use litellm_auth::AuthServices;
use litellm_secrets::source::SecretSource;
pub use prepare::prepare_audio_transcription_provider_call;
use serde_json::Value;
use std::sync::Arc;
use crate::audio_transcription::types::AudioTranscriptionRequest;
pub async fn audio_transcription(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
secrets: &dyn SecretSource,
request: AudioTranscriptionRequest<'_>,
) -> Result<Value, Error> {
let request = prepare_audio_transcription_provider_call(request, secrets).await?;
let http = resources.pool.client(config, ClientVariant::Provider)?;
execute_audio_transcription_provider_call(&http, &resources.auth, request).await
#[derive(Clone)]
pub struct AudioTranscriptionRoute {
http: litellm_http::Client,
auth: Arc<AuthServices>,
secrets: Arc<dyn SecretSource>,
}
impl AudioTranscriptionRoute {
pub fn new(
http: litellm_http::Client,
auth: Arc<AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Self {
Self {
http,
auth,
secrets,
}
}
pub async fn execute(&self, request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
let request =
prepare_audio_transcription_provider_call(request, self.secrets.as_ref()).await?;
execute_audio_transcription_provider_call(&self.http, &self.auth, request).await
}
}

View file

@ -1,4 +1,3 @@
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use litellm_http::request::string_headers;
use litellm_http::request::with_default_headers;
use litellm_llms::{
@ -14,36 +13,35 @@ use super::Error;
use crate::audio_transcription::types::{
AudioTranscriptionRequest, ProviderAudioTranscriptionRequest,
};
use crate::provider::{LlmProviders, resolve_llm_provider};
fn provider_config(provider: &str) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
if provider == "bedrock" {
return Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG);
fn provider_config(provider: LlmProviders) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
match provider {
LlmProviders::Bedrock => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG),
LlmProviders::Anthropic
| LlmProviders::AwsTextract
| LlmProviders::AzureAi
| LlmProviders::Cohere
| LlmProviders::Mistral
| LlmProviders::Openai
| LlmProviders::OpenaiLike
| LlmProviders::Reducto
| LlmProviders::VertexAi => None,
}
let _ = provider;
None
}
pub async fn prepare_audio_transcription_provider_call(
request: AudioTranscriptionRequest<'_>,
secrets: &dyn SecretSource,
) -> Result<ProviderAudioTranscriptionRequest, Error> {
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.or_else(|| {
request
.custom_llm_provider
.map(|provider| CustomLlmProvider {
model: request.model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(
"unable to resolve custom_llm_provider for audio transcription request".to_string(),
)
})?;
let provider_info = resolve_llm_provider(
request.model,
request.custom_llm_provider,
"audio transcription",
)?;
let model = provider_info.model.to_string();
let config = provider_config(provider_info.custom_llm_provider)
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
let config = provider_config(provider_info.provider)
.ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))?;
let snapshot = secrets.resolve(&config.secret_names()).await?;
let env_lookup = |key: &str| snapshot.get(key);
let forwarded = string_headers("audio transcription", request.extra_headers)?;
@ -64,7 +62,7 @@ pub async fn prepare_audio_transcription_provider_call(
config.transform_audio_transcription_request(&model, request.audio, filtered_params)?;
Ok(ProviderAudioTranscriptionRequest {
model,
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
custom_llm_provider: <&str>::from(provider_info.provider).to_string(),
config,
url,
body: transformed.body,

View file

@ -8,15 +8,38 @@ use litellm_llms::{
use serde_json::{Map, Value};
use super::Error;
use crate::provider::LlmProviders;
const HEADER_CONTEXT: &str = "chat completions";
pub(super) fn chat_completions_provider_config(provider: &str) -> Option<&'static dyn BaseConfig> {
pub(super) enum ChatProvider {
Anthropic,
Bedrock,
OpenaiLike,
}
impl ChatProvider {
pub(super) fn config(self) -> &'static dyn BaseConfig {
match self {
Self::Anthropic => &ANTHROPIC_CHAT_COMPLETIONS_CONFIG,
Self::Bedrock => &BEDROCK_CHAT_COMPLETIONS_CONFIG,
Self::OpenaiLike => &OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG,
}
}
}
pub(super) fn chat_completions_provider(provider: LlmProviders) -> Option<ChatProvider> {
match provider {
"anthropic" => Some(&ANTHROPIC_CHAT_COMPLETIONS_CONFIG),
"bedrock" => Some(&BEDROCK_CHAT_COMPLETIONS_CONFIG),
"openai_like" => Some(&OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG),
_ => None,
LlmProviders::Anthropic => Some(ChatProvider::Anthropic),
LlmProviders::Bedrock => Some(ChatProvider::Bedrock),
LlmProviders::OpenaiLike => Some(ChatProvider::OpenaiLike),
LlmProviders::AwsTextract
| LlmProviders::AzureAi
| LlmProviders::Cohere
| LlmProviders::Mistral
| LlmProviders::Openai
| LlmProviders::Reducto
| LlmProviders::VertexAi => None,
}
}

View file

@ -46,7 +46,7 @@ pub(super) async fn execute(
};
let authenticated = resolve_auth(auth, environment, &|key| secrets.get(key)).await?;
let wire = hooks
.before_send(
.before_provider_request(
WireRequest {
url,
headers: authenticated.headers,
@ -65,7 +65,7 @@ pub(super) async fn execute(
timeout,
)?;
let response = outbound.send(http).await.map_err(|err| {
let response = crate::outbound::send(outbound, http).await.map_err(|err| {
// Failing to establish the connection means the request never went out,
// so the host can still serve it. Everything else here, a timeout
// above all, may have reached the provider and been answered.
@ -88,10 +88,11 @@ pub(super) async fn execute(
}));
}
hooks
.emit(MachineEvent::ResponseReceived {
.on_event(MachineEvent::ResponseReceived {
raw: RawResponse { body: text.clone() },
})
.await?;
.await
.map_err(Error::post_call)?;
let body: Value = serde_json::from_str(&text).map_err(|err| {
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
@ -168,7 +169,7 @@ mod tests {
}
impl RouteHooks<Error> for RecordingHooks {
async fn before_send(
async fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
@ -187,7 +188,7 @@ mod tests {
})
}
async fn emit(&self, event: MachineEvent) -> Result<(), Error> {
async fn on_event(&self, event: MachineEvent) -> Result<(), Error> {
let MachineEvent::ResponseReceived { raw } = event;
self.raw.lock().unwrap().push(raw.body);
Ok(())
@ -240,7 +241,7 @@ mod tests {
assert_eq!(request.headers["x-host"], "seen");
assert_eq!(request.headers["x-api-key"], "sk-test");
let [context] = <[RequestContext; 1]>::try_from(hooks.contexts.into_inner().unwrap())
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
.unwrap_or_else(|seen| panic!("before_provider_request runs once, saw {}", seen.len()));
assert_eq!(
(context.model.as_str(), context.custom_llm_provider.as_str()),
("claude-sonnet-4-5", "anthropic")
@ -275,7 +276,7 @@ mod tests {
assert!(hooks.raw.into_inner().unwrap().is_empty());
}
#[test]
#[rstest::rstest]
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
for original in [
Error::MissingField("usage"),

View file

@ -1,35 +1,22 @@
//! The `/chat/completions` call, the Rust equivalent of Python's
//! `litellm.completion()`.
//!
//! [`chat_completions`] is the top-level entrypoint: give it a model, the
//! OpenAI-shaped message list, the provider-mapped optional params, and
//! credentials, and it resolves the provider, translates the conversation,
//! calls the provider, and returns a typed OpenAI-shaped response.
use litellm_secrets::source::SecretSource;
pub mod types;
pub use crate::error::RouteError as Error;
mod common_utils;
pub(crate) mod handler;
mod prepare;
use litellm_http::{ClientVariant, HttpClientConfig};
use litellm_types::utils::ChatCompletionsResponse;
use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request};
use serde_json::{Map, Value};
use crate::chat_completions::types::ChatCompletionsRequest;
use litellm_auth::AuthServices;
use litellm_secrets::source::SecretSource;
use std::sync::Arc;
pub async fn chat_completions(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
secrets: &dyn SecretSource,
request: ChatCompletionsRequest<'_>,
) -> Result<ChatCompletionsResponse, Error> {
let http = resources.pool.client(config, ClientVariant::Provider)?;
let resolved = resolve_request(request)?;
let snapshot = secrets.resolve(&resolved.config.secret_names()).await?;
let request = prepare_provider_request(resolved, snapshot)?;
handler::execute(&http, &resources.auth, request, &()).await
#[derive(Clone)]
pub struct ChatCompletionsRoute {
http: litellm_http::Client,
auth: Arc<AuthServices>,
secrets: Arc<dyn SecretSource>,
}
/// Whether the core would accept this request, without resolving credentials or
@ -58,3 +45,39 @@ pub fn chat_completions_decline_reason(
.unsupported_reason(&messages, optional_params)
.map(|reason| reason.0)
}
impl ChatCompletionsRoute {
pub fn new(
http: litellm_http::Client,
auth: Arc<AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Self {
Self {
http,
auth,
secrets,
}
}
pub async fn execute(
&self,
request: ChatCompletionsRequest<'_>,
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
) -> Result<ChatCompletionsResponse, Error> {
litellm_host::lifecycle::observe_unary(hooks.observer(), self.run(request, hooks)).await
}
async fn run(
&self,
request: ChatCompletionsRequest<'_>,
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
) -> Result<ChatCompletionsResponse, Error> {
let resolved = resolve_request(request)?;
let snapshot = self
.secrets
.resolve(&resolved.config.secret_names())
.await?;
let prepared = prepare_provider_request(resolved, snapshot)?;
handler::execute(&self.http, &self.auth, prepared, hooks).await
}
}

View file

@ -1,5 +1,4 @@
use litellm_auth::SecretValue;
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use litellm_core_utils::settings::Lookup;
use litellm_http::request::with_default_headers;
use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig};
@ -9,11 +8,12 @@ use serde_json::Value;
use super::{
Error,
common_utils::{chat_completions_provider_config, string_headers},
common_utils::{chat_completions_provider, string_headers},
};
use crate::chat_completions::types::{
ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
};
use crate::provider::resolve_llm_provider;
pub(super) struct ResolvedProvider {
pub(super) model: String,
@ -25,23 +25,13 @@ pub(super) fn resolve_provider_config<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,
) -> Result<ResolvedProvider, Error> {
let provider_info = get_custom_llm_provider(model, custom_llm_provider)
.or_else(|| {
custom_llm_provider.map(|provider| CustomLlmProvider {
model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(
"unable to resolve custom_llm_provider for chat completions request".to_string(),
)
})?;
let config = chat_completions_provider_config(provider_info.custom_llm_provider)
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
let provider_info = resolve_llm_provider(model, custom_llm_provider, "chat completions")?;
let config = chat_completions_provider(provider_info.provider)
.ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))?
.config();
Ok(ResolvedProvider {
model: provider_info.model.to_string(),
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
custom_llm_provider: <&str>::from(provider_info.provider).to_string(),
config,
})
}

View file

@ -37,6 +37,8 @@ pub enum RouteError {
Http(#[from] litellm_http::Error),
#[error(transparent)]
Secret(#[from] SecretError),
#[error("post-call hook failed: {0}")]
PostCallHook(#[source] Arc<RouteError>),
}
/// Whether the provider had already been called when the route failed. Before the send, a
@ -47,10 +49,25 @@ pub enum Phase {
AfterSend,
}
impl From<litellm_host::machine::MachineFault> for RouteError {
fn from(fault: litellm_host::machine::MachineFault) -> Self {
use litellm_host::machine::MachineFault;
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("host {message}").into(),
})
}
}
impl RouteError {
pub(crate) fn post_call(error: Self) -> Self {
Self::PostCallHook(Arc::new(error))
}
pub fn phase(&self) -> Phase {
match self {
Self::InvalidResponse(_)
| Self::PostCallHook(_)
| Self::Transport(TransportError::Http { .. } | TransportError::Network(_)) => {
Phase::AfterSend
}
@ -78,9 +95,11 @@ impl RouteError {
| Self::Unsupported(_)
| Self::Headers(_) => true,
Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }),
Self::InvalidResponse(_) | Self::Transport(_) | Self::Http(_) | Self::Secret(_) => {
false
}
Self::InvalidResponse(_)
| Self::Transport(_)
| Self::Http(_)
| Self::Secret(_)
| Self::PostCallHook(_) => false,
}
}
}

View file

@ -5,6 +5,7 @@ pub mod error;
pub mod messages;
pub mod ocr;
mod outbound;
mod provider;
pub mod resources;
pub mod responses;

View file

@ -7,14 +7,13 @@ use litellm_llms::{
bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG,
};
use serde_json::{Map, Value};
use strum::{EnumString, IntoStaticStr};
use super::Error;
use crate::provider::LlmProviders;
const HEADER_CONTEXT: &str = "messages";
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
#[strum(serialize_all = "snake_case")]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum MessagesProvider {
Anthropic,
AzureAi,
@ -23,7 +22,12 @@ pub(crate) enum MessagesProvider {
impl MessagesProvider {
pub(crate) fn as_str(self) -> &'static str {
self.into()
match self {
Self::Anthropic => LlmProviders::Anthropic,
Self::AzureAi => LlmProviders::AzureAi,
Self::Bedrock => LlmProviders::Bedrock,
}
.into()
}
pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig {
@ -35,6 +39,21 @@ impl MessagesProvider {
}
}
pub(crate) fn messages_provider(provider: LlmProviders) -> Option<MessagesProvider> {
match provider {
LlmProviders::Anthropic => Some(MessagesProvider::Anthropic),
LlmProviders::AzureAi => Some(MessagesProvider::AzureAi),
LlmProviders::Bedrock => Some(MessagesProvider::Bedrock),
LlmProviders::AwsTextract
| LlmProviders::Cohere
| LlmProviders::Mistral
| LlmProviders::Openai
| LlmProviders::OpenaiLike
| LlmProviders::Reducto
| LlmProviders::VertexAi => None,
}
}
pub(super) fn string_headers(
extra_headers: Option<Map<String, Value>>,
) -> Result<Vec<(String, String)>, Error> {
@ -47,8 +66,9 @@ mod tests {
use rstest::rstest;
use super::{MessagesProvider, string_headers, truncate_error_body};
use super::{MessagesProvider, messages_provider, string_headers, truncate_error_body};
use crate::messages::Error;
use crate::provider::LlmProviders;
#[rstest]
#[case::anthropic("anthropic", MessagesProvider::Anthropic)]
@ -58,13 +78,16 @@ mod tests {
#[case] name: &str,
#[case] provider: MessagesProvider,
) {
assert_eq!(name.parse::<MessagesProvider>(), Ok(provider));
assert_eq!(
messages_provider(name.parse::<LlmProviders>().unwrap()),
Some(provider)
);
assert_eq!(provider.as_str(), name);
}
#[test]
fn provider_without_a_messages_config_is_rejected() {
assert!("openai".parse::<MessagesProvider>().is_err());
assert_eq!(messages_provider(LlmProviders::Openai), None);
}
#[test]

View file

@ -15,7 +15,7 @@ use litellm_llms::base_llm::{
transformation::BaseAnthropicMessagesConfig,
},
};
use litellm_tracing::{ByteChunk, debug};
use litellm_tracing::ByteChunk;
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use serde_json::Value;
@ -48,7 +48,7 @@ pub(super) async fn execute(
};
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
let wire = hooks
.before_send(
.before_provider_request(
WireRequest {
url,
headers: authenticated.headers,
@ -58,7 +58,7 @@ pub(super) async fn execute(
)
.await?;
let provider_name = provider.as_str();
debug!(provider = provider_name, stream, body = %wire.body, "provider request");
log_request_body(provider_name, stream, &wire.body);
let response = send(
http,
Authenticated {
@ -70,11 +70,6 @@ pub(super) async fn execute(
timeout,
)
.await?;
debug!(
provider = provider_name,
status = response.status().as_u16(),
"provider response headers"
);
if !response.status().is_success() {
return Err(provider_error(response).await);
}
@ -87,14 +82,15 @@ pub(super) async fn execute(
));
}
let text = response.text().await.map_err(network)?;
debug!(body = text.as_str(), "provider response body");
log_response_body(&text);
hooks
.emit(MachineEvent::ResponseReceived {
.on_event(MachineEvent::ResponseReceived {
raw: RawResponse { body: text.clone() },
})
.await?;
.await
.map_err(Error::post_call)?;
decode_response(config, &body.model, &text)
.map(|message| MessagesResponse::Message(Box::new(message)))
.map(|message| MessagesResponse::Complete(Box::new(message)))
}
fn serialize_failure(err: serde_json::Error) -> Error {
@ -121,14 +117,14 @@ async fn send(
body,
Some(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))),
)?;
request.send(http).await.map_err(network)
crate::outbound::send(request, http).await.map_err(network)
}
async fn provider_error(response: reqwest::Response) -> Error {
let status = response.status().as_u16();
match response.text().await {
Ok(text) => {
litellm_tracing::debug!(status, body = text.as_str(), "provider error body");
log_error_body(status, &text);
Error::Transport(TransportError::Http {
status,
body: truncate_error_body(&text),
@ -175,7 +171,10 @@ fn streaming_response(
.boxed(),
Some(decode) => decoded_chunks(response, decode, provider),
};
MessagesResponse::Stream { headers, chunks }
MessagesResponse::Stream {
head: super::route::MessagesStreamHead { headers },
chunks,
}
}
fn decoded_chunks(
@ -199,9 +198,21 @@ fn decoded_chunks(
.boxed()
}
fn log_chunk(provider: &str, stage: &str, data: &Bytes) {
fn log_request_body(provider: &str, stream: bool, body: &serde_json::Value) {
litellm_tracing::debug!(provider, stream, body = %body, "provider request");
}
fn log_response_body(body: &str) {
litellm_tracing::debug!(body, "provider response body");
}
fn log_error_body(status: u16, body: &str) {
litellm_tracing::debug!(status, body, "provider error body");
}
fn log_chunk(provider: &str, stage: &str, data: &bytes::Bytes) {
let chunk = ByteChunk::new(data);
debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
litellm_tracing::debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
}
#[cfg(test)]
@ -217,6 +228,7 @@ mod tests {
"data: {\"type\":\"ping\"}\n\n",
Some("event: ping\ndata: {\"type\":\"ping\"}\n\n")
)]
#[rstest::rstest]
#[case::invalid_event("data: invalid\n\ndata: {\"type\":\"ping\"}\n\n", None)]
#[tokio::test]
async fn decoded_streams_encode_events_and_stop_at_the_first_error(

View file

@ -1,27 +1,50 @@
//! The Anthropic Messages call, the Rust equivalent of Python's `litellm.messages()`.
//!
//! [`messages`] prepares the provider request and sends it in process. [`route`] runs the
//! same two steps as a machine for a host that answers the call's operations itself.
mod common_utils;
mod handler;
mod prepare;
pub mod route;
mod types;
use litellm_http::{ClientVariant, HttpClientConfig};
use litellm_auth::AuthServices;
use litellm_secrets::source::SecretSource;
use std::sync::Arc;
pub use crate::error::RouteError as Error;
pub use types::{MessagesCall, MessagesResponse, MessagesShaping, messages_body};
pub async fn messages(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
secrets: &dyn SecretSource,
call: MessagesCall,
) -> Result<MessagesResponse, Error> {
let http = resources.pool.client(config, ClientVariant::Provider)?;
let request = prepare::prepare(call, secrets).await?;
handler::execute(&http, &resources.auth, request, &()).await
#[derive(Clone)]
pub struct MessagesRoute {
http: litellm_http::Client,
auth: Arc<AuthServices>,
secrets: Arc<dyn SecretSource>,
}
impl MessagesRoute {
pub fn new(
http: litellm_http::Client,
auth: Arc<AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Self {
Self {
http,
auth,
secrets,
}
}
pub async fn execute(
&self,
call: MessagesCall,
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
) -> Result<MessagesResponse, Error> {
litellm_host::lifecycle::observe_call(hooks.observer(), self.run(call, hooks)).await
}
async fn run(
&self,
call: MessagesCall,
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
) -> Result<MessagesResponse, Error> {
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
handler::execute(&self.http, &self.auth, request, hooks).await
}
}

View file

@ -3,9 +3,7 @@ use std::time::Duration;
use litellm_auth::SecretValue;
use litellm_core_utils::{
dot_notation_indexing::delete_nested_value,
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
get_provider_specific_headers::get_provider_specific_headers,
settings::Lookup,
get_provider_specific_headers::get_provider_specific_headers, settings::Lookup,
};
use litellm_http::request::with_default_headers;
use litellm_llms::base_llm::{
@ -16,9 +14,10 @@ use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessage
use super::{
Error, MessagesCall,
common_utils::{MessagesProvider, string_headers},
common_utils::{MessagesProvider, messages_provider, string_headers},
types::invalid_request,
};
use crate::provider::resolve_llm_provider;
struct ResolvedProvider {
model: String,
@ -50,26 +49,11 @@ fn resolve_provider(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<ResolvedProvider, Error> {
let CustomLlmProvider {
model,
custom_llm_provider: provider,
} = get_custom_llm_provider(model, custom_llm_provider)
.or_else(|| {
custom_llm_provider.map(|provider| CustomLlmProvider {
model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(
"unable to resolve custom_llm_provider for messages request".to_string(),
)
})?;
let provider = provider
.parse()
.map_err(|_| Error::InvalidProvider(provider.to_string()))?;
let resolved = resolve_llm_provider(model, custom_llm_provider, "messages")?;
let provider = messages_provider(resolved.provider)
.ok_or_else(|| Error::InvalidProvider(<&str>::from(resolved.provider).to_string()))?;
Ok(ResolvedProvider {
model: model.to_string(),
model: resolved.model.to_string(),
provider,
})
}

View file

@ -1,26 +1,15 @@
use std::{
convert::Infallible,
sync::{Arc, Mutex},
};
use std::convert::Infallible;
use bytes::Bytes;
use futures_util::TryStreamExt;
use litellm_host::{
host::{Demand, Host},
machine::{CallMachine, HostChannel, MachineFault},
call::{HostedCompletion, HostedMachine, hosted_call},
protocol::Protocol,
};
use litellm_http::{Client, ClientVariant, HttpClientConfig};
use litellm_secrets::source::SecretSource;
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare};
use super::{Error, MessagesCall};
pub enum MessagesOutput {
Message(Box<AnthropicMessagesResponse>),
/// Every chunk already reached the host through `Deliver`.
Streamed,
}
pub type MessagesOutput = HostedCompletion<Box<AnthropicMessagesResponse>>;
/// The upstream response as the caller sees it at stream hand-off, before any chunk.
pub struct MessagesStreamHead {
@ -30,91 +19,20 @@ pub struct MessagesStreamHead {
pub struct Messages;
impl Protocol for Messages {
type Response = MessagesOutput;
type Response = Box<AnthropicMessagesResponse>;
type Error = Error;
type Projection = MessagesCall;
type Op = Infallible;
type Request = MessagesCall;
type HostCall = Infallible;
type Chunk = Bytes;
type StreamHead = MessagesStreamHead;
}
impl From<MachineFault> for Error {
fn from(fault: MachineFault) -> Self {
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "messages host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("messages {message}").into(),
pub type MessagesMachine = HostedMachine<Messages>;
impl super::MessagesRoute {
pub fn machine(self, request: super::MessagesCall) -> MessagesMachine {
hosted_call(request, move |call, _, hooks| async move {
self.run(call, &hooks).await
})
}
}
pub type MessagesHost = HostChannel<Messages>;
pub type MessagesMachine = CallMachine<Messages>;
/// The in-process host for a request already in hand. It answers projection once and
/// observes nothing.
pub struct LocalMessagesHost {
call: Mutex<Option<MessagesCall>>,
}
impl LocalMessagesHost {
pub fn new(call: MessagesCall) -> Self {
Self {
call: Mutex::new(Some(call)),
}
}
}
impl Host<Messages> for LocalMessagesHost {
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(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesMachine, litellm_http::Error> {
let http = resources.pool.client(config, ClientVariant::Provider)?;
let auth = resources.auth.clone();
Ok(CallMachine::new(move |host| {
Box::pin(drive(host, http, auth, secrets))
}))
}
/// The call as its host sees it: projection first, then the same prepare and execute as
/// [`super::messages`], with each chunk of a stream handed over as it arrives.
async fn drive(
host: MessagesHost,
http: Client,
auth: Arc<litellm_auth::AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesOutput, Error> {
let call = host.project().await?;
let request = prepare(call, secrets.as_ref()).await?;
match execute(&http, &auth, request, &host).await? {
MessagesResponse::Message(message) => Ok(MessagesOutput::Message(message)),
MessagesResponse::Stream {
headers,
mut chunks,
} => {
if host.open(MessagesStreamHead { headers }).await? == Demand::Detached {
return Ok(MessagesOutput::Streamed);
}
while let Some(chunk) = chunks.try_next().await? {
if host.deliver(chunk).await? == Demand::Detached {
break;
}
}
Ok(MessagesOutput::Streamed)
}
}
}

View file

@ -1,8 +1,8 @@
use std::time::Duration;
use bytes::Bytes;
use futures_util::stream::BoxStream;
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
use litellm_host::call::CallOutput;
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
use litellm_types::{
llms::anthropic_messages::{
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
@ -30,24 +30,16 @@ pub fn messages_body(body: Map<String, Value>) -> Result<AnthropicMessagesReques
}
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
Error::InvalidRequest(litellm_llms::ErrorDetail::invalid(
"Anthropic messages request",
err,
))
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}").into())
}
pub enum MessagesResponse {
Message(Box<AnthropicMessagesResponse>),
Stream {
headers: Vec<(String, String)>,
chunks: BoxStream<'static, Result<Bytes, Error>>,
},
}
pub type MessagesResponse =
CallOutput<Box<AnthropicMessagesResponse>, super::route::MessagesStreamHead, Bytes, Error>;
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct MessagesShaping {
#[serde(default)]
pub capabilities: MessagesModelCapabilities,
pub capabilities: AnthropicModelCapabilities,
#[serde(default)]
pub drop_params: bool,
#[serde(default)]
@ -84,9 +76,9 @@ mod tests {
#[case::partial_capabilities(
json!({"capabilities": {"supports_reasoning": true}}),
MessagesShaping {
capabilities: MessagesModelCapabilities {
capabilities: AnthropicModelCapabilities {
supports_reasoning: true,
..MessagesModelCapabilities::default()
..AnthropicModelCapabilities::default()
},
..MessagesShaping::default()
},
@ -108,7 +100,7 @@ mod tests {
"additional_drop_params": ["metadata.user_id", "thinking"]
}),
MessagesShaping {
capabilities: MessagesModelCapabilities {
capabilities: AnthropicModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
thinking_always_on: false,

View file

@ -1,15 +1,58 @@
use std::sync::Arc;
use litellm_host::hooks::RouteHooks;
use litellm_llms::base_llm::ocr::{
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
};
use crate::ocr::{
route::{LocalOcrHost, ocr_machine},
types::LiteLLMOcrRequest,
use super::{
handler::perform_ocr_request,
types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest},
};
pub async fn perform(
client: &OcrClient,
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::run::run(ocr_machine(client.clone()), &LocalOcrHost::new(request)).await
#[derive(Clone)]
pub struct OcrRoute {
client: OcrClient,
}
impl OcrRoute {
pub fn new(client: OcrClient) -> Self {
Self { client }
}
pub async fn execute(
&self,
request: LiteLLMOcrRequest,
hooks: &impl RouteHooks<Error>,
) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::lifecycle::observe_unary(hooks.observer(), self.run(request, hooks)).await
}
pub(super) async fn run(
&self,
request: LiteLLMOcrRequest,
hooks: &impl RouteHooks<Error>,
) -> Result<LiteLLMOcrResponse, Error> {
let caller_document = matches!(&request.document, OcrDocumentInput::Document(_));
let prepared = prepare_request_document(request).await?;
let execute: futures_util::future::BoxFuture<'_, Result<LiteLLMOcrResponse, Error>> =
Box::pin(perform_ocr_request(
&self.client,
prepared,
hooks,
caller_document,
));
execute.await
}
}
async fn prepare_request_document(
request: LiteLLMOcrRequest<OcrDocumentInput>,
) -> Result<ResolvedOcrRequest, Error> {
if let OcrDocumentInput::Document(_) = &request.document {
return request.map_document(super::document::prepare_document);
}
tokio::task::spawn_blocking(move || request.map_document(super::document::prepare_document))
.await
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
}

View file

@ -1,6 +1,8 @@
use futures_util::future::BoxFuture;
use litellm_auth::SecretValue;
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
hooks::RouteHooks,
};
use litellm_llms::base_llm::ocr::{
error::Error,
handler::{CallHooks, OcrClient},
@ -8,16 +10,13 @@ use litellm_llms::base_llm::ocr::{
};
use serde_json::Value;
use super::{
arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind,
route::OcrHost,
};
use super::{arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind};
use crate::ocr::types::ResolvedOcrRequest;
pub(crate) async fn perform_ocr_request(
client: &OcrClient,
request: ResolvedOcrRequest,
host: &OcrHost,
host: &impl RouteHooks<Error>,
caller_document: bool,
) -> Result<LiteLLMOcrResponse, Error> {
request.response_format()?;
@ -28,53 +27,48 @@ pub(crate) async fn perform_ocr_request(
.await
.map_err(|error| Error::Secret(std::sync::Arc::new(error)))?;
let request = prepare_request(request, caller_document, client, secrets);
let hooks = OcrCallHooks::new(host.clone(), &request, config);
let hooks = OcrCallHooks::new(host, &request, config);
config.ocr(client, &request, &hooks).await
}
/// Lets provider code reach the host mid-call, filling in the request context only the
/// route knows.
pub(crate) struct OcrCallHooks {
host: OcrHost,
model: String,
custom_llm_provider: &'static str,
optional_params: Value,
secret_fields: Vec<String>,
api_key: Option<SecretValue>,
struct OcrCallHooks<'a, H> {
hooks: &'a H,
context: RequestContext,
}
impl OcrCallHooks {
pub(crate) fn new(host: OcrHost, request: &PreparedOcrRequest, config: OcrConfigKind) -> Self {
impl<'a, H> OcrCallHooks<'a, H> {
fn new(hooks: &'a H, request: &PreparedOcrRequest, config: OcrConfigKind) -> Self {
Self {
host,
model: request.model.clone(),
custom_llm_provider: config.provider().into(),
optional_params: Value::Object(request.optional_params.clone().into()),
secret_fields: request
.optional_params
.keys()
.filter(|name| is_secret_param(name))
.cloned()
.collect(),
api_key: request.connection.api_key.clone(),
hooks,
context: RequestContext {
model: request.model.clone(),
custom_llm_provider: <&str>::from(config.provider()).to_owned(),
optional_params: Value::Object(request.optional_params.clone().into()),
secret_fields: request
.optional_params
.keys()
.filter(|name| is_secret_param(name))
.cloned()
.collect(),
api_key: request.connection.api_key.clone(),
},
}
}
}
impl CallHooks<Error> for OcrCallHooks {
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
let context = RequestContext {
model: self.model.clone(),
custom_llm_provider: self.custom_llm_provider.into(),
optional_params: self.optional_params.clone(),
secret_fields: self.secret_fields.clone(),
api_key: self.api_key.clone(),
};
Box::pin(self.host.before_send(wire, context))
impl<H: RouteHooks<Error>> CallHooks<Error> for OcrCallHooks<'_, H> {
fn before_provider_request(
&self,
wire: WireRequest,
) -> BoxFuture<'_, Result<WireRequest, Error>> {
Box::pin(
self.hooks
.before_provider_request(wire, self.context.clone()),
)
}
fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
Box::pin(self.host.emit(MachineEvent::ResponseReceived {
Box::pin(self.hooks.on_event(MachineEvent::ResponseReceived {
raw: RawResponse {
body: String::from_utf8_lossy(body).into_owned(),
},

View file

@ -1,5 +1,6 @@
pub mod arguments;
pub mod client;
mod client;
pub use client::OcrRoute;
pub mod document;
pub(crate) mod handler;
pub(crate) mod prepare;

View file

@ -5,8 +5,8 @@ use litellm_llms::base_llm::ocr::{
};
use litellm_secrets::source::Secrets;
use super::provider_config::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, ResolvedOcrRequest};
use crate::provider::LlmProviders;
pub(crate) fn prepare_request(
request: ResolvedOcrRequest,
@ -16,15 +16,19 @@ pub(crate) fn prepare_request(
) -> PreparedOcrRequest {
let credentials = request.credentials.clone();
let (preferred_api_key_env, api_base_env) = match request.config.provider() {
OcrProvider::Mistral => (
LlmProviders::Mistral => (
Some("MISTRAL_AZURE_API_KEY"),
Some("MISTRAL_AZURE_API_BASE"),
),
OcrProvider::AzureAi => (None, Some("AZURE_AI_API_BASE")),
OcrProvider::AwsTextract
| OcrProvider::Cohere
| OcrProvider::Reducto
| OcrProvider::VertexAi => (None, None),
LlmProviders::AzureAi => (None, Some("AZURE_AI_API_BASE")),
LlmProviders::Anthropic
| LlmProviders::AwsTextract
| LlmProviders::Bedrock
| LlmProviders::Cohere
| LlmProviders::Openai
| LlmProviders::OpenaiLike
| LlmProviders::Reducto
| LlmProviders::VertexAi => (None, None),
};
let secret = |name: &str| secrets.truthy(name);
let dynamic_api_key = credentials.dynamic_api_key.or_else(|| {
@ -100,7 +104,10 @@ mod tests {
struct NoHooks;
impl CallHooks<Error> for NoHooks {
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
fn before_provider_request(
&self,
wire: WireRequest,
) -> BoxFuture<'_, Result<WireRequest, Error>> {
Box::pin(async move { Ok(wire) })
}
@ -144,6 +151,7 @@ mod tests {
json!({"type": "image_url", "image_url": url})
}
#[rstest::rstest]
#[tokio::test]
async fn cohere_body_keeps_native_document_fields_and_untyped_overrides() {
let request = request(
@ -176,6 +184,7 @@ mod tests {
);
}
#[rstest::rstest]
#[tokio::test]
async fn explicit_null_options_use_defaults_before_http() {
let request = request(
@ -199,6 +208,7 @@ mod tests {
assert!(body.get("req_format").is_none());
}
#[rstest::rstest]
#[tokio::test]
async fn direct_and_vertex_mistral_build_the_same_request_and_share_normalization() {
let options = json!({
@ -276,7 +286,7 @@ mod tests {
pages: Option<Vec<i64>>,
}
#[test]
#[rstest::rstest]
fn parsed_provider_params_separates_known_and_extra_params() {
let arguments: CallArguments = serde_json::from_value(json!({
"pages": [0, 2],

View file

@ -1,3 +1,4 @@
use crate::provider::LlmProviders;
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use litellm_llms::{
aws_textract::ocr::{
@ -24,7 +25,6 @@ use litellm_llms::{
deepseek_transformation::VertexAIDeepSeekOCRConfig, transformation::VertexAiOcrConfig,
},
};
use strum::{EnumString, IntoStaticStr};
macro_rules! with_config {
($kind:expr, $config:ident => $body:expr) => {
@ -93,16 +93,16 @@ pub(crate) enum OcrConfigKind {
}
impl OcrConfigKind {
pub(crate) const fn provider(self) -> OcrProvider {
pub(crate) const fn provider(self) -> LlmProviders {
match self {
Self::AwsTextract | Self::AwsTextractAnalyze => OcrProvider::AwsTextract,
Self::Cohere => OcrProvider::Cohere,
Self::Mistral => OcrProvider::Mistral,
Self::AwsTextract | Self::AwsTextractAnalyze => LlmProviders::AwsTextract,
Self::Cohere => LlmProviders::Cohere,
Self::Mistral => LlmProviders::Mistral,
Self::AzureAi | Self::AzureCohere | Self::AzureDocumentIntelligence => {
OcrProvider::AzureAi
LlmProviders::AzureAi
}
Self::ReductoLegacy | Self::ReductoV3 => OcrProvider::Reducto,
Self::VertexAi | Self::VertexDeepSeek => OcrProvider::VertexAi,
Self::ReductoLegacy | Self::ReductoV3 => LlmProviders::Reducto,
Self::VertexAi | Self::VertexDeepSeek => LlmProviders::VertexAi,
}
}
@ -187,17 +187,6 @@ pub fn passthrough_response(
.map(Some)
}
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
#[strum(serialize_all = "snake_case")]
pub(crate) enum OcrProvider {
AwsTextract,
Cohere,
Mistral,
AzureAi,
Reducto,
VertexAi,
}
pub(crate) fn resolve_provider_config(
model: &str,
custom_llm_provider: Option<&str>,
@ -205,37 +194,45 @@ pub(crate) fn resolve_provider_config(
let provider =
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
model,
custom_llm_provider: OcrProvider::Mistral.into(),
custom_llm_provider: LlmProviders::Mistral.into(),
});
let ocr_provider = provider
let llm_provider = provider
.custom_llm_provider
.parse::<OcrProvider>()
.parse::<LlmProviders>()
.map_err(|_| Error::InvalidProvider(provider.custom_llm_provider.to_string()))?;
let config = match ocr_provider {
OcrProvider::AwsTextract => match TextractOperation::from_model(provider.model)? {
let config = match llm_provider {
LlmProviders::AwsTextract => match TextractOperation::from_model(provider.model)? {
TextractOperation::DetectDocumentText => OcrConfigKind::AwsTextract,
TextractOperation::AnalyzeDocument => OcrConfigKind::AwsTextractAnalyze,
},
OcrProvider::Cohere => OcrConfigKind::Cohere,
OcrProvider::Mistral => OcrConfigKind::Mistral,
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
LlmProviders::Cohere => OcrConfigKind::Cohere,
LlmProviders::Mistral => OcrConfigKind::Mistral,
LlmProviders::AzureAi if is_document_intelligence_model(provider.model) => {
OcrConfigKind::AzureDocumentIntelligence
}
OcrProvider::AzureAi
LlmProviders::AzureAi
if provider.model.to_ascii_lowercase().contains("cohere")
&& provider.model.to_ascii_lowercase().contains("parse") =>
{
OcrConfigKind::AzureCohere
}
OcrProvider::AzureAi => OcrConfigKind::AzureAi,
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
LlmProviders::AzureAi => OcrConfigKind::AzureAi,
LlmProviders::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
OcrConfigKind::ReductoLegacy
}
OcrProvider::Reducto => OcrConfigKind::ReductoV3,
OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
LlmProviders::Reducto => OcrConfigKind::ReductoV3,
LlmProviders::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
OcrConfigKind::VertexDeepSeek
}
OcrProvider::VertexAi => OcrConfigKind::VertexAi,
LlmProviders::VertexAi => OcrConfigKind::VertexAi,
LlmProviders::Anthropic
| LlmProviders::Bedrock
| LlmProviders::Openai
| LlmProviders::OpenaiLike => {
return Err(Error::InvalidProvider(
provider.custom_llm_provider.to_string(),
));
}
};
Ok((provider.model.to_string(), config))
}

View file

@ -1,25 +1,20 @@
use std::sync::{Arc, Mutex};
use litellm_auth::ResolvedCredential;
use litellm_host::{
event::{CallEvent, RequestContext, WireRequest},
host::Reply,
machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol},
call::{CallOutput, HostedMachine, hosted_call},
machine::{HostTokenProvider, TokenProtocol},
protocol::Protocol,
protocol::Reply,
};
use litellm_llms::base_llm::ocr::{
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
};
use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse};
use super::handler::perform_ocr_request;
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest};
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput};
pub enum OcrOp {
AcquireAzureAdToken(Reply<ResolvedCredential>),
}
/// The caller's request as the host projects it.
pub struct OcrProjection {
pub struct OcrCall {
pub request: LiteLLMOcrRequest<OcrDocumentInput>,
/// The caller passed its own Azure AD token provider, which the host keeps.
pub caller_token: bool,
@ -30,8 +25,8 @@ pub struct Ocr;
impl Protocol for Ocr {
type Response = LiteLLMOcrResponse;
type Error = Error;
type Projection = OcrProjection;
type Op = OcrOp;
type Request = OcrCall;
type HostCall = OcrOp;
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
@ -42,122 +37,22 @@ impl TokenProtocol for Ocr {
}
}
pub type OcrHost = HostChannel<Ocr>;
pub type OcrMachine = CallMachine<Ocr>;
pub type OcrMachine = HostedMachine<Ocr>;
/// 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 {
CallMachine::new(move |host| Box::pin(execute(client, host)))
}
async fn execute(client: OcrClient, host: OcrHost) -> Result<LiteLLMOcrResponse, Error> {
let OcrProjection {
request,
caller_token,
} = host.project().await?;
let request = LiteLLMOcrRequest {
azure_ad_token_provider: caller_token
.then(|| HostTokenProvider::handle(host.clone()))
.or(request.azure_ad_token_provider),
..request
};
let caller_document = matches!(request.document, OcrDocumentInput::Document(_));
let request = prepare_request_document(request).await?;
perform_ocr_request(&client, request, &host, caller_document).await
}
async fn prepare_request_document(
request: LiteLLMOcrRequest<OcrDocumentInput>,
) -> Result<ResolvedOcrRequest, Error> {
if let OcrDocumentInput::Document(_) = &request.document {
return request.map_document(super::document::prepare_document);
}
tokio::task::spawn_blocking(move || request.map_document(super::document::prepare_document))
.await
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
}
type BeforeSend =
Box<dyn Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error> + Send + Sync>;
type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
/// The in-process host for a request that is already in hand: the request answers
/// projection, and the optional observer sees and may rewrite the wire request.
pub struct LocalOcrHost {
request: Mutex<Option<LiteLLMOcrRequest<OcrDocumentInput>>>,
before_send: Option<BeforeSend>,
observer: Option<Observer>,
}
impl LocalOcrHost {
pub fn new(request: LiteLLMOcrRequest<OcrDocumentInput>) -> Self {
Self {
request: Mutex::new(Some(request)),
before_send: None,
observer: None,
}
}
pub fn with_before_send(
self,
before_send: impl Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error>
+ Send
+ Sync
+ 'static,
) -> Self {
Self {
before_send: Some(Box::new(before_send)),
..self
}
}
pub fn with_observer(self, observer: impl Fn(&CallEvent) + Send + Sync + 'static) -> Self {
Self {
observer: Some(Box::new(observer)),
..self
}
}
}
impl litellm_host::host::Host<Ocr> for LocalOcrHost {
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::AcquireAzureAdToken(_) => {
Err(Error::Auth(litellm_auth::Error::CredentialAcquisition(
"OCR host has no Azure AD token provider".into(),
)))
}
}
}
async fn before_send(
&self,
wire: WireRequest,
context: &RequestContext,
) -> Result<WireRequest, Error> {
match &self.before_send {
Some(before_send) => before_send(wire, context),
None => Ok(wire),
}
}
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
if let Some(observer) = &self.observer {
observer(event);
}
Ok(())
impl crate::ocr::OcrRoute {
pub fn machine(self, request: OcrCall) -> OcrMachine {
hosted_call(
request,
move |projection: OcrCall, services, hooks| async move {
let request = LiteLLMOcrRequest {
azure_ad_token_provider: projection
.caller_token
.then(|| HostTokenProvider::handle(services))
.or(projection.request.azure_ad_token_provider),
..projection.request
};
self.run(request, &hooks).await.map(CallOutput::Complete)
},
)
}
}

View file

@ -4,6 +4,13 @@ use litellm_http::outbound::OutboundRequest;
use litellm_llms::base_llm::auth::Authenticated;
use serde_json::Value;
pub(crate) async fn send(
request: OutboundRequest,
client: &litellm_http::Client,
) -> Result<reqwest::Response, reqwest::Error> {
request.send(client).await
}
/// Header credentials are already in `headers`; SigV4 is applied here, over the
/// bytes that are sent.
pub(crate) fn outbound_request(

View file

@ -0,0 +1,36 @@
pub use litellm_core_utils::get_llm_provider_logic::LlmProviders;
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use crate::error::RouteError as Error;
#[derive(Debug)]
pub(crate) struct ResolvedProvider<'a> {
pub(crate) model: &'a str,
pub(crate) provider: LlmProviders,
}
pub(crate) fn resolve_llm_provider<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,
route: &'static str,
) -> Result<ResolvedProvider<'a>, Error> {
let CustomLlmProvider {
model,
custom_llm_provider,
} = get_custom_llm_provider(model, custom_llm_provider)
.or_else(|| {
custom_llm_provider.map(|provider| CustomLlmProvider {
model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(format!(
"unable to resolve custom_llm_provider for {route} request"
))
})?;
let provider = custom_llm_provider
.parse()
.map_err(|_| Error::InvalidProvider(custom_llm_provider.to_string()))?;
Ok(ResolvedProvider { model, provider })
}

View file

@ -1,9 +1,7 @@
use std::sync::Arc;
use litellm_auth::AuthServices;
use litellm_http::{HttpClientConfig, HttpClientPool, media::UrlPolicy};
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
use litellm_secrets::source::SecretSource;
use litellm_http::HttpClientPool;
#[derive(Clone)]
pub struct CoreResources {
@ -18,21 +16,4 @@ impl CoreResources {
auth: Arc::new(AuthServices::default()),
}
}
pub fn ocr_client(
&self,
config: &HttpClientConfig,
url_policy: UrlPolicy,
settings: OcrSettings,
secrets: Arc<dyn SecretSource>,
) -> Result<OcrClient, litellm_http::Error> {
OcrClient::new(
&self.pool,
config,
url_policy,
self.auth.clone(),
settings,
secrets,
)
}
}

View file

@ -1,6 +1,4 @@
use litellm_core::audio_transcription::{
Error, audio_transcription, types::AudioTranscriptionRequest,
};
use litellm_core::audio_transcription::{Error, types::AudioTranscriptionRequest};
use rstest::{fixture, rstest};
use serde_json::{Map, Value, json};
use wiremock::ResponseTemplate;
@ -11,19 +9,11 @@ use support::*;
const MODEL: &str = "mistral.voxtral-mini-3b-2507";
async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
audio_transcription(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
request,
)
.await
audio_transcription_route().execute(request).await
}
fn transcript_response(text: &str) -> ResponseTemplate {
json_response(
json!({"output": {"message": {"content": [{"text": text}]}}, "usage": {"inputTokens": 1, "outputTokens": 1}}),
)
json_response(json!({"output": {"message": {"content": [{"text": text}]}}}))
}
fn aws_params(region: &str) -> Map<String, Value> {
@ -70,7 +60,6 @@ async fn bedrock_converse_request_is_signed_for_the_requested_region(
assert_eq!(response, json!({"text": "hello"}));
let sent = only_request(&upstream).await;
assert_eq!(sent.method.as_str(), "POST");
assert_eq!(sent.header("content-type"), Some("application/json"));
assert_eq!(sent.url.path(), format!("/model/{MODEL}/converse"));
let authorization = sent.header("authorization").expect("request is signed");
assert!(
@ -261,61 +250,3 @@ async fn an_unreadable_success_body_is_an_invalid_response(
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
}
#[rstest]
#[tokio::test]
async fn injected_secrets_supply_signing_credentials_and_region(
request: AudioTranscriptionRequest<'static>,
) {
let upstream = upstream([transcript_response("hello")]).await;
let base = upstream.uri();
let secrets = RecordingSecrets::new([
("AWS_ACCESS_KEY_ID", "injected-access-key"),
("AWS_SECRET_ACCESS_KEY", "injected-secret-key"),
("AWS_REGION_NAME", "eu-west-1"),
("AWS_SESSION_TOKEN", "injected-session-token"),
]);
let response = audio_transcription(
&support::resources(),
&http_config(),
&secrets,
AudioTranscriptionRequest {
api_base: Some(&base),
optional_params: Map::new(),
..request
},
)
.await
.unwrap();
assert_eq!(response, json!({"text": "hello"}));
let sent = only_request(&upstream).await;
let authorization = sent.header("authorization").unwrap();
assert!(authorization.contains("Credential=injected-access-key/"));
assert!(authorization.contains("/eu-west-1/bedrock/aws4_request"));
assert_eq!(
sent.header("x-amz-security-token"),
Some("injected-session-token")
);
assert!(!sent.body_text().contains("injected-secret-key"));
}
#[rstest]
#[tokio::test]
async fn secret_resolution_failure_prevents_transcription(
request: AudioTranscriptionRequest<'static>,
) {
let upstream = upstream([transcript_response("hello")]).await;
let base = upstream.uri();
let result = audio_transcription(
&support::resources(),
&http_config(),
&RecordingSecrets::failing(),
AudioTranscriptionRequest {
api_base: Some(&base),
..request
},
)
.await;
assert!(matches!(result, Err(Error::Secret(_))));
assert!(received(&upstream).await.is_empty());
}

View file

@ -1,7 +1,7 @@
use std::time::Duration;
use litellm_core::chat_completions::{
Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest,
Error, chat_completions_decline_reason, types::ChatCompletionsRequest,
};
use litellm_http::transport::Error as TransportError;
use litellm_types::utils::ChatCompletionsResponse;
@ -15,13 +15,7 @@ use support::*;
const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
async fn complete(request: ChatCompletionsRequest<'_>) -> Result<ChatCompletionsResponse, Error> {
chat_completions(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
request,
)
.await
chat_completions_route().execute(request, &()).await
}
fn object(value: Value) -> Map<String, Value> {
@ -329,164 +323,3 @@ async fn a_declined_request_fails_the_call_before_sending(
assert_eq!(error, Error::Unsupported("streaming"));
assert!(received(&upstream).await.is_empty());
}
#[rstest]
#[case::source_key(None, "source-key")]
#[case::explicit_key(Some("explicit-key"), "explicit-key")]
#[tokio::test]
async fn injected_secrets_supply_credentials_and_endpoint(
request: ChatCompletionsRequest<'static>,
#[case] api_key: Option<&'static str>,
#[case] expected_key: &str,
) {
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let secrets = RecordingSecrets::new([
("ANTHROPIC_API_KEY", "source-key"),
("ANTHROPIC_API_BASE", upstream.uri().as_str()),
]);
let response = chat_completions(
&support::resources(),
&http_config(),
&secrets,
ChatCompletionsRequest {
api_key,
api_base: None,
..request
},
)
.await
.unwrap();
let sent = only_request(&upstream).await;
assert_eq!(sent.header("x-api-key"), Some(expected_key));
assert_eq!(sent.url.path(), "/v1/messages");
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
}
#[rstest]
#[case::accepted(false)]
#[case::declined(true)]
#[tokio::test]
async fn secret_failure_stops_before_sending_and_declines_skip_resolution(
request: ChatCompletionsRequest<'static>,
#[case] declined: bool,
) {
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
let secrets = RecordingSecrets::failing();
let result = chat_completions(
&support::resources(),
&http_config(),
&secrets,
ChatCompletionsRequest {
api_base: Some(&base),
optional_params: if declined {
object(json!({"stream": true}))
} else {
request.optional_params.clone()
},
..request
},
)
.await;
if declined {
assert!(matches!(result, Err(Error::Unsupported(_))));
assert!(secrets.requested().is_empty());
} else {
assert!(matches!(result, Err(Error::Secret(_))));
assert!(!secrets.requested().is_empty());
}
assert!(received(&upstream).await.is_empty());
}
#[rstest]
#[case::bearer(true)]
#[case::signed(false)]
#[tokio::test]
async fn bedrock_chat_uses_the_injected_credential_source(
request: ChatCompletionsRequest<'static>,
#[case] bearer: bool,
) {
let upstream = upstream([json_response(json!({
"output": {"message": {"content": [{"text": "hello"}]}},
"usage": {"inputTokens": 1, "outputTokens": 1}
}))])
.await;
let base = upstream.uri();
let secrets = RecordingSecrets::new(
[
("AWS_ACCESS_KEY_ID", "injected-access-key"),
("AWS_SECRET_ACCESS_KEY", "injected-secret-key"),
("AWS_REGION_NAME", "eu-west-1"),
]
.into_iter()
.chain(bearer.then_some(("AWS_BEARER_TOKEN_BEDROCK", "injected-bearer"))),
);
let response = chat_completions(
&support::resources(),
&http_config(),
&secrets,
ChatCompletionsRequest {
model: "test-model",
custom_llm_provider: Some("bedrock"),
api_key: None,
api_base: Some(&base),
optional_params: Map::new(),
..request
},
)
.await
.unwrap();
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
let sent = only_request(&upstream).await;
let authorization = sent.header("authorization").unwrap();
if bearer {
assert_eq!(authorization, "Bearer injected-bearer");
} else {
assert!(authorization.contains("Credential=injected-access-key/"));
assert!(authorization.contains("/eu-west-1/bedrock/aws4_request"));
}
assert!(!sent.body_text().contains("injected-secret-key"));
}
#[rstest]
#[tokio::test]
async fn openai_compatible_chat_resolves_its_injected_endpoint_and_key(
request: ChatCompletionsRequest<'static>,
) {
let upstream = upstream([json_response(json!({
"id": "test-response",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
}))]).await;
let secrets = RecordingSecrets::new([
("OPENAI_LIKE_API_BASE", upstream.uri().as_str()),
("OPENAI_LIKE_API_KEY", "injected-key"),
]);
let response = chat_completions(
&support::resources(),
&http_config(),
&secrets,
ChatCompletionsRequest {
model: "test-model",
custom_llm_provider: Some("openai_like"),
api_key: None,
api_base: None,
..request
},
)
.await
.unwrap();
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
let sent = only_request(&upstream).await;
assert_eq!(sent.url.path(), "/chat/completions");
assert_eq!(sent.header("authorization"), Some("Bearer injected-key"));
}

View file

@ -1,18 +1,15 @@
use std::{convert::Infallible, sync::Mutex};
use std::sync::Mutex;
use litellm_core::messages::route::Messages;
use litellm_host::{
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
host::Host,
};
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
use litellm_host::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
use rstest::rstest;
use super::*;
type Rewrite = Box<dyn Fn(WireRequest) -> Result<WireRequest, Error> + Send + Sync>;
/// Projects like `LocalMessagesHost`, answers `before_send` through `rewrite`, and keeps
/// Projects like `LocalMessagesHost`, answers `before_provider_request` through `rewrite`, and keeps
/// every event the driver emits.
struct RecordingHost {
call: LocalMessagesHost,
@ -50,19 +47,32 @@ impl RecordingHost {
}
}
impl Host<Messages> for RecordingHost {
async fn project(&self) -> Result<MessagesCall, Error> {
self.call.project().await
impl RecordingHost {
pub fn request(&self) -> Result<MessagesCall, Error> {
self.call.request()
}
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
match op {}
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, ()> {
litellm_host::in_process::Host {
services: &(),
hooks: self,
stream: &(),
observer: Some(self),
}
}
}
async fn before_send(
impl litellm_host::lifecycle::CallObserver for RecordingHost {
fn observe(&self, event: litellm_host::event::CallEvent) {
self.events.lock().unwrap().push(event.clone());
}
}
impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protocol>::Error>
for RecordingHost
{
async fn before_provider_request(
&self,
wire: WireRequest,
context: &RequestContext,
context: RequestContext,
) -> Result<WireRequest, Error> {
self.optional_params
.lock()
@ -70,15 +80,24 @@ impl Host<Messages> for RecordingHost {
.push(context.optional_params.clone());
(self.rewrite)(wire)
}
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
self.events.lock().unwrap().push(event.clone());
async fn on_event(
&self,
event: litellm_host::event::MachineEvent,
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
);
Ok(())
}
}
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
litellm_host::in_process::run_hosted(
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
host.runtime(),
)
.await
}
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {
@ -129,8 +148,7 @@ async fn a_before_send_failure_never_sends(call: MessagesCall) {
let error = run_through(&host)
.await
.err()
.expect("the host failure fails the call");
.expect_err("the host failure fails the call");
assert_eq!(error, Error::InvalidRequest("vetoed by the host".into()));
assert!(received(&upstream).await.is_empty());
@ -146,7 +164,7 @@ async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall)
let output = run_through(&host).await.expect("messages call succeeds");
assert!(matches!(output, MessagesOutput::Message(_)));
assert!(matches!(output, MessagesOutput::Complete(_)));
let [emitted] = <[String; 1]>::try_from(host.raw_responses())
.unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len()));
assert_eq!(serde_json::from_str::<Value>(&emitted).unwrap(), raw);
@ -183,9 +201,9 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
let host = RecordingHost::passthrough(authenticated(
MessagesCall {
shaping: MessagesShaping {
capabilities: MessagesModelCapabilities {
capabilities: AnthropicModelCapabilities {
supports_sampling_params: false,
..MessagesModelCapabilities::default()
..AnthropicModelCapabilities::default()
},
drop_params: true,
..MessagesShaping::default()
@ -198,6 +216,6 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
run_through(&host).await.expect("messages call succeeds");
let [optional_params] = <[Value; 1]>::try_from(host.optional_params.into_inner().unwrap())
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
.unwrap_or_else(|seen| panic!("before_provider_request runs once, saw {}", seen.len()));
assert_eq!(optional_params, json!({"max_tokens": 16}));
}

View file

@ -1,8 +1,11 @@
use std::{sync::Arc, time::Duration};
use std::{
sync::{Arc, Mutex},
time::Duration,
};
use litellm_core::messages::{
Error, MessagesCall, MessagesShaping,
route::{LocalMessagesHost, MessagesMachine, MessagesOutput, messages_machine},
route::{Messages, MessagesMachine, MessagesOutput},
};
use litellm_http::{HttpSettings, Resolution};
use litellm_secrets::source::SecretSource;
@ -95,16 +98,16 @@ fn headers<'a>(pairs: impl IntoIterator<Item = (&'a str, &'a str)>) -> Option<Ma
)
}
fn machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
messages_machine(&support::resources(), &http_config(), secrets)
.expect("default HTTP settings build a client")
fn machine(secrets: Arc<dyn SecretSource>) -> impl FnOnce(MessagesCall) -> MessagesMachine {
move |request| messages_route(secrets).machine(request)
}
async fn run_with(
secrets: Arc<RecordingSecrets>,
call: MessagesCall,
) -> Result<MessagesOutput, Error> {
litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await
let host = LocalMessagesHost::new(call);
litellm_host::in_process::run_hosted(machine(secrets)(host.request()?), host.runtime()).await
}
/// Runs the route with a secret source that knows nothing, so no environment leaks in.
@ -114,7 +117,67 @@ async fn run(call: MessagesCall) -> Result<MessagesOutput, Error> {
async fn run_message(call: MessagesCall) -> AnthropicMessagesResponse {
match run(call).await.expect("messages call succeeds") {
MessagesOutput::Message(message) => *message,
MessagesOutput::Streamed => panic!("a non-streaming call returned a stream"),
MessagesOutput::Complete(message) => *message,
MessagesOutput::StreamEnded | MessagesOutput::Detached => {
panic!("a non-streaming call returned a stream")
}
}
}
struct LocalMessagesHost {
call: Mutex<Option<MessagesCall>>,
}
impl LocalMessagesHost {
fn new(call: MessagesCall) -> Self {
Self {
call: Mutex::new(Some(call)),
}
}
}
impl LocalMessagesHost {
pub fn request(&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()))
}
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, ()> {
litellm_host::in_process::Host {
services: &(),
hooks: self,
stream: &(),
observer: Some(self),
}
}
}
impl litellm_host::lifecycle::CallObserver for LocalMessagesHost {
fn observe(&self, _: litellm_host::event::CallEvent) {}
}
impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protocol>::Error>
for LocalMessagesHost
{
async fn before_provider_request(
&self,
wire: litellm_host::event::WireRequest,
_: litellm_host::event::RequestContext,
) -> Result<
litellm_host::event::WireRequest,
<Messages as litellm_host::protocol::Protocol>::Error,
> {
Ok(wire)
}
async fn on_event(
&self,
event: litellm_host::event::MachineEvent,
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
);
Ok(())
}
}

View file

@ -87,8 +87,7 @@ async fn a_call_without_credentials_fails_before_sending(
..call
})
.await
.err()
.expect("a call without credentials fails");
.expect_err("a call without credentials fails");
assert!(
matches!(
@ -158,8 +157,7 @@ async fn unsupported_providers_are_rejected_before_sending(
..with_model(call, model)
})
.await
.err()
.expect("unsupported provider errors");
.expect_err("unsupported provider errors");
assert_eq!(error, Error::InvalidProvider(reported.into()));
}
@ -405,8 +403,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
let error = run(shaped(false))
.await
.err()
.expect("an unsupported param is rejected without drop_params");
.expect_err("an unsupported param is rejected without drop_params");
assert!(
matches!(&error, Error::InvalidRequest(message) if message.to_string().contains(rejected_as)),
"{error:?}"
@ -601,8 +598,7 @@ async fn an_invalid_request_fails_before_sending(call: MessagesCall, #[case] fie
fields,
))
.await
.err()
.expect("the request is rejected");
.expect_err("the request is rejected");
assert!(error.is_request(), "{error:?}");
assert!(received(&upstream).await.is_empty());
@ -705,7 +701,7 @@ async fn provider_validation_runs_before_caller_parameter_removal(
json!({"metadata": {"user_id": 7}}),
))
.await;
let error = result.err().expect("metadata is validated before removal");
let error = result.expect_err("metadata is validated before removal");
assert!(
error
.to_string()

View file

@ -1,12 +1,65 @@
use litellm_core::{
Phase,
messages::{MessagesResponse, messages, messages_body},
messages::{MessagesResponse, messages_body},
};
use litellm_http::transport::Error as TransportError;
use rstest::rstest;
use super::*;
#[rstest]
#[case::without_hooks(false)]
#[case::with_hooks(true)]
#[tokio::test]
async fn calls_defer_execution_until_polled(call: MessagesCall, #[case] with_hooks: bool) {
use futures_util::future::BoxFuture;
use litellm_host::event::CallEvent;
let upstream = upstream([message_response()]).await;
let secrets = Arc::new(RecordingSecrets::new([("ANTHROPIC_API_KEY", "test-key")]));
let route = messages_route(secrets.clone());
let host = RecordingCall::<Messages>::new(MessagesCall {
api_base: Some(upstream.uri()),
..call
});
let request = host.request().unwrap();
let future: BoxFuture<'_, Result<MessagesResponse, Error>> = if with_hooks {
Box::pin(route.execute(request, &host))
} else {
Box::pin(route.execute(request, &()))
};
assert!(secrets.requested().is_empty());
assert!(host.events.0.lock().unwrap().is_empty());
assert!(received(&upstream).await.is_empty());
let MessagesResponse::Complete(response) = future.await.unwrap() else {
panic!("expected a completed message");
};
assert_eq!(
response.content,
message_body()["content"].as_array().unwrap().as_slice()
);
assert!(secrets.requested().contains(&"ANTHROPIC_API_KEY".into()));
let sent = only_request(&upstream).await;
assert_eq!(sent.header("x-api-key"), Some("test-key"));
assert_eq!(sent.header("x-hook"), with_hooks.then_some("called"));
let events = host.events.0.lock().unwrap();
if with_hooks {
assert!(matches!(
&events[..],
[
CallEvent::Started { .. },
CallEvent::Machine(_),
CallEvent::Succeeded { .. }
]
));
} else {
assert!(events.is_empty());
}
}
#[rstest]
#[case::anthropic("anthropic")]
#[case::azure_ai("azure_ai")]
@ -75,8 +128,7 @@ async fn a_json_error_envelope_is_kept_verbatim(call: MessagesCall) {
..call
})
.await
.err()
.expect("upstream error propagates");
.expect_err("upstream error propagates");
let Error::Transport(TransportError::Http { status, body }) = error else {
panic!("{error:?}");
@ -97,8 +149,7 @@ async fn a_long_error_body_is_truncated_at_the_documented_cap(call: MessagesCall
..call
})
.await
.err()
.expect("upstream error propagates");
.expect_err("upstream error propagates");
assert_eq!(
error,
@ -126,8 +177,7 @@ async fn an_upstream_error_keeps_its_status_and_body(call: MessagesCall, #[case]
..call
})
.await
.err()
.expect("upstream error propagates");
.expect_err("upstream error propagates");
assert_eq!(
error,
@ -154,8 +204,7 @@ async fn an_unreadable_success_body_is_an_invalid_response(
..call
})
.await
.err()
.expect("an unreadable body fails");
.expect_err("an unreadable body fails");
assert_eq!(error.phase(), Phase::AfterSend, "{error:?}");
}
@ -172,8 +221,7 @@ async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) {
..call
})
.await
.err()
.expect("the call times out");
.expect_err("the call times out");
assert!(matches!(error, Error::Transport(_)), "{error:?}");
}
@ -188,20 +236,24 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes
..HttpSettings::default()
};
let response = messages(
&support::resources(),
&Resolution::from(&settings).config,
&RecordingSecrets::empty(),
let resources = support::resources();
let response = litellm_core::messages::MessagesRoute::new(
provider_http(&resources, &Resolution::from(&settings).config),
resources.auth,
no_secrets(),
)
.execute(
MessagesCall {
api_key: Some("sk-ant".into()),
api_base: Some(base),
..call
},
&(),
)
.await
.expect("messages request succeeds");
let MessagesResponse::Message(message) = response else {
let MessagesResponse::Complete(message) = response else {
panic!("a non-streaming request returns a message");
};
assert_eq!(message.id, "msg_1");

View file

@ -43,7 +43,7 @@ async fn the_credential_and_base_come_from_the_secret_source(
.await
.expect("messages call succeeds");
assert!(matches!(output, MessagesOutput::Message(_)));
assert!(matches!(output, MessagesOutput::Complete(_)));
let request = only_request(&upstream).await;
assert_eq!(request.url.path(), path);
assert_eq!(request.header("x-api-key"), Some("sk-from-manager"));
@ -90,8 +90,7 @@ async fn a_secret_manager_failure_fails_the_call_before_sending(call: MessagesCa
},
)
.await
.err()
.expect("a secret manager failure fails the call");
.expect_err("a secret manager failure fails the call");
assert!(
matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)),
@ -193,8 +192,7 @@ async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall)
},
)
.await
.err()
.expect("azure needs a base");
.expect_err("azure needs a base");
assert_eq!(
error,

View file

@ -1,16 +1,12 @@
use std::{
convert::Infallible,
sync::{Mutex, mpsc},
};
use std::sync::Mutex;
use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt};
use litellm_core::messages::{
MessagesResponse, messages,
MessagesResponse,
route::{Messages, MessagesStreamHead},
};
use litellm_host::host::{Demand, Host};
use litellm_tracing::{Logger, Metadata, Record, Sink};
use litellm_host::protocol::Demand;
use rstest::rstest;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
@ -32,20 +28,6 @@ enum Seen {
Deliver(Bytes),
}
struct TraceSink(mpsc::Sender<(String, Value)>);
impl Sink for TraceSink {
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
metadata.target().starts_with("litellm_core::messages")
}
fn emit(&self, record: &Record) {
self.0
.send((record.message.clone(), Value::Object(record.fields.clone())))
.unwrap();
}
}
/// Projects like `LocalMessagesHost`, records every stream op in the order the route
/// performs it, and detaches after `detach_after` ops.
struct RecordingStreamHost {
@ -73,23 +55,55 @@ impl RecordingStreamHost {
}
}
impl Host<Messages> for RecordingStreamHost {
async fn project(&self) -> Result<MessagesCall, Error> {
self.call.project().await
impl RecordingStreamHost {
pub fn request(&self) -> Result<MessagesCall, Error> {
self.call.request()
}
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
match op {}
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, Self> {
litellm_host::in_process::Host {
services: &(),
hooks: self,
stream: self,
observer: Some(self),
}
}
}
async fn open(&self, head: MessagesStreamHead) -> Result<Demand, Error> {
impl litellm_host::in_process::StreamConsumer<Messages> for RecordingStreamHost {
async fn open_stream(&self, head: MessagesStreamHead) -> Result<Demand, Error> {
Ok(self.record(Seen::Open(head.headers)))
}
async fn deliver(&self, chunk: Bytes) -> Result<Demand, Error> {
async fn send_chunk(&self, chunk: Bytes) -> Result<Demand, Error> {
Ok(self.record(Seen::Deliver(chunk)))
}
}
impl litellm_host::lifecycle::CallObserver for RecordingStreamHost {
fn observe(&self, _: litellm_host::event::CallEvent) {}
}
impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protocol>::Error>
for RecordingStreamHost
{
async fn before_provider_request(
&self,
wire: litellm_host::event::WireRequest,
_: litellm_host::event::RequestContext,
) -> Result<
litellm_host::event::WireRequest,
<Messages as litellm_host::protocol::Protocol>::Error,
> {
Ok(wire)
}
async fn on_event(
&self,
event: litellm_host::event::MachineEvent,
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
);
Ok(())
}
}
fn streaming(call: MessagesCall, api_base: String) -> MessagesCall {
MessagesCall {
@ -107,7 +121,11 @@ fn sse_response() -> ResponseTemplate {
}
async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Error> {
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
litellm_host::in_process::run_hosted(
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
host.runtime(),
)
.await
}
#[rstest]
@ -118,7 +136,7 @@ async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: Me
let outcome = stream_through(&host).await.expect("streamed call succeeds");
assert!(matches!(outcome, MessagesOutput::Streamed));
assert!(matches!(outcome, MessagesOutput::StreamEnded));
let seen = host.seen.into_inner().unwrap();
let [Seen::Open(headers), chunks @ ..] = seen.as_slice() else {
panic!("the stream opens before any chunk is delivered");
@ -143,37 +161,6 @@ async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: Me
assert_eq!(delivered, SSE_BODY.as_bytes());
}
#[rstest]
#[tokio::test]
async fn debug_trace_keeps_provider_input_and_every_stream_chunk(call: MessagesCall) {
let upstream = upstream([sse_response()]).await;
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
let (sender, receiver) = mpsc::channel();
Logger::new(TraceSink(sender))
.instrument(stream_through(&host))
.await
.unwrap();
let records: Vec<(String, Value)> = receiver.try_iter().collect();
let request = records
.iter()
.find(|(message, _)| message == "provider request")
.unwrap();
let body: Value = serde_json::from_str(request.1["body"].as_str().unwrap()).unwrap();
assert_eq!(body["messages"][0]["content"], "hi");
assert_eq!(request.1["stream"], true);
let chunks: String = records
.iter()
.filter(|(message, fields)| {
message == "stream chunk" && fields["stage"] == "provider_response"
})
.map(|(_, fields)| fields["chunk"].as_str().unwrap())
.collect();
assert_eq!(chunks, SSE_BODY);
assert!(!format!("{records:?}").contains("sk-ant"));
}
#[rstest]
#[case::at_open(1)]
#[case::after_the_first_chunk(2)]
@ -186,7 +173,7 @@ async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] det
.await
.expect("a detached stream still completes");
assert!(matches!(outcome, MessagesOutput::Streamed));
assert!(matches!(outcome, MessagesOutput::Detached));
assert_eq!(host.seen.into_inner().unwrap().len(), detach_after);
}
@ -196,6 +183,7 @@ async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] det
status_response(429, json!({"type": "error", "error": {"type": "rate_limit_error", "message": "slow down"}})),
r#"{"type":"error","error":{"type":"rate_limit_error","message":"slow down"}}"#
)]
#[rstest::rstest]
#[tokio::test]
async fn an_upstream_error_fails_the_call_without_opening_the_stream(
call: MessagesCall,
@ -207,8 +195,7 @@ async fn an_upstream_error_fails_the_call_without_opening_the_stream(
let error = stream_through(&host)
.await
.err()
.expect("upstream error propagates");
.expect_err("upstream error propagates");
assert_eq!(
error,
@ -280,8 +267,7 @@ async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
let error = tokio::time::timeout(Duration::from_secs(5), stream_through(&host))
.await
.expect("the stalled stream gives up within the timeout")
.err()
.expect("a stalled body fails the call");
.expect_err("a stalled body fails the call");
assert!(matches!(error, Error::Transport(_)), "{error:?}");
let seen = host.seen.into_inner().unwrap();
@ -305,23 +291,22 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte(
#[case] provider: &str,
) {
let upstream = upstream([sse_response()]).await;
let response = messages(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
MessagesCall {
custom_llm_provider: Some(provider.into()),
..streaming(call, upstream.uri())
},
)
.await
.unwrap();
let response = messages_route(no_secrets())
.execute(
MessagesCall {
custom_llm_provider: Some(provider.into()),
..streaming(call, upstream.uri())
},
&(),
)
.await
.unwrap();
let MessagesResponse::Stream { headers, chunks } = response else {
let MessagesResponse::Stream { head, chunks } = response else {
panic!("a streaming request returns a stream");
};
for (name, value) in UPSTREAM_HEADERS {
assert!(headers.contains(&(name.into(), value.into())));
assert!(head.headers.contains(&(name.into(), value.into())));
}
let delivered = chunks.try_collect::<Vec<_>>().await.unwrap().concat();
assert_eq!(delivered, SSE_BODY.as_bytes());
@ -332,15 +317,11 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte(
#[tokio::test]
async fn the_sdk_returns_http_errors_before_opening_a_stream(call: MessagesCall) {
let upstream = upstream([ResponseTemplate::new(429).set_body_string("slow down")]).await;
let error = messages(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
streaming(call, upstream.uri()),
)
.await
.err()
.expect("upstream failure is returned by messages()");
let error = messages_route(no_secrets())
.execute(streaming(call, upstream.uri()), &())
.await
.err()
.expect("upstream failure is returned by messages()");
assert_eq!(
error,
@ -362,14 +343,12 @@ async fn dropping_the_sdk_stream_closes_the_unfinished_upstream(
let (base, connection) = stalling_upstream().await;
let response = tokio::time::timeout(
Duration::from_secs(5),
messages(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
messages_route(no_secrets()).execute(
MessagesCall {
timeout: Some(Duration::from_secs(30)),
..streaming(call, base)
},
&(),
),
)
.await
@ -399,17 +378,16 @@ async fn dropping_the_sdk_stream_closes_the_unfinished_upstream(
#[tokio::test]
async fn the_sdk_yields_a_body_error_once_after_delivered_chunks(call: MessagesCall) {
let (base, connection) = stalling_upstream().await;
let response = messages(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
MessagesCall {
timeout: Some(Duration::from_millis(300)),
..streaming(call, base)
},
)
.await
.unwrap();
let response = messages_route(no_secrets())
.execute(
MessagesCall {
timeout: Some(Duration::from_millis(300)),
..streaming(call, base)
},
&(),
)
.await
.unwrap();
let MessagesResponse::Stream { mut chunks, .. } = response else {
panic!("a streaming request returns a stream");
@ -445,7 +423,7 @@ async fn a_host_on_anthropic_sse_is_relayed_byte_for_byte(call: MessagesCall) {
let outcome = stream_through(&host).await.expect("azure streams");
assert!(matches!(outcome, MessagesOutput::Streamed));
assert!(matches!(outcome, MessagesOutput::StreamEnded));
let seen = host.seen.into_inner().unwrap();
let delivered: Vec<u8> = seen
.iter()

View file

@ -183,16 +183,16 @@ async fn client_settings_choose_the_api_version_and_the_inch_to_pixel_dpi() {
"analyzeResult": {"pages": [{"pageNumber": 1, "width": 8.5, "height": 11, "unit": "inch"}]}
}))])
.await;
let client = ocr_client().with_settings(OcrSettings {
let route = ocr_route_with(OcrSettings {
document_intelligence_api_version: "2099-01-01".into(),
document_intelligence_dpi: 72,
..OcrSettings::default()
});
let result =
litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({})))
.await
.unwrap();
let result = route
.execute(read_request(&upstream.uri(), json!({})), &())
.await
.unwrap();
assert_eq!(
only_request(&upstream)
@ -347,14 +347,14 @@ async fn the_polling_deadline_bounds_the_retry_delay() {
],
)
.await;
let client = ocr_client().with_settings(OcrSettings {
let route = ocr_route_with(OcrSettings {
poll_timeout: Duration::from_millis(100),
..OcrSettings::default()
});
let error = tokio::time::timeout(
Duration::from_secs(1),
litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({}))),
route.execute(read_request(&upstream.uri(), json!({})), &()),
)
.await
.expect("the deadline cuts the retry delay short")

View file

@ -45,7 +45,7 @@ impl Route {
}
}
/// What the host does to the wire request in `before_send`.
/// What the host does to the wire request in `before_provider_request`.
#[derive(Clone, Copy, Debug)]
enum Guardrail {
Detached,
@ -53,7 +53,7 @@ enum Guardrail {
}
impl Guardrail {
fn before_send(self, wire: WireRequest) -> WireRequest {
fn before_provider_request(self, wire: WireRequest) -> WireRequest {
let Value::Object(fields) = wire.body else {
return wire;
};
@ -96,8 +96,8 @@ async fn provider_document(route: Route, guardrail: Guardrail) -> Value {
json!({"type": document_type, document_type: format!("{}/scan.png", documents.uri())}),
route.options(),
);
let host =
LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(guardrail.before_send(wire)));
let host = LocalOcrHost::new(request)
.with_before_send(move |wire, _| Ok(guardrail.before_provider_request(wire)));
perform_with(host).await.unwrap();
@ -139,6 +139,7 @@ async fn a_document_replaced_by_the_host_reaches_the_provider(
);
}
#[rstest::rstest]
#[tokio::test]
async fn an_empty_byte_document_fails_before_sending() {
let upstream = upstream([pages_response()]).await;
@ -156,6 +157,7 @@ async fn an_empty_byte_document_fails_before_sending() {
assert!(received(&upstream).await.is_empty());
}
#[rstest::rstest]
#[tokio::test]
async fn a_missing_path_document_fails_before_sending() {
let upstream = upstream([pages_response()]).await;
@ -180,3 +182,51 @@ async fn a_missing_path_document_fails_before_sending() {
);
assert!(received(&upstream).await.is_empty());
}
#[rstest]
#[case::blocked(false)]
#[case::allowed(true)]
#[tokio::test]
async fn configured_client_preserves_document_url_policy(#[case] allowed: bool) {
let documents = document_server().await;
let upstream = upstream([pages_response()]).await;
let document_url = format!("{}/scan.png", documents.uri());
let authority = documents.address().to_string();
let route = build_ocr_route(
&resources(),
&http_config(),
litellm_http::media::UrlPolicy {
validate: true,
allowed_hosts: allowed.then_some(authority).into_iter().collect(),
},
Default::default(),
no_secrets(),
);
let host = LocalOcrHost::new(ocr_request_with_document(
"azure_ai/model",
&upstream.uri(),
json!({"type": "document_url", "document_url": document_url}),
json!({}),
));
let result = litellm_host::in_process::run_hosted(
route.machine(host.request().unwrap()),
host.runtime(),
)
.await;
if !allowed {
assert!(matches!(result, Err(Error::BlockedDocumentUrl)));
assert!(received(&documents).await.is_empty());
assert!(received(&upstream).await.is_empty());
return;
}
result.unwrap();
assert_eq!(only_request(&documents).await.url.path(), "/scan.png");
assert_eq!(
only_request(&upstream).await.json()["document"]["document_url"],
format!(
"data:image/png;base64,{}",
base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT)
)
);
}

View file

@ -1,13 +1,10 @@
use std::sync::{Arc, Mutex};
use litellm_core::ocr::{
route::{Ocr, OcrOp, OcrProjection, ocr_machine},
route::{Ocr, OcrCall, OcrOp},
types::OcrDocumentInput,
};
use litellm_host::{
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
host::Host,
};
use litellm_host::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
use rstest::rstest;
use super::*;
@ -18,6 +15,7 @@ pub(crate) fn event_name(event: &CallEvent) -> &'static str {
CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response",
CallEvent::Succeeded { .. } => "success",
CallEvent::Failed { .. } => "failure",
CallEvent::Cancelled { .. } => "cancelled",
}
}
@ -29,7 +27,10 @@ fn recording_host(
let before_send_events = events.clone();
LocalOcrHost::new(request)
.with_before_send(move |wire, _| {
before_send_events.lock().unwrap().push("before_send");
before_send_events
.lock()
.unwrap()
.push("before_provider_request");
match block {
true => Err(Error::InvalidRequest("blocked".into())),
false => Ok(wire),
@ -38,6 +39,7 @@ fn recording_host(
.with_observer(move |event| events.lock().unwrap().push(event_name(event)))
}
#[rstest::rstest]
#[tokio::test]
async fn hooks_run_in_order_and_one_success_is_emitted() {
let upstream = upstream([pages_response()]).await;
@ -53,11 +55,12 @@ async fn hooks_run_in_order_and_one_success_is_emitted() {
assert_eq!(
*events.lock().unwrap(),
["started", "before_send", "response", "success"]
["started", "before_provider_request", "response", "success"]
);
assert_eq!(received(&upstream).await.len(), 1);
}
#[rstest::rstest]
#[tokio::test]
async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() {
let upstream = upstream([pages_response()]).await;
@ -77,11 +80,12 @@ async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() {
);
assert_eq!(
*events.lock().unwrap(),
["started", "before_send", "failure"]
["started", "before_provider_request", "failure"]
);
assert!(received(&upstream).await.is_empty());
}
#[rstest::rstest]
#[tokio::test]
async fn an_upstream_failure_emits_one_terminal_failure() {
let upstream = upstream([status_response(500, json!({"error": "failed"}))]).await;
@ -97,11 +101,12 @@ async fn an_upstream_failure_emits_one_terminal_failure() {
assert!(result.is_err());
assert_eq!(
*events.lock().unwrap(),
["started", "before_send", "failure"]
["started", "before_provider_request", "failure"]
);
assert_eq!(received(&upstream).await.len(), 1);
}
#[rstest::rstest]
#[tokio::test]
async fn an_invalid_provider_response_is_observed_before_normalization_fails() {
let upstream = upstream([json_response(json!({"pages": "invalid"}))]).await;
@ -120,6 +125,7 @@ async fn an_invalid_provider_response_is_observed_before_normalization_fails() {
assert_eq!(*observed.lock().unwrap(), [r#"{"pages":"invalid"}"#]);
}
#[rstest::rstest]
#[tokio::test]
async fn headers_returned_by_before_send_are_sent() {
let upstream = upstream([pages_response()]).await;
@ -147,9 +153,10 @@ async fn before_send_context(request: LiteLLMOcrRequest) -> (WireRequest, Reques
});
perform_with(host).await.unwrap();
let context = observed.lock().unwrap().take();
context.expect("before_send ran")
context.expect("before_provider_request ran")
}
#[rstest::rstest]
#[tokio::test]
async fn before_send_sees_the_route_its_params_and_the_body() {
let upstream = upstream([pages_response()]).await;
@ -187,22 +194,31 @@ async fn before_send_names_the_secret_params(#[case] options: Value, #[case] sec
assert_eq!(context.secret_fields, secrets);
}
/// Hands the route a caller-owned Azure token and rewrites the bearer in `before_send`.
/// Hands the route a caller-owned Azure token and rewrites the bearer in `before_provider_request`.
struct CallerTokenHost {
request: Mutex<Option<LiteLLMOcrRequest>>,
trace: Mutex<Vec<String>>,
}
impl Host<Ocr> for CallerTokenHost {
async fn project(&self) -> Result<OcrProjection, Error> {
impl CallerTokenHost {
pub fn request(&self) -> Result<OcrCall, Error> {
self.trace.lock().unwrap().push("project".into());
Ok(OcrProjection {
Ok(OcrCall {
request: self.request.lock().unwrap().take().unwrap(),
caller_token: true,
})
}
async fn custom_op(&self, op: OcrOp) -> Result<(), Error> {
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, Self, Self, ()> {
litellm_host::in_process::Host {
services: self,
hooks: self,
stream: &(),
observer: Some(self),
}
}
}
impl litellm_host::services::HostCallHandler<Ocr> for CallerTokenHost {
async fn handle_host_call(&self, op: OcrOp) -> Result<(), Error> {
match op {
OcrOp::AcquireAzureAdToken(reply) => {
self.trace.lock().unwrap().push("token".into());
@ -213,11 +229,18 @@ impl Host<Ocr> for CallerTokenHost {
}
}
}
}
async fn before_send(
impl litellm_host::lifecycle::CallObserver for CallerTokenHost {
fn observe(&self, _: litellm_host::event::CallEvent) {}
}
impl litellm_host::hooks::RouteHooks<<Ocr as litellm_host::protocol::Protocol>::Error>
for CallerTokenHost
{
async fn before_provider_request(
&self,
wire: WireRequest,
_: &RequestContext,
_: RequestContext,
) -> Result<WireRequest, Error> {
let is_authorization = |name: &str| name.eq_ignore_ascii_case("authorization");
let authorization = wire
@ -229,7 +252,7 @@ impl Host<Ocr> for CallerTokenHost {
self.trace
.lock()
.unwrap()
.push(format!("before_send:{authorization}"));
.push(format!("before_provider_request:{authorization}"));
let headers = wire
.headers
.into_iter()
@ -240,8 +263,19 @@ impl Host<Ocr> for CallerTokenHost {
.collect();
Ok(WireRequest { headers, ..wire })
}
async fn on_event(
&self,
event: litellm_host::event::MachineEvent,
) -> Result<(), <Ocr as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
);
Ok(())
}
}
#[rstest::rstest]
#[tokio::test]
async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_replace_it() {
let upstream = upstream([pages_response()]).await;
@ -254,16 +288,85 @@ async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_
trace: Mutex::new(Vec::new()),
};
litellm_host::run::run(ocr_machine(ocr_client()), &host)
.await
.unwrap();
litellm_host::in_process::run_hosted(
ocr_route().machine(host.request().unwrap()),
host.runtime(),
)
.await
.unwrap();
assert_eq!(
*host.trace.lock().unwrap(),
["project", "token", "before_send:Bearer caller-token"]
[
"project",
"token",
"before_provider_request:Bearer caller-token"
]
);
assert_eq!(
only_request(&upstream).await.header_values("authorization"),
["Bearer edited"]
);
}
#[rstest]
#[tokio::test]
async fn direct_execution_uses_hooks_without_a_machine() {
use litellm_host::{hooks::RouteHooks, lifecycle::CallObserver};
struct Hooks(Arc<super::support::CallEvents>);
impl RouteHooks<Error> for Hooks {
fn observer(&self) -> Option<Arc<dyn CallObserver>> {
Some(self.0.clone())
}
async fn before_provider_request(
&self,
wire: WireRequest,
_: RequestContext,
) -> Result<WireRequest, Error> {
Ok(WireRequest {
headers: wire
.headers
.into_iter()
.chain([("x-direct-hook".into(), "called".into())])
.collect(),
..wire
})
}
async fn on_event(&self, event: MachineEvent) -> Result<(), Error> {
self.0.observe(CallEvent::Machine(event));
Ok(())
}
}
let upstream = upstream([json_response(
json!({"pages":[{"index":0,"markdown":"direct"}]}),
)])
.await;
let events = Arc::new(super::support::CallEvents::default());
let route = ocr_route();
let hooks = Hooks(events.clone());
let builder = route.execute(
ocr_request("mistral/model", &upstream.uri(), json!({})),
&hooks,
);
assert!(events.0.lock().unwrap().is_empty());
assert!(received(&upstream).await.is_empty());
let result = builder.await.unwrap();
assert_eq!(result.pages[0].markdown, "direct");
assert_eq!(
only_request(&upstream).await.header("x-direct-hook"),
Some("called")
);
assert!(matches!(
&events.0.lock().unwrap()[..],
[
CallEvent::Started { .. },
CallEvent::Machine(_),
CallEvent::Succeeded { .. }
]
));
}

View file

@ -1,3 +1,4 @@
use litellm_host::protocol::HookRequest;
use std::{
sync::{
Arc,
@ -7,13 +8,15 @@ use std::{
};
use litellm_core::ocr::{
route::{OcrMachine, OcrOp, OcrProjection},
route::{OcrCall, OcrMachine, OcrOp},
types::OcrDocumentInput,
};
use litellm_host::{
event::{CallEvent, WireRequest},
host::{Host, HostOp},
hooks::RouteHooks,
machine::{HostFailure, Machine, MachineStep},
protocol::Suspension,
services::HostCallHandler,
};
use litellm_llms::base_llm::ocr::transformation::OcrTransportConfig;
use rstest::rstest;
@ -21,7 +24,7 @@ use tokio::{io::AsyncReadExt, net::TcpListener, sync::Notify};
use super::{lifecycle::event_name, *};
/// Drives the machine by hand, answering every op through `host` except `before_send`,
/// Drives the machine by hand, answering every op through `host` except `before_provider_request`,
/// which `intercept` answers so a test can fail or cancel exactly there.
async fn drive_until(
host: &LocalOcrHost,
@ -31,43 +34,39 @@ async fn drive_until(
Vec<&'static str>,
OcrMachine,
) {
let mut machine = ocr_machine(ocr_client());
let mut machine = ocr_route().machine(host.request().unwrap());
let mut ops = Vec::new();
let outcome = loop {
let op = match machine.resume().await {
Ok(MachineStep::Host(op)) => op,
Ok(MachineStep::Complete(response)) => break Ok(response),
Ok(MachineStep::Suspended(op)) => op,
Ok(MachineStep::Complete(response)) => break Ok(completed(response)),
Err(error) => break Err(error),
};
let answer = match op {
HostOp::Project(reply) => {
ops.push("Project");
host.project()
.await
.map(|projection| reply.send(projection))
.map_err(HostFailure::Error)
}
HostOp::Custom(op) => {
Suspension::Stream(stream) => match stream {
litellm_host::protocol::StreamDelivery::Open(head, _) => match head {},
litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {},
},
Suspension::HostCall(op) => {
ops.push(match op {
OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken",
});
host.custom_op(op).await.map_err(HostFailure::Error)
host.handle_host_call(op).await.map_err(HostFailure::Error)
}
HostOp::BeforeSend { wire, reply, .. } => {
Suspension::Hook(HookRequest::BeforeProviderRequest { 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)
Suspension::Hook(HookRequest::Event(event, reply)) => {
ops.push(event_name(&CallEvent::Machine(event.clone())));
host.on_event(event)
.await
.map(|()| reply.send(()))
.map_err(HostFailure::Error)
}
};
if let Err(failure) = answer {
break machine.interrupt(failure).await;
break machine.interrupt(failure).await.map(completed);
}
};
(outcome, ops, machine)
@ -81,10 +80,13 @@ async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, sto
_ = stop.notified() => break,
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::Suspended(Suspension::HostCall(op)) => host.handle_host_call(op).await.unwrap(),
MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest { wire, reply, .. })) => reply.send(*wire),
MachineStep::Suspended(Suspension::Hook(HookRequest::Event(_, reply))) => reply.send(()),
MachineStep::Suspended(Suspension::Stream(stream)) => match stream {
litellm_host::protocol::StreamDelivery::Open(head, _) => match head {},
litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {},
},
MachineStep::Complete(_) => panic!("the stalled call completed"),
}
}
@ -95,6 +97,7 @@ async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, sto
.expect("the call reached the stall point");
}
#[rstest::rstest]
#[tokio::test]
async fn a_hand_driven_machine_performs_the_same_call() {
let upstream = upstream([json_response(json!({
@ -107,13 +110,14 @@ async fn a_hand_driven_machine_performs_the_same_call() {
assert_eq!(outcome.unwrap().pages[0].markdown, "native");
assert_eq!(received(&upstream).await.len(), 1);
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
assert_eq!(ops, ["BeforeSend", "response"]);
assert!(matches!(
machine.resume().await,
Err(Error::InvalidRequest(_))
));
}
#[rstest::rstest]
#[tokio::test]
async fn a_path_document_is_read_by_core_without_a_host_operation() {
let upstream = upstream([json_response(json!({
@ -135,7 +139,7 @@ async fn a_path_document_is_read_by_core_without_a_host_operation() {
std::fs::remove_dir_all(&dir).unwrap();
assert_eq!(response.unwrap().pages[0].markdown, "path");
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
assert_eq!(ops, ["BeforeSend", "response"]);
assert_eq!(
only_request(&upstream).await.json()["document"]["image_url"],
"data:image/png;base64,YWJj"
@ -143,7 +147,7 @@ async fn a_path_document_is_read_by_core_without_a_host_operation() {
}
#[rstest]
#[case::failed(HostFailure::Error(Error::InvalidRequest("before_send failed".into())), "before_send failed")]
#[case::failed(HostFailure::Error(Error::InvalidRequest("before_provider_request failed".into())), "before_provider_request failed")]
#[case::cancelled(HostFailure::Cancelled(Error::InvalidRequest("cancelled".into())), "cancelled")]
#[tokio::test]
async fn a_before_send_failure_ends_the_call_without_reaching_transport(
@ -159,7 +163,7 @@ async fn a_before_send_failure_ends_the_call_without_reaching_transport(
.lock()
.unwrap()
.take()
.expect("before_send is asked once"))
.expect("before_provider_request is asked once"))
})
.await;
@ -167,27 +171,35 @@ async fn a_before_send_failure_ends_the_call_without_reaching_transport(
matches!(&outcome, Err(Error::InvalidRequest(actual)) if actual == message),
"{outcome:?}"
);
assert_eq!(ops, ["Project", "BeforeSend"]);
assert_eq!(ops, ["BeforeSend"]);
assert!(machine.resume().await.is_err());
assert!(received(&upstream).await.is_empty());
}
#[rstest::rstest]
#[tokio::test]
async fn resuming_before_answering_keeps_the_pending_operation() {
let request = ocr_request("mistral/model", UNREACHABLE_BASE, 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 {
let upstream = upstream([pages_response()]).await;
let request = ocr_request("mistral/model", &upstream.uri(), json!({}));
let mut machine = ocr_route().machine(OcrCall {
request,
caller_token: false,
});
let Ok(MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest {
wire,
reply,
..
}))) = machine.resume().await
else {
panic!("expected the provider request hook");
};
assert!(machine.resume().await.is_err());
reply.send(*wire);
assert!(matches!(
machine.resume().await,
Ok(MachineStep::Host(HostOp::BeforeSend { .. }))
Ok(MachineStep::Suspended(Suspension::Hook(
HookRequest::Event(_, _)
)))
));
}
@ -215,6 +227,7 @@ impl litellm_auth::TokenProvider for PendingToken {
}
}
#[rstest::rstest]
#[tokio::test]
async fn interrupt_drops_provider_captures_before_returning() {
let entered = Arc::new(Notify::new());
@ -231,7 +244,7 @@ async fn interrupt_drops_provider_captures_before_returning() {
},
)));
let host = LocalOcrHost::new(request);
let mut machine = ocr_machine(ocr_client());
let mut machine = ocr_route().machine(host.request().unwrap());
drive_until_notified(&mut machine, &host, &entered).await;
assert!(!dropped.load(Ordering::SeqCst));
@ -248,6 +261,7 @@ async fn interrupt_drops_provider_captures_before_returning() {
);
}
#[rstest::rstest]
#[tokio::test]
async fn interrupting_an_in_flight_provider_request_closes_its_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
@ -266,7 +280,7 @@ async fn interrupting_an_in_flight_provider_request_closes_its_connection() {
while socket.read(&mut buffer).await.unwrap() != 0 {}
});
let host = LocalOcrHost::new(ocr_request("mistral/model", &base, json!({})));
let mut machine = ocr_machine(ocr_client());
let mut machine = ocr_route().machine(host.request().unwrap());
drive_until_notified(&mut machine, &host, &received).await;
let cancelled = Error::InvalidRequest("cancelled".into());

View file

@ -1,16 +1,18 @@
use litellm_core::ocr::{
OcrRoute,
document::prepare_document,
route::{LocalOcrHost, ocr_machine},
types::LiteLLMOcrRequest,
route::{Ocr, OcrCall, OcrOp},
types::{LiteLLMOcrRequest, OcrDocumentInput},
wire::{OcrWireRequest, decode_request},
};
use litellm_http::Client;
use litellm_host::event::{CallEvent, RequestContext, WireRequest};
use litellm_llms::base_llm::ocr::{
error::Error,
handler::OcrClient,
settings::OcrSettings,
transformation::{LiteLLMOcrResponse, OcrDocument},
};
use serde_json::{Map, Value, json};
use std::sync::Mutex;
use wiremock::{MockServer, ResponseTemplate};
#[path = "../support/mod.rs"]
@ -37,16 +39,31 @@ fn object(value: Value) -> Map<String, Value> {
map
}
fn ocr_client() -> OcrClient {
OcrClient::for_test(Client::plain_for_test(), Client::no_redirect_for_test())
fn ocr_route() -> OcrRoute {
ocr_route_with(OcrSettings::default())
}
fn ocr_route_with(settings: OcrSettings) -> OcrRoute {
build_ocr_route(
&resources(),
&http_config(),
litellm_http::media::UrlPolicy {
validate: false,
allowed_hosts: Vec::new(),
},
settings,
no_secrets(),
)
}
async fn perform(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
litellm_core::ocr::client::perform(&ocr_client(), request).await
ocr_route().execute(request, &()).await
}
async fn perform_with(host: LocalOcrHost) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::run::run(ocr_machine(ocr_client()), &host).await
litellm_host::in_process::run_hosted(ocr_route().machine(host.request()?), host.runtime())
.await
.map(completed)
}
fn wire(model: &str, base: &str, document: Value, options: Value) -> OcrWireRequest {
@ -120,3 +137,117 @@ fn accepted(server: &MockServer, body: Value) -> ResponseTemplate {
.insert_header("Operation-Location", format!("{}/operation", server.uri()))
.set_body_json(body)
}
fn completed(
result: litellm_host::call::HostedCompletion<LiteLLMOcrResponse>,
) -> LiteLLMOcrResponse {
match result {
litellm_host::call::HostedCompletion::Complete(response) => response,
other => panic!("unexpected OCR completion: {other:?}"),
}
}
type BeforeSend =
Box<dyn Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error> + Send + Sync>;
type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
struct LocalOcrHost {
request: Mutex<Option<LiteLLMOcrRequest<OcrDocumentInput>>>,
before_provider_request: Option<BeforeSend>,
observer: Option<Observer>,
}
impl LocalOcrHost {
fn new(request: LiteLLMOcrRequest<OcrDocumentInput>) -> Self {
Self {
request: Mutex::new(Some(request)),
before_provider_request: None,
observer: None,
}
}
fn with_before_send(
self,
before_provider_request: impl Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error>
+ Send
+ Sync
+ 'static,
) -> Self {
Self {
before_provider_request: Some(Box::new(before_provider_request)),
..self
}
}
fn with_observer(self, observer: impl Fn(&CallEvent) + Send + Sync + 'static) -> Self {
Self {
observer: Some(Box::new(observer)),
..self
}
}
}
impl LocalOcrHost {
pub fn request(&self) -> Result<OcrCall, Error> {
self.request
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.map(|request| OcrCall {
request,
caller_token: false,
})
.ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into()))
}
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, Self, Self, ()> {
litellm_host::in_process::Host {
services: self,
hooks: self,
stream: &(),
observer: Some(self),
}
}
}
impl litellm_host::services::HostCallHandler<Ocr> for LocalOcrHost {
async fn handle_host_call(&self, op: OcrOp) -> Result<(), Error> {
match op {
OcrOp::AcquireAzureAdToken(_) => {
Err(Error::Auth(litellm_auth::Error::CredentialAcquisition(
"OCR host has no Azure AD token provider".into(),
)))
}
}
}
}
impl litellm_host::lifecycle::CallObserver for LocalOcrHost {
fn observe(&self, event: litellm_host::event::CallEvent) {
if let Some(observer) = &self.observer {
observer(&event);
}
}
}
impl litellm_host::hooks::RouteHooks<<Ocr as litellm_host::protocol::Protocol>::Error>
for LocalOcrHost
{
async fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, Error> {
match &self.before_provider_request {
Some(before_provider_request) => before_provider_request(wire, &context),
None => Ok(wire),
}
}
async fn on_event(
&self,
event: litellm_host::event::MachineEvent,
) -> Result<(), <Ocr as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
);
Ok(())
}
}

View file

@ -1,11 +1,8 @@
use std::sync::Arc;
use litellm_http::{HttpSettings, Resolution, media::UrlPolicy};
use litellm_http::{HttpSettings, Resolution};
use litellm_llms::{
base_llm::ocr::{
settings::OcrSettings,
transformation::{BaseOcrConfig, OCR_RESPONSE_MAX_BYTES},
},
base_llm::ocr::transformation::{BaseOcrConfig, OCR_RESPONSE_MAX_BYTES},
mistral::ocr::transformation::MistralOcrConfig,
};
use rstest::rstest;
@ -156,7 +153,13 @@ async fn missing_credentials_come_from_the_injected_secret_source(
.copied()
.chain([("MISTRAL_AZURE_API_BASE", base.as_str())]),
));
let client = ocr_client().with_secrets(source.clone());
let route = build_ocr_route(
&resources(),
&http_config(),
Default::default(),
Default::default(),
source.clone(),
);
let request = decode_request(OcrWireRequest {
api_key: None,
api_base: None,
@ -169,9 +172,7 @@ async fn missing_credentials_come_from_the_injected_secret_source(
})
.unwrap();
litellm_core::ocr::client::perform(&client, request)
.await
.unwrap();
route.execute(request, &()).await.unwrap();
assert_eq!(source.requested(), MistralOcrConfig.secret_names());
assert_eq!(
@ -188,25 +189,21 @@ async fn the_client_uses_the_injected_http_pool_configuration() {
user_agent: Some("host-owned/1".into()),
..HttpSettings::default()
};
let client = resources()
.ocr_client(
&Resolution::from(&settings).config,
UrlPolicy::default(),
OcrSettings::default(),
Arc::new(
litellm_secrets::source::EnvironmentSecrets::python_compatible(
litellm_http::Client::plain_for_test(),
),
),
)
.unwrap();
let route = build_ocr_route(
&resources(),
&Resolution::from(&settings).config,
Default::default(),
Default::default(),
no_secrets(),
);
litellm_core::ocr::client::perform(
&client,
ocr_request("mistral/model", &upstream.uri(), json!({})),
)
.await
.unwrap();
route
.execute(
ocr_request("mistral/model", &upstream.uri(), json!({})),
&(),
)
.await
.unwrap();
assert_eq!(
only_request(&upstream).await.header("user-agent"),

View file

@ -45,18 +45,19 @@ async fn mistral_is_served_at_the_resolved_project_and_location() {
#[tokio::test]
async fn configured_project_and_location_apply_when_the_call_sets_neither() {
let upstream = upstream([pages_response()]).await;
let client = ocr_client().with_settings(OcrSettings {
let route = ocr_route_with(OcrSettings {
vertex_project: Some("configured-project".into()),
vertex_location: Some("europe-west4".into()),
..OcrSettings::default()
});
litellm_core::ocr::client::perform(
&client,
ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})),
)
.await
.unwrap();
route
.execute(
ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})),
&(),
)
.await
.unwrap();
assert_eq!(
only_request(&upstream).await.url.path(),

View file

@ -10,10 +10,7 @@ use litellm_auth_gcp::{
CredentialSource, VertexAuth, VertexAuthFuture, VertexProviderLoader, VertexTokenSource,
};
use litellm_core::{
ocr::{
client::perform,
wire::{OcrWireRequest, decode_request},
},
ocr::wire::{OcrWireRequest, decode_request},
resources::CoreResources,
};
use litellm_http::{HttpSettings, Resolution};
@ -105,17 +102,16 @@ async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets(
..HttpSettings::default()
})
.config;
let client = owner
.ocr_client(
&http,
Default::default(),
OcrSettings {
vertex_location: Some(location.into()),
..OcrSettings::default()
},
Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])),
)
.unwrap();
let route = support::build_ocr_route(
owner,
&http,
Default::default(),
OcrSettings {
vertex_location: Some(location.into()),
..OcrSettings::default()
},
Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])),
);
let request = decode_request(OcrWireRequest {
model: "vertex_ai/mistral-ocr-maas".into(),
document: json!({"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}),
@ -127,7 +123,7 @@ async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets(
input_sources: Default::default(),
timeout_seconds: Some(5.0),
}).unwrap();
let result = perform(&client, request).await.unwrap();
let result = route.execute(request, &()).await.unwrap();
assert!(!result.pages.is_empty());
}
let requests = upstream.received_requests().await.unwrap();

View file

@ -7,7 +7,8 @@ use std::sync::{Arc, Mutex};
use futures_util::future::BoxFuture;
use litellm_http::{
HttpClientConfig, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver,
ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, Resolution,
media::PublicDnsResolver,
};
use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::Value;
@ -24,6 +25,67 @@ pub fn resources() -> litellm_core::resources::CoreResources {
litellm_core::resources::CoreResources::new(Arc::new(http_pool()))
}
pub fn no_secrets() -> Arc<dyn SecretSource> {
Arc::new(RecordingSecrets::empty())
}
pub fn provider_http(
resources: &litellm_core::resources::CoreResources,
config: &HttpClientConfig,
) -> litellm_http::Client {
resources
.pool
.client(config, ClientVariant::Provider)
.unwrap()
}
pub fn messages_route(secrets: Arc<dyn SecretSource>) -> litellm_core::messages::MessagesRoute {
let resources = resources();
litellm_core::messages::MessagesRoute::new(
provider_http(&resources, &http_config()),
resources.auth,
secrets,
)
}
pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletionsRoute {
let resources = resources();
litellm_core::chat_completions::ChatCompletionsRoute::new(
provider_http(&resources, &http_config()),
resources.auth,
no_secrets(),
)
}
pub fn audio_transcription_route() -> litellm_core::audio_transcription::AudioTranscriptionRoute {
let resources = resources();
litellm_core::audio_transcription::AudioTranscriptionRoute::new(
provider_http(&resources, &http_config()),
resources.auth,
no_secrets(),
)
}
pub fn build_ocr_route(
resources: &litellm_core::resources::CoreResources,
config: &HttpClientConfig,
url_policy: litellm_http::media::UrlPolicy,
settings: litellm_llms::base_llm::ocr::settings::OcrSettings,
secrets: Arc<dyn SecretSource>,
) -> litellm_core::ocr::OcrRoute {
litellm_core::ocr::OcrRoute::new(
litellm_llms::base_llm::ocr::handler::OcrClient::new(
&resources.pool,
config,
url_policy,
resources.auth.clone(),
settings,
secrets,
)
.unwrap(),
)
}
pub fn http_config() -> HttpClientConfig {
Resolution::from(&HttpSettings::default()).config
}
@ -168,3 +230,114 @@ impl SecretSource for RecordingSecrets {
})
}
}
pub struct RecordingCall<P: litellm_host::protocol::Protocol> {
pub request: Mutex<Option<P::Request>>,
pub events: Arc<CallEvents>,
pub chunks: Mutex<Vec<P::Chunk>>,
pub head: Mutex<Option<P::StreamHead>>,
}
#[derive(Default)]
pub struct CallEvents(pub Mutex<Vec<litellm_host::event::CallEvent>>);
impl litellm_host::lifecycle::CallObserver for CallEvents {
fn observe(&self, event: litellm_host::event::CallEvent) {
self.0.lock().unwrap().push(event);
}
}
impl<P: litellm_host::protocol::Protocol> RecordingCall<P> {
pub fn new(request: P::Request) -> Self {
Self {
request: Mutex::new(Some(request)),
events: Arc::new(CallEvents::default()),
chunks: Mutex::new(Vec::new()),
head: Mutex::new(None),
}
}
}
impl<P: litellm_host::protocol::Protocol> litellm_host::hooks::RouteHooks<P::Error>
for RecordingCall<P>
{
fn observer(&self) -> Option<Arc<dyn litellm_host::lifecycle::CallObserver>> {
Some(self.events.clone())
}
async fn before_provider_request(
&self,
wire: litellm_host::event::WireRequest,
_: litellm_host::event::RequestContext,
) -> Result<litellm_host::event::WireRequest, P::Error> {
Ok(litellm_host::event::WireRequest {
headers: wire
.headers
.into_iter()
.chain([("x-hook".into(), "called".into())])
.collect(),
..wire
})
}
async fn on_event(&self, event: litellm_host::event::MachineEvent) -> Result<(), P::Error> {
self.events
.0
.lock()
.unwrap()
.push(litellm_host::event::CallEvent::Machine(event));
Ok(())
}
}
impl<P> RecordingCall<P>
where
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
P::Error: From<litellm_host::machine::MachineFault>,
{
pub fn request(&self) -> Result<P::Request, P::Error> {
self.request
.lock()
.unwrap()
.take()
.ok_or_else(|| litellm_host::machine::MachineFault::Abandoned.into())
}
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, Self> {
litellm_host::in_process::Host {
services: &(),
hooks: self,
stream: self,
observer: Some(self),
}
}
}
impl<P> litellm_host::in_process::StreamConsumer<P> for RecordingCall<P>
where
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
P::Error: From<litellm_host::machine::MachineFault>,
{
async fn open_stream(
&self,
head: P::StreamHead,
) -> Result<litellm_host::protocol::Demand, P::Error> {
*self.head.lock().unwrap() = Some(head);
Ok(litellm_host::protocol::Demand::More)
}
async fn send_chunk(
&self,
chunk: P::Chunk,
) -> Result<litellm_host::protocol::Demand, P::Error> {
self.chunks.lock().unwrap().push(chunk);
Ok(litellm_host::protocol::Demand::More)
}
}
impl<P> litellm_host::lifecycle::CallObserver for RecordingCall<P>
where
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
P::Error: From<litellm_host::machine::MachineFault>,
{
fn observe(&self, event: litellm_host::event::CallEvent) {
self.events.0.lock().unwrap().push(event.clone());
}
}

View file

@ -6,7 +6,7 @@ use axum::{
response::{IntoResponse, Response},
};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_core::audio_transcription::{audio_transcription, types::AudioTranscriptionRequest};
use litellm_core::audio_transcription::types::AudioTranscriptionRequest;
use serde_json::{Value, json};
use crate::{Error, Gateway, request};
@ -38,11 +38,9 @@ async fn handle(gateway: &Gateway, request: Request) -> Result<Value, Error> {
.cloned()
.ok_or_else(|| Error::InvalidBody("audio is required".into()))?,
};
Ok(audio_transcription(
&gateway.resources,
&gateway.http,
gateway.secrets.as_ref(),
AudioTranscriptionRequest {
Ok(gateway
.audio_transcription
.execute(AudioTranscriptionRequest {
model: &deployment.model,
audio,
api_key: deployment.api_key.as_deref(),
@ -54,7 +52,6 @@ async fn handle(gateway: &Gateway, request: Request) -> Result<Value, Error> {
.filter(|(name, _)| !matches!(name.as_str(), "model" | "audio"))
.collect(),
timeout: deployment.timeout,
},
)
.await?)
})
.await?)
}

View file

@ -7,7 +7,7 @@ use axum::{
http::StatusCode,
response::{IntoResponse, Response},
};
use litellm_core::chat_completions::{chat_completions, types::ChatCompletionsRequest};
use litellm_core::chat_completions::types::ChatCompletionsRequest;
use serde_json::{Map, Value};
use crate::{Error, Gateway, request};
@ -61,24 +61,24 @@ async fn handle(gateway: &Gateway, body: Map<String, Value>) -> Result<Response,
.get("messages")
.cloned()
.ok_or_else(|| Error::InvalidBody("messages is required".into()))?;
let response = chat_completions(
&gateway.resources,
&gateway.http,
gateway.secrets.as_ref(),
ChatCompletionsRequest {
model: &deployment.model,
messages,
optional_params: body
.into_iter()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages" | "stream"))
.collect(),
api_key: deployment.api_key.as_deref(),
api_base: deployment.api_base.as_deref(),
custom_llm_provider: deployment.custom_llm_provider.as_deref(),
extra_headers: None,
timeout: deployment.timeout,
},
)
.await?;
let response = gateway
.chat_completions
.execute(
ChatCompletionsRequest {
model: &deployment.model,
messages,
optional_params: body
.into_iter()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages" | "stream"))
.collect(),
api_key: deployment.api_key.as_deref(),
api_base: deployment.api_base.as_deref(),
custom_llm_provider: deployment.custom_llm_provider.as_deref(),
extra_headers: None,
timeout: deployment.timeout,
},
&(),
)
.await?;
Ok(Json(response).into_response())
}

View file

@ -13,20 +13,63 @@ mod request;
use std::sync::Arc;
use axum::{Router, routing::post};
use litellm_core::resources::CoreResources;
use litellm_http::HttpClientConfig;
use litellm_llms::base_llm::ocr::handler::OcrClient;
use litellm_core::{
audio_transcription::AudioTranscriptionRoute, chat_completions::ChatCompletionsRoute,
messages::MessagesRoute, ocr::OcrRoute, resources::CoreResources,
};
use litellm_http::{ClientVariant, HttpClientConfig, media::UrlPolicy};
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
use litellm_secrets::source::SecretSource;
pub use error::Error;
pub use litellm_router::{Deployment, Router as ModelList};
pub struct Gateway {
pub audio_transcription: AudioTranscriptionRoute,
pub chat_completions: ChatCompletionsRoute,
pub messages: MessagesRoute,
pub ocr: OcrRoute,
pub models: ModelList,
pub secrets: Arc<dyn SecretSource>,
pub resources: CoreResources,
pub http: HttpClientConfig,
pub secrets: Arc<dyn SecretSource>,
pub models: ModelList,
pub ocr: OcrClient,
}
impl Gateway {
pub fn new(
resources: CoreResources,
http: HttpClientConfig,
secrets: Arc<dyn SecretSource>,
models: ModelList,
) -> Result<Self, litellm_http::Error> {
let provider = resources.pool.client(&http, ClientVariant::Provider)?;
let auth = resources.auth.clone();
Ok(Self {
audio_transcription: AudioTranscriptionRoute::new(
provider.clone(),
auth.clone(),
secrets.clone(),
),
chat_completions: ChatCompletionsRoute::new(
provider.clone(),
auth.clone(),
secrets.clone(),
),
messages: MessagesRoute::new(provider, auth.clone(), secrets.clone()),
ocr: OcrRoute::new(OcrClient::new(
&resources.pool,
&http,
UrlPolicy::default(),
auth,
OcrSettings::default(),
secrets.clone(),
)?),
models,
secrets,
resources,
http,
})
}
}
pub fn router(gateway: Arc<Gateway>) -> Router {

View file

@ -10,9 +10,7 @@ use axum::{
response::{IntoResponse, Response},
};
use futures_util::{StreamExt, stream::BoxStream};
use litellm_core::messages::{
Error as RouteError, MessagesCall, MessagesResponse, messages, messages_body,
};
use litellm_core::messages::{Error as RouteError, MessagesCall, MessagesResponse, messages_body};
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
use serde_json::{Map, Value};
@ -52,15 +50,8 @@ async fn handle(gateway: &Gateway, headers: &HeaderMap, body: &[u8]) -> Result<R
.get(model_name)
.ok_or_else(|| Error::UnknownModel(model_name.to_owned()))?;
let call = project(deployment, body, headers)?;
match messages(
&gateway.resources,
&gateway.http,
gateway.secrets.as_ref(),
call,
)
.await?
{
MessagesResponse::Message(message) => Ok(Json(message).into_response()),
match gateway.messages.execute(call, &()).await? {
MessagesResponse::Complete(message) => Ok(Json(message).into_response()),
MessagesResponse::Stream { chunks, .. } => Ok(stream(chunks)),
}
}

View file

@ -6,10 +6,7 @@ use axum::{
response::{IntoResponse, Response},
};
use litellm_auth::SecretValue;
use litellm_core::ocr::{
client::perform,
types::{LiteLLMOcrRequest, OcrConnectionInputs, OcrDocumentInput},
};
use litellm_core::ocr::types::{LiteLLMOcrRequest, OcrConnectionInputs, OcrDocumentInput};
use litellm_llms::base_llm::ocr::transformation::OcrDocument;
use serde_json::Value;
@ -69,7 +66,7 @@ async fn handle(gateway: &Gateway, request: Request) -> Result<Value, Error> {
..Default::default()
},
)?;
let response = perform(&gateway.ocr, call).await?;
let response = gateway.ocr.execute(call, &()).await?;
match response.provider_native_response {
Some(native) => Ok(Value::Object(native)),
None => Ok(response.into_json()),

View file

@ -10,7 +10,6 @@ use futures_util::future::BoxFuture;
use litellm_core::resources::CoreResources;
use litellm_gateway_inference::{Deployment, Gateway, router};
use litellm_http::{HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver};
use litellm_llms::base_llm::ocr::settings::OcrSettings;
use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::Value;
use tower::ServiceExt;
@ -31,32 +30,26 @@ pub fn app(model: &str, api_base: &str) -> Router {
let http = Resolution::from(&HttpSettings::default()).config;
let secrets = Arc::new(NoSecrets);
let resources = CoreResources::new(pool);
let ocr = resources
.ocr_client(
&http,
Default::default(),
OcrSettings::default(),
secrets.clone(),
router(Arc::new(
Gateway::new(
resources,
http,
secrets,
[(
"public/model".into(),
Deployment {
model: model.into(),
api_base: Some(api_base.into()),
api_key: Some("test-key".into()),
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
)]
.into_iter()
.collect(),
)
.unwrap();
router(Arc::new(Gateway {
resources,
http,
secrets,
ocr,
models: [(
"public/model".into(),
Deployment {
model: model.into(),
api_base: Some(api_base.into()),
api_key: Some("test-key".into()),
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
)]
.into_iter()
.collect(),
}))
.unwrap(),
))
}
pub async fn post(app: Router, path: &str, body: Value) -> Response {

View file

@ -16,7 +16,6 @@ use litellm_gateway_inference::{Gateway, ModelList};
use litellm_http::{
ClientVariant, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver,
};
use litellm_llms::base_llm::ocr::settings::OcrSettings;
use litellm_secrets::source::EnvironmentSecrets;
use litellm_tracing::ByteChunk;
use uuid::Uuid;
@ -27,20 +26,12 @@ pub fn build_inference(config: &Config) -> Result<Arc<Gateway>, litellm_http::Er
let client = pool.client(&http, ClientVariant::Provider)?;
let secrets = Arc::new(EnvironmentSecrets::python_compatible(client));
let resources = CoreResources::new(pool);
let ocr = resources.ocr_client(
&http,
Default::default(),
OcrSettings::default(),
secrets.clone(),
)?;
Ok(Arc::new(Gateway {
Ok(Arc::new(Gateway::new(
resources,
http,
secrets,
models: ModelList::from_model_list(&config.model_list),
ocr,
}))
ModelList::from_model_list(&config.model_list),
)?))
}
pub fn router(inference: Arc<Gateway>, config: &Config) -> Router {

View file

@ -1,10 +1,10 @@
- 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`/`ProtocolHost` 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 `PythonBinding`, `PythonHostCalls`, `PythonCallHooks` and `PythonOwned` 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
- `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, rewritten in place by the route's `Preflight`, not the caller's dict; a protocol host that projects from it inherits the adapter's rewrites (for the legacy adapter: setup, deployment hooks) and the preflight's (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
- The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the lifecycle's business
- `PythonBinding::decode_request` receives the keyword view the hooks' `prepare_arguments` returned, rewritten in place by the route's `Preflight`, not the caller's dict; a binding that decodes from it inherits the lifecycle's rewrites (for the legacy adapter: setup, deployment hooks) and the preflight's (credential inheritance)
- A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the binding's `map_error`; a Python exception raised inside the call, and a failure in `prepare_arguments` or `transform_response`, is raised as is
- A failing `map_error` is raised with the native error's text as its `__context__`, never swallowed
- Use standard PyO3 ownership and conversion APIs
- Prefer `Bound<'py, T>` for attached operations/results, `Py<T>` for retention; binding/unbinding does not copy payloads
- Use `pythonize` for selected Serde data, never a JSON-text round trip; share conversion with `Pythonized<T>`
@ -15,6 +15,6 @@
- Release exclusive class borrows/locks before Python calls or decrements that can invoke finalizers; expose retained Python edges to GC without calling Python during traversal
- Keep coroutine driving in the shared Python driver and the native handle
- Driver: `litellm/rust_bridge/lifecycle.py`; handle: `src/handle.rs`; call driver: `src/driver.rs`; native-backed behavior tests: `tests/lifecycle.py`
- Every adapter suspension is awaited inline in the caller's task; `into_future` creates a separate task and cannot satisfy this contract
- Every lifecycle suspension is awaited inline in the caller's task; `into_future` creates a separate task and cannot satisfy this contract
- References: [ownership](https://pyo3.rs/v0.29.2/types.html), [conversions](https://pyo3.rs/v0.29.2/conversions/traits.html), [pythonize errors](https://docs.rs/pythonize/0.29.0/src/pythonize/error.rs.html)
- [GC](https://pyo3.rs/v0.29.2/class/protocols.html#garbage-collector-integration), [re-entry](https://pyo3.rs/v0.29.2/class/call.html), [parallelism](https://pyo3.rs/v0.29.2/parallelism.html), [async delivery source](https://docs.rs/pyo3-async-runtimes/0.29.0/src/pyo3_async_runtimes/generic.rs.html)

View file

@ -1,168 +0,0 @@
use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest};
use litellm_host::protocol::Protocol;
use pyo3::exceptions::PyRuntimeError;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::types::PyDict;
pub fn missing_state() -> PyErr {
PyRuntimeError::new_err("missing native call state")
}
/// The SDK's request policy, run by the driver on the keyword view `begin` returned and
/// before the protocol host projects from it. It rewrites that view in place, so the
/// lifecycle that returned it sees the rewrite too; a rejection fails the call as a host
/// failure, so the lifecycle still observes it.
pub type Preflight = fn(Python<'_>, &Bound<'_, PyDict>) -> PyResult<()>;
/// What an adapter step produced: either the value the driver asked for, or a Python
/// awaitable the driver hands back to the caller's task before asking again.
pub enum LifecycleStep {
Await(Py<PyAny>),
Arguments(Py<PyDict>),
Wire(Box<WireRequest>),
Response(Py<PyAny>),
Done,
}
/// What a lifecycle observes: the driver's start, the machine's own events, and one
/// terminal event carrying the public value the caller receives.
pub enum LifecycleEvent<'a> {
Started {
start_time: f64,
},
Machine(&'a MachineEvent),
Succeeded {
timing: Timing,
response: &'a Py<PyAny>,
},
Failed {
timing: Timing,
origin: FailureOrigin,
error: &'a PyErr,
},
}
/// One consumer of a call's lifecycle on the Python side. The driver calls the steps in
/// order: `begin` before the machine starts, `before_send` and `emit` while it runs,
/// `after_success` and one terminal `emit` after it completes. Whenever a step returns
/// [`LifecycleStep::Await`], the driver awaits it in the caller's task and continues the
/// same step through `resume`.
///
/// A step that fails with an ordinary exception fails the call with that exception,
/// except on a terminal event, where the adapter is expected to report and swallow its
/// own errors. An exception that is not a `PyException`, such as a cancellation, ends
/// the call without further dispatch.
pub trait PythonLifecycle: Send + Sync {
fn begin(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
started_at: f64,
) -> PyResult<LifecycleStep>;
fn before_send(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<LifecycleStep>;
fn after_success(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
timing: Timing,
) -> PyResult<LifecycleStep>;
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep>;
/// The call streams and its stream was handed to the caller. The caller is not
/// inside an await here, so this step and `delivered` cannot suspend.
fn opened(&mut self, py: Python<'_>) -> PyResult<()>;
/// One chunk of an open stream is about to reach the caller.
fn delivered(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()>;
fn resume(&mut self, py: Python<'_>, result: PyResult<Py<PyAny>>) -> PyResult<LifecycleStep>;
fn close(&mut self, py: Python<'_>);
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
}
/// 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)]
pub enum InvokeError<E> {
Native(E),
Python(PyErr),
}
impl<E> From<PyErr> for InvokeError<E> {
fn from(error: PyErr) -> Self {
Self::Python(error)
}
}
/// The Python side of one protocol: answers its custom operations, builds the public
/// response and classifies native failures into public exceptions.
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>;
/// 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>,
) -> 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::Protocol as Protocol>::Response,
) -> PyResult<Py<PyAny>>;
/// What the stream carries at hand-off, as the caller's stream receives it.
fn head(
&mut self,
py: Python<'_>,
head: <Self::Protocol as Protocol>::StreamHead,
) -> PyResult<Py<PyAny>>;
/// One streamed chunk as the caller receives it.
fn chunk(
&mut self,
py: Python<'_>,
chunk: <Self::Protocol as Protocol>::Chunk,
) -> PyResult<Py<PyAny>>;
fn classify(
&self,
py: Python<'_>,
error: <Self::Protocol as Protocol>::Error,
) -> PyResult<Self::Failure>;
fn host_error(error: &PyErr) -> <Self::Protocol as Protocol>::Error;
fn close(&mut self, py: Python<'_>);
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
}

View file

@ -0,0 +1,51 @@
use crate::{InvokeError, PythonOwned};
use litellm_host::protocol::Protocol;
use pyo3::prelude::*;
use pyo3::types::PyDict;
/// Converts requests, responses, stream values and errors at the Python boundary.
pub trait PythonBinding: PythonOwned {
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>;
/// Decodes the keyword view returned by `prepare_arguments`, including preflight rewrites.
fn decode_request(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> Result<
<Self::Protocol as Protocol>::Request,
InvokeError<<Self::Protocol as Protocol>::Error>,
>;
fn encode_response(
&mut self,
py: Python<'_>,
response: <Self::Protocol as Protocol>::Response,
) -> PyResult<Py<PyAny>>;
/// What the stream carries at hand-off, as the caller's stream receives it.
fn encode_stream_head(
&mut self,
py: Python<'_>,
head: <Self::Protocol as Protocol>::StreamHead,
) -> PyResult<Py<PyAny>>;
/// One streamed chunk as the caller receives it.
fn encode_chunk(
&mut self,
py: Python<'_>,
chunk: <Self::Protocol as Protocol>::Chunk,
) -> PyResult<Py<PyAny>>;
fn map_error(
&self,
py: Python<'_>,
error: <Self::Protocol as Protocol>::Error,
) -> PyResult<Self::Failure>;
fn host_error(error: &PyErr) -> <Self::Protocol as Protocol>::Error;
}

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,20 @@
use pyo3::prelude::*;
/// 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)]
pub enum InvokeError<E> {
Native(E),
Python(PyErr),
}
impl<E> From<PyErr> for InvokeError<E> {
fn from(error: PyErr) -> Self {
Self::Python(error)
}
}
pub fn missing_state() -> PyErr {
pyo3::exceptions::PyRuntimeError::new_err("missing native call state")
}

View file

@ -31,6 +31,10 @@ pub struct Execution {
state: ExecutionState,
}
fn lifecycle(py: Python<'_>) -> PyResult<Bound<'_, PyModule>> {
py.import("litellm.rust_bridge.lifecycle")
}
impl Execution {
pub fn new(body: impl ExecutionBody + 'static) -> Self {
Self {
@ -38,6 +42,21 @@ impl Execution {
}
}
pub fn into_coroutine(self, py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
let execution = Py::new(py, self)?;
lifecycle(py)?.getattr("drive")?.call1((execution,))
}
pub(crate) fn into_sync_stream(
self,
py: Python<'_>,
head: Py<PyAny>,
) -> PyResult<Bound<'_, PyAny>> {
lifecycle(py)?
.getattr("SyncStream")?
.call1((Py::new(py, self)?, head))
}
/// An execution already started elsewhere and now waiting for its next input.
pub fn suspended(body: impl ExecutionBody + 'static) -> Self {
Self {
@ -79,11 +98,7 @@ impl Execution {
ExecutionStep::Yield(value) => ("Yield", value, true),
ExecutionStep::Return(value) => ("Complete", value, false),
};
let step = py
.import("litellm.rust_bridge.lifecycle")?
.getattr(tag)?
.call1((value,))?
.unbind();
let step = lifecycle(py)?.getattr(tag)?.call1((value,))?.unbind();
Ok((step, suspended))
}))
.map_err(panic_to_pyerr)

View file

@ -0,0 +1,79 @@
use crate::PythonOwned;
use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest};
use pyo3::prelude::*;
use pyo3::types::PyDict;
/// The SDK's request policy, run by the driver on the keyword view `prepare_arguments` returned and
/// before the binding decodes from it. It rewrites that view in place, so the
/// hooks that returned it see the rewrite too; a rejection fails the call as a host
/// failure, so the hooks still observe it.
pub type Preflight = fn(Python<'_>, &Bound<'_, PyDict>) -> PyResult<()>;
/// What a hook step produced: either the value the driver asked for, or a Python
/// awaitable the driver hands back to the caller's task before asking again.
pub type HookResume<L, T> = fn(&mut L, Python<'_>, PyResult<Py<PyAny>>) -> PyResult<HookStep<L, T>>;
pub enum HookStep<L, T> {
Await(Py<PyAny>, HookResume<L, T>),
Ready(T),
}
/// Events dispatched to call hooks: the driver's start, the machine's own events, and one
/// terminal event carrying the public value the caller receives.
pub enum HookEvent<'a> {
Started {
start_time: f64,
},
Machine(&'a MachineEvent),
Succeeded {
timing: Timing,
response: &'a Py<PyAny>,
},
Failed {
timing: Timing,
origin: FailureOrigin,
error: &'a PyErr,
},
}
/// Active Python hooks that can transform values or fail execution. The driver calls the steps in
/// order: `prepare_arguments` before the machine starts, `before_provider_request` and `on_event` while it runs,
/// `transform_response` and one terminal `on_event` after it completes. Whenever a step returns
/// [`HookStep::Await`], the driver awaits it in the caller's task and continues the
/// same step through its typed continuation.
///
/// A step that fails with an ordinary exception fails the call with that exception,
/// except on a terminal event, where the hooks are expected to report and swallow their
/// own errors. An exception that is not a `PyException`, such as a cancellation, ends
/// the call without further dispatch.
pub trait PythonCallHooks: Sized + PythonOwned {
fn prepare_arguments(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
started_at: f64,
) -> PyResult<HookStep<Self, Py<PyDict>>>;
fn before_provider_request(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<HookStep<Self, Box<WireRequest>>>;
fn transform_response(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
timing: Timing,
) -> PyResult<HookStep<Self, Py<PyAny>>>;
fn on_event(&mut self, py: Python<'_>, event: HookEvent<'_>) -> PyResult<HookStep<Self, ()>>;
/// The call streams and its stream was handed to the caller. The caller is not
/// inside an await here, so this step and `on_stream_chunk` cannot suspend.
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()>;
/// One chunk of an open stream is about to reach the caller.
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()>;
}

View file

@ -1,38 +1,44 @@
//! 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 protocol host and a Python lifecycle. Everything here is Python-specific by
//! against a Python binding, host services and active call hooks. Everything here is Python-specific by
//! construction; another host language gets its own crate of the same shape.
mod adapter;
mod argument;
mod binding;
mod callable;
mod driver;
mod execution;
mod error;
mod file_reader;
mod fork_gate;
mod gil;
mod handle;
mod hooks;
mod marshal;
mod native;
mod owned;
mod runtime;
mod services;
pub use adapter::{
InvokeError, LifecycleEvent, LifecycleStep, Preflight, ProtocolHost, PythonLifecycle,
missing_state,
};
pub use argument::lookup;
pub use binding::PythonBinding;
pub use callable::wrap_failure;
pub use driver::run_call;
pub use execution::{
ForkedAfterNativeRuntimeStarted, ProcessReservedForForking, enter_native, poll_async_value,
reserve_process_for_forking, run_async, run_async_value, run_sync, run_sync_value,
runtime_started,
};
pub use error::{InvokeError, missing_state};
pub use file_reader::{FileContent, PythonFileReader, py_bytes};
pub use fork_gate::RuntimeAlreadyStarted;
pub use gil::{PythonContext, attach_blocking, release_count, release_gil};
pub use handle::{Execution, ExecutionBody, ExecutionStep};
pub use hooks::{HookEvent, HookResume, HookStep, Preflight, PythonCallHooks};
pub use marshal::{
Pythonized, from_py, from_py_argument, json_loads, json_object_field, panic_to_pyerr, to_py,
};
pub use owned::PythonOwned;
pub use runtime::{
ForkedAfterNativeRuntimeStarted, ProcessReservedForForking, enter_native, poll_async_value,
reserve_process_for_forking, run_async, run_async_value, run_sync, run_sync_value,
runtime_started,
};
pub use services::PythonHostCalls;
/// Starts the interpreter and imports the standard modules the tests share, once, so
/// parallel test threads never race a first import of `asyncio`.

View file

@ -0,0 +1,134 @@
use std::sync::Arc;
use std::task::Poll;
use futures_util::future::{AbortHandle, Abortable};
use litellm_host::call::HostedCompletion;
use litellm_host::machine::{HostFailure, Machine, MachineStep};
use litellm_host::protocol::Protocol;
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::*;
use tokio::sync::Mutex;
use crate::missing_state;
use crate::runtime::{poll_async_value, run_async_value, run_sync_value};
type NativeResult<M> = Result<
MachineStep<
<M as Machine>::Protocol,
HostedCompletion<<<M as Machine>::Protocol as Protocol>::Response>,
>,
<<M as Machine>::Protocol as Protocol>::Error,
>;
type MachineResult<M> = Result<
MachineStep<<M as Machine>::Protocol, <M as Machine>::Complete>,
<<M as Machine>::Protocol as Protocol>::Error,
>;
struct MachineState<M: Machine> {
machine: M,
result: Option<MachineResult<M>>,
}
pub(super) enum NativePoll<T> {
Ready(T),
Suspend(Py<PyAny>),
}
pub(super) struct NativeMachine<M: Machine> {
state: Option<Arc<Mutex<MachineState<M>>>>,
abort: Option<AbortHandle>,
asynchronous: bool,
}
impl<M: Machine + 'static> NativeMachine<M>
where
M::Complete: Into<HostedCompletion<<M::Protocol as Protocol>::Response>>,
{
pub(super) fn new(asynchronous: bool) -> Self {
Self {
state: None,
abort: None,
asynchronous,
}
}
pub(super) fn start(&mut self, machine: M) {
self.state = Some(Arc::new(Mutex::new(MachineState {
machine,
result: None,
})));
}
pub(super) fn resume(
&mut self,
py: Python<'_>,
interruption: Option<HostFailure<<M::Protocol as Protocol>::Error>>,
) -> PyResult<NativePoll<NativeResult<M>>> {
let state = Arc::clone(self.state.as_ref().ok_or_else(missing_state)?);
let future = async move {
let mut state = state.lock().await;
let result = match interruption {
Some(failure) => state
.machine
.interrupt(failure)
.await
.map(MachineStep::Complete),
None => state.machine.resume().await,
};
state.result = Some(result);
Ok(())
};
if self.asynchronous {
let mut future = Box::pin(future);
if let Poll::Ready(()) = poll_async_value(py, future.as_mut())? {
return Ok(NativePoll::Ready(self.take_result()?));
}
let (abort, registration) = AbortHandle::new_pair();
self.abort = Some(abort);
Ok(NativePoll::Suspend(
run_async_value(py, async move {
Abortable::new(future, registration)
.await
.map_err(|_| PyRuntimeError::new_err("native execution closed"))?
})?
.unbind(),
))
} else {
run_sync_value(py, future)?;
Ok(NativePoll::Ready(self.take_result()?))
}
}
pub(super) fn take_result(&self) -> PyResult<NativeResult<M>> {
self.state
.as_ref()
.ok_or_else(missing_state)?
.try_lock()
.map_err(|_| missing_state())?
.result
.take()
.ok_or_else(missing_state)
.map(|result| {
result.map(|step| match step {
MachineStep::Suspended(op) => MachineStep::Suspended(op),
MachineStep::Complete(response) => MachineStep::Complete(response.into()),
})
})
}
}
impl<M: Machine> NativeMachine<M> {
pub(super) fn close(&mut self) {
if let Some(abort) = self.abort.take() {
abort.abort();
}
self.state = None;
}
}
impl<M: Machine> Drop for NativeMachine<M> {
fn drop(&mut self) {
self.close();
}
}

View file

@ -0,0 +1,9 @@
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,
};
pub trait PythonOwned: Send + Sync {
fn close(&mut self, py: Python<'_>);
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
}

View file

@ -0,0 +1,11 @@
use crate::{InvokeError, PythonOwned};
use litellm_host::protocol::Protocol;
use pyo3::prelude::*;
pub trait PythonHostCalls<P: Protocol>: PythonOwned {
fn handle_host_call(
&mut self,
py: Python<'_>,
call: P::HostCall,
) -> Result<(), InvokeError<P::Error>>;
}

View file

@ -0,0 +1,19 @@
`litellm-host` defines typed calls, execution hooks, host services, stream delivery and the resumable machine. HTTP and Python drivers interpret the same suspension protocol in their own runtimes
| Responsibility | Shared contract | HTTP | Python |
| --- | --- | --- | --- |
| Input and output conversion | `Protocol::{Request, Response, StreamHead, Chunk, Error}` | Typed input; `ResponseEncoder` and `StreamEncoder` produce HTTP values | `PythonBinding` decodes prepared arguments, encodes public values and maps native errors |
| Host services | `Protocol::HostCall`, `HostServices::call` | `HostCallHandler` answers typed calls; `()` handles protocols without host calls | `PythonHostCalls` invokes retained Python objects; it may share an owner with the binding |
| Active hooks | `RouteHooks::{before_provider_request, on_event}` | Request interception and fallible execution callbacks | `PythonCallHooks` also prepares arguments, transforms public responses and receives stream callbacks |
| Passive observation | `lifecycle::CallObserver` | Start and terminal observation, retained by the response body | Public Python callbacks remain active hooks with their existing failure policy |
| Runtime driving | `Machine`, `Suspension::{HostCall, Hook, Stream}` | Body polling controls demand | Native polling and inline caller-task Python awaits control progress |
`hosted_call(request, execute)` starts with a typed request. Its route closure receives separate `HostServices` and `ChannelHooks`; it returns `CallOutput`. Hosted-call plumbing alone forwards the returned stream through demand replies. Lower-level `CallMachine` users receive a `CallContext` containing separately named services, hooks and stream delivery
Core route constructors prepare their dependencies and return a closure accepting the typed request. Python starts that closure only after argument preparation, preflight and decoding succeed. These steps remain inside the driver's terminal and error handling. Decoding may retain objects for subsequent host service calls
Each driver owns terminal dispatch. Hooks can change or fail execution; passive observers return no result. HTTP observes success after response conversion or stream exhaustion, failure on errors, and cancellation on body drop. Python preserves exception identity and maps native failures once. Explicit Python stream close reports success for delivered chunks; cancellation stops further callback dispatch
`in_process::Host` is an assembly of services, hooks, stream consumer and optional observer. It is not a trait mirroring every suspension. Use `run_hosted` to preserve the distinction between stream completion and detachment
Keep API policy in gateway-inference and python-bridge, and legacy callback policy in callbacks-legacy-python. Python bindings and hooks expose retained references through `PythonOwned`, with idempotent close and GC traversal. Runtime machinery stays in driver, native, handle and runtime modules

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
futures-util.workspace = true
litellm-auth.workspace = true
litellm-coroutine.workspace = true
serde_json.workspace = true

View file

@ -0,0 +1,65 @@
use std::future::Future;
use futures_util::{TryStreamExt, stream::BoxStream};
use crate::{
machine::{CallMachine, ChannelHooks, HostServices, MachineFault},
protocol::{Demand, Protocol},
};
pub enum CallOutput<Response, Head, Chunk, Error> {
Complete(Response),
Stream {
head: Head,
chunks: BoxStream<'static, Result<Chunk, Error>>,
},
}
#[derive(Debug, PartialEq, Eq)]
pub enum HostedCompletion<Response> {
Complete(Response),
StreamEnded,
Detached,
}
impl<Response> From<Response> for HostedCompletion<Response> {
fn from(response: Response) -> Self {
Self::Complete(response)
}
}
pub type OutputOf<P> = CallOutput<
<P as Protocol>::Response,
<P as Protocol>::StreamHead,
<P as Protocol>::Chunk,
<P as Protocol>::Error,
>;
pub type HostedMachine<P> = CallMachine<P, HostedCompletion<<P as Protocol>::Response>>;
pub fn hosted_call<P, F, Fut>(request: P::Request, execute: F) -> HostedMachine<P>
where
P: Protocol,
P::Error: From<MachineFault>,
F: FnOnce(P::Request, HostServices<P>, ChannelHooks<P>) -> Fut + Send + 'static,
Fut: Future<Output = Result<OutputOf<P>, P::Error>> + Send + 'static,
{
CallMachine::new(move |host| {
Box::pin(async move {
match execute(request, host.services, host.hooks).await? {
CallOutput::Complete(response) => Ok(HostedCompletion::Complete(response)),
CallOutput::Stream { head, mut chunks } => {
if host.stream.open_stream(head).await? == Demand::Detached {
return Ok(HostedCompletion::Detached);
}
while let Some(chunk) = chunks.try_next().await? {
if host.stream.send_chunk(chunk).await? == Demand::Detached {
return Ok(HostedCompletion::Detached);
}
}
Ok(HostedCompletion::StreamEnded)
}
}
})
})
}

View file

@ -73,4 +73,7 @@ pub enum CallEvent {
timing: Timing,
origin: FailureOrigin,
},
Cancelled {
timing: Timing,
},
}

View file

@ -1,51 +1,38 @@
use std::future::Future;
use crate::{
event::{MachineEvent, RequestContext, WireRequest},
machine::{HostChannel, MachineFault},
protocol::Protocol,
};
use crate::event::{MachineEvent, RequestContext, WireRequest};
/// What a route reaches for mid-call: the send-time rewrite and the events it reports.
/// Python's `logging_obj.pre_call` and `post_call`, in that order.
pub trait RouteHooks<E>: Send + Sync {
fn before_send(
fn observer(&self) -> Option<std::sync::Arc<dyn crate::lifecycle::CallObserver>> {
None
}
fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
) -> impl Future<Output = Result<WireRequest, E>> + Send;
fn emit(&self, event: MachineEvent) -> impl Future<Output = Result<(), E>> + Send;
fn on_event(&self, event: MachineEvent) -> impl Future<Output = Result<(), E>> + Send;
}
/// No host: the wire request goes out as prepared and nothing observes the call.
impl<E> RouteHooks<E> for () {
async fn before_send(&self, wire: WireRequest, _: RequestContext) -> Result<WireRequest, E> {
async fn before_provider_request(
&self,
wire: WireRequest,
_: RequestContext,
) -> Result<WireRequest, E> {
Ok(wire)
}
async fn emit(&self, _: MachineEvent) -> Result<(), E> {
async fn on_event(&self, _: MachineEvent) -> Result<(), E> {
Ok(())
}
}
impl<R: Protocol> RouteHooks<R::Error> for HostChannel<R>
where
R::Error: From<MachineFault>,
{
async fn before_send(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, R::Error> {
HostChannel::before_send(self, wire, context).await
}
async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> {
HostChannel::emit(self, event).await
}
}
#[cfg(test)]
mod tests {
use std::convert::Infallible;
@ -53,10 +40,13 @@ mod tests {
use serde_json::json;
use super::*;
use crate::protocol::HookRequest;
use crate::{
event::RawResponse,
host::HostOp,
machine::MachineFault,
machine::{CallMachine, Machine, MachineStep},
protocol::Protocol,
protocol::Suspension,
};
struct Unit;
@ -67,8 +57,8 @@ mod tests {
impl Protocol for Unit {
type Response = (WireRequest, ());
type Error = Fault;
type Projection = ();
type Op = Infallible;
type Request = ();
type HostCall = Infallible;
type Chunk = Infallible;
type StreamHead = Infallible;
}
@ -97,13 +87,19 @@ mod tests {
}
}
#[rstest::rstest]
#[tokio::test]
async fn the_channel_yields_each_hook_as_its_op_and_returns_the_answer() {
let mut machine = CallMachine::<Unit>::new(|channel| {
Box::pin(async move {
let sent = RouteHooks::before_send(&channel, wire("prepared"), context()).await?;
RouteHooks::emit(
&channel,
let sent = RouteHooks::before_provider_request(
&channel.hooks,
wire("prepared"),
context(),
)
.await?;
RouteHooks::on_event(
&channel.hooks,
MachineEvent::ResponseReceived {
raw: RawResponse { body: "raw".into() },
},
@ -113,9 +109,13 @@ mod tests {
})
});
let Ok(MachineStep::Host(HostOp::BeforeSend { wire, reply, .. })) = machine.resume().await
let Ok(MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest {
wire,
reply,
..
}))) = machine.resume().await
else {
panic!("before_send yields BeforeSend");
panic!("before_provider_request yields BeforeSend");
};
assert_eq!(wire.url, "prepared");
reply.send(WireRequest {
@ -123,8 +123,10 @@ mod tests {
..*wire
});
let Ok(MachineStep::Host(HostOp::Emit(event, reply))) = machine.resume().await else {
panic!("emit yields Emit");
let Ok(MachineStep::Suspended(Suspension::Hook(HookRequest::Event(event, reply)))) =
machine.resume().await
else {
panic!("on_event yields Emit");
};
assert!(matches!(event, MachineEvent::ResponseReceived { .. }));
reply.send(());
@ -135,9 +137,10 @@ mod tests {
assert_eq!(sent.url, "rewritten");
}
#[rstest::rstest]
#[tokio::test]
async fn no_hooks_pass_the_wire_request_through() {
let sent = RouteHooks::<Fault>::before_send(&(), wire("prepared"), context())
let sent = RouteHooks::<Fault>::before_provider_request(&(), wire("prepared"), context())
.await
.unwrap();
assert_eq!(sent.url, "prepared");

View file

@ -1,68 +0,0 @@
use std::future::Future;
pub use litellm_coroutine::{Abandoned, Answer, Reply, reply};
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, 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, Reply<Demand>),
/// The next chunk of an open stream, answered once the caller asks for the one after.
Deliver(R::Chunk, Reply<Demand>),
}
/// Whether the caller of a streamed call still reads it.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Demand {
More,
Detached,
}
/// A host answer that is either available now or arrives once the host's own
/// suspension (a Python awaitable, for example) resolves.
pub enum HostStep<V, S> {
Ready(V),
Suspend(S),
}
/// 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: 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,
wire: WireRequest,
_context: &RequestContext,
) -> impl Future<Output = Result<WireRequest, R::Error>> + Send {
async move { Ok(wire) }
}
fn emit(&self, _event: &CallEvent) -> impl Future<Output = Result<(), R::Error>> + Send {
async { Ok(()) }
}
fn open(&self, _head: R::StreamHead) -> impl Future<Output = Result<Demand, R::Error>> + Send {
async { Ok(Demand::More) }
}
fn deliver(&self, _chunk: R::Chunk) -> impl Future<Output = Result<Demand, R::Error>> + Send {
async { Ok(Demand::More) }
}
}

View file

@ -0,0 +1,323 @@
use crate::{
event::{CallEvent, FailureOrigin, Timing, epoch_seconds},
hooks::RouteHooks,
lifecycle::CallObserver,
machine::{HostFailure, Machine, MachineStep},
protocol::{Demand, HookRequest, Protocol, StreamDelivery, Suspension},
services::HostCallHandler,
};
use std::future::Future;
pub trait StreamConsumer<P: Protocol>: Send + Sync {
fn open_stream(
&self,
head: P::StreamHead,
) -> impl Future<Output = Result<Demand, P::Error>> + Send;
fn send_chunk(&self, chunk: P::Chunk) -> impl Future<Output = Result<Demand, P::Error>> + Send;
}
impl<P: Protocol> StreamConsumer<P> for () {
async fn open_stream(&self, _: P::StreamHead) -> Result<Demand, P::Error> {
Ok(Demand::More)
}
async fn send_chunk(&self, _: P::Chunk) -> Result<Demand, P::Error> {
Ok(Demand::More)
}
}
pub struct Host<'a, S, H, C> {
pub services: &'a S,
pub hooks: &'a H,
pub stream: &'a C,
pub observer: Option<&'a dyn CallObserver>,
}
pub async fn run<M, S, H, C>(
machine: M,
host: Host<'_, S, H, C>,
) -> Result<M::Complete, <M::Protocol as Protocol>::Error>
where
M: Machine,
S: HostCallHandler<M::Protocol>,
H: RouteHooks<<M::Protocol as Protocol>::Error>,
C: StreamConsumer<M::Protocol>,
{
run_with_completion(machine, host, |_| false).await
}
pub async fn run_hosted<P, S, H, C>(
machine: crate::call::HostedMachine<P>,
host: Host<'_, S, H, C>,
) -> Result<crate::call::HostedCompletion<P::Response>, P::Error>
where
P: Protocol,
P::Error: From<crate::machine::MachineFault>,
S: HostCallHandler<P>,
H: RouteHooks<P::Error>,
C: StreamConsumer<P>,
{
run_with_completion(machine, host, |completion| {
matches!(completion, crate::call::HostedCompletion::Detached)
})
.await
}
async fn run_with_completion<M, S, H, C>(
mut machine: M,
host: Host<'_, S, H, C>,
detached: impl Fn(&M::Complete) -> bool,
) -> Result<M::Complete, <M::Protocol as Protocol>::Error>
where
M: Machine,
S: HostCallHandler<M::Protocol>,
H: RouteHooks<<M::Protocol as Protocol>::Error>,
C: StreamConsumer<M::Protocol>,
{
let start_time = epoch_seconds();
if let Some(observer) = host.observer {
observer.observe(CallEvent::Started { start_time });
}
let outcome = loop {
let suspension = match machine.resume().await {
Ok(MachineStep::Complete(complete)) => break Ok(complete),
Ok(MachineStep::Suspended(suspension)) => suspension,
Err(error) => break Err(error),
};
let result = match suspension {
Suspension::HostCall(call) => host.services.handle_host_call(call).await,
Suspension::Hook(HookRequest::BeforeProviderRequest {
wire,
context,
reply,
}) => host
.hooks
.before_provider_request(*wire, *context)
.await
.map(|wire| reply.send(wire)),
Suspension::Hook(HookRequest::Event(event, reply)) => {
host.hooks.on_event(event).await.map(|()| reply.send(()))
}
Suspension::Stream(StreamDelivery::Open(head, reply)) => host
.stream
.open_stream(head)
.await
.map(|demand| reply.send(demand)),
Suspension::Stream(StreamDelivery::Chunk(chunk, reply)) => host
.stream
.send_chunk(chunk)
.await
.map(|demand| reply.send(demand)),
};
if let Err(error) = result {
break machine.interrupt(HostFailure::Error(error)).await;
}
};
let timing = Timing {
start_time,
end_time: epoch_seconds(),
};
let terminal = match &outcome {
Ok(completion) if detached(completion) => CallEvent::Cancelled { timing },
Ok(_) => CallEvent::Succeeded { timing },
Err(_) => CallEvent::Failed {
timing,
origin: FailureOrigin::Call,
},
};
if let Some(observer) = host.observer {
observer.observe(terminal);
}
outcome
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use super::*;
use crate::machine::{CallMachine, MachineFault};
use crate::protocol::Reply;
struct Unit;
impl Protocol for Unit {
type Response = ();
type Error = &'static str;
type Request = ();
type HostCall = (&'static str, Reply<()>);
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
impl From<MachineFault> for &'static str {
fn from(_: MachineFault) -> Self {
"machine fault"
}
}
#[derive(Default)]
struct Recording {
seen: Mutex<Vec<String>>,
fail: Option<&'static str>,
}
impl Recording {
pub fn runtime(&self) -> crate::in_process::Host<'_, Self, Self, ()> {
crate::in_process::Host {
services: self,
hooks: self,
stream: &(),
observer: Some(self),
}
}
}
impl crate::services::HostCallHandler<Unit> for Recording {
async fn handle_host_call(
&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(())
}
}
impl crate::lifecycle::CallObserver for Recording {
fn observe(&self, event: crate::event::CallEvent) {
self.seen.lock().unwrap().push(match event {
CallEvent::Started { .. } => "started".into(),
CallEvent::Succeeded { .. } => "succeeded".into(),
CallEvent::Failed { .. } => "failed".into(),
other => format!("{other:?}"),
});
}
}
impl crate::hooks::RouteHooks<<Unit as crate::protocol::Protocol>::Error> for Recording {
async fn before_provider_request(
&self,
wire: crate::event::WireRequest,
_: crate::event::RequestContext,
) -> Result<crate::event::WireRequest, <Unit as crate::protocol::Protocol>::Error> {
Ok(wire)
}
async fn on_event(
&self,
event: crate::event::MachineEvent,
) -> Result<(), <Unit as crate::protocol::Protocol>::Error> {
crate::lifecycle::CallObserver::observe(self, crate::event::CallEvent::Machine(event));
Ok(())
}
}
fn scripted(
ops: &'static [&'static str],
outcome: Result<(), &'static str>,
) -> CallMachine<Unit> {
CallMachine::new(move |host| {
Box::pin(async move {
for op in ops {
host.services.call(|reply| (*op, reply)).await?;
}
outcome
})
})
}
#[rstest::rstest]
#[tokio::test]
async fn forwards_every_op_then_emits_one_succeeded() {
let host = Recording::default();
let outcome = run(scripted(&["sign", "send"], Ok(())), host.runtime()).await;
assert_eq!(outcome, Ok(()));
assert_eq!(
*host.seen.lock().unwrap(),
["started", "op:sign", "op:send", "succeeded"]
);
}
#[rstest::rstest]
#[tokio::test]
async fn errors_and_host_failures_each_emit_failed_once() {
let host = Recording::default();
let outcome = run(scripted(&[], Err("boom")), host.runtime()).await;
assert_eq!(outcome, Err("boom"));
assert_eq!(*host.seen.lock().unwrap(), ["started", "failed"]);
let host = Recording {
fail: Some("send"),
..Recording::default()
};
let outcome = run(scripted(&["sign", "send", "never"], Ok(())), host.runtime()).await;
assert_eq!(outcome, Err("host failed"));
assert_eq!(
*host.seen.lock().unwrap(),
["started", "op:sign", "op:send", "failed"]
);
}
struct StartTimes(Mutex<Vec<f64>>);
impl StartTimes {
pub fn runtime(&self) -> crate::in_process::Host<'_, Self, Self, ()> {
crate::in_process::Host {
services: self,
hooks: self,
stream: &(),
observer: Some(self),
}
}
}
impl crate::services::HostCallHandler<Unit> for StartTimes {
async fn handle_host_call(
&self,
(_, reply): (&'static str, Reply<()>),
) -> Result<(), &'static str> {
reply.send(());
Ok(())
}
}
impl crate::lifecycle::CallObserver for StartTimes {
fn observe(&self, event: crate::event::CallEvent) {
if let CallEvent::Started { start_time }
| CallEvent::Succeeded {
timing: Timing { start_time, .. },
} = event
{
self.0.lock().unwrap().push(start_time);
}
}
}
impl crate::hooks::RouteHooks<<Unit as crate::protocol::Protocol>::Error> for StartTimes {
async fn before_provider_request(
&self,
wire: crate::event::WireRequest,
_: crate::event::RequestContext,
) -> Result<crate::event::WireRequest, <Unit as crate::protocol::Protocol>::Error> {
Ok(wire)
}
async fn on_event(
&self,
event: crate::event::MachineEvent,
) -> Result<(), <Unit as crate::protocol::Protocol>::Error> {
crate::lifecycle::CallObserver::observe(self, crate::event::CallEvent::Machine(event));
Ok(())
}
}
#[rstest::rstest]
#[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(&["send"], Ok(())), host.runtime()).await,
Ok(())
);
let times = host.0.lock().unwrap();
assert_eq!(times.len(), 2);
assert_eq!(times[0], times[1]);
}
}

View file

@ -2,13 +2,16 @@
//!
//! 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::CallMachine`] and never learns
//! which host is on the other end. The machine yields [`host::HostOp`]s; a driver answers
//! each through the typed [`host::Reply`] it carries, observes [`event::CallEvent`]s and
//! which host is on the other end. The machine yields [`protocol::Suspension`]s; a driver answers
//! each through the typed [`protocol::Reply`] it carries, observes [`event::CallEvent`]s and
//! may rewrite the wire request before it is sent.
pub mod call;
pub mod event;
pub mod hooks;
pub mod host;
pub mod in_process;
pub mod lifecycle;
pub mod machine;
pub mod protocol;
pub mod run;
pub mod services;

View file

@ -0,0 +1,121 @@
use std::{future::Future, sync::Arc};
use futures_util::TryStreamExt;
use crate::{
call::CallOutput,
event::{CallEvent, FailureOrigin, Timing, epoch_seconds},
};
pub trait CallObserver: Send + Sync {
fn observe(&self, event: CallEvent);
}
struct CallGuard {
observer: Option<Arc<dyn CallObserver>>,
started_at: f64,
}
impl CallGuard {
fn new(observer: Arc<dyn CallObserver>) -> Self {
let started_at = epoch_seconds();
observer.observe(CallEvent::Started {
start_time: started_at,
});
Self {
observer: Some(observer),
started_at,
}
}
fn timing(&self) -> Timing {
Timing {
start_time: self.started_at,
end_time: epoch_seconds(),
}
}
fn finish(mut self, failed: bool) {
if let Some(observer) = self.observer.take() {
observer.observe(if failed {
CallEvent::Failed {
timing: self.timing(),
origin: FailureOrigin::Call,
}
} else {
CallEvent::Succeeded {
timing: self.timing(),
}
});
}
}
}
impl Drop for CallGuard {
fn drop(&mut self) {
if let Some(observer) = self.observer.take() {
observer.observe(CallEvent::Cancelled {
timing: self.timing(),
});
}
}
}
pub async fn observe_call<R, H, C, E>(
observer: Option<Arc<dyn CallObserver>>,
execute: impl Future<Output = Result<CallOutput<R, H, C, E>, E>>,
) -> Result<CallOutput<R, H, C, E>, E>
where
C: Send + 'static,
E: Send + 'static,
{
let Some(observer) = observer else {
return execute.await;
};
let guard = CallGuard::new(observer);
match execute.await {
Err(error) => {
guard.finish(true);
Err(error)
}
Ok(CallOutput::Complete(response)) => {
guard.finish(false);
Ok(CallOutput::Complete(response))
}
Ok(CallOutput::Stream { head, chunks }) => {
let stream = futures_util::stream::try_unfold(
(chunks, guard),
|(mut chunks, guard)| async move {
match chunks.try_next().await {
Ok(Some(chunk)) => Ok(Some((chunk, (chunks, guard)))),
Ok(None) => {
guard.finish(false);
Ok(None)
}
Err(error) => {
guard.finish(true);
Err(error)
}
}
},
);
Ok(CallOutput::Stream {
head,
chunks: Box::pin(stream),
})
}
}
}
pub async fn observe_unary<R, E>(
observer: Option<Arc<dyn CallObserver>>,
execute: impl Future<Output = Result<R, E>>,
) -> Result<R, E> {
let Some(observer) = observer else {
return execute.await;
};
let guard = CallGuard::new(observer);
let result = execute.await;
guard.finish(result.is_err());
result
}

View file

@ -1,18 +1,18 @@
use std::sync::Arc;
use super::{HostChannel, MachineFault};
use crate::{host::Reply, protocol::Protocol};
use super::{HostServices, MachineFault};
use crate::{protocol::Protocol, protocol::Reply};
use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};
/// 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;
fn acquire_token_op(reply: Reply<ResolvedCredential>) -> Self::HostCall;
}
/// 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: Protocol> {
channel: HostChannel<R>,
channel: HostServices<R>,
}
impl<R: Protocol> std::fmt::Debug for HostTokenProvider<R> {
@ -26,7 +26,7 @@ where
R: TokenProtocol,
R::Error: From<MachineFault> + std::fmt::Display,
{
pub fn handle(channel: HostChannel<R>) -> TokenProviderHandle {
pub fn handle(channel: HostServices<R>) -> TokenProviderHandle {
TokenProviderHandle::new(Arc::new(Self { channel }))
}
}
@ -39,7 +39,7 @@ where
fn acquire(&self) -> TokenFuture<'_> {
Box::pin(async move {
self.channel
.custom_op(R::acquire_token_op)
.call(R::acquire_token_op)
.await
.map_err(|error| Error::CredentialAcquisition(error.to_string().into()))
})

View file

@ -1,7 +1,9 @@
//! 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
//! [`Coroutine`] that yields [`Suspension`]s, each answered through its own typed reply. No
//! task is spawned; dropping the machine drops the in-flight call.
use crate::protocol::HookRequest;
use crate::protocol::StreamDelivery;
use std::{future::Future, pin::Pin};
use litellm_coroutine::{Co, Coroutine, CoroutineState, ResumeError};
@ -9,8 +11,7 @@ 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,
protocol::{Demand, Protocol, Reply, Suspension},
};
/// The machine's own failures, distinct from anything the provider call reports.
@ -22,99 +23,146 @@ pub enum MachineFault {
Protocol(ResumeError),
}
pub type ExecuteFuture<R> =
Pin<Box<dyn Future<Output = Result<<R as Protocol>::Response, <R as Protocol>::Error>> + Send>>;
pub type ExecuteFuture<R, C = <R as Protocol>::Response> =
Pin<Box<dyn Future<Output = Result<C, <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>>,
pub struct CallContext<P: Protocol> {
pub services: HostServices<P>,
pub hooks: ChannelHooks<P>,
pub stream: StreamSender<P>,
}
impl<R: Protocol> Clone for HostChannel<R> {
struct Channel<P: Protocol>(Co<Suspension<P>>);
impl<P: Protocol> Clone for Channel<P> {
fn clone(&self) -> Self {
Self {
co: self.co.clone(),
}
Self(self.0.clone())
}
}
impl<R: Protocol> HostChannel<R>
impl<P: Protocol> Channel<P>
where
R::Error: From<MachineFault>,
P::Error: From<MachineFault>,
{
async fn yield_<A: Send>(
async fn request_reply<A: Send>(
&self,
ask: impl FnOnce(Reply<A>) -> HostOp<R> + Send,
) -> Result<A, R::Error> {
self.co
.yield_(ask)
request: impl FnOnce(Reply<A>) -> Suspension<P> + Send,
) -> Result<A, P::Error> {
self.0
.yield_(request)
.await
.map_err(|_| MachineFault::Abandoned.into())
}
}
pub async fn project(&self) -> Result<R::Projection, R::Error> {
self.yield_(HostOp::Project).await
pub struct HostServices<P: Protocol>(Channel<P>);
impl<P: Protocol> Clone for HostServices<P> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
/// 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>(
impl<P: Protocol> HostServices<P>
where
P::Error: From<MachineFault>,
{
pub async fn call<A: Send>(
&self,
ask: impl FnOnce(Reply<A>) -> R::Op + Send,
) -> Result<A, R::Error> {
self.yield_(|reply| HostOp::Custom(ask(reply))).await
request: impl FnOnce(Reply<A>) -> P::HostCall + Send,
) -> Result<A, P::Error> {
self.0
.request_reply(|reply| Suspension::HostCall(request(reply)))
.await
}
}
pub async fn before_send(
pub struct ChannelHooks<P: Protocol>(Channel<P>);
impl<P: Protocol> Clone for ChannelHooks<P> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl<P: Protocol> crate::hooks::RouteHooks<P::Error> for ChannelHooks<P>
where
P::Error: From<MachineFault>,
{
async fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, R::Error> {
self.yield_(|reply| HostOp::BeforeSend {
wire: Box::new(wire),
context: Box::new(context),
reply,
})
.await
) -> Result<WireRequest, P::Error> {
self.0
.request_reply(|reply| {
Suspension::Hook(HookRequest::BeforeProviderRequest {
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
async fn on_event(&self, event: MachineEvent) -> Result<(), P::Error> {
self.0
.request_reply(|reply| Suspension::Hook(HookRequest::Event(event, reply)))
.await
}
}
type CallCoroutine<R> =
Coroutine<HostOp<R>, Result<<R as Protocol>::Response, <R as Protocol>::Error>>;
pub struct StreamSender<P: Protocol>(Channel<P>);
pub struct CallMachine<R: Protocol> {
coroutine: CallCoroutine<R>,
impl<P: Protocol> StreamSender<P>
where
P::Error: From<MachineFault>,
{
pub async fn open_stream(&self, head: P::StreamHead) -> Result<Demand, P::Error> {
self.0
.request_reply(|reply| Suspension::Stream(StreamDelivery::Open(head, reply)))
.await
}
pub async fn send_chunk(&self, chunk: P::Chunk) -> Result<Demand, P::Error> {
self.0
.request_reply(|reply| Suspension::Stream(StreamDelivery::Chunk(chunk, reply)))
.await
}
}
impl<R: Protocol> CallMachine<R>
type CallCoroutine<R, C> = Coroutine<Suspension<R>, Result<C, <R as Protocol>::Error>>;
pub struct CallMachine<R: Protocol, C = <R as Protocol>::Response> {
coroutine: CallCoroutine<R, C>,
}
impl<R: Protocol, C: Send + 'static> CallMachine<R, C>
where
R::Error: From<MachineFault>,
{
pub fn new(execute: impl FnOnce(HostChannel<R>) -> ExecuteFuture<R> + Send + 'static) -> Self {
pub fn new(
execute: impl FnOnce(CallContext<R>) -> ExecuteFuture<R, C> + Send + 'static,
) -> Self {
Self {
coroutine: Coroutine::new(|co| execute(HostChannel { co })),
coroutine: Coroutine::new(|co| {
let channel = Channel(co);
execute(CallContext {
services: HostServices(channel.clone()),
hooks: ChannelHooks(channel.clone()),
stream: StreamSender(channel),
})
}),
}
}
}
impl<R: Protocol> Machine for CallMachine<R>
impl<R: Protocol, C: Send + 'static> Machine for CallMachine<R, C>
where
R::Error: From<MachineFault>,
{
type Protocol = R;
type Complete = R::Response;
type Complete = C;
fn resume(&mut self) -> Step<'_, Self> {
Box::pin(async move {
@ -124,7 +172,7 @@ where
.await
.map_err(MachineFault::Protocol)?
{
CoroutineState::Yielded(op) => Ok(MachineStep::Host(op)),
CoroutineState::Yielded(op) => Ok(MachineStep::Suspended(op)),
CoroutineState::Complete(outcome) => outcome.map(MachineStep::Complete),
}
})

View file

@ -5,13 +5,14 @@ use std::future::Future;
use std::pin::Pin;
pub use auth::{HostTokenProvider, TokenProtocol};
pub use call_machine::{CallMachine, ExecuteFuture, HostChannel, MachineFault};
pub use call_machine::{
CallContext, CallMachine, ChannelHooks, ExecuteFuture, HostServices, MachineFault, StreamSender,
};
use crate::host::HostOp;
use crate::protocol::Protocol;
use crate::protocol::{Protocol, Suspension};
pub enum MachineStep<R: Protocol, C> {
Host(HostOp<R>),
Suspended(Suspension<R>),
Complete(C),
}

View file

@ -1,17 +1,38 @@
/// 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 use litellm_coroutine::{Abandoned, Answer, Reply, reply};
use crate::event::{MachineEvent, RequestContext, WireRequest};
pub trait Protocol: Send + Sync + 'static {
type Request: Send + '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 HostCall: Send + 'static;
type Chunk: Send + 'static;
/// What the call knows once a streamed response starts, before its first chunk.
type StreamHead: Send + 'static;
}
pub enum Suspension<P: Protocol> {
HostCall(P::HostCall),
Hook(HookRequest),
Stream(StreamDelivery<P>),
}
pub enum HookRequest {
BeforeProviderRequest {
wire: Box<WireRequest>,
context: Box<RequestContext>,
reply: Reply<WireRequest>,
},
Event(MachineEvent, Reply<()>),
}
pub enum StreamDelivery<P: Protocol> {
Open(P::StreamHead, Reply<Demand>),
Chunk(P::Chunk, Reply<Demand>),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Demand {
More,
Detached,
}

View file

@ -1,207 +0,0 @@
use crate::event::{CallEvent, FailureOrigin, Timing, epoch_seconds};
use crate::host::{Host, HostOp};
use crate::machine::{HostFailure, Machine, MachineStep};
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::Protocol as Protocol>::Error>
where
M: Machine,
H: Host<M::Protocol>,
{
let start_time = epoch_seconds();
let _ = host.emit(&CallEvent::Started { start_time }).await;
let outcome = loop {
let op = match machine.resume().await {
Ok(MachineStep::Complete(complete)) => break Ok(complete),
Ok(MachineStep::Host(op)) => op,
Err(error) => break Err(error),
};
if let Err(error) = perform(host, op).await {
break machine.interrupt(HostFailure::Error(error)).await;
}
};
let timing = Timing {
start_time,
end_time: epoch_seconds(),
};
let terminal = match &outcome {
Ok(_) => CallEvent::Succeeded { timing },
Err(_) => CallEvent::Failed {
timing,
origin: FailureOrigin::Call,
},
};
let _ = host.emit(&terminal).await;
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::host::Reply;
use crate::machine::{CallMachine, MachineFault};
struct Unit;
impl Protocol for Unit {
type Response = ();
type Error = &'static str;
type Projection = ();
type Op = (&'static str, Reply<()>);
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
impl From<MachineFault> for &'static str {
fn from(_: MachineFault) -> Self {
"machine fault"
}
}
#[derive(Default)]
struct Recording {
seen: Mutex<Vec<String>>,
fail: Option<&'static str>,
}
impl Host<Unit> for Recording {
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> {
self.seen.lock().unwrap().push(match event {
CallEvent::Started { .. } => "started".into(),
CallEvent::Succeeded { .. } => "succeeded".into(),
CallEvent::Failed { .. } => "failed".into(),
other => format!("{other:?}"),
});
Ok(())
}
}
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(&["sign", "send"], Ok(())), &host).await;
assert_eq!(outcome, Ok(()));
assert_eq!(
*host.seen.lock().unwrap(),
["started", "project", "op:sign", "op:send", "succeeded"]
);
}
#[tokio::test]
async fn errors_and_host_failures_each_emit_failed_once() {
let host = Recording::default();
let outcome = run(scripted(&[], Err("boom")), &host).await;
assert_eq!(outcome, Err("boom"));
assert_eq!(*host.seen.lock().unwrap(), ["started", "project", "failed"]);
let host = Recording {
fail: Some("send"),
..Recording::default()
};
let outcome = run(scripted(&["sign", "send", "never"], Ok(())), &host).await;
assert_eq!(outcome, Err("host failed"));
assert_eq!(
*host.seen.lock().unwrap(),
["started", "project", "op:sign", "op:send", "failed"]
);
}
struct StartTimes(Mutex<Vec<f64>>);
impl Host<Unit> for StartTimes {
async fn project(&self) -> Result<(), &'static str> {
Ok(())
}
async fn custom_op(
&self,
(_, reply): (&'static str, Reply<()>),
) -> Result<(), &'static str> {
reply.send(());
Ok(())
}
async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> {
if let CallEvent::Started { start_time }
| CallEvent::Succeeded {
timing: Timing { start_time, .. },
} = event
{
self.0.lock().unwrap().push(*start_time);
}
Err("observer failed")
}
}
#[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(&["send"], Ok(())), &host).await, Ok(()));
let times = host.0.lock().unwrap();
assert_eq!(times.len(), 2);
assert_eq!(times[0], times[1]);
}
}

View file

@ -0,0 +1,15 @@
use crate::protocol::Protocol;
use std::{convert::Infallible, future::Future};
pub trait HostCallHandler<P: Protocol>: Send + Sync {
fn handle_host_call(
&self,
call: P::HostCall,
) -> impl Future<Output = Result<(), P::Error>> + Send;
}
impl<P: Protocol<HostCall = Infallible>> HostCallHandler<P> for () {
async fn handle_host_call(&self, call: Infallible) -> Result<(), P::Error> {
match call {}
}
}

View file

@ -0,0 +1,224 @@
use litellm_host::protocol::StreamDelivery;
use std::{
convert::Infallible,
sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
},
};
use futures_util::{StreamExt, stream};
use litellm_host::{
call::{CallOutput, HostedCompletion, hosted_call},
event::CallEvent,
lifecycle::{CallObserver, observe_call, observe_unary},
machine::{Machine, MachineFault, MachineStep},
protocol::{Demand, Protocol, Suspension},
};
use rstest::{fixture, rstest};
#[derive(Debug, Clone)]
struct TestError;
struct TestProtocol;
impl Protocol for TestProtocol {
type Response = &'static str;
type Error = TestError;
type Request = usize;
type HostCall = Infallible;
type Chunk = usize;
type StreamHead = &'static str;
}
impl From<MachineFault> for TestError {
fn from(_: MachineFault) -> Self {
Self
}
}
#[rstest]
#[case::end(None, 3, HostedCompletion::StreamEnded)]
#[case::detach_at_open(Some(0), 0, HostedCompletion::Detached)]
#[case::detach_after_chunk(Some(1), 1, HostedCompletion::Detached)]
#[tokio::test]
async fn delivery_obeys_demand_and_distinguishes_detachment(
#[case] detach_after: Option<usize>,
#[case] expected_polls: usize,
#[case] expected: HostedCompletion<&'static str>,
) {
let polls = Arc::new(AtomicUsize::new(0));
let stream_polls = polls.clone();
let mut machine = hosted_call::<TestProtocol, _, _>(3, move |count, _, _| async move {
let chunks = stream::iter((0..count).map(Ok))
.inspect(move |_| {
stream_polls.fetch_add(1, Ordering::SeqCst);
})
.boxed();
Ok(CallOutput::Stream {
head: "headers",
chunks,
})
});
let MachineStep::Suspended(Suspension::Stream(StreamDelivery::Open(head, reply))) =
machine.resume().await.unwrap()
else {
panic!()
};
assert_eq!(head, "headers");
assert_eq!(polls.load(Ordering::SeqCst), 0);
reply.send(if detach_after == Some(0) {
Demand::Detached
} else {
Demand::More
});
let mut delivered = Vec::new();
let completed = loop {
match machine.resume().await.unwrap() {
MachineStep::Suspended(Suspension::Stream(StreamDelivery::Chunk(chunk, reply))) => {
delivered.push(chunk);
assert_eq!(polls.load(Ordering::SeqCst), delivered.len());
reply.send(if detach_after == Some(delivered.len()) {
Demand::Detached
} else {
Demand::More
});
}
MachineStep::Complete(result) => break result,
_ => panic!("unexpected operation"),
}
};
assert_eq!(completed, expected);
assert_eq!(polls.load(Ordering::SeqCst), expected_polls);
assert_eq!(delivered, (0..expected_polls).collect::<Vec<_>>());
}
#[derive(Default)]
struct Observer(Mutex<Vec<CallEvent>>);
impl CallObserver for Observer {
fn observe(&self, event: CallEvent) {
self.0.lock().unwrap().push(event);
}
}
#[fixture]
fn observer() -> Arc<Observer> {
Arc::new(Observer::default())
}
type Output = CallOutput<(), (), usize, &'static str>;
#[rstest]
#[case::success(false)]
#[case::failure(true)]
#[tokio::test]
async fn unary_calls_emit_one_terminal_event(observer: Arc<Observer>, #[case] fail: bool) {
let expected = if fail { Err("provider") } else { Ok(7) };
assert_eq!(
observe_unary(Some(observer.clone()), async { expected }).await,
expected
);
let events = observer.0.lock().unwrap();
assert_eq!(events.len(), 2);
assert!(matches!(events[0], CallEvent::Started { .. }));
assert_eq!(matches!(events[1], CallEvent::Failed { .. }), fail);
assert_eq!(matches!(events[1], CallEvent::Succeeded { .. }), !fail);
}
#[rstest]
#[case::end(false)]
#[case::error(true)]
#[tokio::test]
async fn streams_finish_only_when_consumed(observer: Arc<Observer>, #[case] fail: bool) {
let chunks = stream::iter([Ok(1), if fail { Err("provider") } else { Ok(2) }]).boxed();
let output = observe_call(Some(observer.clone()), async {
Ok::<Output, _>(CallOutput::Stream { head: (), chunks })
})
.await
.unwrap();
assert_eq!(observer.0.lock().unwrap().len(), 1);
let CallOutput::Stream { mut chunks, .. } = output else {
panic!()
};
assert_eq!(chunks.next().await, Some(Ok(1)));
assert_eq!(observer.0.lock().unwrap().len(), 1);
let last = chunks.next().await;
if fail {
assert_eq!(last, Some(Err("provider")));
} else {
assert_eq!(last, Some(Ok(2)));
assert_eq!(chunks.next().await, None);
}
drop(chunks);
let events = observer.0.lock().unwrap();
assert_eq!(events.len(), 2);
assert_eq!(matches!(events[1], CallEvent::Failed { .. }), fail);
assert_eq!(matches!(events[1], CallEvent::Succeeded { .. }), !fail);
}
#[rstest]
#[tokio::test]
async fn dropping_a_stream_cancels_without_success(observer: Arc<Observer>) {
let chunks = stream::pending().boxed();
let output = observe_call(Some(observer.clone()), async {
Ok::<Output, _>(CallOutput::Stream { head: (), chunks })
})
.await
.unwrap();
drop(output);
let events = observer.0.lock().unwrap();
assert_eq!(events.len(), 2);
assert!(matches!(events[1], CallEvent::Cancelled { .. }));
}
#[rstest]
#[tokio::test]
async fn cancelling_provider_execution_releases_the_lifecycle(observer: Arc<Observer>) {
let mut call = Box::pin(observe_unary(
Some(observer.clone()),
std::future::pending::<Result<(), ()>>(),
));
assert!(futures_util::poll!(&mut call).is_pending());
drop(call);
let events = observer.0.lock().unwrap();
assert_eq!(events.len(), 2);
assert!(matches!(events[1], CallEvent::Cancelled { .. }));
}
#[rstest]
#[tokio::test]
async fn a_hosted_detachment_is_reported_as_cancellation(observer: Arc<Observer>) {
struct DetachingConsumer;
impl litellm_host::in_process::StreamConsumer<TestProtocol> for DetachingConsumer {
async fn open_stream(&self, _: &'static str) -> Result<Demand, TestError> {
Ok(Demand::Detached)
}
async fn send_chunk(&self, _: usize) -> Result<Demand, TestError> {
panic!("detached consumers must not receive chunks")
}
}
let machine = hosted_call::<TestProtocol, _, _>(1, |count, _, _| async move {
Ok(CallOutput::Stream {
head: "headers",
chunks: stream::iter((0..count).map(Ok)).boxed(),
})
});
let completion = litellm_host::in_process::run_hosted(
machine,
litellm_host::in_process::Host {
services: &(),
hooks: &(),
stream: &DetachingConsumer,
observer: Some(observer.as_ref()),
},
)
.await
.unwrap();
assert_eq!(completion, HostedCompletion::Detached);
assert!(matches!(
&observer.0.lock().unwrap()[..],
[CallEvent::Started { .. }, CallEvent::Cancelled { .. }]
));
}

View file

@ -26,7 +26,7 @@ use litellm_secrets::source::SecretSource;
/// The route's view of one call, handed to provider code that has to reach the
/// caller's hooks mid-flight (guardrails on the outgoing body, raw response events).
pub trait CallHooks<E>: Send + Sync {
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, E>>;
fn before_provider_request(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, E>>;
fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), E>>;
}
@ -226,7 +226,7 @@ pub async fn transform_request_body<C: BaseOcrConfig, B: Serialize>(
)?;
config.validate_request_body(&composed)?;
let changed = hooks
.before_send(wire_request(url, headers, composed))
.before_provider_request(wire_request(url, headers, composed))
.await?;
if !changed.body.is_object() {
return Err(Error::RequestField {
@ -278,7 +278,9 @@ pub async fn guardrail_document(
let body = serde_json::to_value(&request.document).map_err(|_| Error::RequestField {
path: "document".into(),
})?;
let changed = hooks.before_send(wire_request(url, headers, body)).await?;
let changed = hooks
.before_provider_request(wire_request(url, headers, body))
.await?;
let document = decode_request_value(changed.body, "guardrail.document")?;
Ok((document, changed.headers))
}
@ -304,6 +306,7 @@ mod tests {
use super::*;
#[rstest::rstest]
#[tokio::test]
async fn request_timeout_has_an_http_408_status() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();

View file

@ -7,7 +7,7 @@
- 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 and tokio invariants, each pinned by a test in `host-python` (`runtime.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

View file

@ -207,8 +207,5 @@ impl ExecutionBody for SemanticExecution {
}
pub(super) fn drive(py: Python<'_>, body: SemanticExecution) -> PyResult<Bound<'_, PyAny>> {
let execution = Py::new(py, Execution::new(body))?;
py.import("litellm.rust_bridge.lifecycle")?
.getattr("drive")?
.call1((execution,))
Execution::new(body).into_coroutine(py)
}

View file

@ -117,6 +117,15 @@ pub(crate) fn call_config(
Ok(resolution.config)
}
pub(crate) fn provider_client(
py: Python<'_>,
kwargs: &Bound<'_, PyDict>,
asynchronous: bool,
) -> PyResult<Result<Client, litellm_http::Error>> {
let config = call_config(py, kwargs, asynchronous)?;
Ok(pool().client(&config, ClientVariant::Provider))
}
pub(crate) fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult<Client> {
let config = call_config(py, &PyDict::new(py), true)?;
pool().client(&config, variant).map_err(client_error)

View file

@ -12,8 +12,8 @@ struct DiagnosticMachine;
impl Protocol for DiagnosticMachine {
type Response = ();
type Error = String;
type Projection = ();
type Op = ();
type Request = ();
type HostCall = ();
type Chunk = ();
type StreamHead = ();
}

View file

@ -1,13 +1,10 @@
use crate::logger::{run_async, run_sync};
use litellm_core::audio_transcription::{
Error, audio_transcription as run_audio_transcription, types::AudioTranscriptionRequest,
AudioTranscriptionRoute, Error, types::AudioTranscriptionRequest,
};
use litellm_host_python::from_py_argument;
use litellm_http::HttpClientConfig;
use litellm_secrets::source::SecretSource;
use pyo3::{prelude::*, types::PyDict};
use serde_json::{Map, Value};
use std::sync::Arc;
use crate::{
errors::route_error_to_pyerr,
@ -15,8 +12,8 @@ use crate::{
};
async fn execute(
config: HttpClientConfig,
secrets: Arc<dyn SecretSource>,
http: Result<litellm_http::Client, litellm_http::Error>,
secrets: std::sync::Arc<dyn litellm_secrets::source::SecretSource>,
audio: Value,
optional_params: Map<String, Value>,
options: RouteOptions,
@ -29,11 +26,8 @@ async fn execute(
extra_headers,
timeout,
} = options;
run_audio_transcription(
crate::http::resources(),
&config,
secrets.as_ref(),
AudioTranscriptionRequest {
AudioTranscriptionRoute::new(http?, crate::http::resources().auth.clone(), secrets)
.execute(AudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
@ -42,9 +36,8 @@ async fn execute(
extra_headers,
optional_params,
timeout,
},
)
.await
})
.await
}
#[pyfunction]
@ -72,12 +65,13 @@ pub(crate) fn transcription(
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let config = crate::http::call_config(py, &PyDict::new(py), false)?;
let http = crate::http::provider_client(py, &PyDict::new(py), false)?;
let secrets = crate::secrets::source(py)?;
run_sync(
py,
execute(
config,
crate::secrets::source(py)?,
http,
secrets,
audio,
optional_params.unwrap_or_default(),
options,
@ -111,12 +105,13 @@ pub(crate) fn atranscription<'py>(
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let config = crate::http::call_config(py, &PyDict::new(py), true)?;
let http = crate::http::provider_client(py, &PyDict::new(py), true)?;
let secrets = crate::secrets::source(py)?;
run_async(
py,
execute(
config,
crate::secrets::source(py)?,
http,
secrets,
audio,
optional_params.unwrap_or_default(),
options,

View file

@ -1,15 +1,11 @@
use litellm_secrets::source::SecretSource;
use pyo3::types::{PyDict, PyTuple};
use std::sync::Arc;
use crate::errors::RustBridgeDeclined;
use crate::logger::{run_async, run_sync};
use litellm_core::chat_completions::{
Error, chat_completions as run_chat_completions, chat_completions_decline_reason,
types::ChatCompletionsRequest,
ChatCompletionsRoute, Error, chat_completions_decline_reason, types::ChatCompletionsRequest,
};
use litellm_host_python::from_py_argument;
use litellm_http::HttpClientConfig;
use litellm_types::utils::ChatCompletionsResponse;
use pyo3::prelude::*;
use serde_json::{Map, Value};
@ -23,8 +19,8 @@ use crate::{
};
async fn execute(
config: HttpClientConfig,
secrets: Arc<dyn SecretSource>,
http: Result<litellm_http::Client, litellm_http::Error>,
secrets: std::sync::Arc<dyn litellm_secrets::source::SecretSource>,
messages: Vec<Value>,
optional_params: Map<String, Value>,
options: RouteOptions,
@ -37,22 +33,21 @@ async fn execute(
extra_headers,
timeout,
} = options;
run_chat_completions(
crate::http::resources(),
&config,
secrets.as_ref(),
ChatCompletionsRequest {
model: &model,
messages: Value::Array(messages),
optional_params,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
},
)
.await
ChatCompletionsRoute::new(http?, crate::http::resources().auth.clone(), secrets)
.execute(
ChatCompletionsRequest {
model: &model,
messages: Value::Array(messages),
optional_params,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
},
&(),
)
.await
}
#[pyfunction]
@ -97,12 +92,13 @@ pub(crate) fn chat_completions(
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let config = crate::http::call_config(py, &PyDict::new(py), false)?;
let http = crate::http::provider_client(py, &PyDict::new(py), false)?;
let secrets = crate::secrets::source(py)?;
run_sync(
py,
execute(
config,
crate::secrets::source(py)?,
http,
secrets,
messages,
optional_params.unwrap_or_default(),
options,
@ -136,12 +132,13 @@ pub(crate) fn achat_completions<'py>(
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let config = crate::http::call_config(py, &PyDict::new(py), true)?;
let http = crate::http::provider_client(py, &PyDict::new(py), true)?;
let secrets = crate::secrets::source(py)?;
run_async(
py,
execute(
config,
crate::secrets::source(py)?,
http,
secrets,
messages,
optional_params.unwrap_or_default(),
options,

View file

@ -1,11 +1,12 @@
use litellm_host_python::{PythonHostCalls, PythonOwned};
use std::convert::Infallible;
use bytes::Bytes;
use litellm_core::messages::{
Error, MessagesCall, MessagesShaping, messages_body,
route::{Messages, MessagesOutput, MessagesStreamHead},
route::{Messages, MessagesStreamHead},
};
use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py};
use litellm_host_python::{InvokeError, PythonBinding, from_py, lookup, to_py};
use litellm_http::transport::Error as TransportError;
use litellm_types::utils::ProviderSpecificHeaders;
use pyo3::{
@ -221,11 +222,11 @@ impl MessagesPythonHost {
}
}
impl ProtocolHost for MessagesPythonHost {
impl PythonBinding for MessagesPythonHost {
type Protocol = Messages;
type Failure = PyErr;
fn project(
fn decode_request(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
@ -235,33 +236,35 @@ impl ProtocolHost for MessagesPythonHost {
.map_err(InvokeError::Native)
}
fn invoke(&mut self, _: Python<'_>, op: Infallible) -> Result<(), InvokeError<Error>> {
match op {}
fn encode_response(
&mut self,
py: Python<'_>,
response: Box<
litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
>,
) -> PyResult<Py<PyAny>> {
py.import(ROUTE_HOST_MODULE)?
.getattr("response")?
.call1((to_py(py, response.as_ref())?,))
.map(Bound::unbind)
}
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {
match response {
MessagesOutput::Message(message) => py
.import(ROUTE_HOST_MODULE)?
.getattr("response")?
.call1((to_py(py, message.as_ref())?,))
.map(Bound::unbind),
MessagesOutput::Streamed => Ok(py.None()),
}
}
fn head(&mut self, py: Python<'_>, head: MessagesStreamHead) -> PyResult<Py<PyAny>> {
fn encode_stream_head(
&mut self,
py: Python<'_>,
head: MessagesStreamHead,
) -> PyResult<Py<PyAny>> {
py.import(ROUTE_HOST_MODULE)?
.getattr("stream_hidden_params")?
.call1((to_py(py, &head.headers)?,))
.map(Bound::unbind)
}
fn chunk(&mut self, py: Python<'_>, chunk: Bytes) -> PyResult<Py<PyAny>> {
fn encode_chunk(&mut self, py: Python<'_>, chunk: Bytes) -> PyResult<Py<PyAny>> {
Ok(PyBytes::new(py, &chunk).into_any().unbind())
}
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
fn map_error(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
if let Error::Secret(source) = &error
&& let Some(original) = crate::secrets::python_error(py, source.source_error())
{
@ -273,9 +276,20 @@ impl ProtocolHost for MessagesPythonHost {
fn host_error(error: &PyErr) -> Error {
Error::InvalidRequest(error.to_string().into())
}
}
impl PythonHostCalls<Messages> for MessagesPythonHost {
fn handle_host_call(
&mut self,
_: Python<'_>,
op: Infallible,
) -> Result<(), InvokeError<Error>> {
match op {}
}
}
impl PythonOwned for MessagesPythonHost {
fn close(&mut self, _: Python<'_>) {}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.request)
}

View file

@ -4,7 +4,6 @@ use host::MessagesPythonHost;
use litellm_callbacks_legacy_python::{
LegacySurface, PassThroughStream, PublicCall, run_legacy_call,
};
use litellm_core::messages::route::messages_machine;
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
@ -26,15 +25,17 @@ fn run_messages(
kwargs: Bound<'_, PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let secrets = crate::secrets::source(py)?;
let config = crate::http::call_config(py, &kwargs, asynchronous)?;
let machine = messages_machine(crate::http::resources(), &config, secrets)
.map_err(crate::http::client_error)?;
let route = litellm_core::messages::MessagesRoute::new(
crate::http::provider_client(py, &kwargs, asynchronous)?
.map_err(crate::http::client_error)?,
crate::http::resources().auth.clone(),
crate::secrets::source(py)?,
);
run_legacy_call(
py,
SURFACE,
PublicCall::capture(&request, &args, &kwargs)?,
crate::logger::LoggedMachine::new(machine),
move |request| crate::logger::LoggedMachine::new(route.machine(request)),
MessagesPythonHost::new(request.unbind()),
crate::preflight::sdk_preflight,
asynchronous,

View file

@ -1,6 +1,7 @@
use litellm_auth::ResolvedCredential;
use litellm_core::ocr::route::{Ocr, OcrOp, OcrProjection};
use litellm_host_python::{InvokeError, ProtocolHost, missing_state, to_py};
use litellm_core::ocr::route::{Ocr, OcrCall, OcrOp};
use litellm_host_python::{InvokeError, PythonBinding, missing_state, to_py};
use litellm_host_python::{PythonHostCalls, PythonOwned};
use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse};
use pyo3::{
exceptions::{PyBaseException, PyException},
@ -51,18 +52,14 @@ impl OcrPythonHost {
.acquire(py)
}
fn projection(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<OcrProjection> {
fn projection(&mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<OcrCall> {
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 {
Ok(OcrCall {
request,
caller_token,
})
@ -88,44 +85,47 @@ impl OcrPythonHost {
}
}
impl ProtocolHost for OcrPythonHost {
impl PythonBinding for OcrPythonHost {
type Protocol = Ocr;
type Failure = PyErr;
fn project(
fn decode_request(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> Result<OcrProjection, InvokeError<Error>> {
) -> Result<OcrCall, 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>> {
fn encode_response(
&mut self,
py: Python<'_>,
response: LiteLLMOcrResponse,
) -> PyResult<Py<PyAny>> {
py.import("litellm.rust_bridge.ocr.route_host")?
.getattr("response")?
.call1((to_py(py, &response)?,))
.map(Bound::unbind)
}
fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult<Py<PyAny>> {
fn encode_stream_head(
&mut self,
_: Python<'_>,
head: std::convert::Infallible,
) -> PyResult<Py<PyAny>> {
match head {}
}
fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult<Py<PyAny>> {
fn encode_chunk(
&mut self,
_: Python<'_>,
chunk: std::convert::Infallible,
) -> PyResult<Py<PyAny>> {
match chunk {}
}
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
fn map_error(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
if let Error::Secret(source) = &error
&& let Some(original) = crate::secrets::python_error(py, source)
{
@ -137,11 +137,23 @@ impl ProtocolHost for OcrPythonHost {
fn host_error(error: &PyErr) -> Error {
Error::InvalidRequest(error.to_string())
}
}
impl PythonHostCalls<Ocr> for OcrPythonHost {
fn handle_host_call(&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))),
}
}
}
impl PythonOwned for OcrPythonHost {
fn close(&mut self, _: Python<'_>) {
self.data = OcrHostData::Released;
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.request)?;
if let OcrHostData::Projected(handles) = &self.data
@ -199,12 +211,13 @@ del provider
.cast_into::<PyDict>()
.unwrap();
let mut host = OcrPythonHost::new(py.None());
assert!(host.project(py, &kwargs).unwrap().caller_token);
assert!(host.decode_request(py, &kwargs).unwrap().caller_token);
locals.del_item("kwargs").unwrap();
drop(kwargs);
let (reply, _) = litellm_host::host::reply();
let (reply, _) = litellm_host::protocol::reply();
assert_eq!(
host.invoke(py, OcrOp::AcquireAzureAdToken(reply)).is_ok(),
host.handle_host_call(py, OcrOp::AcquireAzureAdToken(reply))
.is_ok(),
succeeds
);
let alive = || {

View file

@ -5,7 +5,7 @@ mod project;
use host::OcrPythonHost;
use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call};
use litellm_core::ocr::{provider_config, route::ocr_machine};
use litellm_core::ocr::provider_config;
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_host_python::to_py;
use litellm_llms::base_llm::ocr::settings::OcrSettings;
@ -18,7 +18,6 @@ use crate::{
coercion::FieldSpec,
http,
python_settings::{PythonSettings, Snapshot},
secrets,
};
const VERTEX_PROJECT: FieldSpec<Option<String>> =
@ -48,16 +47,22 @@ fn run_ocr(
kwargs: Bound<'_, PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let secrets = secrets::source(py)?;
let config = http::call_config(py, &kwargs, asynchronous)?;
let client = http::resources()
.ocr_client(&config, http::url_policy(py)?, ocr_settings(py)?, secrets)
.map_err(http::client_error)?;
let client = litellm_llms::base_llm::ocr::handler::OcrClient::new(
&http::resources().pool,
&config,
http::url_policy(py)?,
http::resources().auth.clone(),
ocr_settings(py)?,
crate::secrets::source(py)?,
)
.map_err(http::client_error)?;
let route = litellm_core::ocr::OcrRoute::new(client);
run_legacy_call(
py,
if asynchronous { ASYNC_SURFACE } else { SURFACE },
PublicCall::capture(&request, &args, &kwargs)?,
crate::logger::LoggedMachine::new(ocr_machine(client)),
move |request| crate::logger::LoggedMachine::new(route.machine(request)),
OcrPythonHost::new(request.unbind()),
crate::preflight::sdk_preflight,
asynchronous,
@ -127,7 +132,7 @@ mod tests {
use crate::python_settings::PythonSettings;
#[test]
#[rstest::rstest]
fn provider_defaults_distinguish_falsey_values_and_exact_true() {
Python::initialize();
Python::attach(|py| {

View file

@ -12,7 +12,7 @@ mod vault;
use std::sync::Arc;
pub(crate) use error::python_error;
use litellm_secrets::source::SecretSource;
use litellm_secrets::source::{EnvironmentSecrets, SecretSource};
use pyo3::prelude::*;
use python::PythonSecrets;
use resolved::ResolvedSecrets;
@ -23,15 +23,15 @@ const NATIVE: FieldSpec<bool> = FieldSpec::new("native", |field| field.schema_bo
/// 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.
/// native secret manager. A bare extension module without the litellm package reads the process
/// environment.
pub(crate) fn source(py: Python<'_>) -> PyResult<Arc<dyn SecretSource>> {
let Some(settings) = PythonSettings::SecretManager.read_or_unset(py)? else {
let client = crate::http::host_client(py, litellm_http::ClientVariant::Provider)?;
return Ok(Arc::new(
litellm_secrets::source::EnvironmentSecrets::python_compatible(client),
));
let Some(snapshot) = PythonSettings::SecretManager.read_or_unset(py)? else {
return Ok(Arc::new(EnvironmentSecrets::python_compatible(
crate::http::host_client(py, litellm_http::ClientVariant::Provider)?,
)));
};
if settings.read(&NATIVE)? {
if snapshot.read(&NATIVE)? {
let context = litellm_host_python::PythonContext::capture(py)?;
let client = crate::http::host_client(py, litellm_http::ClientVariant::Provider)?;
return Ok(Arc::new(ResolvedSecrets::new(