diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index dd773528622..0350aa2f24a 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -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" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index d6b7dcac433..32919e23927 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -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" } diff --git a/litellm-rust/crates/auth-types/src/token.rs b/litellm-rust/crates/auth-types/src/token.rs index 4175641ce10..a3126052065 100644 --- a/litellm-rust/crates/auth-types/src/token.rs +++ b/litellm-rust/crates/auth-types/src/token.rs @@ -39,7 +39,33 @@ impl TokenProviderHandle { Self(caller) } + pub fn from_callback(acquire: F) -> Self + where + F: Fn() -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + Self::new(Arc::new(CallbackTokenProvider(acquire))) + } + pub async fn acquire(&self) -> Result { self.0.acquire().await } } + +struct CallbackTokenProvider(F); + +impl std::fmt::Debug for CallbackTokenProvider { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("CallbackTokenProvider") + } +} + +impl TokenProvider for CallbackTokenProvider +where + F: Fn() -> Fut + Send + Sync, + Fut: Future> + Send + 'static, +{ + fn acquire(&self) -> TokenFuture<'_> { + Box::pin((self.0)()) + } +} diff --git a/litellm-rust/crates/auth-types/tests/token.rs b/litellm-rust/crates/auth-types/tests/token.rs new file mode 100644 index 00000000000..d7b23531596 --- /dev/null +++ b/litellm-rust/crates/auth-types/tests/token.rs @@ -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::()) + .unwrap(); + assert_eq!(source.kind(), std::io::ErrorKind::PermissionDenied); +} + +struct Release(Arc); + +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)); +} diff --git a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md index 56ef51e4758..e76a15099dc 100644 --- a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md +++ b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md @@ -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 diff --git a/litellm-rust/crates/callbacks-legacy-python/Cargo.toml b/litellm-rust/crates/callbacks-legacy-python/Cargo.toml index ed5e0fb9691..8ee795092b4 100644 --- a/litellm-rust/crates/callbacks-legacy-python/Cargo.toml +++ b/litellm-rust/crates/callbacks-legacy-python/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-types.workspace = true litellm-host.workspace = true litellm-host-python.workspace = true diff --git a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs index 1fa7f3dcdd4..00d92168285 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs @@ -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, -} - -/// 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, start: Py, @@ -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 { + 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) { + 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>> { 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> { - 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, @@ -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, @@ -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, @@ -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> { - 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> { + 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, + ) -> PyResult> { + 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> { + 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) -> PyResult<()> { + pub(crate) fn stream_chunk(&mut self, py: Python<'_>, chunk: &Py) -> 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::() + .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::() + .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, &[])] diff --git a/litellm-rust/crates/callbacks-legacy-python/src/call.rs b/litellm-rust/crates/callbacks-legacy-python/src/call.rs index 52b64d8b181..8e7645cae7b 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/call.rs @@ -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, kwargs: Py, @@ -34,6 +31,10 @@ impl PublicCall { }) } + pub fn arguments(&self, py: Python<'_>) -> Py { + self.kwargs.clone_ref(py) + } + pub(crate) fn args(&self) -> &Py { &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( - py: Python<'_>, - surface: LegacySurface, - call: PublicCall, - start: impl FnOnce(::Request) -> M + Send + Sync + 'static, - host: H, - preflight: Preflight, - asynchronous: bool, -) -> PyResult> -where - H: PythonBinding + PythonHostCalls + 'static, - M: Machine + 'static, - M::Complete: Into::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::*; diff --git a/litellm-rust/crates/callbacks-legacy-python/src/callbacks.rs b/litellm-rust/crates/callbacks-legacy-python/src/callbacks.rs index 7caa787dd9d..f185f21a9c3 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/callbacks.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/callbacks.rs @@ -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}; diff --git a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs index a5d5c6762c2..bce186380b8 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs @@ -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; diff --git a/litellm-rust/crates/callbacks-legacy-python/src/mapping.rs b/litellm-rust/crates/callbacks-legacy-python/src/mapping.rs new file mode 100644 index 00000000000..8321e63e196 --- /dev/null +++ b/litellm-rust/crates/callbacks-legacy-python/src/mapping.rs @@ -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 { + boundary: CallBoundary, + invoke: H, + callbacks: &'static [&'static str], +} + +impl Binding { + fn mappings(&self) -> impl Iterator { + self.callbacks.iter().map(|callback| CallbackMapping { + callback, + dispatch: Dispatch::Call(self.boundary), + }) + } +} + +type Step = PyResult>; +type Prepare = fn(&mut LegacyLogging, Python<'_>, Py, f64) -> Step>; +type Before = + fn(&mut LegacyLogging, Python<'_>, Box, &RequestContext) -> Step>; +type After = fn(&mut LegacyLogging, Python<'_>, &RawResponse) -> Step<()>; +type Transform = fn(&mut LegacyLogging, Python<'_>, Py, Timing) -> Step>; +type Success = fn(&mut LegacyLogging, Python<'_>, Timing, &Py) -> 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) -> PyResult<()>; + +const PREPARE: Binding = Binding { + boundary: CallBoundary::PrepareArguments, + invoke: LegacyLogging::prepare_call, + callbacks: &["async_pre_call_deployment_hook"], +}; + +const BEFORE: Binding = Binding { + boundary: CallBoundary::BeforeProviderRequest, + invoke: LegacyLogging::pre_call, + callbacks: &["log_pre_api_call", "log_input_event"], +}; + +const AFTER: Binding = Binding { + boundary: CallBoundary::AfterProviderResponse, + invoke: LegacyLogging::post_call, + callbacks: &["log_post_api_call"], +}; + +const TRANSFORM: Binding = Binding { + boundary: CallBoundary::TransformResponse, + invoke: LegacyLogging::transform_public_response, + callbacks: &["async_post_call_success_deployment_hook"], +}; + +const SUCCESS: Binding = 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 = 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 = Binding { + boundary: CallBoundary::StreamOpened, + invoke: LegacyLogging::stream_opened, + callbacks: &[], +}; + +const CHUNK: Binding = Binding { + boundary: CallBoundary::StreamChunk, + invoke: LegacyLogging::stream_chunk, + callbacks: &[], +}; + +pub fn callback_mappings() -> impl Iterator { + 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 for LegacyLogging { + fn prepare_arguments( + &mut self, + py: Python<'_>, + arguments: Py, + started_at: f64, + ) -> Step> { + (PREPARE.invoke)(self, py, arguments, started_at) + } + + fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py) -> PyResult<()> { + self.adopt_arguments(py, arguments); + Ok(()) + } + + fn before_provider_request( + &mut self, + py: Python<'_>, + wire: Box, + context: &RequestContext, + ) -> Step> { + (BEFORE.invoke)(self, py, wire, context) + } + + fn transform_response( + &mut self, + py: Python<'_>, + response: Py, + timing: Timing, + ) -> Step> { + (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) -> 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", + ], +} diff --git a/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs index e7973a7e1a0..39a879f9ff6 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs @@ -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::().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) } diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 0a10dac1bdd..f316e6f7799 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -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//transformation.rs`, `//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 diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 365f806142f..023267d56ef 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index 2d4be71463c..dd54c3057f1 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -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, + interceptors: &impl Interceptors, + observers: Option<&ObservationSender>, ) -> Result { 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>, } - impl RouteHooks for RecordingHooks { + impl Interceptors 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] diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index ed3909870fa..9e350e757d3 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -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, + interceptors: &impl litellm_host::interceptors::Interceptors, + observers: Option, ) -> Result { - 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, + interceptors: &impl litellm_host::interceptors::Interceptors, + observers: Option<&ObservationSender>, ) -> Result { crate::diagnostic::unary(async { let resolved = resolve_request(request)?; @@ -64,7 +71,13 @@ impl ChatCompletionsRoute { let execute: futures_util::future::BoxFuture< '_, Result, - > = Box::pin(handler::execute(&self.http, &self.auth, prepared, hooks)); + > = Box::pin(handler::execute( + &self.http, + &self.auth, + prepared, + interceptors, + observers, + )); execute.await }) .await diff --git a/litellm-rust/crates/core/src/chat_completions/route.rs b/litellm-rust/crates/core/src/chat_completions/route.rs index 3f10a588692..0816850bcc1 100644 --- a/litellm-rust/crates/core/src/chat_completions/route.rs +++ b/litellm-rust/crates/core/src/chat_completions/route.rs @@ -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 { + pub fn machine( + self, + call: ChatCompletionsCall, + observers: Option, + ) -> HostedMachine { 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) }, ) } diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index d015c238043..d061f456b2a 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -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, + interceptors: &impl Interceptors, + observers: Option<&ObservationSender>, ) -> Result { 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) diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index f74f2c54771..07f1fff8d45 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -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, + interceptors: &impl litellm_host::interceptors::Interceptors, + observers: Option, ) -> Result { - 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, + interceptors: &impl litellm_host::interceptors::Interceptors, + observers: Option<&ObservationSender>, ) -> Result { 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> = - Box::pin(handler::execute(&self.http, &self.auth, request, hooks)); + Box::pin(handler::execute( + &self.http, + &self.auth, + request, + interceptors, + observers, + )); execute.await }) .await diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 5a1880d0a99..954c31be9d8 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -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; 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, + ) -> MessagesMachine { + hosted_call( + request, + observers, + move |call, _, interceptors, observers| async move { + self.run(call, &interceptors, observers.as_ref()).await + }, + ) } } diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index f39a8379a3c..df1fd1cda92 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -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, + interceptors: &impl Interceptors, + observers: Option, ) -> Result { - 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, + interceptors: &impl Interceptors, + observers: Option<&ObservationSender>, ) -> Result { 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 }) diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index 1c3705a58ac..7aed179e07a 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -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, + host: &impl Interceptors, caller_document: bool, + observers: Option<&ObservationSender>, ) -> Result { 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> CallHooks for OcrCallHooks<'_, H> { +impl> CallHooks for OcrCallHooks<'_, H> { fn before_provider_request( &self, wire: WireRequest, ) -> BoxFuture<'_, Result> { 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)) } } diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index cf2bae1ab9f..c2b401a7ab9 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -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 for NoHooks { diff --git a/litellm-rust/crates/core/src/ocr/provider_config.rs b/litellm-rust/crates/core/src/ocr/provider_config.rs index 3b4e9d88da4..1e27b83c4c5 100644 --- a/litellm-rust/crates/core/src/ocr/provider_config.rs +++ b/litellm-rust/crates/core/src/ocr/provider_config.rs @@ -133,9 +133,9 @@ impl OcrConfigKind { self, client: &OcrClient, request: &PreparedOcrRequest, - hooks: &dyn CallHooks, + interceptors: &dyn CallHooks, ) -> Result { - with_config!(self, config => handler::ocr(&config, client, request, hooks).await) + with_config!(self, config => handler::ocr(&config, client, request, interceptors).await) } } diff --git a/litellm-rust/crates/core/src/ocr/route.rs b/litellm-rust/crates/core/src/ocr/route.rs index 4e96fd2ebd9..b576585049a 100644 --- a/litellm-rust/crates/core/src/ocr/route.rs +++ b/litellm-rust/crates/core/src/ocr/route.rs @@ -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) -> OcrOp { - OcrOp::AcquireAzureAdToken(reply) - } -} - pub type OcrMachine = HostedMachine; +fn caller_token_provider(services: HostServices) -> 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) -> 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) }, ) } diff --git a/litellm-rust/crates/core/src/responses/handler.rs b/litellm-rust/crates/core/src/responses/handler.rs index ad85a750b4c..71b3b268d74 100644 --- a/litellm-rust/crates/core/src/responses/handler.rs +++ b/litellm-rust/crates/core/src/responses/handler.rs @@ -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, + interceptors: &impl Interceptors, + observers: Option<&ObservationSender>, ) -> Result { 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) diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs index cba9313d0ad..8997aea8269 100644 --- a/litellm-rust/crates/core/src/responses/mod.rs +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -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, + interceptors: &impl Interceptors, + observers: Option, ) -> Result { - 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, + interceptors: &impl Interceptors, + observers: Option<&ObservationSender>, ) -> Result { 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> = - Box::pin(handler::execute(&self.http, &self.auth, request, hooks)); + Box::pin(handler::execute( + &self.http, + &self.auth, + request, + interceptors, + observers, + )); execute.await }) .await diff --git a/litellm-rust/crates/core/src/responses/prepare.rs b/litellm-rust/crates/core/src/responses/prepare.rs index cf7eb11813d..2d58d9097d8 100644 --- a/litellm-rust/crates/core/src/responses/prepare.rs +++ b/litellm-rust/crates/core/src/responses/prepare.rs @@ -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, diff --git a/litellm-rust/crates/core/src/responses/route.rs b/litellm-rust/crates/core/src/responses/route.rs index a2d2788d519..3d7f545f443 100644 --- a/litellm-rust/crates/core/src/responses/route.rs +++ b/litellm-rust/crates/core/src/responses/route.rs @@ -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 { - hosted_call(call, move |call, _, hooks| async move { - self.run(call, &hooks).await - }) + pub fn machine( + self, + call: ResponsesCall, + observers: Option, + ) -> HostedMachine { + hosted_call( + call, + observers, + move |call, _, interceptors, observers| async move { + self.run(call, &interceptors, observers.as_ref()).await + }, + ) } } diff --git a/litellm-rust/crates/core/src/responses/types.rs b/litellm-rust/crates/core/src/responses/types.rs index ee93f28ac3a..ce634a9862f 100644 --- a/litellm-rust/crates/core/src/responses/types.rs +++ b/litellm-rust/crates/core/src/responses/types.rs @@ -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, } diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs index 681670b56f6..bb62fd000a8 100644 --- a/litellm-rust/crates/core/tests/chat_completions.rs +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -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 { - chat_completions_route().execute(request, &()).await + chat_completions_route().execute(request, &(), None).await } fn object(value: Value) -> Map { @@ -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 for FailingHook { + impl Interceptors 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 { 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(); diff --git a/litellm-rust/crates/core/tests/messages/host.rs b/litellm-rust/crates/core/tests/messages/host.rs index bcea70c91fb..6c6bf144238 100644 --- a/litellm-rust/crates/core/tests/messages/host.rs +++ b/litellm-rust/crates/core/tests/messages/host.rs @@ -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 Result + Send + Sy struct RecordingHost { call: LocalMessagesHost, rewrite: Rewrite, - events: Mutex>, + events: super::support::Observations, optional_params: Mutex>, } @@ -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 { 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<::Error> +impl litellm_host::interceptors::Interceptors<::Error> for RecordingHost { async fn before_provider_request( @@ -80,20 +84,22 @@ impl litellm_host::hooks::RouteHooks< Result<(), ::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 { - litellm_host::in_process::run_hosted( + litellm_host_native::in_process::run_hosted( machine(Arc::new(RecordingSecrets::empty()))(host.request()?), host.runtime(), ) diff --git a/litellm-rust/crates/core/tests/messages/main.rs b/litellm-rust/crates/core/tests/messages/main.rs index 7de463fb8fb..05e9aadd351 100644 --- a/litellm-rust/crates/core/tests/messages/main.rs +++ b/litellm-rust/crates/core/tests/messages/main.rs @@ -99,7 +99,7 @@ fn headers<'a>(pairs: impl IntoIterator) -> Option) -> 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 { 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<::Error> +impl litellm_host::interceptors::Interceptors<::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, ::Error, > { Ok(wire) } - async fn on_event( + async fn after_provider_response( &self, - event: litellm_host::event::MachineEvent, + raw: litellm_host::interceptors::RawResponse, ) -> Result<(), ::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(()) } diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index c9bcc210126..7d2fffc5beb 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -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 = + with_observer.then(|| host.events.0.sender.clone()); let future: BoxFuture<'_, Result> = 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"); diff --git a/litellm-rust/crates/core/tests/messages/stream.rs b/litellm-rust/crates/core/tests/messages/stream.rs index 8b51704d968..0fb85920077 100644 --- a/litellm-rust/crates/core/tests/messages/stream.rs +++ b/litellm-rust/crates/core/tests/messages/stream.rs @@ -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 { 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 for RecordingStreamHost { - async fn open_stream(&self, head: MessagesStreamHead) -> Result { +impl litellm_host_native::in_process::StreamConsumer for RecordingStreamHost { + async fn open_stream(&self, head: MessagesStreamHead) -> Result, Error> { Ok(self.record(Seen::Open(head.headers))) } - async fn send_chunk(&self, chunk: Bytes) -> Result { + async fn send_chunk(&self, chunk: Bytes) -> Result, 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<::Error> +impl litellm_host::interceptors::Interceptors<::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, ::Error, > { Ok(wire) } - async fn on_event( + async fn after_provider_response( &self, - event: litellm_host::event::MachineEvent, + raw: litellm_host::interceptors::RawResponse, ) -> Result<(), ::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 { - 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(); diff --git a/litellm-rust/crates/core/tests/ocr/azure_document_intelligence.rs b/litellm-rust/crates/core/tests/ocr/azure_document_intelligence.rs index e7b30e5a7e4..cec43e2ed0e 100644 --- a/litellm-rust/crates/core/tests/ocr/azure_document_intelligence.rs +++ b/litellm-rust/crates/core/tests/ocr/azure_document_intelligence.rs @@ -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") diff --git a/litellm-rust/crates/core/tests/ocr/documents.rs b/litellm-rust/crates/core/tests/ocr/documents.rs index f0bbfc5e0be..d99d99ead6e 100644 --- a/litellm-rust/crates/core/tests/ocr/documents.rs +++ b/litellm-rust/crates/core/tests/ocr/documents.rs @@ -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; diff --git a/litellm-rust/crates/core/tests/ocr/lifecycle.rs b/litellm-rust/crates/core/tests/ocr/lifecycle.rs index b138fafff17..d9492c4173a 100644 --- a/litellm-rust/crates/core/tests/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/tests/ocr/lifecycle.rs @@ -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>>, + interceptions: Arc, 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 for CallerTokenHost { +impl litellm_host_native::services::HostCallHandler 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 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<::Error> +impl litellm_host::interceptors::Interceptors<::Error> for CallerTokenHost { async fn before_provider_request( @@ -263,13 +268,15 @@ impl litellm_host::hooks::RouteHooks<:: .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<(), ::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); - - impl RouteHooks for Hooks { - fn observer(&self) -> Option> { - Some(self.0.clone()) - } + struct Hooks; + impl Interceptors 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 { .. } ] )); diff --git a/litellm-rust/crates/core/tests/ocr/machine.rs b/litellm-rust/crates/core/tests/ocr/machine.rs index f4eaa735fc1..b3653b65f2b 100644 --- a/litellm-rust/crates/core/tests/ocr/machine.rs +++ b/litellm-rust/crates/core/tests/ocr/machine.rs @@ -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()); diff --git a/litellm-rust/crates/core/tests/ocr/main.rs b/litellm-rust/crates/core/tests/ocr/main.rs index d180edaf105..bf0752707fb 100644 --- a/litellm-rust/crates/core/tests/ocr/main.rs +++ b/litellm-rust/crates/core/tests/ocr/main.rs @@ -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 { - ocr_route().execute(request, &()).await + ocr_route().execute(request, &(), None).await } async fn perform_with(host: LocalOcrHost) -> Result { - 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>>, before_provider_request: Option, observer: Option, + 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 for LocalOcrHost { +impl litellm_host_native::services::HostCallHandler 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 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<::Error> +impl litellm_host::interceptors::Interceptors<::Error> for LocalOcrHost { async fn before_provider_request( @@ -240,13 +252,15 @@ impl litellm_host::hooks::RouteHooks<:: None => Ok(wire), } } - async fn on_event( + async fn after_provider_response( &self, - event: litellm_host::event::MachineEvent, + raw: litellm_host::interceptors::RawResponse, ) -> Result<(), ::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(()) } diff --git a/litellm-rust/crates/core/tests/ocr/mistral.rs b/litellm-rust/crates/core/tests/ocr/mistral.rs index 6c5463ce035..34dec577c05 100644 --- a/litellm-rust/crates/core/tests/ocr/mistral.rs +++ b/litellm-rust/crates/core/tests/ocr/mistral.rs @@ -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(); diff --git a/litellm-rust/crates/core/tests/ocr/reducto.rs b/litellm-rust/crates/core/tests/ocr/reducto.rs index 8ccab27e58d..649dca3d690 100644 --- a/litellm-rust/crates/core/tests/ocr/reducto.rs +++ b/litellm-rust/crates/core/tests/ocr/reducto.rs @@ -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()); } }); diff --git a/litellm-rust/crates/core/tests/ocr/vertex_ai.rs b/litellm-rust/crates/core/tests/ocr/vertex_ai.rs index d77c065686e..9981c5362e2 100644 --- a/litellm-rust/crates/core/tests/ocr/vertex_ai.rs +++ b/litellm-rust/crates/core/tests/ocr/vertex_ai.rs @@ -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(); diff --git a/litellm-rust/crates/core/tests/resources.rs b/litellm-rust/crates/core/tests/resources.rs index 15cfc063506..d910792ff72 100644 --- a/litellm-rust/crates/core/tests/resources.rs +++ b/litellm-rust/crates/core/tests/resources.rs @@ -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(); diff --git a/litellm-rust/crates/core/tests/responses.rs b/litellm-rust/crates/core/tests/responses.rs index 84a0a7d14f3..73d3208afdf 100644 --- a/litellm-rust/crates/core/tests/responses.rs +++ b/litellm-rust/crates/core/tests/responses.rs @@ -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 }) diff --git a/litellm-rust/crates/core/tests/support/mod.rs b/litellm-rust/crates/core/tests/support/mod.rs index 82663fbb418..1dd53114293 100644 --- a/litellm-rust/crates/core/tests/support/mod.rs +++ b/litellm-rust/crates/core/tests/support/mod.rs @@ -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 { } #[derive(Default)] -pub struct CallEvents(pub Mutex>); +pub struct CallEvents(pub Observations); +pub struct Observations { + pub sender: litellm_host::observation::ObservationSender, + receiver: Mutex>, + recorded: Mutex>, +} + +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>> + { + 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 RecordingCall

{ } } -impl litellm_host::hooks::RouteHooks +impl litellm_host::interceptors::Interceptors for RecordingCall

{ - fn observer(&self) -> Option> { - Some(self.events.clone()) - } - async fn before_provider_request( &self, - wire: litellm_host::event::WireRequest, - _: litellm_host::event::RequestContext, - ) -> Result { - Ok(litellm_host::event::WireRequest { + wire: litellm_host::interceptors::WireRequest, + _: litellm_host::interceptors::RequestContext, + ) -> Result { + Ok(litellm_host::interceptors::WireRequest { headers: wire .headers .into_iter() @@ -289,12 +320,10 @@ impl litellm_host::hooks::RouteHooks 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

litellm_host::in_process::StreamConsumer

for RecordingCall

+impl

litellm_host_native::in_process::StreamConsumer

for RecordingCall

where P: litellm_host::protocol::Protocol, P::Error: From, { - async fn open_stream( - &self, - head: P::StreamHead, - ) -> Result { + async fn open_stream(&self, head: P::StreamHead) -> Result, 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 { + async fn send_chunk(&self, chunk: P::Chunk) -> Result, P::Error> { self.chunks.lock().unwrap().push(chunk); - Ok(litellm_host::protocol::Demand::More) + Ok(ControlFlow::Continue(())) } } impl

litellm_host::lifecycle::CallObserver for RecordingCall

@@ -346,8 +369,8 @@ where P: litellm_host::protocol::Protocol, P::Error: From, { - 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); } } diff --git a/litellm-rust/crates/gateway-inference/src/chat_completions.rs b/litellm-rust/crates/gateway-inference/src/chat_completions.rs index 069bbece5fe..85b1f990a40 100644 --- a/litellm-rust/crates/gateway-inference/src/chat_completions.rs +++ b/litellm-rust/crates/gateway-inference/src/chat_completions.rs @@ -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) diff --git a/litellm-rust/crates/gateway-inference/src/messages.rs b/litellm-rust/crates/gateway-inference/src/messages.rs index 87572156502..14fa0865b6e 100644 --- a/litellm-rust/crates/gateway-inference/src/messages.rs +++ b/litellm-rust/crates/gateway-inference/src/messages.rs @@ -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::::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( diff --git a/litellm-rust/crates/gateway-inference/src/ocr.rs b/litellm-rust/crates/gateway-inference/src/ocr.rs index a7b9d26d656..3a62e006dce 100644 --- a/litellm-rust/crates/gateway-inference/src/ocr.rs +++ b/litellm-rust/crates/gateway-inference/src/ocr.rs @@ -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()), diff --git a/litellm-rust/crates/gateway-inference/src/responses.rs b/litellm-rust/crates/gateway-inference/src/responses.rs index a6ce4bad4ac..3aca2a0c7f7 100644 --- a/litellm-rust/crates/gateway-inference/src/responses.rs +++ b/litellm-rust/crates/gateway-inference/src/responses.rs @@ -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::::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?) } diff --git a/litellm-rust/crates/gateway-inference/tests/routes.rs b/litellm-rust/crates/gateway-inference/tests/routes.rs index 69450255086..d5233cadad8 100644 --- a/litellm-rust/crates/gateway-inference/tests/routes.rs +++ b/litellm-rust/crates/gateway-inference/tests/routes.rs @@ -110,6 +110,7 @@ async fn chat_errors_come_from_core( timeout: None, }, &(), + None, ) .await .unwrap_err(); diff --git a/litellm-rust/crates/host-http/AGENTS.md b/litellm-rust/crates/host-http/AGENTS.md index f41af54f6d0..e6b52d9ec1b 100644 --- a/litellm-rust/crates/host-http/AGENTS.md +++ b/litellm-rust/crates/host-http/AGENTS.md @@ -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 diff --git a/litellm-rust/crates/host-http/Cargo.toml b/litellm-rust/crates/host-http/Cargo.toml index 4555d2fba49..bd3347baa97 100644 --- a/litellm-rust/crates/host-http/Cargo.toml +++ b/litellm-rust/crates/host-http/Cargo.toml @@ -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] diff --git a/litellm-rust/crates/host-http/src/driver.rs b/litellm-rust/crates/host-http/src/driver.rs index 704b93c6b6b..bd2e919d1db 100644 --- a/litellm-rust/crates/host-http/src/driver.rs +++ b/litellm-rust/crates/host-http/src/driver.rs @@ -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

= MachineStep::Response>>; type Output = CallOutput, Bytes, E>; +type HostedDriver = Driver, S, H>; pub async fn serve_unary( machine: HostedMachine

, services: S, - hooks: H, + interceptors: H, encoder: A, + observers: Option, ) -> Result> where P: Protocol, P::Error: From, - H: RouteHooks, + H: Interceptors, S: HostCallHandler

, A: ResponseEncoder, { - 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( machine: HostedMachine

, services: S, - hooks: H, + interceptors: H, encoder: A, + observers: Option, ) -> Result> where P: Protocol, P::Error: From, A: StreamEncoder, - H: RouteHooks + 'static, + H: Interceptors + 'static, S: HostCallHandler

+ '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 { - machine: HostedMachine

, - services: S, - hooks: H, - demand: Option>, -} - -impl Driver +async fn start( + mut driver: HostedDriver, + encoder: Arc, +) -> Result>, Error> where P: Protocol, P::Error: From, - H: RouteHooks, - S: HostCallHandler

, + A: StreamEncoder, + H: Interceptors + 'static, + S: HostCallHandler

+ 'static, { - fn new(machine: HostedMachine

, services: S, hooks: H) -> Self { - Self { - machine, - services, - hooks, - demand: None, - } - } - - async fn advance(&mut self) -> Result, 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(mut self, encoder: Arc) -> Result>, Error> - where - A: StreamEncoder, - 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), } } diff --git a/litellm-rust/crates/host-http/tests/serve.rs b/litellm-rust/crates/host-http/tests/serve.rs index b3ae805a7f6..2ba39ad4cd2 100644 --- a/litellm-rust/crates/host-http/tests/serve.rs +++ b/litellm-rust/crates/host-http/tests/serve.rs @@ -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 for Adapter { +impl litellm_host_native::services::HostCallHandler 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>); +struct Observer(Observations); +struct Observations { + sender: litellm_host::observation::ObservationSender, + receiver: Mutex>, + recorded: Mutex>, +} + +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>> { + 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 for Hooks { - fn observer(&self) -> Option> { - Some(self.observer.clone()) - } - +impl Interceptors for Hooks { async fn before_provider_request( &self, wire: WireRequest, @@ -141,11 +165,13 @@ impl RouteHooks 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 { } #[fixture] -fn hooks(observer: Arc) -> Hooks { +fn interceptors(observer: Arc) -> Hooks { Hooks { observer, reject: false, @@ -165,11 +191,12 @@ fn hooks(observer: Arc) -> 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::( "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, #[case] drop_after: Option, ) { - 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::("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::( + "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::("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::( + "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::("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::( + "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::("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::( + "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::("input", move |_, _, route_hooks| async move { - if event { - route_hooks - .on_event(MachineEvent::ResponseReceived { - raw: RawResponse { + let machine = hosted_call::( + "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::>::new(move |host| { + let observer = interceptors.observer.clone(); + let machine = CallMachine::>::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::("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::( + "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::("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::( + "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::("input", move |_, _, _| async move { - let _release = release; - std::future::pending().await - }); + let machine = hosted_call::( + "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 for Credentials { +impl litellm_host_native::services::HostCallHandler 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 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, #[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::( "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); diff --git a/litellm-rust/crates/host-http/tests/sse.rs b/litellm-rust/crates/host-http/tests/sse.rs index d7bce0615ca..af2bb404207 100644 --- a/litellm-rust/crates/host-http/tests/sse.rs +++ b/litellm-rust/crates/host-http/tests/sse.rs @@ -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::((), 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::((), 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::((), |(), _, _| async { + let machine = hosted_call::((), 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!( diff --git a/litellm-rust/crates/host-native/AGENTS.md b/litellm-rust/crates/host-native/AGENTS.md new file mode 100644 index 00000000000..6882d3aaca2 --- /dev/null +++ b/litellm-rust/crates/host-native/AGENTS.md @@ -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>` 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 diff --git a/litellm-rust/crates/host-native/Cargo.toml b/litellm-rust/crates/host-native/Cargo.toml new file mode 100644 index 00000000000..d1e6c0f4349 --- /dev/null +++ b/litellm-rust/crates/host-native/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/host-native/src/driver.rs b/litellm-rust/crates/host-native/src/driver.rs new file mode 100644 index 00000000000..b559de6e29f --- /dev/null +++ b/litellm-rust/crates/host-native/src/driver.rs @@ -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 = ::Protocol; +type ErrorOf = as Protocol>::Error; + +pub enum Boundary { + Complete(M::Complete), + Open( as Protocol>::StreamHead), + Chunk( 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 { + machine: M, + services: S, + interceptors: H, + demand: Option>>, +} + +impl Driver +where + M: Machine, + S: HostCallHandler>, + H: Interceptors>, +{ + pub fn new(machine: M, services: S, interceptors: H) -> Self { + Self { + machine, + services, + interceptors, + demand: None, + } + } + + pub async fn advance(&mut self) -> Result, ErrorOf> { + self.resume(ControlFlow::Continue(())).await + } + + pub async fn detach(&mut self) -> Result, ErrorOf> { + 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) -> Result> { + self.demand = None; + self.machine.interrupt(HostFailure::Error(error)).await + } + + async fn resume(&mut self, demand: ControlFlow<()>) -> Result, ErrorOf> { + 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); + } + } + } +} diff --git a/litellm-rust/crates/host-native/src/in_process.rs b/litellm-rust/crates/host-native/src/in_process.rs new file mode 100644 index 00000000000..9962ddfcf7d --- /dev/null +++ b/litellm-rust/crates/host-native/src/in_process.rs @@ -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: Send + Sync { + fn open_stream( + &self, + head: P::StreamHead, + ) -> impl Future, P::Error>> + Send; + fn send_chunk( + &self, + chunk: P::Chunk, + ) -> impl Future, P::Error>> + Send; +} + +impl StreamConsumer

for () { + async fn open_stream(&self, _: P::StreamHead) -> Result, P::Error> { + Ok(ControlFlow::Continue(())) + } + async fn send_chunk(&self, _: P::Chunk) -> Result, 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( + machine: M, + host: Host<'_, S, H, C>, +) -> Result::Error> +where + M: Machine, + S: HostCallHandler, + H: Interceptors<::Error>, + C: StreamConsumer, +{ + run_with_completion(machine, host, |_| false).await +} + +pub async fn run_hosted( + machine: HostedMachine

, + host: Host<'_, S, H, C>, +) -> Result, P::Error> +where + P: Protocol, + P::Error: From, + S: HostCallHandler

, + H: Interceptors, + C: StreamConsumer

, +{ + run_with_completion(machine, host, |completion| { + matches!(completion, HostedCompletion::Detached) + }) + .await +} + +async fn run_with_completion( + machine: M, + host: Host<'_, S, H, C>, + detached: impl Fn(&M::Complete) -> bool, +) -> Result::Error> +where + M: Machine, + S: HostCallHandler, + H: Interceptors<::Error>, + C: StreamConsumer, +{ + 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( + mut driver: Driver, + stream: &C, +) -> Result::Error> +where + M: Machine, + S: HostCallHandler, + H: Interceptors<::Error>, + C: StreamConsumer, +{ + 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, + }; + } +} diff --git a/litellm-rust/crates/host-native/src/lib.rs b/litellm-rust/crates/host-native/src/lib.rs new file mode 100644 index 00000000000..ee8980f80d5 --- /dev/null +++ b/litellm-rust/crates/host-native/src/lib.rs @@ -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}; diff --git a/litellm-rust/crates/host/src/services.rs b/litellm-rust/crates/host-native/src/services.rs similarity index 58% rename from litellm-rust/crates/host/src/services.rs rename to litellm-rust/crates/host-native/src/services.rs index 48e046ea69c..5d43c0aeeb9 100644 --- a/litellm-rust/crates/host/src/services.rs +++ b/litellm-rust/crates/host-native/src/services.rs @@ -1,6 +1,7 @@ -use crate::protocol::Protocol; use std::{convert::Infallible, future::Future}; +use litellm_host::protocol::Protocol; + pub trait HostCallHandler: Send + Sync { fn handle_host_call( &self, @@ -8,6 +9,15 @@ pub trait HostCallHandler: Send + Sync { ) -> impl Future> + Send; } +impl + ?Sized> HostCallHandler

for &T { + fn handle_host_call( + &self, + call: P::HostCall, + ) -> impl Future> + Send { + (**self).handle_host_call(call) + } +} + impl> HostCallHandler

for () { async fn handle_host_call(&self, call: Infallible) -> Result<(), P::Error> { match call {} diff --git a/litellm-rust/crates/host-native/tests/driver.rs b/litellm-rust/crates/host-native/tests/driver.rs new file mode 100644 index 00000000000..c41c1e5a43b --- /dev/null +++ b/litellm-rust/crates/host-native/tests/driver.rs @@ -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 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; + +struct Services { + reject: bool, +} + +impl HostCallHandler 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>, + recorded: Mutex>, +} + +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>> { + 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, + reject: bool, +} + +impl Interceptors for Hooks { + async fn before_provider_request( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + 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 { + Arc::new(Observer::default()) +} + +#[fixture] +fn interceptors(observer: Arc) -> Hooks { + Hooks { + observer, + reject: false, + } +} + +fn dispatching_call() -> TestMachine { + hosted_call::( + "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); + +impl Drop for Release { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } +} + +struct Streaming { + polls: Arc, + released: Arc, + machine: TestMachine, +} + +fn streaming_call(chunks: Vec>) -> 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::( + "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, + #[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>>>, +} + +impl Machine for Interruptible { + type Protocol = TestProtocol; + type Complete = HostedCompletion; + + fn resume(&mut self) -> Step<'_, Self> { + self.inner.resume() + } + + fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { + self.interrupted.lock().unwrap().push(failure.clone()); + self.inner.interrupt(failure) + } +} + +struct Consumer { + detach_after: Option, + fail_after: Option, + delivered: Mutex>, +} + +impl Consumer { + fn demand_after(&self, delivered: usize) -> Result, 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 for Consumer { + async fn open_stream(&self, head: &'static str) -> Result, TestError> { + assert_eq!(head, "headers"); + self.demand_after(0) + } + + async fn send_chunk(&self, chunk: usize) -> Result, 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, + #[case] detach_after: Option, + #[case] fail_after: Option, + #[case] expected: Result, 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::>() + ); + 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>, + 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 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 { + 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::>() + ); + 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) { + 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 for ResponseGate { + async fn before_provider_request( + &self, + wire: WireRequest, + _: RequestContext, + ) -> Result { + 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::( + "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"), + } +} diff --git a/litellm-rust/crates/host-python/AGENTS.md b/litellm-rust/crates/host-python/AGENTS.md index 4b4cfc1a772..9cd65c4d0a6 100644 --- a/litellm-rust/crates/host-python/AGENTS.md +++ b/litellm-rust/crates/host-python/AGENTS.md @@ -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) diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 123840151cb..07586e9db07 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -1,18 +1,22 @@ -use crate::PythonHostCalls; -use litellm_host::call::HostedCompletion; -use litellm_host::event::WireRequest; -use litellm_host::event::{FailureOrigin, Timing, epoch_seconds}; -use litellm_host::machine::{HostFailure, Machine, MachineStep}; -use litellm_host::protocol::HookRequest; -use litellm_host::protocol::StreamDelivery; -use litellm_host::protocol::{Demand, Protocol, Reply, Suspension}; +use std::ops::ControlFlow; + use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; use pyo3::types::PyDict; -use crate::handle::{Execution, ExecutionBody, ExecutionStep}; -use crate::hooks::{HookEvent, HookResume, HookStep, Preflight, PythonCallHooks}; +use litellm_host::{ + call::HostedCompletion, + interceptors::WireRequest, + lifecycle::{CallEvent, ExecutionEvent, FailureOrigin, Timing, epoch_seconds}, + machine::{HostFailure, Machine, MachineStep}, + observation::ObservationSender, + protocol::{HostRequest, InterceptRequest, Protocol, Reply, StreamDelivery}, +}; + +use crate::PythonHostCalls; +use crate::handle::{Execution, ExecutionBody, ExecutionStep, PythonLifecycle}; +use crate::hooks::{HookResume, HookStep, PythonCallEvent, PythonCallHooks}; use crate::native::{NativeMachine, NativePoll}; use crate::{InvokeError, PythonBinding, missing_state}; @@ -22,7 +26,27 @@ type ResponseOf = as Protocol>::Response; type NativeStep = MachineStep, HostedCompletion>>; type NativeResult = Result, ErrorOf>; type Interruption = Option>>; -type StartMachine = Box::Request) -> M + Send + Sync>; +type StartMachine = Box< + dyn FnOnce(Python<'_>, &Bound<'_, PyDict>,

::Request) -> PyResult + + Send + + Sync, +>; + +pub struct CallOptions { + pub asynchronous: bool, + pub lifecycle: PythonLifecycle, + pub observers: Option, +} + +impl CallOptions { + pub fn new(asynchronous: bool, lifecycle: PythonLifecycle) -> Self { + Self { + asynchronous, + lifecycle, + observers: None, + } + } +} enum Stage { Begin, @@ -45,7 +69,7 @@ enum Pending { Wire(HookResume>, Reply), Response(HookResume>), Event(HookResume, EventNext), - Consumer(Reply), + Consumer(Reply>), } /// A route answer as the driver resumes on it: a Python exception interrupts the call as @@ -72,7 +96,6 @@ where { binding: H, hooks: L, - preflight: Preflight, native: NativeMachine, start: Option, M>>, closed: bool, @@ -82,19 +105,25 @@ where stage: Stage, pending: Option>, interrupted: Option>, + observers: Option, + observing: bool, + terminal_observation: Option, } -/// Runs one native call for Python: synchronously, or as a coroutine that awaits every -/// host suspension inline in the caller's task. `preflight` runs once, on the keyword view -/// the hooks' `prepare_arguments` returned, before the binding decodes the request. pub fn run_call( py: Python<'_>, - start: impl FnOnce(::Request) -> M + Send + Sync + 'static, + start: impl FnOnce( + Python<'_>, + &Bound<'_, PyDict>, + ::Request, + ) -> PyResult + + Send + + Sync + + 'static, binding: H, hooks: L, - preflight: Preflight, arguments: Py, - asynchronous: bool, + options: CallOptions, ) -> PyResult> where L: PythonCallHooks + 'static, @@ -105,8 +134,7 @@ where let mut driver = PythonDriver { binding, hooks, - preflight, - native: NativeMachine::new(asynchronous), + native: NativeMachine::new(options.asynchronous), start: Some(Box::new(start)), closed: false, arguments: Some(arguments), @@ -115,13 +143,18 @@ where stage: Stage::Begin, pending: None, interrupted: None, + observers: options.observers, + observing: false, + terminal_observation: None, }; - if asynchronous { - return Execution::new(driver).into_coroutine(py).map(Bound::unbind); + if options.asynchronous { + return Execution::new(driver, options.lifecycle) + .into_coroutine(py) + .map(Bound::unbind); } match driver.resume(None)? { ExecutionStep::Return(value) => Ok(value), - ExecutionStep::Open(head) => Execution::suspended(driver) + ExecutionStep::Open(head) => Execution::suspended(driver, options.lifecycle) .into_sync_stream(py, head) .map(Bound::unbind), ExecutionStep::Await(_) | ExecutionStep::Yield(_) => { @@ -156,9 +189,11 @@ where match (self.pending.take(), result) { (None, None) => { self.started_at = epoch_seconds(); - let started = HookEvent::Started { + let started = PythonCallEvent::Started { start_time: self.started_at, }; + self.observing = true; + self.observe(&started); match self.hooks.on_event(py, started) { Ok(step) => self.on_event(py, step, EventNext::Started), Err(error) => self.hook_failed(py, error), @@ -171,9 +206,9 @@ where (Some(Pending::Native), Some(Err(error))) => self.interrupt(py, error), (Some(Pending::Consumer(reply)), Some(read)) => { reply.send(if read.is_ok() { - Demand::More + ControlFlow::Continue(()) } else { - Demand::Detached + ControlFlow::Break(()) }); self.resume_machine(py, None) } @@ -220,7 +255,7 @@ where Ok(ExecutionStep::Await(awaitable)) } HookStep::Ready(arguments) => { - if let Err(error) = (self.preflight)(py, arguments.bind(py)) { + if let Err(error) = self.hooks.arguments_prepared(py, &arguments) { return self.hook_failed(py, error); } let decoded = self.binding.decode_request(py, arguments.bind(py)); @@ -233,7 +268,12 @@ where } }; let start = self.start.take().ok_or_else(missing_state)?; - self.native.start(start(request)); + let arguments = self.arguments.as_ref().ok_or_else(missing_state)?; + let machine = match start(py, arguments.bind(py), request) { + Ok(machine) => machine, + Err(error) => return self.failure(py, error, FailureOrigin::Host), + }; + self.native.start(machine); self.stage = Stage::Call; self.resume_machine(py, None) } @@ -289,13 +329,23 @@ where reply.send(()); self.resume_machine(py, None) } - EventNext::Terminal => match &self.stage { - Stage::Succeeded(response) => Ok(ExecutionStep::Return(response.clone_ref(py))), - Stage::Failed(error) => { - Err(PyErr::from_value(error.bind(py).clone().into_any())) + EventNext::Terminal => { + if let Some(event) = self.terminal_observation.take() + && let Some(observers) = &self.observers + { + observers.emit(event); } - _ => Err(missing_state()), - }, + self.observing = false; + match &self.stage { + Stage::Succeeded(response) => { + Ok(ExecutionStep::Return(response.clone_ref(py))) + } + Stage::Failed(error) => { + Err(PyErr::from_value(error.bind(py).clone().into_any())) + } + _ => Err(missing_state()), + } + } }, } } @@ -355,8 +405,8 @@ where Err(error) => return self.machine_failed(py, error).map(Next::Return), }; let answered = match op { - Suspension::HostCall(op) => answered(self.binding.handle_host_call(py, op)), - Suspension::Hook(HookRequest::BeforeProviderRequest { + HostRequest::HostCall(op) => answered(self.binding.handle_host_call(py, op)), + HostRequest::Intercept(InterceptRequest::BeforeProviderRequest { wire, context, reply, @@ -371,14 +421,18 @@ where } Err(error) => Err(error), }, - Suspension::Stream(StreamDelivery::Open(head, reply)) => { + HostRequest::Stream(StreamDelivery::Open(head, reply)) => { return self.opened(py, head, reply).map(Next::Return); } - Suspension::Stream(StreamDelivery::Chunk(chunk, reply)) => { + HostRequest::Stream(StreamDelivery::Chunk(chunk, reply)) => { return self.delivered(py, chunk, reply).map(Next::Return); } - Suspension::Hook(HookRequest::Event(event, reply)) => { - match self.hooks.on_event(py, HookEvent::Machine(&event)) { + HostRequest::Intercept(InterceptRequest::AfterProviderResponse { raw, reply }) => { + let event = PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { + raw: &raw, + }); + self.observe(&event); + match self.hooks.on_event(py, event) { Ok(HookStep::Ready(())) => { reply.send(()); Ok(Ok(())) @@ -405,7 +459,7 @@ where &mut self, py: Python<'_>, head: as Protocol>::StreamHead, - reply: Reply, + reply: Reply>, ) -> PyResult { self.stage = Stage::Streaming; let head = match self.binding.encode_stream_head(py, head) { @@ -425,7 +479,7 @@ where &mut self, py: Python<'_>, chunk: as Protocol>::Chunk, - reply: Reply, + reply: Reply>, ) -> PyResult { let chunk = match self.binding.encode_chunk(py, chunk) { Ok(chunk) => chunk, @@ -500,10 +554,11 @@ where } fn succeeded(&mut self, py: Python<'_>, response: Py) -> PyResult { - let event = HookEvent::Succeeded { + let event = PythonCallEvent::Succeeded { timing: self.timing(), response: &response, }; + self.terminal_observation = self.observers.as_ref().map(|_| event.snapshot()); let step = self.hooks.on_event(py, event)?; self.stage = Stage::Succeeded(response); self.on_event(py, step, EventNext::Terminal) @@ -519,19 +574,32 @@ where if is_cancellation(py, &error) { return Err(error); } - let event = HookEvent::Failed { + let event = PythonCallEvent::Failed { timing: self.timing(), origin, error: &error, }; + self.terminal_observation = self.observers.as_ref().map(|_| event.snapshot()); let step = self.hooks.on_event(py, event)?; self.stage = Stage::Failed(error.into_value(py)); self.on_event(py, step, EventNext::Terminal) } + fn observe(&self, event: &PythonCallEvent<'_>) { + if let Some(observers) = &self.observers { + observers.emit(event.snapshot()); + } + } + fn clear(&mut self) { if !self.closed { self.closed = true; + if self.observing { + self.observe(&PythonCallEvent::Cancelled { + timing: self.timing(), + }); + self.observing = false; + } self.native.close(); self.start = None; Python::attach(|py| { @@ -550,7 +618,28 @@ where M::Complete: Into>>, { fn resume(&mut self, result: Option>>) -> PyResult { - Python::attach(|py| self.drive(py, result)) + Python::attach(|py| { + let outcome = self.drive(py, result); + if let Err(error) = &outcome + && self.observing + { + let event = if is_cancellation(py, error) { + PythonCallEvent::Cancelled { + timing: self.timing(), + } + } else { + PythonCallEvent::Failed { + timing: self.timing(), + origin: FailureOrigin::Host, + error, + } + }; + self.observe(&event); + self.observing = false; + self.terminal_observation = None; + } + outcome + }) } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { @@ -582,32 +671,29 @@ where mod tests { use std::sync::{Arc, Mutex}; - use litellm_host::event::{MachineEvent, RawResponse, RequestContext}; - use litellm_host::machine::{CallMachine, MachineFault}; + use litellm_host::{ + interceptors::{RawResponse, RequestContext}, + machine::{CallMachine, MachineFault}, + }; use pyo3::exceptions::{PyBaseException, PyValueError}; use pyo3::types::PyDict; use super::*; - use crate::PythonOwned; - use litellm_host::hooks::RouteHooks; + use crate::{PythonOwned, PythonRuntime}; + use litellm_host::hooks::CallHooks; + use litellm_host::interceptors::Interceptors; static PYTHON_GLOBALS: Mutex<()> = Mutex::new(()); - fn install_lifecycle_module(py: Python<'_>) -> Bound<'_, PyModule> { - py.run( - pyo3::ffi::c_str!( - r#" -import sys -import types + fn lifecycle_binding(py: Python<'_>) -> PyResult> { + py.import("host_test_lifecycle") + } -sys.modules.setdefault('litellm', types.ModuleType('litellm')) -sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bridge')) -"# - ), - None, - None, - ) - .unwrap(); + fn call_options(asynchronous: bool) -> CallOptions { + CallOptions::new(asynchronous, lifecycle_binding) + } + + fn install_lifecycle_module(py: Python<'_>) -> Bound<'_, PyModule> { let source = std::ffi::CString::new(include_str!("../../../../litellm/rust_bridge/lifecycle.py")) .unwrap(); @@ -615,7 +701,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py, &source, pyo3::ffi::c_str!("lifecycle.py"), - pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"), + pyo3::ffi::c_str!("host_test_lifecycle"), ) .unwrap() } @@ -795,9 +881,13 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri #[derive(Clone, Copy)] enum HookScript { Plain, + RewriteArguments, + ObserveArguments, FailBegin, ReplaceResponse, FailAfterSuccess, + CancelTerminal, + FailTerminal, } struct SyntheticHooks { @@ -805,7 +895,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri script: HookScript, } - impl PythonCallHooks for SyntheticHooks { + impl CallHooks for SyntheticHooks { fn prepare_arguments( &mut self, _: Python<'_>, @@ -816,9 +906,30 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri if matches!(self.script, HookScript::FailBegin) { return Err(PyValueError::new_err("begin failed")); } + if matches!(self.script, HookScript::RewriteArguments) { + let prepared = Python::attach(|py| -> PyResult> { + let prepared = arguments.bind(py).copy()?; + prepared.set_item("prepared", "hook")?; + prepared.set_item("api_key", "hook-key")?; + Ok(prepared.unbind()) + })?; + return Ok(HookStep::Ready(prepared)); + } Ok(HookStep::Ready(arguments)) } + fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py) -> PyResult<()> { + if matches!(self.script, HookScript::ObserveArguments) { + let value: String = arguments + .bind(py) + .get_item("prepared")? + .unwrap() + .extract()?; + self.log.push(format!("adopted:{value}")); + } + Ok(()) + } + fn before_provider_request( &mut self, _: Python<'_>, @@ -844,24 +955,48 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "replaced".into_pyobject(py)?.into_any().unbind(), )), HookScript::FailAfterSuccess => Err(PyValueError::new_err("after_success failed")), - HookScript::Plain | HookScript::FailBegin => Ok(HookStep::Ready(response)), + HookScript::Plain + | HookScript::ObserveArguments + | HookScript::RewriteArguments + | HookScript::FailBegin + | HookScript::CancelTerminal + | HookScript::FailTerminal => Ok(HookStep::Ready(response)), } } fn on_event( &mut self, py: Python<'_>, - event: HookEvent<'_>, + event: PythonCallEvent<'_>, ) -> PyResult> { + if matches!(self.script, HookScript::CancelTerminal) + && matches!( + event, + PythonCallEvent::Succeeded { .. } | PythonCallEvent::Failed { .. } + ) + { + return Err(pyo3::exceptions::asyncio::CancelledError::new_err( + "callback cancelled", + )); + } + if matches!(self.script, HookScript::FailTerminal) + && matches!( + event, + PythonCallEvent::Succeeded { .. } | PythonCallEvent::Failed { .. } + ) + { + return Err(PyValueError::new_err("callback failed")); + } self.log.push(match event { - HookEvent::Started { .. } => "started".into(), - HookEvent::Machine(MachineEvent::ResponseReceived { raw }) => { + PythonCallEvent::Started { .. } => "started".into(), + PythonCallEvent::Cancelled { .. } => "cancelled".into(), + PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => { format!("response:{}", raw.body) } - HookEvent::Succeeded { response, .. } => { + PythonCallEvent::Succeeded { response, .. } => { format!("succeeded:{}", response.bind(py)) } - HookEvent::Failed { origin, error, .. } => { + PythonCallEvent::Failed { origin, error, .. } => { format!("failed:{origin:?}:{}", error.value(py)) } }); @@ -915,21 +1050,93 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri script: HookScript, asynchronous: bool, ) -> (PyResult>, Vec) { - run_preflighted(py, machine, host, script, no_preflight, asynchronous) + run_composed( + py, + machine, + host, + script, + std::convert::identity, + call_options(asynchronous), + ) } - fn no_preflight(_: Python<'_>, _: &Bound<'_, PyDict>) -> PyResult<()> { - Ok(()) + #[rstest::rstest] + #[case::synchronous(false)] + #[case::asynchronous(true)] + fn composed_hooks_share_prepared_arguments_and_one_execution(#[case] asynchronous: bool) { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + let log = Log::default(); + let hook = |script| SyntheticHooks { + log: Log(log.0.clone()), + script, + }; + let hooks = crate::HookChain::new() + .with(hook(HookScript::ObserveArguments)) + .with(hook(HookScript::RewriteArguments)) + .with(hook(HookScript::ReplaceResponse)); + let result = run_call( + py, + move |_, _, request| Ok(success_machine()(request)), + SyntheticBinding { + log: Log(log.0.clone()), + op: OpScript::Answer, + classifier_fails: false, + }, + hooks, + PyDict::new(py).unbind(), + call_options(asynchronous), + ) + .unwrap(); + let response = if asynchronous { + assert!(log.entries().is_empty()); + let completion = result.call_method1(py, "send", (py.None(),)).unwrap_err(); + assert!(completion.is_instance_of::(py)); + completion.value(py).getattr("value").unwrap().unbind() + } else { + result + }; + assert_eq!(response.extract::(py).unwrap(), "replaced"); + let entries = log.entries(); + assert!(entries.iter().any(|entry| entry == "adopted:hook")); + assert_eq!( + entries + .iter() + .filter(|entry| entry.as_str() == "project") + .count(), + 1 + ); + assert_eq!( + entries + .iter() + .filter(|entry| entry.as_str() == "op:sign") + .count(), + 1 + ); + assert_eq!( + entries + .iter() + .filter(|entry| entry.as_str() == "succeeded:replaced") + .count(), + 3 + ); + assert!(!entries.iter().any(|entry| entry.starts_with("failed:"))); + }); } - fn run_preflighted( + fn run_composed( py: Python<'_>, machine: impl FnOnce(String) -> CallMachine + Send + Sync + 'static, host: SyntheticBinding, script: HookScript, - preflight: Preflight, - asynchronous: bool, + compose: impl FnOnce(SyntheticHooks) -> L, + options: CallOptions, ) -> (PyResult>, Vec) { + let asynchronous = options.asynchronous; let log = Log(host.log.0.clone()); let adapter = SyntheticHooks { log: Log(log.0.clone()), @@ -939,12 +1146,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri arguments.set_item("model", "m").unwrap(); let result = run_call( py, - machine, + move |_, _, request| Ok(machine(request)), host, - adapter, - preflight, + compose(adapter), arguments.unbind(), - asynchronous, + options, ); let result = if asynchronous { result.and_then(|coroutine| { @@ -966,17 +1172,15 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri /// response, so a driver that misroutes a reply changes what the call returns. fn success_machine() -> impl FnOnce(String) -> CallMachine + Send + Sync { move |projected| { - CallMachine::new(move |host| { + CallMachine::new(None, move |host| { Box::pin(async move { let signed = host.services.call(|reply| ("sign", reply)).await?; let wire = host - .hooks + .interceptors .before_provider_request(wire(), context()) .await?; - host.hooks - .on_event(MachineEvent::ResponseReceived { - raw: RawResponse { body: "raw".into() }, - }) + host.interceptors + .after_provider_response(RawResponse { body: "raw".into() }) .await?; Ok(format!("{projected}|{signed}|{}", wire.url)) }) @@ -984,6 +1188,155 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } + #[rstest::rstest] + #[case::draining(4, false)] + #[case::full(1, false)] + #[case::closed(4, true)] + fn observation_delivery_preserves_callback_order_and_response( + #[case] capacity: usize, + #[case] closed: bool, + #[values(false, true)] asynchronous: bool, + ) { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + let (expected, expected_log) = run_scripted( + py, + success_machine(), + OpScript::Answer, + HookScript::ReplaceResponse, + asynchronous, + ); + let (sender, mut receiver) = litellm_host::observation::observation_channel( + std::num::NonZeroUsize::new(capacity).unwrap(), + ); + if closed { + receiver.close(); + } + let (actual, actual_log) = run_composed( + py, + success_machine(), + SyntheticBinding { + log: Log::default(), + op: OpScript::Answer, + classifier_fails: false, + }, + HookScript::ReplaceResponse, + std::convert::identity, + CallOptions { + asynchronous, + lifecycle: lifecycle_binding, + observers: Some(sender.clone()), + }, + ); + assert_eq!( + actual.unwrap().extract::(py).unwrap(), + expected.unwrap().extract::(py).unwrap() + ); + assert_eq!(actual_log, expected_log); + let events: Vec<_> = std::iter::from_fn(|| receiver.try_recv().ok()).collect(); + assert_eq!(events.len() as u64 + sender.dropped_events(), 3); + if !closed && capacity >= 3 { + let [ + CallEvent::Started { start_time }, + CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }), + CallEvent::Succeeded { timing, .. }, + ] = events.as_slice() + else { + panic!("expected start, provider response and success: {events:?}"); + }; + assert_eq!(raw.body, "raw"); + assert_eq!(*start_time, timing.start_time); + assert!(timing.end_time >= timing.start_time); + } + }); + } + + #[rstest::rstest] + #[case::preparation(HookScript::FailBegin, false, 1)] + #[case::transformation(HookScript::FailAfterSuccess, false, 1)] + #[case::terminal_cancellation(HookScript::CancelTerminal, true, 0)] + #[case::terminal_failure(HookScript::FailTerminal, false, 0)] + fn observation_reports_the_final_outcome_without_replaying_callbacks( + #[case] script: HookScript, + #[case] cancelled: bool, + #[case] failure_callbacks: usize, + #[values(false, true)] asynchronous: bool, + ) { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + let (sender, mut receiver) = litellm_host::observation::observation_channel( + std::num::NonZeroUsize::new(4).unwrap(), + ); + let (result, log) = run_composed( + py, + success_machine(), + SyntheticBinding { + log: Log::default(), + op: OpScript::Answer, + classifier_fails: false, + }, + script, + std::convert::identity, + CallOptions { + asynchronous, + lifecycle: lifecycle_binding, + observers: Some(sender), + }, + ); + let error = result.unwrap_err(); + assert_eq!( + error.is_instance_of::(py), + cancelled + ); + let events: Vec<_> = std::iter::from_fn(|| receiver.try_recv().ok()).collect(); + assert!(matches!(events.first(), Some(CallEvent::Started { .. }))); + assert_eq!( + matches!(events.last(), Some(CallEvent::Cancelled { .. })), + cancelled + ); + assert_eq!( + matches!( + events.last(), + Some(CallEvent::Failed { + origin: FailureOrigin::Host, + .. + }) + ), + !cancelled + ); + assert_eq!( + events + .iter() + .filter(|event| matches!( + event, + CallEvent::Failed { .. } + | CallEvent::Succeeded { .. } + | CallEvent::Cancelled { .. } + )) + .count(), + 1 + ); + assert_eq!( + log.iter() + .filter(|entry| entry.starts_with("failed:")) + .count(), + failure_callbacks + ); + assert!( + !log.iter() + .any(|entry| entry.starts_with("succeeded:") || entry == "cancelled") + ); + }); + } + #[rstest::rstest] #[case::preparation(HookScript::FailBegin, OpScript::Answer)] #[case::native_decode(HookScript::Plain, OpScript::RejectRequestNatively)] @@ -1154,7 +1507,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn streaming_machine() -> impl FnOnce(()) -> litellm_host::call::HostedMachine + Send + Sync { move |()| { - litellm_host::call::hosted_call((), |(), _, _| async { + litellm_host::call::hosted_call((), None, |(), _, _, _observations| async { Ok(litellm_host::call::CallOutput::Stream { head: vec![("request-id", "req_1")], chunks: Box::pin(futures_util::stream::iter([Ok("first"), Ok("second")])), @@ -1176,17 +1529,23 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri Python::attach(|py| { install_lifecycle_module(py); let log = Log::default(); + let (sender, mut receiver) = litellm_host::observation::observation_channel( + std::num::NonZeroUsize::new(3).unwrap(), + ); let handed = run_call( py, - streaming_machine(), + |_, _, request| Ok(streaming_machine()(request)), StreamingBinding, SyntheticHooks { log: Log(log.0.clone()), script: HookScript::Plain, }, - no_preflight, PyDict::new(py).unbind(), - asynchronous, + CallOptions { + asynchronous, + lifecycle: lifecycle_binding, + observers: Some(sender), + }, ) .unwrap(); let stream = if asynchronous { @@ -1216,6 +1575,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri stream.call_method0("close").unwrap(); stream.call_method0("close").unwrap(); } + let events: Vec<_> = std::iter::from_fn(|| receiver.try_recv().ok()).collect(); + assert!(matches!( + events.as_slice(), + [CallEvent::Started { .. }, CallEvent::Succeeded { .. }] + )); assert_eq!( log.entries(), [ @@ -1255,7 +1619,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } #[rstest::rstest] - fn a_stream_carries_its_head_as_hidden_params_before_the_first_chunk() { + fn a_stream_carries_its_head_before_the_first_chunk() { let _guard = PYTHON_GLOBALS .lock() .unwrap_or_else(|error| error.into_inner()); @@ -1270,12 +1634,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }; let handed = run_call( py, - streaming_machine(), + |_, _, request| Ok(streaming_machine()(request)), StreamingBinding, adapter, - no_preflight, PyDict::new(py).unbind(), - asynchronous, + call_options(asynchronous), ) .unwrap(); let stream = if asynchronous { @@ -1287,7 +1650,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let hidden: std::collections::HashMap< String, std::collections::HashMap, - > = stream.getattr("_hidden_params").unwrap().extract().unwrap(); + > = stream.getattr("head").unwrap().extract().unwrap(); assert_eq!( hidden["additional_headers"], std::collections::HashMap::from([( @@ -1303,7 +1666,9 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn failing_machine() -> impl FnOnce(String) -> CallMachine + Send + Sync { move |_| { - CallMachine::new(|_| Box::pin(async move { Err(Error("provider exploded".into())) })) + CallMachine::new(None, |_| { + Box::pin(async move { Err(Error("provider exploded".into())) }) + }) } } @@ -1476,84 +1841,240 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - /// The rejection a preflight raised, kept so a test can check the caller receives that - /// exact object. A `Preflight` is a plain `fn`, so it cannot capture one itself. - static REJECTION: Mutex>> = Mutex::new(None); - - fn rejecting_preflight(py: Python<'_>, _: &Bound<'_, PyDict>) -> PyResult<()> { - let error = PyValueError::new_err("over budget"); - *REJECTION.lock().unwrap() = Some(error.value(py).clone().unbind()); - Err(error) + enum ArgumentPolicy { + Inherit, + Reject(Py), } - fn inheriting_preflight(_: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<()> { - arguments.set_item("api_key", "inherited") + impl CallHooks for ArgumentPolicy { + fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py) -> PyResult<()> { + match self { + Self::Inherit => arguments.bind(py).set_item("api_key", "inherited"), + Self::Reject(error) => Err(PyErr::from_value(error.bind(py).clone().into_any())), + } + } + } + + impl PythonOwned for ArgumentPolicy { + fn close(&mut self, _: Python<'_>) {} + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + match self { + Self::Inherit => Ok(()), + Self::Reject(error) => visit.call(error), + } + } } #[rstest::rstest] - fn a_preflight_rejection_is_the_callers_error_and_the_machine_never_starts() { + #[case::sync_success(false, false)] + #[case::async_success(true, false)] + #[case::sync_failure(false, true)] + #[case::async_failure(true, true)] + fn resource_setup_uses_prepared_arguments_and_failures_are_terminal( + #[case] asynchronous: bool, + #[case] fail_setup: bool, + ) { let _guard = PYTHON_GLOBALS .lock() .unwrap_or_else(|error| error.into_inner()); crate::initialize_python(); Python::attach(|py| { install_lifecycle_module(py); - for asynchronous in [false, true] { - let (result, log) = run_preflighted( - py, - success_machine(), - SyntheticBinding { - log: Log::default(), - op: OpScript::Answer, - classifier_fails: false, - }, - HookScript::Plain, - rejecting_preflight, - asynchronous, - ); - let error = result.unwrap_err(); - let raised = REJECTION.lock().unwrap().take().unwrap(); - assert!(error.value(py).is(&raised)); + let log = Log::default(); + let setup_log = Log(log.0.clone()); + let error = PyValueError::new_err("setup failed").into_value(py); + let setup_error = error.clone_ref(py); + let arguments = PyDict::new(py); + arguments.set_item("api_key", "original").unwrap(); + let result = run_call( + py, + move |py, prepared, request| { + let source: String = prepared.get_item("prepared")?.unwrap().extract()?; + let key: String = prepared.get_item("api_key")?.unwrap().extract()?; + setup_log.push(format!("setup:{source}:{key}")); + if fail_setup { + return Err(PyErr::from_value(setup_error.into_bound(py).into_any())); + } + Ok(success_machine()(request)) + }, + SyntheticBinding { + log: Log(log.0.clone()), + op: OpScript::Answer, + classifier_fails: false, + }, + crate::HookChain::new() + .with(SyntheticHooks { + log: Log(log.0.clone()), + script: HookScript::RewriteArguments, + }) + .with(ArgumentPolicy::Inherit), + arguments.clone().unbind(), + call_options(asynchronous), + ); + let settled = if asynchronous { + assert!(log.entries().is_empty()); + let completion = result + .unwrap() + .call_method1(py, "send", (py.None(),)) + .unwrap_err(); + if completion.is_instance_of::(py) { + completion.value(py).getattr("value").map(Bound::unbind) + } else { + Err(completion) + } + } else { + result + }; + assert_eq!( + arguments + .get_item("api_key") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + "original" + ); + assert_eq!( + &log.entries()[..4], + ["started", "begin", "project", "setup:hook:inherited"] + ); + if fail_setup { + assert!(settled.unwrap_err().value(py).is(error.bind(py))); assert_eq!( - log, + log.entries(), [ "started", "begin", - "failed:Host:over budget", + "project", + "setup:hook:inherited", + "failed:Host:setup failed", "adapter.close", - "host.close" + "host.close", ] ); + } else { + assert!(settled.is_ok()); + assert!( + log.entries() + .iter() + .any(|entry| entry.starts_with("succeeded:")) + ); } }); } #[rstest::rstest] - fn the_host_projects_from_the_keyword_view_the_preflight_rewrote() { + fn closing_an_unstarted_call_never_initializes_resources() { let _guard = PYTHON_GLOBALS .lock() .unwrap_or_else(|error| error.into_inner()); crate::initialize_python(); Python::attach(|py| { install_lifecycle_module(py); - for asynchronous in [false, true] { - let (result, _) = run_preflighted( - py, - success_machine(), - SyntheticBinding { - log: Log::default(), - op: OpScript::Answer, - classifier_fails: false, - }, - HookScript::Plain, - inheriting_preflight, - asynchronous, - ); - assert_eq!( - result.unwrap().extract::(py).unwrap(), - "project:2|sign|rewritten" - ); - } + let log = Log::default(); + let setup_log = Log(log.0.clone()); + let pending = run_call( + py, + move |_, _, request| { + setup_log.push("setup"); + Ok(success_machine()(request)) + }, + SyntheticBinding { + log: Log(log.0.clone()), + op: OpScript::Answer, + classifier_fails: false, + }, + SyntheticHooks { + log: Log(log.0.clone()), + script: HookScript::Plain, + }, + PyDict::new(py).unbind(), + call_options(true), + ) + .unwrap(); + assert!(log.entries().is_empty()); + pending.call_method0(py, "close").unwrap(); + drop(pending); + assert_eq!(log.entries(), ["adapter.close", "host.close"]); + }); + } + + #[rstest::rstest] + #[case::synchronous(false)] + #[case::asynchronous(true)] + fn an_argument_policy_rejection_is_the_callers_error_and_the_machine_never_starts( + #[case] asynchronous: bool, + ) { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + let raised = PyValueError::new_err("over budget").into_value(py); + let (result, log) = run_composed( + py, + success_machine(), + SyntheticBinding { + log: Log::default(), + op: OpScript::Answer, + classifier_fails: false, + }, + HookScript::Plain, + |hooks| { + crate::HookChain::new() + .with(hooks) + .with(ArgumentPolicy::Reject(raised.clone_ref(py))) + }, + call_options(asynchronous), + ); + assert!(result.unwrap_err().value(py).is(raised.bind(py))); + assert_eq!( + log, + [ + "started", + "begin", + "failed:Host:over budget", + "adapter.close", + "host.close" + ] + ); + }); + } + + #[rstest::rstest] + #[case::synchronous(false)] + #[case::asynchronous(true)] + fn the_host_projects_from_the_keyword_view_the_argument_policy_rewrote( + #[case] asynchronous: bool, + ) { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + let (result, _) = run_composed( + py, + success_machine(), + SyntheticBinding { + log: Log::default(), + op: OpScript::Answer, + classifier_fails: false, + }, + HookScript::Plain, + |hooks| { + crate::HookChain::new() + .with(hooks) + .with(ArgumentPolicy::Inherit) + }, + call_options(asynchronous), + ); + assert_eq!( + result.unwrap().extract::(py).unwrap(), + "project:2|sign|rewritten" + ); }); } @@ -1682,12 +2203,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }; let error = run_call( py, - success_machine(), + |_, _, request| Ok(success_machine()(request)), host, adapter, - no_preflight, PyDict::new(py).unbind(), - false, + call_options(false), ) .unwrap_err(); assert!(!error.is_instance_of::(py)); @@ -1745,9 +2265,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn retaining_coroutine(py: Python<'_>, retained: Py) -> PyResult> { Py::new( py, - Execution::new(RetainingHost { - retained: Some(retained), - }), + Execution::new( + RetainingHost { + retained: Some(retained), + }, + lifecycle_binding, + ), ) } @@ -1770,7 +2293,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri #[pyfunction] fn await_execution(awaitable: Py) -> Execution { - Execution::new(AwaitBody(Some(awaitable))) + Execution::new(AwaitBody(Some(awaitable)), lifecycle_binding) } struct CallingBody(Py); @@ -1787,7 +2310,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri #[pyfunction] fn calling_execution(callback: Py) -> Execution { - Execution::new(CallingBody(callback)) + Execution::new(CallingBody(callback), lifecycle_binding) } struct ErrorBody(Option>); @@ -1808,7 +2331,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri #[pyfunction] fn error_execution(error: Bound<'_, PyBaseException>) -> Execution { - Execution::new(ErrorBody(Some(error.unbind()))) + Execution::new(ErrorBody(Some(error.unbind())), lifecycle_binding) } #[rstest::rstest] diff --git a/litellm-rust/crates/host-python/src/handle.rs b/litellm-rust/crates/host-python/src/handle.rs index 91d003f32b9..cca9059341a 100644 --- a/litellm-rust/crates/host-python/src/handle.rs +++ b/litellm-rust/crates/host-python/src/handle.rs @@ -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>; + pub enum ExecutionStep { Return(Py), Await(Py), @@ -29,22 +31,21 @@ enum ExecutionState { #[pyclass] pub struct Execution { state: ExecutionState, -} - -fn lifecycle(py: Python<'_>) -> PyResult> { - 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> { + 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, ) -> PyResult> { - 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 { diff --git a/litellm-rust/crates/host-python/src/hooks.rs b/litellm-rust/crates/host-python/src/hooks.rs index 5a0f95d9c89..bab628bb7c6 100644 --- a/litellm-rust/crates/host-python/src/hooks.rs +++ b/litellm-rust/crates/host-python/src/hooks.rs @@ -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 = fn(&mut L, Python<'_>, PyResult>) -> PyResult>; @@ -18,62 +15,23 @@ pub enum HookStep { 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, - }, - Failed { - timing: Timing, - origin: FailureOrigin, - error: &'a PyErr, - }, +pub struct PythonRuntime; + +impl HookRuntime for PythonRuntime { + type Context<'a> = Python<'a>; + type Arguments = Py; + type Response = Py; + type Chunk = Py; + type Error = PyErr; + type Step = HookStep; + + fn ready(value: T) -> Self::Step { + 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, - started_at: f64, - ) -> PyResult>>; +pub type PythonCallEvent<'a> = RuntimeCallEvent<'a, PythonRuntime>; - fn before_provider_request( - &mut self, - py: Python<'_>, - wire: Box, - context: &RequestContext, - ) -> PyResult>>; +pub trait PythonCallHooks: CallHooks + PythonOwned {} - fn transform_response( - &mut self, - py: Python<'_>, - response: Py, - timing: Timing, - ) -> PyResult>>; - - fn on_event(&mut self, py: Python<'_>, event: HookEvent<'_>) -> PyResult>; - - /// 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) -> PyResult<()>; -} +impl + PythonOwned> PythonCallHooks for H {} diff --git a/litellm-rust/crates/host-python/src/hooks/chain/adapter.rs b/litellm-rust/crates/host-python/src/hooks/chain/adapter.rs new file mode 100644 index 00000000000..87f77c8428e --- /dev/null +++ b/litellm-rust/crates/host-python/src/hooks/chain/adapter.rs @@ -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 { + Ready(T), + Await(Py), +} + +enum Continuation { + Arguments(HookResume>), + Wire(HookResume>), + Response(HookResume>), + Event(HookResume), +} + +pub(super) struct HookAdapter { + hooks: H, + continuation: Option>, +} + +impl HookAdapter { + pub(super) fn new(hooks: H) -> Self { + Self { + hooks, + continuation: None, + } + } + + fn step( + &mut self, + step: HookStep, + continuation: impl FnOnce(HookResume) -> Continuation, + ) -> ChainStep { + 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, + started_at: f64, + ) -> PyResult>>; + fn resume_arguments( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult>>; + fn before_provider_request( + &mut self, + py: Python<'_>, + wire: Box, + context: &RequestContext, + ) -> PyResult>>; + fn resume_wire( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult>>; + fn transform_response( + &mut self, + py: Python<'_>, + response: Py, + timing: Timing, + ) -> PyResult>>; + fn resume_response( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult>>; + fn on_event(&mut self, py: Python<'_>, event: PythonCallEvent<'_>) -> PyResult>; + fn resume_event( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult>; + fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py) -> PyResult<()>; + fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()>; + fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py) -> PyResult<()>; +} + +impl ChainHooks for HookAdapter { + fn prepare_arguments( + &mut self, + py: Python<'_>, + arguments: Py, + started_at: f64, + ) -> PyResult>> { + 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>, + ) -> PyResult>> { + 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, + context: &RequestContext, + ) -> PyResult>> { + 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>, + ) -> PyResult>> { + 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, + timing: Timing, + ) -> PyResult>> { + let step = self.hooks.transform_response(py, response, timing)?; + Ok(self.step(step, Continuation::Response)) + } + + fn resume_response( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult>> { + 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> { + let step = self.hooks.on_event(py, event)?; + Ok(self.step(step, Continuation::Event)) + } + + fn resume_event( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult> { + 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) -> 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) -> PyResult<()> { + self.hooks.on_stream_chunk(py, chunk) + } +} + +impl PythonOwned for HookAdapter { + fn close(&mut self, py: Python<'_>) { + self.continuation = None; + self.hooks.close(py); + } + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + self.hooks.traverse(visit) + } +} diff --git a/litellm-rust/crates/host-python/src/hooks/chain/dispatch.rs b/litellm-rust/crates/host-python/src/hooks/chain/dispatch.rs new file mode 100644 index 00000000000..2f7ed0d778a --- /dev/null +++ b/litellm-rust/crates/host-python/src/hooks/chain/dispatch.rs @@ -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, RawResponse>; + +#[derive(Default)] +pub struct HookChain { + hooks: Vec>, + 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, + started_at: f64, + ) -> PyResult>> { + 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>, + started_at: f64, + ) -> PyResult>> { + 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>, + ) -> PyResult>> { + 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, + context: &RequestContext, + ) -> PyResult>> { + 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>, + context: &RequestContext, + ) -> PyResult>> { + 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>, + ) -> PyResult>> { + 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, + timing: Timing, + ) -> PyResult>> { + 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>, + timing: Timing, + ) -> PyResult>> { + 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>, + ) -> PyResult>> { + 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> { + 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> { + 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>, + ) -> PyResult> { + 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>, +) -> PyResult>> { + match result { + Err(error) if !error.is_instance_of::(py) => Err(error), + result => Ok(result), + } +} + +fn is_terminal(event: &CallEvent) -> bool { + matches!( + event, + CallEvent::Succeeded { .. } | CallEvent::Failed { .. } + ) +} + +fn notification_result( + py: Python<'_>, + terminal: bool, + result: PyResult>, +) -> PyResult> { + match result { + Err(error) if terminal && error.is_instance_of::(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> { + 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 for HookChain { + fn prepare_arguments( + &mut self, + py: Python<'_>, + arguments: Py, + started_at: f64, + ) -> PyResult>> { + self.arguments_from(py, 0, arguments, started_at) + } + + fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py) -> PyResult<()> { + self.hooks + .iter_mut() + .try_for_each(|hooks| hooks.arguments_prepared(py, arguments)) + } + + fn before_provider_request( + &mut self, + py: Python<'_>, + wire: Box, + context: &RequestContext, + ) -> PyResult>> { + self.wire_from(py, 0, wire, context) + } + + fn transform_response( + &mut self, + py: Python<'_>, + response: Py, + timing: Timing, + ) -> PyResult>> { + self.response_from(py, 0, response, timing) + } + + fn on_event( + &mut self, + py: Python<'_>, + event: PythonCallEvent<'_>, + ) -> PyResult> { + 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) -> 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(()), + } + } +} diff --git a/litellm-rust/crates/host-python/src/hooks/chain/mod.rs b/litellm-rust/crates/host-python/src/hooks/chain/mod.rs new file mode 100644 index 00000000000..c6267c385ff --- /dev/null +++ b/litellm-rust/crates/host-python/src/hooks/chain/mod.rs @@ -0,0 +1,4 @@ +mod adapter; +mod dispatch; + +pub use dispatch::HookChain; diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index cbb7fe3507b..4de404e3624 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -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; diff --git a/litellm-rust/crates/host-python/src/runtime.rs b/litellm-rust/crates/host-python/src/runtime.rs index c948c15d8eb..9cb576dc156 100644 --- a/litellm-rust/crates/host-python/src/runtime.rs +++ b/litellm-rust/crates/host-python/src/runtime.rs @@ -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> { + 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( 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> { + 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::(py)); + }); + } + #[derive(Debug)] struct Error(String); diff --git a/litellm-rust/crates/host-python/tests/hook_chain.rs b/litellm-rust/crates/host-python/tests/hook_chain.rs new file mode 100644 index 00000000000..c1efa941227 --- /dev/null +++ b/litellm-rust/crates/host-python/tests/hook_chain.rs @@ -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, + asynchronous: bool, + wire: Option>, +} + +impl ScriptHooks { + fn invoke(&self, py: Python<'_>, name: &str, value: Py) -> PyResult> { + 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>, + ) -> PyResult>> { + Ok(HookStep::Ready( + result?.into_bound(py).cast_into::()?.unbind(), + )) + } + + fn response( + &mut self, + _: Python<'_>, + result: PyResult>, + ) -> PyResult>> { + result.map(HookStep::Ready) + } + + fn notification( + &mut self, + _: Python<'_>, + result: PyResult>, + ) -> PyResult> { + result.map(|_| HookStep::Ready(())) + } + + fn request( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult>> { + 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 for ScriptHooks { + fn prepare_arguments( + &mut self, + py: Python<'_>, + arguments: Py, + _: f64, + ) -> PyResult>> { + 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) -> PyResult<()> { + self.object.call_method1(py, "adopt", (arguments,))?; + Ok(()) + } + + fn before_provider_request( + &mut self, + py: Python<'_>, + wire: Box, + context: &RequestContext, + ) -> PyResult>> { + 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, + _: Timing, + ) -> PyResult>> { + 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> { + 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) -> PyResult<()> { + self.object.call_method1(py, "stream", (chunk,))?; + Ok(()) + } +} + +#[fixture] +fn scripts() -> Py { + 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, 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(py: Python<'_>, hooks: &mut H, step: HookStep) -> PyResult { + 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 for PreparedPolicy { + fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py) -> 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, + #[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, + #[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, #[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, #[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::(py) + ); + py.run( + c"assert len(log) == 1 and log[0][0] == 'a'", + Some(locals), + Some(locals), + ) + .unwrap(); + }); +} + +struct Preparing { + hooks: HookChain, + arguments: Option>, + resume: Option>>, +} + +impl litellm_host_python::ExecutionBody for Preparing { + fn resume( + &mut self, + result: Option>>, + ) -> PyResult { + 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> { + 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) { + 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, + #[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, + #[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::() + .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(); + }); +} diff --git a/litellm-rust/crates/host/AGENTS.md b/litellm-rust/crates/host/AGENTS.md index 47f3a15575b..e2034e76c6f 100644 --- a/litellm-rust/crates/host/AGENTS.md +++ b/litellm-rust/crates/host/AGENTS.md @@ -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` | 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 diff --git a/litellm-rust/crates/host/src/call.rs b/litellm-rust/crates/host/src/call.rs index 6468cf319ae..d8187a216e5 100644 --- a/litellm-rust/crates/host/src/call.rs +++ b/litellm-rust/crates/host/src/call.rs @@ -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 { @@ -37,23 +38,34 @@ pub type OutputOf

= CallOutput< pub type HostedMachine

= CallMachine::Response>>; -pub fn hosted_call(request: P::Request, execute: F) -> HostedMachine

+pub fn hosted_call( + request: P::Request, + observers: Option, + execute: F, +) -> HostedMachine

where P: Protocol, P::Error: From, - F: FnOnce(P::Request, HostServices

, ChannelHooks

) -> Fut + Send + 'static, + F: FnOnce( + P::Request, + HostServices

, + ChannelInterceptors

, + Option, + ) -> Fut + + Send + + 'static, Fut: Future, 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); } } diff --git a/litellm-rust/crates/host/src/event.rs b/litellm-rust/crates/host/src/event.rs deleted file mode 100644 index dc888ddf7e7..00000000000 --- a/litellm-rust/crates/host/src/event.rs +++ /dev/null @@ -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, - /// The credential the route resolved for the provider call. - pub api_key: Option, -} - -#[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, - }, -} diff --git a/litellm-rust/crates/host/src/hooks.rs b/litellm-rust/crates/host/src/hooks.rs index fdf6e05bfa9..4444aeec551 100644 --- a/litellm-rust/crates/host/src/hooks.rs +++ b/litellm-rust/crates/host/src/hooks.rs @@ -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; -/// 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: Send + Sync { - fn observer(&self) -> Option> { - None + fn ready(value: T) -> Self::Step; +} + +pub type RuntimeCallEvent<'a, R> = + CallEvent<&'a ::Response, &'a ::Error, &'a RawResponse>; + +pub trait CallHooks: Sized { + fn prepare_arguments( + &mut self, + _runtime: R::Context<'_>, + arguments: R::Arguments, + _started_at: f64, + ) -> Result, 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> + Send; - - fn on_event(&self, event: MachineEvent) -> impl Future> + Send; -} - -/// No host: the wire request goes out as prepared and nothing observes the call. -impl RouteHooks for () { - async fn before_provider_request( - &self, - wire: WireRequest, - _: RequestContext, - ) -> Result { - Ok(wire) + &mut self, + _runtime: R::Context<'_>, + wire: Box, + _context: &RequestContext, + ) -> Result>, 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::Error> { + Ok(R::ready(response)) + } + + fn on_event( + &mut self, + _runtime: R::Context<'_>, + _event: RuntimeCallEvent<'_, R>, + ) -> Result, 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 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::::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::::before_provider_request(&(), wire("prepared"), context()) - .await - .unwrap(); - assert_eq!(sent.url, "prepared"); - } -} diff --git a/litellm-rust/crates/host/src/in_process.rs b/litellm-rust/crates/host/src/in_process.rs deleted file mode 100644 index ab2025fb289..00000000000 --- a/litellm-rust/crates/host/src/in_process.rs +++ /dev/null @@ -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: Send + Sync { - fn open_stream( - &self, - head: P::StreamHead, - ) -> impl Future> + Send; - fn send_chunk(&self, chunk: P::Chunk) -> impl Future> + Send; -} - -impl StreamConsumer

for () { - async fn open_stream(&self, _: P::StreamHead) -> Result { - Ok(Demand::More) - } - async fn send_chunk(&self, _: P::Chunk) -> Result { - 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( - machine: M, - host: Host<'_, S, H, C>, -) -> Result::Error> -where - M: Machine, - S: HostCallHandler, - H: RouteHooks<::Error>, - C: StreamConsumer, -{ - run_with_completion(machine, host, |_| false).await -} - -pub async fn run_hosted( - machine: crate::call::HostedMachine

, - host: Host<'_, S, H, C>, -) -> Result, P::Error> -where - P: Protocol, - P::Error: From, - S: HostCallHandler

, - H: RouteHooks, - C: StreamConsumer

, -{ - run_with_completion(machine, host, |completion| { - matches!(completion, crate::call::HostedCompletion::Detached) - }) - .await -} - -async fn run_with_completion( - mut machine: M, - host: Host<'_, S, H, C>, - detached: impl Fn(&M::Complete) -> bool, -) -> Result::Error> -where - M: Machine, - S: HostCallHandler, - H: RouteHooks<::Error>, - C: StreamConsumer, -{ - 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 for &'static str { - fn from(_: MachineFault) -> Self { - "machine fault" - } - } - - #[derive(Default)] - struct Recording { - seen: Mutex>, - 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 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<::Error> for Recording { - async fn before_provider_request( - &self, - wire: crate::event::WireRequest, - _: crate::event::RequestContext, - ) -> Result::Error> { - Ok(wire) - } - async fn on_event( - &self, - event: crate::event::MachineEvent, - ) -> Result<(), ::Error> { - crate::lifecycle::CallObserver::observe(self, crate::event::CallEvent::Machine(event)); - Ok(()) - } - } - - fn scripted( - ops: &'static [&'static str], - outcome: Result<(), &'static str>, - ) -> CallMachine { - 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>); - - 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 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<::Error> for StartTimes { - async fn before_provider_request( - &self, - wire: crate::event::WireRequest, - _: crate::event::RequestContext, - ) -> Result::Error> { - Ok(wire) - } - async fn on_event( - &self, - event: crate::event::MachineEvent, - ) -> Result<(), ::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]); - } -} diff --git a/litellm-rust/crates/host/src/interceptors.rs b/litellm-rust/crates/host/src/interceptors.rs new file mode 100644 index 00000000000..0044e842e4a --- /dev/null +++ b/litellm-rust/crates/host/src/interceptors.rs @@ -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, + /// The credential the route resolved for the provider call. + pub api_key: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct RawResponse { + pub body: String, +} + +pub trait Interceptors: Send + Sync { + fn before_provider_request( + &self, + wire: WireRequest, + context: RequestContext, + ) -> impl Future> + Send; + + fn after_provider_response( + &self, + raw: RawResponse, + ) -> impl Future> + Send; +} + +impl + ?Sized> Interceptors for &T { + fn before_provider_request( + &self, + wire: WireRequest, + context: RequestContext, + ) -> impl Future> + Send { + (**self).before_provider_request(wire, context) + } + + fn after_provider_response( + &self, + raw: RawResponse, + ) -> impl Future> + Send { + (**self).after_provider_response(raw) + } +} + +impl Interceptors for () { + async fn before_provider_request( + &self, + wire: WireRequest, + _: RequestContext, + ) -> Result { + 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 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::::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::::before_provider_request(&(), wire("prepared"), context()) + .await + .unwrap(); + assert_eq!(sent.url, "prepared"); + } +} diff --git a/litellm-rust/crates/host/src/lib.rs b/litellm-rust/crates/host/src/lib.rs index da130002449..76fce582fb3 100644 --- a/litellm-rust/crates/host/src/lib.rs +++ b/litellm-rust/crates/host/src/lib.rs @@ -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; diff --git a/litellm-rust/crates/host/src/lifecycle.rs b/litellm-rust/crates/host/src/lifecycle.rs index 25ea4bff8fc..f16a99fbcef 100644 --- a/litellm-rust/crates/host/src/lifecycle.rs +++ b/litellm-rust/crates/host/src/lifecycle.rs @@ -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 { + Started { + start_time: f64, + }, + Execution(ExecutionEvent), + Succeeded { + timing: Timing, + response: Response, + }, + Failed { + timing: Timing, + origin: FailureOrigin, + error: Error, + }, + Cancelled { + timing: Timing, + }, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum ExecutionEvent { + ProviderResponseReceived { raw: Raw }, +} + +impl> CallEvent { + 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>, + observers: Option, started_at: f64, } impl CallGuard { - fn new(observer: Arc) -> 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( - observer: Option>, + observers: Option, execute: impl Future, E>>, ) -> Result, 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( - observer: Option>, + observers: Option, execute: impl Future>, ) -> Result { - 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 diff --git a/litellm-rust/crates/host/src/machine/AGENTS.md b/litellm-rust/crates/host/src/machine/AGENTS.md new file mode 100644 index 00000000000..048f9aacf3a --- /dev/null +++ b/litellm-rust/crates/host/src/machine/AGENTS.md @@ -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 diff --git a/litellm-rust/crates/host/src/machine/auth.rs b/litellm-rust/crates/host/src/machine/auth.rs deleted file mode 100644 index 21ba1dc570f..00000000000 --- a/litellm-rust/crates/host/src/machine/auth.rs +++ /dev/null @@ -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) -> 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 { - channel: HostServices, -} - -impl std::fmt::Debug for HostTokenProvider { - fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - formatter.write_str("HostTokenProvider") - } -} - -impl HostTokenProvider -where - R: TokenProtocol, - R::Error: From + std::fmt::Display, -{ - pub fn handle(channel: HostServices) -> TokenProviderHandle { - TokenProviderHandle::new(Arc::new(Self { channel })) - } -} - -impl TokenProvider for HostTokenProvider -where - R: TokenProtocol, - R::Error: From + 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())) - }) - } -} diff --git a/litellm-rust/crates/host/src/machine/call_machine.rs b/litellm-rust/crates/host/src/machine/call_machine.rs deleted file mode 100644 index 3d57b024a55..00000000000 --- a/litellm-rust/crates/host/src/machine/call_machine.rs +++ /dev/null @@ -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::Response> = - Pin::Error>> + Send>>; - -pub struct CallContext { - pub services: HostServices

, - pub hooks: ChannelHooks

, - pub stream: StreamSender

, -} - -struct Channel(Co>); - -impl Clone for Channel

{ - fn clone(&self) -> Self { - Self(self.0.clone()) - } -} - -impl Channel

-where - P::Error: From, -{ - async fn request_reply( - &self, - request: impl FnOnce(Reply) -> Suspension

+ Send, - ) -> Result { - self.0 - .yield_(request) - .await - .map_err(|_| MachineFault::Abandoned.into()) - } -} - -pub struct HostServices(Channel

); - -impl Clone for HostServices

{ - fn clone(&self) -> Self { - Self(self.0.clone()) - } -} - -impl HostServices

-where - P::Error: From, -{ - pub async fn call( - &self, - request: impl FnOnce(Reply) -> P::HostCall + Send, - ) -> Result { - self.0 - .request_reply(|reply| Suspension::HostCall(request(reply))) - .await - } -} - -pub struct ChannelHooks(Channel

); - -impl Clone for ChannelHooks

{ - fn clone(&self) -> Self { - Self(self.0.clone()) - } -} - -impl crate::hooks::RouteHooks for ChannelHooks

-where - P::Error: From, -{ - async fn before_provider_request( - &self, - wire: WireRequest, - context: RequestContext, - ) -> Result { - 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(Channel

); - -impl StreamSender

-where - P::Error: From, -{ - pub async fn open_stream(&self, head: P::StreamHead) -> Result { - self.0 - .request_reply(|reply| Suspension::Stream(StreamDelivery::Open(head, reply))) - .await - } - - pub async fn send_chunk(&self, chunk: P::Chunk) -> Result { - self.0 - .request_reply(|reply| Suspension::Stream(StreamDelivery::Chunk(chunk, reply))) - .await - } -} - -type CallCoroutine = Coroutine, Result::Error>>; - -pub struct CallMachine::Response> { - coroutine: CallCoroutine, -} - -impl CallMachine -where - R::Error: From, -{ - pub fn new( - execute: impl FnOnce(CallContext) -> ExecuteFuture + 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 Machine for CallMachine -where - R::Error: From, -{ - 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) -> Interrupted<'_, Self> { - self.coroutine.cancel(); - Box::pin(async move { Err(failure.into_error()) }) - } -} diff --git a/litellm-rust/crates/host/src/machine/context.rs b/litellm-rust/crates/host/src/machine/context.rs new file mode 100644 index 00000000000..39b2b82d502 --- /dev/null +++ b/litellm-rust/crates/host/src/machine/context.rs @@ -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 { + pub services: HostServices

, + pub interceptors: ChannelInterceptors

, + pub stream: StreamSender

, + pub observers: Option, +} + +impl CallContext

{ + pub(super) fn new(co: Co>, observers: Option) -> Self { + let channel = Channel(co); + Self { + services: HostServices(channel.clone()), + interceptors: ChannelInterceptors(channel.clone()), + stream: StreamSender(channel), + observers, + } + } +} + +struct Channel(Co>); + +impl Clone for Channel

{ + fn clone(&self) -> Self { + Self(self.0.clone()) + } +} + +impl Channel

+where + P::Error: From, +{ + async fn request_reply( + &self, + request: impl FnOnce(Reply) -> HostRequest

+ Send, + ) -> Result { + self.0 + .yield_(request) + .await + .map_err(|_| MachineFault::Abandoned.into()) + } +} + +pub struct HostServices(Channel

); + +impl Clone for HostServices

{ + fn clone(&self) -> Self { + Self(self.0.clone()) + } +} + +impl HostServices

+where + P::Error: From, +{ + pub async fn call( + &self, + request: impl FnOnce(Reply) -> P::HostCall + Send, + ) -> Result { + self.0 + .request_reply(|reply| HostRequest::HostCall(request(reply))) + .await + } +} + +pub struct ChannelInterceptors(Channel

); + +impl Clone for ChannelInterceptors

{ + fn clone(&self) -> Self { + Self(self.0.clone()) + } +} + +impl crate::interceptors::Interceptors for ChannelInterceptors

+where + P::Error: From, +{ + async fn before_provider_request( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + 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(Channel

); + +impl StreamSender

+where + P::Error: From, +{ + pub async fn open_stream(&self, head: P::StreamHead) -> Result, P::Error> { + self.0 + .request_reply(|reply| HostRequest::Stream(StreamDelivery::Open(head, reply))) + .await + } + + pub async fn send_chunk(&self, chunk: P::Chunk) -> Result, P::Error> { + self.0 + .request_reply(|reply| HostRequest::Stream(StreamDelivery::Chunk(chunk, reply))) + .await + } +} diff --git a/litellm-rust/crates/host/src/machine/contract.rs b/litellm-rust/crates/host/src/machine/contract.rs new file mode 100644 index 00000000000..cfb602bdb5a --- /dev/null +++ b/litellm-rust/crates/host/src/machine/contract.rs @@ -0,0 +1,60 @@ +use std::{future::Future, pin::Pin}; + +use crate::protocol::{HostRequest, Protocol}; + +pub enum MachineStep { + Suspended(HostRequest), + Complete(C), +} + +pub type Step<'a, M> = Pin< + Box< + dyn Future< + Output = Result< + MachineStep<::Protocol, ::Complete>, + <::Protocol as Protocol>::Error, + >, + > + Send + + 'a, + >, +>; + +pub type Interrupted<'a, M> = Pin< + Box< + dyn Future< + Output = Result< + ::Complete, + <::Protocol as Protocol>::Error, + >, + > + Send + + 'a, + >, +>; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum HostFailure { + Error(E), + Cancelled(E), +} + +impl HostFailure { + 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<::Error>, + ) -> Interrupted<'_, Self>; +} diff --git a/litellm-rust/crates/host/src/machine/coroutine.rs b/litellm-rust/crates/host/src/machine/coroutine.rs new file mode 100644 index 00000000000..825536e1d9b --- /dev/null +++ b/litellm-rust/crates/host/src/machine/coroutine.rs @@ -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::Response> = + Pin::Error>> + Send>>; + +type CallCoroutine = Coroutine, Result::Error>>; + +pub struct CallMachine::Response> { + coroutine: CallCoroutine, +} + +impl CallMachine +where + R::Error: From, +{ + pub fn new( + observers: Option, + execute: impl FnOnce(CallContext) -> ExecuteFuture + Send + 'static, + ) -> Self { + Self { + coroutine: Coroutine::new(move |co| execute(CallContext::new(co, observers))), + } + } +} + +impl Machine for CallMachine +where + R::Error: From, +{ + 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) -> Interrupted<'_, Self> { + self.coroutine.cancel(); + Box::pin(async move { Err(failure.into_error()) }) + } +} diff --git a/litellm-rust/crates/host/src/machine/mod.rs b/litellm-rust/crates/host/src/machine/mod.rs index d307c9cec7d..fda876cee4a 100644 --- a/litellm-rust/crates/host/src/machine/mod.rs +++ b/litellm-rust/crates/host/src/machine/mod.rs @@ -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 { - Suspended(Suspension), - Complete(C), -} - -pub type Step<'a, M> = Pin< - Box< - dyn Future< - Output = Result< - MachineStep<::Protocol, ::Complete>, - <::Protocol as Protocol>::Error, - >, - > + Send - + 'a, - >, ->; - -pub type Interrupted<'a, M> = Pin< - Box< - dyn Future< - Output = Result< - ::Complete, - <::Protocol as Protocol>::Error, - >, - > + Send - + 'a, - >, ->; - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum HostFailure { - Error(E), - Cancelled(E), -} - -impl HostFailure { - 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<::Error>, - ) -> Interrupted<'_, Self>; -} +pub use context::{CallContext, ChannelInterceptors, HostServices, StreamSender}; +pub use contract::{HostFailure, Interrupted, Machine, MachineStep, Step}; +pub use coroutine::{CallMachine, ExecuteFuture, MachineFault}; diff --git a/litellm-rust/crates/host/src/observation.rs b/litellm-rust/crates/host/src/observation.rs new file mode 100644 index 00000000000..904855dd6d0 --- /dev/null +++ b/litellm-rust/crates/host/src/observation.rs @@ -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, + dropped: Arc, +} + +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) { + let (sender, receiver) = mpsc::channel(capacity.get()); + ( + ObservationSender { + sender, + dropped: Arc::new(AtomicU64::new(0)), + }, + receiver, + ) +} diff --git a/litellm-rust/crates/host/src/protocol.rs b/litellm-rust/crates/host/src/protocol.rs index 1715d185c04..c9db375cde3 100644 --- a/litellm-rust/crates/host/src/protocol.rs +++ b/litellm-rust/crates/host/src/protocol.rs @@ -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 { +pub enum HostRequest { HostCall(P::HostCall), - Hook(HookRequest), + Intercept(InterceptRequest), Stream(StreamDelivery

), } -pub enum HookRequest { +pub enum InterceptRequest { BeforeProviderRequest { wire: Box, context: Box, reply: Reply, }, - Event(MachineEvent, Reply<()>), + AfterProviderResponse { + raw: RawResponse, + reply: Reply<()>, + }, } pub enum StreamDelivery { - Open(P::StreamHead, Reply), - Chunk(P::Chunk, Reply), -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum Demand { - More, - Detached, + Open(P::StreamHead, Reply>), + Chunk(P::Chunk, Reply>), } diff --git a/litellm-rust/crates/host/tests/call.rs b/litellm-rust/crates/host/tests/call.rs index 3d94f9046e4..a8ba71e7b07 100644 --- a/litellm-rust/crates/host/tests/call.rs +++ b/litellm-rust/crates/host/tests/call.rs @@ -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::(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::(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>); +struct Observer(Observations); +struct Observations { + sender: litellm_host::observation::ObservationSender, + receiver: Mutex>, + recorded: Mutex>, +} + +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>> { + 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, #[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, #[case] fa #[tokio::test] async fn streams_finish_only_when_consumed(observer: Arc, #[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::(CallOutput::Stream { head: (), chunks }) }) .await @@ -161,7 +191,7 @@ async fn streams_finish_only_when_consumed(observer: Arc, #[case] fail #[tokio::test] async fn dropping_a_stream_cancels_without_success(observer: Arc) { let chunks = stream::pending().boxed(); - let output = observe_call(Some(observer.clone()), async { + let output = observe_call(Some(observer.0.sender.clone()), async { Ok::(CallOutput::Stream { head: (), chunks }) }) .await @@ -176,7 +206,7 @@ async fn dropping_a_stream_cancels_without_success(observer: Arc) { #[tokio::test] async fn cancelling_provider_execution_releases_the_lifecycle(observer: Arc) { let mut call = Box::pin(observe_unary( - Some(observer.clone()), + Some(observer.0.sender.clone()), std::future::pending::>(), )); assert!(futures_util::poll!(&mut call).is_pending()); @@ -185,40 +215,3 @@ async fn cancelling_provider_execution_releases_the_lifecycle(observer: Arc) { - struct DetachingConsumer; - impl litellm_host::in_process::StreamConsumer for DetachingConsumer { - async fn open_stream(&self, _: &'static str) -> Result { - Ok(Demand::Detached) - } - async fn send_chunk(&self, _: usize) -> Result { - panic!("detached consumers must not receive chunks") - } - } - - let machine = hosted_call::(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 { .. }] - )); -} diff --git a/litellm-rust/crates/host/tests/observation.rs b/litellm-rust/crates/host/tests/observation.rs new file mode 100644 index 00000000000..22872047d81 --- /dev/null +++ b/litellm-rust/crates/host/tests/observation.rs @@ -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>> = 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::::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); +} diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs index 6b34fc3ef7c..f34af90df0e 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs @@ -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}, diff --git a/litellm-rust/crates/python-bridge/AGENTS.md b/litellm-rust/crates/python-bridge/AGENTS.md index 507224e3446..162b0bdb07b 100644 --- a/litellm-rust/crates/python-bridge/AGENTS.md +++ b/litellm-rust/crates/python-bridge/AGENTS.md @@ -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` - 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 diff --git a/litellm-rust/crates/python-bridge/src/cache/AGENTS.md b/litellm-rust/crates/python-bridge/src/cache/AGENTS.md new file mode 100644 index 00000000000..8c6f31a7780 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/AGENTS.md @@ -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 diff --git a/litellm-rust/crates/python-bridge/src/cache/future.rs b/litellm-rust/crates/python-bridge/src/cache/future.rs index 42593eee1f4..242a23242cc 100644 --- a/litellm-rust/crates/python-bridge/src/cache/future.rs +++ b/litellm-rust/crates/python-bridge/src/cache/future.rs @@ -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> { @@ -9,10 +9,5 @@ pub(super) fn ready_value<'py, T: serde::Serialize>( py: Python<'py>, value: &T, ) -> PyResult> { - 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)) } diff --git a/litellm-rust/crates/python-bridge/src/cache/semantic.rs b/litellm-rust/crates/python-bridge/src/cache/semantic.rs index f05635c03df..4a1cc0dfe8e 100644 --- a/litellm-rust/crates/python-bridge/src/cache/semantic.rs +++ b/litellm-rust/crates/python-bridge/src/cache/semantic.rs @@ -207,5 +207,5 @@ impl ExecutionBody for SemanticExecution { } pub(super) fn drive(py: Python<'_>, body: SemanticExecution) -> PyResult> { - Execution::new(body).into_coroutine(py) + Execution::new(body, crate::lifecycle::binding).into_coroutine(py) } diff --git a/litellm-rust/crates/host-python/src/callable.rs b/litellm-rust/crates/python-bridge/src/callable.rs similarity index 94% rename from litellm-rust/crates/host-python/src/callable.rs rename to litellm-rust/crates/python-bridge/src/callable.rs index 2e454422e95..d0ed76a3e8b 100644 --- a/litellm-rust/crates/host-python/src/callable.rs +++ b/litellm-rust/crates/python-bridge/src/callable.rs @@ -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(py: Python<'_>, template: &str, result: PyResult) -> PyResult { +pub(crate) fn wrap_failure(py: Python<'_>, template: &str, result: PyResult) -> PyResult { result.map_err(|error| { if error.is_instance_of::(py) || !error.is_instance_of::(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); }); diff --git a/litellm-rust/crates/python-bridge/src/coercion.rs b/litellm-rust/crates/python-bridge/src/coercion.rs index 32a474ccec3..77c8cabc22b 100644 --- a/litellm-rust/crates/python-bridge/src/coercion.rs +++ b/litellm-rust/crates/python-bridge/src/coercion.rs @@ -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, ()>, #[case] tuning: Result, ()>, ) { - Python::initialize(); - Python::attach(|py| { + with_python(|py| { let field = Field::new("test", "string", evaluate(py, source)); let owned = |expected: Result, ()>| expected.map(|value| value.map(str::to_owned)); @@ -351,8 +354,7 @@ mod tests { #[case] source: &str, #[case] expected: Option, ) { - 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'}]", diff --git a/litellm-rust/crates/python-bridge/src/credentials.rs b/litellm-rust/crates/python-bridge/src/credentials.rs index 775e21d686d..04d8e94c8fe 100644 --- a/litellm-rust/crates/python-bridge/src/credentials.rs +++ b/litellm-rust/crates/python-bridge/src/credentials.rs @@ -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}, diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index e74e3556dd7..b90b313e799 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -1,9 +1,11 @@ mod cache; +mod callable; mod coercion; mod credentials; mod diagnostics; mod errors; mod http; +mod lifecycle; mod logger; mod marshal; mod preflight; diff --git a/litellm-rust/crates/python-bridge/src/lifecycle.rs b/litellm-rust/crates/python-bridge/src/lifecycle.rs new file mode 100644 index 00000000000..b69976278d4 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/lifecycle.rs @@ -0,0 +1,10 @@ +use litellm_host_python::CallOptions; +use pyo3::prelude::*; + +pub(crate) fn binding(py: Python<'_>) -> PyResult> { + py.import("litellm.rust_bridge.streams") +} + +pub(crate) fn call_options(asynchronous: bool) -> CallOptions { + CallOptions::new(asynchronous, binding) +} diff --git a/litellm-rust/crates/python-bridge/src/preflight.rs b/litellm-rust/crates/python-bridge/src/preflight.rs index 34813672c09..bc692498880 100644 --- a/litellm-rust/crates/python-bridge/src/preflight.rs +++ b/litellm-rust/crates/python-bridge/src/preflight.rs @@ -1,9 +1,7 @@ -//! The SDK's request policy the driver runs on every route's keyword view before the host -//! projects from it: credential-name inheritance from `litellm.credential_list`, then the -//! budget and retry-count limits. It is the `@client` prologue after `function_setup` and the -//! deployment hook, and belongs to no callback contract. - +use litellm_host::hooks::CallHooks; +use litellm_host_python::{PythonOwned, PythonRuntime}; use pyo3::{ + gc::{PyTraverseError, PyVisit}, prelude::*, types::{PyDict, PyList}, }; @@ -35,16 +33,26 @@ impl PythonPreflight { #[cfg(test)] pub(crate) const PYTHON_CONTRACT: &str = include_str!("../preflight_contract.json"); -/// Rewrites `arguments` in place, in the order the Python wrapper runs: credentials first, -/// so the limits see the same view the provider request is built from. -pub(crate) fn sdk_preflight(py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<()> { - inherit_credentials(py, arguments, || { - Ok(PythonPreflight::CredentialList - .call(py, ())? - .cast_into::()?) - })?; - PythonPreflight::CheckLimits.call(py, (arguments,))?; - Ok(()) +pub(crate) struct SdkPolicy; + +impl CallHooks for SdkPolicy { + fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py) -> PyResult<()> { + inherit_credentials(py, arguments.bind(py), || { + Ok(PythonPreflight::CredentialList + .call(py, ())? + .cast_into::()?) + })?; + PythonPreflight::CheckLimits.call(py, (arguments,))?; + Ok(()) + } +} + +impl PythonOwned for SdkPolicy { + fn close(&mut self, _: Python<'_>) {} + + fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { + Ok(()) + } } struct CredentialEntry<'py>(Bound<'py, PyAny>); @@ -400,7 +408,7 @@ arguments = {'litellm_credential_name': 'ocr-test'} }); } - #[test] + #[rstest::rstest] fn an_unknown_name_is_reported_with_the_loaded_count_and_leaves_the_arguments_alone() { let _guard = PREFLIGHT_MODULE .lock() @@ -417,7 +425,9 @@ preflight.credential_list = lambda: [Credential(), Credential()] arguments = {'litellm_credential_name': 'missing'} ", ); - sdk_preflight(py, &argument_dict(&locals)).unwrap(); + SdkPolicy + .arguments_prepared(py, &argument_dict(&locals).unbind()) + .unwrap(); py.run( c" assert arguments == {'litellm_credential_name': 'missing'}, arguments @@ -431,7 +441,7 @@ assert preflight.checked == [arguments] }); } - #[test] + #[rstest::rstest] fn limits_are_checked_on_the_arguments_after_credentials_are_inherited() { let _guard = PREFLIGHT_MODULE .lock() @@ -453,7 +463,9 @@ preflight.check_limits = check_limits arguments = {'litellm_credential_name': 'ocr-test'} ", ); - let error = sdk_preflight(py, &argument_dict(&locals)).unwrap_err(); + let error = SdkPolicy + .arguments_prepared(py, &argument_dict(&locals).unbind()) + .unwrap_err(); assert!( error .value(py) diff --git a/litellm-rust/crates/python-bridge/src/routes/AGENTS.md b/litellm-rust/crates/python-bridge/src/routes/AGENTS.md new file mode 100644 index 00000000000..76578447bba --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/AGENTS.md @@ -0,0 +1,13 @@ +# Route boundary + +These invariants apply to lifecycle-bearing public calls; value-oriented APIs that return an already running Future retain their explicit contract + +Route modules own public argument projection, route-specific host operations, response and error construction, and composition of the core route with the Python host. Provider dispatch, transport execution and normalization belong to core and provider crates. Runtime waiting, cancellation mechanics and execution state validation belong to `litellm-host-python` + +Before execution starts, perform only admission checks needed to select native execution or legacy fallback. Do not fully project a request just to decide admission. Keep effectful settings reads, HTTP client acquisition, secret-source construction and caller context capture inside execution, after argument preparation and SDK preflight. Configure resources from that prepared view, never the original entrypoint kwargs + +The host driver owns sequencing and terminal events; the bridge supplies fallible resource composition without exposing route types to the driver. An unstarted async call performs no resource setup. Setup errors after start follow the terminal failure contract and never authorize fallback or provider replay + +Use the shared `run_public_call` boundary with hooks supplied by bridge composition. `callbacks-legacy-python` owns legacy argument sharing and `Logging` dispatch behind `PublicCall` and `LegacyLogging`. Route bindings identify their neutral `Operation` and may retain the request needed for projection, but must not duplicate the legacy callback contract + +Regression tests must observe that an unstarted call does no setup, hook and preflight rewrites affect resource configuration, setup failures reach the selected failure handler once, and provider work is not replayed. Retain existing read-point and object-identity guarantees while changing setup timing diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 0f9db3b6305..850d6526d7f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -44,6 +44,7 @@ async fn execute( timeout, }, &(), + None, ) .await } @@ -136,32 +137,33 @@ fn run_public( asynchronous: bool, ) -> PyResult> { use super::inference::InferenceHost; - use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; + use litellm_types::Operation; let host = InferenceHost::new( request.clone().unbind(), "litellm.rust_bridge.chat_completions.route_host", ); - let route = ChatCompletionsRoute::new( - crate::http::provider_client(py, &kwargs, asynchronous)? - .map_err(crate::http::client_error)?, - crate::http::resources().auth.clone(), - crate::secrets::source(py)?, - ); - run_legacy_call( + let (arguments, hooks) = crate::routes::call_hooks( py, - LegacySurface { - call_type: if asynchronous { - "acompletion" - } else { - "completion" - }, - input_description: "Chat completions", - stream: None, + Operation::Completion, + &request, + &args, + &kwargs, + asynchronous, + )?; + crate::routes::run_public_call( + py, + arguments, + move |py, arguments, request| { + let route = ChatCompletionsRoute::new( + crate::http::provider_client(py, arguments, asynchronous)? + .map_err(crate::http::client_error)?, + crate::http::resources().auth.clone(), + crate::secrets::source(py)?, + ); + Ok(route.machine(request, None)) }, - PublicCall::capture(&request, &args, &kwargs)?, - move |request| crate::logger::LoggedMachine::new(route.machine(request)), host::ChatCompletionsPythonHost(host), - crate::preflight::sdk_preflight, + hooks, asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/inference.rs b/litellm-rust/crates/python-bridge/src/routes/inference.rs index 13499bb13cb..42015802f01 100644 --- a/litellm-rust/crates/python-bridge/src/routes/inference.rs +++ b/litellm-rust/crates/python-bridge/src/routes/inference.rs @@ -33,32 +33,11 @@ impl InferenceHost { arguments: &Bound<'_, PyDict>, input: &str, ) -> PyResult { - let request = self.request.bind(py); - let argument = |name: &str| -> PyResult>> { - if let Some(value) = lookup(arguments, request, name)? { - return Ok((!value.is_none()).then_some(value)); - } - let parameter = request - .getattr("parameters")? - .call_method1("get", (name,))?; - if !parameter.is_none() { - return Ok(Some(parameter)); - } - let extra = request.getattr("kwargs")?.call_method1("get", (name,))?; - Ok((!extra.is_none()).then_some(extra)) - }; + let argument = |name: &str| self.argument(py, arguments, name); let string = |name: &str| -> PyResult> { argument(name)?.map(|value| value.extract()).transpose() }; - let names: Vec = py.import(self.module)?.getattr("PARAMETERS")?.extract()?; - let params = names - .iter() - .filter_map(|name| match argument(name) { - Ok(Some(value)) => Some(from_py(&value).map(|value| (name.clone(), value))), - Ok(None) => None, - Err(error) => Some(Err(error)), - }) - .collect::>>()?; + let params = self.parameters(py, arguments)?; let timeout = argument("timeout")? .or(argument("request_timeout")?) .map(|value| python_timeout_seconds(py, value.unbind())) @@ -97,6 +76,42 @@ impl InferenceHost { }) } + pub fn argument<'py>( + &self, + py: Python<'py>, + arguments: &Bound<'py, PyDict>, + name: &str, + ) -> PyResult>> { + let request = self.request.bind(py); + if let Some(value) = lookup(arguments, request, name)? { + return Ok((!value.is_none()).then_some(value)); + } + let parameter = request + .getattr("parameters")? + .call_method1("get", (name,))?; + if !parameter.is_none() { + return Ok(Some(parameter)); + } + let extra = request.getattr("kwargs")?.call_method1("get", (name,))?; + Ok((!extra.is_none()).then_some(extra)) + } + + pub fn parameters( + &self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> PyResult> { + let names: Vec = py.import(self.module)?.getattr("PARAMETERS")?.extract()?; + names + .iter() + .filter_map(|name| match self.argument(py, arguments, name) { + Ok(Some(value)) => Some(from_py(&value).map(|value| (name.clone(), value))), + Ok(None) => None, + Err(error) => Some(Err(error)), + }) + .collect() + } + pub fn response(&self, py: Python<'_>, response: &impl Serialize) -> PyResult> { py.import(self.module)? .getattr("response")? diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index f288b71dbde..d13be311b10 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -1,23 +1,12 @@ mod host; use host::MessagesPythonHost; -use litellm_callbacks_legacy_python::{ - LegacySurface, PassThroughStream, PublicCall, run_legacy_call, -}; +use litellm_types::Operation; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, }; -const SURFACE: LegacySurface = LegacySurface { - call_type: "anthropic_messages", - input_description: "Messages", - stream: Some(PassThroughStream { - url_route: "/v1/messages", - endpoint_type: "anthropic", - }), -}; - fn run_messages( py: Python<'_>, request: Bound<'_, PyAny>, @@ -25,19 +14,28 @@ fn run_messages( kwargs: Bound<'_, PyDict>, asynchronous: bool, ) -> PyResult> { - let route = litellm_core::messages::MessagesRoute::new( - crate::http::provider_client(py, &kwargs, asynchronous)? - .map_err(crate::http::client_error)?, - crate::http::resources().auth.clone(), - crate::secrets::source(py)?, - ); - run_legacy_call( + let (arguments, hooks) = crate::routes::call_hooks( py, - SURFACE, - PublicCall::capture(&request, &args, &kwargs)?, - move |request| crate::logger::LoggedMachine::new(route.machine(request)), + Operation::Messages, + &request, + &args, + &kwargs, + asynchronous, + )?; + crate::routes::run_public_call( + py, + arguments, + move |py, arguments, request| { + let route = litellm_core::messages::MessagesRoute::new( + crate::http::provider_client(py, arguments, asynchronous)? + .map_err(crate::http::client_error)?, + crate::http::resources().auth.clone(), + crate::secrets::source(py)?, + ); + Ok(route.machine(request, None)) + }, MessagesPythonHost::new(request.unbind()), - crate::preflight::sdk_preflight, + hooks, asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 313bbb3945c..6542983016e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -7,6 +7,65 @@ pub(crate) mod ocr; pub(crate) mod responses; pub(crate) mod token_counter; +use litellm_callbacks_legacy_python::{LegacyLogging, PublicCall}; +use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol}; +use litellm_host_python::{HookChain, PythonBinding, PythonCallHooks, PythonHostCalls}; +use litellm_types::Operation; +use pyo3::{ + prelude::*, + types::{PyDict, PyTuple}, +}; + +fn call_hooks( + py: Python<'_>, + operation: Operation, + request: &Bound<'_, PyAny>, + args: &Bound<'_, PyTuple>, + kwargs: &Bound<'_, PyDict>, + asynchronous: bool, +) -> PyResult<(Py, impl PythonCallHooks + use<>)> { + let call = PublicCall::capture(request, args, kwargs)?; + let arguments = call.arguments(py); + Ok(( + arguments, + LegacyLogging::new(py, operation, call, asynchronous), + )) +} + +fn run_public_call( + py: Python<'_>, + arguments: Py, + start: impl FnOnce( + Python<'_>, + &Bound<'_, PyDict>, + ::Request, + ) -> PyResult + + Send + + Sync + + 'static, + host: H, + hooks: impl PythonCallHooks + 'static, + asynchronous: bool, +) -> PyResult> +where + H: PythonBinding + PythonHostCalls + 'static, + M: Machine + 'static, + M::Complete: Into::Response>>, +{ + litellm_host_python::run_call( + py, + move |py, arguments, request| { + start(py, arguments, request).map(crate::logger::LoggedMachine::new) + }, + host, + HookChain::new() + .with(hooks) + .with(crate::preflight::SdkPolicy), + arguments, + crate::lifecycle::call_options(asynchronous), + ) +} + #[cfg(test)] mod tests { use pyo3::{ diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index 82ee3277f97..2c732c3b1a3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -4,11 +4,11 @@ mod host; mod project; use host::OcrPythonHost; -use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; use litellm_core::ocr::provider_config; use litellm_core_utils::settings::ProcessEnvironment; use litellm_host_python::to_py; use litellm_llms::base_llm::ocr::settings::OcrSettings; +use litellm_types::Operation; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, @@ -29,17 +29,6 @@ const ENABLE_AZURE_AD_TOKEN_REFRESH: FieldSpec = Ok(field.exact_true()) }); -const SURFACE: LegacySurface = LegacySurface { - call_type: "ocr", - input_description: "OCR document processing", - stream: None, -}; - -const ASYNC_SURFACE: LegacySurface = LegacySurface { - call_type: "aocr", - ..SURFACE -}; - fn run_ocr( py: Python<'_>, request: Bound<'_, PyAny>, @@ -47,24 +36,27 @@ fn run_ocr( kwargs: Bound<'_, PyDict>, asynchronous: bool, ) -> PyResult> { - let config = http::call_config(py, &kwargs, asynchronous)?; - let client = litellm_llms::base_llm::ocr::handler::OcrClient::new( - &http::resources().pool, - &config, - http::url_policy(py)?, - http::resources().auth.clone(), - ocr_settings(py)?, - crate::secrets::source(py)?, - ) - .map_err(http::client_error)?; - let route = litellm_core::ocr::OcrRoute::new(client); - run_legacy_call( + let (arguments, hooks) = + crate::routes::call_hooks(py, Operation::Ocr, &request, &args, &kwargs, asynchronous)?; + crate::routes::run_public_call( py, - if asynchronous { ASYNC_SURFACE } else { SURFACE }, - PublicCall::capture(&request, &args, &kwargs)?, - move |request| crate::logger::LoggedMachine::new(route.machine(request)), + arguments, + move |py, arguments, request| { + let config = http::call_config(py, arguments, asynchronous)?; + let client = litellm_llms::base_llm::ocr::handler::OcrClient::new( + &http::resources().pool, + &config, + http::url_policy(py)?, + http::resources().auth.clone(), + ocr_settings(py)?, + crate::secrets::source(py)?, + ) + .map_err(http::client_error)?; + let route = litellm_core::ocr::OcrRoute::new(client); + Ok(route.machine(request, None)) + }, OcrPythonHost::new(request.unbind()), - crate::preflight::sdk_preflight, + hooks, asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses.rs index 4cacb850f81..e65f74ec1b4 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses.rs @@ -20,7 +20,7 @@ fn run_public( asynchronous: bool, ) -> PyResult> { use super::inference::InferenceHost; - use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; + use litellm_types::Operation; let host = InferenceHost::new( request.clone().unbind(), "litellm.rust_bridge.responses.route_host", @@ -33,51 +33,58 @@ fn run_public( { return Err(RustBridgeDeclined::new_err(reason)); } - let admission = host::project(&host, py, &kwargs)?; - if admission - .custom_llm_provider + let model = host + .argument(py, &kwargs, "model")? + .ok_or_else(|| pyo3::exceptions::PyValueError::new_err("model is required"))? + .extract::()?; + let provider = host + .argument(py, &kwargs, "custom_llm_provider")? + .map(|value| value.extract::()) + .transpose()?; + if provider .as_deref() .is_some_and(|provider| provider != "openai") - || admission - .model + || model .strip_prefix("openai/") - .unwrap_or(&admission.model) + .unwrap_or(&model) .contains('/') { return Err(RustBridgeDeclined::new_err( "native HTTP responses provider", )); } - if admission - .optional_params - .get("stream") - .is_some_and(|value| value == &serde_json::Value::Bool(true)) + if host + .argument(py, &kwargs, "stream")? + .map(|value| litellm_host_python::from_py::(&value)) + .transpose()? + .is_some_and(|value| value == Value::Bool(true)) { return Err(RustBridgeDeclined::new_err( "native Python responses streaming", )); } - let route = litellm_core::responses::ResponsesRoute::new( - crate::http::provider_client(py, &kwargs, asynchronous)? - .map_err(crate::http::client_error)?, - crate::http::resources().auth.clone(), - crate::secrets::source(py)?, - ); - run_legacy_call( + let (arguments, hooks) = crate::routes::call_hooks( py, - LegacySurface { - call_type: if asynchronous { - "aresponses" - } else { - "responses" - }, - input_description: "Responses", - stream: None, + Operation::Responses, + &request, + &args, + &kwargs, + asynchronous, + )?; + crate::routes::run_public_call( + py, + arguments, + move |py, arguments, request| { + let route = litellm_core::responses::ResponsesRoute::new( + crate::http::provider_client(py, arguments, asynchronous)? + .map_err(crate::http::client_error)?, + crate::http::resources().auth.clone(), + crate::secrets::source(py)?, + ); + Ok(route.machine(request, None)) }, - PublicCall::capture(&request, &args, &kwargs)?, - move |request| crate::logger::LoggedMachine::new(route.machine(request)), host::ResponsesPythonHost(host), - crate::preflight::sdk_preflight, + hooks, asynchronous, ) } diff --git a/litellm-rust/crates/types/src/lib.rs b/litellm-rust/crates/types/src/lib.rs index baa729072f0..4460e60d51c 100644 --- a/litellm-rust/crates/types/src/lib.rs +++ b/litellm-rust/crates/types/src/lib.rs @@ -4,3 +4,11 @@ pub mod messages; pub mod recognized; pub mod responses; pub mod utils; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Operation { + Completion, + Responses, + Messages, + Ocr, +} diff --git a/litellm/rust_bridge/AGENTS.md b/litellm/rust_bridge/AGENTS.md new file mode 100644 index 00000000000..a3fa17217fc --- /dev/null +++ b/litellm/rust_bridge/AGENTS.md @@ -0,0 +1,13 @@ +# Python boundary + +This package owns native rollout and fallback selection, Python public API compatibility, settings projection and the Python bindings supplied to the Rust bridge. Rust core owns provider execution; `litellm-host-python` owns CPython runtime mechanics; `callbacks-legacy-python` owns legacy callback sharing and dispatch policy + +`lifecycle.py` owns generic inline execution and stream iteration. `streams.py` supplies the product binding and public stream wrappers through the bridge rather than let the generic Rust host import this package by name. Keep one driver implementing `start`, `resume_value`, `resume_error` and idempotent `close`; do not create a second implementation + +Generic execution steps and inline suspension handling must not depend on LiteLLM response metadata. Public stream construction, `_hidden_params` and header compatibility remain product responsibilities. Keep generic stream heads opaque and preserve public header behavior in the product wrappers + +Await each selected suspension in the caller's task so hooks retain thread, loop and context identity. A final awaitable value is returned as a value, not awaited implicitly. Closing, cancellation and `GeneratorExit` release the execution without replaying provider work or starting further callbacks + +Native admission may select legacy before execution starts. A failure after execution starts is terminal and never authorizes fallback. Creating an unstarted native coroutine must not acquire clients or credentials or capture execution context + +Verify cross-language behavior using a fresh installed extension, including unstarted-call release, setup ordering, exception identity, sync and async streams, cancellation and GC. Keep `_native.pyi` accurate about Future-returning and coroutine-returning APIs diff --git a/litellm/rust_bridge/lifecycle.py b/litellm/rust_bridge/lifecycle.py index 7a6485a5f2c..fe95e9ca6a2 100644 --- a/litellm/rust_bridge/lifecycle.py +++ b/litellm/rust_bridge/lifecycle.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import AsyncIterator, Awaitable, Iterator, Mapping +from collections.abc import AsyncIterator, Awaitable, Callable, Iterator from dataclasses import dataclass from typing import Final, Protocol @@ -17,7 +17,7 @@ class Complete: @dataclass(frozen=True, slots=True) class Open: - value: Mapping[str, object] | None + value: object @dataclass(frozen=True, slots=True) @@ -62,13 +62,15 @@ def _settled(step: Step) -> Settled: return step -async def drive(execution: Execution) -> object: +async def drive(execution: Execution, stream_factory: Callable[[Execution, object], object] | None = None) -> object: handed_off = False # rebind-ok: set once the execution belongs to the returned stream try: step: Final = await _settle(execution, execution.start()) if isinstance(step, Open): + factory: Final = Stream if stream_factory is None else stream_factory + stream: Final = factory(execution, step.value) handed_off = True - return Stream(execution, step.value) + return stream return step.value finally: if not handed_off: @@ -78,10 +80,10 @@ async def drive(execution: Execution) -> object: class Stream(AsyncIterator[object]): """A streamed native call: each read resumes the execution until its next chunk.""" - def __init__(self, execution: Execution, hidden_params: Mapping[str, object] | None = None) -> None: + def __init__(self, execution: Execution, head: object = None) -> None: self._execution: Final = execution self._done = False - self._hidden_params: dict[str, object] = dict(hidden_params or {}) # mutable-ok: header writers mutate it + self.head: Final = head def __aiter__(self) -> Stream: return self @@ -115,10 +117,10 @@ class Stream(AsyncIterator[object]): class SyncStream(Iterator[object]): """The sync form of `Stream`; its execution never suspends on an awaitable.""" - def __init__(self, execution: Execution, hidden_params: Mapping[str, object] | None = None) -> None: + def __init__(self, execution: Execution, head: object = None) -> None: self._execution: Final = execution self._done = False - self._hidden_params: dict[str, object] = dict(hidden_params or {}) # mutable-ok: header writers mutate it + self.head: Final = head def __iter__(self) -> SyncStream: return self diff --git a/litellm/rust_bridge/streams.py b/litellm/rust_bridge/streams.py new file mode 100644 index 00000000000..6ead17777d2 --- /dev/null +++ b/litellm/rust_bridge/streams.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +from collections.abc import AsyncIterator, Iterator, Mapping +from typing import Final + +from pydantic import TypeAdapter + +from litellm.rust_bridge import lifecycle +from litellm.rust_bridge.lifecycle import Execution + +Await: Final = lifecycle.Await +Complete: Final = lifecycle.Complete +Open: Final = lifecycle.Open +Yield: Final = lifecycle.Yield + +_HEADERS: Final = TypeAdapter(Mapping[str, object]) + + +async def drive(execution: Execution) -> object: + return await lifecycle.drive(execution, Stream) + + +class Stream(AsyncIterator[object]): + def __init__(self, execution: Execution, hidden_params: object = None) -> None: + self._stream: Final = lifecycle.Stream(execution, hidden_params) + self._hidden_params: dict[str, object] = dict(_headers(hidden_params)) # mutable-ok: header writers mutate it + + def __aiter__(self) -> Stream: + return self + + async def __anext__(self) -> object: + return await self._stream.__anext__() + + async def aclose(self) -> None: + await self._stream.aclose() + + +class SyncStream(Iterator[object]): + def __init__(self, execution: Execution, hidden_params: object = None) -> None: + self._stream: Final = lifecycle.SyncStream(execution, hidden_params) + self._hidden_params: dict[str, object] = dict(_headers(hidden_params)) # mutable-ok: header writers mutate it + + def __iter__(self) -> SyncStream: + return self + + def __next__(self) -> object: + return next(self._stream) + + def close(self) -> None: + self._stream.close() + + +def _headers(value: object) -> Mapping[str, object]: + if value is None: + return {} + return _HEADERS.validate_python(value) diff --git a/tests/test_litellm_rust/test_inference.py b/tests/test_litellm_rust/test_inference.py index 8f071c9aeda..f1b6071a844 100644 --- a/tests/test_litellm_rust/test_inference.py +++ b/tests/test_litellm_rust/test_inference.py @@ -8,12 +8,13 @@ from pydantic import JsonValue, TypeAdapter import litellm from litellm import RateLimitError from litellm.integrations.custom_logger import CustomLogger +from litellm.models.credentials import CredentialItem from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.rust_bridge import _native from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest from litellm.types.llms.openai import ResponsesAPIResponse -from litellm.types.utils import ModelResponse +from litellm.types.utils import CallTypes, ModelResponse from tests.test_litellm_rust.support.callback_recorder import RecordingLogger from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_MODEL, MESSAGES_RESPONSE, request_body @@ -137,6 +138,62 @@ async def test_native_inference_pre_call_edits_reach_the_provider( assert _OBJECT.validate_python(recording_server.requests[0].body)["temperature"] == 0.75 +@pytest.mark.asyncio +@pytest.mark.parametrize("from_credentials", (False, True)) +async def test_native_resource_setup_uses_deployment_hook_arguments( + route: Route, + recording_server: RecordingServer, + monkeypatch: pytest.MonkeyPatch, + from_credentials: bool, +) -> None: + invalid_settings: Final = {"ssl_verify": object()} + credential: Final = CredentialItem( + credential_name="resource-settings", credential_info={}, credential_values=invalid_settings + ) + monkeypatch.setattr(litellm, "credential_list", [credential]) + + class Prepare(CustomLogger): + async def async_pre_call_deployment_hook( + self, kwargs: dict[str, object], call_type: CallTypes | None + ) -> dict[str, object]: + return { + **kwargs, + **({"litellm_credential_name": credential.credential_name} if from_credentials else invalid_settings), + } + + litellm.callbacks.append(Prepare()) + recorder: Final = RecordingLogger() + recording_server.expected_requests = 0 + with pytest.raises(ValueError, match=r"request\.ssl_verify") as caught: + await execute(route, True, recording_server, {"callbacks": [recorder]}) + failure: Final = await recorder.wait_for_async("async_log_failure_event") + assert len(failure) == 1 + assert _OBJECT.validate_python(failure[0].kwargs)["exception"] is caught.value + assert not recording_server.requests + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_sdk_policy_rejection_precedes_resource_setup_and_is_logged_once( + route: Route, + asynchronous: bool, + recording_server: RecordingServer, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "max_budget", 1.0) + monkeypatch.setattr(litellm, "_current_cost", 2.0) + monkeypatch.setattr(litellm, "ssl_verify", object()) + recorder: Final = RecordingLogger() + recording_server.expected_requests = 0 + with pytest.raises(litellm.BudgetExceededError) as caught: + await execute(route, asynchronous, recording_server, {"callbacks": [recorder]}) + failure: Final = await recorder.wait_for_async("async_log_failure_event" if asynchronous else "log_failure_event") + assert len(failure) == 1 + assert _OBJECT.validate_python(failure[0].kwargs)["exception"] is caught.value + assert not recording_server.requests + assert not any("success" in name for name in recorder.names) + + @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", (False, True)) async def test_native_inference_provider_failure_is_terminal_and_shared_with_callbacks( @@ -162,9 +219,13 @@ async def test_native_inference_provider_failure_is_terminal_and_shared_with_cal async def test_unstarted_native_inference_has_no_provider_or_callback_effects( route: Route, recording_server: RecordingServer, + monkeypatch: pytest.MonkeyPatch, ) -> None: recorder: Final = RecordingLogger() recording_server.expected_requests = 0 + monkeypatch.setattr(litellm, "ssl_verify", object()) + monkeypatch.setattr(litellm, "max_budget", 1.0) + monkeypatch.setattr(litellm, "_current_cost", 2.0) pending: Final = native_call(route, True, recording_server, {"callbacks": [recorder]}) assert asyncio.iscoroutine(pending) pending.close() diff --git a/tests/unit/rust_bridge/test_lifecycle.py b/tests/unit/rust_bridge/test_lifecycle.py index 4a5a741ba8a..021a4aad85a 100644 --- a/tests/unit/rust_bridge/test_lifecycle.py +++ b/tests/unit/rust_bridge/test_lifecycle.py @@ -4,25 +4,27 @@ import asyncio from collections.abc import Sequence from typing import Final -from litellm.rust_bridge.lifecycle import Await, Complete, drive +import pytest + +from litellm.rust_bridge.lifecycle import Await, Complete, Execution, Open, Step, drive class ScriptedExecution: """Plays scripted steps and records how it was resumed and whether it was closed.""" - def __init__(self, steps: Sequence[Await | Complete]) -> None: + def __init__(self, steps: Sequence[Step]) -> None: self._steps: Final = list(steps) self.resumed: list[tuple[str, object]] = [] self.closed = False - def start(self) -> Await | Complete: + def start(self) -> Step: return self._steps.pop(0) - def resume_value(self, value: object) -> Await | Complete: + def resume_value(self, value: object) -> Step: self.resumed.append(("value", value)) return self._steps.pop(0) - def resume_error(self, error: BaseException) -> Await | Complete: + def resume_error(self, error: BaseException) -> Step: self.resumed.append(("error", type(error))) return self._steps.pop(0) @@ -45,3 +47,28 @@ def test_drive_resumes_each_await_with_its_result_or_error_and_returns_the_compl assert execution.resumed == [("value", 1), ("error", ValueError)] assert execution.closed + + +@pytest.mark.parametrize("factory_fails", (False, True)) +def test_stream_handoff_preserves_head_identity_and_closes_on_construction_failure(factory_fails: bool) -> None: + head: Final = object() + stream: Final = object() + execution: Final = ScriptedExecution([Open(head)]) + failure: Final = ValueError("stream construction failed") + + def construct(owner: Execution, received: object) -> object: + assert owner is execution + assert received is head + if factory_fails: + raise failure + return stream + + if factory_fails: + with pytest.raises(ValueError, match="stream construction failed") as caught: + asyncio.run(drive(execution, construct)) + assert caught.value is failure + assert execution.closed + else: + assert asyncio.run(drive(execution, construct)) is stream + assert not execution.closed + execution.close() diff --git a/tests/unit/rust_bridge/test_runtime.py b/tests/unit/rust_bridge/test_runtime.py index 2f166631667..fd84d629937 100644 --- a/tests/unit/rust_bridge/test_runtime.py +++ b/tests/unit/rust_bridge/test_runtime.py @@ -12,7 +12,8 @@ from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_di from litellm.rust_bridge import bindings, configuration, runtime from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule from litellm.rust_bridge.configuration import Rollout -from litellm.rust_bridge.lifecycle import Complete, Open, Stream, SyncStream, Yield +from litellm.rust_bridge.lifecycle import Complete, Open, Yield +from litellm.rust_bridge.streams import Stream, SyncStream class RustBridgeDeclined(Exception):