refactor(rust): centralize host execution and compose callbacks (#43515)

* refactor(rust): extract litellm-host-native as the shared Rust host driver

Move service and hook dispatch out of host-http into a Driver that owns the
machine and Rust handlers, returning at completion or a stream boundary and
holding the demand reply until the consumer advances. Move the in-process
runner onto the same driver. host-http now layers encoding, SSE, body polling
and lifecycle observation over it. host-python keeps driving litellm-host
directly

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

* fix(rust): interrupt the machine when the in-process stream consumer fails

Restores the pre-refactor interruption path for StreamConsumer errors via
Driver::fail and ports the generic run lifecycle tests into host-native.

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

* refactor(rust): separate the machine contract from coroutine execution

* auth update

* refactor(rust): use standard flow control for host requests

* style(rust): keep host driver imports formatted

* chores

* mostly relocation

* refactor(rust): separate interceptors from queued observers

* refactor(rust): centralize legacy callback mappings and lifecycle

* docs: define Python host boundaries and migration plan

* refactor: enforce Python host and bridge boundaries

* refactor(rust): separate operations from callback composition

* refactor(rust): compose SDK policy through call hooks

---------

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-28 19:20:27 +00:00 • committed by GitHub
parent 37be82e45e
commit e4190d86a6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
117 changed files with 6181 additions and 2191 deletions

View file

@ -3616,6 +3616,7 @@ dependencies = [
"litellm-auth",
"litellm-host",
"litellm-host-python",
"litellm-types",
"proptest",
"pyo3",
"rstest",
@ -3647,6 +3648,7 @@ dependencies = [
"litellm-auth-gcp",
"litellm-core-utils",
"litellm-host",
"litellm-host-native",
"litellm-http",
"litellm-llms",
"litellm-secrets",
@ -3907,12 +3909,24 @@ dependencies = [
"futures-util",
"http 1.4.2",
"litellm-host",
"litellm-host-native",
"rstest",
"serde_json",
"thiserror 2.0.19",
"tokio",
]
[[package]]
name = "litellm-host-native"
version = "0.1.0"
dependencies = [
"futures-util",
"litellm-host",
"rstest",
"serde_json",
"tokio",
]
[[package]]
name = "litellm-host-python"
version = "0.1.0"

View file

@ -22,6 +22,7 @@ litellm-gateway-ui = { path = "crates/gateway-ui" }
litellm-coroutine = { path = "crates/coroutine" }
litellm-host = { path = "crates/host" }
litellm-host-http = { path = "crates/host-http" }
litellm-host-native = { path = "crates/host-native" }
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
litellm-framing = { path = "crates/framer" }
litellm-auth = { path = "crates/auth" }

View file

@ -39,7 +39,33 @@ impl TokenProviderHandle {
Self(caller)
}
pub fn from_callback<F, Fut>(acquire: F) -> Self
where
F: Fn() -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<ResolvedCredential, Error>> + Send + 'static,
{
Self::new(Arc::new(CallbackTokenProvider(acquire)))
}
pub async fn acquire(&self) -> Result<ResolvedCredential, Error> {
self.0.acquire().await
}
}
struct CallbackTokenProvider<F>(F);
impl<F> std::fmt::Debug for CallbackTokenProvider<F> {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("CallbackTokenProvider")
}
}
impl<F, Fut> TokenProvider for CallbackTokenProvider<F>
where
F: Fn() -> Fut + Send + Sync,
Fut: Future<Output = Result<ResolvedCredential, Error>> + Send + 'static,
{
fn acquire(&self) -> TokenFuture<'_> {
Box::pin((self.0)())
}
}

View file

@ -0,0 +1,104 @@
use std::{
error::Error as StdError,
future::{Future, poll_fn},
sync::{
Arc,
atomic::{AtomicBool, AtomicUsize, Ordering},
},
task::Poll,
time::{Duration, SystemTime},
};
use litellm_auth_types::{
Error, ErrorDetail, ResolvedCredential, SecretValue, TokenProviderHandle,
};
use rstest::rstest;
fn credential(index: usize, access_token: bool) -> ResolvedCredential {
let token = SecretValue::new(format!("credential-{index}"));
if access_token {
return ResolvedCredential::AccessToken {
token,
expires_on: Some(SystemTime::UNIX_EPOCH + Duration::from_secs(index as u64)),
};
}
ResolvedCredential::Static(token)
}
#[rstest]
#[case::static_secret(false)]
#[case::access_token(true)]
#[tokio::test]
async fn callbacks_acquire_fresh_credentials_on_demand(#[case] access_token: bool) {
let calls = Arc::new(AtomicUsize::new(0));
let callback_calls = calls.clone();
let provider = TokenProviderHandle::from_callback(move || {
let index = callback_calls.fetch_add(1, Ordering::SeqCst);
async move {
tokio::task::yield_now().await;
Ok(credential(index, access_token))
}
});
let cloned = provider.clone();
assert_eq!(calls.load(Ordering::SeqCst), 0);
assert_eq!(
provider.acquire().await.unwrap(),
credential(0, access_token)
);
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert_eq!(cloned.acquire().await.unwrap(), credential(1, access_token));
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[rstest]
#[tokio::test]
async fn callback_errors_preserve_the_original_source() {
let provider = TokenProviderHandle::from_callback(|| async {
Err(Error::CredentialAcquisition(ErrorDetail::failed(
"caller credential",
std::io::Error::from(std::io::ErrorKind::PermissionDenied),
)))
});
let error = provider.acquire().await.unwrap_err();
assert!(matches!(error, Error::CredentialAcquisition(_)));
let source = std::iter::successors(Some(&error as &(dyn StdError + 'static)), |error| {
(*error).source()
})
.find_map(|error| error.downcast_ref::<std::io::Error>())
.unwrap();
assert_eq!(source.kind(), std::io::ErrorKind::PermissionDenied);
}
struct Release(Arc<AtomicBool>);
impl Drop for Release {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
#[rstest]
#[tokio::test]
async fn cancelling_acquisition_drops_the_callback_future() {
let released = Arc::new(AtomicBool::new(false));
let callback_released = released.clone();
let provider = TokenProviderHandle::from_callback(move || {
let released = callback_released.clone();
async move {
let _release = Release(released);
std::future::pending().await
}
});
let mut acquisition = Box::pin(provider.acquire());
poll_fn(|context| {
assert!(acquisition.as_mut().poll(context).is_pending());
assert!(!released.load(Ordering::SeqCst));
Poll::Ready(())
})
.await;
drop(acquisition);
assert!(released.load(Ordering::SeqCst));
}

View file

@ -1,12 +1,12 @@
- Target invariants, not completion claims
- 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)
- This crate owns compatibility for all existing Python callbacks and loggers, including `CustomLogger`. `mapping.rs` owns the executable call bindings and the inventory of Python-owned hooks. A Python-owned entry records an existing path, never permission to invoke it a second time. The native call adapter preserves 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 `PythonCallHooks`; they never learn which Python objects consume a call
- SDK request policy (credential inheritance, the budget and retry-count limits) is a separate hook supplied by `python-bridge`; compose it after this adapter so logging adopts the final keyword view before policy mutates or rejects it
- The driver in `litellm-host-python`, the routes and core see one `PythonCallHooks` using the shared `CallEvent`; 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`
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; routes hand it over through `run_legacy_call` and keep no copy
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; shared bridge composition hands it to `LegacyLogging`; routes use the neutral call boundary
- `setup` reuses a `Logging` passed as `litellm_logging_obj` (the proxy and Router) and otherwise builds one through `function_setup`; which callbacks run is `Logging`'s decision, never this crate's
- Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation
- Retain complete boundary arguments, opaque values, aliases, omitted/default distinctions and deliberate copies; preserve the deployment-hook kwargs view

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
litellm-types.workspace = true
litellm-host.workspace = true
litellm-host-python.workspace = true

View file

@ -3,11 +3,13 @@
//! `@client` path makes them.
use litellm_host_python::PythonOwned;
use litellm_types::Operation;
use litellm_host::event::{
FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest, epoch_seconds,
use litellm_host::{
interceptors::{RawResponse, RequestContext, WireRequest},
lifecycle::{FailureOrigin, Timing, epoch_seconds},
};
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, from_py, missing_state, to_py};
use litellm_host_python::{HookStep, from_py, missing_state, to_py};
use pyo3::{
exceptions::{PyBaseException, PyException},
gc::{PyTraverseError, PyVisit},
@ -24,22 +26,10 @@ use crate::{
setup,
};
/// What the legacy contract needs to know about the route it is logging.
#[derive(Clone, Copy, Debug)]
pub struct LegacySurface {
pub call_type: &'static str,
/// What `Logging.pre_call` is told the input was.
pub input_description: &'static str,
/// How a streamed response is billed; `None` for a route that never streams.
pub stream: Option<PassThroughStream>,
}
/// The pass-through billing a streamed response goes through once its chunks are in.
#[derive(Clone, Copy, Debug)]
pub struct PassThroughStream {
pub url_route: &'static str,
/// A value of Python's `EndpointType`.
pub endpoint_type: &'static str,
struct PassThroughStream {
url_route: &'static str,
endpoint_type: &'static str,
}
/// What the Messages stream iterator keeps for its end-of-stream billing.
@ -55,7 +45,7 @@ struct LoggedRequest {
}
pub struct LegacyLogging {
surface: LegacySurface,
operation: Operation,
call: PublicCall,
logger: Option<PythonLogger>,
start: Py<PyAny>,
@ -77,14 +67,9 @@ fn is_cancellation(py: Python<'_>, error: &PyErr) -> bool {
}
impl LegacyLogging {
pub fn new(
py: Python<'_>,
surface: LegacySurface,
call: PublicCall,
asynchronous: bool,
) -> Self {
pub fn new(py: Python<'_>, operation: Operation, call: PublicCall, asynchronous: bool) -> Self {
Self {
surface,
operation,
call,
logger: None,
start: py.None(),
@ -98,6 +83,41 @@ impl LegacyLogging {
}
}
fn call_type(&self) -> &'static str {
match (self.operation, self.asynchronous) {
(Operation::Completion, false) => "completion",
(Operation::Completion, true) => "acompletion",
(Operation::Responses, false) => "responses",
(Operation::Responses, true) => "aresponses",
(Operation::Messages, _) => "anthropic_messages",
(Operation::Ocr, false) => "ocr",
(Operation::Ocr, true) => "aocr",
}
}
fn input_description(&self) -> &'static str {
match self.operation {
Operation::Completion => "Chat completions",
Operation::Responses => "Responses",
Operation::Messages => "Messages",
Operation::Ocr => "OCR document processing",
}
}
fn stream_billing(&self) -> Option<PassThroughStream> {
match self.operation {
Operation::Messages => Some(PassThroughStream {
url_route: "/v1/messages",
endpoint_type: "anthropic",
}),
Operation::Completion | Operation::Responses | Operation::Ocr => None,
}
}
pub(crate) fn adopt_arguments(&mut self, py: Python<'_>, arguments: &Py<PyDict>) {
self.call.set_kwargs(arguments.clone_ref(py));
}
/// Deployment hooks are awaited, and Python's synchronous `@client` wrapper never
/// runs them.
fn runs_deployment_hooks(&self) -> bool {
@ -112,7 +132,6 @@ 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<HookStep<Self, Py<PyDict>>> {
let prepared = self.call.kwargs().bind(py).copy()?;
prepared.set_item("litellm_logging_obj", self.logger()?.object(py))?;
@ -181,7 +200,7 @@ impl LegacyLogging {
fn stream_success(&self, py: Python<'_>, stream: &DeliveredStream) -> PyResult<()> {
let logger = self.logger()?;
let billing = self.surface.stream.ok_or_else(missing_state)?;
let billing = self.stream_billing().ok_or_else(missing_state)?;
let billed = Streaming::Success.call(
py,
(
@ -211,9 +230,12 @@ impl LegacyLogging {
/// 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<HookStep<Self, ()>> {
let (Some(logger), Some(error), Some(stream), Some(billing)) =
(&self.logger, &self.error, &self.stream, self.surface.stream)
else {
let (Some(logger), Some(error), Some(stream), Some(billing)) = (
&self.logger,
&self.error,
&self.stream,
self.stream_billing(),
) else {
return Ok(HookStep::Ready(()));
};
if !self.asynchronous {
@ -309,8 +331,8 @@ impl LegacyLogging {
}
}
impl PythonCallHooks for LegacyLogging {
fn prepare_arguments(
impl LegacyLogging {
pub(crate) fn prepare_call(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
@ -321,7 +343,7 @@ impl PythonCallHooks for LegacyLogging {
self.internal = is_internal_call(py)?;
let result = setup(
py,
self.surface.call_type,
self.call_type(),
self.call.args(),
self.call.kwargs(),
&self.start,
@ -331,14 +353,14 @@ impl PythonCallHooks for LegacyLogging {
self.call.set_kwargs(result.kwargs()?);
if self.runs_deployment_hooks() {
return Ok(HookStep::Await(
DeploymentHooks::before_call(py, self.call.kwargs(), self.surface.call_type)?,
DeploymentHooks::before_call(py, self.call.kwargs(), self.call_type())?,
Self::resume_begin,
));
}
self.prepare(py)
}
fn before_provider_request(
pub(crate) fn pre_call(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
@ -367,7 +389,7 @@ impl PythonCallHooks for LegacyLogging {
});
self.logger()?.pre_call(
py,
self.surface.input_description,
self.input_description(),
context.api_key.as_ref().map(|api_key| api_key.expose()),
&body,
&headers,
@ -384,7 +406,7 @@ impl PythonCallHooks for LegacyLogging {
})))
}
fn transform_response(
pub(crate) fn transform_public_response(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
@ -398,7 +420,7 @@ impl PythonCallHooks for LegacyLogging {
py,
self.call.kwargs(),
&self.response,
self.surface.call_type,
self.call_type(),
)?,
Self::resume_after_success,
));
@ -406,65 +428,65 @@ impl PythonCallHooks for LegacyLogging {
self.finalize(py)
}
fn on_event(&mut self, py: Python<'_>, event: HookEvent<'_>) -> PyResult<HookStep<Self, ()>> {
match event {
HookEvent::Started { .. } => Ok(HookStep::Ready(())),
HookEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
let api_key = self
.request
.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.request.as_ref().map(|request| &request.body),
self.request.as_ref().map(|request| &request.headers),
)?;
Ok(HookStep::Ready(()))
}
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(HookStep::Ready(()))
}
HookEvent::Failed {
timing,
origin,
error,
} => {
self.end = Some(datetime(py, timing.end_time)?);
self.error = Some(error.clone_ref(py).into_value(py));
if self.stream.is_some() {
return self.stream_failure(py);
}
if origin == FailureOrigin::Call
&& self.logger.is_some()
&& self.runs_deployment_hooks()
{
let error = self.error.as_ref().ok_or_else(missing_state)?;
return Ok(HookStep::Await(
DeploymentHooks::after_failure(
py,
self.call.kwargs(),
error,
self.surface.call_type,
)?,
Self::resume_deployment_failure,
));
}
self.dispatch_failure(py)
}
}
pub(crate) fn post_call(
&mut self,
py: Python<'_>,
raw: &RawResponse,
) -> PyResult<HookStep<Self, ()>> {
let api_key = self
.request
.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.request.as_ref().map(|request| &request.body),
self.request.as_ref().map(|request| &request.headers),
)?;
Ok(HookStep::Ready(()))
}
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
if self.surface.stream.is_none() {
pub(crate) fn succeeded(
&mut self,
py: Python<'_>,
timing: Timing,
response: &Py<PyAny>,
) -> PyResult<HookStep<Self, ()>> {
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(HookStep::Ready(()))
}
pub(crate) fn failed(
&mut self,
py: Python<'_>,
timing: Timing,
origin: FailureOrigin,
error: &PyErr,
) -> PyResult<HookStep<Self, ()>> {
self.end = Some(datetime(py, timing.end_time)?);
self.error = Some(error.clone_ref(py).into_value(py));
if self.stream.is_some() {
return self.stream_failure(py);
}
if origin == FailureOrigin::Call && self.logger.is_some() && self.runs_deployment_hooks() {
let error = self.error.as_ref().ok_or_else(missing_state)?;
return Ok(HookStep::Await(
DeploymentHooks::after_failure(py, self.call.kwargs(), error, self.call_type())?,
Self::resume_deployment_failure,
));
}
self.dispatch_failure(py)
}
pub(crate) fn stream_opened(&mut self, py: Python<'_>) -> PyResult<()> {
if self.stream_billing().is_none() {
return Err(missing_state());
}
Streaming::Opened.call(py, (self.logger()?.object(py),))?;
@ -475,7 +497,7 @@ impl PythonCallHooks for LegacyLogging {
Ok(())
}
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
pub(crate) fn 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())?);
@ -519,8 +541,9 @@ impl PythonOwned for LegacyLogging {
mod deployment_hooks_tests {
use std::ffi::CStr;
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks};
use litellm_host::hooks::CallHooks;
use litellm_host::lifecycle::{FailureOrigin, Timing};
use litellm_host_python::{HookStep, PythonCallEvent};
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
use pyo3::types::PyDict;
@ -579,6 +602,47 @@ kwargs = {'logger': logger, 'document': document}
matches!(step, HookStep::Await(_, _))
}
#[rstest]
#[case::sync_completion(litellm_types::Operation::Completion, false, "completion")]
#[case::async_completion(litellm_types::Operation::Completion, true, "acompletion")]
#[case::sync_responses(litellm_types::Operation::Responses, false, "responses")]
#[case::async_responses(litellm_types::Operation::Responses, true, "aresponses")]
#[case::sync_messages(litellm_types::Operation::Messages, false, "anthropic_messages")]
#[case::async_messages(litellm_types::Operation::Messages, true, "anthropic_messages")]
#[case::sync_ocr(litellm_types::Operation::Ocr, false, "ocr")]
#[case::async_ocr(litellm_types::Operation::Ocr, true, "aocr")]
fn operation_selects_the_legacy_setup_and_deployment_hook_contract(
#[case] operation: litellm_types::Operation,
#[case] asynchronous: bool,
#[case] expected: &str,
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, CALL);
let mut logging = LegacyLogging {
operation,
..legacy_call(py, &locals, asynchronous)
};
let kwargs = local(&locals, "kwargs")
.cast_into::<PyDict>()
.unwrap()
.unbind();
let step = logging.prepare_arguments(py, kwargs, 0.0).unwrap();
assert_eq!(awaits_deployment_hook(&step), asynchronous);
locals.set_item("expected", expected).unwrap();
locals.set_item("asynchronous", asynchronous).unwrap();
run(
py,
&locals,
c"
assert logger.setup_call_type == expected
if asynchronous:
assert logger.calls == [('pre_hook', expected)]
",
);
});
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
@ -775,7 +839,7 @@ assert finalized is replacement
)
.unwrap();
let failure = PyErr::from_value(local(&locals, "failure"));
let failed = HookEvent::Failed {
let failed = PythonCallEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Call,
error: &failure,
@ -808,8 +872,10 @@ mod payload_tests {
use std::ffi::CStr;
use litellm_auth::SecretValue;
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, PythonOwned, to_py};
use litellm_host::hooks::CallHooks;
use litellm_host::interceptors::{RawResponse, RequestContext, WireRequest};
use litellm_host::lifecycle::ExecutionEvent;
use litellm_host_python::{HookStep, PythonCallEvent, PythonOwned, to_py};
use proptest::prelude::*;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
@ -833,6 +899,7 @@ class PayloadLogger(StubLogger):
def pre_call(self, input, api_key, additional_args):
self.record('pre_call', None)
self.pre = additional_args
self.pre_input = input
self.pre_api_key = api_key
on_pre_call(additional_args)
@ -924,13 +991,18 @@ check = lambda: None
let step = logging
.before_provider_request(py, Box::new(wire), context)
.unwrap();
let raw = MachineEvent::ResponseReceived {
raw: RawResponse {
body: "raw response".into(),
},
let raw = RawResponse {
body: "raw response".into(),
};
assert!(matches!(
logging.on_event(py, HookEvent::Machine(&raw)).unwrap(),
logging
.on_event(
py,
PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived {
raw: &raw
})
)
.unwrap(),
HookStep::Ready(())
));
(logging, step)
@ -975,6 +1047,57 @@ check = lambda: None
}
}
#[rstest]
#[case::completion(litellm_types::Operation::Completion, "Chat completions")]
#[case::responses(litellm_types::Operation::Responses, "Responses")]
#[case::messages(litellm_types::Operation::Messages, "Messages")]
#[case::ocr(litellm_types::Operation::Ocr, "OCR document processing")]
fn prepared_arguments_replace_the_legacy_view_without_losing_callback_aliases(
#[case] operation: litellm_types::Operation,
#[case] description: &str,
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, PAYLOAD_LOGGER);
run(
py,
&locals,
c"
original = [0]
replacement = [1]
kwargs['pages'] = original
prepared = {'pages': replacement}
",
);
let mut logging = LegacyLogging {
operation,
logger: Some(PythonLogger::new(local(&locals, "logger").unbind())),
..legacy_call(py, &locals, false)
};
let prepared = local(&locals, "prepared")
.cast_into::<pyo3::types::PyDict>()
.unwrap()
.unbind();
logging.arguments_prepared(py, &prepared).unwrap();
let wire = WireRequest {
body: json!({"pages": [1]}),
..route_wire()
};
let (_, step) = send_and_receive(py, &mut logging, wire, &route_context());
assert!(matches!(step, HookStep::Ready(_)));
locals.set_item("description", description).unwrap();
run(
py,
&locals,
c"
assert logger.pre['complete_input_dict']['pages'] is replacement
assert logger.pre_input == description
assert original == [0]
",
);
});
}
#[rstest::rstest]
fn a_cycle_through_the_retained_headers_is_collected() {
Python::initialize();
@ -1459,8 +1582,9 @@ def check():
mod terminal_tests {
use std::ffi::CStr;
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, PythonOwned};
use litellm_host::hooks::CallHooks;
use litellm_host::lifecycle::{FailureOrigin, Timing};
use litellm_host_python::{HookStep, PythonCallEvent, PythonOwned};
use pyo3::exceptions::PyRuntimeError;
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
@ -1492,7 +1616,7 @@ mod terminal_tests {
logging
.on_event(
py,
HookEvent::Succeeded {
PythonCallEvent::Succeeded {
timing: TIMING,
response: &response,
},
@ -1509,7 +1633,7 @@ mod terminal_tests {
logging
.on_event(
py,
HookEvent::Failed {
PythonCallEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Host,
error: &failure,
@ -1558,6 +1682,75 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy
});
}
#[rstest]
fn dropped_observations_preserve_deferred_success_and_response_identity() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"response = object()\nlogger._defer_async_logging = True",
);
let mut logging = logged(py, &locals, true);
let response = local(&locals, "response").unbind();
let event = PythonCallEvent::Succeeded {
timing: TIMING,
response: &response,
};
let (sender, receiver) = litellm_host::observation::observation_channel(
std::num::NonZeroUsize::new(1).unwrap(),
);
drop(receiver);
sender.emit(event.snapshot());
assert!(matches!(
logging.on_event(py, event).unwrap(),
HookStep::Ready(())
));
assert_eq!(sender.dropped_events(), 1);
run(py, &locals, c"
assert logger.names() == ['sync_success_for_async_call'], logger.calls
logger._native_pending_logging.release(True)
logger._native_pending_logging.release(True)
assert logger.names() == ['sync_success_for_async_call', 'async_success_handler', 'enqueued'], logger.calls
assert logger.calls[0][1] is response
assert logger.calls[1][1] is response
");
});
}
#[rstest]
fn stream_bindings_deliver_collected_chunks_in_order_without_success_fan_out() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"first = b'first'\nlast = b'last'\nresponse = None");
let mut logging = LegacyLogging {
operation: litellm_types::Operation::Messages,
..logged(py, &locals, true)
};
logging.on_stream_open(py).unwrap();
logging
.on_stream_chunk(py, &local(&locals, "first").unbind())
.unwrap();
logging
.on_stream_chunk(py, &local(&locals, "last").unbind())
.unwrap();
assert!(matches!(
succeed(py, &locals, &mut logging),
HookStep::Ready(())
));
run(
py,
&locals,
c"
assert logger.names() == ['stream_opened', 'stream_success'], logger.calls
chunks = logger.calls[1][1]
assert len(chunks) == 2
assert chunks[0] is first
assert chunks[1] is last
",
);
});
}
#[rstest]
#[case::synchronous(false, &["failure_handler"])]
#[case::asynchronous(true, &[])]

View file

@ -3,16 +3,13 @@
//! lifetime. No other callback host has that obligation, which is why nothing outside
//! this crate holds them.
use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol};
use litellm_host_python::{Preflight, PythonBinding, PythonHostCalls, lookup, run_call};
use litellm_host_python::lookup;
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,
types::{PyDict, PyTuple},
};
use crate::{LegacyLogging, LegacySurface};
pub struct PublicCall {
args: Py<PyTuple>,
kwargs: Py<PyDict>,
@ -34,6 +31,10 @@ impl PublicCall {
})
}
pub fn arguments(&self, py: Python<'_>) -> Py<PyDict> {
self.kwargs.clone_ref(py)
}
pub(crate) fn args(&self) -> &Py<PyTuple> {
&self.args
}
@ -64,35 +65,6 @@ impl PublicCall {
}
}
/// Runs one native call under the legacy `Logging` contract: the protocol host projects from
/// the keyword view the contract prepares and `preflight` rewrites, and the contract
/// observes the call.
pub fn run_legacy_call<H, M>(
py: Python<'_>,
surface: LegacySurface,
call: PublicCall,
start: impl FnOnce(<H::Protocol as Protocol>::Request) -> M + Send + Sync + 'static,
host: H,
preflight: Preflight,
asynchronous: bool,
) -> PyResult<Py<PyAny>>
where
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,
start,
host,
LegacyLogging::new(py, surface, call, asynchronous),
preflight,
arguments,
asynchronous,
)
}
#[cfg(test)]
mod tests {
use super::*;

View file

@ -2,7 +2,7 @@
//! the deferred and worker-submitted success paths, and the sync-callbacks-for-async-calls
//! duplication. All of it expires with the legacy callback contract.
use litellm_host::event::{RequestContext, WireRequest};
use litellm_host::interceptors::{RequestContext, WireRequest};
use litellm_host_python::to_py;
use pyo3::{exceptions::PyBaseException, prelude::*, types::PyDict};

View file

@ -2,25 +2,23 @@
//! sync and async callback registries it fans out to, the deployment hooks and the deferred
//! proxy release. All of it sits behind one
//! [`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.
//! core never learn which Python object is on the other end.
//!
//! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`]
//! is where those objects live, and [`run_legacy_call`] is how a route hands them over
//! without keeping a copy.
//! is where those objects live.
mod adapter;
mod call;
mod callbacks;
mod deferred;
mod logger;
mod mapping;
mod python;
pub(crate) use adapter::LegacyLogging;
pub use adapter::{LegacySurface, PassThroughStream};
pub use call::{PublicCall, run_legacy_call};
pub use adapter::LegacyLogging;
pub use call::PublicCall;
pub(crate) use callbacks::{LegacyCallbacks, is_internal_call};
pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup};
pub use mapping::{CallBoundary, CallbackMapping, Dispatch, callback_mappings};
#[cfg(test)]
mod test_support;

View file

@ -0,0 +1,285 @@
use litellm_host::{
hooks::CallHooks,
interceptors::{RawResponse, RequestContext, WireRequest},
lifecycle::{ExecutionEvent, FailureOrigin, Timing},
};
use litellm_host_python::{HookStep, PythonCallEvent, PythonRuntime};
use pyo3::{prelude::*, types::PyDict};
use crate::LegacyLogging;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CallBoundary {
PrepareArguments,
BeforeProviderRequest,
AfterProviderResponse,
TransformResponse,
Succeeded,
Failed,
StreamOpened,
StreamChunk,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Dispatch {
Call(CallBoundary),
Python(&'static str),
DeclarationOnly,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct CallbackMapping {
pub callback: &'static str,
pub dispatch: Dispatch,
}
struct Binding<H> {
boundary: CallBoundary,
invoke: H,
callbacks: &'static [&'static str],
}
impl<H> Binding<H> {
fn mappings(&self) -> impl Iterator<Item = CallbackMapping> {
self.callbacks.iter().map(|callback| CallbackMapping {
callback,
dispatch: Dispatch::Call(self.boundary),
})
}
}
type Step<T> = PyResult<HookStep<LegacyLogging, T>>;
type Prepare = fn(&mut LegacyLogging, Python<'_>, Py<PyDict>, f64) -> Step<Py<PyDict>>;
type Before =
fn(&mut LegacyLogging, Python<'_>, Box<WireRequest>, &RequestContext) -> Step<Box<WireRequest>>;
type After = fn(&mut LegacyLogging, Python<'_>, &RawResponse) -> Step<()>;
type Transform = fn(&mut LegacyLogging, Python<'_>, Py<PyAny>, Timing) -> Step<Py<PyAny>>;
type Success = fn(&mut LegacyLogging, Python<'_>, Timing, &Py<PyAny>) -> Step<()>;
type Failure = fn(&mut LegacyLogging, Python<'_>, Timing, FailureOrigin, &PyErr) -> Step<()>;
type Open = fn(&mut LegacyLogging, Python<'_>) -> PyResult<()>;
type Chunk = fn(&mut LegacyLogging, Python<'_>, &Py<PyAny>) -> PyResult<()>;
const PREPARE: Binding<Prepare> = Binding {
boundary: CallBoundary::PrepareArguments,
invoke: LegacyLogging::prepare_call,
callbacks: &["async_pre_call_deployment_hook"],
};
const BEFORE: Binding<Before> = Binding {
boundary: CallBoundary::BeforeProviderRequest,
invoke: LegacyLogging::pre_call,
callbacks: &["log_pre_api_call", "log_input_event"],
};
const AFTER: Binding<After> = Binding {
boundary: CallBoundary::AfterProviderResponse,
invoke: LegacyLogging::post_call,
callbacks: &["log_post_api_call"],
};
const TRANSFORM: Binding<Transform> = Binding {
boundary: CallBoundary::TransformResponse,
invoke: LegacyLogging::transform_public_response,
callbacks: &["async_post_call_success_deployment_hook"],
};
const SUCCESS: Binding<Success> = Binding {
boundary: CallBoundary::Succeeded,
invoke: LegacyLogging::succeeded,
callbacks: &[
"log_success_event",
"async_log_success_event",
"logging_hook",
"async_logging_hook",
"redact_standard_logging_payload_from_model_call_details",
"log_event",
"async_log_event",
],
};
const FAILURE: Binding<Failure> = Binding {
boundary: CallBoundary::Failed,
invoke: LegacyLogging::failed,
callbacks: &[
"async_post_call_failure_deployment_hook",
"log_failure_event",
"async_log_failure_event",
"log_model_group_rate_limit_error",
"log_event",
"async_log_event",
],
};
const OPEN: Binding<Open> = Binding {
boundary: CallBoundary::StreamOpened,
invoke: LegacyLogging::stream_opened,
callbacks: &[],
};
const CHUNK: Binding<Chunk> = Binding {
boundary: CallBoundary::StreamChunk,
invoke: LegacyLogging::stream_chunk,
callbacks: &[],
};
pub fn callback_mappings() -> impl Iterator<Item = CallbackMapping> {
PREPARE
.mappings()
.chain(BEFORE.mappings())
.chain(AFTER.mappings())
.chain(TRANSFORM.mappings())
.chain(SUCCESS.mappings())
.chain(FAILURE.mappings())
.chain(OPEN.mappings())
.chain(CHUNK.mappings())
.chain(PYTHON_CALLBACKS.iter().copied())
}
impl CallHooks<PythonRuntime> for LegacyLogging {
fn prepare_arguments(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
started_at: f64,
) -> Step<Py<PyDict>> {
(PREPARE.invoke)(self, py, arguments, started_at)
}
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()> {
self.adopt_arguments(py, arguments);
Ok(())
}
fn before_provider_request(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> Step<Box<WireRequest>> {
(BEFORE.invoke)(self, py, wire, context)
}
fn transform_response(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
timing: Timing,
) -> Step<Py<PyAny>> {
(TRANSFORM.invoke)(self, py, response, timing)
}
fn on_event(&mut self, py: Python<'_>, event: PythonCallEvent<'_>) -> Step<()> {
match event {
PythonCallEvent::Started { .. } | PythonCallEvent::Cancelled { .. } => {
Ok(HookStep::Ready(()))
}
PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
(AFTER.invoke)(self, py, raw)
}
PythonCallEvent::Succeeded { timing, response } => {
(SUCCESS.invoke)(self, py, timing, response)
}
PythonCallEvent::Failed {
timing,
origin,
error,
} => (FAILURE.invoke)(self, py, timing, origin, error),
}
}
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
(OPEN.invoke)(self, py)
}
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
(CHUNK.invoke)(self, py, chunk)
}
}
macro_rules! python_callbacks {
($($dispatch:expr => [$($callback:literal),* $(,)?]),* $(,)?) => {
const PYTHON_CALLBACKS: &[CallbackMapping] = &[
$($(CallbackMapping { callback: $callback, dispatch: $dispatch },)*)*
];
};
}
python_callbacks! {
Dispatch::Python("litellm.router") => [
"async_pre_routing_hook",
"async_filter_deployments",
"pre_call_check",
"async_pre_call_check",
],
Dispatch::Python("litellm.router_utils.fallback_event_handlers") => [
"log_success_fallback_event",
"log_failure_fallback_event",
],
Dispatch::Python("litellm.proxy.utils") => [
"async_pre_call_hook",
"async_post_call_response_headers_hook",
"async_post_call_failure_hook",
"async_post_call_success_hook",
"async_moderation_hook",
"async_post_call_streaming_hook",
"async_post_call_streaming_iterator_hook",
"async_filter_listed_models",
],
Dispatch::Python("litellm.litellm_core_utils.litellm_logging") => [
"async_get_chat_completion_prompt",
"get_chat_completion_prompt",
"log_stream_event",
"async_log_stream_event",
"async_post_mcp_tool_call_hook",
],
Dispatch::Python("litellm.llms.anthropic.pass_through.messages.handler") => [
"async_pre_request_hook",
],
Dispatch::Python("litellm.litellm_core_utils.streaming_handler") => [
"async_post_call_streaming_deployment_hook",
],
Dispatch::Python("litellm.responses.streaming_iterator") => [
"async_post_call_streaming_deployment_hook",
],
Dispatch::Python("litellm.main") => [
"translate_completion_input_params",
"translate_completion_output_params",
"translate_completion_output_params_streaming",
],
Dispatch::Python("litellm.integrations.argilla") => ["async_dataset_hook"],
Dispatch::Python("litellm.proxy.management_helpers.audit_logs") => ["async_log_audit_log_event"],
Dispatch::Python("litellm.llms.custom_httpx.llm_http_handler") => [
"async_should_run_agentic_loop",
"async_run_agentic_loop",
"async_build_agentic_loop_plan",
"async_post_agentic_loop_response_hook",
"async_agentic_loop_cleanup_hook",
"async_should_run_chat_completion_agentic_loop",
"async_run_chat_completion_agentic_loop",
"async_build_chat_completion_agentic_loop_plan",
],
Dispatch::Python("litellm.litellm_core_utils.chat_completion_agentic_loop") => [
"async_should_run_agentic_loop",
"async_run_agentic_loop",
"async_build_agentic_loop_plan",
"async_post_agentic_loop_response_hook",
"async_agentic_loop_cleanup_hook",
],
Dispatch::Python("litellm.llms.openai.openai") => [
"async_should_run_chat_completion_agentic_loop",
"async_run_chat_completion_agentic_loop",
],
Dispatch::Python("litellm.proxy.spend_tracking.cold_storage_handler") => [
"get_proxy_server_request_from_cold_storage_with_object_key",
],
Dispatch::Python("litellm.integrations.custom_logger") => [
"truncate_standard_logging_payload_content",
"redacts_messages_itself",
"handle_callback_failure",
"get_callback_env_vars",
],
Dispatch::DeclarationOnly => [
"async_log_pre_api_call",
"async_log_input_event",
],
}

View file

@ -3,7 +3,7 @@ use std::ffi::CStr;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use crate::{LegacyLogging, LegacySurface, PublicCall};
use crate::{LegacyLogging, PublicCall};
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
/// `tests/unit/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
@ -45,11 +45,14 @@ def contracted(name, fake):
if not hasattr(legacy, 'is_internal'):
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
def setup(call_type, args, kwargs, start, asynchronous):
logger = kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger']
logger.setup_call_type = call_type
return types.SimpleNamespace(logger=logger, kwargs=kwargs)
FAKES = {
'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
kwargs=kwargs,
),
'setup': setup,
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
kwargs=kwargs,
@ -82,10 +85,10 @@ FAKES = {
),
'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type),
'stream_opened': lambda logger: logger.record('stream_opened', None),
'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record(
'stream_success': lambda logger, url_route, endpoint_type, request_body, chunks, start, end, first_chunk: logger.record(
'stream_success', list(chunks)
),
'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error),
'stream_failure': lambda logger, endpoint_type, request_body, chunks, error: logger.record('stream_failure', error),
}
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
for name, fake in FAKES.items():
@ -186,14 +189,5 @@ pub(crate) fn legacy_call(
.map(|kwargs| kwargs.cast_into::<PyDict>().unwrap())
.unwrap_or_else(|| PyDict::new(py));
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
LegacyLogging::new(
py,
LegacySurface {
call_type: "test",
input_description: "test input",
stream: None,
},
call,
asynchronous,
)
LegacyLogging::new(py, litellm_types::Operation::Ocr, call, asynchronous)
}

View file

@ -2,7 +2,7 @@ litellm-core owns route orchestration. Messages and HTTP Responses return `litel
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
Chat Completions, Messages, Responses, and OCR execute through their route objects. Calls pass `Interceptors` and an optional `ObservationSender` separately; use `&()` for no hooks and `None` for no observer. Construction does no work; preparation and lifecycle observation begin when the future is polled. Handlers accept `Interceptors`, never a concrete `ChannelInterceptors`. Native observers receive start and terminal events through the shared call runner; a stream retains its lifecycle until exhaustion, error, or drop. Hosted routes leave terminal observation to their driver
`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
@ -20,7 +20,7 @@ Each crate mirrors one top-level Python package, so a Rust path reads as its Pyt
- `litellm-llms` mirrors `litellm/llms/`: `base_llm/<api>/transformation.rs`, `<provider>/<api>/transformation.rs`, and `base_llm/ocr/handler.rs` (the OCR request handler)
- `litellm-core` mirrors the route packages (`litellm/ocr/`, `litellm/messages/`, ...): entrypoints, route request types, provider dispatch, the route machine, and hooks
A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, the provider-level hooks OCR implements over its host until it folds into `litellm_host::hooks::RouteHooks`. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, the provider-level hooks OCR implements over its host until it folds into `litellm_host::interceptors::Interceptors`. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
## Error placement

View file

@ -38,6 +38,7 @@ veil.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
litellm-auth-gcp.workspace = true
litellm-host-native.workspace = true
litellm-llms = { workspace = true, features = ["test-support"] }
rstest.workspace = true
rstest_reuse.workspace = true

View file

@ -1,10 +1,9 @@
use litellm_host::lifecycle::ExecutionEvent;
use litellm_host::observation::ObservationSender;
use std::time::Duration;
use litellm_auth::AuthServices;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
hooks::RouteHooks,
};
use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body};
use litellm_llms::base_llm::{
auth::{Authenticated, resolve_auth},
@ -23,7 +22,8 @@ pub(super) async fn execute(
http: &Client,
auth: &AuthServices,
request: ProviderChatCompletionsRequest,
hooks: &impl RouteHooks<Error>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {
let ProviderChatCompletionsRequest {
model,
@ -45,7 +45,7 @@ pub(super) async fn execute(
api_key,
};
let authenticated = resolve_auth(auth, environment, &|key| secrets.get(key)).await?;
let wire = hooks
let wire = interceptors
.before_provider_request(
WireRequest {
url,
@ -87,10 +87,14 @@ pub(super) async fn execute(
body: truncate_error_body(&text),
}));
}
hooks
.on_event(MachineEvent::ResponseReceived {
raw: RawResponse { body: text.clone() },
})
let raw = RawResponse { body: text.clone() };
if let Some(observers) = observers {
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
));
}
interceptors
.after_provider_response(raw)
.await
.map_err(Error::post_call)?;
@ -168,7 +172,7 @@ mod tests {
raw: Mutex<Vec<String>>,
}
impl RouteHooks<Error> for RecordingHooks {
impl Interceptors<Error> for RecordingHooks {
async fn before_provider_request(
&self,
wire: WireRequest,
@ -188,8 +192,7 @@ mod tests {
})
}
async fn on_event(&self, event: MachineEvent) -> Result<(), Error> {
let MachineEvent::ResponseReceived { raw } = event;
async fn after_provider_response(&self, raw: RawResponse) -> Result<(), Error> {
self.raw.lock().unwrap().push(raw.body);
Ok(())
}
@ -223,13 +226,14 @@ mod tests {
)
.mount(&upstream)
.await;
let hooks = RecordingHooks::default();
let interceptors = RecordingHooks::default();
execute(
&Client::plain_for_test(),
&AuthServices::default(),
prepared(&upstream.uri()),
&hooks,
&interceptors,
None,
)
.await
.expect("chat completions call succeeds");
@ -240,14 +244,17 @@ mod tests {
assert_eq!(sent["system"], "added by the host");
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_provider_request runs once, saw {}", seen.len()));
let [context] =
<[RequestContext; 1]>::try_from(interceptors.contexts.into_inner().unwrap())
.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")
);
assert_eq!(context.optional_params, json!({"max_tokens": 16}));
assert_eq!(hooks.raw.into_inner().unwrap(), [ANTHROPIC_MESSAGE]);
assert_eq!(interceptors.raw.into_inner().unwrap(), [ANTHROPIC_MESSAGE]);
}
#[rstest]
@ -258,13 +265,14 @@ mod tests {
.respond_with(ResponseTemplate::new(500).set_body_string("boom"))
.mount(&upstream)
.await;
let hooks = RecordingHooks::default();
let interceptors = RecordingHooks::default();
let error = execute(
&Client::plain_for_test(),
&AuthServices::default(),
prepared(&upstream.uri()),
&hooks,
&interceptors,
None,
)
.await
.expect_err("the upstream failure fails the call");
@ -273,7 +281,7 @@ mod tests {
error,
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
));
assert!(hooks.raw.into_inner().unwrap().is_empty());
assert!(interceptors.raw.into_inner().unwrap().is_empty());
}
#[rstest::rstest]

View file

@ -1,3 +1,4 @@
use litellm_host::observation::ObservationSender;
pub mod route;
pub mod types;
pub use crate::error::RouteError as Error;
@ -35,9 +36,14 @@ impl ChatCompletionsRoute {
pub async fn execute(
&self,
request: ChatCompletionsRequest<'_>,
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {
litellm_host::lifecycle::observe_unary(hooks.observer(), self.run(request, hooks)).await
litellm_host::lifecycle::observe_unary(
observers.clone(),
self.run(request, interceptors, observers.as_ref()),
)
.await
}
#[tracing::instrument(name = "litellm.route", skip_all, fields(
@ -51,7 +57,8 @@ impl ChatCompletionsRoute {
async fn run(
&self,
request: ChatCompletionsRequest<'_>,
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {
crate::diagnostic::unary(async {
let resolved = resolve_request(request)?;
@ -64,7 +71,13 @@ impl ChatCompletionsRoute {
let execute: futures_util::future::BoxFuture<
'_,
Result<ChatCompletionsResponse, Error>,
> = Box::pin(handler::execute(&self.http, &self.auth, prepared, hooks));
> = Box::pin(handler::execute(
&self.http,
&self.auth,
prepared,
interceptors,
observers,
));
execute.await
})
.await

View file

@ -1,3 +1,4 @@
use litellm_host::observation::ObservationSender;
use std::convert::Infallible;
use litellm_host::{
@ -23,10 +24,15 @@ impl Protocol for ChatCompletions {
}
impl ChatCompletionsRoute {
pub fn machine(self, call: ChatCompletionsCall) -> HostedMachine<ChatCompletions> {
pub fn machine(
self,
call: ChatCompletionsCall,
observers: Option<ObservationSender>,
) -> HostedMachine<ChatCompletions> {
hosted_call(
call,
move |call: ChatCompletionsCall, _, hooks| async move {
observers,
move |call: ChatCompletionsCall, _, interceptors, observers| async move {
let request = ChatCompletionsRequest {
model: &call.model,
messages: call.messages,
@ -37,7 +43,9 @@ impl ChatCompletionsRoute {
extra_headers: call.extra_headers,
timeout: call.timeout,
};
self.run(request, &hooks).await.map(CallOutput::Complete)
self.run(request, &interceptors, observers.as_ref())
.await
.map(CallOutput::Complete)
},
)
}

View file

@ -1,12 +1,11 @@
use litellm_host::lifecycle::ExecutionEvent;
use litellm_host::observation::ObservationSender;
use std::time::Duration;
use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt, stream::BoxStream};
use litellm_auth::AuthServices;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
hooks::RouteHooks,
};
use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
use litellm_http::transport::Error as TransportError;
use litellm_llms::base_llm::{
auth::{Authenticated, resolve_auth},
@ -28,7 +27,8 @@ pub(super) async fn execute(
http: &litellm_http::Client,
auth: &AuthServices,
request: ProviderMessagesRequest,
hooks: &impl RouteHooks<Error>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<MessagesResponse, Error> {
let ProviderMessagesRequest {
provider,
@ -47,7 +47,7 @@ pub(super) async fn execute(
api_key,
};
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
let wire = hooks
let wire = interceptors
.before_provider_request(
WireRequest {
url,
@ -83,10 +83,14 @@ pub(super) async fn execute(
}
let text = response.text().await.map_err(network)?;
log_response_body(&text);
hooks
.on_event(MachineEvent::ResponseReceived {
raw: RawResponse { body: text.clone() },
})
let raw = RawResponse { body: text.clone() };
if let Some(observers) = observers {
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
));
}
interceptors
.after_provider_response(raw)
.await
.map_err(Error::post_call)?;
decode_response(config, &body.model, &text)

View file

@ -1,3 +1,4 @@
use litellm_host::observation::ObservationSender;
mod common_utils;
mod handler;
mod prepare;
@ -34,9 +35,14 @@ impl MessagesRoute {
pub async fn execute(
&self,
call: MessagesCall,
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<ObservationSender>,
) -> Result<MessagesResponse, Error> {
litellm_host::lifecycle::observe_call(hooks.observer(), self.run(call, hooks)).await
litellm_host::lifecycle::observe_call(
observers.clone(),
self.run(call, interceptors, observers.as_ref()),
)
.await
}
#[tracing::instrument(name = "litellm.route", skip_all, fields(
@ -50,13 +56,20 @@ impl MessagesRoute {
async fn run(
&self,
call: MessagesCall,
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<MessagesResponse, Error> {
crate::diagnostic::call(async {
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
crate::diagnostic::provider(&request.body.model, request.provider.as_str());
let execute: futures_util::future::BoxFuture<'_, Result<MessagesResponse, Error>> =
Box::pin(handler::execute(&self.http, &self.auth, request, hooks));
Box::pin(handler::execute(
&self.http,
&self.auth,
request,
interceptors,
observers,
));
execute.await
})
.await

View file

@ -1,3 +1,4 @@
use litellm_host::observation::ObservationSender;
use std::convert::Infallible;
use bytes::Bytes;
@ -30,9 +31,17 @@ impl Protocol for Messages {
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 fn machine(
self,
request: super::MessagesCall,
observers: Option<ObservationSender>,
) -> MessagesMachine {
hosted_call(
request,
observers,
move |call, _, interceptors, observers| async move {
self.run(call, &interceptors, observers.as_ref()).await
},
)
}
}

View file

@ -1,6 +1,7 @@
use litellm_host::observation::ObservationSender;
use std::sync::Arc;
use litellm_host::hooks::RouteHooks;
use litellm_host::interceptors::Interceptors;
use litellm_llms::base_llm::ocr::{
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
};
@ -23,9 +24,14 @@ impl OcrRoute {
pub async fn execute(
&self,
request: LiteLLMOcrRequest,
hooks: &impl RouteHooks<Error>,
interceptors: &impl Interceptors<Error>,
observers: Option<ObservationSender>,
) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::lifecycle::observe_unary(hooks.observer(), self.run(request, hooks)).await
litellm_host::lifecycle::observe_unary(
observers.clone(),
self.run(request, interceptors, observers.as_ref()),
)
.await
}
#[tracing::instrument(name = "litellm.route", skip_all, fields(
@ -39,7 +45,8 @@ impl OcrRoute {
pub(super) async fn run(
&self,
request: LiteLLMOcrRequest,
hooks: &impl RouteHooks<Error>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<LiteLLMOcrResponse, Error> {
crate::diagnostic::unary(async {
let caller_document = matches!(&request.document, OcrDocumentInput::Document(_));
@ -48,8 +55,9 @@ impl OcrRoute {
Box::pin(perform_ocr_request(
&self.client,
prepared,
hooks,
interceptors,
caller_document,
observers,
));
execute.await
})

View file

@ -1,8 +1,7 @@
use futures_util::future::BoxFuture;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
hooks::RouteHooks,
};
use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
use litellm_host::lifecycle::ExecutionEvent;
use litellm_host::observation::ObservationSender;
use litellm_llms::base_llm::ocr::{
error::Error,
handler::{CallHooks, OcrClient},
@ -16,8 +15,9 @@ use crate::ocr::types::ResolvedOcrRequest;
pub(crate) async fn perform_ocr_request(
client: &OcrClient,
request: ResolvedOcrRequest,
host: &impl RouteHooks<Error>,
host: &impl Interceptors<Error>,
caller_document: bool,
observers: Option<&ObservationSender>,
) -> Result<LiteLLMOcrResponse, Error> {
request.response_format()?;
let config = request.config;
@ -27,19 +27,26 @@ 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, &request, config);
config.ocr(client, &request, &hooks).await
let interceptors = OcrCallHooks::new(host, &request, config, observers);
config.ocr(client, &request, &interceptors).await
}
struct OcrCallHooks<'a, H> {
hooks: &'a H,
interceptors: &'a H,
context: RequestContext,
observers: Option<&'a ObservationSender>,
}
impl<'a, H> OcrCallHooks<'a, H> {
fn new(hooks: &'a H, request: &PreparedOcrRequest, config: OcrConfigKind) -> Self {
fn new(
interceptors: &'a H,
request: &PreparedOcrRequest,
config: OcrConfigKind,
observers: Option<&'a ObservationSender>,
) -> Self {
Self {
hooks,
interceptors,
observers,
context: RequestContext {
model: request.model.clone(),
custom_llm_provider: <&str>::from(config.provider()).to_owned(),
@ -56,22 +63,26 @@ impl<'a, H> OcrCallHooks<'a, H> {
}
}
impl<H: RouteHooks<Error>> CallHooks<Error> for OcrCallHooks<'_, H> {
impl<H: Interceptors<Error>> CallHooks<Error> for OcrCallHooks<'_, H> {
fn before_provider_request(
&self,
wire: WireRequest,
) -> BoxFuture<'_, Result<WireRequest, Error>> {
Box::pin(
self.hooks
self.interceptors
.before_provider_request(wire, self.context.clone()),
)
}
fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
Box::pin(self.hooks.on_event(MachineEvent::ResponseReceived {
raw: RawResponse {
body: String::from_utf8_lossy(body).into_owned(),
},
}))
let raw = RawResponse {
body: String::from_utf8_lossy(body).into_owned(),
};
if let Some(observers) = self.observers {
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
));
}
Box::pin(self.interceptors.after_provider_response(raw))
}
}

View file

@ -80,7 +80,7 @@ mod tests {
use futures_util::future::BoxFuture;
use litellm_core_utils::call_arguments::{CallArguments, compose_body, parse_options};
use litellm_host::event::WireRequest;
use litellm_host::interceptors::WireRequest;
use litellm_llms::{
base_llm::ocr::{
error::Error,
@ -100,7 +100,7 @@ mod tests {
wire::{OcrWireRequest, decode_request},
};
/// Stands in for a host with no hooks registered.
/// Stands in for a host with no interceptors registered.
struct NoHooks;
impl CallHooks<Error> for NoHooks {

View file

@ -133,9 +133,9 @@ impl OcrConfigKind {
self,
client: &OcrClient,
request: &PreparedOcrRequest,
hooks: &dyn CallHooks<Error>,
interceptors: &dyn CallHooks<Error>,
) -> Result<LiteLLMOcrResponse, Error> {
with_config!(self, config => handler::ocr(&config, client, request, hooks).await)
with_config!(self, config => handler::ocr(&config, client, request, interceptors).await)
}
}

View file

@ -1,7 +1,8 @@
use litellm_auth::ResolvedCredential;
use litellm_auth::{ResolvedCredential, TokenProviderHandle};
use litellm_host::observation::ObservationSender;
use litellm_host::{
call::{CallOutput, HostedMachine, hosted_call},
machine::{HostTokenProvider, TokenProtocol},
machine::HostServices,
protocol::Protocol,
protocol::Reply,
};
@ -31,27 +32,38 @@ impl Protocol for Ocr {
type StreamHead = std::convert::Infallible;
}
impl TokenProtocol for Ocr {
fn acquire_token_op(reply: Reply<ResolvedCredential>) -> OcrOp {
OcrOp::AcquireAzureAdToken(reply)
}
}
pub type OcrMachine = HostedMachine<Ocr>;
fn caller_token_provider(services: HostServices<Ocr>) -> TokenProviderHandle {
TokenProviderHandle::from_callback(move || {
let host_services = services.clone();
async move {
host_services
.call(OcrOp::AcquireAzureAdToken)
.await
.map_err(|error| {
litellm_auth::Error::CredentialAcquisition(error.to_string().into())
})
}
})
}
impl crate::ocr::OcrRoute {
pub fn machine(self, request: OcrCall) -> OcrMachine {
pub fn machine(self, request: OcrCall, observers: Option<ObservationSender>) -> OcrMachine {
hosted_call(
request,
move |projection: OcrCall, services, hooks| async move {
observers,
move |projection: OcrCall, services, interceptors, observers| async move {
let request = LiteLLMOcrRequest {
azure_ad_token_provider: projection
.caller_token
.then(|| HostTokenProvider::handle(services))
.then(|| caller_token_provider(services))
.or(projection.request.azure_ad_token_provider),
..projection.request
};
self.run(request, &hooks).await.map(CallOutput::Complete)
self.run(request, &interceptors, observers.as_ref())
.await
.map(CallOutput::Complete)
},
)
}

View file

@ -1,10 +1,9 @@
use litellm_host::lifecycle::ExecutionEvent;
use litellm_host::observation::ObservationSender;
use std::time::Duration;
use futures_util::StreamExt;
use litellm_host::{
event::{MachineEvent, RawResponse, WireRequest},
hooks::RouteHooks,
};
use litellm_host::interceptors::{Interceptors, RawResponse, WireRequest};
use litellm_llms::base_llm::auth::{Authenticated, resolve_auth};
use super::{
@ -16,10 +15,11 @@ pub(super) async fn execute(
http: &litellm_http::Client,
auth: &litellm_auth::AuthServices,
request: ProviderResponsesRequest,
hooks: &impl RouteHooks<Error>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {
let authenticated = resolve_auth(auth, request.environment, &|_| None).await?;
let wire = hooks
let wire = interceptors
.before_provider_request(
WireRequest {
url: request.url,
@ -71,10 +71,14 @@ pub(super) async fn execute(
});
}
let body = response.text().await.map_err(network)?;
hooks
.on_event(MachineEvent::ResponseReceived {
raw: RawResponse { body: body.clone() },
})
let raw = RawResponse { body: body.clone() };
if let Some(observers) = observers {
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
));
}
interceptors
.after_provider_response(raw)
.await
.map_err(Error::post_call)?;
let value = serde_json::from_str(&body)

View file

@ -1,4 +1,5 @@
pub use crate::error::RouteError as Error;
use litellm_host::observation::ObservationSender;
pub mod websocket;
mod handler;
@ -9,7 +10,7 @@ pub mod types;
use std::sync::Arc;
use litellm_auth::AuthServices;
use litellm_host::hooks::RouteHooks;
use litellm_host::interceptors::Interceptors;
use litellm_secrets::source::SecretSource;
use types::{ResponsesCall, ResponsesOutput};
@ -36,9 +37,14 @@ impl ResponsesRoute {
pub async fn execute(
&self,
call: ResponsesCall,
hooks: &impl RouteHooks<Error>,
interceptors: &impl Interceptors<Error>,
observers: Option<ObservationSender>,
) -> Result<ResponsesOutput, Error> {
litellm_host::lifecycle::observe_call(hooks.observer(), self.run(call, hooks)).await
litellm_host::lifecycle::observe_call(
observers.clone(),
self.run(call, interceptors, observers.as_ref()),
)
.await
}
#[tracing::instrument(name = "litellm.route", skip_all, fields(
@ -52,7 +58,8 @@ impl ResponsesRoute {
async fn run(
&self,
call: ResponsesCall,
hooks: &impl RouteHooks<Error>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {
crate::diagnostic::call(async {
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
@ -61,7 +68,13 @@ impl ResponsesRoute {
&request.context.custom_llm_provider,
);
let execute: futures_util::future::BoxFuture<'_, Result<ResponsesOutput, Error>> =
Box::pin(handler::execute(&self.http, &self.auth, request, hooks));
Box::pin(handler::execute(
&self.http,
&self.auth,
request,
interceptors,
observers,
));
execute.await
})
.await

View file

@ -1,4 +1,4 @@
use litellm_host::event::RequestContext;
use litellm_host::interceptors::RequestContext;
use litellm_llms::{
base_llm::responses::transformation::BaseResponsesApiConfig,
openai::responses::transformation::OpenAiResponsesApiConfig,

View file

@ -1,3 +1,4 @@
use litellm_host::observation::ObservationSender;
use std::convert::Infallible;
use bytes::Bytes;
@ -24,9 +25,17 @@ impl Protocol for Responses {
}
impl ResponsesRoute {
pub fn machine(self, call: ResponsesCall) -> HostedMachine<Responses> {
hosted_call(call, move |call, _, hooks| async move {
self.run(call, &hooks).await
})
pub fn machine(
self,
call: ResponsesCall,
observers: Option<ObservationSender>,
) -> HostedMachine<Responses> {
hosted_call(
call,
observers,
move |call, _, interceptors, observers| async move {
self.run(call, &interceptors, observers.as_ref()).await
},
)
}
}

View file

@ -32,6 +32,6 @@ pub(super) struct ProviderResponsesRequest {
pub environment: ValidatedEnvironment,
pub url: String,
pub body: Value,
pub context: litellm_host::event::RequestContext,
pub context: litellm_host::interceptors::RequestContext,
pub timeout: Option<Duration>,
}

View file

@ -1,3 +1,4 @@
use litellm_host::interceptors::RawResponse;
use std::time::Duration;
use litellm_core::chat_completions::{Error, types::ChatCompletionsRequest};
@ -13,7 +14,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_route().execute(request, &()).await
chat_completions_route().execute(request, &(), None).await
}
fn object(value: Value) -> Map<String, Value> {
@ -253,7 +254,7 @@ async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
#[case] hosted: bool,
) {
use litellm_core::chat_completions::route::ChatCompletions;
use litellm_host::{call::HostedCompletion, event::CallEvent};
use litellm_host::{call::HostedCompletion, lifecycle::CallEvent};
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
@ -265,8 +266,9 @@ async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
.into(),
);
let response = if hosted {
let result = litellm_host::in_process::run_hosted(
chat_completions_route().machine(host.request().unwrap()),
let result = litellm_host_native::in_process::run_hosted(
chat_completions_route()
.machine(host.request().unwrap(), Some(host.events.0.sender.clone())),
host.runtime(),
)
.await
@ -290,6 +292,7 @@ async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
timeout: call.timeout,
},
&host,
Some(host.events.0.sender.clone()),
)
.await
.unwrap()
@ -307,7 +310,7 @@ async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
&events[..],
[
CallEvent::Started { .. },
CallEvent::Machine(_),
CallEvent::Execution(_),
CallEvent::Succeeded { .. }
]
));
@ -318,12 +321,9 @@ async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
async fn a_post_call_hook_failure_never_looks_safe_to_retry(
request: ChatCompletionsRequest<'static>,
) {
use litellm_host::{
event::{MachineEvent, RequestContext, WireRequest},
hooks::RouteHooks,
};
use litellm_host::interceptors::{Interceptors, RequestContext, WireRequest};
struct FailingHook;
impl RouteHooks<Error> for FailingHook {
impl Interceptors<Error> for FailingHook {
async fn before_provider_request(
&self,
wire: WireRequest,
@ -331,7 +331,7 @@ async fn a_post_call_hook_failure_never_looks_safe_to_retry(
) -> Result<WireRequest, Error> {
Ok(wire)
}
async fn on_event(&self, _: MachineEvent) -> Result<(), Error> {
async fn after_provider_response(&self, _: RawResponse) -> Result<(), Error> {
Err(Error::InvalidRequest("callback rejected".into()))
}
}
@ -344,6 +344,7 @@ async fn a_post_call_hook_failure_never_looks_safe_to_retry(
..request
},
&FailingHook,
None,
)
.await
.unwrap_err();

View file

@ -1,7 +1,11 @@
use litellm_host::lifecycle::ExecutionEvent;
use std::sync::Mutex;
use litellm_core::messages::route::Messages;
use litellm_host::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
use litellm_host::{
interceptors::{RequestContext, WireRequest},
lifecycle::CallEvent,
};
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
use rstest::rstest;
@ -14,7 +18,7 @@ type Rewrite = Box<dyn Fn(WireRequest) -> Result<WireRequest, Error> + Send + Sy
struct RecordingHost {
call: LocalMessagesHost,
rewrite: Rewrite,
events: Mutex<Vec<CallEvent>>,
events: super::support::Observations,
optional_params: Mutex<Vec<Value>>,
}
@ -23,7 +27,7 @@ impl RecordingHost {
Self {
call: LocalMessagesHost::new(call),
rewrite,
events: Mutex::new(Vec::new()),
events: super::support::Observations::default(),
optional_params: Mutex::new(Vec::new()),
}
}
@ -38,7 +42,7 @@ impl RecordingHost {
.unwrap()
.iter()
.filter_map(|event| match event {
CallEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
Some(raw.body.clone())
}
_ => None,
@ -51,22 +55,22 @@ impl RecordingHost {
pub fn request(&self) -> Result<MessagesCall, Error> {
self.call.request()
}
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, ()> {
litellm_host::in_process::Host {
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, ()> {
litellm_host_native::in_process::Host {
services: &(),
hooks: self,
interceptors: self,
stream: &(),
observer: Some(self),
observers: Some(&self.events.sender),
}
}
}
impl litellm_host::lifecycle::CallObserver for RecordingHost {
fn observe(&self, event: litellm_host::event::CallEvent) {
self.events.lock().unwrap().push(event.clone());
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
self.events.sender.emit(event);
}
}
impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protocol>::Error>
impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protocol::Protocol>::Error>
for RecordingHost
{
async fn before_provider_request(
@ -80,20 +84,22 @@ impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protoc
.push(context.optional_params.clone());
(self.rewrite)(wire)
}
async fn on_event(
async fn after_provider_response(
&self,
event: litellm_host::event::MachineEvent,
raw: litellm_host::interceptors::RawResponse,
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
litellm_host::lifecycle::CallEvent::Execution(
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
),
);
Ok(())
}
}
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
litellm_host::in_process::run_hosted(
litellm_host_native::in_process::run_hosted(
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
host.runtime(),
)

View file

@ -99,7 +99,7 @@ fn headers<'a>(pairs: impl IntoIterator<Item = (&'a str, &'a str)>) -> Option<Ma
}
fn machine(secrets: Arc<dyn SecretSource>) -> impl FnOnce(MessagesCall) -> MessagesMachine {
move |request| messages_route(secrets).machine(request)
move |request| messages_route(secrets).machine(request, None)
}
async fn run_with(
@ -107,7 +107,8 @@ async fn run_with(
call: MessagesCall,
) -> Result<MessagesOutput, Error> {
let host = LocalMessagesHost::new(call);
litellm_host::in_process::run_hosted(machine(secrets)(host.request()?), host.runtime()).await
litellm_host_native::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.
@ -144,39 +145,41 @@ impl LocalMessagesHost {
.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 {
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, ()> {
litellm_host_native::in_process::Host {
services: &(),
hooks: self,
interceptors: self,
stream: &(),
observer: Some(self),
observers: None,
}
}
}
impl litellm_host::lifecycle::CallObserver for LocalMessagesHost {
fn observe(&self, _: litellm_host::event::CallEvent) {}
fn observe(&self, _: litellm_host::lifecycle::CallEvent) {}
}
impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protocol>::Error>
impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protocol::Protocol>::Error>
for LocalMessagesHost
{
async fn before_provider_request(
&self,
wire: litellm_host::event::WireRequest,
_: litellm_host::event::RequestContext,
wire: litellm_host::interceptors::WireRequest,
_: litellm_host::interceptors::RequestContext,
) -> Result<
litellm_host::event::WireRequest,
litellm_host::interceptors::WireRequest,
<Messages as litellm_host::protocol::Protocol>::Error,
> {
Ok(wire)
}
async fn on_event(
async fn after_provider_response(
&self,
event: litellm_host::event::MachineEvent,
raw: litellm_host::interceptors::RawResponse,
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
litellm_host::lifecycle::CallEvent::Execution(
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
),
);
Ok(())
}

View file

@ -5,13 +5,19 @@ use rstest::rstest;
use super::*;
#[rstest]
#[case::without_hooks(false)]
#[case::with_hooks(true)]
#[case::neither(false, false)]
#[case::hooks_only(true, false)]
#[case::observer_only(false, true)]
#[case::both(true, true)]
#[tokio::test]
async fn calls_defer_execution_until_polled(call: MessagesCall, #[case] with_hooks: bool) {
async fn calls_defer_execution_until_polled(
call: MessagesCall,
#[case] with_hooks: bool,
#[case] with_observer: bool,
) {
use futures_util::future::BoxFuture;
use litellm_host::event::CallEvent;
use litellm_host::lifecycle::CallEvent;
let upstream = upstream([message_response()]).await;
let secrets = Arc::new(RecordingSecrets::new([("ANTHROPIC_API_KEY", "test-key")]));
@ -21,10 +27,12 @@ async fn calls_defer_execution_until_polled(call: MessagesCall, #[case] with_hoo
..call
});
let request = host.request().unwrap();
let observer: Option<litellm_host::observation::ObservationSender> =
with_observer.then(|| host.events.0.sender.clone());
let future: BoxFuture<'_, Result<MessagesResponse, Error>> = if with_hooks {
Box::pin(route.execute(request, &host))
Box::pin(route.execute(request, &host, observer))
} else {
Box::pin(route.execute(request, &()))
Box::pin(route.execute(request, &(), observer))
};
assert!(secrets.requested().is_empty());
@ -43,18 +51,29 @@ async fn calls_defer_execution_until_polled(call: MessagesCall, #[case] with_hoo
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());
}
assert!(matches!(
(with_hooks, with_observer, events.as_slice()),
(false, false, [])
| (true, false, [])
| (
false,
true,
[
CallEvent::Started { .. },
CallEvent::Execution(_),
CallEvent::Succeeded { .. }
]
)
| (
true,
true,
[
CallEvent::Started { .. },
CallEvent::Execution(_),
CallEvent::Succeeded { .. }
]
)
));
}
#[rstest]
@ -246,6 +265,7 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes
..call
},
&(),
None,
)
.await
.expect("messages request succeeds");

View file

@ -1,4 +1,7 @@
use std::sync::{Mutex, mpsc};
use std::{
ops::ControlFlow,
sync::{Mutex, mpsc},
};
use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt};
@ -6,7 +9,6 @@ use litellm_core::messages::{
MessagesResponse,
route::{Messages, MessagesStreamHead},
};
use litellm_host::protocol::Demand;
use litellm_tracing::{Logger, Metadata, Record, Sink};
use rstest::rstest;
use tokio::{
@ -60,12 +62,12 @@ impl RecordingStreamHost {
}
}
fn record(&self, op: Seen) -> Demand {
fn record(&self, op: Seen) -> ControlFlow<()> {
let mut seen = self.seen.lock().unwrap();
seen.push(op);
match seen.len() < self.detach_after {
true => Demand::More,
false => Demand::Detached,
true => ControlFlow::Continue(()),
false => ControlFlow::Break(()),
}
}
}
@ -74,47 +76,49 @@ impl RecordingStreamHost {
pub fn request(&self) -> Result<MessagesCall, Error> {
self.call.request()
}
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, Self> {
litellm_host::in_process::Host {
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, Self> {
litellm_host_native::in_process::Host {
services: &(),
hooks: self,
interceptors: self,
stream: self,
observer: Some(self),
observers: None,
}
}
}
impl litellm_host::in_process::StreamConsumer<Messages> for RecordingStreamHost {
async fn open_stream(&self, head: MessagesStreamHead) -> Result<Demand, Error> {
impl litellm_host_native::in_process::StreamConsumer<Messages> for RecordingStreamHost {
async fn open_stream(&self, head: MessagesStreamHead) -> Result<ControlFlow<()>, Error> {
Ok(self.record(Seen::Open(head.headers)))
}
async fn send_chunk(&self, chunk: Bytes) -> Result<Demand, Error> {
async fn send_chunk(&self, chunk: Bytes) -> Result<ControlFlow<()>, Error> {
Ok(self.record(Seen::Deliver(chunk)))
}
}
impl litellm_host::lifecycle::CallObserver for RecordingStreamHost {
fn observe(&self, _: litellm_host::event::CallEvent) {}
fn observe(&self, _: litellm_host::lifecycle::CallEvent) {}
}
impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protocol>::Error>
impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protocol::Protocol>::Error>
for RecordingStreamHost
{
async fn before_provider_request(
&self,
wire: litellm_host::event::WireRequest,
_: litellm_host::event::RequestContext,
wire: litellm_host::interceptors::WireRequest,
_: litellm_host::interceptors::RequestContext,
) -> Result<
litellm_host::event::WireRequest,
litellm_host::interceptors::WireRequest,
<Messages as litellm_host::protocol::Protocol>::Error,
> {
Ok(wire)
}
async fn on_event(
async fn after_provider_response(
&self,
event: litellm_host::event::MachineEvent,
raw: litellm_host::interceptors::RawResponse,
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
litellm_host::lifecycle::CallEvent::Execution(
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
),
);
Ok(())
}
@ -136,7 +140,7 @@ fn sse_response() -> ResponseTemplate {
}
async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Error> {
litellm_host::in_process::run_hosted(
litellm_host_native::in_process::run_hosted(
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
host.runtime(),
)
@ -344,6 +348,7 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte(
..streaming(call, upstream.uri())
},
&(),
None,
)
.await
.unwrap();
@ -364,7 +369,7 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte(
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_route(no_secrets())
.execute(streaming(call, upstream.uri()), &())
.execute(streaming(call, upstream.uri()), &(), None)
.await
.err()
.expect("upstream failure is returned by messages()");
@ -395,6 +400,7 @@ async fn dropping_the_sdk_stream_closes_the_unfinished_upstream(
..streaming(call, base)
},
&(),
None,
),
)
.await
@ -431,6 +437,7 @@ async fn the_sdk_yields_a_body_error_once_after_delivered_chunks(call: MessagesC
..streaming(call, base)
},
&(),
None,
)
.await
.unwrap();

View file

@ -3,8 +3,10 @@ use std::{
time::Duration,
};
use litellm_host::event::{CallEvent, MachineEvent};
use litellm_host::lifecycle::CallEvent;
use litellm_host::lifecycle::ExecutionEvent;
use litellm_llms::base_llm::ocr::settings::OcrSettings;
use rstest::rstest;
use super::*;
@ -190,7 +192,7 @@ async fn client_settings_choose_the_api_version_and_the_inch_to_pixel_dpi() {
});
let result = route
.execute(read_request(&upstream.uri(), json!({})), &())
.execute(read_request(&upstream.uri(), json!({})), &(), None)
.await
.unwrap();
@ -281,7 +283,7 @@ async fn response_received_fires_for_the_submission_and_the_completed_poll() {
let recorder = observed.clone();
let host =
LocalOcrHost::new(read_request(&upstream.uri(), json!({}))).with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
if let CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) = event {
recorder.lock().unwrap().push(raw.body.clone());
}
});
@ -354,7 +356,7 @@ async fn the_polling_deadline_bounds_the_retry_delay() {
let error = tokio::time::timeout(
Duration::from_secs(1),
route.execute(read_request(&upstream.uri(), json!({})), &()),
route.execute(read_request(&upstream.uri(), json!({})), &(), None),
)
.await
.expect("the deadline cuts the retry delay short")

View file

@ -1,6 +1,6 @@
use base64::Engine;
use litellm_core::ocr::types::OcrDocumentInput;
use litellm_host::event::WireRequest;
use litellm_host::interceptors::WireRequest;
use rstest::rstest;
use wiremock::{Mock, matchers::any};
@ -208,8 +208,8 @@ async fn configured_client_preserves_document_url_policy(#[case] allowed: bool)
json!({"type": "document_url", "document_url": document_url}),
json!({}),
));
let result = litellm_host::in_process::run_hosted(
route.machine(host.request().unwrap()),
let result = litellm_host_native::in_process::run_hosted(
route.machine(host.request().unwrap(), None),
host.runtime(),
)
.await;

View file

@ -1,10 +1,18 @@
use std::sync::{Arc, Mutex};
use litellm_host::interceptors::RawResponse;
use litellm_host::lifecycle::ExecutionEvent;
use std::sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
};
use litellm_core::ocr::{
route::{Ocr, OcrCall, OcrOp},
types::OcrDocumentInput,
};
use litellm_host::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
use litellm_host::{
interceptors::{RequestContext, WireRequest},
lifecycle::CallEvent,
};
use rstest::rstest;
use super::*;
@ -12,7 +20,7 @@ use super::*;
pub(crate) fn event_name(event: &CallEvent) -> &'static str {
match event {
CallEvent::Started { .. } => "started",
CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response",
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }) => "response",
CallEvent::Succeeded { .. } => "success",
CallEvent::Failed { .. } => "failure",
CallEvent::Cancelled { .. } => "cancelled",
@ -22,15 +30,12 @@ pub(crate) fn event_name(event: &CallEvent) -> &'static str {
fn recording_host(
request: LiteLLMOcrRequest,
events: Arc<Mutex<Vec<&'static str>>>,
interceptions: Arc<AtomicUsize>,
block: bool,
) -> LocalOcrHost {
let before_send_events = events.clone();
LocalOcrHost::new(request)
.with_before_send(move |wire, _| {
before_send_events
.lock()
.unwrap()
.push("before_provider_request");
interceptions.fetch_add(1, Ordering::SeqCst);
match block {
true => Err(Error::InvalidRequest("blocked".into())),
false => Ok(wire),
@ -41,22 +46,22 @@ fn recording_host(
#[rstest::rstest]
#[tokio::test]
async fn hooks_run_in_order_and_one_success_is_emitted() {
async fn interception_runs_once_and_observers_receive_ordered_success_events() {
let upstream = upstream([pages_response()]).await;
let events = Arc::new(Mutex::new(Vec::new()));
let interceptions = Arc::new(AtomicUsize::new(0));
perform_with(recording_host(
ocr_request("mistral/model", &upstream.uri(), json!({})),
events.clone(),
interceptions.clone(),
false,
))
.await
.unwrap();
assert_eq!(
*events.lock().unwrap(),
["started", "before_provider_request", "response", "success"]
);
assert_eq!(interceptions.load(Ordering::SeqCst), 1);
assert_eq!(*events.lock().unwrap(), ["started", "response", "success"]);
assert_eq!(received(&upstream).await.len(), 1);
}
@ -65,10 +70,12 @@ async fn hooks_run_in_order_and_one_success_is_emitted() {
async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() {
let upstream = upstream([pages_response()]).await;
let events = Arc::new(Mutex::new(Vec::new()));
let interceptions = Arc::new(AtomicUsize::new(0));
let error = perform_with(recording_host(
ocr_request("mistral/model", &upstream.uri(), json!({})),
events.clone(),
interceptions.clone(),
true,
))
.await
@ -78,10 +85,8 @@ async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() {
matches!(&error, Error::InvalidRequest(message) if message == "blocked"),
"{error:?}"
);
assert_eq!(
*events.lock().unwrap(),
["started", "before_provider_request", "failure"]
);
assert_eq!(interceptions.load(Ordering::SeqCst), 1);
assert_eq!(*events.lock().unwrap(), ["started", "failure"]);
assert!(received(&upstream).await.is_empty());
}
@ -90,19 +95,19 @@ async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() {
async fn an_upstream_failure_emits_one_terminal_failure() {
let upstream = upstream([status_response(500, json!({"error": "failed"}))]).await;
let events = Arc::new(Mutex::new(Vec::new()));
let interceptions = Arc::new(AtomicUsize::new(0));
let result = perform_with(recording_host(
ocr_request("mistral/model", &upstream.uri(), json!({})),
events.clone(),
interceptions.clone(),
false,
))
.await;
assert!(result.is_err());
assert_eq!(
*events.lock().unwrap(),
["started", "before_provider_request", "failure"]
);
assert_eq!(interceptions.load(Ordering::SeqCst), 1);
assert_eq!(*events.lock().unwrap(), ["started", "failure"]);
assert_eq!(received(&upstream).await.len(), 1);
}
@ -114,7 +119,7 @@ async fn an_invalid_provider_response_is_observed_before_normalization_fails() {
let recorder = observed.clone();
let host = LocalOcrHost::new(ocr_request("mistral/model", &upstream.uri(), json!({})))
.with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
if let CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) = event {
recorder.lock().unwrap().push(raw.body.clone());
}
});
@ -208,16 +213,16 @@ impl CallerTokenHost {
caller_token: true,
})
}
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, Self, Self, ()> {
litellm_host::in_process::Host {
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, Self, Self, ()> {
litellm_host_native::in_process::Host {
services: self,
hooks: self,
interceptors: self,
stream: &(),
observer: Some(self),
observers: None,
}
}
}
impl litellm_host::services::HostCallHandler<Ocr> for CallerTokenHost {
impl litellm_host_native::services::HostCallHandler<Ocr> for CallerTokenHost {
async fn handle_host_call(&self, op: OcrOp) -> Result<(), Error> {
match op {
OcrOp::AcquireAzureAdToken(reply) => {
@ -232,9 +237,9 @@ impl litellm_host::services::HostCallHandler<Ocr> for CallerTokenHost {
}
impl litellm_host::lifecycle::CallObserver for CallerTokenHost {
fn observe(&self, _: litellm_host::event::CallEvent) {}
fn observe(&self, _: litellm_host::lifecycle::CallEvent) {}
}
impl litellm_host::hooks::RouteHooks<<Ocr as litellm_host::protocol::Protocol>::Error>
impl litellm_host::interceptors::Interceptors<<Ocr as litellm_host::protocol::Protocol>::Error>
for CallerTokenHost
{
async fn before_provider_request(
@ -263,13 +268,15 @@ impl litellm_host::hooks::RouteHooks<<Ocr as litellm_host::protocol::Protocol>::
.collect();
Ok(WireRequest { headers, ..wire })
}
async fn on_event(
async fn after_provider_response(
&self,
event: litellm_host::event::MachineEvent,
raw: litellm_host::interceptors::RawResponse,
) -> Result<(), <Ocr as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
litellm_host::lifecycle::CallEvent::Execution(
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
),
);
Ok(())
}
@ -288,8 +295,8 @@ async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_
trace: Mutex::new(Vec::new()),
};
litellm_host::in_process::run_hosted(
ocr_route().machine(host.request().unwrap()),
litellm_host_native::in_process::run_hosted(
ocr_route().machine(host.request().unwrap(), None),
host.runtime(),
)
.await
@ -312,15 +319,11 @@ async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_
#[rstest]
#[tokio::test]
async fn direct_execution_uses_hooks_without_a_machine() {
use litellm_host::{hooks::RouteHooks, lifecycle::CallObserver};
use litellm_host::interceptors::Interceptors;
struct Hooks(Arc<super::support::CallEvents>);
impl RouteHooks<Error> for Hooks {
fn observer(&self) -> Option<Arc<dyn CallObserver>> {
Some(self.0.clone())
}
struct Hooks;
impl Interceptors<Error> for Hooks {
async fn before_provider_request(
&self,
wire: WireRequest,
@ -336,8 +339,7 @@ async fn direct_execution_uses_hooks_without_a_machine() {
})
}
async fn on_event(&self, event: MachineEvent) -> Result<(), Error> {
self.0.observe(CallEvent::Machine(event));
async fn after_provider_response(&self, _: RawResponse) -> Result<(), Error> {
Ok(())
}
}
@ -348,10 +350,11 @@ async fn direct_execution_uses_hooks_without_a_machine() {
.await;
let events = Arc::new(super::support::CallEvents::default());
let route = ocr_route();
let hooks = Hooks(events.clone());
let interceptors = Hooks;
let builder = route.execute(
ocr_request("mistral/model", &upstream.uri(), json!({})),
&hooks,
&interceptors,
Some(events.0.sender.clone()),
);
assert!(events.0.lock().unwrap().is_empty());
assert!(received(&upstream).await.is_empty());
@ -365,7 +368,7 @@ async fn direct_execution_uses_hooks_without_a_machine() {
&events.0.lock().unwrap()[..],
[
CallEvent::Started { .. },
CallEvent::Machine(_),
CallEvent::Execution(_),
CallEvent::Succeeded { .. }
]
));

View file

@ -1,4 +1,3 @@
use litellm_host::protocol::HookRequest;
use std::{
sync::{
Arc,
@ -12,17 +11,16 @@ use litellm_core::ocr::{
types::OcrDocumentInput,
};
use litellm_host::{
event::{CallEvent, WireRequest},
hooks::RouteHooks,
interceptors::{Interceptors, WireRequest},
machine::{HostFailure, Machine, MachineStep},
protocol::Suspension,
services::HostCallHandler,
protocol::{HostRequest, InterceptRequest},
};
use litellm_host_native::services::HostCallHandler;
use litellm_llms::base_llm::ocr::transformation::OcrTransportConfig;
use rstest::rstest;
use tokio::{io::AsyncReadExt, net::TcpListener, sync::Notify};
use super::{lifecycle::event_name, *};
use super::*;
/// 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.
@ -34,7 +32,7 @@ async fn drive_until(
Vec<&'static str>,
OcrMachine,
) {
let mut machine = ocr_route().machine(host.request().unwrap());
let mut machine = ocr_route().machine(host.request().unwrap(), None);
let mut ops = Vec::new();
let outcome = loop {
let op = match machine.resume().await {
@ -43,23 +41,25 @@ async fn drive_until(
Err(error) => break Err(error),
};
let answer = match op {
Suspension::Stream(stream) => match stream {
HostRequest::Stream(stream) => match stream {
litellm_host::protocol::StreamDelivery::Open(head, _) => match head {},
litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {},
},
Suspension::HostCall(op) => {
HostRequest::HostCall(op) => {
ops.push(match op {
OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken",
});
host.handle_host_call(op).await.map_err(HostFailure::Error)
}
Suspension::Hook(HookRequest::BeforeProviderRequest { wire, reply, .. }) => {
HostRequest::Intercept(InterceptRequest::BeforeProviderRequest {
wire, reply, ..
}) => {
ops.push("BeforeSend");
intercept(*wire).map(|wire| reply.send(wire))
}
Suspension::Hook(HookRequest::Event(event, reply)) => {
ops.push(event_name(&CallEvent::Machine(event.clone())));
host.on_event(event)
HostRequest::Intercept(InterceptRequest::AfterProviderResponse { raw, reply }) => {
ops.push("response");
host.after_provider_response(raw)
.await
.map(|()| reply.send(()))
.map_err(HostFailure::Error)
@ -80,10 +80,10 @@ async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, sto
_ = stop.notified() => break,
step = machine.resume() => {
match step.unwrap() {
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 {
MachineStep::Suspended(HostRequest::HostCall(op)) => host.handle_host_call(op).await.unwrap(),
MachineStep::Suspended(HostRequest::Intercept(InterceptRequest::BeforeProviderRequest { wire, reply, .. })) => reply.send(*wire),
MachineStep::Suspended(HostRequest::Intercept(InterceptRequest::AfterProviderResponse { reply, .. })) => reply.send(()),
MachineStep::Suspended(HostRequest::Stream(stream)) => match stream {
litellm_host::protocol::StreamDelivery::Open(head, _) => match head {},
litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {},
},
@ -181,15 +181,16 @@ async fn a_before_send_failure_ends_the_call_without_reaching_transport(
async fn resuming_before_answering_keeps_the_pending_operation() {
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
let mut machine = ocr_route().machine(
OcrCall {
request,
caller_token: false,
},
None,
);
let Ok(MachineStep::Suspended(HostRequest::Intercept(
InterceptRequest::BeforeProviderRequest { wire, reply, .. },
))) = machine.resume().await
else {
panic!("expected the provider request hook");
};
@ -197,8 +198,8 @@ async fn resuming_before_answering_keeps_the_pending_operation() {
reply.send(*wire);
assert!(matches!(
machine.resume().await,
Ok(MachineStep::Suspended(Suspension::Hook(
HookRequest::Event(_, _)
Ok(MachineStep::Suspended(HostRequest::Intercept(
InterceptRequest::AfterProviderResponse { .. }
)))
));
}
@ -244,7 +245,7 @@ async fn interrupt_drops_provider_captures_before_returning() {
},
)));
let host = LocalOcrHost::new(request);
let mut machine = ocr_route().machine(host.request().unwrap());
let mut machine = ocr_route().machine(host.request().unwrap(), None);
drive_until_notified(&mut machine, &host, &entered).await;
assert!(!dropped.load(Ordering::SeqCst));
@ -280,7 +281,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_route().machine(host.request().unwrap());
let mut machine = ocr_route().machine(host.request().unwrap(), None);
drive_until_notified(&mut machine, &host, &received).await;
let cancelled = Error::InvalidRequest("cancelled".into());

View file

@ -5,7 +5,10 @@ use litellm_core::ocr::{
types::{LiteLLMOcrRequest, OcrDocumentInput},
wire::{OcrWireRequest, decode_request},
};
use litellm_host::event::{CallEvent, RequestContext, WireRequest};
use litellm_host::{
interceptors::{RequestContext, WireRequest},
lifecycle::CallEvent,
};
use litellm_llms::base_llm::ocr::{
error::Error,
settings::OcrSettings,
@ -57,13 +60,22 @@ fn ocr_route_with(settings: OcrSettings) -> OcrRoute {
}
async fn perform(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
ocr_route().execute(request, &()).await
ocr_route().execute(request, &(), None).await
}
async fn perform_with(host: LocalOcrHost) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::in_process::run_hosted(ocr_route().machine(host.request()?), host.runtime())
.await
.map(completed)
let result = litellm_host_native::in_process::run_hosted(
ocr_route().machine(host.request()?, None),
host.runtime(),
)
.await
.map(completed);
if let Some(observer) = &host.observer {
for event in host.events.0.lock().unwrap().iter() {
observer(event);
}
}
result
}
fn wire(model: &str, base: &str, document: Value, options: Value) -> OcrWireRequest {
@ -155,6 +167,7 @@ struct LocalOcrHost {
request: Mutex<Option<LiteLLMOcrRequest<OcrDocumentInput>>>,
before_provider_request: Option<BeforeSend>,
observer: Option<Observer>,
events: support::CallEvents,
}
impl LocalOcrHost {
@ -163,6 +176,7 @@ impl LocalOcrHost {
request: Mutex::new(Some(request)),
before_provider_request: None,
observer: None,
events: support::CallEvents::default(),
}
}
@ -199,16 +213,16 @@ impl LocalOcrHost {
})
.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 {
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, Self, Self, ()> {
litellm_host_native::in_process::Host {
services: self,
hooks: self,
interceptors: self,
stream: &(),
observer: Some(self),
observers: Some(&self.events.0.sender),
}
}
}
impl litellm_host::services::HostCallHandler<Ocr> for LocalOcrHost {
impl litellm_host_native::services::HostCallHandler<Ocr> for LocalOcrHost {
async fn handle_host_call(&self, op: OcrOp) -> Result<(), Error> {
match op {
OcrOp::AcquireAzureAdToken(_) => {
@ -221,13 +235,11 @@ impl litellm_host::services::HostCallHandler<Ocr> for LocalOcrHost {
}
impl litellm_host::lifecycle::CallObserver for LocalOcrHost {
fn observe(&self, event: litellm_host::event::CallEvent) {
if let Some(observer) = &self.observer {
observer(&event);
}
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
self.events.0.sender.emit(event);
}
}
impl litellm_host::hooks::RouteHooks<<Ocr as litellm_host::protocol::Protocol>::Error>
impl litellm_host::interceptors::Interceptors<<Ocr as litellm_host::protocol::Protocol>::Error>
for LocalOcrHost
{
async fn before_provider_request(
@ -240,13 +252,15 @@ impl litellm_host::hooks::RouteHooks<<Ocr as litellm_host::protocol::Protocol>::
None => Ok(wire),
}
}
async fn on_event(
async fn after_provider_response(
&self,
event: litellm_host::event::MachineEvent,
raw: litellm_host::interceptors::RawResponse,
) -> Result<(), <Ocr as litellm_host::protocol::Protocol>::Error> {
litellm_host::lifecycle::CallObserver::observe(
self,
litellm_host::event::CallEvent::Machine(event),
litellm_host::lifecycle::CallEvent::Execution(
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
),
);
Ok(())
}

View file

@ -172,7 +172,7 @@ async fn missing_credentials_come_from_the_injected_secret_source(
})
.unwrap();
route.execute(request, &()).await.unwrap();
route.execute(request, &(), None).await.unwrap();
assert_eq!(source.requested(), MistralOcrConfig.secret_names());
assert_eq!(
@ -201,6 +201,7 @@ async fn the_client_uses_the_injected_http_pool_configuration() {
.execute(
ocr_request("mistral/model", &upstream.uri(), json!({})),
&(),
None,
)
.await
.unwrap();

View file

@ -1,6 +1,7 @@
use litellm_host::lifecycle::ExecutionEvent;
use std::sync::{Arc, Mutex};
use litellm_host::event::{CallEvent, MachineEvent, WireRequest};
use litellm_host::{interceptors::WireRequest, lifecycle::CallEvent};
use rstest::rstest;
use super::*;
@ -147,7 +148,7 @@ async fn response_received_fires_once_for_the_parse_response() {
let recorder = observed.clone();
let host = LocalOcrHost::new(ocr_request("reducto/parse-v3", &upstream.uri(), json!({})))
.with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
if let CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) = event {
recorder.lock().unwrap().push(raw.body.clone());
}
});

View file

@ -55,6 +55,7 @@ async fn configured_project_and_location_apply_when_the_call_sets_neither() {
.execute(
ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})),
&(),
None,
)
.await
.unwrap();

View file

@ -123,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 = route.execute(request, &()).await.unwrap();
let result = route.execute(request, &(), None).await.unwrap();
assert!(!result.pages.is_empty());
}
let requests = upstream.received_requests().await.unwrap();

View file

@ -5,7 +5,7 @@ use litellm_core::responses::{
route::Responses,
types::{ResponsesCall, ResponsesOutput},
};
use litellm_host::{call::HostedCompletion, event::CallEvent};
use litellm_host::{call::HostedCompletion, lifecycle::CallEvent};
use rstest::{fixture, rstest};
use serde_json::json;
use wiremock::ResponseTemplate;
@ -39,8 +39,9 @@ async fn http_responses_share_execution_and_hooks(call: ResponsesCall, #[case] h
..call
});
let response = if hosted {
let HostedCompletion::Complete(response) = litellm_host::in_process::run_hosted(
responses_route(no_secrets()).machine(host.request().unwrap()),
let HostedCompletion::Complete(response) = litellm_host_native::in_process::run_hosted(
responses_route(no_secrets())
.machine(host.request().unwrap(), Some(host.events.0.sender.clone())),
host.runtime(),
)
.await
@ -51,7 +52,7 @@ async fn http_responses_share_execution_and_hooks(call: ResponsesCall, #[case] h
} else {
let call = host.request.lock().unwrap().take().unwrap();
let ResponsesOutput::Complete(response) = responses_route(no_secrets())
.execute(call, &host)
.execute(call, &host, Some(host.events.0.sender.clone()))
.await
.unwrap()
else {
@ -69,7 +70,7 @@ async fn http_responses_share_execution_and_hooks(call: ResponsesCall, #[case] h
&host.events.0.lock().unwrap()[..],
[
CallEvent::Started { .. },
CallEvent::Machine(_),
CallEvent::Execution(_),
CallEvent::Succeeded { .. }
]
));
@ -95,8 +96,8 @@ async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption(
});
let (headers, bytes) = if hosted {
assert_eq!(
litellm_host::in_process::run_hosted(
responses_route(no_secrets()).machine(host.request().unwrap()),
litellm_host_native::in_process::run_hosted(
responses_route(no_secrets()).machine(host.request().unwrap(), None,),
host.runtime(),
)
.await
@ -110,7 +111,7 @@ async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption(
} else {
let call = host.request.lock().unwrap().take().unwrap();
let ResponsesOutput::Stream { head, chunks } = responses_route(no_secrets())
.execute(call, &host)
.execute(call, &host, Some(host.events.0.sender.clone()))
.await
.unwrap()
else {
@ -147,7 +148,7 @@ async fn provider_failures_emit_failure_once(
let call = host.request.lock().unwrap().take().unwrap();
assert!(
responses_route(no_secrets())
.execute(call, &host)
.execute(call, &host, Some(host.events.0.sender.clone()))
.await
.is_err()
);
@ -190,7 +191,7 @@ async fn credentials_and_endpoint_are_resolved_only_when_needed(
..call
};
responses_route(secrets.clone())
.execute(call, &())
.execute(call, &(), None)
.await
.unwrap();
assert_eq!(
@ -224,7 +225,7 @@ async fn unsupported_providers_fail_before_secrets_or_transport(
};
assert!(
responses_route(secrets.clone())
.execute(call, &())
.execute(call, &(), None)
.await
.is_err()
);
@ -257,15 +258,17 @@ async fn route_tracing_covers_native_and_hosted_outcomes(
.logger()
.instrument(async {
if hosted {
litellm_host::in_process::run_hosted(
route.clone().machine(host.request().unwrap()),
litellm_host_native::in_process::run_hosted(
route
.clone()
.machine(host.request().unwrap(), Some(host.events.0.sender.clone())),
host.runtime(),
)
.await
.map(|_| ())
} else {
route
.execute(host.request().unwrap(), &())
.execute(host.request().unwrap(), &(), None)
.await
.map(|_| ())
}
@ -311,6 +314,7 @@ async fn stream_trace_survives_handoff_and_closes_before_the_stream_object_is_dr
..call
},
&(),
None,
)
.await
})
@ -350,6 +354,7 @@ async fn preparation_failure_is_traced_but_unpolled_builders_are_not(
..call
},
&(),
None,
))
});
assert!(traces.records().is_empty());
@ -363,6 +368,7 @@ async fn preparation_failure_is_traced_but_unpolled_builders_are_not(
..self::call()
},
&(),
None,
)
.await
})

View file

@ -3,7 +3,10 @@
#![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset
use std::sync::{Arc, Mutex};
use std::{
ops::ControlFlow,
sync::{Arc, Mutex},
};
use futures_util::future::BoxFuture;
use litellm_http::{
@ -248,11 +251,43 @@ pub struct RecordingCall<P: litellm_host::protocol::Protocol> {
}
#[derive(Default)]
pub struct CallEvents(pub Mutex<Vec<litellm_host::event::CallEvent>>);
pub struct CallEvents(pub Observations);
pub struct Observations {
pub sender: litellm_host::observation::ObservationSender,
receiver: Mutex<tokio::sync::mpsc::Receiver<litellm_host::lifecycle::CallEvent>>,
recorded: Mutex<Vec<litellm_host::lifecycle::CallEvent>>,
}
impl Default for Observations {
fn default() -> Self {
let (sender, receiver) = litellm_host::observation::observation_channel(
std::num::NonZeroUsize::new(128).unwrap(),
);
Self {
sender,
receiver: Mutex::new(receiver),
recorded: Mutex::new(Vec::new()),
}
}
}
impl Observations {
pub fn lock(
&self,
) -> std::sync::LockResult<std::sync::MutexGuard<'_, Vec<litellm_host::lifecycle::CallEvent>>>
{
let mut events = self.recorded.lock()?;
let mut receiver = self.receiver.lock().unwrap();
while let Ok(event) = receiver.try_recv() {
events.push(event);
}
Ok(events)
}
}
impl litellm_host::lifecycle::CallObserver for CallEvents {
fn observe(&self, event: litellm_host::event::CallEvent) {
self.0.lock().unwrap().push(event);
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
self.0.sender.emit(event);
}
}
@ -267,19 +302,15 @@ impl<P: litellm_host::protocol::Protocol> RecordingCall<P> {
}
}
impl<P: litellm_host::protocol::Protocol> litellm_host::hooks::RouteHooks<P::Error>
impl<P: litellm_host::protocol::Protocol> litellm_host::interceptors::Interceptors<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 {
wire: litellm_host::interceptors::WireRequest,
_: litellm_host::interceptors::RequestContext,
) -> Result<litellm_host::interceptors::WireRequest, P::Error> {
Ok(litellm_host::interceptors::WireRequest {
headers: wire
.headers
.into_iter()
@ -289,12 +320,10 @@ impl<P: litellm_host::protocol::Protocol> litellm_host::hooks::RouteHooks<P::Err
})
}
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));
async fn after_provider_response(
&self,
_: litellm_host::interceptors::RawResponse,
) -> Result<(), P::Error> {
Ok(())
}
}
@ -311,34 +340,28 @@ where
.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 {
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, Self> {
litellm_host_native::in_process::Host {
services: &(),
hooks: self,
interceptors: self,
stream: self,
observer: Some(self),
observers: Some(&self.events.0.sender),
}
}
}
impl<P> litellm_host::in_process::StreamConsumer<P> for RecordingCall<P>
impl<P> litellm_host_native::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> {
async fn open_stream(&self, head: P::StreamHead) -> Result<ControlFlow<()>, P::Error> {
*self.head.lock().unwrap() = Some(head);
Ok(litellm_host::protocol::Demand::More)
Ok(ControlFlow::Continue(()))
}
async fn send_chunk(
&self,
chunk: P::Chunk,
) -> Result<litellm_host::protocol::Demand, P::Error> {
async fn send_chunk(&self, chunk: P::Chunk) -> Result<ControlFlow<()>, P::Error> {
self.chunks.lock().unwrap().push(chunk);
Ok(litellm_host::protocol::Demand::More)
Ok(ControlFlow::Continue(()))
}
}
impl<P> litellm_host::lifecycle::CallObserver for RecordingCall<P>
@ -346,8 +369,8 @@ 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());
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
self.events.0.sender.emit(event);
}
}

View file

@ -44,10 +44,8 @@ async fn handle(
request::authorize_model(identity, deployment, &body).await?;
let messages = body.get("messages").cloned().unwrap_or_default();
let response = litellm_host_http::serve_unary(
gateway
.chat_completions
.clone()
.machine(ChatCompletionsCall {
gateway.chat_completions.clone().machine(
ChatCompletionsCall {
model: deployment.model.clone(),
messages,
optional_params: body
@ -59,10 +57,13 @@ async fn handle(
custom_llm_provider: deployment.custom_llm_provider.clone(),
extra_headers: None,
timeout: deployment.timeout,
}),
},
None,
),
(),
(),
litellm_host_http::Unary::new(Json),
None,
)
.await?;
Ok(response)

View file

@ -44,10 +44,10 @@ async fn handle(
let deployment = request::resolve_deployment(gateway, &body)?;
request::authorize_model(identity, deployment, &body).await?;
let call = project(deployment, body, headers)?;
let machine = gateway.messages.clone().machine(call);
let machine = gateway.messages.clone().machine(call, None);
let stream =
Sse::<Messages, _, _>::new(Json, |error| Bytes::from(Error::from(error).sse_frame()));
Ok(litellm_host_http::serve(machine, (), (), stream).await?)
Ok(litellm_host_http::serve(machine, (), (), stream, None).await?)
}
fn project(

View file

@ -70,7 +70,7 @@ async fn handle(
..Default::default()
},
)?;
let response = gateway.ocr.execute(call, &()).await?;
let response = gateway.ocr.execute(call, &(), None).await?;
match response.provider_native_response {
Some(native) => Ok(Value::Object(native)),
None => Ok(response.into_json()),

View file

@ -28,7 +28,7 @@ pub(crate) async fn create(
extra_headers: None,
timeout: deployment.timeout,
};
let machine = gateway.responses.clone().machine(call);
let machine = gateway.responses.clone().machine(call, None);
let stream = Sse::<Responses, _, _>::new(Json, |error| {
let error = Error::from(error);
Bytes::from(format!(
@ -36,5 +36,5 @@ pub(crate) async fn create(
json!({"type": "error", "code": error.status().as_u16().to_string(), "message": error.to_string(), "param": null})
))
});
Ok(litellm_host_http::serve(machine, (), (), stream).await?)
Ok(litellm_host_http::serve(machine, (), (), stream, None).await?)
}

View file

@ -110,6 +110,7 @@ async fn chat_errors_come_from_core(
timeout: None,
},
&(),
None,
)
.await
.unwrap_err();

View file

@ -1,5 +1,7 @@
Own the HTTP driver for hosted calls, including response-body demand, cancellation, and lifecycle observation
`serve` and `serve_unary` receive an optional `ObservationSender` separately from active hooks. Observation works with `()` hooks and covers response conversion and body delivery
Keep endpoint paths, request parsing, deployment selection, and API-specific response and error formats in gateway-inference
Depend on the neutral host protocol, never on core routes, Python, or gateway crates

View file

@ -11,6 +11,7 @@ bytes.workspace = true
futures-util.workspace = true
http.workspace = true
litellm-host.workspace = true
litellm-host-native.workspace = true
thiserror.workspace = true
[dev-dependencies]

View file

@ -1,3 +1,4 @@
use litellm_host::observation::ObservationSender;
use std::{convert::Infallible, sync::Arc};
use axum::{body::Body, response::Response};
@ -6,36 +7,36 @@ use futures_util::{StreamExt, stream};
use litellm_host::{
call::{CallOutput, HostedCompletion, HostedMachine},
hooks::RouteHooks,
interceptors::Interceptors,
lifecycle::{observe_call, observe_unary},
machine::{Machine, MachineFault, MachineStep},
protocol::{Demand, HookRequest, Protocol, Reply, StreamDelivery, Suspension},
services::HostCallHandler,
machine::MachineFault,
protocol::Protocol,
};
use litellm_host_native::{Boundary, Driver, services::HostCallHandler};
use crate::{Error, ResponseEncoder, StreamEncoder};
type StepOf<P> = MachineStep<P, HostedCompletion<<P as Protocol>::Response>>;
type Output<E> = CallOutput<Response, http::Response<()>, Bytes, E>;
type HostedDriver<P, S, H> = Driver<HostedMachine<P>, S, H>;
pub async fn serve_unary<P, A, H, S>(
machine: HostedMachine<P>,
services: S,
hooks: H,
interceptors: H,
encoder: A,
observers: Option<ObservationSender>,
) -> Result<Response, Error<P::Error>>
where
P: Protocol<Chunk = Infallible, StreamHead = Infallible>,
P::Error: From<MachineFault>,
H: RouteHooks<P::Error>,
H: Interceptors<P::Error>,
S: HostCallHandler<P>,
A: ResponseEncoder<Protocol = P>,
{
let observer = hooks.observer();
let mut driver = Driver::new(machine, services, hooks);
observe_unary(observer, async move {
match driver.advance().await? {
MachineStep::Complete(HostedCompletion::Complete(value)) => {
let mut driver = Driver::new(machine, services, interceptors);
observe_unary(observers, async move {
match driver.advance().await.map_err(Error::Call)? {
Boundary::Complete(HostedCompletion::Complete(value)) => {
encoder.encode_response(value).map_err(Error::Call)
}
_ => Err(Error::Protocol),
@ -47,20 +48,20 @@ where
pub async fn serve<P, A, H, S>(
machine: HostedMachine<P>,
services: S,
hooks: H,
interceptors: H,
encoder: A,
observers: Option<ObservationSender>,
) -> Result<Response, Error<P::Error>>
where
P: Protocol,
P::Error: From<MachineFault>,
A: StreamEncoder<Protocol = P>,
H: RouteHooks<P::Error> + 'static,
H: Interceptors<P::Error> + 'static,
S: HostCallHandler<P> + 'static,
{
let encoder = Arc::new(encoder);
let observer = hooks.observer();
let driver = Driver::new(machine, services, hooks);
match observe_call(observer, driver.start(encoder.clone())).await? {
let driver = Driver::new(machine, services, interceptors);
match observe_call(observers, start(driver, encoder.clone())).await? {
CallOutput::Complete(response) => Ok(response),
CallOutput::Stream { head, chunks } => {
let body = chunks.map(move |chunk| {
@ -73,93 +74,38 @@ where
}
}
struct Driver<P: Protocol, H, S> {
machine: HostedMachine<P>,
services: S,
hooks: H,
demand: Option<Reply<Demand>>,
}
impl<P, H, S> Driver<P, H, S>
async fn start<P, A, H, S>(
mut driver: HostedDriver<P, S, H>,
encoder: Arc<A>,
) -> Result<Output<Error<P::Error>>, Error<P::Error>>
where
P: Protocol,
P::Error: From<MachineFault>,
H: RouteHooks<P::Error>,
S: HostCallHandler<P>,
A: StreamEncoder<Protocol = P>,
H: Interceptors<P::Error> + 'static,
S: HostCallHandler<P> + 'static,
{
fn new(machine: HostedMachine<P>, services: S, hooks: H) -> Self {
Self {
machine,
services,
hooks,
demand: None,
}
}
async fn advance(&mut self) -> Result<StepOf<P>, Error<P::Error>> {
if let Some(reply) = self.demand.take() {
reply.send(Demand::More);
}
loop {
match self.machine.resume().await.map_err(Error::Call)? {
MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest {
wire,
context,
reply,
})) => reply.send(
self.hooks
.before_provider_request(*wire, *context)
.await
.map_err(Error::Call)?,
),
MachineStep::Suspended(Suspension::Hook(HookRequest::Event(event, reply))) => {
self.hooks.on_event(event).await.map_err(Error::Call)?;
reply.send(());
}
MachineStep::Suspended(Suspension::HostCall(op)) => {
self.services
.handle_host_call(op)
.await
.map_err(Error::Call)?;
}
boundary => return Ok(boundary),
}
}
}
async fn start<A>(mut self, encoder: Arc<A>) -> Result<Output<Error<P::Error>>, Error<P::Error>>
where
A: StreamEncoder<Protocol = P>,
H: 'static,
S: 'static,
{
match self.advance().await? {
MachineStep::Complete(HostedCompletion::Complete(value)) => encoder
.encode_response(value)
.map(CallOutput::Complete)
.map_err(Error::Call),
MachineStep::Suspended(Suspension::Stream(StreamDelivery::Open(head, reply))) => {
let head = encoder.encode_stream_head(head).map_err(Error::Call)?;
self.demand = Some(reply);
let chunks =
stream::try_unfold((self, encoder), |(mut driver, encoder)| async move {
match driver.advance().await? {
MachineStep::Suspended(Suspension::Stream(StreamDelivery::Chunk(
chunk,
reply,
))) => {
let bytes = encoder.encode_chunk(chunk).map_err(Error::Call)?;
driver.demand = Some(reply);
Ok(Some((bytes, (driver, encoder))))
}
MachineStep::Complete(HostedCompletion::StreamEnded) => Ok(None),
_ => Err(Error::Protocol),
match driver.advance().await.map_err(Error::Call)? {
Boundary::Complete(HostedCompletion::Complete(value)) => encoder
.encode_response(value)
.map(CallOutput::Complete)
.map_err(Error::Call),
Boundary::Open(head) => {
let head = encoder.encode_stream_head(head).map_err(Error::Call)?;
let chunks =
stream::try_unfold((driver, encoder), |(mut driver, encoder)| async move {
match driver.advance().await.map_err(Error::Call)? {
Boundary::Chunk(chunk) => {
let bytes = encoder.encode_chunk(chunk).map_err(Error::Call)?;
Ok(Some((bytes, (driver, encoder))))
}
})
.boxed();
Ok(CallOutput::Stream { head, chunks })
}
_ => Err(Error::Protocol),
Boundary::Complete(HostedCompletion::StreamEnded) => Ok(None),
_ => Err(Error::Protocol),
}
})
.boxed();
Ok(CallOutput::Stream { head, chunks })
}
_ => Err(Error::Protocol),
}
}

View file

@ -1,3 +1,4 @@
use litellm_host::lifecycle::ExecutionEvent;
use std::sync::{
Arc, Mutex,
atomic::{AtomicBool, AtomicUsize, Ordering},
@ -13,12 +14,10 @@ use futures_util::{StreamExt, stream};
use http::{StatusCode, header::CONTENT_TYPE};
use litellm_host::{
call::{CallOutput, hosted_call},
event::{CallEvent, MachineEvent, RawResponse, RequestContext, WireRequest},
hooks::RouteHooks,
lifecycle::CallObserver,
interceptors::{Interceptors, RawResponse, RequestContext, WireRequest},
lifecycle::{CallEvent, CallObserver},
machine::MachineFault,
protocol::Protocol,
protocol::Reply,
protocol::{Protocol, Reply},
};
use litellm_host_http::{Error, ResponseEncoder, StreamEncoder, Unary, serve, serve_unary};
use rstest::{fixture, rstest};
@ -71,7 +70,7 @@ impl ResponseEncoder for Adapter {
}
}
impl litellm_host::services::HostCallHandler<TestProtocol> for Adapter {
impl litellm_host_native::services::HostCallHandler<TestProtocol> for Adapter {
async fn handle_host_call(&self, reply: Reply<&'static str>) -> Result<(), TestError> {
if self.0 == Rejection::Custom {
return Err(TestError::Adapter);
@ -109,11 +108,40 @@ impl StreamEncoder for Adapter {
}
#[derive(Default)]
struct Observer(Mutex<Vec<CallEvent>>);
struct Observer(Observations);
struct Observations {
sender: litellm_host::observation::ObservationSender,
receiver: Mutex<tokio::sync::mpsc::Receiver<CallEvent>>,
recorded: Mutex<Vec<CallEvent>>,
}
impl Default for Observations {
fn default() -> Self {
let (sender, receiver) = litellm_host::observation::observation_channel(
std::num::NonZeroUsize::new(128).unwrap(),
);
Self {
sender,
receiver: Mutex::new(receiver),
recorded: Mutex::new(Vec::new()),
}
}
}
impl Observations {
fn lock(&self) -> std::sync::LockResult<std::sync::MutexGuard<'_, Vec<CallEvent>>> {
let mut events = self.recorded.lock()?;
let mut receiver = self.receiver.lock().unwrap();
while let Ok(event) = receiver.try_recv() {
events.push(event);
}
Ok(events)
}
}
impl CallObserver for Observer {
fn observe(&self, event: CallEvent) {
self.0.lock().unwrap().push(event);
self.0.sender.emit(event);
}
}
@ -122,11 +150,7 @@ struct Hooks {
reject: bool,
}
impl RouteHooks<TestError> for Hooks {
fn observer(&self) -> Option<Arc<dyn CallObserver>> {
Some(self.observer.clone())
}
impl Interceptors<TestError> for Hooks {
async fn before_provider_request(
&self,
wire: WireRequest,
@ -141,11 +165,13 @@ impl RouteHooks<TestError> for Hooks {
})
}
async fn on_event(&self, event: MachineEvent) -> Result<(), TestError> {
async fn after_provider_response(&self, raw: RawResponse) -> Result<(), TestError> {
if self.reject {
return Err(TestError::Hook);
}
self.observer.observe(CallEvent::Machine(event));
self.observer.observe(CallEvent::Execution(
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
));
Ok(())
}
}
@ -156,7 +182,7 @@ fn observer() -> Arc<Observer> {
}
#[fixture]
fn hooks(observer: Arc<Observer>) -> Hooks {
fn interceptors(observer: Arc<Observer>) -> Hooks {
Hooks {
observer,
reject: false,
@ -165,11 +191,12 @@ fn hooks(observer: Arc<Observer>) -> Hooks {
#[rstest]
#[tokio::test]
async fn projection_custom_operations_and_hooks_feed_the_http_response(hooks: Hooks) {
let observer = hooks.observer.clone();
async fn projection_custom_operations_and_hooks_feed_the_http_response(interceptors: Hooks) {
let observer = interceptors.observer.clone();
let machine = hosted_call::<TestProtocol, _, _>(
"projected",
|request, services, route_hooks| async move {
None,
|request, services, route_hooks, _observations| async move {
let custom = services.call(|reply| reply).await?;
let wire = route_hooks
.before_provider_request(
@ -188,10 +215,8 @@ async fn projection_custom_operations_and_hooks_feed_the_http_response(hooks: Ho
)
.await?;
route_hooks
.on_event(MachineEvent::ResponseReceived {
raw: RawResponse {
body: wire.url.clone(),
},
.after_provider_response(RawResponse {
body: wire.url.clone(),
})
.await?;
Ok(CallOutput::Complete(Bytes::from(wire.url)))
@ -200,8 +225,9 @@ async fn projection_custom_operations_and_hooks_feed_the_http_response(hooks: Ho
let response = serve(
machine,
Adapter(Rejection::None),
hooks,
interceptors,
Adapter(Rejection::None),
Some(observer.0.sender.clone()),
)
.await
.unwrap();
@ -214,7 +240,7 @@ async fn projection_custom_operations_and_hooks_feed_the_http_response(hooks: Ho
let events = observer.0.lock().unwrap();
assert!(matches!(events.as_slice(), [
CallEvent::Started { .. },
CallEvent::Machine(MachineEvent::ResponseReceived { raw }),
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }),
CallEvent::Succeeded { .. },
] if raw.body == "projected/custom"));
}
@ -234,32 +260,36 @@ impl Drop for Release {
#[case::dropped_before_eof(Some(2))]
#[tokio::test]
async fn body_demand_controls_polling_and_lifecycle(
hooks: Hooks,
observer: Arc<Observer>,
#[case] drop_after: Option<usize>,
) {
let observer = hooks.observer.clone();
let polls = Arc::new(AtomicUsize::new(0));
let released = Arc::new(AtomicBool::new(false));
let provider_polls = polls.clone();
let release = Release(released.clone());
let machine = hosted_call::<TestProtocol, _, _>("input", move |_, _, _| async move {
let chunks = stream::unfold((0, release), move |(index, release)| {
provider_polls.fetch_add(1, Ordering::SeqCst);
async move {
(index < 2).then(|| (Ok(Bytes::from(index.to_string())), (index + 1, release)))
}
})
.boxed();
Ok(CallOutput::Stream {
head: "text/event-stream",
chunks,
})
});
let machine = hosted_call::<TestProtocol, _, _>(
"input",
None,
move |_, _, _, _observations| async move {
let chunks = stream::unfold((0, release), move |(index, release)| {
provider_polls.fetch_add(1, Ordering::SeqCst);
async move {
(index < 2).then(|| (Ok(Bytes::from(index.to_string())), (index + 1, release)))
}
})
.boxed();
Ok(CallOutput::Stream {
head: "text/event-stream",
chunks,
})
},
);
let response = serve(
machine,
Adapter(Rejection::None),
hooks,
(),
Adapter(Rejection::None),
Some(observer.0.sender.clone()),
)
.await
.unwrap();
@ -299,32 +329,42 @@ async fn body_demand_controls_polling_and_lifecycle(
#[case::encoding(Rejection::Chunk, TestError::Adapter, 1)]
#[tokio::test]
async fn stream_failure_emits_one_error_frame_and_stops(
hooks: Hooks,
interceptors: Hooks,
#[case] rejection: Rejection,
#[case] expected: TestError,
#[case] expected_polls: usize,
) {
let observer = hooks.observer.clone();
let observer = interceptors.observer.clone();
let polls = Arc::new(AtomicUsize::new(0));
let provider_polls = polls.clone();
let machine = hosted_call::<TestProtocol, _, _>("input", move |_, _, _| async move {
let chunks = stream::iter([
Ok(Bytes::from_static(b"first")),
Err(TestError::Provider),
Ok(Bytes::from_static(b"must not be delivered")),
])
.inspect(move |_| {
provider_polls.fetch_add(1, Ordering::SeqCst);
})
.boxed();
Ok(CallOutput::Stream {
head: "text/event-stream",
chunks,
})
});
let response = serve(machine, Adapter(rejection), hooks, Adapter(rejection))
.await
.unwrap();
let machine = hosted_call::<TestProtocol, _, _>(
"input",
None,
move |_, _, _, _observations| async move {
let chunks = stream::iter([
Ok(Bytes::from_static(b"first")),
Err(TestError::Provider),
Ok(Bytes::from_static(b"must not be delivered")),
])
.inspect(move |_| {
provider_polls.fetch_add(1, Ordering::SeqCst);
})
.boxed();
Ok(CallOutput::Stream {
head: "text/event-stream",
chunks,
})
},
);
let response = serve(
machine,
Adapter(rejection),
interceptors,
Adapter(rejection),
Some(observer.0.sender.clone()),
)
.await
.unwrap();
let body = to_bytes(response.into_body(), 1024).await.unwrap();
let prefix = if rejection == Rejection::Chunk {
""
@ -351,28 +391,38 @@ async fn stream_failure_emits_one_error_frame_and_stops(
#[case::custom_operation(Rejection::Custom)]
#[case::response_conversion(Rejection::Complete)]
#[tokio::test]
async fn failures_before_open_return_an_error(hooks: Hooks, #[case] rejection: Rejection) {
let observer = hooks.observer.clone();
let machine = hosted_call::<TestProtocol, _, _>("input", move |_, services, _| async move {
services.call(|reply| reply).await?;
match rejection {
Rejection::Head => Ok(CallOutput::Stream {
head: "text/event-stream",
chunks: stream::pending().boxed(),
}),
Rejection::None => Err(TestError::Provider),
_ => Ok(CallOutput::Complete(Bytes::new())),
}
});
async fn failures_before_open_return_an_error(interceptors: Hooks, #[case] rejection: Rejection) {
let observer = interceptors.observer.clone();
let machine = hosted_call::<TestProtocol, _, _>(
"input",
None,
move |_, services, _, _observations| async move {
services.call(|reply| reply).await?;
match rejection {
Rejection::Head => Ok(CallOutput::Stream {
head: "text/event-stream",
chunks: stream::pending().boxed(),
}),
Rejection::None => Err(TestError::Provider),
_ => Ok(CallOutput::Complete(Bytes::new())),
}
},
);
let expected = if rejection == Rejection::None {
TestError::Provider
} else {
TestError::Adapter
};
assert_eq!(
serve(machine, Adapter(rejection), hooks, Adapter(rejection))
.await
.unwrap_err(),
serve(
machine,
Adapter(rejection),
interceptors,
Adapter(rejection),
Some(observer.0.sender.clone())
)
.await
.unwrap_err(),
Error::Call(expected)
);
assert!(matches!(
@ -385,30 +435,38 @@ async fn failures_before_open_return_an_error(hooks: Hooks, #[case] rejection: R
#[case::before_headers(false)]
#[case::awaiting_chunk(true)]
#[tokio::test]
async fn cancelling_pending_work_releases_the_machine(hooks: Hooks, #[case] streaming: bool) {
let observer = hooks.observer.clone();
async fn cancelling_pending_work_releases_the_machine(
interceptors: Hooks,
#[case] streaming: bool,
) {
let observer = interceptors.observer.clone();
let released = Arc::new(AtomicBool::new(false));
let release = Release(released.clone());
let machine = hosted_call::<TestProtocol, _, _>("input", move |_, _, _| async move {
if !streaming {
let _release = release;
return std::future::pending().await;
}
let chunks = stream::once(async move {
let _release = release;
std::future::pending().await
})
.boxed();
Ok(CallOutput::Stream {
head: "text/event-stream",
chunks,
})
});
let machine = hosted_call::<TestProtocol, _, _>(
"input",
None,
move |_, _, _, _observations| async move {
if !streaming {
let _release = release;
return std::future::pending().await;
}
let chunks = stream::once(async move {
let _release = release;
std::future::pending().await
})
.boxed();
Ok(CallOutput::Stream {
head: "text/event-stream",
chunks,
})
},
);
let mut response = Box::pin(serve(
machine,
Adapter(Rejection::None),
hooks,
interceptors,
Adapter(Rejection::None),
Some(observer.0.sender.clone()),
));
if streaming {
let mut body = response.await.unwrap().into_body().into_data_stream();
@ -437,37 +495,39 @@ async fn hook_rejection_stops_execution_and_is_reported_once(
) {
let continued = Arc::new(AtomicBool::new(false));
let executed = continued.clone();
let machine = hosted_call::<TestProtocol, _, _>("input", move |_, _, route_hooks| async move {
if event {
route_hooks
.on_event(MachineEvent::ResponseReceived {
raw: RawResponse {
let machine = hosted_call::<TestProtocol, _, _>(
"input",
None,
move |_, _, route_hooks, _observations| async move {
if event {
route_hooks
.after_provider_response(RawResponse {
body: "response".into(),
},
})
.await?;
} else {
route_hooks
.before_provider_request(
WireRequest {
url: "url".into(),
headers: Vec::new(),
body: json!({}),
},
RequestContext {
model: "model".into(),
custom_llm_provider: "provider".into(),
optional_params: json!({}),
secret_fields: Vec::new(),
api_key: None,
},
)
.await?;
}
executed.store(true, Ordering::SeqCst);
Ok(CallOutput::Complete(Bytes::new()))
});
let hooks = Hooks {
})
.await?;
} else {
route_hooks
.before_provider_request(
WireRequest {
url: "url".into(),
headers: Vec::new(),
body: json!({}),
},
RequestContext {
model: "model".into(),
custom_llm_provider: "provider".into(),
optional_params: json!({}),
secret_fields: Vec::new(),
api_key: None,
},
)
.await?;
}
executed.store(true, Ordering::SeqCst);
Ok(CallOutput::Complete(Bytes::new()))
},
);
let interceptors = Hooks {
observer: observer.clone(),
reject: true,
};
@ -475,8 +535,9 @@ async fn hook_rejection_stops_execution_and_is_reported_once(
serve(
machine,
Adapter(Rejection::None),
hooks,
Adapter(Rejection::None)
interceptors,
Adapter(Rejection::None),
Some(observer.0.sender.clone()),
)
.await
.unwrap_err(),
@ -499,19 +560,22 @@ enum InvalidFlow {
#[case::deliver_before_open(InvalidFlow::DeliverBeforeOpen)]
#[case::open_twice(InvalidFlow::OpenTwice)]
#[tokio::test]
async fn invalid_host_operations_fail_without_panicking(hooks: Hooks, #[case] flow: InvalidFlow) {
async fn invalid_host_operations_fail_without_panicking(
interceptors: Hooks,
#[case] flow: InvalidFlow,
) {
use litellm_host::{call::HostedCompletion, machine::CallMachine};
let observer = hooks.observer.clone();
let machine = CallMachine::<TestProtocol, HostedCompletion<Bytes>>::new(move |host| {
let observer = interceptors.observer.clone();
let machine = CallMachine::<TestProtocol, HostedCompletion<Bytes>>::new(None, move |host| {
Box::pin(async move {
match flow {
InvalidFlow::DeliverBeforeOpen => {
host.stream.send_chunk(Bytes::new()).await?;
let _ = host.stream.send_chunk(Bytes::new()).await?;
}
InvalidFlow::OpenTwice => {
host.stream.open_stream("text/event-stream").await?;
host.stream.open_stream("text/event-stream").await?;
let _ = host.stream.open_stream("text/event-stream").await?;
let _ = host.stream.open_stream("text/event-stream").await?;
}
}
Ok(HostedCompletion::StreamEnded)
@ -520,8 +584,9 @@ async fn invalid_host_operations_fail_without_panicking(hooks: Hooks, #[case] fl
let result = serve(
machine,
Adapter(Rejection::None),
hooks,
interceptors,
Adapter(Rejection::None),
Some(observer.0.sender.clone()),
)
.await;
if matches!(flow, InvalidFlow::OpenTwice) {
@ -549,10 +614,12 @@ impl Protocol for UnaryProtocol {
#[rstest]
#[tokio::test]
async fn unary_calls_use_into_response_after_hooks_and_before_success(hooks: Hooks) {
let observer = hooks.observer.clone();
let machine =
hosted_call::<UnaryProtocol, _, _>("projected", |request, _, route_hooks| async move {
async fn unary_calls_use_into_response_after_hooks_and_before_success(interceptors: Hooks) {
let observer = interceptors.observer.clone();
let machine = hosted_call::<UnaryProtocol, _, _>(
"projected",
None,
|request, _, route_hooks, _observations| async move {
let wire = route_hooks
.before_provider_request(
WireRequest {
@ -570,25 +637,25 @@ async fn unary_calls_use_into_response_after_hooks_and_before_success(hooks: Hoo
)
.await?;
route_hooks
.on_event(MachineEvent::ResponseReceived {
raw: RawResponse {
body: wire.url.clone(),
},
.after_provider_response(RawResponse {
body: wire.url.clone(),
})
.await?;
Ok(CallOutput::Complete(json!({"url": wire.url})))
});
},
);
let response = serve_unary(
machine,
(),
hooks,
interceptors,
Unary::new(|value| {
assert!(matches!(
observer.0.lock().unwrap().as_slice(),
[CallEvent::Started { .. }, CallEvent::Machine(_),]
[CallEvent::Started { .. }, CallEvent::Execution(_),]
));
(StatusCode::CREATED, [("x-converted", "yes")], Json(value))
}),
Some(observer.0.sender.clone()),
)
.await
.unwrap();
@ -602,7 +669,7 @@ async fn unary_calls_use_into_response_after_hooks_and_before_success(hooks: Hoo
observer.0.lock().unwrap().as_slice(),
[
CallEvent::Started { .. },
CallEvent::Machine(_),
CallEvent::Execution(_),
CallEvent::Succeeded { .. },
]
));
@ -617,15 +684,17 @@ async fn unary_failure_preserves_the_error_without_converting(
#[case] reject_hook: bool,
#[case] expected: TestError,
) {
let machine = hosted_call::<UnaryProtocol, _, _>("input", |_, _, route_hooks| async move {
route_hooks
.on_event(MachineEvent::ResponseReceived {
raw: RawResponse { body: "raw".into() },
})
.await?;
Err(TestError::Provider)
});
let hooks = Hooks {
let machine = hosted_call::<UnaryProtocol, _, _>(
"input",
None,
|_, _, route_hooks, _observations| async move {
route_hooks
.after_provider_response(RawResponse { body: "raw".into() })
.await?;
Err(TestError::Provider)
},
);
let interceptors = Hooks {
observer: observer.clone(),
reject: reject_hook,
};
@ -633,11 +702,12 @@ async fn unary_failure_preserves_the_error_without_converting(
let result = serve_unary(
machine,
(),
hooks,
interceptors,
Unary::new(|value| {
converted.store(true, Ordering::SeqCst);
Json(value)
}),
Some(observer.0.sender.clone()),
)
.await;
assert_eq!(result.unwrap_err(), Error::Call(expected));
@ -650,23 +720,28 @@ async fn unary_failure_preserves_the_error_without_converting(
#[rstest]
#[tokio::test]
async fn cancelling_unary_execution_releases_work_without_converting(hooks: Hooks) {
let observer = hooks.observer.clone();
async fn cancelling_unary_execution_releases_work_without_converting(interceptors: Hooks) {
let observer = interceptors.observer.clone();
let released = Arc::new(AtomicBool::new(false));
let release = Release(released.clone());
let machine = hosted_call::<UnaryProtocol, _, _>("input", move |_, _, _| async move {
let _release = release;
std::future::pending().await
});
let machine = hosted_call::<UnaryProtocol, _, _>(
"input",
None,
move |_, _, _, _observations| async move {
let _release = release;
std::future::pending().await
},
);
let converted = AtomicBool::new(false);
let mut call = Box::pin(serve_unary(
machine,
(),
hooks,
interceptors,
Unary::new(|value| {
converted.store(true, Ordering::SeqCst);
Json(value)
}),
Some(observer.0.sender.clone()),
));
assert!(futures_util::poll!(&mut call).is_pending());
assert!(!released.load(Ordering::SeqCst));
@ -709,7 +784,7 @@ impl ResponseEncoder for CustomUnaryAdapter {
struct Credentials(bool);
impl litellm_host::services::HostCallHandler<CustomUnaryProtocol> for Credentials {
impl litellm_host_native::services::HostCallHandler<CustomUnaryProtocol> for Credentials {
async fn handle_host_call(&self, reply: Reply<&'static str>) -> Result<(), TestError> {
if self.0 {
return Err(TestError::Adapter);
@ -725,11 +800,10 @@ impl litellm_host::services::HostCallHandler<CustomUnaryProtocol> for Credential
#[case::response_conversion_fails(false, true)]
#[tokio::test]
async fn unary_custom_operations_and_conversion_finish_before_terminal_observation(
hooks: Hooks,
observer: Arc<Observer>,
#[case] reject_op: bool,
#[case] reject_response: bool,
) {
let observer = hooks.observer.clone();
let continued = Arc::new(AtomicBool::new(false));
let executed = continued.clone();
let released = Arc::new(AtomicBool::new(false));
@ -737,7 +811,8 @@ async fn unary_custom_operations_and_conversion_finish_before_terminal_observati
let converted = Arc::new(AtomicBool::new(false));
let machine = hosted_call::<CustomUnaryProtocol, _, _>(
"request",
move |request, services, _| async move {
None,
move |request, services, _, _observations| async move {
let _release = release;
let credential = services.call(|reply| reply).await?;
executed.store(true, Ordering::SeqCst);
@ -749,11 +824,12 @@ async fn unary_custom_operations_and_conversion_finish_before_terminal_observati
let result = serve_unary(
machine,
Credentials(reject_op),
hooks,
(),
CustomUnaryAdapter {
reject_response,
converted: converted.clone(),
},
Some(observer.0.sender.clone()),
)
.await;
assert_eq!(continued.load(Ordering::SeqCst), !reject_op);

View file

@ -46,25 +46,26 @@ impl Protocol for TestProtocol {
async fn sse_preserves_encoded_chunks_and_uses_the_supplied_error_format(#[case] fail: bool) {
let first = Bytes::from_static(b"event: custom\ndata: first\n\n");
let last = Bytes::from_static(b"data: [DONE]\n\n");
let machine = hosted_call::<TestProtocol, _, _>((), move |(), _, _| async move {
let chunks = stream::iter([
Ok(first),
if fail {
Err(TestError::Upstream)
} else {
Ok(last)
},
])
.boxed();
Ok(CallOutput::Stream { head: (), chunks })
});
let machine =
hosted_call::<TestProtocol, _, _>((), None, move |(), _, _, _observations| async move {
let chunks = stream::iter([
Ok(first),
if fail {
Err(TestError::Upstream)
} else {
Ok(last)
},
])
.boxed();
Ok(CallOutput::Stream { head: (), chunks })
});
let errors = Arc::new(AtomicUsize::new(0));
let formatted_errors = errors.clone();
let adapter = Sse::new(std::convert::identity, move |error| {
formatted_errors.fetch_add(1, Ordering::SeqCst);
Bytes::from(format!("event: custom_error\ndata: {error:?}\n\n"))
});
let response = serve(machine, (), (), adapter).await.unwrap();
let response = serve(machine, (), (), adapter, None).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers()[CONTENT_TYPE], "text/event-stream");
assert_eq!(errors.load(Ordering::SeqCst), 0);
@ -85,14 +86,14 @@ async fn sse_preserves_encoded_chunks_and_uses_the_supplied_error_format(#[case]
async fn completed_calls_use_the_response_converter_without_sse_headers() {
use axum::response::IntoResponse;
let machine = hosted_call::<TestProtocol, _, _>((), |(), _, _| async {
let machine = hosted_call::<TestProtocol, _, _>((), None, |(), _, _, _observations| async {
Ok(CallOutput::Complete(Bytes::from_static(b"completed")))
});
let adapter = Sse::new(
|response| (StatusCode::CREATED, [("x-converted", "yes")], response).into_response(),
|_| panic!("a completed call cannot format a stream error"),
);
let response = serve(machine, (), (), adapter).await.unwrap();
let response = serve(machine, (), (), adapter, None).await.unwrap();
assert_eq!(response.status(), StatusCode::CREATED);
assert_eq!(response.headers()["x-converted"], "yes");
assert_ne!(

View file

@ -0,0 +1,7 @@
`litellm-host-native` is the Rust driver for hosted calls. `Driver` owns the machine, a `HostCallHandler` and a `Interceptors`; `advance()` answers services and hooks inline and returns at completion or at the next stream boundary, holding the `Reply<ControlFlow<()>>` until the consumer calls `advance()` or `detach()` again. Dropping the driver drops the machine and so cancels the call
The consumer decides demand, so the driver never spawns a producer task and never buffers chunks ahead of demand. `litellm-host-http` polls it from the response body; `in_process::run_hosted` polls it on behalf of a `StreamConsumer`. Both observe lifecycle terminals themselves, the driver reports none
Depend on `litellm-host` only. HTTP encoding stays in `litellm-host-http`; `litellm-host-python` drives the machine directly so Python callbacks stay in the caller's asyncio task
`services.rs` owns `HostCallHandler` and its borrowed and no-service implementations. This is the Rust driver's handler contract; the shared host crate owns the service request protocol

View file

@ -0,0 +1,15 @@
[package]
name = "litellm-host-native"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-host.workspace = true
[dev-dependencies]
futures-util.workspace = true
rstest.workspace = true
serde_json.workspace = true
tokio.workspace = true

View file

@ -0,0 +1,99 @@
use std::ops::ControlFlow;
use litellm_host::{
interceptors::Interceptors,
machine::{HostFailure, Machine, MachineStep},
protocol::{HostRequest, InterceptRequest, Protocol, Reply, StreamDelivery},
};
use crate::services::HostCallHandler;
type ProtocolOf<M> = <M as Machine>::Protocol;
type ErrorOf<M> = <ProtocolOf<M> as Protocol>::Error;
pub enum Boundary<M: Machine> {
Complete(M::Complete),
Open(<ProtocolOf<M> as Protocol>::StreamHead),
Chunk(<ProtocolOf<M> as Protocol>::Chunk),
}
/// Answers host calls and interceptors inline and stops at each stream delivery, holding its demand
/// reply until the consumer advances again. Dropping the driver drops the in-flight call.
pub struct Driver<M: Machine, S, H> {
machine: M,
services: S,
interceptors: H,
demand: Option<Reply<ControlFlow<()>>>,
}
impl<M, S, H> Driver<M, S, H>
where
M: Machine,
S: HostCallHandler<ProtocolOf<M>>,
H: Interceptors<ErrorOf<M>>,
{
pub fn new(machine: M, services: S, interceptors: H) -> Self {
Self {
machine,
services,
interceptors,
demand: None,
}
}
pub async fn advance(&mut self) -> Result<Boundary<M>, ErrorOf<M>> {
self.resume(ControlFlow::Continue(())).await
}
pub async fn detach(&mut self) -> Result<Boundary<M>, ErrorOf<M>> {
self.resume(ControlFlow::Break(())).await
}
/// Interrupts the machine with a failure the consumer hit at the last stream boundary,
/// dropping the held demand reply unanswered
pub async fn fail(&mut self, error: ErrorOf<M>) -> Result<M::Complete, ErrorOf<M>> {
self.demand = None;
self.machine.interrupt(HostFailure::Error(error)).await
}
async fn resume(&mut self, demand: ControlFlow<()>) -> Result<Boundary<M>, ErrorOf<M>> {
if let Some(reply) = self.demand.take() {
reply.send(demand);
}
loop {
let request = match self.machine.resume().await? {
MachineStep::Complete(complete) => return Ok(Boundary::Complete(complete)),
MachineStep::Suspended(request) => request,
};
let answered = match request {
HostRequest::HostCall(call) => self.services.handle_host_call(call).await,
HostRequest::Intercept(InterceptRequest::BeforeProviderRequest {
wire,
context,
reply,
}) => self
.interceptors
.before_provider_request(*wire, *context)
.await
.map(|wire| reply.send(wire)),
HostRequest::Intercept(InterceptRequest::AfterProviderResponse { raw, reply }) => {
self.interceptors
.after_provider_response(raw)
.await
.map(|()| reply.send(()))
}
HostRequest::Stream(StreamDelivery::Open(head, reply)) => {
self.demand = Some(reply);
return Ok(Boundary::Open(head));
}
HostRequest::Stream(StreamDelivery::Chunk(chunk, reply)) => {
self.demand = Some(reply);
return Ok(Boundary::Chunk(chunk));
}
};
if let Err(error) = answered {
return self.fail(error).await.map(Boundary::Complete);
}
}
}
}

View file

@ -0,0 +1,142 @@
use litellm_host::observation::ObservationSender;
use std::{future::Future, ops::ControlFlow};
use litellm_host::{
call::{HostedCompletion, HostedMachine},
interceptors::Interceptors,
lifecycle::{CallEvent, FailureOrigin, Timing, epoch_seconds},
machine::{Machine, MachineFault},
protocol::Protocol,
};
use crate::{
driver::{Boundary, Driver},
services::HostCallHandler,
};
pub trait StreamConsumer<P: Protocol>: Send + Sync {
fn open_stream(
&self,
head: P::StreamHead,
) -> impl Future<Output = Result<ControlFlow<()>, P::Error>> + Send;
fn send_chunk(
&self,
chunk: P::Chunk,
) -> impl Future<Output = Result<ControlFlow<()>, P::Error>> + Send;
}
impl<P: Protocol> StreamConsumer<P> for () {
async fn open_stream(&self, _: P::StreamHead) -> Result<ControlFlow<()>, P::Error> {
Ok(ControlFlow::Continue(()))
}
async fn send_chunk(&self, _: P::Chunk) -> Result<ControlFlow<()>, P::Error> {
Ok(ControlFlow::Continue(()))
}
}
pub struct Host<'a, S, H, C> {
pub services: &'a S,
pub interceptors: &'a H,
pub stream: &'a C,
pub observers: Option<&'a ObservationSender>,
}
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: Interceptors<<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: HostedMachine<P>,
host: Host<'_, S, H, C>,
) -> Result<HostedCompletion<P::Response>, P::Error>
where
P: Protocol,
P::Error: From<MachineFault>,
S: HostCallHandler<P>,
H: Interceptors<P::Error>,
C: StreamConsumer<P>,
{
run_with_completion(machine, host, |completion| {
matches!(completion, HostedCompletion::Detached)
})
.await
}
async fn run_with_completion<M, S, H, C>(
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: Interceptors<<M::Protocol as Protocol>::Error>,
C: StreamConsumer<M::Protocol>,
{
let start_time = epoch_seconds();
if let Some(observers) = host.observers {
observers.emit(CallEvent::Started { start_time });
}
let outcome = consume(
Driver::new(machine, host.services, host.interceptors),
host.stream,
)
.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,
response: (),
},
Err(_) => CallEvent::Failed {
timing,
origin: FailureOrigin::Call,
error: (),
},
};
if let Some(observers) = host.observers {
observers.emit(terminal);
}
outcome
}
async fn consume<M, S, H, C>(
mut driver: Driver<M, &S, &H>,
stream: &C,
) -> Result<M::Complete, <M::Protocol as Protocol>::Error>
where
M: Machine,
S: HostCallHandler<M::Protocol>,
H: Interceptors<<M::Protocol as Protocol>::Error>,
C: StreamConsumer<M::Protocol>,
{
let mut demand = ControlFlow::Continue(());
loop {
let boundary = match demand {
ControlFlow::Continue(()) => driver.advance().await?,
ControlFlow::Break(()) => driver.detach().await?,
};
let delivered = match boundary {
Boundary::Complete(complete) => return Ok(complete),
Boundary::Open(head) => stream.open_stream(head).await,
Boundary::Chunk(chunk) => stream.send_chunk(chunk).await,
};
demand = match delivered {
Ok(demand) => demand,
Err(error) => return driver.fail(error).await,
};
}
}

View file

@ -0,0 +1,9 @@
//! The Rust driver for hosted calls: it answers host services and interceptors with Rust handlers
//! and hands stream deliveries to whichever consumer sits on top, HTTP body polling or an
//! in-process stream consumer.
mod driver;
pub mod in_process;
pub mod services;
pub use driver::{Boundary, Driver};

View file

@ -1,6 +1,7 @@
use crate::protocol::Protocol;
use std::{convert::Infallible, future::Future};
use litellm_host::protocol::Protocol;
pub trait HostCallHandler<P: Protocol>: Send + Sync {
fn handle_host_call(
&self,
@ -8,6 +9,15 @@ pub trait HostCallHandler<P: Protocol>: Send + Sync {
) -> impl Future<Output = Result<(), P::Error>> + Send;
}
impl<P: Protocol, T: HostCallHandler<P> + ?Sized> HostCallHandler<P> for &T {
fn handle_host_call(
&self,
call: P::HostCall,
) -> impl Future<Output = Result<(), P::Error>> + Send {
(**self).handle_host_call(call)
}
}
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,695 @@
use litellm_host::lifecycle::ExecutionEvent;
use std::{
ops::ControlFlow,
sync::{
Arc, Mutex,
atomic::{AtomicBool, AtomicUsize, Ordering},
},
};
use futures_util::{StreamExt, stream};
use litellm_host::{
call::{CallOutput, HostedCompletion, hosted_call},
interceptors::{Interceptors, RawResponse, RequestContext, WireRequest},
lifecycle::{CallEvent, CallObserver},
machine::{CallMachine, HostFailure, Interrupted, Machine, MachineFault, Step},
protocol::{Protocol, Reply},
};
use litellm_host_native::{
Boundary, Driver,
in_process::{Host, StreamConsumer, run, run_hosted},
services::HostCallHandler,
};
use rstest::{fixture, rstest};
use serde_json::json;
#[derive(Clone, Debug, PartialEq)]
enum TestError {
Provider,
Service,
Hook,
Consumer,
Machine,
}
impl From<MachineFault> for TestError {
fn from(_: MachineFault) -> Self {
Self::Machine
}
}
struct TestProtocol;
impl Protocol for TestProtocol {
type Request = &'static str;
type Response = String;
type Error = TestError;
type HostCall = Reply<&'static str>;
type Chunk = usize;
type StreamHead = &'static str;
}
type TestMachine = litellm_host::call::HostedMachine<TestProtocol>;
struct Services {
reject: bool,
}
impl HostCallHandler<TestProtocol> for Services {
async fn handle_host_call(&self, reply: Reply<&'static str>) -> Result<(), TestError> {
if self.reject {
return Err(TestError::Service);
}
reply.send("custom");
Ok(())
}
}
#[derive(Default)]
struct Observer(Observations);
struct Observations {
sender: litellm_host::observation::ObservationSender,
receiver: Mutex<tokio::sync::mpsc::Receiver<CallEvent>>,
recorded: Mutex<Vec<CallEvent>>,
}
impl Default for Observations {
fn default() -> Self {
let (sender, receiver) = litellm_host::observation::observation_channel(
std::num::NonZeroUsize::new(128).unwrap(),
);
Self {
sender,
receiver: Mutex::new(receiver),
recorded: Mutex::new(Vec::new()),
}
}
}
impl Observations {
fn lock(&self) -> std::sync::LockResult<std::sync::MutexGuard<'_, Vec<CallEvent>>> {
let mut events = self.recorded.lock()?;
let mut receiver = self.receiver.lock().unwrap();
while let Ok(event) = receiver.try_recv() {
events.push(event);
}
Ok(events)
}
}
impl CallObserver for Observer {
fn observe(&self, event: CallEvent) {
self.0.sender.emit(event);
}
}
struct Hooks {
observer: Arc<Observer>,
reject: bool,
}
impl Interceptors<TestError> for Hooks {
async fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, TestError> {
if self.reject {
return Err(TestError::Hook);
}
Ok(WireRequest {
url: format!("{}/{}", wire.url, context.model),
..wire
})
}
async fn after_provider_response(&self, raw: RawResponse) -> Result<(), TestError> {
if self.reject {
return Err(TestError::Hook);
}
self.observer.observe(CallEvent::Execution(
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
));
Ok(())
}
}
#[fixture]
fn observer() -> Arc<Observer> {
Arc::new(Observer::default())
}
#[fixture]
fn interceptors(observer: Arc<Observer>) -> Hooks {
Hooks {
observer,
reject: false,
}
}
fn dispatching_call() -> TestMachine {
hosted_call::<TestProtocol, _, _>(
"projected",
None,
|request, services, route_hooks, _observations| async move {
let custom = services.call(|reply| reply).await?;
let wire = route_hooks
.before_provider_request(
WireRequest {
url: request.into(),
headers: Vec::new(),
body: json!({}),
},
RequestContext {
model: custom.into(),
custom_llm_provider: "test".into(),
optional_params: json!({}),
secret_fields: Vec::new(),
api_key: None,
},
)
.await?;
route_hooks
.after_provider_response(RawResponse {
body: wire.url.clone(),
})
.await?;
Ok(CallOutput::Complete(wire.url))
},
)
}
struct Release(Arc<AtomicBool>);
impl Drop for Release {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
struct Streaming {
polls: Arc<AtomicUsize>,
released: Arc<AtomicBool>,
machine: TestMachine,
}
fn streaming_call(chunks: Vec<Result<usize, TestError>>) -> Streaming {
let polls = Arc::new(AtomicUsize::new(0));
let released = Arc::new(AtomicBool::new(false));
let provider_polls = polls.clone();
let release = Release(released.clone());
let machine = hosted_call::<TestProtocol, _, _>(
"input",
None,
move |_, _, _, _observations| async move {
let chunks = stream::iter(chunks)
.inspect(move |_| {
let _held = &release;
provider_polls.fetch_add(1, Ordering::SeqCst);
})
.boxed();
Ok(CallOutput::Stream {
head: "headers",
chunks,
})
},
);
Streaming {
polls,
released,
machine,
}
}
#[rstest]
#[tokio::test]
async fn services_and_hooks_answer_the_machine_inline(interceptors: Hooks) {
let observer = interceptors.observer.clone();
let mut driver = Driver::new(dispatching_call(), Services { reject: false }, interceptors);
let Boundary::Complete(HostedCompletion::Complete(response)) = driver.advance().await.unwrap()
else {
panic!("a unary call completes at the first boundary")
};
assert_eq!(response, "projected/custom");
assert!(matches!(
observer.0.lock().unwrap().as_slice(),
[CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw })] if raw.body == "projected/custom"
));
}
#[rstest]
#[case::service(true, false, TestError::Service)]
#[case::hook(false, true, TestError::Hook)]
#[tokio::test]
async fn handler_failures_interrupt_the_machine(
observer: Arc<Observer>,
#[case] reject_service: bool,
#[case] reject_hook: bool,
#[case] expected: TestError,
) {
let mut driver = Driver::new(
dispatching_call(),
Services {
reject: reject_service,
},
Hooks {
observer: observer.clone(),
reject: reject_hook,
},
);
assert!(matches!(driver.advance().await, Err(error) if error == expected));
assert!(observer.0.lock().unwrap().is_empty());
}
#[rstest]
#[tokio::test]
async fn advancing_delivers_one_chunk_per_demand() {
let Streaming { polls, machine, .. } = streaming_call(vec![Ok(0), Ok(1)]);
let mut driver = Driver::new(machine, Services { reject: false }, ());
assert!(matches!(
driver.advance().await,
Ok(Boundary::Open("headers"))
));
assert_eq!(polls.load(Ordering::SeqCst), 0);
for index in 0..2 {
assert!(matches!(driver.advance().await, Ok(Boundary::Chunk(chunk)) if chunk == index));
assert_eq!(polls.load(Ordering::SeqCst), index + 1);
}
assert!(matches!(
driver.advance().await,
Ok(Boundary::Complete(HostedCompletion::StreamEnded))
));
assert_eq!(polls.load(Ordering::SeqCst), 2);
}
#[rstest]
#[tokio::test]
async fn provider_stream_errors_surface_at_the_failing_chunk() {
let Streaming { polls, machine, .. } =
streaming_call(vec![Ok(0), Err(TestError::Provider), Ok(2)]);
let mut driver = Driver::new(machine, Services { reject: false }, ());
assert!(matches!(driver.advance().await, Ok(Boundary::Open(_))));
assert!(matches!(driver.advance().await, Ok(Boundary::Chunk(0))));
assert_eq!(driver.advance().await.err(), Some(TestError::Provider));
assert_eq!(polls.load(Ordering::SeqCst), 2);
}
#[rstest]
#[tokio::test]
async fn detaching_completes_without_pulling_more_chunks() {
let Streaming { polls, machine, .. } = streaming_call(vec![Ok(0), Ok(1)]);
let mut driver = Driver::new(machine, Services { reject: false }, ());
assert!(matches!(driver.advance().await, Ok(Boundary::Open(_))));
assert!(matches!(driver.advance().await, Ok(Boundary::Chunk(0))));
assert!(matches!(
driver.detach().await,
Ok(Boundary::Complete(HostedCompletion::Detached))
));
assert_eq!(polls.load(Ordering::SeqCst), 1);
}
#[rstest]
#[case::at_open(0)]
#[case::after_chunk(1)]
#[tokio::test]
async fn dropping_the_driver_drops_the_call(#[case] chunks_before_drop: usize) {
let Streaming {
polls,
released,
machine,
} = streaming_call(vec![Ok(0), Ok(1)]);
let mut driver = Driver::new(machine, Services { reject: false }, ());
assert!(matches!(driver.advance().await, Ok(Boundary::Open(_))));
for _ in 0..chunks_before_drop {
assert!(matches!(driver.advance().await, Ok(Boundary::Chunk(_))));
}
assert!(!released.load(Ordering::SeqCst));
drop(driver);
assert!(released.load(Ordering::SeqCst));
assert_eq!(polls.load(Ordering::SeqCst), chunks_before_drop);
}
struct Interruptible {
inner: TestMachine,
interrupted: Arc<Mutex<Vec<HostFailure<TestError>>>>,
}
impl Machine for Interruptible {
type Protocol = TestProtocol;
type Complete = HostedCompletion<String>;
fn resume(&mut self) -> Step<'_, Self> {
self.inner.resume()
}
fn interrupt(&mut self, failure: HostFailure<TestError>) -> Interrupted<'_, Self> {
self.interrupted.lock().unwrap().push(failure.clone());
self.inner.interrupt(failure)
}
}
struct Consumer {
detach_after: Option<usize>,
fail_after: Option<usize>,
delivered: Mutex<Vec<usize>>,
}
impl Consumer {
fn demand_after(&self, delivered: usize) -> Result<ControlFlow<()>, TestError> {
if self.fail_after == Some(delivered) {
return Err(TestError::Consumer);
}
Ok(if self.detach_after == Some(delivered) {
ControlFlow::Break(())
} else {
ControlFlow::Continue(())
})
}
}
impl StreamConsumer<TestProtocol> for Consumer {
async fn open_stream(&self, head: &'static str) -> Result<ControlFlow<()>, TestError> {
assert_eq!(head, "headers");
self.demand_after(0)
}
async fn send_chunk(&self, chunk: usize) -> Result<ControlFlow<()>, TestError> {
let mut delivered = self.delivered.lock().unwrap();
delivered.push(chunk);
self.demand_after(delivered.len())
}
}
#[rstest]
#[case::consumed(None, None, Ok(HostedCompletion::StreamEnded), 3, 3)]
#[case::detach_at_open(Some(0), None, Ok(HostedCompletion::Detached), 0, 0)]
#[case::detach_after_chunk(Some(1), None, Ok(HostedCompletion::Detached), 1, 1)]
#[case::consumer_fails(None, Some(1), Err(TestError::Consumer), 1, 1)]
#[tokio::test]
async fn in_process_runner_follows_consumer_demand(
observer: Arc<Observer>,
#[case] detach_after: Option<usize>,
#[case] fail_after: Option<usize>,
#[case] expected: Result<HostedCompletion<String>, TestError>,
#[case] expected_polls: usize,
#[case] expected_delivered: usize,
) {
let Streaming {
polls,
released,
machine,
} = streaming_call(vec![Ok(0), Ok(1), Ok(2)]);
let consumer = Consumer {
detach_after,
fail_after,
delivered: Mutex::new(Vec::new()),
};
let outcome = run_hosted(
machine,
Host {
services: &Services { reject: false },
interceptors: &(),
stream: &consumer,
observers: Some(&observer.0.sender),
},
)
.await;
assert_eq!(outcome, expected);
assert_eq!(polls.load(Ordering::SeqCst), expected_polls);
assert_eq!(
*consumer.delivered.lock().unwrap(),
(0..expected_delivered).collect::<Vec<_>>()
);
assert!(released.load(Ordering::SeqCst));
let events = observer.0.lock().unwrap();
assert_eq!(events.len(), 2);
assert!(matches!(events[0], CallEvent::Started { .. }));
match &expected {
Ok(HostedCompletion::Detached) => {
assert!(matches!(events[1], CallEvent::Cancelled { .. }))
}
Ok(_) => assert!(matches!(events[1], CallEvent::Succeeded { .. })),
Err(_) => assert!(matches!(events[1], CallEvent::Failed { .. })),
}
}
#[rstest]
#[case::at_open(0)]
#[case::after_chunk(1)]
#[tokio::test]
async fn consumer_failures_interrupt_the_machine(#[case] fail_after: usize) {
let Streaming { polls, machine, .. } = streaming_call(vec![Ok(0), Ok(1), Ok(2)]);
let interrupted = Arc::new(Mutex::new(Vec::new()));
let consumer = Consumer {
detach_after: None,
fail_after: Some(fail_after),
delivered: Mutex::new(Vec::new()),
};
let outcome = run(
Interruptible {
inner: machine,
interrupted: interrupted.clone(),
},
Host {
services: &Services { reject: false },
interceptors: &(),
stream: &consumer,
observers: None,
},
)
.await;
assert_eq!(outcome, Err(TestError::Consumer));
assert_eq!(
*interrupted.lock().unwrap(),
[HostFailure::Error(TestError::Consumer)]
);
assert_eq!(polls.load(Ordering::SeqCst), fail_after);
}
struct Recording {
ops: &'static [&'static str],
calls: AtomicUsize,
seen: Mutex<Vec<String>>,
fail: Option<&'static str>,
events: Observations,
}
impl Recording {
fn runtime(&self) -> Host<'_, Self, (), ()> {
Host {
services: self,
interceptors: &(),
stream: &(),
observers: Some(&self.events.sender),
}
}
}
impl HostCallHandler<TestProtocol> for Recording {
async fn handle_host_call(&self, reply: Reply<&'static str>) -> Result<(), TestError> {
let op = self.ops[self.calls.fetch_add(1, Ordering::SeqCst)];
self.seen.lock().unwrap().push(format!("op:{op}"));
if self.fail == Some(op) {
return Err(TestError::Service);
}
reply.send(op);
Ok(())
}
}
impl CallObserver for Recording {
fn observe(&self, event: CallEvent) {
self.seen.lock().unwrap().push(match event {
CallEvent::Started { .. } => "started".into(),
CallEvent::Succeeded { .. } => "succeeded".into(),
CallEvent::Failed { .. } => "failed".into(),
other => format!("{other:?}"),
});
}
}
fn scripted(
ops: &'static [&'static str],
outcome: Result<(), TestError>,
) -> CallMachine<TestProtocol, ()> {
CallMachine::new(None, move |host| {
Box::pin(async move {
for op in ops {
let answered = host.services.call(|reply| reply).await?;
assert_eq!(answered, *op);
}
outcome
})
})
}
#[rstest]
#[case::succeeds(&["sign", "send"], Ok(()), None, Ok(()), &["started", "op:sign", "op:send", "succeeded"])]
#[case::call_fails(&[], Err(TestError::Provider), None, Err(TestError::Provider), &["started", "failed"])]
#[case::service_fails(&["sign", "send", "never"], Ok(()), Some("send"), Err(TestError::Service), &["started", "op:sign", "op:send", "failed"])]
#[tokio::test]
async fn generic_runner_forwards_ops_and_emits_one_terminal(
#[case] ops: &'static [&'static str],
#[case] call_outcome: Result<(), TestError>,
#[case] fail: Option<&'static str>,
#[case] expected: Result<(), TestError>,
#[case] seen: &[&str],
) {
let host = Recording {
ops,
calls: AtomicUsize::new(0),
seen: Mutex::new(Vec::new()),
fail,
events: Observations::default(),
};
let outcome = run(scripted(ops, call_outcome), host.runtime()).await;
assert_eq!(outcome, expected);
assert_eq!(
*host.seen.lock().unwrap(),
seen.iter()
.filter(|item| item.starts_with("op:"))
.copied()
.collect::<Vec<_>>()
);
let events = host.events.lock().unwrap();
assert!(matches!(events.as_slice(), [CallEvent::Started { .. }, _]));
assert_eq!(
matches!(events[1], CallEvent::Failed { .. }),
expected.is_err()
);
assert_eq!(
matches!(events[1], CallEvent::Succeeded { .. }),
expected.is_ok()
);
}
#[rstest]
#[tokio::test]
async fn generic_runner_success_keeps_the_start_time(observer: Arc<Observer>) {
let services = Recording {
ops: &["send"],
calls: AtomicUsize::new(0),
seen: Mutex::new(Vec::new()),
fail: None,
events: Observations::default(),
};
let outcome = run(
scripted(&["send"], Ok(())),
Host {
services: &services,
interceptors: &(),
stream: &(),
observers: Some(&observer.0.sender),
},
)
.await;
assert_eq!(outcome, Ok(()));
let events = observer.0.lock().unwrap();
let [
CallEvent::Started { start_time },
CallEvent::Succeeded { timing, .. },
] = events.as_slice()
else {
panic!("unexpected events {events:?}");
};
assert_eq!(*start_time, timing.start_time);
}
#[rstest]
#[tokio::test]
async fn in_process_runner_dispatches_services_and_hooks(interceptors: Hooks) {
let observer = interceptors.observer.clone();
let completion = run_hosted(
dispatching_call(),
Host {
services: &Services { reject: false },
interceptors: &interceptors,
stream: &(),
observers: Some(&observer.0.sender),
},
)
.await
.unwrap();
assert_eq!(
completion,
HostedCompletion::Complete("projected/custom".into())
);
assert!(matches!(
observer.0.lock().unwrap().as_slice(),
[
CallEvent::Started { .. },
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }),
CallEvent::Succeeded { .. },
]
));
}
struct ResponseGate {
ready: tokio::sync::Notify,
reject: bool,
}
impl Interceptors<TestError> for ResponseGate {
async fn before_provider_request(
&self,
wire: WireRequest,
_: RequestContext,
) -> Result<WireRequest, TestError> {
Ok(wire)
}
async fn after_provider_response(&self, raw: RawResponse) -> Result<(), TestError> {
assert_eq!(raw.body, "provider response");
self.ready.notified().await;
if self.reject {
Err(TestError::Hook)
} else {
Ok(())
}
}
}
#[rstest]
#[case::accept(false)]
#[case::reject(true)]
#[tokio::test]
async fn response_interception_waits_and_can_reject_after_observation(#[case] reject: bool) {
let (sender, mut receiver) =
litellm_host::observation::observation_channel(std::num::NonZeroUsize::new(1).unwrap());
let machine = hosted_call::<TestProtocol, _, _>(
"input",
Some(sender),
|_, _, interceptors, observers| async move {
let raw = RawResponse {
body: "provider response".into(),
};
observers.unwrap().emit(CallEvent::Execution(
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
));
interceptors.after_provider_response(raw).await?;
Ok(CallOutput::Complete("accepted".into()))
},
);
let interceptor = ResponseGate {
ready: tokio::sync::Notify::new(),
reject,
};
let mut driver = Driver::new(machine, Services { reject: false }, &interceptor);
let mut advance = Box::pin(driver.advance());
assert!(futures_util::poll!(&mut advance).is_pending());
assert!(
matches!(receiver.try_recv(), Ok(CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw })) if raw.body == "provider response")
);
interceptor.ready.notify_one();
match advance.await {
Err(error) => {
assert!(reject);
assert_eq!(error, TestError::Hook);
}
Ok(Boundary::Complete(HostedCompletion::Complete(value))) => {
assert!(!reject);
assert_eq!(value, "accepted");
}
_ => panic!("expected completion or rejection"),
}
}

View file

@ -1,8 +1,31 @@
- 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 `PythonBinding`, `PythonHostCalls`, `PythonCallHooks` and `PythonOwned` traits
## Boundary with Python consumers
This crate owns CPython execution mechanics for generic `litellm-host` machines and hooks. Consumers supply domain bindings, host operations, public result construction and callback policy. Neither Rust dependencies nor Python imports may require LiteLLM route modules or legacy `Logging`
`src/native.rs` belongs here: it runs a generic machine through the Python runtime and owns its pending execution and abort handle. Keep provider selection, request projection and public exception policy out of it. A rename to `machine_runner.rs` is optional and must not change behavior
The execution handle receives its Python lifecycle binding from its consumer through `PythonLifecycle` rather than import a fixed `litellm.rust_bridge` module. Generic suspension and execution state validation belong here; public stream wrappers and `_hidden_params` conventions belong to the consumer
Creating a resolved asyncio Future from an already constructed Python value belongs here, alongside runtime waiting, interpreter detachment and panic containment. Choosing which callable exceptions become a public `RuntimeError` belongs to the consumer; `python-bridge::callable::wrap_failure` owns that policy
The driver owns ordering: start, argument preparation, prepared-argument hooks, binding decode and machine start. Fallible per-call resource setup supplied by the consumer runs after all argument hooks, using the prepared argument view, and before provider work. Setup failure follows the existing terminal failure path. Creating or discarding an unstarted coroutine must not initialize clients, acquire credentials or capture execution context
Boundary tests exercise behavior with a supplied lifecycle binding without importing the LiteLLM Python package. Pin inline awaiting, awaitable final values, exception identity, cancellation, re-entry and release of retained objects, rather than module names or source layout
`HookChain` composes Python runtime hooks in order. Each argument, wire-request and response transformation feeds its result to the next hook. After all argument transformations, the driver calls `arguments_prepared` on every hook in order. Retained callback views must adopt that dictionary before later policy hooks can mutate or reject it. SDK policy is supplied by bridge composition as a hook, never a separate driver phase or parameter. Hooks implement only the stages they need; default stages preserve the supplied values
Terminal notifications share the selected response or exception. An ordinary notification error is reported as unraisable and does not skip the next hook or replace the selected outcome. Preparation, interception and transformation errors stop the chain. Cancellation stops all further hook dispatch. Suspensions stay inline in the existing driver, and the chain traverses retained event values for GC
## Existing runtime invariants
- 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` and `PythonOwned` traits, and the `PythonRuntime` specialization of `host::hooks::CallHooks`
- 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 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)
- `PythonCallHooks` only constrains the shared call-stage interface to `PythonRuntime` and Python ownership. It must not redeclare the stages
- `PythonCallEvent` is a specialization of the shared `CallEvent`, never a separately defined lifecycle. The driver emits `Succeeded` or `Failed` exactly once and never dispatches callbacks after a cancellation; which Python objects consume those events is the legacy adapter's business
- `CallOptions` can publish snapshots independently of callback delivery. Terminal snapshots follow completed hook dispatch, and cancellation never calls a Python callback. For Python-driven calls, leave the machine's observation publisher unset so the driver is the sole publisher of intercepted provider-response snapshots
- `PythonBinding::decode_request` receives the keyword view returned by `prepare_arguments` and updated by `arguments_prepared`, not the caller's dict; a binding that decodes from it inherits all composed argument rewrites
- 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
@ -14,7 +37,7 @@
- Keep diagnostic counters in the consumer; wrapper invocations do not measure every interpreter release
- 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`
- Shared driver implementation: `litellm/rust_bridge/lifecycle.py`; handle: `src/handle.rs`; call driver: `src/driver.rs`; native-backed behavior tests: `tests/lifecycle.py`. The consumer supplies the lifecycle binding
- 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)

File diff suppressed because it is too large Load diff

View file

@ -5,6 +5,8 @@ use pyo3::exceptions::{PyBaseException, PyRuntimeError};
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
pub type PythonLifecycle = for<'py> fn(Python<'py>) -> PyResult<Bound<'py, PyModule>>;
pub enum ExecutionStep {
Return(Py<PyAny>),
Await(Py<PyAny>),
@ -29,22 +31,21 @@ enum ExecutionState {
#[pyclass]
pub struct Execution {
state: ExecutionState,
}
fn lifecycle(py: Python<'_>) -> PyResult<Bound<'_, PyModule>> {
py.import("litellm.rust_bridge.lifecycle")
lifecycle: PythonLifecycle,
}
impl Execution {
pub fn new(body: impl ExecutionBody + 'static) -> Self {
pub fn new(body: impl ExecutionBody + 'static, lifecycle: PythonLifecycle) -> Self {
Self {
state: ExecutionState::Created(Box::new(body)),
lifecycle,
}
}
pub fn into_coroutine(self, py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
let binding = (self.lifecycle)(py)?;
let execution = Py::new(py, self)?;
lifecycle(py)?.getattr("drive")?.call1((execution,))
binding.getattr("drive")?.call1((execution,))
}
pub(crate) fn into_sync_stream(
@ -52,15 +53,16 @@ impl Execution {
py: Python<'_>,
head: Py<PyAny>,
) -> PyResult<Bound<'_, PyAny>> {
lifecycle(py)?
(self.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 {
pub fn suspended(body: impl ExecutionBody + 'static, lifecycle: PythonLifecycle) -> Self {
Self {
state: ExecutionState::Suspended(Box::new(body)),
lifecycle,
}
}
@ -90,6 +92,7 @@ impl Execution {
_ => unreachable!(),
}
};
let lifecycle = slf.borrow().lifecycle;
let outcome = catch_unwind(AssertUnwindSafe(|| {
let step = body.resume(result)?;
let (tag, value, suspended) = match step {

View file

@ -1,14 +1,11 @@
mod chain;
pub use chain::HookChain;
use crate::PythonOwned;
use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest};
use litellm_host::hooks::{CallHooks, HookRuntime, RuntimeCallEvent};
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>>;
@ -18,62 +15,23 @@ pub enum HookStep<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,
},
pub struct PythonRuntime;
impl HookRuntime for PythonRuntime {
type Context<'a> = Python<'a>;
type Arguments = Py<PyDict>;
type Response = Py<PyAny>;
type Chunk = Py<PyAny>;
type Error = PyErr;
type Step<H, T> = HookStep<H, T>;
fn ready<H, T>(value: T) -> Self::Step<H, T> {
HookStep::Ready(value)
}
}
/// 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>>>;
pub type PythonCallEvent<'a> = RuntimeCallEvent<'a, PythonRuntime>;
fn before_provider_request(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<HookStep<Self, Box<WireRequest>>>;
pub trait PythonCallHooks: CallHooks<PythonRuntime> + PythonOwned {}
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<()>;
}
impl<H: CallHooks<PythonRuntime> + PythonOwned> PythonCallHooks for H {}

View file

@ -0,0 +1,202 @@
use litellm_host::interceptors::{RequestContext, WireRequest};
use litellm_host::lifecycle::Timing;
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,
types::PyDict,
};
use crate::{HookResume, HookStep, PythonCallEvent, PythonCallHooks, PythonOwned, missing_state};
pub(super) enum ChainStep<T> {
Ready(T),
Await(Py<PyAny>),
}
enum Continuation<H> {
Arguments(HookResume<H, Py<PyDict>>),
Wire(HookResume<H, Box<WireRequest>>),
Response(HookResume<H, Py<PyAny>>),
Event(HookResume<H, ()>),
}
pub(super) struct HookAdapter<H> {
hooks: H,
continuation: Option<Continuation<H>>,
}
impl<H> HookAdapter<H> {
pub(super) fn new(hooks: H) -> Self {
Self {
hooks,
continuation: None,
}
}
fn step<T>(
&mut self,
step: HookStep<H, T>,
continuation: impl FnOnce(HookResume<H, T>) -> Continuation<H>,
) -> ChainStep<T> {
match step {
HookStep::Ready(value) => ChainStep::Ready(value),
HookStep::Await(awaitable, resume) => {
self.continuation = Some(continuation(resume));
ChainStep::Await(awaitable)
}
}
}
}
pub(super) trait ChainHooks: PythonOwned {
fn prepare_arguments(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
started_at: f64,
) -> PyResult<ChainStep<Py<PyDict>>>;
fn resume_arguments(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<ChainStep<Py<PyDict>>>;
fn before_provider_request(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<ChainStep<Box<WireRequest>>>;
fn resume_wire(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<ChainStep<Box<WireRequest>>>;
fn transform_response(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
timing: Timing,
) -> PyResult<ChainStep<Py<PyAny>>>;
fn resume_response(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<ChainStep<Py<PyAny>>>;
fn on_event(&mut self, py: Python<'_>, event: PythonCallEvent<'_>) -> PyResult<ChainStep<()>>;
fn resume_event(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<ChainStep<()>>;
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()>;
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()>;
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()>;
}
impl<H: PythonCallHooks> ChainHooks for HookAdapter<H> {
fn prepare_arguments(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
started_at: f64,
) -> PyResult<ChainStep<Py<PyDict>>> {
let step = self.hooks.prepare_arguments(py, arguments, started_at)?;
Ok(self.step(step, Continuation::Arguments))
}
fn resume_arguments(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<ChainStep<Py<PyDict>>> {
let Some(Continuation::Arguments(resume)) = self.continuation.take() else {
return Err(missing_state());
};
let step = resume(&mut self.hooks, py, result)?;
Ok(self.step(step, Continuation::Arguments))
}
fn before_provider_request(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<ChainStep<Box<WireRequest>>> {
let step = self.hooks.before_provider_request(py, wire, context)?;
Ok(self.step(step, Continuation::Wire))
}
fn resume_wire(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<ChainStep<Box<WireRequest>>> {
let Some(Continuation::Wire(resume)) = self.continuation.take() else {
return Err(missing_state());
};
let step = resume(&mut self.hooks, py, result)?;
Ok(self.step(step, Continuation::Wire))
}
fn transform_response(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
timing: Timing,
) -> PyResult<ChainStep<Py<PyAny>>> {
let step = self.hooks.transform_response(py, response, timing)?;
Ok(self.step(step, Continuation::Response))
}
fn resume_response(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<ChainStep<Py<PyAny>>> {
let Some(Continuation::Response(resume)) = self.continuation.take() else {
return Err(missing_state());
};
let step = resume(&mut self.hooks, py, result)?;
Ok(self.step(step, Continuation::Response))
}
fn on_event(&mut self, py: Python<'_>, event: PythonCallEvent<'_>) -> PyResult<ChainStep<()>> {
let step = self.hooks.on_event(py, event)?;
Ok(self.step(step, Continuation::Event))
}
fn resume_event(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<ChainStep<()>> {
let Some(Continuation::Event(resume)) = self.continuation.take() else {
return Err(missing_state());
};
let step = resume(&mut self.hooks, py, result)?;
Ok(self.step(step, Continuation::Event))
}
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()> {
self.hooks.arguments_prepared(py, arguments)
}
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
self.hooks.on_stream_open(py)
}
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
self.hooks.on_stream_chunk(py, chunk)
}
}
impl<H: PythonOwned> PythonOwned for HookAdapter<H> {
fn close(&mut self, py: Python<'_>) {
self.continuation = None;
self.hooks.close(py);
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
self.hooks.traverse(visit)
}
}

View file

@ -0,0 +1,387 @@
use litellm_host::{
hooks::CallHooks,
interceptors::{RawResponse, RequestContext, WireRequest},
lifecycle::{CallEvent, ExecutionEvent, Timing},
};
use pyo3::{
exceptions::{PyBaseException, PyException},
gc::{PyTraverseError, PyVisit},
prelude::*,
types::PyDict,
};
use crate::{
HookStep, PythonCallEvent, PythonCallHooks, PythonOwned, PythonRuntime, missing_state,
};
use super::adapter::{ChainHooks, ChainStep, HookAdapter};
type OwnedEvent = CallEvent<Py<PyAny>, Py<PyBaseException>, RawResponse>;
#[derive(Default)]
pub struct HookChain {
hooks: Vec<Box<dyn ChainHooks>>,
arguments: Option<(usize, f64)>,
wire: Option<(usize, RequestContext)>,
response: Option<(usize, Timing)>,
event: Option<(usize, OwnedEvent)>,
}
impl HookChain {
pub fn new() -> Self {
Self::default()
}
pub fn with(mut self, hooks: impl PythonCallHooks + 'static) -> Self {
self.hooks.push(Box::new(HookAdapter::new(hooks)));
self
}
fn arguments_from(
&mut self,
py: Python<'_>,
index: usize,
arguments: Py<PyDict>,
started_at: f64,
) -> PyResult<HookStep<Self, Py<PyDict>>> {
let Some(hooks) = self.hooks.get_mut(index) else {
return Ok(HookStep::Ready(arguments));
};
let step = hooks.prepare_arguments(py, arguments, started_at)?;
self.arguments_step(py, index, step, started_at)
}
fn arguments_step(
&mut self,
py: Python<'_>,
index: usize,
step: ChainStep<Py<PyDict>>,
started_at: f64,
) -> PyResult<HookStep<Self, Py<PyDict>>> {
match step {
ChainStep::Ready(value) => self.arguments_from(py, index + 1, value, started_at),
ChainStep::Await(awaitable) => {
self.arguments = Some((index, started_at));
Ok(HookStep::Await(awaitable, Self::resume_arguments))
}
}
}
fn resume_arguments(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, Py<PyDict>>> {
let (index, context) = self.arguments.take().ok_or_else(missing_state)?;
let result = resume_unless_cancelled(py, result)?;
let step = self.hooks[index].resume_arguments(py, result)?;
self.arguments_step(py, index, step, context)
}
fn wire_from(
&mut self,
py: Python<'_>,
index: usize,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<HookStep<Self, Box<WireRequest>>> {
let Some(hooks) = self.hooks.get_mut(index) else {
return Ok(HookStep::Ready(wire));
};
let step = hooks.before_provider_request(py, wire, context)?;
self.wire_step(py, index, step, context)
}
fn wire_step(
&mut self,
py: Python<'_>,
index: usize,
step: ChainStep<Box<WireRequest>>,
context: &RequestContext,
) -> PyResult<HookStep<Self, Box<WireRequest>>> {
match step {
ChainStep::Ready(value) => self.wire_from(py, index + 1, value, context),
ChainStep::Await(awaitable) => {
self.wire = Some((index, context.clone()));
Ok(HookStep::Await(awaitable, Self::resume_wire))
}
}
}
fn resume_wire(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, Box<WireRequest>>> {
let (index, context) = self.wire.take().ok_or_else(missing_state)?;
let result = resume_unless_cancelled(py, result)?;
let step = self.hooks[index].resume_wire(py, result)?;
self.wire_step(py, index, step, &context)
}
fn response_from(
&mut self,
py: Python<'_>,
index: usize,
response: Py<PyAny>,
timing: Timing,
) -> PyResult<HookStep<Self, Py<PyAny>>> {
let Some(hooks) = self.hooks.get_mut(index) else {
return Ok(HookStep::Ready(response));
};
let step = hooks.transform_response(py, response, timing)?;
self.response_step(py, index, step, timing)
}
fn response_step(
&mut self,
py: Python<'_>,
index: usize,
step: ChainStep<Py<PyAny>>,
timing: Timing,
) -> PyResult<HookStep<Self, Py<PyAny>>> {
match step {
ChainStep::Ready(value) => self.response_from(py, index + 1, value, timing),
ChainStep::Await(awaitable) => {
self.response = Some((index, timing));
Ok(HookStep::Await(awaitable, Self::resume_response))
}
}
}
fn resume_response(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, Py<PyAny>>> {
let (index, context) = self.response.take().ok_or_else(missing_state)?;
let result = resume_unless_cancelled(py, result)?;
let step = self.hooks[index].resume_response(py, result)?;
self.response_step(py, index, step, context)
}
fn event_from(
&mut self,
py: Python<'_>,
index: usize,
event: OwnedEvent,
) -> PyResult<HookStep<Self, ()>> {
let Some(hooks) = self.hooks.get_mut(index) else {
return Ok(HookStep::Ready(()));
};
let result = dispatch(py, hooks.as_mut(), &event);
let step = notification_result(py, is_terminal(&event), result)?;
self.event_step(py, index, step, event)
}
fn event_step(
&mut self,
py: Python<'_>,
index: usize,
step: ChainStep<()>,
event: OwnedEvent,
) -> PyResult<HookStep<Self, ()>> {
match step {
ChainStep::Ready(()) => self.event_from(py, index + 1, event),
ChainStep::Await(awaitable) => {
self.event = Some((index, event));
Ok(HookStep::Await(awaitable, Self::resume_event))
}
}
}
fn resume_event(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, ()>> {
let (index, event) = self.event.take().ok_or_else(missing_state)?;
let result = resume_unless_cancelled(py, result)?;
let result = self.hooks[index].resume_event(py, result);
let step = notification_result(py, is_terminal(&event), result)?;
self.event_step(py, index, step, event)
}
}
fn resume_unless_cancelled(
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<PyResult<Py<PyAny>>> {
match result {
Err(error) if !error.is_instance_of::<PyException>(py) => Err(error),
result => Ok(result),
}
}
fn is_terminal<Response, Error, Raw>(event: &CallEvent<Response, Error, Raw>) -> bool {
matches!(
event,
CallEvent::Succeeded { .. } | CallEvent::Failed { .. }
)
}
fn notification_result(
py: Python<'_>,
terminal: bool,
result: PyResult<ChainStep<()>>,
) -> PyResult<ChainStep<()>> {
match result {
Err(error) if terminal && error.is_instance_of::<PyException>(py) => {
error.write_unraisable(py, None);
Ok(ChainStep::Ready(()))
}
result => result,
}
}
fn retain_event(py: Python<'_>, event: PythonCallEvent<'_>) -> OwnedEvent {
match event {
CallEvent::Started { start_time } => CallEvent::Started { start_time },
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw: raw.clone() })
}
CallEvent::Succeeded { timing, response } => CallEvent::Succeeded {
timing,
response: response.clone_ref(py),
},
CallEvent::Failed {
timing,
origin,
error,
} => CallEvent::Failed {
timing,
origin,
error: error.clone_ref(py).into_value(py),
},
CallEvent::Cancelled { timing } => CallEvent::Cancelled { timing },
}
}
fn dispatch(
py: Python<'_>,
hooks: &mut dyn ChainHooks,
event: &OwnedEvent,
) -> PyResult<ChainStep<()>> {
match event {
CallEvent::Started { start_time } => hooks.on_event(
py,
CallEvent::Started {
start_time: *start_time,
},
),
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => hooks.on_event(
py,
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }),
),
CallEvent::Succeeded { timing, response } => hooks.on_event(
py,
CallEvent::Succeeded {
timing: *timing,
response,
},
),
CallEvent::Failed {
timing,
origin,
error,
} => hooks.on_event(
py,
CallEvent::Failed {
timing: *timing,
origin: *origin,
error: &PyErr::from_value(error.bind(py).clone().into_any()),
},
),
CallEvent::Cancelled { timing } => {
hooks.on_event(py, CallEvent::Cancelled { timing: *timing })
}
}
}
impl CallHooks<PythonRuntime> for HookChain {
fn prepare_arguments(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
started_at: f64,
) -> PyResult<HookStep<Self, Py<PyDict>>> {
self.arguments_from(py, 0, arguments, started_at)
}
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()> {
self.hooks
.iter_mut()
.try_for_each(|hooks| hooks.arguments_prepared(py, arguments))
}
fn before_provider_request(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<HookStep<Self, Box<WireRequest>>> {
self.wire_from(py, 0, wire, context)
}
fn transform_response(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
timing: Timing,
) -> PyResult<HookStep<Self, Py<PyAny>>> {
self.response_from(py, 0, response, timing)
}
fn on_event(
&mut self,
py: Python<'_>,
event: PythonCallEvent<'_>,
) -> PyResult<HookStep<Self, ()>> {
for (index, hooks) in self.hooks.iter_mut().enumerate() {
let result = hooks.on_event(py, event.clone());
match notification_result(py, is_terminal(&event), result)? {
ChainStep::Ready(()) => {}
ChainStep::Await(awaitable) => {
self.event = Some((index, retain_event(py, event)));
return Ok(HookStep::Await(awaitable, Self::resume_event));
}
}
}
Ok(HookStep::Ready(()))
}
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
self.hooks
.iter_mut()
.try_for_each(|hooks| hooks.on_stream_open(py))
}
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
self.hooks
.iter_mut()
.try_for_each(|hooks| hooks.on_stream_chunk(py, chunk))
}
}
impl PythonOwned for HookChain {
fn close(&mut self, py: Python<'_>) {
self.arguments = None;
self.wire = None;
self.response = None;
self.event = None;
for hooks in &mut self.hooks {
hooks.close(py);
}
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
for hooks in &self.hooks {
hooks.traverse(visit)?;
}
match &self.event {
Some((_, CallEvent::Succeeded { response, .. })) => visit.call(response),
Some((_, CallEvent::Failed { error, .. })) => visit.call(error),
_ => Ok(()),
}
}
}

View file

@ -0,0 +1,4 @@
mod adapter;
mod dispatch;
pub use dispatch::HookChain;

View file

@ -5,7 +5,6 @@
mod argument;
mod binding;
mod callable;
mod driver;
mod error;
mod file_reader;
@ -21,22 +20,21 @@ mod services;
pub use argument::lookup;
pub use binding::PythonBinding;
pub use callable::wrap_failure;
pub use driver::run_call;
pub use driver::{CallOptions, run_call};
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 handle::{Execution, ExecutionBody, ExecutionStep, PythonLifecycle};
pub use hooks::{HookChain, HookResume, HookStep, PythonCallEvent, PythonCallHooks, PythonRuntime};
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,
ready_future, reserve_process_for_forking, run_async, run_async_value, run_sync,
run_sync_value, runtime_started,
};
pub use services::PythonHostCalls;

View file

@ -73,6 +73,18 @@ where
pyo3_async_runtimes::tokio::future_into_py(py, future)
}
pub fn ready_future<'py>(
py: Python<'py>,
value: &Bound<'py, PyAny>,
) -> PyResult<Bound<'py, PyAny>> {
let future = py
.import("asyncio")?
.call_method0("get_running_loop")?
.call_method0("create_future")?;
future.call_method1("set_result", (value,))?;
Ok(future)
}
pub fn run_sync<T, E, F>(
py: Python<'_>,
future: F,
@ -232,6 +244,58 @@ mod tests {
use super::*;
use crate::{InitializedPython, initialized_python};
#[pyfunction]
fn completed_future<'py>(
py: Python<'py>,
value: Bound<'py, PyAny>,
) -> PyResult<Bound<'py, PyAny>> {
ready_future(py, &value)
}
#[rstest]
fn a_ready_future_preserves_identity_and_the_callers_loop(
#[from(initialized_python)] python: &InitializedPython,
) {
python.attach(|py| {
let locals = PyDict::new(py);
locals
.set_item(
"completed_future",
wrap_pyfunction!(completed_future, py).unwrap(),
)
.unwrap();
py.run(
c"
import asyncio
async def exercise():
value = object()
future = completed_future(value)
assert isinstance(future, asyncio.Future)
assert future.done()
assert future.get_loop() is asyncio.get_running_loop()
assert future.result() is value
assert await future is value
asyncio.run(exercise())
",
Some(&locals),
Some(&locals),
)
.unwrap();
});
}
#[rstest]
fn a_ready_future_requires_a_running_loop(
#[from(initialized_python)] python: &InitializedPython,
) {
python.attach(|py| {
let error = ready_future(py, py.None().bind(py)).unwrap_err();
assert!(error.is_instance_of::<PyRuntimeError>(py));
});
}
#[derive(Debug)]
struct Error(String);

View file

@ -0,0 +1,734 @@
use litellm_host::{
hooks::CallHooks,
interceptors::{RequestContext, WireRequest},
lifecycle::{FailureOrigin, Timing},
};
use litellm_host_python::{HookChain, HookStep, PythonCallEvent, PythonOwned, PythonRuntime};
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,
types::{PyDict, PyTuple},
};
use rstest::{fixture, rstest};
struct ScriptHooks {
object: Py<PyAny>,
asynchronous: bool,
wire: Option<Box<WireRequest>>,
}
impl ScriptHooks {
fn invoke(&self, py: Python<'_>, name: &str, value: Py<PyAny>) -> PyResult<Py<PyAny>> {
if self.asynchronous {
self.object.call_method1(py, "invoke", (name, value))
} else {
self.object.call_method1(py, name, (value,))
}
}
fn arguments(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, Py<PyDict>>> {
Ok(HookStep::Ready(
result?.into_bound(py).cast_into::<PyDict>()?.unbind(),
))
}
fn response(
&mut self,
_: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, Py<PyAny>>> {
result.map(HookStep::Ready)
}
fn notification(
&mut self,
_: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, ()>> {
result.map(|_| HookStep::Ready(()))
}
fn request(
&mut self,
py: Python<'_>,
result: PyResult<Py<PyAny>>,
) -> PyResult<HookStep<Self, Box<WireRequest>>> {
let wire = self.wire.take().unwrap();
Ok(HookStep::Ready(Box::new(WireRequest {
url: result?.extract(py)?,
..*wire
})))
}
}
impl PythonOwned for ScriptHooks {
fn close(&mut self, py: Python<'_>) {
self.object = py.None();
self.wire = None;
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.object)
}
}
impl CallHooks<PythonRuntime> for ScriptHooks {
fn prepare_arguments(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
_: f64,
) -> PyResult<HookStep<Self, Py<PyDict>>> {
let value = self.invoke(py, "prepare", arguments.into_any())?;
if self.asynchronous {
Ok(HookStep::Await(value, Self::arguments))
} else {
self.arguments(py, Ok(value))
}
}
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()> {
self.object.call_method1(py, "adopt", (arguments,))?;
Ok(())
}
fn before_provider_request(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<HookStep<Self, Box<WireRequest>>> {
self.object.bind(py).setattr("model", &context.model)?;
let value = self.invoke(
py,
"before",
wire.url.clone().into_pyobject(py)?.into_any().unbind(),
)?;
self.wire = Some(wire);
if self.asynchronous {
Ok(HookStep::Await(value, Self::request))
} else {
self.request(py, Ok(value))
}
}
fn transform_response(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
_: Timing,
) -> PyResult<HookStep<Self, Py<PyAny>>> {
let value = self.invoke(py, "transform", response)?;
if self.asynchronous {
Ok(HookStep::Await(value, Self::response))
} else {
self.response(py, Ok(value))
}
}
fn on_event(
&mut self,
py: Python<'_>,
event: PythonCallEvent<'_>,
) -> PyResult<HookStep<Self, ()>> {
let (name, value) = match event {
PythonCallEvent::Succeeded { response, .. } => ("success", response.clone_ref(py)),
PythonCallEvent::Failed { error, .. } => {
("failure", error.clone_ref(py).into_value(py).into_any())
}
PythonCallEvent::Started { .. } => ("started", py.None()),
PythonCallEvent::Cancelled { .. } => ("cancelled", py.None()),
PythonCallEvent::Execution(_) => ("provider", py.None()),
};
let value = self.invoke(
py,
"event",
PyTuple::new(py, [name.into_pyobject(py)?.into_any().unbind(), value])?
.into_any()
.unbind(),
)?;
if self.asynchronous {
Ok(HookStep::Await(value, Self::notification))
} else {
self.notification(py, Ok(value))
}
}
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
self.object.call_method1(py, "stream", (py.None(),))?;
Ok(())
}
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
self.object.call_method1(py, "stream", (chunk,))?;
Ok(())
}
}
#[fixture]
fn scripts() -> Py<PyDict> {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"
import asyncio
log = []
class Hooks:
def __init__(self, name, delegate=None):
self.name = name
self.delegate = delegate
self.error = None
self.adopted = None
async def invoke(self, name, value):
if self.delegate is not None:
return await self.delegate(name, value)
await asyncio.sleep(0)
return getattr(self, name)(value)
def prepare(self, value):
log.append((self.name, 'prepare', value))
if self.error:
raise self.error
return {**value, 'order': value.get('order', '') + self.name}
def adopt(self, value):
self.adopted = value
def before(self, value):
return value + self.name
def transform(self, value):
return (value, self.name)
def event(self, value):
log.append((self.name, *value))
if self.error:
raise self.error
def stream(self, value):
log.append((self.name, 'stream', value))
first = Hooks('a')
second = Hooks('b')
third = Hooks('c')
fourth = Hooks('d')
",
Some(&locals),
Some(&locals),
)
.unwrap();
locals.unbind()
})
}
fn chain(py: Python<'_>, scripts: &Py<PyDict>, asynchronous: bool) -> HookChain {
let hook = |name| ScriptHooks {
object: scripts.bind(py).get_item(name).unwrap().unwrap().unbind(),
asynchronous,
wire: None,
};
HookChain::new().with(hook("first")).with(hook("second"))
}
fn finish<H, T>(py: Python<'_>, hooks: &mut H, step: HookStep<H, T>) -> PyResult<T> {
match step {
HookStep::Ready(value) => Ok(value),
HookStep::Await(awaitable, resume) => {
let value = py
.import("asyncio")?
.call_method1("run", (awaitable,))
.map(Bound::unbind);
let next = resume(hooks, py, value)?;
finish(py, hooks, next)
}
}
}
const TIMING: Timing = Timing {
start_time: 0.0,
end_time: 1.0,
};
struct PreparedPolicy;
impl CallHooks<PythonRuntime> for PreparedPolicy {
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()> {
arguments.bind(py).set_item("policy", "configured")
}
}
impl PythonOwned for PreparedPolicy {
fn close(&mut self, _: Python<'_>) {}
fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> {
Ok(())
}
}
#[rstest]
#[case::sync(false)]
#[case::suspended(true)]
fn transformations_feed_each_other_and_notifications_share_final_values(
scripts: Py<PyDict>,
#[case] asynchronous: bool,
) {
Python::attach(|py| {
let mut hooks = chain(py, &scripts, asynchronous).with(PreparedPolicy);
let original = PyDict::new(py).unbind();
let step = hooks
.prepare_arguments(py, original.clone_ref(py), 0.0)
.unwrap();
let arguments = finish(py, &mut hooks, step).unwrap();
hooks.arguments_prepared(py, &arguments).unwrap();
let context = RequestContext {
model: "model".into(),
custom_llm_provider: "provider".into(),
optional_params: serde_json::Value::Null,
secret_fields: vec![],
api_key: None,
};
let wire = Box::new(WireRequest {
url: "url".into(),
headers: vec![],
body: serde_json::Value::Null,
});
let step = hooks.before_provider_request(py, wire, &context).unwrap();
assert_eq!(finish(py, &mut hooks, step).unwrap().url, "urlab");
let response = PyDict::new(py).unbind().into_any();
let step = hooks
.transform_response(py, response.clone_ref(py), TIMING)
.unwrap();
let final_response = finish(py, &mut hooks, step).unwrap();
let step = hooks
.on_event(
py,
PythonCallEvent::Succeeded {
timing: TIMING,
response: &final_response,
},
)
.unwrap();
finish(py, &mut hooks, step).unwrap();
hooks.on_stream_open(py).unwrap();
hooks.on_stream_chunk(py, &response).unwrap();
let locals = scripts.bind(py);
locals.set_item("arguments", arguments).unwrap();
locals.set_item("original", original).unwrap();
locals.set_item("response", response).unwrap();
locals.set_item("final_response", final_response).unwrap();
py.run(
c"
assert arguments['order'] == 'ab'
assert original == {}
assert first.adopted is arguments and second.adopted is arguments
assert first.adopted['policy'] == second.adopted['policy'] == 'configured'
assert first.model == second.model == 'model'
assert final_response == ((response, 'a'), 'b')
assert [(name, kind) for name, kind, value in log] == [
('a', 'prepare'), ('b', 'prepare'), ('a', 'success'), ('b', 'success'),
('a', 'stream'), ('b', 'stream'), ('a', 'stream'), ('b', 'stream'),
]
assert log[2][2] is final_response and log[3][2] is final_response
assert log[6][2] is response and log[7][2] is response
",
Some(locals),
Some(locals),
)
.unwrap();
});
}
#[rstest]
#[case::sync_success(false, false)]
#[case::async_success(true, false)]
#[case::sync_failure(false, true)]
#[case::async_failure(true, true)]
fn terminal_failure_does_not_skip_later_hooks(
scripts: Py<PyDict>,
#[case] asynchronous: bool,
#[case] failed: bool,
#[values("first", "second")] failing_hook: &str,
) {
Python::attach(|py| {
let mut hooks = chain(py, &scripts, asynchronous).with(ScriptHooks {
object: scripts
.bind(py)
.get_item("third")
.unwrap()
.unwrap()
.unbind(),
asynchronous,
wire: None,
});
let locals = scripts.bind(py);
locals
.get_item(failing_hook)
.unwrap()
.unwrap()
.setattr(
"error",
pyo3::exceptions::PyRuntimeError::new_err("callback failed").into_value(py),
)
.unwrap();
let error = pyo3::exceptions::PyValueError::new_err("provider failed");
let response = py.None();
let event = if failed {
PythonCallEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Call,
error: &error,
}
} else {
PythonCallEvent::Succeeded {
timing: TIMING,
response: &response,
}
};
let step = hooks.on_event(py, event).unwrap();
finish(py, &mut hooks, step).unwrap();
locals
.set_item(
"selected",
if failed {
error.into_value(py).into_any()
} else {
response
},
)
.unwrap();
py.run(
c"
assert [name for name, kind, value in log] == ['a', 'b', 'c']
assert all(kind == log[0][1] for name, kind, value in log)
assert all(value is selected for name, kind, value in log)
",
Some(locals),
Some(locals),
)
.unwrap();
});
}
#[rstest]
#[case::sync(false)]
#[case::async_(true)]
fn transformation_failure_stops_the_chain(scripts: Py<PyDict>, #[case] asynchronous: bool) {
Python::attach(|py| {
let mut hooks = chain(py, &scripts, asynchronous);
let locals = scripts.bind(py);
py.run(
c"first.error = ValueError('prepare failed')",
Some(locals),
Some(locals),
)
.unwrap();
let result = hooks
.prepare_arguments(py, PyDict::new(py).unbind(), 0.0)
.and_then(|step| finish(py, &mut hooks, step));
assert!(
result.unwrap_err().value(py).is(locals
.get_item("first")
.unwrap()
.unwrap()
.getattr("error")
.unwrap())
);
py.run(
c"assert len(log) == 1 and log[0][0] == 'a'",
Some(locals),
Some(locals),
)
.unwrap();
});
}
#[rstest]
#[case::sync(false)]
#[case::async_(true)]
fn cancellation_stops_notification_dispatch(scripts: Py<PyDict>, #[case] asynchronous: bool) {
Python::attach(|py| {
let mut hooks = chain(py, &scripts, asynchronous);
let locals = scripts.bind(py);
py.run(
c"first.error = asyncio.CancelledError()",
Some(locals),
Some(locals),
)
.unwrap();
let response = py.None();
let result = hooks
.on_event(
py,
PythonCallEvent::Succeeded {
timing: TIMING,
response: &response,
},
)
.and_then(|step| finish(py, &mut hooks, step));
assert!(
result
.unwrap_err()
.is_instance_of::<pyo3::exceptions::asyncio::CancelledError>(py)
);
py.run(
c"assert len(log) == 1 and log[0][0] == 'a'",
Some(locals),
Some(locals),
)
.unwrap();
});
}
struct Preparing {
hooks: HookChain,
arguments: Option<Py<PyDict>>,
resume: Option<litellm_host_python::HookResume<HookChain, Py<PyDict>>>,
}
impl litellm_host_python::ExecutionBody for Preparing {
fn resume(
&mut self,
result: Option<PyResult<Py<PyAny>>>,
) -> PyResult<litellm_host_python::ExecutionStep> {
Python::attach(|py| {
let step = match self.resume.take() {
Some(resume) => resume(&mut self.hooks, py, result.unwrap())?,
None => self
.hooks
.prepare_arguments(py, self.arguments.take().unwrap(), 0.0)?,
};
match step {
HookStep::Ready(arguments) => {
self.hooks.arguments_prepared(py, &arguments)?;
Ok(litellm_host_python::ExecutionStep::Return(
arguments.into_any(),
))
}
HookStep::Await(awaitable, resume) => {
self.resume = Some(resume);
Ok(litellm_host_python::ExecutionStep::Await(awaitable))
}
}
})
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.arguments)?;
self.hooks.traverse(visit)
}
}
fn lifecycle(py: Python<'_>) -> PyResult<Bound<'_, PyModule>> {
let source =
std::ffi::CString::new(include_str!("../../../../litellm/rust_bridge/lifecycle.py"))
.unwrap();
PyModule::from_code(py, &source, c"lifecycle.py", c"hook_chain_test_lifecycle")
}
#[rstest]
fn suspended_hooks_keep_the_callers_task_and_context(scripts: Py<PyDict>) {
Python::attach(|py| {
let locals = scripts.bind(py);
py.run(
c"
import contextvars
import threading
state = contextvars.ContextVar('state')
async def invoke(name, value):
assert asyncio.current_task() is caller
assert threading.get_ident() == thread
if name == 'prepare':
state.set(state.get() + 'x')
await asyncio.sleep(0)
assert asyncio.current_task() is caller
return {**value, 'context': state.get()}
first = Hooks('a', invoke)
second = Hooks('b', invoke)
",
Some(locals),
Some(locals),
)
.unwrap();
let hooks = chain(py, &scripts, true);
let execution = litellm_host_python::Execution::new(
Preparing {
hooks,
arguments: Some(PyDict::new(py).unbind()),
resume: None,
},
lifecycle,
);
locals
.set_item("call", execution.into_coroutine(py).unwrap())
.unwrap();
py.run(
c"
async def exercise():
global caller, thread
caller = asyncio.current_task()
thread = threading.get_ident()
state.set('caller')
result = await call
assert result['context'] == 'callerxx'
assert state.get() == 'callerxx'
assert first.adopted is result and second.adopted is result
asyncio.run(exercise())
",
Some(locals),
Some(locals),
)
.unwrap();
});
}
#[pyclass(weakref)]
struct HookOwner {
hooks: HookChain,
}
#[pymethods]
impl HookOwner {
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
self.hooks.traverse(&visit)
}
fn __clear__(&mut self, py: Python<'_>) {
self.hooks.close(py);
}
}
#[rstest]
#[case::first_response(false, false)]
#[case::first_exception(true, false)]
#[case::second_response(false, true)]
#[case::second_exception(true, true)]
fn suspended_notification_cycles_are_collectable(
scripts: Py<PyDict>,
#[case] failed: bool,
#[case] first_ready: bool,
) {
Python::attach(|py| {
let hook = |name, asynchronous| ScriptHooks {
object: scripts.bind(py).get_item(name).unwrap().unwrap().unbind(),
asynchronous,
wire: None,
};
let owner = Py::new(
py,
HookOwner {
hooks: HookChain::new()
.with(hook("first", !first_ready))
.with(hook("second", true)),
},
)
.unwrap();
let locals = scripts.bind(py);
locals.set_item("owner", &owner).unwrap();
py.run(
c"
import gc
import weakref
class Payload(Exception):
pass
payload = Payload()
payload.owner = owner
owner_ref = weakref.ref(owner)
payload_ref = weakref.ref(payload)
",
Some(locals),
Some(locals),
)
.unwrap();
let payload = locals.get_item("payload").unwrap().unwrap().unbind();
let step = if failed {
owner.borrow_mut(py).hooks.on_event(
py,
PythonCallEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Call,
error: &PyErr::from_value(payload.bind(py).clone()),
},
)
} else {
owner.borrow_mut(py).hooks.on_event(
py,
PythonCallEvent::Succeeded {
timing: TIMING,
response: &payload,
},
)
}
.unwrap();
let HookStep::Await(awaitable, _) = step else {
panic!("notification must suspend")
};
awaitable.call_method0(py, "close").unwrap();
drop(awaitable);
drop(payload);
drop(owner);
py.run(
c"
log.clear()
del owner, payload
gc.collect()
assert owner_ref() is None
assert payload_ref() is None
",
Some(locals),
Some(locals),
)
.unwrap();
});
}
#[rstest]
#[case::empty(0, "")]
#[case::single(1, "a")]
#[case::pair(2, "ab")]
#[case::three(3, "abc")]
#[case::four(4, "abcd")]
fn builder_runs_hooks_in_append_order(
scripts: Py<PyDict>,
#[case] count: usize,
#[case] expected: &str,
#[values(false, true)] asynchronous: bool,
) {
Python::attach(|py| {
let mut hooks = ["first", "second", "third", "fourth"]
.into_iter()
.take(count)
.fold(HookChain::new(), |chain, name| {
chain.with(ScriptHooks {
object: scripts.bind(py).get_item(name).unwrap().unwrap().unbind(),
asynchronous,
wire: None,
})
});
let original = PyDict::new(py).unbind();
let step = hooks
.prepare_arguments(py, original.clone_ref(py), 0.0)
.unwrap();
let result = finish(py, &mut hooks, step).unwrap();
if count == 0 {
assert!(result.bind(py).is(original.bind(py)));
} else {
assert_eq!(
result
.bind(py)
.get_item("order")
.unwrap()
.unwrap()
.extract::<String>()
.unwrap(),
expected,
);
}
assert!(original.bind(py).is_empty());
let locals = scripts.bind(py);
locals.set_item("expected", expected).unwrap();
py.run(
c"assert ''.join(name for name, kind, value in log) == expected",
Some(locals),
Some(locals),
)
.unwrap();
});
}

View file

@ -3,17 +3,25 @@
| 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 |
| Host services | `Protocol::HostCall`, `HostServices::call` | `host-native::services::HostCallHandler` answers typed calls; `()` handles protocols without host calls | `PythonHostCalls` invokes retained Python objects; it may share an owner with the binding |
| Active hooks | `Interceptors::{before_provider_request, after_provider_response}` and `hooks::CallHooks<Runtime>` | Request interception and fallible execution callbacks | `PythonCallHooks` also prepares arguments, transforms public responses and receives stream callbacks |
| Passive observation | `observation::ObservationSender` | Queued execution and lifecycle snapshots, retained by the response body | Public Python callbacks remain active hooks with their existing failure policy |
| Runtime driving | `Machine`, `HostRequest::{HostCall, Intercept, 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
`hosted_call(request, observers, execute)` starts with a typed request. Its route closure receives separate `HostServices`, `ChannelInterceptors` and optional observation publisher; 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, interceptors, observers 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
`interceptors.rs` owns `Interceptors` and its request/response payload types. `lifecycle.rs` owns `CallObserver`, `CallEvent`, `ExecutionEvent`, timing, failure origin, and the observation wrappers. Event payloads are generic so a runtime can retain its own response, exception and raw-response references without introducing a language dependency. `snapshot()` projects them into the owned observation contract without retaining runtime objects. Pass interceptors and observers separately at direct route and HTTP entrypoints. Routes publish execution events independently of interception
`hooks.rs` owns the call-stage interface and its runtime-associated types. It contains no Python types or legacy callback policy. A runtime supplies its context and continuation representation through `HookRuntime`
`protocol.rs` owns `Protocol` and suspension messages, including `InterceptRequest` and `StreamDelivery`. `call.rs` owns route outputs and their adaptation into a hosted machine. Rust service handling belongs in `host-native::services`; coroutine channel handles stay in `machine/context.rs`
Rust handlers answer suspensions through `litellm-host-native::Driver`, which `litellm-host-http` and `litellm_host_native::in_process` share. `in_process::Host` is an assembly of services, interceptors, stream consumer and optional observation publisher. 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
Interceptors run inline and can rewrite values or fail execution. Observers consume owned `CallEvent` snapshots from `observation_channel`; its bounded `ObservationSender` never waits for delivery and counts events dropped when the queue is full or closed. The host owns receiver processing and draining. Pass the same publisher to machine construction and the driver when one receiver should collect execution and lifecycle events. Legacy Python callbacks retain their existing awaited, fallible behavior through the Python adapter

View file

@ -1,10 +1,11 @@
use crate::observation::ObservationSender;
use std::future::Future;
use futures_util::{TryStreamExt, stream::BoxStream};
use crate::{
machine::{CallMachine, ChannelHooks, HostServices, MachineFault},
protocol::{Demand, Protocol},
machine::{CallMachine, ChannelInterceptors, HostServices, MachineFault},
protocol::Protocol,
};
pub enum CallOutput<Response, Head, Chunk, Error> {
@ -37,23 +38,34 @@ pub type OutputOf<P> = CallOutput<
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>
pub fn hosted_call<P, F, Fut>(
request: P::Request,
observers: Option<ObservationSender>,
execute: F,
) -> HostedMachine<P>
where
P: Protocol,
P::Error: From<MachineFault>,
F: FnOnce(P::Request, HostServices<P>, ChannelHooks<P>) -> Fut + Send + 'static,
F: FnOnce(
P::Request,
HostServices<P>,
ChannelInterceptors<P>,
Option<ObservationSender>,
) -> Fut
+ Send
+ 'static,
Fut: Future<Output = Result<OutputOf<P>, P::Error>> + Send + 'static,
{
CallMachine::new(move |host| {
CallMachine::new(observers, move |host| {
Box::pin(async move {
match execute(request, host.services, host.hooks).await? {
match execute(request, host.services, host.interceptors, host.observers).await? {
CallOutput::Complete(response) => Ok(HostedCompletion::Complete(response)),
CallOutput::Stream { head, mut chunks } => {
if host.stream.open_stream(head).await? == Demand::Detached {
if host.stream.open_stream(head).await?.is_break() {
return Ok(HostedCompletion::Detached);
}
while let Some(chunk) = chunks.try_next().await? {
if host.stream.send_chunk(chunk).await? == Demand::Detached {
if host.stream.send_chunk(chunk).await?.is_break() {
return Ok(HostedCompletion::Detached);
}
}

View file

@ -1,79 +0,0 @@
use std::time::{SystemTime, UNIX_EPOCH};
use serde_json::Value;
/// Seconds since the Unix epoch, on one clock for every host.
pub fn epoch_seconds() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs_f64())
.unwrap_or(0.0)
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Timing {
pub start_time: f64,
pub end_time: f64,
}
/// The provider request as it is about to leave, offered to the host for rewriting.
#[derive(Clone, Debug, PartialEq)]
pub struct WireRequest {
pub url: String,
pub headers: Vec<(String, String)>,
pub body: Value,
}
/// What the route knows about the request it is sending, for a host that logs it. The
/// route owns these facts; a host reads them beside the wire request and never rewrites
/// them.
#[derive(Clone, Debug, PartialEq)]
pub struct RequestContext {
pub model: String,
pub custom_llm_provider: String,
/// The route's parameters before the provider transformation.
pub optional_params: Value,
/// Optional-param names that carry credentials and must be redacted when logged.
pub secret_fields: Vec<String>,
/// The credential the route resolved for the provider call.
pub api_key: Option<litellm_auth::SecretValue>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RawResponse {
pub body: String,
}
/// Whether a failure surfaced inside the call, including a host op the call asked for,
/// or in a host step around it (preparing the arguments, finalizing the response).
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum FailureOrigin {
Call,
Host,
}
/// What a machine reports while it runs.
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum MachineEvent {
ResponseReceived { raw: RawResponse },
}
/// What an in-process host observes: the machine's own events between the driver's
/// start and terminal ones.
#[derive(Clone, Debug, PartialEq)]
pub enum CallEvent {
Started {
start_time: f64,
},
Machine(MachineEvent),
Succeeded {
timing: Timing,
},
Failed {
timing: Timing,
origin: FailureOrigin,
},
Cancelled {
timing: Timing,
},
}

View file

@ -1,148 +1,75 @@
use std::future::Future;
use crate::{
interceptors::{RawResponse, RequestContext, WireRequest},
lifecycle::{CallEvent, Timing},
};
use crate::event::{MachineEvent, RequestContext, WireRequest};
pub trait HookRuntime {
type Context<'a>;
type Arguments;
type Response;
type Chunk;
type Error;
type Step<H, T>;
/// 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 observer(&self) -> Option<std::sync::Arc<dyn crate::lifecycle::CallObserver>> {
None
fn ready<H, T>(value: T) -> Self::Step<H, T>;
}
pub type RuntimeCallEvent<'a, R> =
CallEvent<&'a <R as HookRuntime>::Response, &'a <R as HookRuntime>::Error, &'a RawResponse>;
pub trait CallHooks<R: HookRuntime>: Sized {
fn prepare_arguments(
&mut self,
_runtime: R::Context<'_>,
arguments: R::Arguments,
_started_at: f64,
) -> Result<R::Step<Self, R::Arguments>, R::Error> {
Ok(R::ready(arguments))
}
fn arguments_prepared(
&mut self,
_runtime: R::Context<'_>,
_arguments: &R::Arguments,
) -> Result<(), R::Error> {
Ok(())
}
fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
) -> impl Future<Output = Result<WireRequest, 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_provider_request(
&self,
wire: WireRequest,
_: RequestContext,
) -> Result<WireRequest, E> {
Ok(wire)
&mut self,
_runtime: R::Context<'_>,
wire: Box<WireRequest>,
_context: &RequestContext,
) -> Result<R::Step<Self, Box<WireRequest>>, R::Error> {
Ok(R::ready(wire))
}
async fn on_event(&self, _: MachineEvent) -> Result<(), E> {
fn transform_response(
&mut self,
_runtime: R::Context<'_>,
response: R::Response,
_timing: Timing,
) -> Result<R::Step<Self, R::Response>, R::Error> {
Ok(R::ready(response))
}
fn on_event(
&mut self,
_runtime: R::Context<'_>,
_event: RuntimeCallEvent<'_, R>,
) -> Result<R::Step<Self, ()>, R::Error> {
Ok(R::ready(()))
}
fn on_stream_open(&mut self, _runtime: R::Context<'_>) -> Result<(), R::Error> {
Ok(())
}
fn on_stream_chunk(
&mut self,
_runtime: R::Context<'_>,
_chunk: &R::Chunk,
) -> Result<(), R::Error> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::convert::Infallible;
use serde_json::json;
use super::*;
use crate::protocol::HookRequest;
use crate::{
event::RawResponse,
machine::MachineFault,
machine::{CallMachine, Machine, MachineStep},
protocol::Protocol,
protocol::Suspension,
};
struct Unit;
#[derive(Clone, Debug)]
struct Fault;
impl Protocol for Unit {
type Response = (WireRequest, ());
type Error = Fault;
type Request = ();
type HostCall = Infallible;
type Chunk = Infallible;
type StreamHead = Infallible;
}
impl From<MachineFault> for Fault {
fn from(_: MachineFault) -> Self {
Fault
}
}
fn wire(url: &str) -> WireRequest {
WireRequest {
url: url.into(),
headers: Vec::new(),
body: json!({}),
}
}
fn context() -> RequestContext {
RequestContext {
model: "m".into(),
custom_llm_provider: "p".into(),
optional_params: json!({}),
secret_fields: Vec::new(),
api_key: None,
}
}
#[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_provider_request(
&channel.hooks,
wire("prepared"),
context(),
)
.await?;
RouteHooks::on_event(
&channel.hooks,
MachineEvent::ResponseReceived {
raw: RawResponse { body: "raw".into() },
},
)
.await?;
Ok((sent, ()))
})
});
let Ok(MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest {
wire,
reply,
..
}))) = machine.resume().await
else {
panic!("before_provider_request yields BeforeSend");
};
assert_eq!(wire.url, "prepared");
reply.send(WireRequest {
url: "rewritten".into(),
..*wire
});
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(());
let Ok(MachineStep::Complete((sent, ()))) = machine.resume().await else {
panic!("the call completes with the answers");
};
assert_eq!(sent.url, "rewritten");
}
#[rstest::rstest]
#[tokio::test]
async fn no_hooks_pass_the_wire_request_through() {
let sent = RouteHooks::<Fault>::before_provider_request(&(), wire("prepared"), context())
.await
.unwrap();
assert_eq!(sent.url, "prepared");
}
}

View file

@ -1,323 +0,0 @@
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

@ -0,0 +1,183 @@
use std::future::Future;
use serde_json::Value;
/// The provider request as it is about to leave, offered to the host for rewriting.
#[derive(Clone, Debug, PartialEq)]
pub struct WireRequest {
pub url: String,
pub headers: Vec<(String, String)>,
pub body: Value,
}
/// What the route knows about the request it is sending, for a host that logs it. The
/// route owns these facts; a host reads them beside the wire request and never rewrites
/// them.
#[derive(Clone, Debug, PartialEq)]
pub struct RequestContext {
pub model: String,
pub custom_llm_provider: String,
/// The route's parameters before the provider transformation.
pub optional_params: Value,
/// Optional-param names that carry credentials and must be redacted when logged.
pub secret_fields: Vec<String>,
/// The credential the route resolved for the provider call.
pub api_key: Option<litellm_auth::SecretValue>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RawResponse {
pub body: String,
}
pub trait Interceptors<E>: Send + Sync {
fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
) -> impl Future<Output = Result<WireRequest, E>> + Send;
fn after_provider_response(
&self,
raw: RawResponse,
) -> impl Future<Output = Result<(), E>> + Send;
}
impl<E, T: Interceptors<E> + ?Sized> Interceptors<E> for &T {
fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
) -> impl Future<Output = Result<WireRequest, E>> + Send {
(**self).before_provider_request(wire, context)
}
fn after_provider_response(
&self,
raw: RawResponse,
) -> impl Future<Output = Result<(), E>> + Send {
(**self).after_provider_response(raw)
}
}
impl<E> Interceptors<E> for () {
async fn before_provider_request(
&self,
wire: WireRequest,
_: RequestContext,
) -> Result<WireRequest, E> {
Ok(wire)
}
async fn after_provider_response(&self, _: RawResponse) -> Result<(), E> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::convert::Infallible;
use serde_json::json;
use super::*;
use crate::protocol::InterceptRequest;
use crate::{
machine::{CallMachine, Machine, MachineFault, MachineStep},
protocol::{HostRequest, Protocol},
};
struct Unit;
#[derive(Clone, Debug)]
struct Fault;
impl Protocol for Unit {
type Response = (WireRequest, ());
type Error = Fault;
type Request = ();
type HostCall = Infallible;
type Chunk = Infallible;
type StreamHead = Infallible;
}
impl From<MachineFault> for Fault {
fn from(_: MachineFault) -> Self {
Fault
}
}
fn wire(url: &str) -> WireRequest {
WireRequest {
url: url.into(),
headers: Vec::new(),
body: json!({}),
}
}
fn context() -> RequestContext {
RequestContext {
model: "m".into(),
custom_llm_provider: "p".into(),
optional_params: json!({}),
secret_fields: Vec::new(),
api_key: None,
}
}
#[rstest::rstest]
#[tokio::test]
async fn the_channel_yields_each_hook_as_its_op_and_returns_the_answer() {
let mut machine = CallMachine::<Unit>::new(None, |channel| {
Box::pin(async move {
let sent = Interceptors::before_provider_request(
&channel.interceptors,
wire("prepared"),
context(),
)
.await?;
Interceptors::after_provider_response(
&channel.interceptors,
RawResponse { body: "raw".into() },
)
.await?;
Ok((sent, ()))
})
});
let Ok(MachineStep::Suspended(HostRequest::Intercept(
InterceptRequest::BeforeProviderRequest { wire, reply, .. },
))) = machine.resume().await
else {
panic!("before_provider_request yields BeforeSend");
};
assert_eq!(wire.url, "prepared");
reply.send(WireRequest {
url: "rewritten".into(),
..*wire
});
let Ok(MachineStep::Suspended(HostRequest::Intercept(
InterceptRequest::AfterProviderResponse { raw, reply },
))) = machine.resume().await
else {
panic!("on_event yields Emit");
};
assert_eq!(raw.body, "raw");
reply.send(());
let Ok(MachineStep::Complete((sent, ()))) = machine.resume().await else {
panic!("the call completes with the answers");
};
assert_eq!(sent.url, "rewritten");
}
#[rstest::rstest]
#[tokio::test]
async fn no_hooks_pass_the_wire_request_through() {
let sent = Interceptors::<Fault>::before_provider_request(&(), wire("prepared"), context())
.await
.unwrap();
assert_eq!(sent.url, "prepared");
}
}

View file

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

View file

@ -1,29 +1,102 @@
use std::{future::Future, sync::Arc};
use crate::observation::ObservationSender;
use std::{
future::Future,
time::{SystemTime, UNIX_EPOCH},
};
use futures_util::TryStreamExt;
use crate::{
call::CallOutput,
event::{CallEvent, FailureOrigin, Timing, epoch_seconds},
};
use crate::{call::CallOutput, interceptors::RawResponse};
/// Seconds since the Unix epoch, on one clock for every host.
pub fn epoch_seconds() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs_f64())
.unwrap_or(0.0)
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Timing {
pub start_time: f64,
pub end_time: f64,
}
/// Whether a failure surfaced inside the call, including a host op the call asked for,
/// or in a host step around it (preparing the arguments, finalizing the response).
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum FailureOrigin {
Call,
Host,
}
#[derive(Clone, Debug, PartialEq)]
pub enum CallEvent<Response = (), Error = (), Raw = RawResponse> {
Started {
start_time: f64,
},
Execution(ExecutionEvent<Raw>),
Succeeded {
timing: Timing,
response: Response,
},
Failed {
timing: Timing,
origin: FailureOrigin,
error: Error,
},
Cancelled {
timing: Timing,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ExecutionEvent<Raw = RawResponse> {
ProviderResponseReceived { raw: Raw },
}
impl<Response, Error, Raw: std::borrow::Borrow<RawResponse>> CallEvent<Response, Error, Raw> {
pub fn snapshot(&self) -> CallEvent {
match self {
Self::Started { start_time } => CallEvent::Started {
start_time: *start_time,
},
Self::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived {
raw: raw.borrow().clone(),
})
}
Self::Succeeded { timing, .. } => CallEvent::Succeeded {
timing: *timing,
response: (),
},
Self::Failed { timing, origin, .. } => CallEvent::Failed {
timing: *timing,
origin: *origin,
error: (),
},
Self::Cancelled { timing } => CallEvent::Cancelled { timing: *timing },
}
}
}
pub trait CallObserver: Send + Sync {
fn observe(&self, event: CallEvent);
}
struct CallGuard {
observer: Option<Arc<dyn CallObserver>>,
observers: Option<ObservationSender>,
started_at: f64,
}
impl CallGuard {
fn new(observer: Arc<dyn CallObserver>) -> Self {
fn new(observers: ObservationSender) -> Self {
let started_at = epoch_seconds();
observer.observe(CallEvent::Started {
observers.emit(CallEvent::Started {
start_time: started_at,
});
Self {
observer: Some(observer),
observers: Some(observers),
started_at,
}
}
@ -36,15 +109,17 @@ impl CallGuard {
}
fn finish(mut self, failed: bool) {
if let Some(observer) = self.observer.take() {
observer.observe(if failed {
if let Some(observers) = self.observers.take() {
observers.emit(if failed {
CallEvent::Failed {
timing: self.timing(),
origin: FailureOrigin::Call,
error: (),
}
} else {
CallEvent::Succeeded {
timing: self.timing(),
response: (),
}
});
}
@ -53,8 +128,8 @@ impl CallGuard {
impl Drop for CallGuard {
fn drop(&mut self) {
if let Some(observer) = self.observer.take() {
observer.observe(CallEvent::Cancelled {
if let Some(observers) = self.observers.take() {
observers.emit(CallEvent::Cancelled {
timing: self.timing(),
});
}
@ -62,17 +137,17 @@ impl Drop for CallGuard {
}
pub async fn observe_call<R, H, C, E>(
observer: Option<Arc<dyn CallObserver>>,
observers: Option<ObservationSender>,
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 {
let Some(observers) = observers else {
return execute.await;
};
let guard = CallGuard::new(observer);
let guard = CallGuard::new(observers);
match execute.await {
Err(error) => {
guard.finish(true);
@ -108,13 +183,13 @@ where
}
pub async fn observe_unary<R, E>(
observer: Option<Arc<dyn CallObserver>>,
observers: Option<ObservationSender>,
execute: impl Future<Output = Result<R, E>>,
) -> Result<R, E> {
let Some(observer) = observer else {
let Some(observers) = observers else {
return execute.await;
};
let guard = CallGuard::new(observer);
let guard = CallGuard::new(observers);
let result = execute.await;
guard.finish(result.is_err());
result

View file

@ -0,0 +1,32 @@
# Resumable execution
`Machine` is the driver-facing contract for a resumable execution. `CallMachine` implements that contract with `litellm_coroutine::Coroutine`, holding the async execution of a route invocation. Resuming continues that same execution, which may request host services, invoke hooks, and deliver many stream chunks
## File ownership
| File | Responsibility |
| --- | --- |
| `mod.rs` | Module declarations and public exports |
| `contract.rs` | `Machine`, its step and interruption futures, and host failure values |
| `coroutine.rs` | `CallMachine`, coroutine state conversion, execution futures, and machine faults |
| `context.rs` | The route's services, hooks, and stream handles, sharing one coroutine channel |
Keep the contract independent of the coroutine implementation. Drivers and wrappers can implement `Machine` without constructing a coroutine. Keep existing public imports through `litellm_host::machine` stable when reorganizing private modules
Credential acquisition contracts and reusable adapters belong in `litellm-auth-types`. A route can use `TokenProviderHandle::from_callback` to request a credential through `HostServices::call`. Keep authentication policy and token-specific traits out of the machine layer
## Execution and replies
The coroutine polls the route future until it completes or yields a `HostRequest`. Each request carries a typed `Reply` that its driver must answer before resuming, or abandon when interrupting or dropping the execution. A pending network future is an ordinary async wait, not a host suspension
`CallContext` gives the route separate capabilities: `HostServices` requests host operations, `ChannelInterceptors` requests interception, `ObservationSender` publishes events, and `StreamSender` delivers stream values. Keep their yield-and-reply mechanics in `context.rs`. The actual service and hook implementations belong to the host. This follows the effect-handler pattern: the route requests an operation, the driver handles it, and the route continues with the reply. The suspended computation stays in the coroutine; `Reply` only supplies its result
Stream replies use `std::ops::ControlFlow<()>`. `ControlFlow::Continue(())` permits stream execution to continue. `ControlFlow::Break(())` tells it that the consumer stopped reading. Holding the reply applies backpressure until the consumer advances. Keep stream forwarding and the distinction between stream exhaustion and detachment in `crate::call::hosted_call`
`CallMachine::interrupt` cancels the coroutine and returns the supplied failure. Dropping `CallMachine` drops its execution future. Preserve both behaviors and do not spawn a producer task or poll stream chunks ahead of consumer demand
## Host responsibilities
Rust drivers live in `litellm-host-native` and `litellm-host-http`. The Python driver lives in `litellm-host-python` and awaits Python hooks in the caller's task. Keep runtime scheduling, encoding, terminal observation, and callback policy in those layers and their adapters
Public behavior tests belong in `crates/host/tests`; private behavior tests stay inline with their owning implementation. Test suspension answers, interruption, resource release, and stream demand through behavior, rather than asserting file layout

View file

@ -1,47 +0,0 @@
use std::sync::Arc;
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::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: HostServices<R>,
}
impl<R: Protocol> std::fmt::Debug for HostTokenProvider<R> {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("HostTokenProvider")
}
}
impl<R> HostTokenProvider<R>
where
R: TokenProtocol,
R::Error: From<MachineFault> + std::fmt::Display,
{
pub fn handle(channel: HostServices<R>) -> TokenProviderHandle {
TokenProviderHandle::new(Arc::new(Self { channel }))
}
}
impl<R> TokenProvider for HostTokenProvider<R>
where
R: TokenProtocol,
R::Error: From<MachineFault> + std::fmt::Display,
{
fn acquire(&self) -> TokenFuture<'_> {
Box::pin(async move {
self.channel
.call(R::acquire_token_op)
.await
.map_err(|error| Error::CredentialAcquisition(error.to_string().into()))
})
}
}

View file

@ -1,185 +0,0 @@
//! The one machine every route runs on: the route's provider future as a
//! [`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};
use super::{HostFailure, Interrupted, Machine, MachineStep, Step};
use crate::{
event::{MachineEvent, RequestContext, WireRequest},
protocol::{Demand, Protocol, Reply, Suspension},
};
/// The machine's own failures, distinct from anything the provider call reports.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MachineFault {
/// The host dropped an op's reply unanswered, or went away while the call waited.
Abandoned,
/// The host resumed the call out of turn.
Protocol(ResumeError),
}
pub type ExecuteFuture<R, C = <R as Protocol>::Response> =
Pin<Box<dyn Future<Output = Result<C, <R as Protocol>::Error>> + Send>>;
pub struct CallContext<P: Protocol> {
pub services: HostServices<P>,
pub hooks: ChannelHooks<P>,
pub stream: StreamSender<P>,
}
struct Channel<P: Protocol>(Co<Suspension<P>>);
impl<P: Protocol> Clone for Channel<P> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl<P: Protocol> Channel<P>
where
P::Error: From<MachineFault>,
{
async fn request_reply<A: Send>(
&self,
request: impl FnOnce(Reply<A>) -> Suspension<P> + Send,
) -> Result<A, P::Error> {
self.0
.yield_(request)
.await
.map_err(|_| MachineFault::Abandoned.into())
}
}
pub struct HostServices<P: Protocol>(Channel<P>);
impl<P: Protocol> Clone for HostServices<P> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl<P: Protocol> HostServices<P>
where
P::Error: From<MachineFault>,
{
pub async fn call<A: Send>(
&self,
request: impl FnOnce(Reply<A>) -> P::HostCall + Send,
) -> Result<A, P::Error> {
self.0
.request_reply(|reply| Suspension::HostCall(request(reply)))
.await
}
}
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, P::Error> {
self.0
.request_reply(|reply| {
Suspension::Hook(HookRequest::BeforeProviderRequest {
wire: Box::new(wire),
context: Box::new(context),
reply,
})
})
.await
}
async fn on_event(&self, event: MachineEvent) -> Result<(), P::Error> {
self.0
.request_reply(|reply| Suspension::Hook(HookRequest::Event(event, reply)))
.await
}
}
pub struct StreamSender<P: Protocol>(Channel<P>);
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
}
}
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(CallContext<R>) -> ExecuteFuture<R, C> + Send + 'static,
) -> Self {
Self {
coroutine: Coroutine::new(|co| {
let channel = Channel(co);
execute(CallContext {
services: HostServices(channel.clone()),
hooks: ChannelHooks(channel.clone()),
stream: StreamSender(channel),
})
}),
}
}
}
impl<R: Protocol, C: Send + 'static> Machine for CallMachine<R, C>
where
R::Error: From<MachineFault>,
{
type Protocol = R;
type Complete = C;
fn resume(&mut self) -> Step<'_, Self> {
Box::pin(async move {
match self
.coroutine
.resume()
.await
.map_err(MachineFault::Protocol)?
{
CoroutineState::Yielded(op) => Ok(MachineStep::Suspended(op)),
CoroutineState::Complete(outcome) => outcome.map(MachineStep::Complete),
}
})
}
fn interrupt(&mut self, failure: HostFailure<R::Error>) -> Interrupted<'_, Self> {
self.coroutine.cancel();
Box::pin(async move { Err(failure.into_error()) })
}
}

View file

@ -0,0 +1,130 @@
use crate::observation::ObservationSender;
use std::ops::ControlFlow;
use litellm_coroutine::Co;
use super::coroutine::MachineFault;
use crate::{
interceptors::{RawResponse, RequestContext, WireRequest},
protocol::{HostRequest, InterceptRequest, Protocol, Reply, StreamDelivery},
};
pub struct CallContext<P: Protocol> {
pub services: HostServices<P>,
pub interceptors: ChannelInterceptors<P>,
pub stream: StreamSender<P>,
pub observers: Option<ObservationSender>,
}
impl<P: Protocol> CallContext<P> {
pub(super) fn new(co: Co<HostRequest<P>>, observers: Option<ObservationSender>) -> Self {
let channel = Channel(co);
Self {
services: HostServices(channel.clone()),
interceptors: ChannelInterceptors(channel.clone()),
stream: StreamSender(channel),
observers,
}
}
}
struct Channel<P: Protocol>(Co<HostRequest<P>>);
impl<P: Protocol> Clone for Channel<P> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl<P: Protocol> Channel<P>
where
P::Error: From<MachineFault>,
{
async fn request_reply<A: Send>(
&self,
request: impl FnOnce(Reply<A>) -> HostRequest<P> + Send,
) -> Result<A, P::Error> {
self.0
.yield_(request)
.await
.map_err(|_| MachineFault::Abandoned.into())
}
}
pub struct HostServices<P: Protocol>(Channel<P>);
impl<P: Protocol> Clone for HostServices<P> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl<P: Protocol> HostServices<P>
where
P::Error: From<MachineFault>,
{
pub async fn call<A: Send>(
&self,
request: impl FnOnce(Reply<A>) -> P::HostCall + Send,
) -> Result<A, P::Error> {
self.0
.request_reply(|reply| HostRequest::HostCall(request(reply)))
.await
}
}
pub struct ChannelInterceptors<P: Protocol>(Channel<P>);
impl<P: Protocol> Clone for ChannelInterceptors<P> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl<P: Protocol> crate::interceptors::Interceptors<P::Error> for ChannelInterceptors<P>
where
P::Error: From<MachineFault>,
{
async fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, P::Error> {
self.0
.request_reply(|reply| {
HostRequest::Intercept(InterceptRequest::BeforeProviderRequest {
wire: Box::new(wire),
context: Box::new(context),
reply,
})
})
.await
}
async fn after_provider_response(&self, raw: RawResponse) -> Result<(), P::Error> {
self.0
.request_reply(|reply| {
HostRequest::Intercept(InterceptRequest::AfterProviderResponse { raw, reply })
})
.await
}
}
pub struct StreamSender<P: Protocol>(Channel<P>);
impl<P: Protocol> StreamSender<P>
where
P::Error: From<MachineFault>,
{
pub async fn open_stream(&self, head: P::StreamHead) -> Result<ControlFlow<()>, P::Error> {
self.0
.request_reply(|reply| HostRequest::Stream(StreamDelivery::Open(head, reply)))
.await
}
pub async fn send_chunk(&self, chunk: P::Chunk) -> Result<ControlFlow<()>, P::Error> {
self.0
.request_reply(|reply| HostRequest::Stream(StreamDelivery::Chunk(chunk, reply)))
.await
}
}

View file

@ -0,0 +1,60 @@
use std::{future::Future, pin::Pin};
use crate::protocol::{HostRequest, Protocol};
pub enum MachineStep<R: Protocol, C> {
Suspended(HostRequest<R>),
Complete(C),
}
pub type Step<'a, M> = Pin<
Box<
dyn Future<
Output = Result<
MachineStep<<M as Machine>::Protocol, <M as Machine>::Complete>,
<<M as Machine>::Protocol as Protocol>::Error,
>,
> + Send
+ 'a,
>,
>;
pub type Interrupted<'a, M> = Pin<
Box<
dyn Future<
Output = Result<
<M as Machine>::Complete,
<<M as Machine>::Protocol as Protocol>::Error,
>,
> + Send
+ 'a,
>,
>;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum HostFailure<E> {
Error(E),
Cancelled(E),
}
impl<E> HostFailure<E> {
pub fn into_error(self) -> E {
match self {
Self::Error(error) | Self::Cancelled(error) => error,
}
}
}
pub trait Machine: Send {
type Protocol: Protocol;
type Complete: Send + 'static;
fn resume(&mut self) -> Step<'_, Self>;
/// The host failed to perform the pending op, or the caller cancelled. The call
/// yields no further ops.
fn interrupt(
&mut self,
failure: HostFailure<<Self::Protocol as Protocol>::Error>,
) -> Interrupted<'_, Self>;
}

View file

@ -0,0 +1,69 @@
use crate::observation::ObservationSender;
use std::{future::Future, pin::Pin};
use litellm_coroutine::{Coroutine, CoroutineState, ResumeError};
use super::{
context::CallContext,
contract::{HostFailure, Interrupted, Machine, MachineStep, Step},
};
use crate::protocol::{HostRequest, Protocol};
/// The machine's own failures, distinct from anything the provider call reports.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MachineFault {
/// The host dropped an op's reply unanswered, or went away while the call waited.
Abandoned,
/// The host resumed the call out of turn.
Protocol(ResumeError),
}
pub type ExecuteFuture<R, C = <R as Protocol>::Response> =
Pin<Box<dyn Future<Output = Result<C, <R as Protocol>::Error>> + Send>>;
type CallCoroutine<R, C> = Coroutine<HostRequest<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(
observers: Option<ObservationSender>,
execute: impl FnOnce(CallContext<R>) -> ExecuteFuture<R, C> + Send + 'static,
) -> Self {
Self {
coroutine: Coroutine::new(move |co| execute(CallContext::new(co, observers))),
}
}
}
impl<R: Protocol, C: Send + 'static> Machine for CallMachine<R, C>
where
R::Error: From<MachineFault>,
{
type Protocol = R;
type Complete = C;
fn resume(&mut self) -> Step<'_, Self> {
Box::pin(async move {
match self
.coroutine
.resume()
.await
.map_err(MachineFault::Protocol)?
{
CoroutineState::Yielded(op) => Ok(MachineStep::Suspended(op)),
CoroutineState::Complete(outcome) => outcome.map(MachineStep::Complete),
}
})
}
fn interrupt(&mut self, failure: HostFailure<R::Error>) -> Interrupted<'_, Self> {
self.coroutine.cancel();
Box::pin(async move { Err(failure.into_error()) })
}
}

View file

@ -1,72 +1,7 @@
mod auth;
mod call_machine;
mod context;
mod contract;
mod coroutine;
use std::future::Future;
use std::pin::Pin;
pub use auth::{HostTokenProvider, TokenProtocol};
pub use call_machine::{
CallContext, CallMachine, ChannelHooks, ExecuteFuture, HostServices, MachineFault, StreamSender,
};
use crate::protocol::{Protocol, Suspension};
pub enum MachineStep<R: Protocol, C> {
Suspended(Suspension<R>),
Complete(C),
}
pub type Step<'a, M> = Pin<
Box<
dyn Future<
Output = Result<
MachineStep<<M as Machine>::Protocol, <M as Machine>::Complete>,
<<M as Machine>::Protocol as Protocol>::Error,
>,
> + Send
+ 'a,
>,
>;
pub type Interrupted<'a, M> = Pin<
Box<
dyn Future<
Output = Result<
<M as Machine>::Complete,
<<M as Machine>::Protocol as Protocol>::Error,
>,
> + Send
+ 'a,
>,
>;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum HostFailure<E> {
Error(E),
Cancelled(E),
}
impl<E> HostFailure<E> {
pub fn into_error(self) -> E {
match self {
Self::Error(error) | Self::Cancelled(error) => error,
}
}
}
/// A resumable call. Core implements it per route; a host drives it. Every suspension
/// point is an op the host performs and answers through the op's own reply before it
/// resumes the call again.
pub trait Machine: Send {
type Protocol: Protocol;
type Complete: Send + 'static;
fn resume(&mut self) -> Step<'_, Self>;
/// The host failed to perform the pending op, or the caller cancelled. The call
/// yields no further ops.
fn interrupt(
&mut self,
failure: HostFailure<<Self::Protocol as Protocol>::Error>,
) -> Interrupted<'_, Self>;
}
pub use context::{CallContext, ChannelInterceptors, HostServices, StreamSender};
pub use contract::{HostFailure, Interrupted, Machine, MachineStep, Step};
pub use coroutine::{CallMachine, ExecuteFuture, MachineFault};

View file

@ -0,0 +1,42 @@
use std::{
num::NonZeroUsize,
sync::{
Arc,
atomic::{AtomicU64, Ordering},
},
};
use tokio::sync::mpsc;
use crate::lifecycle::CallEvent;
#[derive(Clone)]
pub struct ObservationSender {
sender: mpsc::Sender<CallEvent>,
dropped: Arc<AtomicU64>,
}
impl ObservationSender {
pub fn emit(&self, event: CallEvent) {
if self.sender.try_send(event).is_err() {
self.dropped.fetch_add(1, Ordering::Relaxed);
}
}
pub fn dropped_events(&self) -> u64 {
self.dropped.load(Ordering::Relaxed)
}
}
pub fn observation_channel(
capacity: NonZeroUsize,
) -> (ObservationSender, mpsc::Receiver<CallEvent>) {
let (sender, receiver) = mpsc::channel(capacity.get());
(
ObservationSender {
sender,
dropped: Arc::new(AtomicU64::new(0)),
},
receiver,
)
}

View file

@ -1,6 +1,8 @@
use std::ops::ControlFlow;
pub use litellm_coroutine::{Abandoned, Answer, Reply, reply};
use crate::event::{MachineEvent, RequestContext, WireRequest};
use crate::interceptors::{RawResponse, RequestContext, WireRequest};
pub trait Protocol: Send + Sync + 'static {
type Request: Send + 'static;
@ -11,28 +13,25 @@ pub trait Protocol: Send + Sync + 'static {
type StreamHead: Send + 'static;
}
pub enum Suspension<P: Protocol> {
pub enum HostRequest<P: Protocol> {
HostCall(P::HostCall),
Hook(HookRequest),
Intercept(InterceptRequest),
Stream(StreamDelivery<P>),
}
pub enum HookRequest {
pub enum InterceptRequest {
BeforeProviderRequest {
wire: Box<WireRequest>,
context: Box<RequestContext>,
reply: Reply<WireRequest>,
},
Event(MachineEvent, Reply<()>),
AfterProviderResponse {
raw: RawResponse,
reply: 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,
Open(P::StreamHead, Reply<ControlFlow<()>>),
Chunk(P::Chunk, Reply<ControlFlow<()>>),
}

View file

@ -1,6 +1,7 @@
use litellm_host::protocol::StreamDelivery;
use std::{
convert::Infallible,
ops::ControlFlow,
sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
@ -10,10 +11,9 @@ use std::{
use futures_util::{StreamExt, stream};
use litellm_host::{
call::{CallOutput, HostedCompletion, hosted_call},
event::CallEvent,
lifecycle::{CallObserver, observe_call, observe_unary},
lifecycle::{CallEvent, CallObserver, observe_call, observe_unary},
machine::{Machine, MachineFault, MachineStep},
protocol::{Demand, Protocol, Suspension},
protocol::{HostRequest, Protocol},
};
use rstest::{fixture, rstest};
@ -49,18 +49,19 @@ async fn delivery_obeys_demand_and_distinguishes_detachment(
) {
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);
let mut machine =
hosted_call::<TestProtocol, _, _>(3, None, move |count, _, _, _observations| 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,
})
.boxed();
Ok(CallOutput::Stream {
head: "headers",
chunks,
})
});
let MachineStep::Suspended(Suspension::Stream(StreamDelivery::Open(head, reply))) =
});
let MachineStep::Suspended(HostRequest::Stream(StreamDelivery::Open(head, reply))) =
machine.resume().await.unwrap()
else {
panic!()
@ -68,20 +69,20 @@ async fn delivery_obeys_demand_and_distinguishes_detachment(
assert_eq!(head, "headers");
assert_eq!(polls.load(Ordering::SeqCst), 0);
reply.send(if detach_after == Some(0) {
Demand::Detached
ControlFlow::Break(())
} else {
Demand::More
ControlFlow::Continue(())
});
let mut delivered = Vec::new();
let completed = loop {
match machine.resume().await.unwrap() {
MachineStep::Suspended(Suspension::Stream(StreamDelivery::Chunk(chunk, reply))) => {
MachineStep::Suspended(HostRequest::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
ControlFlow::Break(())
} else {
Demand::More
ControlFlow::Continue(())
});
}
MachineStep::Complete(result) => break result,
@ -94,11 +95,40 @@ async fn delivery_obeys_demand_and_distinguishes_detachment(
}
#[derive(Default)]
struct Observer(Mutex<Vec<CallEvent>>);
struct Observer(Observations);
struct Observations {
sender: litellm_host::observation::ObservationSender,
receiver: Mutex<tokio::sync::mpsc::Receiver<CallEvent>>,
recorded: Mutex<Vec<CallEvent>>,
}
impl Default for Observations {
fn default() -> Self {
let (sender, receiver) = litellm_host::observation::observation_channel(
std::num::NonZeroUsize::new(128).unwrap(),
);
Self {
sender,
receiver: Mutex::new(receiver),
recorded: Mutex::new(Vec::new()),
}
}
}
impl Observations {
fn lock(&self) -> std::sync::LockResult<std::sync::MutexGuard<'_, Vec<CallEvent>>> {
let mut events = self.recorded.lock()?;
let mut receiver = self.receiver.lock().unwrap();
while let Ok(event) = receiver.try_recv() {
events.push(event);
}
Ok(events)
}
}
impl CallObserver for Observer {
fn observe(&self, event: CallEvent) {
self.0.lock().unwrap().push(event);
self.0.sender.emit(event);
}
}
@ -116,7 +146,7 @@ type Output = CallOutput<(), (), usize, &'static str>;
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,
observe_unary(Some(observer.0.sender.clone()), async { expected }).await,
expected
);
let events = observer.0.lock().unwrap();
@ -132,7 +162,7 @@ async fn unary_calls_emit_one_terminal_event(observer: Arc<Observer>, #[case] fa
#[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 {
let output = observe_call(Some(observer.0.sender.clone()), async {
Ok::<Output, _>(CallOutput::Stream { head: (), chunks })
})
.await
@ -161,7 +191,7 @@ async fn streams_finish_only_when_consumed(observer: Arc<Observer>, #[case] fail
#[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 {
let output = observe_call(Some(observer.0.sender.clone()), async {
Ok::<Output, _>(CallOutput::Stream { head: (), chunks })
})
.await
@ -176,7 +206,7 @@ async fn dropping_a_stream_cancels_without_success(observer: Arc<Observer>) {
#[tokio::test]
async fn cancelling_provider_execution_releases_the_lifecycle(observer: Arc<Observer>) {
let mut call = Box::pin(observe_unary(
Some(observer.clone()),
Some(observer.0.sender.clone()),
std::future::pending::<Result<(), ()>>(),
));
assert!(futures_util::poll!(&mut call).is_pending());
@ -185,40 +215,3 @@ async fn cancelling_provider_execution_releases_the_lifecycle(observer: Arc<Obse
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

@ -0,0 +1,175 @@
use std::{cell::Cell, convert::Infallible, num::NonZeroUsize, rc::Rc};
use litellm_host::{
interceptors::RawResponse,
lifecycle::{CallEvent, ExecutionEvent, FailureOrigin, Timing, observe_unary},
machine::{CallMachine, Machine, MachineFault, MachineStep},
observation::observation_channel,
protocol::Protocol,
};
use rstest::{fixture, rstest};
struct TestProtocol;
impl Protocol for TestProtocol {
type Request = ();
type Response = usize;
type Error = MachineFault;
type HostCall = Infallible;
type Chunk = Infallible;
type StreamHead = Infallible;
}
#[fixture]
fn capacity() -> NonZeroUsize {
NonZeroUsize::new(2).unwrap()
}
#[rstest]
#[case::success(false)]
#[case::failure(true)]
fn snapshots_do_not_retain_runtime_objects(capacity: NonZeroUsize, #[case] failed: bool) {
let payload = Rc::new(Cell::new(7));
let retained = Rc::downgrade(&payload);
let timing = Timing {
start_time: 11.0,
end_time: 19.0,
};
let event: CallEvent<Rc<Cell<u8>>, Rc<Cell<u8>>> = if failed {
CallEvent::Failed {
timing,
origin: FailureOrigin::Host,
error: payload,
}
} else {
CallEvent::Succeeded {
timing,
response: payload,
}
};
let (sender, mut receiver) = observation_channel(capacity);
sender.emit(event.snapshot());
match &event {
CallEvent::Succeeded { response, .. } => response.set(9),
CallEvent::Failed { error, .. } => error.set(9),
_ => unreachable!(),
}
assert_eq!(retained.upgrade().unwrap().get(), 9);
drop(event);
assert!(retained.upgrade().is_none());
let expected = if failed {
CallEvent::Failed {
timing,
origin: FailureOrigin::Host,
error: (),
}
} else {
CallEvent::Succeeded {
timing,
response: (),
}
};
assert_eq!(receiver.try_recv().unwrap(), expected);
}
#[rstest]
fn provider_snapshots_own_the_response_body(capacity: NonZeroUsize) {
let mut raw = RawResponse {
body: "provider response".into(),
};
let event: CallEvent<(), (), &RawResponse> =
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw: &raw });
let expected = raw.clone();
let (sender, mut receiver) = observation_channel(capacity);
sender.emit(event.snapshot());
raw.body.clear();
assert_eq!(
receiver.try_recv().unwrap(),
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw: expected })
);
}
#[rstest]
#[tokio::test]
async fn observations_do_not_suspend_the_machine(capacity: NonZeroUsize) {
let (sender, mut receiver) = observation_channel(capacity);
let mut machine = CallMachine::<TestProtocol>::new(Some(sender), |host| {
Box::pin(async move {
host.observers
.unwrap()
.emit(CallEvent::Started { start_time: 1.0 });
Ok(42)
})
});
assert!(matches!(
machine.resume().await,
Ok(MachineStep::Complete(42))
));
assert!(matches!(
receiver.recv().await,
Some(CallEvent::Started { start_time: 1.0 })
));
assert_eq!(receiver.recv().await, None);
}
#[rstest]
#[case::full(false)]
#[case::closed(true)]
#[tokio::test]
async fn unavailable_observers_do_not_change_the_call_outcome(
capacity: NonZeroUsize,
#[case] closed: bool,
) {
let (sender, mut receiver) = observation_channel(capacity);
sender.emit(CallEvent::Started { start_time: 1.0 });
sender.emit(CallEvent::Started { start_time: 2.0 });
if closed {
receiver.close();
}
let outcome = observe_unary(Some(sender.clone()), async {
Err::<(), _>("provider failed")
})
.await;
assert_eq!(outcome, Err("provider failed"));
assert_eq!(sender.dropped_events(), 2);
assert_eq!(
receiver.try_recv().unwrap(),
CallEvent::Started { start_time: 1.0 }
);
assert_eq!(
receiver.try_recv().unwrap(),
CallEvent::Started { start_time: 2.0 }
);
assert!(receiver.try_recv().is_err());
if !closed {
sender.emit(CallEvent::Started { start_time: 3.0 });
assert_eq!(
receiver.try_recv().unwrap(),
CallEvent::Started { start_time: 3.0 }
);
assert_eq!(sender.dropped_events(), 2);
}
}
#[rstest]
#[tokio::test]
async fn the_receiver_drains_after_all_publishers_are_dropped(capacity: NonZeroUsize) {
let (sender, mut receiver) = observation_channel(capacity);
let other = sender.clone();
sender.emit(CallEvent::Started { start_time: 1.0 });
other.emit(CallEvent::Started { start_time: 2.0 });
other.emit(CallEvent::Started { start_time: 3.0 });
assert_eq!(sender.dropped_events(), 1);
assert_eq!(other.dropped_events(), 1);
drop(sender);
drop(other);
assert_eq!(
receiver.recv().await,
Some(CallEvent::Started { start_time: 1.0 })
);
assert_eq!(
receiver.recv().await,
Some(CallEvent::Started { start_time: 2.0 })
);
assert_eq!(receiver.recv().await, None);
}

View file

@ -3,7 +3,7 @@ use std::sync::Arc;
use bytes::{Bytes, BytesMut};
use futures_util::future::BoxFuture;
use litellm_auth::AuthServices;
use litellm_host::event::WireRequest;
use litellm_host::interceptors::WireRequest;
use litellm_http::{
Client, ClientVariant, HttpClientConfig, HttpClientPool,
media::{MediaFetcher, UrlPolicy},

View file

@ -1,7 +1,22 @@
- Target invariants, not completion claims; these supersede the crate guidance below where they conflict
## Boundary migration
Keep domain composition here and execution mechanics in `litellm-host-python`. A helper does not belong in the runtime adapter merely because it uses PyO3. LiteLLM argument rules, provider defaults, public responses, public exception policy and cache or secret-manager compatibility remain product responsibilities
The bridge supplies the Python lifecycle binding and public stream construction to the runtime adapter. Preserve the single inline coroutine driver; removing the host's hardcoded import must not introduce another driver or a separate asyncio task for caller hooks
`src/callable.rs::wrap_failure` owns callable exception policy here. Resolved-Future construction uses `litellm-host-python::ready_future`, passing an already constructed Python value. Keep cache-specific serialization and disabled-cache results here
Implement the migration in separate steps that preserve public API contracts: first defer route resource setup until prepared arguments and preflight are available, then supply the lifecycle binding and separate public stream construction, then relocate the two helpers. Change the host interface and its consumers together in each step. The `native.rs` rename is optional and comes last
Each step needs focused regression tests in the owning crate and Python integration coverage where the public contract crosses crates. Verify deferred setup and setup failure ordering, caller task and context identity, sync and async streams, exception provenance, cancellation and GC. Use a fresh installed extension to verify Python behavior and update `_native.pyi` when public signatures change. Do not treat these instructions as evidence that the migration is complete
## Existing bridge invariants
- Keep this crate the product-specific PyO3 consumer of `litellm-host-python`
- Own registration, input projection, the route host and the caller callables it answers operations with (file readers, token providers), public response/error construction and the per-call composition of machine, route host and callback contract
- Legacy callback sharing (the caller's args, kwargs and request object, body/header roots, re-aliasing unchanged body keys) lives in `litellm-callbacks-legacy-python` behind `PublicCall` and `run_legacy_call`; the bridge hands the public call over and keeps no copy
- Legacy callback sharing (the caller's args, kwargs and request object, body/header roots, re-aliasing unchanged body keys) lives in `litellm-callbacks-legacy-python` behind `PublicCall` and `LegacyLogging`; the bridge hands the public call over and keeps no copy
- Value-oriented execution, sync waiting, nested-runtime checks, signal polling and panic containment live in `litellm-host-python`; native async work uses `pyo3-async-runtimes`, Serde output uses `Pythonized<T>`
- Core owns typed native state, the route machine, provider preparation/I/O and normalization; the host driver owns terminal events; the legacy adapter in `litellm-callbacks-legacy-python` owns `Logging` dispatch policy
- Python, Rust SDK and gateway use one lifecycle-bearing core route entrypoint; provider helpers stay private, never bridge-accessible transport drivers

View file

@ -0,0 +1,9 @@
# Cache boundary
This folder owns Python cache API compatibility: argument projection, facade identity, public result construction, Python embedding calls and per-operation composition of native backends. Cache algorithms, storage protocols and response-cache semantics belong to their cache crates
`SemanticExecution` belongs here because its steps select cache operations and invoke the Python embedder. Use the shared `Execution` handle and inline lifecycle driver; do not duplicate coroutine state validation, runtime waiting or GIL machinery. Python embedding awaits stay in the caller's task, and cancellation must prevent later backend or batch operations from starting
Resolved asyncio Future construction is generic host machinery. Use `litellm-host-python::ready_future` with an already constructed Python value. Keep cache-specific conversion and disabled-cache return values here. Preserve the Future-returning API and running-loop requirement
Tests for Future mechanics belong in `host-python`; tests for disabled-cache values, embedding failure policy, batch sequencing and cancellation belong with this cache adapter. Assert public behavior rather than the location or name of a helper

View file

@ -1,4 +1,4 @@
use litellm_host_python::to_py;
use litellm_host_python::{ready_future, to_py};
use pyo3::prelude::*;
pub(super) fn ready_none(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
@ -9,10 +9,5 @@ pub(super) fn ready_value<'py, T: serde::Serialize>(
py: Python<'py>,
value: &T,
) -> PyResult<Bound<'py, PyAny>> {
let future = py
.import("asyncio")?
.call_method0("get_running_loop")?
.call_method0("create_future")?;
future.call_method1("set_result", (to_py(py, value)?,))?;
Ok(future)
ready_future(py, to_py(py, value)?.bind(py))
}

View file

@ -207,5 +207,5 @@ impl ExecutionBody for SemanticExecution {
}
pub(super) fn drive(py: Python<'_>, body: SemanticExecution) -> PyResult<Bound<'_, PyAny>> {
Execution::new(body).into_coroutine(py)
Execution::new(body, crate::lifecycle::binding).into_coroutine(py)
}

View file

@ -12,7 +12,7 @@ use pyo3::types::PyString;
/// `__context__`, with the message rendered by Python so the exception's own `__format__`
/// is honored. A `__format__` that raises surfaces as that failure instead, with the
/// original attached as its context.
pub fn wrap_failure<T>(py: Python<'_>, template: &str, result: PyResult<T>) -> PyResult<T> {
pub(crate) fn wrap_failure<T>(py: Python<'_>, template: &str, result: PyResult<T>) -> PyResult<T> {
result.map_err(|error| {
if error.is_instance_of::<PyTypeError>(py) || !error.is_instance_of::<PyException>(py) {
return error;
@ -48,9 +48,9 @@ mod tests {
Err(PyErr::from_value(error.clone()))
}
#[test]
#[rstest::rstest]
fn only_ordinary_exceptions_are_reported_under_the_template() {
crate::initialize_python();
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
@ -87,9 +87,9 @@ abort = KeyboardInterrupt('cancelled')
});
}
#[test]
#[rstest::rstest]
fn a_raising_format_surfaces_instead_of_the_report_and_keeps_the_original_as_context() {
crate::initialize_python();
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
@ -113,9 +113,9 @@ original = Unformattable('cannot render')
});
}
#[test]
#[rstest::rstest]
fn successful_results_pass_through_untouched() {
crate::initialize_python();
Python::initialize();
Python::attach(|py| {
assert_eq!(wrap_failure(py, TEMPLATE, Ok(7)).unwrap(), 7);
});

View file

@ -267,6 +267,11 @@ mod tests {
py.eval(&CString::new(source).unwrap(), None, None).unwrap()
}
fn with_python(f: impl for<'py> FnOnce(Python<'py>)) {
Python::initialize();
Python::attach(f);
}
#[rstest]
#[case("None", false, false)]
#[case("False", false, false)]
@ -284,8 +289,7 @@ mod tests {
#[case] truth: bool,
#[case] exact: bool,
) {
Python::initialize();
Python::attach(|py| {
with_python(|py| {
let value = evaluate(py, source);
let field = Field::new("test", "flag", value.clone());
assert_eq!(field.truthy().unwrap(), truth);
@ -323,8 +327,7 @@ mod tests {
#[case] fallback: Result<Option<&str>, ()>,
#[case] tuning: Result<Option<&str>, ()>,
) {
Python::initialize();
Python::attach(|py| {
with_python(|py| {
let field = Field::new("test", "string", evaluate(py, source));
let owned =
|expected: Result<Option<&str>, ()>| expected.map(|value| value.map(str::to_owned));
@ -351,8 +354,7 @@ mod tests {
#[case] source: &str,
#[case] expected: Option<bool>,
) {
Python::initialize();
Python::attach(|py| {
with_python(|py| {
assert_eq!(
Field::new("test", "flag", evaluate(py, source))
.str_bool()
@ -364,8 +366,7 @@ mod tests {
#[test]
fn protocol_errors_preserve_exception_identity_traceback_cause_and_context() {
Python::initialize();
Python::attach(|py| {
with_python(|py| {
let locals = PyDict::new(py);
py.run(
c"
@ -442,8 +443,7 @@ descriptor = Descriptor()
#[test]
fn identity_and_string_contents_do_not_invoke_unrelated_protocols() {
Python::initialize();
Python::attach(|py| {
with_python(|py| {
let locals = PyDict::new(py);
py.run(
c"
@ -476,8 +476,7 @@ text = Text(' False ')
#[test]
fn missing_snapshot_fields_and_descriptor_attribute_errors_are_distinct() {
Python::initialize();
Python::attach(|py| {
with_python(|py| {
let locals = PyDict::new(py);
py.run(
c"
@ -521,8 +520,7 @@ intercepted = Intercepted()
#[test]
fn configuration_errors_name_fields_without_exposing_values() {
Python::initialize();
Python::attach(|py| {
with_python(|py| {
for source in [
"{'secret': 'do-not-print'}",
"['host.test', {'secret': 'do-not-print'}]",

View file

@ -1,8 +1,8 @@
//! Credentials the caller supplies as Python callables, projected out of a route's
//! keyword arguments and acquired on the host's own thread when the call asks for one.
use crate::callable::wrap_failure;
use litellm_auth::{ResolvedCredential, SecretValue};
use litellm_host_python::wrap_failure;
use pyo3::{
exceptions::PyTypeError,
gc::{PyTraverseError, PyVisit},

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