diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index bee4421f1e9..a74ec097148 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3279,6 +3279,7 @@ dependencies = [ name = "litellm-host" version = "0.1.0" dependencies = [ + "futures-util", "litellm-auth", "litellm-coroutine", "rstest", diff --git a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md index de6e0c1b225..56ef51e4758 100644 --- a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md +++ b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md @@ -2,7 +2,7 @@ - This crate is the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on) - Smell test: if a future callback host (`callbacks-v1-python`, WASM, in-process Rust) could share a piece of this crate, it does not belong here - SDK request policy (credential inheritance, the budget and retry-count limits) is the driver's preflight, supplied by `python-bridge`; this crate only adopts the keyword view it produces - - The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call + - The driver in `litellm-host-python`, the routes and core see one `PythonCallHooks`; they never learn which Python objects consume a call - Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json` - The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it - Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython` diff --git a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs index fe0f6c7dd45..1fa7f3dcdd4 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs @@ -2,12 +2,12 @@ //! raises is answered with the same `Logging` calls, in the same order, as the Python //! `@client` path makes them. +use litellm_host_python::PythonOwned; + use litellm_host::event::{ FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest, epoch_seconds, }; -use litellm_host_python::{ - LifecycleEvent, LifecycleStep, PythonLifecycle, from_py, missing_state, to_py, -}; +use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, from_py, missing_state, to_py}; use pyo3::{ exceptions::{PyBaseException, PyException}, gc::{PyTraverseError, PyVisit}, @@ -48,11 +48,10 @@ struct DeliveredStream { first_chunk: Option>, } -enum Pending { - DeploymentPreCall, - DeploymentPostCall, - DeploymentFailure, - AsyncFailure, +struct LoggedRequest { + body: Py, + headers: Py, + context: RequestContext, } pub struct LegacyLogging { @@ -63,13 +62,10 @@ pub struct LegacyLogging { end: Option>, response: Option>, error: Option>, - body: Option>, - headers: Option>, - context: Option, + request: Option, stream: Option, asynchronous: bool, internal: bool, - pending: Option, } fn datetime(py: Python<'_>, epoch_seconds: f64) -> PyResult> { @@ -95,13 +91,10 @@ impl LegacyLogging { end: None, response: None, error: None, - body: None, - headers: None, - context: None, + request: None, stream: None, asynchronous, internal: false, - pending: None, } } @@ -120,14 +113,14 @@ impl LegacyLogging { /// The keyword view the rest of the call reads: a copy, so the deployment hook's own /// dict is left as the hook returned it, carrying the logger as `@client` injects it. /// The driver's preflight rewrites this same dict before the host projects from it. - fn prepare(&mut self, py: Python<'_>) -> PyResult { + 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))?; self.call.set_kwargs(prepared.unbind()); - Ok(LifecycleStep::Arguments(self.call.kwargs().clone_ref(py))) + Ok(HookStep::Ready(self.call.kwargs().clone_ref(py))) } - fn finalize(&mut self, py: Python<'_>) -> PyResult { + fn finalize(&mut self, py: Python<'_>) -> PyResult>> { finalize( py, &self.response, @@ -138,7 +131,7 @@ impl LegacyLogging { )?; self.response .as_ref() - .map(|response| LifecycleStep::Response(response.clone_ref(py))) + .map(|response| HookStep::Ready(response.clone_ref(py))) .ok_or_else(missing_state) } @@ -195,7 +188,10 @@ impl LegacyLogging { logger.object(py), billing.url_route, billing.endpoint_type, - &self.body, + &self + .request + .as_ref() + .map(|request| request.body.clone_ref(py)), &stream.chunks, &self.start, &self.end, @@ -214,11 +210,11 @@ impl LegacyLogging { /// A failure after the stream reached the caller bills the delivered chunks as /// partial usage. The sync path has no loop to schedule that on, so it falls back to /// the plain failure handler. - fn stream_failure(&mut self, py: Python<'_>) -> PyResult { + 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 { - return Ok(LifecycleStep::Done); + return Ok(HookStep::Ready(())); }; if !self.asynchronous { return self.dispatch_failure(py); @@ -228,30 +224,33 @@ impl LegacyLogging { ( logger.object(py), billing.endpoint_type, - &self.body, + &self + .request + .as_ref() + .map(|request| request.body.clone_ref(py)), &stream.chunks, error, ), ); match scheduled { - Ok(awaitable) => { - self.pending = Some(Pending::AsyncFailure); - Ok(LifecycleStep::Await(awaitable.unbind())) - } + Ok(awaitable) => Ok(HookStep::Await( + awaitable.unbind(), + Self::resume_async_failure, + )), Err(failure) if is_cancellation(py, &failure) => Err(failure), - Err(_) => Ok(LifecycleStep::Done), + Err(_) => Ok(HookStep::Ready(())), } } /// The sync failure handler, then the async one for async calls. Ordinary handler /// errors never replace the selected failure or suppress the other family; a /// cancellation does end the call. - fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult { + fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult> { let (Some(logger), Some(error)) = (&self.logger, &self.error) else { - return Ok(LifecycleStep::Done); + return Ok(HookStep::Ready(())); }; if self.asynchronous && self.internal { - return Ok(LifecycleStep::Done); + return Ok(HookStep::Ready(())); } if let Err(failure) = logger.failure(py, error, &self.start, &self.end, false) && is_cancellation(py, &failure) @@ -259,27 +258,64 @@ impl LegacyLogging { return Err(failure); } if !self.asynchronous { - return Ok(LifecycleStep::Done); + return Ok(HookStep::Ready(())); } match logger.failure(py, error, &self.start, &self.end, true) { - Ok(Some(awaitable)) => { - self.pending = Some(Pending::AsyncFailure); - Ok(LifecycleStep::Await(awaitable)) - } - Ok(None) => Ok(LifecycleStep::Done), + Ok(Some(awaitable)) => Ok(HookStep::Await(awaitable, Self::resume_async_failure)), + Ok(None) => Ok(HookStep::Ready(())), Err(failure) if is_cancellation(py, &failure) => Err(failure), - Err(_) => Ok(LifecycleStep::Done), + Err(_) => Ok(HookStep::Ready(())), } } } -impl PythonLifecycle for LegacyLogging { - fn begin( +impl LegacyLogging { + fn resume_begin( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult>> { + self.call + .set_kwargs(result?.into_bound(py).cast_into::()?.unbind()); + self.prepare(py) + } + + fn resume_after_success( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult>> { + self.response = Some(result?); + self.finalize(py) + } + + fn resume_deployment_failure( + &mut self, + py: Python<'_>, + _: PyResult>, + ) -> PyResult> { + self.dispatch_failure(py) + } + + fn resume_async_failure( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult> { + match result { + Err(error) if is_cancellation(py, &error) => Err(error), + _ => Ok(HookStep::Ready(())), + } + } +} + +impl PythonCallHooks for LegacyLogging { + fn prepare_arguments( &mut self, py: Python<'_>, arguments: Py, started_at: f64, - ) -> PyResult { + ) -> PyResult>> { self.call.set_kwargs(arguments); self.start = datetime(py, started_at)?; self.internal = is_internal_call(py)?; @@ -294,22 +330,20 @@ impl PythonLifecycle for LegacyLogging { self.logger = Some(result.logger()?); self.call.set_kwargs(result.kwargs()?); if self.runs_deployment_hooks() { - self.pending = Some(Pending::DeploymentPreCall); - return Ok(LifecycleStep::Await(DeploymentHooks::before_call( - py, - self.call.kwargs(), - self.surface.call_type, - )?)); + return Ok(HookStep::Await( + DeploymentHooks::before_call(py, self.call.kwargs(), self.surface.call_type)?, + Self::resume_begin, + )); } self.prepare(py) } - fn before_send( + fn before_provider_request( &mut self, py: Python<'_>, wire: Box, context: &RequestContext, - ) -> PyResult { + ) -> PyResult>> { let logger = self.logger()?; logger.update_from_kwargs(py, self.call.kwargs(), &wire, context)?; let body = to_py(py, &wire.body)? @@ -326,9 +360,11 @@ impl PythonLifecycle for LegacyLogging { for (name, value) in &wire.headers { headers.set_item(name, value)?; } - self.body = Some(body.clone().unbind()); - self.headers = Some(headers.clone().unbind()); - self.context = Some(context.clone()); + self.request = Some(LoggedRequest { + body: body.clone().unbind(), + headers: headers.clone().unbind(), + context: context.clone(), + }); self.logger()?.pre_call( py, self.surface.input_description, @@ -341,61 +377,63 @@ impl PythonLifecycle for LegacyLogging { .iter() .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) .collect::>>()?; - Ok(LifecycleStep::Wire(Box::new(WireRequest { + Ok(HookStep::Ready(Box::new(WireRequest { body: from_py(&body)?, headers, ..*wire }))) } - fn after_success( + fn transform_response( &mut self, py: Python<'_>, response: Py, timing: Timing, - ) -> PyResult { + ) -> PyResult>> { self.end = Some(datetime(py, timing.end_time)?); self.response = Some(response); if self.runs_deployment_hooks() { - self.pending = Some(Pending::DeploymentPostCall); - return Ok(LifecycleStep::Await(DeploymentHooks::after_success( - py, - self.call.kwargs(), - &self.response, - self.surface.call_type, - )?)); + return Ok(HookStep::Await( + DeploymentHooks::after_success( + py, + self.call.kwargs(), + &self.response, + self.surface.call_type, + )?, + Self::resume_after_success, + )); } self.finalize(py) } - fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult { + fn on_event(&mut self, py: Python<'_>, event: HookEvent<'_>) -> PyResult> { match event { - LifecycleEvent::Started { .. } => Ok(LifecycleStep::Done), - LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => { + HookEvent::Started { .. } => Ok(HookStep::Ready(())), + HookEvent::Machine(MachineEvent::ResponseReceived { raw }) => { let api_key = self - .context + .request .as_ref() - .and_then(|context| context.api_key.as_ref()) + .and_then(|request| request.context.api_key.as_ref()) .map(|api_key| api_key.expose()); self.logger()?.post_call( py, &raw.body, api_key, - self.body.as_ref(), - self.headers.as_ref(), + self.request.as_ref().map(|request| &request.body), + self.request.as_ref().map(|request| &request.headers), )?; - Ok(LifecycleStep::Done) + Ok(HookStep::Ready(())) } - LifecycleEvent::Succeeded { timing, response } => { + HookEvent::Succeeded { timing, response } => { self.end = Some(datetime(py, timing.end_time)?); self.response = Some(response.clone_ref(py)); match &self.stream { Some(stream) => self.stream_success(py, stream)?, None => self.dispatch_success(py)?, } - Ok(LifecycleStep::Done) + Ok(HookStep::Ready(())) } - LifecycleEvent::Failed { + HookEvent::Failed { timing, origin, error, @@ -410,20 +448,22 @@ impl PythonLifecycle for LegacyLogging { && self.runs_deployment_hooks() { let error = self.error.as_ref().ok_or_else(missing_state)?; - self.pending = Some(Pending::DeploymentFailure); - return Ok(LifecycleStep::Await(DeploymentHooks::after_failure( - py, - self.call.kwargs(), - error, - self.surface.call_type, - )?)); + return Ok(HookStep::Await( + DeploymentHooks::after_failure( + py, + self.call.kwargs(), + error, + self.surface.call_type, + )?, + Self::resume_deployment_failure, + )); } self.dispatch_failure(py) } } } - fn opened(&mut self, py: Python<'_>) -> PyResult<()> { + fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> { if self.surface.stream.is_none() { return Err(missing_state()); } @@ -435,45 +475,25 @@ impl PythonLifecycle for LegacyLogging { Ok(()) } - fn delivered(&mut self, py: Python<'_>, chunk: &Py) -> PyResult<()> { + fn on_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())?); } stream.chunks.bind(py).append(chunk) } +} - fn resume(&mut self, py: Python<'_>, result: PyResult>) -> PyResult { - match self.pending.take().ok_or_else(missing_state)? { - Pending::DeploymentPreCall => { - self.call - .set_kwargs(result?.into_bound(py).cast_into::()?.unbind()); - self.prepare(py) - } - Pending::DeploymentPostCall => { - self.response = Some(result?); - self.finalize(py) - } - Pending::DeploymentFailure => self.dispatch_failure(py), - Pending::AsyncFailure => match result { - Err(failure) if is_cancellation(py, &failure) => Err(failure), - _ => Ok(LifecycleStep::Done), - }, - } - } - +impl PythonOwned for LegacyLogging { fn close(&mut self, py: Python<'_>) { if let Some(logger) = self.logger.take() && let Err(error) = logger.restore_context(py) { error.write_unraisable(py, None); } - self.body = None; - self.headers = None; - self.context = None; + self.request = None; self.stream = None; } - fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { self.call.traverse(visit)?; if let Some(logger) = &self.logger { @@ -487,8 +507,11 @@ impl PythonLifecycle for LegacyLogging { visit.call(&stream.chunks)?; visit.call(&stream.first_chunk)?; } - visit.call(&self.body)?; - visit.call(&self.headers) + if let Some(request) = &self.request { + visit.call(&request.body)?; + visit.call(&request.headers)?; + } + Ok(()) } } @@ -497,7 +520,7 @@ mod deployment_hooks_tests { use std::ffi::CStr; use litellm_host::event::{FailureOrigin, Timing}; - use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle}; + use litellm_host_python::{HookEvent, HookStep, PythonCallHooks}; use pyo3::exceptions::asyncio::CancelledError; use pyo3::prelude::*; use pyo3::types::PyDict; @@ -506,6 +529,18 @@ mod deployment_hooks_tests { use super::LegacyLogging; use crate::test_support::{legacy_call, local, namespace, run}; + fn resume( + logging: &mut LegacyLogging, + step: HookStep, + py: Python<'_>, + value: PyResult>, + ) -> PyResult> { + let HookStep::Await(_, continuation) = step else { + panic!("expected suspension") + }; + continuation(logging, py, value) + } + const CALL: &CStr = c" document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'} kwargs = {'logger': logger, 'document': document} @@ -520,25 +555,28 @@ kwargs = {'logger': logger, 'document': document} py: Python<'py>, locals: &Bound<'py, PyDict>, asynchronous: bool, - ) -> (LegacyLogging, LifecycleStep) { + ) -> (LegacyLogging, HookStep>) { let mut logging = legacy_call(py, locals, asynchronous); let kwargs = local(locals, "kwargs") .cast_into::() .unwrap() .unbind(); - let step = logging.begin(py, kwargs, 0.0).unwrap(); + let step = logging.prepare_arguments(py, kwargs, 0.0).unwrap(); (logging, step) } - fn arguments<'py>(py: Python<'py>, step: LifecycleStep) -> Bound<'py, PyDict> { - let LifecycleStep::Arguments(arguments) = step else { + fn arguments<'py>( + py: Python<'py>, + step: HookStep>, + ) -> Bound<'py, PyDict> { + let HookStep::Ready(arguments) = step else { panic!("expected the prepared arguments"); }; arguments.into_bound(py) } - fn awaits_deployment_hook(step: &LifecycleStep) -> bool { - matches!(step, LifecycleStep::Await(_)) + fn awaits_deployment_hook(step: &HookStep) -> bool { + matches!(step, HookStep::Await(_, _)) } #[rstest] @@ -559,7 +597,7 @@ kwargs = {'logger': logger, 'document': document} }); } - #[test] + #[rstest::rstest] fn kwargs_returned_by_the_pre_call_hook_are_what_the_call_prepares() { Python::initialize(); Python::attach(|py| { @@ -574,9 +612,13 @@ replaced_kwargs = {'logger': logger, 'document': replacement, 'pages': [0]} ); let (mut logging, step) = begin(py, &locals, true); assert!(awaits_deployment_hook(&step)); - let step = logging - .resume(py, Ok(local(&locals, "replaced_kwargs").unbind())) - .unwrap(); + let step = resume( + &mut logging, + step, + py, + Ok(local(&locals, "replaced_kwargs").unbind()), + ) + .unwrap(); locals.set_item("prepared", arguments(py, step)).unwrap(); run( py, @@ -610,7 +652,9 @@ kwargs = {'logger': logger, 'vendor_extension': opaque} ); let (mut logging, step) = begin(py, &locals, asynchronous); let step = match step { - LifecycleStep::Await(hook_result) => logging.resume(py, Ok(hook_result)).unwrap(), + HookStep::Await(hook_result, resume) => { + resume(&mut logging, py, Ok(hook_result)).unwrap() + } step => step, }; locals.set_item("prepared", arguments(py, step)).unwrap(); @@ -626,7 +670,7 @@ assert hooked == ([opaque] if asynchronous else []), hooked }); } - #[test] + #[rstest::rstest] fn response_returned_by_the_post_call_hook_is_finalized_and_returned() { Python::initialize(); Python::attach(|py| { @@ -639,18 +683,26 @@ replacement = object() logger.hooks = {'pre': lambda kwargs: kwargs} ", ); - let (mut logging, _) = begin(py, &locals, true); - logging - .resume(py, Ok(local(&locals, "kwargs").unbind())) - .unwrap(); + let (mut logging, step) = begin(py, &locals, true); + resume( + &mut logging, + step, + py, + Ok(local(&locals, "kwargs").unbind()), + ) + .unwrap(); let step = logging - .after_success(py, local(&locals, "response").unbind(), TIMING) + .transform_response(py, local(&locals, "response").unbind(), TIMING) .unwrap(); assert!(awaits_deployment_hook(&step)); - let step = logging - .resume(py, Ok(local(&locals, "replacement").unbind())) - .unwrap(); - let LifecycleStep::Response(returned) = step else { + let step = resume( + &mut logging, + step, + py, + Ok(local(&locals, "replacement").unbind()), + ) + .unwrap(); + let HookStep::Ready(returned) = step else { panic!("expected the finalized response"); }; assert!(returned.bind(py).is(local(&locals, "replacement"))); @@ -672,18 +724,28 @@ assert finalized is replacement Python::initialize(); Python::attach(|py| { let locals = namespace(py, c"kwargs = {'logger': logger}\nresponse = object()"); - let (mut logging, _) = begin(py, &locals, true); - if post_call { - logging - .resume(py, Ok(local(&locals, "kwargs").unbind())) - .unwrap(); - logging - .after_success(py, local(&locals, "response").unbind(), TIMING) - .unwrap(); - } + let (mut logging, step) = begin(py, &locals, true); let cancellation = CancelledError::new_err("cancelled"); let cancelled = cancellation.value(py).clone(); - let error = logging.resume(py, Err(cancellation)).err().unwrap(); + let error = if post_call { + resume( + &mut logging, + step, + py, + Ok(local(&locals, "kwargs").unbind()), + ) + .unwrap(); + let step = logging + .transform_response(py, local(&locals, "response").unbind(), TIMING) + .unwrap(); + resume(&mut logging, step, py, Err(cancellation)) + .err() + .unwrap() + } else { + resume(&mut logging, step, py, Err(cancellation)) + .err() + .unwrap() + }; assert!(error.value(py).is(&cancelled)); let names: Vec = local(&locals, "logger") .call_method0("names") @@ -704,17 +766,21 @@ assert finalized is replacement py, c"kwargs = {'logger': logger}\nfailure = ValueError('provider')", ); - let (mut logging, _) = begin(py, &locals, true); - logging - .resume(py, Ok(local(&locals, "kwargs").unbind())) - .unwrap(); + let (mut logging, step) = begin(py, &locals, true); + resume( + &mut logging, + step, + py, + Ok(local(&locals, "kwargs").unbind()), + ) + .unwrap(); let failure = PyErr::from_value(local(&locals, "failure")); - let failed = LifecycleEvent::Failed { + let failed = HookEvent::Failed { timing: TIMING, origin: FailureOrigin::Call, error: &failure, }; - let step = logging.emit(py, failed).unwrap(); + let step = logging.on_event(py, failed).unwrap(); assert!(awaits_deployment_hook(&step)); let hook_result = if cancelled { Err(CancelledError::new_err("cancelled")) @@ -722,8 +788,8 @@ assert finalized is replacement Ok(py.None()) }; assert!(matches!( - logging.resume(py, hook_result).unwrap(), - LifecycleStep::Await(_) + resume(&mut logging, step, py, hook_result).unwrap(), + HookStep::Await(_, _) )); run( py, @@ -743,7 +809,7 @@ mod payload_tests { use litellm_auth::SecretValue; use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest}; - use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py}; + use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, PythonOwned, to_py}; use proptest::prelude::*; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -788,11 +854,11 @@ check = lambda: None json!({"type": "document_url", "document_url": source}) } - fn before_send(script: &CStr, body: Value) -> WireRequest { + fn before_provider_request(script: &CStr, body: Value) -> WireRequest { before_send_with_secrets(script, json!({}), body, &[]) } - /// Runs `before_send` over `body` for a route whose parameters are `optional_params`, with + /// Runs `before_provider_request` over `body` for a route whose parameters are `optional_params`, with /// the Python objects `script` binds, then delivers the provider's raw response the way the /// driver does and runs the script's `check()`. fn before_send_with_secrets( @@ -837,30 +903,35 @@ check = lambda: None }; let (_, step) = send_and_receive(py, &mut logging, wire, &context); run(py, &locals, c"check()"); - let LifecycleStep::Wire(wire) = step else { - panic!("before_send did not hand back the wire request"); + let HookStep::Ready(wire) = step else { + panic!("before_provider_request did not hand back the wire request"); }; *wire }) } - /// `before_send` over `wire`, then the provider's raw response the way the driver + /// `before_provider_request` over `wire`, then the provider's raw response the way the driver /// delivers it, so `pre_call` and `post_call` have both seen the retained payload. fn send_and_receive<'a>( py: Python<'_>, logging: &'a mut LegacyLogging, wire: WireRequest, context: &RequestContext, - ) -> (&'a mut LegacyLogging, LifecycleStep) { - let step = logging.before_send(py, Box::new(wire), context).unwrap(); + ) -> ( + &'a mut LegacyLogging, + HookStep>, + ) { + let step = logging + .before_provider_request(py, Box::new(wire), context) + .unwrap(); let raw = MachineEvent::ResponseReceived { raw: RawResponse { body: "raw response".into(), }, }; assert!(matches!( - logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(), - LifecycleStep::Done + logging.on_event(py, HookEvent::Machine(&raw)).unwrap(), + HookStep::Ready(()) )); (logging, step) } @@ -904,7 +975,7 @@ check = lambda: None } } - #[test] + #[rstest::rstest] fn a_cycle_through_the_retained_headers_is_collected() { Python::initialize(); Python::attach(|py| { @@ -940,7 +1011,7 @@ assert reference() is None }); } - #[test] + #[rstest::rstest] fn close_releases_the_retained_headers() { Python::initialize(); Python::attach(|py| { @@ -998,13 +1069,13 @@ def check(): ")] fn passthrough_keys_reach_pre_call_as_the_callers_own_objects(#[case] script: &CStr) { let body = json!({"model": "model", "document": document(DOCUMENT), "pages": [0]}); - let wire = before_send(script, body.clone()); + let wire = before_provider_request(script, body.clone()); assert_eq!(wire.body, body); } - #[test] + #[rstest::rstest] fn pre_call_edit_of_a_passthrough_object_reaches_the_caller_and_the_wire() { - let wire = before_send( + let wire = before_provider_request( c" document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'} kwargs = {'document': document} @@ -1018,9 +1089,9 @@ def check(): assert_eq!(wire.body["document"], document(EDITED)); } - #[test] + #[rstest::rstest] fn a_body_key_the_route_rewrote_is_not_the_callers_object() { - let wire = before_send( + let wire = before_provider_request( c" document = {'type': 'document_url', 'document_url': 'https://example.invalid/scan.pdf'} kwargs = {'document': document} @@ -1040,10 +1111,10 @@ def check(): ); } - #[test] + #[rstest::rstest] fn a_caller_value_with_no_json_form_is_left_out_of_realiasing() { let body = json!({"pages": [0]}); - let wire = before_send( + let wire = before_provider_request( c" opaque = object() kwargs = {'pages': opaque} @@ -1072,14 +1143,14 @@ def on_pre_call(args): )] fn rebinding_the_payload_envelope_does_not_reach_the_wire(#[case] script: &CStr) { let body = json!({"document": document(DOCUMENT)}); - let wire = before_send(script, body.clone()); + let wire = before_provider_request(script, body.clone()); assert_eq!(wire.body, body); assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]); } - #[test] + #[rstest::rstest] fn pre_call_header_edit_reaches_the_wire() { - let wire = before_send( + let wire = before_provider_request( c" def on_pre_call(args): args['headers']['x-callback'] = 'edited' @@ -1095,7 +1166,7 @@ def on_pre_call(args): ); } - #[test] + #[rstest::rstest] fn pre_call_receives_the_wire_request_and_the_logger_its_redacted_request() { let body = json!({"model": "model", "document": document(DOCUMENT)}); before_send_with_secrets( @@ -1167,13 +1238,13 @@ def on_pre_call(args): )] fn pre_call_body_edits_reach_the_wire(#[case] script: &CStr, #[case] expected: Value) { let body = json!({"document": document(DOCUMENT)}); - let wire = before_send(script, body); + let wire = before_provider_request(script, body); assert_eq!(wire.body, expected); } - #[test] + #[rstest::rstest] fn retained_headers_edited_after_rebinding_reach_the_wire() { - let wire = before_send( + let wire = before_provider_request( c" def on_pre_call(args): retained = args['headers'] @@ -1191,9 +1262,9 @@ def on_pre_call(args): ); } - #[test] + #[rstest::rstest] fn post_call_receives_the_raw_response_the_route_key_and_the_body_and_headers_pre_call_saw() { - before_send( + before_provider_request( c" def check(): original_response, api_key, additional_args = logger.post @@ -1210,9 +1281,9 @@ def check(): ); } - #[test] + #[rstest::rstest] fn every_request_runs_the_full_pre_call_and_post_call() { - let wire = before_send( + let wire = before_provider_request( c" def on_pre_call(args): args['complete_input_dict']['include_image_base64'] = True @@ -1343,7 +1414,7 @@ def check(): /// For any body, any caller keywords and any callback edit: every keyword the route /// sends unchanged reaches `pre_call` as the caller's own object, and the provider is /// sent exactly what the model says, so a callback that edits nothing changes nothing. - #[test] + #[rstest::rstest] fn the_wire_is_the_body_pre_call_received_as_the_callback_left_it( fields in prop::collection::btree_map(key(), (json_value(), caller()), 0..5), edit in edit(), @@ -1389,7 +1460,7 @@ mod terminal_tests { use std::ffi::CStr; use litellm_host::event::{FailureOrigin, Timing}; - use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle}; + use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, PythonOwned}; use pyo3::exceptions::PyRuntimeError; use pyo3::exceptions::asyncio::CancelledError; use pyo3::prelude::*; @@ -1416,12 +1487,12 @@ mod terminal_tests { py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging, - ) -> LifecycleStep { + ) -> HookStep { let response = local(locals, "response").unbind(); logging - .emit( + .on_event( py, - LifecycleEvent::Succeeded { + HookEvent::Succeeded { timing: TIMING, response: &response, }, @@ -1433,12 +1504,12 @@ mod terminal_tests { py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging, - ) -> LifecycleStep { + ) -> HookStep { let failure = PyErr::from_value(local(locals, "failure")); logging - .emit( + .on_event( py, - LifecycleEvent::Failed { + HookEvent::Failed { timing: TIMING, origin: FailureOrigin::Host, error: &failure, @@ -1468,7 +1539,7 @@ mod terminal_tests { let mut logging = logged(py, &locals, asynchronous); assert!(matches!( succeed(py, &locals, &mut logging), - LifecycleStep::Done + HookStep::Ready(()) )); let names: Vec = local(&locals, "logger") .call_method0("names") @@ -1503,7 +1574,7 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy }; assert!(matches!( fail(py, &locals, &mut logging), - LifecycleStep::Done + HookStep::Ready(()) )); let names: Vec = local(&locals, "logger") .call_method0("names") @@ -1514,7 +1585,7 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy }); } - #[test] + #[rstest::rstest] fn internal_async_calls_skip_the_async_success_fan_out() { Python::initialize(); Python::attach(|py| { @@ -1532,7 +1603,7 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy }); } - #[test] + #[rstest::rstest] fn a_failing_success_callback_is_reported_without_replacing_the_response() { Python::initialize(); Python::attach(|py| { @@ -1552,7 +1623,7 @@ logger = FailingLogger() let mut logging = logged(py, &locals, true); assert!(matches!( succeed(py, &locals, &mut logging), - LifecycleStep::Done + HookStep::Ready(()) )); assert!( logging @@ -1581,10 +1652,7 @@ logger = FailingLogger() let mut logging = logged(py, &locals, asynchronous); let step = fail(py, &locals, &mut logging); let awaits_async_handler = expected.contains(&"async_failure_handler"); - assert_eq!( - matches!(step, LifecycleStep::Await(_)), - awaits_async_handler - ); + assert_eq!(matches!(step, HookStep::Await(_, _)), awaits_async_handler); let names: Vec = local(&locals, "logger") .call_method0("names") .unwrap() @@ -1599,7 +1667,7 @@ logger = FailingLogger() }); } - #[test] + #[rstest::rstest] fn a_failing_sync_failure_callback_keeps_the_error_and_still_runs_the_async_family() { Python::initialize(); Python::attach(|py| { @@ -1619,7 +1687,7 @@ logger = FailingLogger() let mut logging = logged(py, &locals, true); assert!(matches!( fail(py, &locals, &mut logging), - LifecycleStep::Await(_) + HookStep::Await(_, _) )); assert!( logging @@ -1649,15 +1717,17 @@ logger = FailingLogger() Python::attach(|py| { let locals = namespace(py, c"failure = ValueError('provider')"); let mut logging = logged(py, &locals, true); - fail(py, &locals, &mut logging); + let HookStep::Await(_, resume) = fail(py, &locals, &mut logging) else { + panic!("expected async failure handler") + }; let result = match error { None => Ok(py.None()), Some(false) => Err(PyRuntimeError::new_err("handler failed")), Some(true) => Err(CancelledError::new_err("cancelled")), }; let expected = result.as_ref().err().map(|error| error.value(py).clone()); - match logging.resume(py, result) { - Ok(step) => assert!(done && matches!(step, LifecycleStep::Done)), + match resume(&mut logging, py, result) { + Ok(step) => assert!(done && matches!(step, HookStep::Ready(()))), Err(propagated) => { assert!(!done); assert!(propagated.value(py).is(expected.unwrap())); @@ -1666,7 +1736,7 @@ logger = FailingLogger() }); } - #[test] + #[rstest::rstest] fn closing_restores_the_correlation_context_once() { Python::initialize(); Python::attach(|py| { diff --git a/litellm-rust/crates/callbacks-legacy-python/src/call.rs b/litellm-rust/crates/callbacks-legacy-python/src/call.rs index 3fa638ac6d3..52b64d8b181 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/call.rs @@ -3,8 +3,8 @@ //! lifetime. No other callback host has that obligation, which is why nothing outside //! this crate holds them. -use litellm_host::{machine::Machine, protocol::Protocol}; -use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call}; +use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol}; +use litellm_host_python::{Preflight, PythonBinding, PythonHostCalls, lookup, run_call}; use pyo3::{ gc::{PyTraverseError, PyVisit}, prelude::*, @@ -71,21 +71,22 @@ pub fn run_legacy_call( py: Python<'_>, surface: LegacySurface, call: PublicCall, - machine: M, + start: impl FnOnce(::Request) -> M + Send + Sync + 'static, host: H, preflight: Preflight, asynchronous: bool, ) -> PyResult> where - H: ProtocolHost + 'static, - M: Machine::Response> + 'static, + H: PythonBinding + PythonHostCalls + 'static, + M: Machine + 'static, + M::Complete: Into::Response>>, { let arguments = call.kwargs.clone_ref(py); run_call( py, - machine, + start, host, - Box::new(LegacyLogging::new(py, surface, call, asynchronous)), + LegacyLogging::new(py, surface, call, asynchronous), preflight, arguments, asynchronous, @@ -110,7 +111,7 @@ mod tests { (call, locals) } - #[test] + #[rstest::rstest] fn capture_copies_the_keyword_dict_without_copying_its_values() { Python::initialize(); Python::attach(|py| { diff --git a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs index 8fca64d1b0e..a5d5c6762c2 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs @@ -1,7 +1,7 @@ //! The legacy `@client` wrapper as the native call sees it: litellm's `Logging` object, the //! sync and async callback registries it fans out to, the deployment hooks and the deferred //! proxy release. All of it sits behind one -//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and +//! [`PythonCallHooks`](litellm_host_python::PythonCallHooks), so the driver, the routes and //! core never learn which Python object is on the other end. The SDK's own request policy //! (credential inheritance, the budget and retry limits) is the driver's preflight, not this //! crate's. diff --git a/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs b/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs index 6333eedebfc..e71ecdffa44 100644 --- a/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs +++ b/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs @@ -1,9 +1,26 @@ +use strum::{EnumString, IntoStaticStr}; + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct CustomLlmProvider<'a> { pub model: &'a str, pub custom_llm_provider: &'a str, } +#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)] +#[strum(serialize_all = "snake_case")] +pub enum LlmProviders { + Anthropic, + AwsTextract, + AzureAi, + Bedrock, + Cohere, + Mistral, + Openai, + OpenaiLike, + Reducto, + VertexAi, +} + pub fn get_custom_llm_provider<'a>( model: &'a str, custom_llm_provider: Option<&'a str>, diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 912115f4536..0a10dac1bdd 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -1,6 +1,12 @@ -litellm-core is the LiteLLM SDK in Rust. Each top-level call is a module under `src//` exposing a public entrypoint named after the route. `messages::messages()` returns `MessagesResponse::Message` for a completed response or `MessagesResponse::Stream { headers, chunks }` when the request sets `stream: true`. The chunks are Anthropic SSE bytes in a `Stream>`. Dropping the stream cancels the call. The Python bridge drives `messages::route::messages_machine()` instead, because Python has to answer the call's operations on its own thread; the gateway and the Rust SDK call the plain entrypoint +litellm-core owns route orchestration. Messages and HTTP Responses return `litellm_host::call::CallOutput`, containing either a completed response or a stream head and chunks. OCR and currently non-streaming Chat Completions return their completed response directly -A route module has the same five pieces, in the order Python runs them. `types.rs` holds the call, the provider request, and the response. `prepare.rs` resolves the provider and credentials and shapes the request (Python's `validate_environment`, `get_complete_url`, `transform_request`). `handler.rs` resolves auth, offers the wire request to `litellm_host::hooks::RouteHooks::before_send`, sends it, reports the raw response through `emit`, and normalizes the response or stream (`pre_call`, `post`, `post_call`, `transform_response`). `mod.rs` exposes the entrypoint that runs prepare then handler with no hooks (`()`). `route.rs`, where a host needs it, wraps the same two calls in a `CallMachine` whose `HostChannel` is the hooks, and pumps a stream through `open` and `deliver`. A handler takes `&impl RouteHooks` and never a `HostChannel` directly, so it runs without a coroutine. Keep provider transport and transformation details out of the machine driver +Hosts assemble route objects from shared `CoreResources`, HTTP settings, and secret sources. Each route owns its provider client and authentication dependencies. Gateway routes live for the gateway lifetime; Python assembles routes per call from its settings snapshot + +Chat Completions, Messages, and OCR execute through their route objects. Calls pass `RouteHooks` directly; use `&()` when no hooks are needed. Construction does no work; preparation and lifecycle observation begin when the future is polled. Handlers accept `RouteHooks`, never a concrete `ChannelHooks`. Native observers receive start and terminal events through the shared call runner; a stream retains its lifecycle until exhaustion, error, or drop. A host channel has no native observer because its driver owns terminal dispatch + +`route.rs` declares the concrete `Protocol` and implements a route method that accepts a typed request and constructs a `litellm_host::call::HostedMachine` with `hosted_call`. The shared call plumbing owns stream opening, delivery, backpressure, and detachment. Request decoding belongs to the boundary before the machine starts. Route closures only supply execution dependencies and route-specific host capabilities such as an OCR token provider. Use `run_hosted` for a native host so detachment is reported as cancellation. Python uses its own shared driver and preserves caller-task callback execution + +Responses WebSocket sessions remain separate from the HTTP call driver because a connection can accept multiple requests while receiving events ## Crate layering diff --git a/litellm-rust/crates/core/src/audio_transcription/handler.rs b/litellm-rust/crates/core/src/audio_transcription/handler.rs index e7c3f543f86..944dfc75763 100644 --- a/litellm-rust/crates/core/src/audio_transcription/handler.rs +++ b/litellm-rust/crates/core/src/audio_transcription/handler.rs @@ -17,7 +17,7 @@ pub async fn execute_audio_transcription_provider_call( ) -> Result { let env_lookup = |key: &str| request.secrets.get(key); let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?; - let response = crate::outbound::outbound_request( + let outbound = crate::outbound::outbound_request( authenticated, request.url.clone(), &request.body, @@ -26,12 +26,12 @@ pub async fn execute_audio_transcription_provider_call( .timeout .unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)), ), - )? - .send(http) - .await - .map_err(|error| { - Error::Transport(litellm_http::transport::Error::Network(error.to_string())) - })?; + )?; + let response = crate::outbound::send(outbound, http) + .await + .map_err(|error| { + Error::Transport(litellm_http::transport::Error::Network(error.to_string())) + })?; let status = response.status(); let text = response.text().await.map_err(|error| { Error::Transport(litellm_http::transport::Error::Network(error.to_string())) diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs index 736ee7e7434..e387ba2c6df 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -1,22 +1,39 @@ -use litellm_secrets::source::SecretSource; pub mod types; pub use crate::error::RouteError as Error; mod handler; mod prepare; pub use handler::execute_audio_transcription_provider_call; -use litellm_http::{ClientVariant, HttpClientConfig}; +use litellm_auth::AuthServices; +use litellm_secrets::source::SecretSource; pub use prepare::prepare_audio_transcription_provider_call; use serde_json::Value; +use std::sync::Arc; use crate::audio_transcription::types::AudioTranscriptionRequest; -pub async fn audio_transcription( - resources: &crate::resources::CoreResources, - config: &HttpClientConfig, - secrets: &dyn SecretSource, - request: AudioTranscriptionRequest<'_>, -) -> Result { - let request = prepare_audio_transcription_provider_call(request, secrets).await?; - let http = resources.pool.client(config, ClientVariant::Provider)?; - execute_audio_transcription_provider_call(&http, &resources.auth, request).await +#[derive(Clone)] +pub struct AudioTranscriptionRoute { + http: litellm_http::Client, + auth: Arc, + secrets: Arc, +} + +impl AudioTranscriptionRoute { + pub fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self { + Self { + http, + auth, + secrets, + } + } + + pub async fn execute(&self, request: AudioTranscriptionRequest<'_>) -> Result { + let request = + prepare_audio_transcription_provider_call(request, self.secrets.as_ref()).await?; + execute_audio_transcription_provider_call(&self.http, &self.auth, request).await + } } diff --git a/litellm-rust/crates/core/src/audio_transcription/prepare.rs b/litellm-rust/crates/core/src/audio_transcription/prepare.rs index 1ec1b114c29..4c951352db2 100644 --- a/litellm-rust/crates/core/src/audio_transcription/prepare.rs +++ b/litellm-rust/crates/core/src/audio_transcription/prepare.rs @@ -1,4 +1,3 @@ -use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}; use litellm_http::request::string_headers; use litellm_http::request::with_default_headers; use litellm_llms::{ @@ -14,36 +13,35 @@ use super::Error; use crate::audio_transcription::types::{ AudioTranscriptionRequest, ProviderAudioTranscriptionRequest, }; +use crate::provider::{LlmProviders, resolve_llm_provider}; -fn provider_config(provider: &str) -> Option<&'static dyn BaseAudioTranscriptionConfig> { - if provider == "bedrock" { - return Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG); +fn provider_config(provider: LlmProviders) -> Option<&'static dyn BaseAudioTranscriptionConfig> { + match provider { + LlmProviders::Bedrock => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG), + LlmProviders::Anthropic + | LlmProviders::AwsTextract + | LlmProviders::AzureAi + | LlmProviders::Cohere + | LlmProviders::Mistral + | LlmProviders::Openai + | LlmProviders::OpenaiLike + | LlmProviders::Reducto + | LlmProviders::VertexAi => None, } - let _ = provider; - None } pub async fn prepare_audio_transcription_provider_call( request: AudioTranscriptionRequest<'_>, secrets: &dyn SecretSource, ) -> Result { - let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) - .or_else(|| { - request - .custom_llm_provider - .map(|provider| CustomLlmProvider { - model: request.model, - custom_llm_provider: provider, - }) - }) - .ok_or_else(|| { - Error::InvalidProvider( - "unable to resolve custom_llm_provider for audio transcription request".to_string(), - ) - })?; + let provider_info = resolve_llm_provider( + request.model, + request.custom_llm_provider, + "audio transcription", + )?; let model = provider_info.model.to_string(); - let config = provider_config(provider_info.custom_llm_provider) - .ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?; + let config = provider_config(provider_info.provider) + .ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))?; let snapshot = secrets.resolve(&config.secret_names()).await?; let env_lookup = |key: &str| snapshot.get(key); let forwarded = string_headers("audio transcription", request.extra_headers)?; @@ -64,7 +62,7 @@ pub async fn prepare_audio_transcription_provider_call( config.transform_audio_transcription_request(&model, request.audio, filtered_params)?; Ok(ProviderAudioTranscriptionRequest { model, - custom_llm_provider: provider_info.custom_llm_provider.to_string(), + custom_llm_provider: <&str>::from(provider_info.provider).to_string(), config, url, body: transformed.body, diff --git a/litellm-rust/crates/core/src/chat_completions/common_utils.rs b/litellm-rust/crates/core/src/chat_completions/common_utils.rs index 1c7875c3c33..c8001b52891 100644 --- a/litellm-rust/crates/core/src/chat_completions/common_utils.rs +++ b/litellm-rust/crates/core/src/chat_completions/common_utils.rs @@ -8,15 +8,38 @@ use litellm_llms::{ use serde_json::{Map, Value}; use super::Error; +use crate::provider::LlmProviders; const HEADER_CONTEXT: &str = "chat completions"; -pub(super) fn chat_completions_provider_config(provider: &str) -> Option<&'static dyn BaseConfig> { +pub(super) enum ChatProvider { + Anthropic, + Bedrock, + OpenaiLike, +} + +impl ChatProvider { + pub(super) fn config(self) -> &'static dyn BaseConfig { + match self { + Self::Anthropic => &ANTHROPIC_CHAT_COMPLETIONS_CONFIG, + Self::Bedrock => &BEDROCK_CHAT_COMPLETIONS_CONFIG, + Self::OpenaiLike => &OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG, + } + } +} + +pub(super) fn chat_completions_provider(provider: LlmProviders) -> Option { match provider { - "anthropic" => Some(&ANTHROPIC_CHAT_COMPLETIONS_CONFIG), - "bedrock" => Some(&BEDROCK_CHAT_COMPLETIONS_CONFIG), - "openai_like" => Some(&OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG), - _ => None, + LlmProviders::Anthropic => Some(ChatProvider::Anthropic), + LlmProviders::Bedrock => Some(ChatProvider::Bedrock), + LlmProviders::OpenaiLike => Some(ChatProvider::OpenaiLike), + LlmProviders::AwsTextract + | LlmProviders::AzureAi + | LlmProviders::Cohere + | LlmProviders::Mistral + | LlmProviders::Openai + | LlmProviders::Reducto + | LlmProviders::VertexAi => None, } } diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index 651e85789ec..2d4be71463c 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -46,7 +46,7 @@ pub(super) async fn execute( }; let authenticated = resolve_auth(auth, environment, &|key| secrets.get(key)).await?; let wire = hooks - .before_send( + .before_provider_request( WireRequest { url, headers: authenticated.headers, @@ -65,7 +65,7 @@ pub(super) async fn execute( timeout, )?; - let response = outbound.send(http).await.map_err(|err| { + let response = crate::outbound::send(outbound, http).await.map_err(|err| { // Failing to establish the connection means the request never went out, // so the host can still serve it. Everything else here, a timeout // above all, may have reached the provider and been answered. @@ -88,10 +88,11 @@ pub(super) async fn execute( })); } hooks - .emit(MachineEvent::ResponseReceived { + .on_event(MachineEvent::ResponseReceived { raw: RawResponse { body: text.clone() }, }) - .await?; + .await + .map_err(Error::post_call)?; let body: Value = serde_json::from_str(&text).map_err(|err| { Error::InvalidResponse(litellm_llms::ErrorDetail::invalid( @@ -168,7 +169,7 @@ mod tests { } impl RouteHooks for RecordingHooks { - async fn before_send( + async fn before_provider_request( &self, wire: WireRequest, context: RequestContext, @@ -187,7 +188,7 @@ mod tests { }) } - async fn emit(&self, event: MachineEvent) -> Result<(), Error> { + async fn on_event(&self, event: MachineEvent) -> Result<(), Error> { let MachineEvent::ResponseReceived { raw } = event; self.raw.lock().unwrap().push(raw.body); Ok(()) @@ -240,7 +241,7 @@ mod tests { assert_eq!(request.headers["x-host"], "seen"); assert_eq!(request.headers["x-api-key"], "sk-test"); let [context] = <[RequestContext; 1]>::try_from(hooks.contexts.into_inner().unwrap()) - .unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len())); + .unwrap_or_else(|seen| panic!("before_provider_request runs once, saw {}", seen.len())); assert_eq!( (context.model.as_str(), context.custom_llm_provider.as_str()), ("claude-sonnet-4-5", "anthropic") @@ -275,7 +276,7 @@ mod tests { assert!(hooks.raw.into_inner().unwrap().is_empty()); } - #[test] + #[rstest::rstest] fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() { for original in [ Error::MissingField("usage"), diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index da8cdc43c3e..62e9293c56e 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -1,35 +1,22 @@ -//! The `/chat/completions` call, the Rust equivalent of Python's -//! `litellm.completion()`. -//! -//! [`chat_completions`] is the top-level entrypoint: give it a model, the -//! OpenAI-shaped message list, the provider-mapped optional params, and -//! credentials, and it resolves the provider, translates the conversation, -//! calls the provider, and returns a typed OpenAI-shaped response. -use litellm_secrets::source::SecretSource; - pub mod types; pub use crate::error::RouteError as Error; mod common_utils; pub(crate) mod handler; mod prepare; -use litellm_http::{ClientVariant, HttpClientConfig}; use litellm_types::utils::ChatCompletionsResponse; use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request}; use serde_json::{Map, Value}; use crate::chat_completions::types::ChatCompletionsRequest; +use litellm_auth::AuthServices; +use litellm_secrets::source::SecretSource; +use std::sync::Arc; -pub async fn chat_completions( - resources: &crate::resources::CoreResources, - config: &HttpClientConfig, - secrets: &dyn SecretSource, - request: ChatCompletionsRequest<'_>, -) -> Result { - let http = resources.pool.client(config, ClientVariant::Provider)?; - let resolved = resolve_request(request)?; - let snapshot = secrets.resolve(&resolved.config.secret_names()).await?; - let request = prepare_provider_request(resolved, snapshot)?; - handler::execute(&http, &resources.auth, request, &()).await +#[derive(Clone)] +pub struct ChatCompletionsRoute { + http: litellm_http::Client, + auth: Arc, + secrets: Arc, } /// Whether the core would accept this request, without resolving credentials or @@ -58,3 +45,39 @@ pub fn chat_completions_decline_reason( .unsupported_reason(&messages, optional_params) .map(|reason| reason.0) } + +impl ChatCompletionsRoute { + pub fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self { + Self { + http, + auth, + secrets, + } + } + + pub async fn execute( + &self, + request: ChatCompletionsRequest<'_>, + hooks: &impl litellm_host::hooks::RouteHooks, + ) -> Result { + litellm_host::lifecycle::observe_unary(hooks.observer(), self.run(request, hooks)).await + } + + async fn run( + &self, + request: ChatCompletionsRequest<'_>, + hooks: &impl litellm_host::hooks::RouteHooks, + ) -> Result { + let resolved = resolve_request(request)?; + let snapshot = self + .secrets + .resolve(&resolved.config.secret_names()) + .await?; + let prepared = prepare_provider_request(resolved, snapshot)?; + handler::execute(&self.http, &self.auth, prepared, hooks).await + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index a1e1f5fc838..6b3246d44b5 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -1,5 +1,4 @@ use litellm_auth::SecretValue; -use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}; use litellm_core_utils::settings::Lookup; use litellm_http::request::with_default_headers; use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig}; @@ -9,11 +8,12 @@ use serde_json::Value; use super::{ Error, - common_utils::{chat_completions_provider_config, string_headers}, + common_utils::{chat_completions_provider, string_headers}, }; use crate::chat_completions::types::{ ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest, }; +use crate::provider::resolve_llm_provider; pub(super) struct ResolvedProvider { pub(super) model: String, @@ -25,23 +25,13 @@ pub(super) fn resolve_provider_config<'a>( model: &'a str, custom_llm_provider: Option<&'a str>, ) -> Result { - let provider_info = get_custom_llm_provider(model, custom_llm_provider) - .or_else(|| { - custom_llm_provider.map(|provider| CustomLlmProvider { - model, - custom_llm_provider: provider, - }) - }) - .ok_or_else(|| { - Error::InvalidProvider( - "unable to resolve custom_llm_provider for chat completions request".to_string(), - ) - })?; - let config = chat_completions_provider_config(provider_info.custom_llm_provider) - .ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?; + let provider_info = resolve_llm_provider(model, custom_llm_provider, "chat completions")?; + let config = chat_completions_provider(provider_info.provider) + .ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))? + .config(); Ok(ResolvedProvider { model: provider_info.model.to_string(), - custom_llm_provider: provider_info.custom_llm_provider.to_string(), + custom_llm_provider: <&str>::from(provider_info.provider).to_string(), config, }) } diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index ad6dcc0e403..a301c85fbeb 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -37,6 +37,8 @@ pub enum RouteError { Http(#[from] litellm_http::Error), #[error(transparent)] Secret(#[from] SecretError), + #[error("post-call hook failed: {0}")] + PostCallHook(#[source] Arc), } /// Whether the provider had already been called when the route failed. Before the send, a @@ -47,10 +49,25 @@ pub enum Phase { AfterSend, } +impl From for RouteError { + fn from(fault: litellm_host::machine::MachineFault) -> Self { + use litellm_host::machine::MachineFault; + Self::InvalidRequest(match fault { + MachineFault::Abandoned => "host driver was abandoned".into(), + MachineFault::Protocol(message) => format!("host {message}").into(), + }) + } +} + impl RouteError { + pub(crate) fn post_call(error: Self) -> Self { + Self::PostCallHook(Arc::new(error)) + } + pub fn phase(&self) -> Phase { match self { Self::InvalidResponse(_) + | Self::PostCallHook(_) | Self::Transport(TransportError::Http { .. } | TransportError::Network(_)) => { Phase::AfterSend } @@ -78,9 +95,11 @@ impl RouteError { | Self::Unsupported(_) | Self::Headers(_) => true, Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }), - Self::InvalidResponse(_) | Self::Transport(_) | Self::Http(_) | Self::Secret(_) => { - false - } + Self::InvalidResponse(_) + | Self::Transport(_) + | Self::Http(_) + | Self::Secret(_) + | Self::PostCallHook(_) => false, } } } diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index d373262ae7d..66643eb65ff 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -5,6 +5,7 @@ pub mod error; pub mod messages; pub mod ocr; mod outbound; +mod provider; pub mod resources; pub mod responses; diff --git a/litellm-rust/crates/core/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs index d8b30b31c03..8f5ad05d3dc 100644 --- a/litellm-rust/crates/core/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -7,14 +7,13 @@ use litellm_llms::{ bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG, }; use serde_json::{Map, Value}; -use strum::{EnumString, IntoStaticStr}; use super::Error; +use crate::provider::LlmProviders; const HEADER_CONTEXT: &str = "messages"; -#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)] -#[strum(serialize_all = "snake_case")] +#[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(crate) enum MessagesProvider { Anthropic, AzureAi, @@ -23,7 +22,12 @@ pub(crate) enum MessagesProvider { impl MessagesProvider { pub(crate) fn as_str(self) -> &'static str { - self.into() + match self { + Self::Anthropic => LlmProviders::Anthropic, + Self::AzureAi => LlmProviders::AzureAi, + Self::Bedrock => LlmProviders::Bedrock, + } + .into() } pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig { @@ -35,6 +39,21 @@ impl MessagesProvider { } } +pub(crate) fn messages_provider(provider: LlmProviders) -> Option { + match provider { + LlmProviders::Anthropic => Some(MessagesProvider::Anthropic), + LlmProviders::AzureAi => Some(MessagesProvider::AzureAi), + LlmProviders::Bedrock => Some(MessagesProvider::Bedrock), + LlmProviders::AwsTextract + | LlmProviders::Cohere + | LlmProviders::Mistral + | LlmProviders::Openai + | LlmProviders::OpenaiLike + | LlmProviders::Reducto + | LlmProviders::VertexAi => None, + } +} + pub(super) fn string_headers( extra_headers: Option>, ) -> Result, Error> { @@ -47,8 +66,9 @@ mod tests { use rstest::rstest; - use super::{MessagesProvider, string_headers, truncate_error_body}; + use super::{MessagesProvider, messages_provider, string_headers, truncate_error_body}; use crate::messages::Error; + use crate::provider::LlmProviders; #[rstest] #[case::anthropic("anthropic", MessagesProvider::Anthropic)] @@ -58,13 +78,16 @@ mod tests { #[case] name: &str, #[case] provider: MessagesProvider, ) { - assert_eq!(name.parse::(), Ok(provider)); + assert_eq!( + messages_provider(name.parse::().unwrap()), + Some(provider) + ); assert_eq!(provider.as_str(), name); } #[test] fn provider_without_a_messages_config_is_rejected() { - assert!("openai".parse::().is_err()); + assert_eq!(messages_provider(LlmProviders::Openai), None); } #[test] diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index 9fa1a594b76..604a9abaca7 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -15,7 +15,7 @@ use litellm_llms::base_llm::{ transformation::BaseAnthropicMessagesConfig, }, }; -use litellm_tracing::{ByteChunk, debug}; +use litellm_tracing::ByteChunk; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use serde_json::Value; @@ -48,7 +48,7 @@ pub(super) async fn execute( }; let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?; let wire = hooks - .before_send( + .before_provider_request( WireRequest { url, headers: authenticated.headers, @@ -58,7 +58,7 @@ pub(super) async fn execute( ) .await?; let provider_name = provider.as_str(); - debug!(provider = provider_name, stream, body = %wire.body, "provider request"); + log_request_body(provider_name, stream, &wire.body); let response = send( http, Authenticated { @@ -70,11 +70,6 @@ pub(super) async fn execute( timeout, ) .await?; - debug!( - provider = provider_name, - status = response.status().as_u16(), - "provider response headers" - ); if !response.status().is_success() { return Err(provider_error(response).await); } @@ -87,14 +82,15 @@ pub(super) async fn execute( )); } let text = response.text().await.map_err(network)?; - debug!(body = text.as_str(), "provider response body"); + log_response_body(&text); hooks - .emit(MachineEvent::ResponseReceived { + .on_event(MachineEvent::ResponseReceived { raw: RawResponse { body: text.clone() }, }) - .await?; + .await + .map_err(Error::post_call)?; decode_response(config, &body.model, &text) - .map(|message| MessagesResponse::Message(Box::new(message))) + .map(|message| MessagesResponse::Complete(Box::new(message))) } fn serialize_failure(err: serde_json::Error) -> Error { @@ -121,14 +117,14 @@ async fn send( body, Some(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))), )?; - request.send(http).await.map_err(network) + crate::outbound::send(request, http).await.map_err(network) } async fn provider_error(response: reqwest::Response) -> Error { let status = response.status().as_u16(); match response.text().await { Ok(text) => { - litellm_tracing::debug!(status, body = text.as_str(), "provider error body"); + log_error_body(status, &text); Error::Transport(TransportError::Http { status, body: truncate_error_body(&text), @@ -175,7 +171,10 @@ fn streaming_response( .boxed(), Some(decode) => decoded_chunks(response, decode, provider), }; - MessagesResponse::Stream { headers, chunks } + MessagesResponse::Stream { + head: super::route::MessagesStreamHead { headers }, + chunks, + } } fn decoded_chunks( @@ -199,9 +198,21 @@ fn decoded_chunks( .boxed() } -fn log_chunk(provider: &str, stage: &str, data: &Bytes) { +fn log_request_body(provider: &str, stream: bool, body: &serde_json::Value) { + litellm_tracing::debug!(provider, stream, body = %body, "provider request"); +} + +fn log_response_body(body: &str) { + litellm_tracing::debug!(body, "provider response body"); +} + +fn log_error_body(status: u16, body: &str) { + litellm_tracing::debug!(status, body, "provider error body"); +} + +fn log_chunk(provider: &str, stage: &str, data: &bytes::Bytes) { let chunk = ByteChunk::new(data); - debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk"); + litellm_tracing::debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk"); } #[cfg(test)] @@ -217,6 +228,7 @@ mod tests { "data: {\"type\":\"ping\"}\n\n", Some("event: ping\ndata: {\"type\":\"ping\"}\n\n") )] + #[rstest::rstest] #[case::invalid_event("data: invalid\n\ndata: {\"type\":\"ping\"}\n\n", None)] #[tokio::test] async fn decoded_streams_encode_events_and_stop_at_the_first_error( diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index f3a57da4d32..6a08df5b411 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -1,27 +1,50 @@ -//! The Anthropic Messages call, the Rust equivalent of Python's `litellm.messages()`. -//! -//! [`messages`] prepares the provider request and sends it in process. [`route`] runs the -//! same two steps as a machine for a host that answers the call's operations itself. - mod common_utils; mod handler; mod prepare; pub mod route; mod types; -use litellm_http::{ClientVariant, HttpClientConfig}; +use litellm_auth::AuthServices; use litellm_secrets::source::SecretSource; +use std::sync::Arc; pub use crate::error::RouteError as Error; pub use types::{MessagesCall, MessagesResponse, MessagesShaping, messages_body}; -pub async fn messages( - resources: &crate::resources::CoreResources, - config: &HttpClientConfig, - secrets: &dyn SecretSource, - call: MessagesCall, -) -> Result { - let http = resources.pool.client(config, ClientVariant::Provider)?; - let request = prepare::prepare(call, secrets).await?; - handler::execute(&http, &resources.auth, request, &()).await +#[derive(Clone)] +pub struct MessagesRoute { + http: litellm_http::Client, + auth: Arc, + secrets: Arc, +} + +impl MessagesRoute { + pub fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self { + Self { + http, + auth, + secrets, + } + } + + pub async fn execute( + &self, + call: MessagesCall, + hooks: &impl litellm_host::hooks::RouteHooks, + ) -> Result { + litellm_host::lifecycle::observe_call(hooks.observer(), self.run(call, hooks)).await + } + + async fn run( + &self, + call: MessagesCall, + hooks: &impl litellm_host::hooks::RouteHooks, + ) -> Result { + let request = prepare::prepare(call, self.secrets.as_ref()).await?; + handler::execute(&self.http, &self.auth, request, hooks).await + } } diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index c9fb0542101..0e4c72e4270 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -3,9 +3,7 @@ use std::time::Duration; use litellm_auth::SecretValue; use litellm_core_utils::{ dot_notation_indexing::delete_nested_value, - get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}, - get_provider_specific_headers::get_provider_specific_headers, - settings::Lookup, + get_provider_specific_headers::get_provider_specific_headers, settings::Lookup, }; use litellm_http::request::with_default_headers; use litellm_llms::base_llm::{ @@ -16,9 +14,10 @@ use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessage use super::{ Error, MessagesCall, - common_utils::{MessagesProvider, string_headers}, + common_utils::{MessagesProvider, messages_provider, string_headers}, types::invalid_request, }; +use crate::provider::resolve_llm_provider; struct ResolvedProvider { model: String, @@ -50,26 +49,11 @@ fn resolve_provider( model: &str, custom_llm_provider: Option<&str>, ) -> Result { - let CustomLlmProvider { - model, - custom_llm_provider: provider, - } = get_custom_llm_provider(model, custom_llm_provider) - .or_else(|| { - custom_llm_provider.map(|provider| CustomLlmProvider { - model, - custom_llm_provider: provider, - }) - }) - .ok_or_else(|| { - Error::InvalidProvider( - "unable to resolve custom_llm_provider for messages request".to_string(), - ) - })?; - let provider = provider - .parse() - .map_err(|_| Error::InvalidProvider(provider.to_string()))?; + let resolved = resolve_llm_provider(model, custom_llm_provider, "messages")?; + let provider = messages_provider(resolved.provider) + .ok_or_else(|| Error::InvalidProvider(<&str>::from(resolved.provider).to_string()))?; Ok(ResolvedProvider { - model: model.to_string(), + model: resolved.model.to_string(), provider, }) } diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 3762a04a775..5a1880d0a99 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -1,26 +1,15 @@ -use std::{ - convert::Infallible, - sync::{Arc, Mutex}, -}; +use std::convert::Infallible; use bytes::Bytes; -use futures_util::TryStreamExt; use litellm_host::{ - host::{Demand, Host}, - machine::{CallMachine, HostChannel, MachineFault}, + call::{HostedCompletion, HostedMachine, hosted_call}, protocol::Protocol, }; -use litellm_http::{Client, ClientVariant, HttpClientConfig}; -use litellm_secrets::source::SecretSource; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; -use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare}; +use super::{Error, MessagesCall}; -pub enum MessagesOutput { - Message(Box), - /// Every chunk already reached the host through `Deliver`. - Streamed, -} +pub type MessagesOutput = HostedCompletion>; /// The upstream response as the caller sees it at stream hand-off, before any chunk. pub struct MessagesStreamHead { @@ -30,91 +19,20 @@ pub struct MessagesStreamHead { pub struct Messages; impl Protocol for Messages { - type Response = MessagesOutput; + type Response = Box; type Error = Error; - type Projection = MessagesCall; - type Op = Infallible; + type Request = MessagesCall; + type HostCall = Infallible; type Chunk = Bytes; type StreamHead = MessagesStreamHead; } -impl From for Error { - fn from(fault: MachineFault) -> Self { - Self::InvalidRequest(match fault { - MachineFault::Abandoned => "messages host driver was abandoned".into(), - MachineFault::Protocol(message) => format!("messages {message}").into(), +pub type MessagesMachine = HostedMachine; + +impl super::MessagesRoute { + pub fn machine(self, request: super::MessagesCall) -> MessagesMachine { + hosted_call(request, move |call, _, hooks| async move { + self.run(call, &hooks).await }) } } - -pub type MessagesHost = HostChannel; -pub type MessagesMachine = CallMachine; - -/// The in-process host for a request already in hand. It answers projection once and -/// observes nothing. -pub struct LocalMessagesHost { - call: Mutex>, -} - -impl LocalMessagesHost { - pub fn new(call: MessagesCall) -> Self { - Self { - call: Mutex::new(Some(call)), - } - } -} - -impl Host for LocalMessagesHost { - async fn project(&self) -> Result { - self.call - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .ok_or_else(|| Error::InvalidRequest("messages request was already projected".into())) - } - - async fn custom_op(&self, op: Infallible) -> Result<(), Error> { - match op {} - } -} - -pub fn messages_machine( - resources: &crate::resources::CoreResources, - config: &HttpClientConfig, - secrets: Arc, -) -> Result { - let http = resources.pool.client(config, ClientVariant::Provider)?; - let auth = resources.auth.clone(); - Ok(CallMachine::new(move |host| { - Box::pin(drive(host, http, auth, secrets)) - })) -} - -/// The call as its host sees it: projection first, then the same prepare and execute as -/// [`super::messages`], with each chunk of a stream handed over as it arrives. -async fn drive( - host: MessagesHost, - http: Client, - auth: Arc, - secrets: Arc, -) -> Result { - let call = host.project().await?; - let request = prepare(call, secrets.as_ref()).await?; - match execute(&http, &auth, request, &host).await? { - MessagesResponse::Message(message) => Ok(MessagesOutput::Message(message)), - MessagesResponse::Stream { - headers, - mut chunks, - } => { - if host.open(MessagesStreamHead { headers }).await? == Demand::Detached { - return Ok(MessagesOutput::Streamed); - } - while let Some(chunk) = chunks.try_next().await? { - if host.deliver(chunk).await? == Demand::Detached { - break; - } - } - Ok(MessagesOutput::Streamed) - } - } -} diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index 3cb76369c57..bc77b1dbded 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -1,8 +1,8 @@ use std::time::Duration; use bytes::Bytes; -use futures_util::stream::BoxStream; -use litellm_llms::base_llm::messages::context::MessagesModelCapabilities; +use litellm_host::call::CallOutput; +use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities; use litellm_types::{ llms::anthropic_messages::{ anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, @@ -30,24 +30,16 @@ pub fn messages_body(body: Map) -> Result Error { - Error::InvalidRequest(litellm_llms::ErrorDetail::invalid( - "Anthropic messages request", - err, - )) + Error::InvalidRequest(format!("invalid Anthropic messages request: {err}").into()) } -pub enum MessagesResponse { - Message(Box), - Stream { - headers: Vec<(String, String)>, - chunks: BoxStream<'static, Result>, - }, -} +pub type MessagesResponse = + CallOutput, super::route::MessagesStreamHead, Bytes, Error>; #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] pub struct MessagesShaping { #[serde(default)] - pub capabilities: MessagesModelCapabilities, + pub capabilities: AnthropicModelCapabilities, #[serde(default)] pub drop_params: bool, #[serde(default)] @@ -84,9 +76,9 @@ mod tests { #[case::partial_capabilities( json!({"capabilities": {"supports_reasoning": true}}), MessagesShaping { - capabilities: MessagesModelCapabilities { + capabilities: AnthropicModelCapabilities { supports_reasoning: true, - ..MessagesModelCapabilities::default() + ..AnthropicModelCapabilities::default() }, ..MessagesShaping::default() }, @@ -108,7 +100,7 @@ mod tests { "additional_drop_params": ["metadata.user_id", "thinking"] }), MessagesShaping { - capabilities: MessagesModelCapabilities { + capabilities: AnthropicModelCapabilities { supports_reasoning: true, supports_adaptive_thinking: true, thinking_always_on: false, diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index e635f93a294..4ab11f6692e 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -1,15 +1,58 @@ +use std::sync::Arc; + +use litellm_host::hooks::RouteHooks; use litellm_llms::base_llm::ocr::{ error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse, }; -use crate::ocr::{ - route::{LocalOcrHost, ocr_machine}, - types::LiteLLMOcrRequest, +use super::{ + handler::perform_ocr_request, + types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest}, }; -pub async fn perform( - client: &OcrClient, - request: LiteLLMOcrRequest, -) -> Result { - litellm_host::run::run(ocr_machine(client.clone()), &LocalOcrHost::new(request)).await +#[derive(Clone)] +pub struct OcrRoute { + client: OcrClient, +} + +impl OcrRoute { + pub fn new(client: OcrClient) -> Self { + Self { client } + } + + pub async fn execute( + &self, + request: LiteLLMOcrRequest, + hooks: &impl RouteHooks, + ) -> Result { + litellm_host::lifecycle::observe_unary(hooks.observer(), self.run(request, hooks)).await + } + + pub(super) async fn run( + &self, + request: LiteLLMOcrRequest, + hooks: &impl RouteHooks, + ) -> Result { + let caller_document = matches!(&request.document, OcrDocumentInput::Document(_)); + let prepared = prepare_request_document(request).await?; + let execute: futures_util::future::BoxFuture<'_, Result> = + Box::pin(perform_ocr_request( + &self.client, + prepared, + hooks, + caller_document, + )); + execute.await + } +} + +async fn prepare_request_document( + request: LiteLLMOcrRequest, +) -> Result { + if let OcrDocumentInput::Document(_) = &request.document { + return request.map_document(super::document::prepare_document); + } + tokio::task::spawn_blocking(move || request.map_document(super::document::prepare_document)) + .await + .map_err(|error| Error::DocumentTask(Arc::new(error)))? } diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index c1265e1e91c..1c3705a58ac 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -1,6 +1,8 @@ use futures_util::future::BoxFuture; -use litellm_auth::SecretValue; -use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest}; +use litellm_host::{ + event::{MachineEvent, RawResponse, RequestContext, WireRequest}, + hooks::RouteHooks, +}; use litellm_llms::base_llm::ocr::{ error::Error, handler::{CallHooks, OcrClient}, @@ -8,16 +10,13 @@ use litellm_llms::base_llm::ocr::{ }; use serde_json::Value; -use super::{ - arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind, - route::OcrHost, -}; +use super::{arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind}; use crate::ocr::types::ResolvedOcrRequest; pub(crate) async fn perform_ocr_request( client: &OcrClient, request: ResolvedOcrRequest, - host: &OcrHost, + host: &impl RouteHooks, caller_document: bool, ) -> Result { request.response_format()?; @@ -28,53 +27,48 @@ pub(crate) async fn perform_ocr_request( .await .map_err(|error| Error::Secret(std::sync::Arc::new(error)))?; let request = prepare_request(request, caller_document, client, secrets); - let hooks = OcrCallHooks::new(host.clone(), &request, config); + let hooks = OcrCallHooks::new(host, &request, config); config.ocr(client, &request, &hooks).await } -/// Lets provider code reach the host mid-call, filling in the request context only the -/// route knows. -pub(crate) struct OcrCallHooks { - host: OcrHost, - model: String, - custom_llm_provider: &'static str, - optional_params: Value, - secret_fields: Vec, - api_key: Option, +struct OcrCallHooks<'a, H> { + hooks: &'a H, + context: RequestContext, } -impl OcrCallHooks { - pub(crate) fn new(host: OcrHost, request: &PreparedOcrRequest, config: OcrConfigKind) -> Self { +impl<'a, H> OcrCallHooks<'a, H> { + fn new(hooks: &'a H, request: &PreparedOcrRequest, config: OcrConfigKind) -> Self { Self { - host, - model: request.model.clone(), - custom_llm_provider: config.provider().into(), - optional_params: Value::Object(request.optional_params.clone().into()), - secret_fields: request - .optional_params - .keys() - .filter(|name| is_secret_param(name)) - .cloned() - .collect(), - api_key: request.connection.api_key.clone(), + hooks, + context: RequestContext { + model: request.model.clone(), + custom_llm_provider: <&str>::from(config.provider()).to_owned(), + optional_params: Value::Object(request.optional_params.clone().into()), + secret_fields: request + .optional_params + .keys() + .filter(|name| is_secret_param(name)) + .cloned() + .collect(), + api_key: request.connection.api_key.clone(), + }, } } } -impl CallHooks for OcrCallHooks { - fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result> { - let context = RequestContext { - model: self.model.clone(), - custom_llm_provider: self.custom_llm_provider.into(), - optional_params: self.optional_params.clone(), - secret_fields: self.secret_fields.clone(), - api_key: self.api_key.clone(), - }; - Box::pin(self.host.before_send(wire, context)) +impl> CallHooks for OcrCallHooks<'_, H> { + fn before_provider_request( + &self, + wire: WireRequest, + ) -> BoxFuture<'_, Result> { + Box::pin( + self.hooks + .before_provider_request(wire, self.context.clone()), + ) } fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> { - Box::pin(self.host.emit(MachineEvent::ResponseReceived { + Box::pin(self.hooks.on_event(MachineEvent::ResponseReceived { raw: RawResponse { body: String::from_utf8_lossy(body).into_owned(), }, diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index 2a0d20f69c9..45b664961b2 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,5 +1,6 @@ pub mod arguments; -pub mod client; +mod client; +pub use client::OcrRoute; pub mod document; pub(crate) mod handler; pub(crate) mod prepare; diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index f13d6984763..cf2bae1ab9f 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -5,8 +5,8 @@ use litellm_llms::base_llm::ocr::{ }; use litellm_secrets::source::Secrets; -use super::provider_config::OcrProvider; use crate::ocr::types::{LiteLLMOcrRequest, ResolvedOcrRequest}; +use crate::provider::LlmProviders; pub(crate) fn prepare_request( request: ResolvedOcrRequest, @@ -16,15 +16,19 @@ pub(crate) fn prepare_request( ) -> PreparedOcrRequest { let credentials = request.credentials.clone(); let (preferred_api_key_env, api_base_env) = match request.config.provider() { - OcrProvider::Mistral => ( + LlmProviders::Mistral => ( Some("MISTRAL_AZURE_API_KEY"), Some("MISTRAL_AZURE_API_BASE"), ), - OcrProvider::AzureAi => (None, Some("AZURE_AI_API_BASE")), - OcrProvider::AwsTextract - | OcrProvider::Cohere - | OcrProvider::Reducto - | OcrProvider::VertexAi => (None, None), + LlmProviders::AzureAi => (None, Some("AZURE_AI_API_BASE")), + LlmProviders::Anthropic + | LlmProviders::AwsTextract + | LlmProviders::Bedrock + | LlmProviders::Cohere + | LlmProviders::Openai + | LlmProviders::OpenaiLike + | LlmProviders::Reducto + | LlmProviders::VertexAi => (None, None), }; let secret = |name: &str| secrets.truthy(name); let dynamic_api_key = credentials.dynamic_api_key.or_else(|| { @@ -100,7 +104,10 @@ mod tests { struct NoHooks; impl CallHooks for NoHooks { - fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result> { + fn before_provider_request( + &self, + wire: WireRequest, + ) -> BoxFuture<'_, Result> { Box::pin(async move { Ok(wire) }) } @@ -144,6 +151,7 @@ mod tests { json!({"type": "image_url", "image_url": url}) } + #[rstest::rstest] #[tokio::test] async fn cohere_body_keeps_native_document_fields_and_untyped_overrides() { let request = request( @@ -176,6 +184,7 @@ mod tests { ); } + #[rstest::rstest] #[tokio::test] async fn explicit_null_options_use_defaults_before_http() { let request = request( @@ -199,6 +208,7 @@ mod tests { assert!(body.get("req_format").is_none()); } + #[rstest::rstest] #[tokio::test] async fn direct_and_vertex_mistral_build_the_same_request_and_share_normalization() { let options = json!({ @@ -276,7 +286,7 @@ mod tests { pages: Option>, } - #[test] + #[rstest::rstest] fn parsed_provider_params_separates_known_and_extra_params() { let arguments: CallArguments = serde_json::from_value(json!({ "pages": [0, 2], diff --git a/litellm-rust/crates/core/src/ocr/provider_config.rs b/litellm-rust/crates/core/src/ocr/provider_config.rs index 105a4da958e..3b4e9d88da4 100644 --- a/litellm-rust/crates/core/src/ocr/provider_config.rs +++ b/litellm-rust/crates/core/src/ocr/provider_config.rs @@ -1,3 +1,4 @@ +use crate::provider::LlmProviders; use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}; use litellm_llms::{ aws_textract::ocr::{ @@ -24,7 +25,6 @@ use litellm_llms::{ deepseek_transformation::VertexAIDeepSeekOCRConfig, transformation::VertexAiOcrConfig, }, }; -use strum::{EnumString, IntoStaticStr}; macro_rules! with_config { ($kind:expr, $config:ident => $body:expr) => { @@ -93,16 +93,16 @@ pub(crate) enum OcrConfigKind { } impl OcrConfigKind { - pub(crate) const fn provider(self) -> OcrProvider { + pub(crate) const fn provider(self) -> LlmProviders { match self { - Self::AwsTextract | Self::AwsTextractAnalyze => OcrProvider::AwsTextract, - Self::Cohere => OcrProvider::Cohere, - Self::Mistral => OcrProvider::Mistral, + Self::AwsTextract | Self::AwsTextractAnalyze => LlmProviders::AwsTextract, + Self::Cohere => LlmProviders::Cohere, + Self::Mistral => LlmProviders::Mistral, Self::AzureAi | Self::AzureCohere | Self::AzureDocumentIntelligence => { - OcrProvider::AzureAi + LlmProviders::AzureAi } - Self::ReductoLegacy | Self::ReductoV3 => OcrProvider::Reducto, - Self::VertexAi | Self::VertexDeepSeek => OcrProvider::VertexAi, + Self::ReductoLegacy | Self::ReductoV3 => LlmProviders::Reducto, + Self::VertexAi | Self::VertexDeepSeek => LlmProviders::VertexAi, } } @@ -187,17 +187,6 @@ pub fn passthrough_response( .map(Some) } -#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)] -#[strum(serialize_all = "snake_case")] -pub(crate) enum OcrProvider { - AwsTextract, - Cohere, - Mistral, - AzureAi, - Reducto, - VertexAi, -} - pub(crate) fn resolve_provider_config( model: &str, custom_llm_provider: Option<&str>, @@ -205,37 +194,45 @@ pub(crate) fn resolve_provider_config( let provider = get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider { model, - custom_llm_provider: OcrProvider::Mistral.into(), + custom_llm_provider: LlmProviders::Mistral.into(), }); - let ocr_provider = provider + let llm_provider = provider .custom_llm_provider - .parse::() + .parse::() .map_err(|_| Error::InvalidProvider(provider.custom_llm_provider.to_string()))?; - let config = match ocr_provider { - OcrProvider::AwsTextract => match TextractOperation::from_model(provider.model)? { + let config = match llm_provider { + LlmProviders::AwsTextract => match TextractOperation::from_model(provider.model)? { TextractOperation::DetectDocumentText => OcrConfigKind::AwsTextract, TextractOperation::AnalyzeDocument => OcrConfigKind::AwsTextractAnalyze, }, - OcrProvider::Cohere => OcrConfigKind::Cohere, - OcrProvider::Mistral => OcrConfigKind::Mistral, - OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => { + LlmProviders::Cohere => OcrConfigKind::Cohere, + LlmProviders::Mistral => OcrConfigKind::Mistral, + LlmProviders::AzureAi if is_document_intelligence_model(provider.model) => { OcrConfigKind::AzureDocumentIntelligence } - OcrProvider::AzureAi + LlmProviders::AzureAi if provider.model.to_ascii_lowercase().contains("cohere") && provider.model.to_ascii_lowercase().contains("parse") => { OcrConfigKind::AzureCohere } - OcrProvider::AzureAi => OcrConfigKind::AzureAi, - OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => { + LlmProviders::AzureAi => OcrConfigKind::AzureAi, + LlmProviders::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => { OcrConfigKind::ReductoLegacy } - OcrProvider::Reducto => OcrConfigKind::ReductoV3, - OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => { + LlmProviders::Reducto => OcrConfigKind::ReductoV3, + LlmProviders::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => { OcrConfigKind::VertexDeepSeek } - OcrProvider::VertexAi => OcrConfigKind::VertexAi, + LlmProviders::VertexAi => OcrConfigKind::VertexAi, + LlmProviders::Anthropic + | LlmProviders::Bedrock + | LlmProviders::Openai + | LlmProviders::OpenaiLike => { + return Err(Error::InvalidProvider( + provider.custom_llm_provider.to_string(), + )); + } }; Ok((provider.model.to_string(), config)) } diff --git a/litellm-rust/crates/core/src/ocr/route.rs b/litellm-rust/crates/core/src/ocr/route.rs index bbddec9b2c4..4e96fd2ebd9 100644 --- a/litellm-rust/crates/core/src/ocr/route.rs +++ b/litellm-rust/crates/core/src/ocr/route.rs @@ -1,25 +1,20 @@ -use std::sync::{Arc, Mutex}; - use litellm_auth::ResolvedCredential; use litellm_host::{ - event::{CallEvent, RequestContext, WireRequest}, - host::Reply, - machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol}, + call::{CallOutput, HostedMachine, hosted_call}, + machine::{HostTokenProvider, TokenProtocol}, protocol::Protocol, + protocol::Reply, }; -use litellm_llms::base_llm::ocr::{ - error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse, -}; +use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}; -use super::handler::perform_ocr_request; -use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest}; +use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput}; pub enum OcrOp { AcquireAzureAdToken(Reply), } /// The caller's request as the host projects it. -pub struct OcrProjection { +pub struct OcrCall { pub request: LiteLLMOcrRequest, /// The caller passed its own Azure AD token provider, which the host keeps. pub caller_token: bool, @@ -30,8 +25,8 @@ pub struct Ocr; impl Protocol for Ocr { type Response = LiteLLMOcrResponse; type Error = Error; - type Projection = OcrProjection; - type Op = OcrOp; + type Request = OcrCall; + type HostCall = OcrOp; type Chunk = std::convert::Infallible; type StreamHead = std::convert::Infallible; } @@ -42,122 +37,22 @@ impl TokenProtocol for Ocr { } } -pub type OcrHost = HostChannel; -pub type OcrMachine = CallMachine; +pub type OcrMachine = HostedMachine; -/// The OCR call as a machine: projection and token acquisition are host operations; -/// everything else runs in Rust. -pub fn ocr_machine(client: OcrClient) -> OcrMachine { - CallMachine::new(move |host| Box::pin(execute(client, host))) -} - -async fn execute(client: OcrClient, host: OcrHost) -> Result { - let OcrProjection { - request, - caller_token, - } = host.project().await?; - let request = LiteLLMOcrRequest { - azure_ad_token_provider: caller_token - .then(|| HostTokenProvider::handle(host.clone())) - .or(request.azure_ad_token_provider), - ..request - }; - let caller_document = matches!(request.document, OcrDocumentInput::Document(_)); - let request = prepare_request_document(request).await?; - perform_ocr_request(&client, request, &host, caller_document).await -} - -async fn prepare_request_document( - request: LiteLLMOcrRequest, -) -> Result { - if let OcrDocumentInput::Document(_) = &request.document { - return request.map_document(super::document::prepare_document); - } - tokio::task::spawn_blocking(move || request.map_document(super::document::prepare_document)) - .await - .map_err(|error| Error::DocumentTask(Arc::new(error)))? -} - -type BeforeSend = - Box Result + Send + Sync>; -type Observer = Box; - -/// The in-process host for a request that is already in hand: the request answers -/// projection, and the optional observer sees and may rewrite the wire request. -pub struct LocalOcrHost { - request: Mutex>>, - before_send: Option, - observer: Option, -} - -impl LocalOcrHost { - pub fn new(request: LiteLLMOcrRequest) -> Self { - Self { - request: Mutex::new(Some(request)), - before_send: None, - observer: None, - } - } - - pub fn with_before_send( - self, - before_send: impl Fn(WireRequest, &RequestContext) -> Result - + Send - + Sync - + 'static, - ) -> Self { - Self { - before_send: Some(Box::new(before_send)), - ..self - } - } - - pub fn with_observer(self, observer: impl Fn(&CallEvent) + Send + Sync + 'static) -> Self { - Self { - observer: Some(Box::new(observer)), - ..self - } - } -} - -impl litellm_host::host::Host for LocalOcrHost { - async fn project(&self) -> Result { - self.request - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .map(|request| OcrProjection { - request, - caller_token: false, - }) - .ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())) - } - - async fn custom_op(&self, op: OcrOp) -> Result<(), Error> { - match op { - OcrOp::AcquireAzureAdToken(_) => { - Err(Error::Auth(litellm_auth::Error::CredentialAcquisition( - "OCR host has no Azure AD token provider".into(), - ))) - } - } - } - - async fn before_send( - &self, - wire: WireRequest, - context: &RequestContext, - ) -> Result { - match &self.before_send { - Some(before_send) => before_send(wire, context), - None => Ok(wire), - } - } - - async fn emit(&self, event: &CallEvent) -> Result<(), Error> { - if let Some(observer) = &self.observer { - observer(event); - } - Ok(()) +impl crate::ocr::OcrRoute { + pub fn machine(self, request: OcrCall) -> OcrMachine { + hosted_call( + request, + move |projection: OcrCall, services, hooks| async move { + let request = LiteLLMOcrRequest { + azure_ad_token_provider: projection + .caller_token + .then(|| HostTokenProvider::handle(services)) + .or(projection.request.azure_ad_token_provider), + ..projection.request + }; + self.run(request, &hooks).await.map(CallOutput::Complete) + }, + ) } } diff --git a/litellm-rust/crates/core/src/outbound.rs b/litellm-rust/crates/core/src/outbound.rs index 0cdbb465f60..6d16b2d88af 100644 --- a/litellm-rust/crates/core/src/outbound.rs +++ b/litellm-rust/crates/core/src/outbound.rs @@ -4,6 +4,13 @@ use litellm_http::outbound::OutboundRequest; use litellm_llms::base_llm::auth::Authenticated; use serde_json::Value; +pub(crate) async fn send( + request: OutboundRequest, + client: &litellm_http::Client, +) -> Result { + request.send(client).await +} + /// Header credentials are already in `headers`; SigV4 is applied here, over the /// bytes that are sent. pub(crate) fn outbound_request( diff --git a/litellm-rust/crates/core/src/provider.rs b/litellm-rust/crates/core/src/provider.rs new file mode 100644 index 00000000000..e071aee8084 --- /dev/null +++ b/litellm-rust/crates/core/src/provider.rs @@ -0,0 +1,36 @@ +pub use litellm_core_utils::get_llm_provider_logic::LlmProviders; +use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}; + +use crate::error::RouteError as Error; + +#[derive(Debug)] +pub(crate) struct ResolvedProvider<'a> { + pub(crate) model: &'a str, + pub(crate) provider: LlmProviders, +} + +pub(crate) fn resolve_llm_provider<'a>( + model: &'a str, + custom_llm_provider: Option<&'a str>, + route: &'static str, +) -> Result, Error> { + let CustomLlmProvider { + model, + custom_llm_provider, + } = get_custom_llm_provider(model, custom_llm_provider) + .or_else(|| { + custom_llm_provider.map(|provider| CustomLlmProvider { + model, + custom_llm_provider: provider, + }) + }) + .ok_or_else(|| { + Error::InvalidProvider(format!( + "unable to resolve custom_llm_provider for {route} request" + )) + })?; + let provider = custom_llm_provider + .parse() + .map_err(|_| Error::InvalidProvider(custom_llm_provider.to_string()))?; + Ok(ResolvedProvider { model, provider }) +} diff --git a/litellm-rust/crates/core/src/resources.rs b/litellm-rust/crates/core/src/resources.rs index 37a29502649..d9ba5e89097 100644 --- a/litellm-rust/crates/core/src/resources.rs +++ b/litellm-rust/crates/core/src/resources.rs @@ -1,9 +1,7 @@ use std::sync::Arc; use litellm_auth::AuthServices; -use litellm_http::{HttpClientConfig, HttpClientPool, media::UrlPolicy}; -use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings}; -use litellm_secrets::source::SecretSource; +use litellm_http::HttpClientPool; #[derive(Clone)] pub struct CoreResources { @@ -18,21 +16,4 @@ impl CoreResources { auth: Arc::new(AuthServices::default()), } } - - pub fn ocr_client( - &self, - config: &HttpClientConfig, - url_policy: UrlPolicy, - settings: OcrSettings, - secrets: Arc, - ) -> Result { - OcrClient::new( - &self.pool, - config, - url_policy, - self.auth.clone(), - settings, - secrets, - ) - } } diff --git a/litellm-rust/crates/core/tests/audio_transcription.rs b/litellm-rust/crates/core/tests/audio_transcription.rs index daf48b19684..1e41df75168 100644 --- a/litellm-rust/crates/core/tests/audio_transcription.rs +++ b/litellm-rust/crates/core/tests/audio_transcription.rs @@ -1,6 +1,4 @@ -use litellm_core::audio_transcription::{ - Error, audio_transcription, types::AudioTranscriptionRequest, -}; +use litellm_core::audio_transcription::{Error, types::AudioTranscriptionRequest}; use rstest::{fixture, rstest}; use serde_json::{Map, Value, json}; use wiremock::ResponseTemplate; @@ -11,19 +9,11 @@ use support::*; const MODEL: &str = "mistral.voxtral-mini-3b-2507"; async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result { - audio_transcription( - &support::resources(), - &http_config(), - &RecordingSecrets::empty(), - request, - ) - .await + audio_transcription_route().execute(request).await } fn transcript_response(text: &str) -> ResponseTemplate { - json_response( - json!({"output": {"message": {"content": [{"text": text}]}}, "usage": {"inputTokens": 1, "outputTokens": 1}}), - ) + json_response(json!({"output": {"message": {"content": [{"text": text}]}}})) } fn aws_params(region: &str) -> Map { @@ -70,7 +60,6 @@ async fn bedrock_converse_request_is_signed_for_the_requested_region( assert_eq!(response, json!({"text": "hello"})); let sent = only_request(&upstream).await; assert_eq!(sent.method.as_str(), "POST"); - assert_eq!(sent.header("content-type"), Some("application/json")); assert_eq!(sent.url.path(), format!("/model/{MODEL}/converse")); let authorization = sent.header("authorization").expect("request is signed"); assert!( @@ -261,61 +250,3 @@ async fn an_unreadable_success_body_is_an_invalid_response( assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}"); } - -#[rstest] -#[tokio::test] -async fn injected_secrets_supply_signing_credentials_and_region( - request: AudioTranscriptionRequest<'static>, -) { - let upstream = upstream([transcript_response("hello")]).await; - let base = upstream.uri(); - let secrets = RecordingSecrets::new([ - ("AWS_ACCESS_KEY_ID", "injected-access-key"), - ("AWS_SECRET_ACCESS_KEY", "injected-secret-key"), - ("AWS_REGION_NAME", "eu-west-1"), - ("AWS_SESSION_TOKEN", "injected-session-token"), - ]); - let response = audio_transcription( - &support::resources(), - &http_config(), - &secrets, - AudioTranscriptionRequest { - api_base: Some(&base), - optional_params: Map::new(), - ..request - }, - ) - .await - .unwrap(); - assert_eq!(response, json!({"text": "hello"})); - let sent = only_request(&upstream).await; - let authorization = sent.header("authorization").unwrap(); - assert!(authorization.contains("Credential=injected-access-key/")); - assert!(authorization.contains("/eu-west-1/bedrock/aws4_request")); - assert_eq!( - sent.header("x-amz-security-token"), - Some("injected-session-token") - ); - assert!(!sent.body_text().contains("injected-secret-key")); -} - -#[rstest] -#[tokio::test] -async fn secret_resolution_failure_prevents_transcription( - request: AudioTranscriptionRequest<'static>, -) { - let upstream = upstream([transcript_response("hello")]).await; - let base = upstream.uri(); - let result = audio_transcription( - &support::resources(), - &http_config(), - &RecordingSecrets::failing(), - AudioTranscriptionRequest { - api_base: Some(&base), - ..request - }, - ) - .await; - assert!(matches!(result, Err(Error::Secret(_)))); - assert!(received(&upstream).await.is_empty()); -} diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs index c2a47cdcd4f..bbd4b083a4c 100644 --- a/litellm-rust/crates/core/tests/chat_completions.rs +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -1,7 +1,7 @@ use std::time::Duration; use litellm_core::chat_completions::{ - Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest, + Error, chat_completions_decline_reason, types::ChatCompletionsRequest, }; use litellm_http::transport::Error as TransportError; use litellm_types::utils::ChatCompletionsResponse; @@ -15,13 +15,7 @@ use support::*; const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#; async fn complete(request: ChatCompletionsRequest<'_>) -> Result { - chat_completions( - &support::resources(), - &http_config(), - &RecordingSecrets::empty(), - request, - ) - .await + chat_completions_route().execute(request, &()).await } fn object(value: Value) -> Map { @@ -329,164 +323,3 @@ async fn a_declined_request_fails_the_call_before_sending( assert_eq!(error, Error::Unsupported("streaming")); assert!(received(&upstream).await.is_empty()); } - -#[rstest] -#[case::source_key(None, "source-key")] -#[case::explicit_key(Some("explicit-key"), "explicit-key")] -#[tokio::test] -async fn injected_secrets_supply_credentials_and_endpoint( - request: ChatCompletionsRequest<'static>, - #[case] api_key: Option<&'static str>, - #[case] expected_key: &str, -) { - let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; - let secrets = RecordingSecrets::new([ - ("ANTHROPIC_API_KEY", "source-key"), - ("ANTHROPIC_API_BASE", upstream.uri().as_str()), - ]); - let response = chat_completions( - &support::resources(), - &http_config(), - &secrets, - ChatCompletionsRequest { - api_key, - api_base: None, - ..request - }, - ) - .await - .unwrap(); - let sent = only_request(&upstream).await; - assert_eq!(sent.header("x-api-key"), Some(expected_key)); - assert_eq!(sent.url.path(), "/v1/messages"); - assert_eq!( - response.choices[0].message.content.as_deref(), - Some("hello") - ); -} - -#[rstest] -#[case::accepted(false)] -#[case::declined(true)] -#[tokio::test] -async fn secret_failure_stops_before_sending_and_declines_skip_resolution( - request: ChatCompletionsRequest<'static>, - #[case] declined: bool, -) { - let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; - let base = upstream.uri(); - let secrets = RecordingSecrets::failing(); - let result = chat_completions( - &support::resources(), - &http_config(), - &secrets, - ChatCompletionsRequest { - api_base: Some(&base), - optional_params: if declined { - object(json!({"stream": true})) - } else { - request.optional_params.clone() - }, - ..request - }, - ) - .await; - if declined { - assert!(matches!(result, Err(Error::Unsupported(_)))); - assert!(secrets.requested().is_empty()); - } else { - assert!(matches!(result, Err(Error::Secret(_)))); - assert!(!secrets.requested().is_empty()); - } - assert!(received(&upstream).await.is_empty()); -} - -#[rstest] -#[case::bearer(true)] -#[case::signed(false)] -#[tokio::test] -async fn bedrock_chat_uses_the_injected_credential_source( - request: ChatCompletionsRequest<'static>, - #[case] bearer: bool, -) { - let upstream = upstream([json_response(json!({ - "output": {"message": {"content": [{"text": "hello"}]}}, - "usage": {"inputTokens": 1, "outputTokens": 1} - }))]) - .await; - let base = upstream.uri(); - let secrets = RecordingSecrets::new( - [ - ("AWS_ACCESS_KEY_ID", "injected-access-key"), - ("AWS_SECRET_ACCESS_KEY", "injected-secret-key"), - ("AWS_REGION_NAME", "eu-west-1"), - ] - .into_iter() - .chain(bearer.then_some(("AWS_BEARER_TOKEN_BEDROCK", "injected-bearer"))), - ); - let response = chat_completions( - &support::resources(), - &http_config(), - &secrets, - ChatCompletionsRequest { - model: "test-model", - custom_llm_provider: Some("bedrock"), - api_key: None, - api_base: Some(&base), - optional_params: Map::new(), - ..request - }, - ) - .await - .unwrap(); - assert_eq!( - response.choices[0].message.content.as_deref(), - Some("hello") - ); - let sent = only_request(&upstream).await; - let authorization = sent.header("authorization").unwrap(); - if bearer { - assert_eq!(authorization, "Bearer injected-bearer"); - } else { - assert!(authorization.contains("Credential=injected-access-key/")); - assert!(authorization.contains("/eu-west-1/bedrock/aws4_request")); - } - assert!(!sent.body_text().contains("injected-secret-key")); -} - -#[rstest] -#[tokio::test] -async fn openai_compatible_chat_resolves_its_injected_endpoint_and_key( - request: ChatCompletionsRequest<'static>, -) { - let upstream = upstream([json_response(json!({ - "id": "test-response", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}], - "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2} - }))]).await; - let secrets = RecordingSecrets::new([ - ("OPENAI_LIKE_API_BASE", upstream.uri().as_str()), - ("OPENAI_LIKE_API_KEY", "injected-key"), - ]); - let response = chat_completions( - &support::resources(), - &http_config(), - &secrets, - ChatCompletionsRequest { - model: "test-model", - custom_llm_provider: Some("openai_like"), - api_key: None, - api_base: None, - ..request - }, - ) - .await - .unwrap(); - assert_eq!( - response.choices[0].message.content.as_deref(), - Some("hello") - ); - let sent = only_request(&upstream).await; - assert_eq!(sent.url.path(), "/chat/completions"); - assert_eq!(sent.header("authorization"), Some("Bearer injected-key")); -} diff --git a/litellm-rust/crates/core/tests/messages/host.rs b/litellm-rust/crates/core/tests/messages/host.rs index ab695be90d9..bcea70c91fb 100644 --- a/litellm-rust/crates/core/tests/messages/host.rs +++ b/litellm-rust/crates/core/tests/messages/host.rs @@ -1,18 +1,15 @@ -use std::{convert::Infallible, sync::Mutex}; +use std::sync::Mutex; use litellm_core::messages::route::Messages; -use litellm_host::{ - event::{CallEvent, MachineEvent, RequestContext, WireRequest}, - host::Host, -}; -use litellm_llms::base_llm::messages::context::MessagesModelCapabilities; +use litellm_host::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; +use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities; use rstest::rstest; use super::*; type Rewrite = Box Result + Send + Sync>; -/// Projects like `LocalMessagesHost`, answers `before_send` through `rewrite`, and keeps +/// Projects like `LocalMessagesHost`, answers `before_provider_request` through `rewrite`, and keeps /// every event the driver emits. struct RecordingHost { call: LocalMessagesHost, @@ -50,19 +47,32 @@ impl RecordingHost { } } -impl Host for RecordingHost { - async fn project(&self) -> Result { - self.call.project().await +impl RecordingHost { + pub fn request(&self) -> Result { + self.call.request() } - - async fn custom_op(&self, op: Infallible) -> Result<(), Error> { - match op {} + pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, ()> { + litellm_host::in_process::Host { + services: &(), + hooks: self, + stream: &(), + observer: Some(self), + } } +} - async fn before_send( +impl litellm_host::lifecycle::CallObserver for RecordingHost { + fn observe(&self, event: litellm_host::event::CallEvent) { + self.events.lock().unwrap().push(event.clone()); + } +} +impl litellm_host::hooks::RouteHooks<::Error> + for RecordingHost +{ + async fn before_provider_request( &self, wire: WireRequest, - context: &RequestContext, + context: RequestContext, ) -> Result { self.optional_params .lock() @@ -70,15 +80,24 @@ impl Host for RecordingHost { .push(context.optional_params.clone()); (self.rewrite)(wire) } - - async fn emit(&self, event: &CallEvent) -> Result<(), Error> { - self.events.lock().unwrap().push(event.clone()); + async fn on_event( + &self, + event: litellm_host::event::MachineEvent, + ) -> Result<(), ::Error> { + litellm_host::lifecycle::CallObserver::observe( + self, + litellm_host::event::CallEvent::Machine(event), + ); Ok(()) } } async fn run_through(host: &RecordingHost) -> Result { - litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await + litellm_host::in_process::run_hosted( + machine(Arc::new(RecordingSecrets::empty()))(host.request()?), + host.runtime(), + ) + .await } fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall { @@ -129,8 +148,7 @@ async fn a_before_send_failure_never_sends(call: MessagesCall) { let error = run_through(&host) .await - .err() - .expect("the host failure fails the call"); + .expect_err("the host failure fails the call"); assert_eq!(error, Error::InvalidRequest("vetoed by the host".into())); assert!(received(&upstream).await.is_empty()); @@ -146,7 +164,7 @@ async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall) let output = run_through(&host).await.expect("messages call succeeds"); - assert!(matches!(output, MessagesOutput::Message(_))); + assert!(matches!(output, MessagesOutput::Complete(_))); let [emitted] = <[String; 1]>::try_from(host.raw_responses()) .unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len())); assert_eq!(serde_json::from_str::(&emitted).unwrap(), raw); @@ -183,9 +201,9 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages let host = RecordingHost::passthrough(authenticated( MessagesCall { shaping: MessagesShaping { - capabilities: MessagesModelCapabilities { + capabilities: AnthropicModelCapabilities { supports_sampling_params: false, - ..MessagesModelCapabilities::default() + ..AnthropicModelCapabilities::default() }, drop_params: true, ..MessagesShaping::default() @@ -198,6 +216,6 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages run_through(&host).await.expect("messages call succeeds"); let [optional_params] = <[Value; 1]>::try_from(host.optional_params.into_inner().unwrap()) - .unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len())); + .unwrap_or_else(|seen| panic!("before_provider_request runs once, saw {}", seen.len())); assert_eq!(optional_params, json!({"max_tokens": 16})); } diff --git a/litellm-rust/crates/core/tests/messages/main.rs b/litellm-rust/crates/core/tests/messages/main.rs index 534af6d7d06..7de463fb8fb 100644 --- a/litellm-rust/crates/core/tests/messages/main.rs +++ b/litellm-rust/crates/core/tests/messages/main.rs @@ -1,8 +1,11 @@ -use std::{sync::Arc, time::Duration}; +use std::{ + sync::{Arc, Mutex}, + time::Duration, +}; use litellm_core::messages::{ Error, MessagesCall, MessagesShaping, - route::{LocalMessagesHost, MessagesMachine, MessagesOutput, messages_machine}, + route::{Messages, MessagesMachine, MessagesOutput}, }; use litellm_http::{HttpSettings, Resolution}; use litellm_secrets::source::SecretSource; @@ -95,16 +98,16 @@ fn headers<'a>(pairs: impl IntoIterator) -> Option) -> MessagesMachine { - messages_machine(&support::resources(), &http_config(), secrets) - .expect("default HTTP settings build a client") +fn machine(secrets: Arc) -> impl FnOnce(MessagesCall) -> MessagesMachine { + move |request| messages_route(secrets).machine(request) } async fn run_with( secrets: Arc, call: MessagesCall, ) -> Result { - litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await + let host = LocalMessagesHost::new(call); + litellm_host::in_process::run_hosted(machine(secrets)(host.request()?), host.runtime()).await } /// Runs the route with a secret source that knows nothing, so no environment leaks in. @@ -114,7 +117,67 @@ async fn run(call: MessagesCall) -> Result { async fn run_message(call: MessagesCall) -> AnthropicMessagesResponse { match run(call).await.expect("messages call succeeds") { - MessagesOutput::Message(message) => *message, - MessagesOutput::Streamed => panic!("a non-streaming call returned a stream"), + MessagesOutput::Complete(message) => *message, + MessagesOutput::StreamEnded | MessagesOutput::Detached => { + panic!("a non-streaming call returned a stream") + } + } +} + +struct LocalMessagesHost { + call: Mutex>, +} + +impl LocalMessagesHost { + fn new(call: MessagesCall) -> Self { + Self { + call: Mutex::new(Some(call)), + } + } +} + +impl LocalMessagesHost { + pub fn request(&self) -> Result { + self.call + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .ok_or_else(|| Error::InvalidRequest("messages request was already projected".into())) + } + pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, ()> { + litellm_host::in_process::Host { + services: &(), + hooks: self, + stream: &(), + observer: Some(self), + } + } +} + +impl litellm_host::lifecycle::CallObserver for LocalMessagesHost { + fn observe(&self, _: litellm_host::event::CallEvent) {} +} +impl litellm_host::hooks::RouteHooks<::Error> + for LocalMessagesHost +{ + async fn before_provider_request( + &self, + wire: litellm_host::event::WireRequest, + _: litellm_host::event::RequestContext, + ) -> Result< + litellm_host::event::WireRequest, + ::Error, + > { + Ok(wire) + } + async fn on_event( + &self, + event: litellm_host::event::MachineEvent, + ) -> Result<(), ::Error> { + litellm_host::lifecycle::CallObserver::observe( + self, + litellm_host::event::CallEvent::Machine(event), + ); + Ok(()) } } diff --git a/litellm-rust/crates/core/tests/messages/request.rs b/litellm-rust/crates/core/tests/messages/request.rs index ffab9277975..b76895b9f1a 100644 --- a/litellm-rust/crates/core/tests/messages/request.rs +++ b/litellm-rust/crates/core/tests/messages/request.rs @@ -87,8 +87,7 @@ async fn a_call_without_credentials_fails_before_sending( ..call }) .await - .err() - .expect("a call without credentials fails"); + .expect_err("a call without credentials fails"); assert!( matches!( @@ -158,8 +157,7 @@ async fn unsupported_providers_are_rejected_before_sending( ..with_model(call, model) }) .await - .err() - .expect("unsupported provider errors"); + .expect_err("unsupported provider errors"); assert_eq!(error, Error::InvalidProvider(reported.into())); } @@ -405,8 +403,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i let error = run(shaped(false)) .await - .err() - .expect("an unsupported param is rejected without drop_params"); + .expect_err("an unsupported param is rejected without drop_params"); assert!( matches!(&error, Error::InvalidRequest(message) if message.to_string().contains(rejected_as)), "{error:?}" @@ -601,8 +598,7 @@ async fn an_invalid_request_fails_before_sending(call: MessagesCall, #[case] fie fields, )) .await - .err() - .expect("the request is rejected"); + .expect_err("the request is rejected"); assert!(error.is_request(), "{error:?}"); assert!(received(&upstream).await.is_empty()); @@ -705,7 +701,7 @@ async fn provider_validation_runs_before_caller_parameter_removal( json!({"metadata": {"user_id": 7}}), )) .await; - let error = result.err().expect("metadata is validated before removal"); + let error = result.expect_err("metadata is validated before removal"); assert!( error .to_string() diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 4fac78ca1af..47505aa86a4 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -1,12 +1,65 @@ use litellm_core::{ Phase, - messages::{MessagesResponse, messages, messages_body}, + messages::{MessagesResponse, messages_body}, }; use litellm_http::transport::Error as TransportError; use rstest::rstest; use super::*; +#[rstest] +#[case::without_hooks(false)] +#[case::with_hooks(true)] +#[tokio::test] +async fn calls_defer_execution_until_polled(call: MessagesCall, #[case] with_hooks: bool) { + use futures_util::future::BoxFuture; + + use litellm_host::event::CallEvent; + + let upstream = upstream([message_response()]).await; + let secrets = Arc::new(RecordingSecrets::new([("ANTHROPIC_API_KEY", "test-key")])); + let route = messages_route(secrets.clone()); + let host = RecordingCall::::new(MessagesCall { + api_base: Some(upstream.uri()), + ..call + }); + let request = host.request().unwrap(); + let future: BoxFuture<'_, Result> = if with_hooks { + Box::pin(route.execute(request, &host)) + } else { + Box::pin(route.execute(request, &())) + }; + + assert!(secrets.requested().is_empty()); + assert!(host.events.0.lock().unwrap().is_empty()); + assert!(received(&upstream).await.is_empty()); + + let MessagesResponse::Complete(response) = future.await.unwrap() else { + panic!("expected a completed message"); + }; + assert_eq!( + response.content, + message_body()["content"].as_array().unwrap().as_slice() + ); + assert!(secrets.requested().contains(&"ANTHROPIC_API_KEY".into())); + let sent = only_request(&upstream).await; + assert_eq!(sent.header("x-api-key"), Some("test-key")); + assert_eq!(sent.header("x-hook"), with_hooks.then_some("called")); + let events = host.events.0.lock().unwrap(); + if with_hooks { + assert!(matches!( + &events[..], + [ + CallEvent::Started { .. }, + CallEvent::Machine(_), + CallEvent::Succeeded { .. } + ] + )); + } else { + assert!(events.is_empty()); + } +} + #[rstest] #[case::anthropic("anthropic")] #[case::azure_ai("azure_ai")] @@ -75,8 +128,7 @@ async fn a_json_error_envelope_is_kept_verbatim(call: MessagesCall) { ..call }) .await - .err() - .expect("upstream error propagates"); + .expect_err("upstream error propagates"); let Error::Transport(TransportError::Http { status, body }) = error else { panic!("{error:?}"); @@ -97,8 +149,7 @@ async fn a_long_error_body_is_truncated_at_the_documented_cap(call: MessagesCall ..call }) .await - .err() - .expect("upstream error propagates"); + .expect_err("upstream error propagates"); assert_eq!( error, @@ -126,8 +177,7 @@ async fn an_upstream_error_keeps_its_status_and_body(call: MessagesCall, #[case] ..call }) .await - .err() - .expect("upstream error propagates"); + .expect_err("upstream error propagates"); assert_eq!( error, @@ -154,8 +204,7 @@ async fn an_unreadable_success_body_is_an_invalid_response( ..call }) .await - .err() - .expect("an unreadable body fails"); + .expect_err("an unreadable body fails"); assert_eq!(error.phase(), Phase::AfterSend, "{error:?}"); } @@ -172,8 +221,7 @@ async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) { ..call }) .await - .err() - .expect("the call times out"); + .expect_err("the call times out"); assert!(matches!(error, Error::Transport(_)), "{error:?}"); } @@ -188,20 +236,24 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes ..HttpSettings::default() }; - let response = messages( - &support::resources(), - &Resolution::from(&settings).config, - &RecordingSecrets::empty(), + let resources = support::resources(); + let response = litellm_core::messages::MessagesRoute::new( + provider_http(&resources, &Resolution::from(&settings).config), + resources.auth, + no_secrets(), + ) + .execute( MessagesCall { api_key: Some("sk-ant".into()), api_base: Some(base), ..call }, + &(), ) .await .expect("messages request succeeds"); - let MessagesResponse::Message(message) = response else { + let MessagesResponse::Complete(message) = response else { panic!("a non-streaming request returns a message"); }; assert_eq!(message.id, "msg_1"); diff --git a/litellm-rust/crates/core/tests/messages/secrets.rs b/litellm-rust/crates/core/tests/messages/secrets.rs index 71b6a2bd93c..4b3d0ba705a 100644 --- a/litellm-rust/crates/core/tests/messages/secrets.rs +++ b/litellm-rust/crates/core/tests/messages/secrets.rs @@ -43,7 +43,7 @@ async fn the_credential_and_base_come_from_the_secret_source( .await .expect("messages call succeeds"); - assert!(matches!(output, MessagesOutput::Message(_))); + assert!(matches!(output, MessagesOutput::Complete(_))); let request = only_request(&upstream).await; assert_eq!(request.url.path(), path); assert_eq!(request.header("x-api-key"), Some("sk-from-manager")); @@ -90,8 +90,7 @@ async fn a_secret_manager_failure_fails_the_call_before_sending(call: MessagesCa }, ) .await - .err() - .expect("a secret manager failure fails the call"); + .expect_err("a secret manager failure fails the call"); assert!( matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)), @@ -193,8 +192,7 @@ async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall) }, ) .await - .err() - .expect("azure needs a base"); + .expect_err("azure needs a base"); assert_eq!( error, diff --git a/litellm-rust/crates/core/tests/messages/stream.rs b/litellm-rust/crates/core/tests/messages/stream.rs index f0e55eca8dd..fa4594525c9 100644 --- a/litellm-rust/crates/core/tests/messages/stream.rs +++ b/litellm-rust/crates/core/tests/messages/stream.rs @@ -1,16 +1,12 @@ -use std::{ - convert::Infallible, - sync::{Mutex, mpsc}, -}; +use std::sync::Mutex; use bytes::Bytes; use futures_util::{StreamExt, TryStreamExt}; use litellm_core::messages::{ - MessagesResponse, messages, + MessagesResponse, route::{Messages, MessagesStreamHead}, }; -use litellm_host::host::{Demand, Host}; -use litellm_tracing::{Logger, Metadata, Record, Sink}; +use litellm_host::protocol::Demand; use rstest::rstest; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, @@ -32,20 +28,6 @@ enum Seen { Deliver(Bytes), } -struct TraceSink(mpsc::Sender<(String, Value)>); - -impl Sink for TraceSink { - fn enabled(&self, metadata: &Metadata<'_>) -> bool { - metadata.target().starts_with("litellm_core::messages") - } - - fn emit(&self, record: &Record) { - self.0 - .send((record.message.clone(), Value::Object(record.fields.clone()))) - .unwrap(); - } -} - /// Projects like `LocalMessagesHost`, records every stream op in the order the route /// performs it, and detaches after `detach_after` ops. struct RecordingStreamHost { @@ -73,23 +55,55 @@ impl RecordingStreamHost { } } -impl Host for RecordingStreamHost { - async fn project(&self) -> Result { - self.call.project().await +impl RecordingStreamHost { + pub fn request(&self) -> Result { + self.call.request() } - - async fn custom_op(&self, op: Infallible) -> Result<(), Error> { - match op {} + pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, Self> { + litellm_host::in_process::Host { + services: &(), + hooks: self, + stream: self, + observer: Some(self), + } } +} - async fn open(&self, head: MessagesStreamHead) -> Result { +impl litellm_host::in_process::StreamConsumer for RecordingStreamHost { + async fn open_stream(&self, head: MessagesStreamHead) -> Result { Ok(self.record(Seen::Open(head.headers))) } - - async fn deliver(&self, chunk: Bytes) -> Result { + async fn send_chunk(&self, chunk: Bytes) -> Result { Ok(self.record(Seen::Deliver(chunk))) } } +impl litellm_host::lifecycle::CallObserver for RecordingStreamHost { + fn observe(&self, _: litellm_host::event::CallEvent) {} +} +impl litellm_host::hooks::RouteHooks<::Error> + for RecordingStreamHost +{ + async fn before_provider_request( + &self, + wire: litellm_host::event::WireRequest, + _: litellm_host::event::RequestContext, + ) -> Result< + litellm_host::event::WireRequest, + ::Error, + > { + Ok(wire) + } + async fn on_event( + &self, + event: litellm_host::event::MachineEvent, + ) -> Result<(), ::Error> { + litellm_host::lifecycle::CallObserver::observe( + self, + litellm_host::event::CallEvent::Machine(event), + ); + Ok(()) + } +} fn streaming(call: MessagesCall, api_base: String) -> MessagesCall { MessagesCall { @@ -107,7 +121,11 @@ fn sse_response() -> ResponseTemplate { } async fn stream_through(host: &RecordingStreamHost) -> Result { - litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await + litellm_host::in_process::run_hosted( + machine(Arc::new(RecordingSecrets::empty()))(host.request()?), + host.runtime(), + ) + .await } #[rstest] @@ -118,7 +136,7 @@ async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: Me let outcome = stream_through(&host).await.expect("streamed call succeeds"); - assert!(matches!(outcome, MessagesOutput::Streamed)); + assert!(matches!(outcome, MessagesOutput::StreamEnded)); let seen = host.seen.into_inner().unwrap(); let [Seen::Open(headers), chunks @ ..] = seen.as_slice() else { panic!("the stream opens before any chunk is delivered"); @@ -143,37 +161,6 @@ async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: Me assert_eq!(delivered, SSE_BODY.as_bytes()); } -#[rstest] -#[tokio::test] -async fn debug_trace_keeps_provider_input_and_every_stream_chunk(call: MessagesCall) { - let upstream = upstream([sse_response()]).await; - let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX); - let (sender, receiver) = mpsc::channel(); - - Logger::new(TraceSink(sender)) - .instrument(stream_through(&host)) - .await - .unwrap(); - - let records: Vec<(String, Value)> = receiver.try_iter().collect(); - let request = records - .iter() - .find(|(message, _)| message == "provider request") - .unwrap(); - let body: Value = serde_json::from_str(request.1["body"].as_str().unwrap()).unwrap(); - assert_eq!(body["messages"][0]["content"], "hi"); - assert_eq!(request.1["stream"], true); - let chunks: String = records - .iter() - .filter(|(message, fields)| { - message == "stream chunk" && fields["stage"] == "provider_response" - }) - .map(|(_, fields)| fields["chunk"].as_str().unwrap()) - .collect(); - assert_eq!(chunks, SSE_BODY); - assert!(!format!("{records:?}").contains("sk-ant")); -} - #[rstest] #[case::at_open(1)] #[case::after_the_first_chunk(2)] @@ -186,7 +173,7 @@ async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] det .await .expect("a detached stream still completes"); - assert!(matches!(outcome, MessagesOutput::Streamed)); + assert!(matches!(outcome, MessagesOutput::Detached)); assert_eq!(host.seen.into_inner().unwrap().len(), detach_after); } @@ -196,6 +183,7 @@ async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] det status_response(429, json!({"type": "error", "error": {"type": "rate_limit_error", "message": "slow down"}})), r#"{"type":"error","error":{"type":"rate_limit_error","message":"slow down"}}"# )] +#[rstest::rstest] #[tokio::test] async fn an_upstream_error_fails_the_call_without_opening_the_stream( call: MessagesCall, @@ -207,8 +195,7 @@ async fn an_upstream_error_fails_the_call_without_opening_the_stream( let error = stream_through(&host) .await - .err() - .expect("upstream error propagates"); + .expect_err("upstream error propagates"); assert_eq!( error, @@ -280,8 +267,7 @@ async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) { let error = tokio::time::timeout(Duration::from_secs(5), stream_through(&host)) .await .expect("the stalled stream gives up within the timeout") - .err() - .expect("a stalled body fails the call"); + .expect_err("a stalled body fails the call"); assert!(matches!(error, Error::Transport(_)), "{error:?}"); let seen = host.seen.into_inner().unwrap(); @@ -305,23 +291,22 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte( #[case] provider: &str, ) { let upstream = upstream([sse_response()]).await; - let response = messages( - &support::resources(), - &http_config(), - &RecordingSecrets::empty(), - MessagesCall { - custom_llm_provider: Some(provider.into()), - ..streaming(call, upstream.uri()) - }, - ) - .await - .unwrap(); + let response = messages_route(no_secrets()) + .execute( + MessagesCall { + custom_llm_provider: Some(provider.into()), + ..streaming(call, upstream.uri()) + }, + &(), + ) + .await + .unwrap(); - let MessagesResponse::Stream { headers, chunks } = response else { + let MessagesResponse::Stream { head, chunks } = response else { panic!("a streaming request returns a stream"); }; for (name, value) in UPSTREAM_HEADERS { - assert!(headers.contains(&(name.into(), value.into()))); + assert!(head.headers.contains(&(name.into(), value.into()))); } let delivered = chunks.try_collect::>().await.unwrap().concat(); assert_eq!(delivered, SSE_BODY.as_bytes()); @@ -332,15 +317,11 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte( #[tokio::test] async fn the_sdk_returns_http_errors_before_opening_a_stream(call: MessagesCall) { let upstream = upstream([ResponseTemplate::new(429).set_body_string("slow down")]).await; - let error = messages( - &support::resources(), - &http_config(), - &RecordingSecrets::empty(), - streaming(call, upstream.uri()), - ) - .await - .err() - .expect("upstream failure is returned by messages()"); + let error = messages_route(no_secrets()) + .execute(streaming(call, upstream.uri()), &()) + .await + .err() + .expect("upstream failure is returned by messages()"); assert_eq!( error, @@ -362,14 +343,12 @@ async fn dropping_the_sdk_stream_closes_the_unfinished_upstream( let (base, connection) = stalling_upstream().await; let response = tokio::time::timeout( Duration::from_secs(5), - messages( - &support::resources(), - &http_config(), - &RecordingSecrets::empty(), + messages_route(no_secrets()).execute( MessagesCall { timeout: Some(Duration::from_secs(30)), ..streaming(call, base) }, + &(), ), ) .await @@ -399,17 +378,16 @@ async fn dropping_the_sdk_stream_closes_the_unfinished_upstream( #[tokio::test] async fn the_sdk_yields_a_body_error_once_after_delivered_chunks(call: MessagesCall) { let (base, connection) = stalling_upstream().await; - let response = messages( - &support::resources(), - &http_config(), - &RecordingSecrets::empty(), - MessagesCall { - timeout: Some(Duration::from_millis(300)), - ..streaming(call, base) - }, - ) - .await - .unwrap(); + let response = messages_route(no_secrets()) + .execute( + MessagesCall { + timeout: Some(Duration::from_millis(300)), + ..streaming(call, base) + }, + &(), + ) + .await + .unwrap(); let MessagesResponse::Stream { mut chunks, .. } = response else { panic!("a streaming request returns a stream"); @@ -445,7 +423,7 @@ async fn a_host_on_anthropic_sse_is_relayed_byte_for_byte(call: MessagesCall) { let outcome = stream_through(&host).await.expect("azure streams"); - assert!(matches!(outcome, MessagesOutput::Streamed)); + assert!(matches!(outcome, MessagesOutput::StreamEnded)); let seen = host.seen.into_inner().unwrap(); let delivered: Vec = seen .iter() 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 1921176d158..e7b30e5a7e4 100644 --- a/litellm-rust/crates/core/tests/ocr/azure_document_intelligence.rs +++ b/litellm-rust/crates/core/tests/ocr/azure_document_intelligence.rs @@ -183,16 +183,16 @@ async fn client_settings_choose_the_api_version_and_the_inch_to_pixel_dpi() { "analyzeResult": {"pages": [{"pageNumber": 1, "width": 8.5, "height": 11, "unit": "inch"}]} }))]) .await; - let client = ocr_client().with_settings(OcrSettings { + let route = ocr_route_with(OcrSettings { document_intelligence_api_version: "2099-01-01".into(), document_intelligence_dpi: 72, ..OcrSettings::default() }); - let result = - litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({}))) - .await - .unwrap(); + let result = route + .execute(read_request(&upstream.uri(), json!({})), &()) + .await + .unwrap(); assert_eq!( only_request(&upstream) @@ -347,14 +347,14 @@ async fn the_polling_deadline_bounds_the_retry_delay() { ], ) .await; - let client = ocr_client().with_settings(OcrSettings { + let route = ocr_route_with(OcrSettings { poll_timeout: Duration::from_millis(100), ..OcrSettings::default() }); let error = tokio::time::timeout( Duration::from_secs(1), - litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({}))), + route.execute(read_request(&upstream.uri(), json!({})), &()), ) .await .expect("the deadline cuts the retry delay short") diff --git a/litellm-rust/crates/core/tests/ocr/documents.rs b/litellm-rust/crates/core/tests/ocr/documents.rs index e29ff3e9ee9..f0bbfc5e0be 100644 --- a/litellm-rust/crates/core/tests/ocr/documents.rs +++ b/litellm-rust/crates/core/tests/ocr/documents.rs @@ -45,7 +45,7 @@ impl Route { } } -/// What the host does to the wire request in `before_send`. +/// What the host does to the wire request in `before_provider_request`. #[derive(Clone, Copy, Debug)] enum Guardrail { Detached, @@ -53,7 +53,7 @@ enum Guardrail { } impl Guardrail { - fn before_send(self, wire: WireRequest) -> WireRequest { + fn before_provider_request(self, wire: WireRequest) -> WireRequest { let Value::Object(fields) = wire.body else { return wire; }; @@ -96,8 +96,8 @@ async fn provider_document(route: Route, guardrail: Guardrail) -> Value { json!({"type": document_type, document_type: format!("{}/scan.png", documents.uri())}), route.options(), ); - let host = - LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(guardrail.before_send(wire))); + let host = LocalOcrHost::new(request) + .with_before_send(move |wire, _| Ok(guardrail.before_provider_request(wire))); perform_with(host).await.unwrap(); @@ -139,6 +139,7 @@ async fn a_document_replaced_by_the_host_reaches_the_provider( ); } +#[rstest::rstest] #[tokio::test] async fn an_empty_byte_document_fails_before_sending() { let upstream = upstream([pages_response()]).await; @@ -156,6 +157,7 @@ async fn an_empty_byte_document_fails_before_sending() { assert!(received(&upstream).await.is_empty()); } +#[rstest::rstest] #[tokio::test] async fn a_missing_path_document_fails_before_sending() { let upstream = upstream([pages_response()]).await; @@ -180,3 +182,51 @@ async fn a_missing_path_document_fails_before_sending() { ); assert!(received(&upstream).await.is_empty()); } + +#[rstest] +#[case::blocked(false)] +#[case::allowed(true)] +#[tokio::test] +async fn configured_client_preserves_document_url_policy(#[case] allowed: bool) { + let documents = document_server().await; + let upstream = upstream([pages_response()]).await; + let document_url = format!("{}/scan.png", documents.uri()); + let authority = documents.address().to_string(); + let route = build_ocr_route( + &resources(), + &http_config(), + litellm_http::media::UrlPolicy { + validate: true, + allowed_hosts: allowed.then_some(authority).into_iter().collect(), + }, + Default::default(), + no_secrets(), + ); + let host = LocalOcrHost::new(ocr_request_with_document( + "azure_ai/model", + &upstream.uri(), + json!({"type": "document_url", "document_url": document_url}), + json!({}), + )); + let result = litellm_host::in_process::run_hosted( + route.machine(host.request().unwrap()), + host.runtime(), + ) + .await; + + if !allowed { + assert!(matches!(result, Err(Error::BlockedDocumentUrl))); + assert!(received(&documents).await.is_empty()); + assert!(received(&upstream).await.is_empty()); + return; + } + result.unwrap(); + assert_eq!(only_request(&documents).await.url.path(), "/scan.png"); + assert_eq!( + only_request(&upstream).await.json()["document"]["document_url"], + format!( + "data:image/png;base64,{}", + base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT) + ) + ); +} diff --git a/litellm-rust/crates/core/tests/ocr/lifecycle.rs b/litellm-rust/crates/core/tests/ocr/lifecycle.rs index 65e64cce79b..7febef105c6 100644 --- a/litellm-rust/crates/core/tests/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/tests/ocr/lifecycle.rs @@ -1,13 +1,10 @@ use std::sync::{Arc, Mutex}; use litellm_core::ocr::{ - route::{Ocr, OcrOp, OcrProjection, ocr_machine}, + route::{Ocr, OcrCall, OcrOp}, types::OcrDocumentInput, }; -use litellm_host::{ - event::{CallEvent, MachineEvent, RequestContext, WireRequest}, - host::Host, -}; +use litellm_host::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; use rstest::rstest; use super::*; @@ -18,6 +15,7 @@ pub(crate) fn event_name(event: &CallEvent) -> &'static str { CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response", CallEvent::Succeeded { .. } => "success", CallEvent::Failed { .. } => "failure", + CallEvent::Cancelled { .. } => "cancelled", } } @@ -29,7 +27,10 @@ fn recording_host( let before_send_events = events.clone(); LocalOcrHost::new(request) .with_before_send(move |wire, _| { - before_send_events.lock().unwrap().push("before_send"); + before_send_events + .lock() + .unwrap() + .push("before_provider_request"); match block { true => Err(Error::InvalidRequest("blocked".into())), false => Ok(wire), @@ -38,6 +39,7 @@ fn recording_host( .with_observer(move |event| events.lock().unwrap().push(event_name(event))) } +#[rstest::rstest] #[tokio::test] async fn hooks_run_in_order_and_one_success_is_emitted() { let upstream = upstream([pages_response()]).await; @@ -53,11 +55,12 @@ async fn hooks_run_in_order_and_one_success_is_emitted() { assert_eq!( *events.lock().unwrap(), - ["started", "before_send", "response", "success"] + ["started", "before_provider_request", "response", "success"] ); assert_eq!(received(&upstream).await.len(), 1); } +#[rstest::rstest] #[tokio::test] async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() { let upstream = upstream([pages_response()]).await; @@ -77,11 +80,12 @@ async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() { ); assert_eq!( *events.lock().unwrap(), - ["started", "before_send", "failure"] + ["started", "before_provider_request", "failure"] ); assert!(received(&upstream).await.is_empty()); } +#[rstest::rstest] #[tokio::test] async fn an_upstream_failure_emits_one_terminal_failure() { let upstream = upstream([status_response(500, json!({"error": "failed"}))]).await; @@ -97,11 +101,12 @@ async fn an_upstream_failure_emits_one_terminal_failure() { assert!(result.is_err()); assert_eq!( *events.lock().unwrap(), - ["started", "before_send", "failure"] + ["started", "before_provider_request", "failure"] ); assert_eq!(received(&upstream).await.len(), 1); } +#[rstest::rstest] #[tokio::test] async fn an_invalid_provider_response_is_observed_before_normalization_fails() { let upstream = upstream([json_response(json!({"pages": "invalid"}))]).await; @@ -120,6 +125,7 @@ async fn an_invalid_provider_response_is_observed_before_normalization_fails() { assert_eq!(*observed.lock().unwrap(), [r#"{"pages":"invalid"}"#]); } +#[rstest::rstest] #[tokio::test] async fn headers_returned_by_before_send_are_sent() { let upstream = upstream([pages_response()]).await; @@ -147,9 +153,10 @@ async fn before_send_context(request: LiteLLMOcrRequest) -> (WireRequest, Reques }); perform_with(host).await.unwrap(); let context = observed.lock().unwrap().take(); - context.expect("before_send ran") + context.expect("before_provider_request ran") } +#[rstest::rstest] #[tokio::test] async fn before_send_sees_the_route_its_params_and_the_body() { let upstream = upstream([pages_response()]).await; @@ -187,22 +194,31 @@ async fn before_send_names_the_secret_params(#[case] options: Value, #[case] sec assert_eq!(context.secret_fields, secrets); } -/// Hands the route a caller-owned Azure token and rewrites the bearer in `before_send`. +/// Hands the route a caller-owned Azure token and rewrites the bearer in `before_provider_request`. struct CallerTokenHost { request: Mutex>, trace: Mutex>, } -impl Host for CallerTokenHost { - async fn project(&self) -> Result { +impl CallerTokenHost { + pub fn request(&self) -> Result { self.trace.lock().unwrap().push("project".into()); - Ok(OcrProjection { + Ok(OcrCall { request: self.request.lock().unwrap().take().unwrap(), caller_token: true, }) } - - async fn custom_op(&self, op: OcrOp) -> Result<(), Error> { + pub fn runtime(&self) -> litellm_host::in_process::Host<'_, Self, Self, ()> { + litellm_host::in_process::Host { + services: self, + hooks: self, + stream: &(), + observer: Some(self), + } + } +} +impl litellm_host::services::HostCallHandler for CallerTokenHost { + async fn handle_host_call(&self, op: OcrOp) -> Result<(), Error> { match op { OcrOp::AcquireAzureAdToken(reply) => { self.trace.lock().unwrap().push("token".into()); @@ -213,11 +229,18 @@ impl Host for CallerTokenHost { } } } +} - async fn before_send( +impl litellm_host::lifecycle::CallObserver for CallerTokenHost { + fn observe(&self, _: litellm_host::event::CallEvent) {} +} +impl litellm_host::hooks::RouteHooks<::Error> + for CallerTokenHost +{ + async fn before_provider_request( &self, wire: WireRequest, - _: &RequestContext, + _: RequestContext, ) -> Result { let is_authorization = |name: &str| name.eq_ignore_ascii_case("authorization"); let authorization = wire @@ -229,7 +252,7 @@ impl Host for CallerTokenHost { self.trace .lock() .unwrap() - .push(format!("before_send:{authorization}")); + .push(format!("before_provider_request:{authorization}")); let headers = wire .headers .into_iter() @@ -240,8 +263,19 @@ impl Host for CallerTokenHost { .collect(); Ok(WireRequest { headers, ..wire }) } + async fn on_event( + &self, + event: litellm_host::event::MachineEvent, + ) -> Result<(), ::Error> { + litellm_host::lifecycle::CallObserver::observe( + self, + litellm_host::event::CallEvent::Machine(event), + ); + Ok(()) + } } +#[rstest::rstest] #[tokio::test] async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_replace_it() { let upstream = upstream([pages_response()]).await; @@ -254,16 +288,85 @@ async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_ trace: Mutex::new(Vec::new()), }; - litellm_host::run::run(ocr_machine(ocr_client()), &host) - .await - .unwrap(); + litellm_host::in_process::run_hosted( + ocr_route().machine(host.request().unwrap()), + host.runtime(), + ) + .await + .unwrap(); assert_eq!( *host.trace.lock().unwrap(), - ["project", "token", "before_send:Bearer caller-token"] + [ + "project", + "token", + "before_provider_request:Bearer caller-token" + ] ); assert_eq!( only_request(&upstream).await.header_values("authorization"), ["Bearer edited"] ); } + +#[rstest] +#[tokio::test] +async fn direct_execution_uses_hooks_without_a_machine() { + use litellm_host::{hooks::RouteHooks, lifecycle::CallObserver}; + + struct Hooks(Arc); + + impl RouteHooks for Hooks { + fn observer(&self) -> Option> { + Some(self.0.clone()) + } + + async fn before_provider_request( + &self, + wire: WireRequest, + _: RequestContext, + ) -> Result { + Ok(WireRequest { + headers: wire + .headers + .into_iter() + .chain([("x-direct-hook".into(), "called".into())]) + .collect(), + ..wire + }) + } + + async fn on_event(&self, event: MachineEvent) -> Result<(), Error> { + self.0.observe(CallEvent::Machine(event)); + Ok(()) + } + } + + let upstream = upstream([json_response( + json!({"pages":[{"index":0,"markdown":"direct"}]}), + )]) + .await; + let events = Arc::new(super::support::CallEvents::default()); + let route = ocr_route(); + let hooks = Hooks(events.clone()); + let builder = route.execute( + ocr_request("mistral/model", &upstream.uri(), json!({})), + &hooks, + ); + assert!(events.0.lock().unwrap().is_empty()); + assert!(received(&upstream).await.is_empty()); + let result = builder.await.unwrap(); + assert_eq!(result.pages[0].markdown, "direct"); + assert_eq!( + only_request(&upstream).await.header("x-direct-hook"), + Some("called") + ); + assert!(matches!( + &events.0.lock().unwrap()[..], + [ + CallEvent::Started { .. }, + CallEvent::Machine(_), + CallEvent::Succeeded { .. } + ] + )); +} diff --git a/litellm-rust/crates/core/tests/ocr/machine.rs b/litellm-rust/crates/core/tests/ocr/machine.rs index 073ce67e74b..f4eaa735fc1 100644 --- a/litellm-rust/crates/core/tests/ocr/machine.rs +++ b/litellm-rust/crates/core/tests/ocr/machine.rs @@ -1,3 +1,4 @@ +use litellm_host::protocol::HookRequest; use std::{ sync::{ Arc, @@ -7,13 +8,15 @@ use std::{ }; use litellm_core::ocr::{ - route::{OcrMachine, OcrOp, OcrProjection}, + route::{OcrCall, OcrMachine, OcrOp}, types::OcrDocumentInput, }; use litellm_host::{ event::{CallEvent, WireRequest}, - host::{Host, HostOp}, + hooks::RouteHooks, machine::{HostFailure, Machine, MachineStep}, + protocol::Suspension, + services::HostCallHandler, }; use litellm_llms::base_llm::ocr::transformation::OcrTransportConfig; use rstest::rstest; @@ -21,7 +24,7 @@ use tokio::{io::AsyncReadExt, net::TcpListener, sync::Notify}; use super::{lifecycle::event_name, *}; -/// Drives the machine by hand, answering every op through `host` except `before_send`, +/// Drives the machine by hand, answering every op through `host` except `before_provider_request`, /// which `intercept` answers so a test can fail or cancel exactly there. async fn drive_until( host: &LocalOcrHost, @@ -31,43 +34,39 @@ async fn drive_until( Vec<&'static str>, OcrMachine, ) { - let mut machine = ocr_machine(ocr_client()); + let mut machine = ocr_route().machine(host.request().unwrap()); let mut ops = Vec::new(); let outcome = loop { let op = match machine.resume().await { - Ok(MachineStep::Host(op)) => op, - Ok(MachineStep::Complete(response)) => break Ok(response), + Ok(MachineStep::Suspended(op)) => op, + Ok(MachineStep::Complete(response)) => break Ok(completed(response)), Err(error) => break Err(error), }; let answer = match op { - HostOp::Project(reply) => { - ops.push("Project"); - host.project() - .await - .map(|projection| reply.send(projection)) - .map_err(HostFailure::Error) - } - HostOp::Custom(op) => { + Suspension::Stream(stream) => match stream { + litellm_host::protocol::StreamDelivery::Open(head, _) => match head {}, + litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {}, + }, + Suspension::HostCall(op) => { ops.push(match op { OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken", }); - host.custom_op(op).await.map_err(HostFailure::Error) + host.handle_host_call(op).await.map_err(HostFailure::Error) } - HostOp::BeforeSend { wire, reply, .. } => { + Suspension::Hook(HookRequest::BeforeProviderRequest { wire, reply, .. }) => { ops.push("BeforeSend"); intercept(*wire).map(|wire| reply.send(wire)) } - HostOp::Emit(event, reply) => { - let event = CallEvent::Machine(event); - ops.push(event_name(&event)); - host.emit(&event) + Suspension::Hook(HookRequest::Event(event, reply)) => { + ops.push(event_name(&CallEvent::Machine(event.clone()))); + host.on_event(event) .await .map(|()| reply.send(())) .map_err(HostFailure::Error) } }; if let Err(failure) = answer { - break machine.interrupt(failure).await; + break machine.interrupt(failure).await.map(completed); } }; (outcome, ops, machine) @@ -81,10 +80,13 @@ async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, sto _ = stop.notified() => break, step = machine.resume() => { match step.unwrap() { - MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()), - MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(), - MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire), - MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()), + MachineStep::Suspended(Suspension::HostCall(op)) => host.handle_host_call(op).await.unwrap(), + MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest { wire, reply, .. })) => reply.send(*wire), + MachineStep::Suspended(Suspension::Hook(HookRequest::Event(_, reply))) => reply.send(()), + MachineStep::Suspended(Suspension::Stream(stream)) => match stream { + litellm_host::protocol::StreamDelivery::Open(head, _) => match head {}, + litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {}, + }, MachineStep::Complete(_) => panic!("the stalled call completed"), } } @@ -95,6 +97,7 @@ async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, sto .expect("the call reached the stall point"); } +#[rstest::rstest] #[tokio::test] async fn a_hand_driven_machine_performs_the_same_call() { let upstream = upstream([json_response(json!({ @@ -107,13 +110,14 @@ async fn a_hand_driven_machine_performs_the_same_call() { assert_eq!(outcome.unwrap().pages[0].markdown, "native"); assert_eq!(received(&upstream).await.len(), 1); - assert_eq!(ops, ["Project", "BeforeSend", "response"]); + assert_eq!(ops, ["BeforeSend", "response"]); assert!(matches!( machine.resume().await, Err(Error::InvalidRequest(_)) )); } +#[rstest::rstest] #[tokio::test] async fn a_path_document_is_read_by_core_without_a_host_operation() { let upstream = upstream([json_response(json!({ @@ -135,7 +139,7 @@ async fn a_path_document_is_read_by_core_without_a_host_operation() { std::fs::remove_dir_all(&dir).unwrap(); assert_eq!(response.unwrap().pages[0].markdown, "path"); - assert_eq!(ops, ["Project", "BeforeSend", "response"]); + assert_eq!(ops, ["BeforeSend", "response"]); assert_eq!( only_request(&upstream).await.json()["document"]["image_url"], "data:image/png;base64,YWJj" @@ -143,7 +147,7 @@ async fn a_path_document_is_read_by_core_without_a_host_operation() { } #[rstest] -#[case::failed(HostFailure::Error(Error::InvalidRequest("before_send failed".into())), "before_send failed")] +#[case::failed(HostFailure::Error(Error::InvalidRequest("before_provider_request failed".into())), "before_provider_request failed")] #[case::cancelled(HostFailure::Cancelled(Error::InvalidRequest("cancelled".into())), "cancelled")] #[tokio::test] async fn a_before_send_failure_ends_the_call_without_reaching_transport( @@ -159,7 +163,7 @@ async fn a_before_send_failure_ends_the_call_without_reaching_transport( .lock() .unwrap() .take() - .expect("before_send is asked once")) + .expect("before_provider_request is asked once")) }) .await; @@ -167,27 +171,35 @@ async fn a_before_send_failure_ends_the_call_without_reaching_transport( matches!(&outcome, Err(Error::InvalidRequest(actual)) if actual == message), "{outcome:?}" ); - assert_eq!(ops, ["Project", "BeforeSend"]); + assert_eq!(ops, ["BeforeSend"]); assert!(machine.resume().await.is_err()); assert!(received(&upstream).await.is_empty()); } +#[rstest::rstest] #[tokio::test] async fn resuming_before_answering_keeps_the_pending_operation() { - let request = ocr_request("mistral/model", UNREACHABLE_BASE, json!({})); - let mut machine = ocr_machine(ocr_client()); - let Ok(MachineStep::Host(HostOp::Project(reply))) = machine.resume().await else { - panic!("expected the projection op first"); - }; - - assert!(machine.resume().await.is_err()); - reply.send(OcrProjection { + let upstream = upstream([pages_response()]).await; + let request = ocr_request("mistral/model", &upstream.uri(), json!({})); + let mut machine = ocr_route().machine(OcrCall { request, caller_token: false, }); + let Ok(MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest { + wire, + reply, + .. + }))) = machine.resume().await + else { + panic!("expected the provider request hook"); + }; + assert!(machine.resume().await.is_err()); + reply.send(*wire); assert!(matches!( machine.resume().await, - Ok(MachineStep::Host(HostOp::BeforeSend { .. })) + Ok(MachineStep::Suspended(Suspension::Hook( + HookRequest::Event(_, _) + ))) )); } @@ -215,6 +227,7 @@ impl litellm_auth::TokenProvider for PendingToken { } } +#[rstest::rstest] #[tokio::test] async fn interrupt_drops_provider_captures_before_returning() { let entered = Arc::new(Notify::new()); @@ -231,7 +244,7 @@ async fn interrupt_drops_provider_captures_before_returning() { }, ))); let host = LocalOcrHost::new(request); - let mut machine = ocr_machine(ocr_client()); + let mut machine = ocr_route().machine(host.request().unwrap()); drive_until_notified(&mut machine, &host, &entered).await; assert!(!dropped.load(Ordering::SeqCst)); @@ -248,6 +261,7 @@ async fn interrupt_drops_provider_captures_before_returning() { ); } +#[rstest::rstest] #[tokio::test] async fn interrupting_an_in_flight_provider_request_closes_its_connection() { let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); @@ -266,7 +280,7 @@ async fn interrupting_an_in_flight_provider_request_closes_its_connection() { while socket.read(&mut buffer).await.unwrap() != 0 {} }); let host = LocalOcrHost::new(ocr_request("mistral/model", &base, json!({}))); - let mut machine = ocr_machine(ocr_client()); + let mut machine = ocr_route().machine(host.request().unwrap()); drive_until_notified(&mut machine, &host, &received).await; let cancelled = Error::InvalidRequest("cancelled".into()); diff --git a/litellm-rust/crates/core/tests/ocr/main.rs b/litellm-rust/crates/core/tests/ocr/main.rs index e1f6b8cb5c1..d180edaf105 100644 --- a/litellm-rust/crates/core/tests/ocr/main.rs +++ b/litellm-rust/crates/core/tests/ocr/main.rs @@ -1,16 +1,18 @@ use litellm_core::ocr::{ + OcrRoute, document::prepare_document, - route::{LocalOcrHost, ocr_machine}, - types::LiteLLMOcrRequest, + route::{Ocr, OcrCall, OcrOp}, + types::{LiteLLMOcrRequest, OcrDocumentInput}, wire::{OcrWireRequest, decode_request}, }; -use litellm_http::Client; +use litellm_host::event::{CallEvent, RequestContext, WireRequest}; use litellm_llms::base_llm::ocr::{ error::Error, - handler::OcrClient, + settings::OcrSettings, transformation::{LiteLLMOcrResponse, OcrDocument}, }; use serde_json::{Map, Value, json}; +use std::sync::Mutex; use wiremock::{MockServer, ResponseTemplate}; #[path = "../support/mod.rs"] @@ -37,16 +39,31 @@ fn object(value: Value) -> Map { map } -fn ocr_client() -> OcrClient { - OcrClient::for_test(Client::plain_for_test(), Client::no_redirect_for_test()) +fn ocr_route() -> OcrRoute { + ocr_route_with(OcrSettings::default()) +} + +fn ocr_route_with(settings: OcrSettings) -> OcrRoute { + build_ocr_route( + &resources(), + &http_config(), + litellm_http::media::UrlPolicy { + validate: false, + allowed_hosts: Vec::new(), + }, + settings, + no_secrets(), + ) } async fn perform(request: LiteLLMOcrRequest) -> Result { - litellm_core::ocr::client::perform(&ocr_client(), request).await + ocr_route().execute(request, &()).await } async fn perform_with(host: LocalOcrHost) -> Result { - litellm_host::run::run(ocr_machine(ocr_client()), &host).await + litellm_host::in_process::run_hosted(ocr_route().machine(host.request()?), host.runtime()) + .await + .map(completed) } fn wire(model: &str, base: &str, document: Value, options: Value) -> OcrWireRequest { @@ -120,3 +137,117 @@ fn accepted(server: &MockServer, body: Value) -> ResponseTemplate { .insert_header("Operation-Location", format!("{}/operation", server.uri())) .set_body_json(body) } + +fn completed( + result: litellm_host::call::HostedCompletion, +) -> LiteLLMOcrResponse { + match result { + litellm_host::call::HostedCompletion::Complete(response) => response, + other => panic!("unexpected OCR completion: {other:?}"), + } +} + +type BeforeSend = + Box Result + Send + Sync>; +type Observer = Box; + +struct LocalOcrHost { + request: Mutex>>, + before_provider_request: Option, + observer: Option, +} + +impl LocalOcrHost { + fn new(request: LiteLLMOcrRequest) -> Self { + Self { + request: Mutex::new(Some(request)), + before_provider_request: None, + observer: None, + } + } + + fn with_before_send( + self, + before_provider_request: impl Fn(WireRequest, &RequestContext) -> Result + + Send + + Sync + + 'static, + ) -> Self { + Self { + before_provider_request: Some(Box::new(before_provider_request)), + ..self + } + } + + fn with_observer(self, observer: impl Fn(&CallEvent) + Send + Sync + 'static) -> Self { + Self { + observer: Some(Box::new(observer)), + ..self + } + } +} + +impl LocalOcrHost { + pub fn request(&self) -> Result { + self.request + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|request| OcrCall { + request, + caller_token: false, + }) + .ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())) + } + pub fn runtime(&self) -> litellm_host::in_process::Host<'_, Self, Self, ()> { + litellm_host::in_process::Host { + services: self, + hooks: self, + stream: &(), + observer: Some(self), + } + } +} +impl litellm_host::services::HostCallHandler for LocalOcrHost { + async fn handle_host_call(&self, op: OcrOp) -> Result<(), Error> { + match op { + OcrOp::AcquireAzureAdToken(_) => { + Err(Error::Auth(litellm_auth::Error::CredentialAcquisition( + "OCR host has no Azure AD token provider".into(), + ))) + } + } + } +} + +impl litellm_host::lifecycle::CallObserver for LocalOcrHost { + fn observe(&self, event: litellm_host::event::CallEvent) { + if let Some(observer) = &self.observer { + observer(&event); + } + } +} +impl litellm_host::hooks::RouteHooks<::Error> + for LocalOcrHost +{ + async fn before_provider_request( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + match &self.before_provider_request { + Some(before_provider_request) => before_provider_request(wire, &context), + None => Ok(wire), + } + } + async fn on_event( + &self, + event: litellm_host::event::MachineEvent, + ) -> Result<(), ::Error> { + litellm_host::lifecycle::CallObserver::observe( + self, + litellm_host::event::CallEvent::Machine(event), + ); + Ok(()) + } +} diff --git a/litellm-rust/crates/core/tests/ocr/mistral.rs b/litellm-rust/crates/core/tests/ocr/mistral.rs index d542eeaf03a..6c5463ce035 100644 --- a/litellm-rust/crates/core/tests/ocr/mistral.rs +++ b/litellm-rust/crates/core/tests/ocr/mistral.rs @@ -1,11 +1,8 @@ use std::sync::Arc; -use litellm_http::{HttpSettings, Resolution, media::UrlPolicy}; +use litellm_http::{HttpSettings, Resolution}; use litellm_llms::{ - base_llm::ocr::{ - settings::OcrSettings, - transformation::{BaseOcrConfig, OCR_RESPONSE_MAX_BYTES}, - }, + base_llm::ocr::transformation::{BaseOcrConfig, OCR_RESPONSE_MAX_BYTES}, mistral::ocr::transformation::MistralOcrConfig, }; use rstest::rstest; @@ -156,7 +153,13 @@ async fn missing_credentials_come_from_the_injected_secret_source( .copied() .chain([("MISTRAL_AZURE_API_BASE", base.as_str())]), )); - let client = ocr_client().with_secrets(source.clone()); + let route = build_ocr_route( + &resources(), + &http_config(), + Default::default(), + Default::default(), + source.clone(), + ); let request = decode_request(OcrWireRequest { api_key: None, api_base: None, @@ -169,9 +172,7 @@ async fn missing_credentials_come_from_the_injected_secret_source( }) .unwrap(); - litellm_core::ocr::client::perform(&client, request) - .await - .unwrap(); + route.execute(request, &()).await.unwrap(); assert_eq!(source.requested(), MistralOcrConfig.secret_names()); assert_eq!( @@ -188,25 +189,21 @@ async fn the_client_uses_the_injected_http_pool_configuration() { user_agent: Some("host-owned/1".into()), ..HttpSettings::default() }; - let client = resources() - .ocr_client( - &Resolution::from(&settings).config, - UrlPolicy::default(), - OcrSettings::default(), - Arc::new( - litellm_secrets::source::EnvironmentSecrets::python_compatible( - litellm_http::Client::plain_for_test(), - ), - ), - ) - .unwrap(); + let route = build_ocr_route( + &resources(), + &Resolution::from(&settings).config, + Default::default(), + Default::default(), + no_secrets(), + ); - litellm_core::ocr::client::perform( - &client, - ocr_request("mistral/model", &upstream.uri(), json!({})), - ) - .await - .unwrap(); + route + .execute( + ocr_request("mistral/model", &upstream.uri(), json!({})), + &(), + ) + .await + .unwrap(); assert_eq!( only_request(&upstream).await.header("user-agent"), diff --git a/litellm-rust/crates/core/tests/ocr/vertex_ai.rs b/litellm-rust/crates/core/tests/ocr/vertex_ai.rs index f0b2488e494..d77c065686e 100644 --- a/litellm-rust/crates/core/tests/ocr/vertex_ai.rs +++ b/litellm-rust/crates/core/tests/ocr/vertex_ai.rs @@ -45,18 +45,19 @@ async fn mistral_is_served_at_the_resolved_project_and_location() { #[tokio::test] async fn configured_project_and_location_apply_when_the_call_sets_neither() { let upstream = upstream([pages_response()]).await; - let client = ocr_client().with_settings(OcrSettings { + let route = ocr_route_with(OcrSettings { vertex_project: Some("configured-project".into()), vertex_location: Some("europe-west4".into()), ..OcrSettings::default() }); - litellm_core::ocr::client::perform( - &client, - ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})), - ) - .await - .unwrap(); + route + .execute( + ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})), + &(), + ) + .await + .unwrap(); assert_eq!( only_request(&upstream).await.url.path(), diff --git a/litellm-rust/crates/core/tests/resources.rs b/litellm-rust/crates/core/tests/resources.rs index 9764e50de1b..15cfc063506 100644 --- a/litellm-rust/crates/core/tests/resources.rs +++ b/litellm-rust/crates/core/tests/resources.rs @@ -10,10 +10,7 @@ use litellm_auth_gcp::{ CredentialSource, VertexAuth, VertexAuthFuture, VertexProviderLoader, VertexTokenSource, }; use litellm_core::{ - ocr::{ - client::perform, - wire::{OcrWireRequest, decode_request}, - }, + ocr::wire::{OcrWireRequest, decode_request}, resources::CoreResources, }; use litellm_http::{HttpSettings, Resolution}; @@ -105,17 +102,16 @@ async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets( ..HttpSettings::default() }) .config; - let client = owner - .ocr_client( - &http, - Default::default(), - OcrSettings { - vertex_location: Some(location.into()), - ..OcrSettings::default() - }, - Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])), - ) - .unwrap(); + let route = support::build_ocr_route( + owner, + &http, + Default::default(), + OcrSettings { + vertex_location: Some(location.into()), + ..OcrSettings::default() + }, + Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])), + ); let request = decode_request(OcrWireRequest { model: "vertex_ai/mistral-ocr-maas".into(), document: json!({"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}), @@ -127,7 +123,7 @@ async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets( input_sources: Default::default(), timeout_seconds: Some(5.0), }).unwrap(); - let result = perform(&client, request).await.unwrap(); + let result = route.execute(request, &()).await.unwrap(); assert!(!result.pages.is_empty()); } let requests = upstream.received_requests().await.unwrap(); diff --git a/litellm-rust/crates/core/tests/support/mod.rs b/litellm-rust/crates/core/tests/support/mod.rs index 5443437df09..d45a537b19d 100644 --- a/litellm-rust/crates/core/tests/support/mod.rs +++ b/litellm-rust/crates/core/tests/support/mod.rs @@ -7,7 +7,8 @@ use std::sync::{Arc, Mutex}; use futures_util::future::BoxFuture; use litellm_http::{ - HttpClientConfig, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver, + ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, Resolution, + media::PublicDnsResolver, }; use litellm_secrets::{SecretValue, source::SecretSource}; use serde_json::Value; @@ -24,6 +25,67 @@ pub fn resources() -> litellm_core::resources::CoreResources { litellm_core::resources::CoreResources::new(Arc::new(http_pool())) } +pub fn no_secrets() -> Arc { + Arc::new(RecordingSecrets::empty()) +} + +pub fn provider_http( + resources: &litellm_core::resources::CoreResources, + config: &HttpClientConfig, +) -> litellm_http::Client { + resources + .pool + .client(config, ClientVariant::Provider) + .unwrap() +} + +pub fn messages_route(secrets: Arc) -> litellm_core::messages::MessagesRoute { + let resources = resources(); + litellm_core::messages::MessagesRoute::new( + provider_http(&resources, &http_config()), + resources.auth, + secrets, + ) +} + +pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletionsRoute { + let resources = resources(); + litellm_core::chat_completions::ChatCompletionsRoute::new( + provider_http(&resources, &http_config()), + resources.auth, + no_secrets(), + ) +} + +pub fn audio_transcription_route() -> litellm_core::audio_transcription::AudioTranscriptionRoute { + let resources = resources(); + litellm_core::audio_transcription::AudioTranscriptionRoute::new( + provider_http(&resources, &http_config()), + resources.auth, + no_secrets(), + ) +} + +pub fn build_ocr_route( + resources: &litellm_core::resources::CoreResources, + config: &HttpClientConfig, + url_policy: litellm_http::media::UrlPolicy, + settings: litellm_llms::base_llm::ocr::settings::OcrSettings, + secrets: Arc, +) -> litellm_core::ocr::OcrRoute { + litellm_core::ocr::OcrRoute::new( + litellm_llms::base_llm::ocr::handler::OcrClient::new( + &resources.pool, + config, + url_policy, + resources.auth.clone(), + settings, + secrets, + ) + .unwrap(), + ) +} + pub fn http_config() -> HttpClientConfig { Resolution::from(&HttpSettings::default()).config } @@ -168,3 +230,114 @@ impl SecretSource for RecordingSecrets { }) } } + +pub struct RecordingCall { + pub request: Mutex>, + pub events: Arc, + pub chunks: Mutex>, + pub head: Mutex>, +} + +#[derive(Default)] +pub struct CallEvents(pub Mutex>); + +impl litellm_host::lifecycle::CallObserver for CallEvents { + fn observe(&self, event: litellm_host::event::CallEvent) { + self.0.lock().unwrap().push(event); + } +} + +impl RecordingCall

{ + pub fn new(request: P::Request) -> Self { + Self { + request: Mutex::new(Some(request)), + events: Arc::new(CallEvents::default()), + chunks: Mutex::new(Vec::new()), + head: Mutex::new(None), + } + } +} + +impl litellm_host::hooks::RouteHooks + 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 { + headers: wire + .headers + .into_iter() + .chain([("x-hook".into(), "called".into())]) + .collect(), + ..wire + }) + } + + async fn on_event(&self, event: litellm_host::event::MachineEvent) -> Result<(), P::Error> { + self.events + .0 + .lock() + .unwrap() + .push(litellm_host::event::CallEvent::Machine(event)); + Ok(()) + } +} + +impl

RecordingCall

+where + P: litellm_host::protocol::Protocol, + P::Error: From, +{ + pub fn request(&self) -> Result { + self.request + .lock() + .unwrap() + .take() + .ok_or_else(|| litellm_host::machine::MachineFault::Abandoned.into()) + } + pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, Self> { + litellm_host::in_process::Host { + services: &(), + hooks: self, + stream: self, + observer: Some(self), + } + } +} + +impl

litellm_host::in_process::StreamConsumer

for RecordingCall

+where + P: litellm_host::protocol::Protocol, + P::Error: From, +{ + async fn open_stream( + &self, + head: P::StreamHead, + ) -> Result { + *self.head.lock().unwrap() = Some(head); + Ok(litellm_host::protocol::Demand::More) + } + async fn send_chunk( + &self, + chunk: P::Chunk, + ) -> Result { + self.chunks.lock().unwrap().push(chunk); + Ok(litellm_host::protocol::Demand::More) + } +} +impl

litellm_host::lifecycle::CallObserver for RecordingCall

+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()); + } +} diff --git a/litellm-rust/crates/gateway-inference/src/audio_transcription.rs b/litellm-rust/crates/gateway-inference/src/audio_transcription.rs index 39229303adf..1f3d32dac80 100644 --- a/litellm-rust/crates/gateway-inference/src/audio_transcription.rs +++ b/litellm-rust/crates/gateway-inference/src/audio_transcription.rs @@ -6,7 +6,7 @@ use axum::{ response::{IntoResponse, Response}, }; use base64::{Engine, engine::general_purpose::STANDARD}; -use litellm_core::audio_transcription::{audio_transcription, types::AudioTranscriptionRequest}; +use litellm_core::audio_transcription::types::AudioTranscriptionRequest; use serde_json::{Value, json}; use crate::{Error, Gateway, request}; @@ -38,11 +38,9 @@ async fn handle(gateway: &Gateway, request: Request) -> Result { .cloned() .ok_or_else(|| Error::InvalidBody("audio is required".into()))?, }; - Ok(audio_transcription( - &gateway.resources, - &gateway.http, - gateway.secrets.as_ref(), - AudioTranscriptionRequest { + Ok(gateway + .audio_transcription + .execute(AudioTranscriptionRequest { model: &deployment.model, audio, api_key: deployment.api_key.as_deref(), @@ -54,7 +52,6 @@ async fn handle(gateway: &Gateway, request: Request) -> Result { .filter(|(name, _)| !matches!(name.as_str(), "model" | "audio")) .collect(), timeout: deployment.timeout, - }, - ) - .await?) + }) + .await?) } diff --git a/litellm-rust/crates/gateway-inference/src/chat_completions.rs b/litellm-rust/crates/gateway-inference/src/chat_completions.rs index 8c66e4c49da..814a368ea24 100644 --- a/litellm-rust/crates/gateway-inference/src/chat_completions.rs +++ b/litellm-rust/crates/gateway-inference/src/chat_completions.rs @@ -7,7 +7,7 @@ use axum::{ http::StatusCode, response::{IntoResponse, Response}, }; -use litellm_core::chat_completions::{chat_completions, types::ChatCompletionsRequest}; +use litellm_core::chat_completions::types::ChatCompletionsRequest; use serde_json::{Map, Value}; use crate::{Error, Gateway, request}; @@ -61,24 +61,24 @@ async fn handle(gateway: &Gateway, body: Map) -> Result, pub resources: CoreResources, pub http: HttpClientConfig, - pub secrets: Arc, - pub models: ModelList, - pub ocr: OcrClient, +} + +impl Gateway { + pub fn new( + resources: CoreResources, + http: HttpClientConfig, + secrets: Arc, + models: ModelList, + ) -> Result { + let provider = resources.pool.client(&http, ClientVariant::Provider)?; + let auth = resources.auth.clone(); + Ok(Self { + audio_transcription: AudioTranscriptionRoute::new( + provider.clone(), + auth.clone(), + secrets.clone(), + ), + chat_completions: ChatCompletionsRoute::new( + provider.clone(), + auth.clone(), + secrets.clone(), + ), + messages: MessagesRoute::new(provider, auth.clone(), secrets.clone()), + ocr: OcrRoute::new(OcrClient::new( + &resources.pool, + &http, + UrlPolicy::default(), + auth, + OcrSettings::default(), + secrets.clone(), + )?), + models, + secrets, + resources, + http, + }) + } } pub fn router(gateway: Arc) -> Router { diff --git a/litellm-rust/crates/gateway-inference/src/messages/mod.rs b/litellm-rust/crates/gateway-inference/src/messages/mod.rs index 5d6a8faa0e8..837757aefa8 100644 --- a/litellm-rust/crates/gateway-inference/src/messages/mod.rs +++ b/litellm-rust/crates/gateway-inference/src/messages/mod.rs @@ -10,9 +10,7 @@ use axum::{ response::{IntoResponse, Response}, }; use futures_util::{StreamExt, stream::BoxStream}; -use litellm_core::messages::{ - Error as RouteError, MessagesCall, MessagesResponse, messages, messages_body, -}; +use litellm_core::messages::{Error as RouteError, MessagesCall, MessagesResponse, messages_body}; use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders}; use serde_json::{Map, Value}; @@ -52,15 +50,8 @@ async fn handle(gateway: &Gateway, headers: &HeaderMap, body: &[u8]) -> Result Ok(Json(message).into_response()), + match gateway.messages.execute(call, &()).await? { + MessagesResponse::Complete(message) => Ok(Json(message).into_response()), MessagesResponse::Stream { chunks, .. } => Ok(stream(chunks)), } } diff --git a/litellm-rust/crates/gateway-inference/src/ocr.rs b/litellm-rust/crates/gateway-inference/src/ocr.rs index d666223e037..eb8f0e7307e 100644 --- a/litellm-rust/crates/gateway-inference/src/ocr.rs +++ b/litellm-rust/crates/gateway-inference/src/ocr.rs @@ -6,10 +6,7 @@ use axum::{ response::{IntoResponse, Response}, }; use litellm_auth::SecretValue; -use litellm_core::ocr::{ - client::perform, - types::{LiteLLMOcrRequest, OcrConnectionInputs, OcrDocumentInput}, -}; +use litellm_core::ocr::types::{LiteLLMOcrRequest, OcrConnectionInputs, OcrDocumentInput}; use litellm_llms::base_llm::ocr::transformation::OcrDocument; use serde_json::Value; @@ -69,7 +66,7 @@ async fn handle(gateway: &Gateway, request: Request) -> Result { ..Default::default() }, )?; - let response = perform(&gateway.ocr, call).await?; + let response = gateway.ocr.execute(call, &()).await?; match response.provider_native_response { Some(native) => Ok(Value::Object(native)), None => Ok(response.into_json()), diff --git a/litellm-rust/crates/gateway-inference/tests/support/mod.rs b/litellm-rust/crates/gateway-inference/tests/support/mod.rs index d56489d28cd..5fa0b044e3c 100644 --- a/litellm-rust/crates/gateway-inference/tests/support/mod.rs +++ b/litellm-rust/crates/gateway-inference/tests/support/mod.rs @@ -10,7 +10,6 @@ use futures_util::future::BoxFuture; use litellm_core::resources::CoreResources; use litellm_gateway_inference::{Deployment, Gateway, router}; use litellm_http::{HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver}; -use litellm_llms::base_llm::ocr::settings::OcrSettings; use litellm_secrets::{SecretValue, source::SecretSource}; use serde_json::Value; use tower::ServiceExt; @@ -31,32 +30,26 @@ pub fn app(model: &str, api_base: &str) -> Router { let http = Resolution::from(&HttpSettings::default()).config; let secrets = Arc::new(NoSecrets); let resources = CoreResources::new(pool); - let ocr = resources - .ocr_client( - &http, - Default::default(), - OcrSettings::default(), - secrets.clone(), + router(Arc::new( + Gateway::new( + resources, + http, + secrets, + [( + "public/model".into(), + Deployment { + model: model.into(), + api_base: Some(api_base.into()), + api_key: Some("test-key".into()), + timeout: Some(Duration::from_secs(5)), + ..Default::default() + }, + )] + .into_iter() + .collect(), ) - .unwrap(); - router(Arc::new(Gateway { - resources, - http, - secrets, - ocr, - models: [( - "public/model".into(), - Deployment { - model: model.into(), - api_base: Some(api_base.into()), - api_key: Some("test-key".into()), - timeout: Some(Duration::from_secs(5)), - ..Default::default() - }, - )] - .into_iter() - .collect(), - })) + .unwrap(), + )) } pub async fn post(app: Router, path: &str, body: Value) -> Response { diff --git a/litellm-rust/crates/gateway/src/lib.rs b/litellm-rust/crates/gateway/src/lib.rs index fc16dc0de67..33d96234015 100644 --- a/litellm-rust/crates/gateway/src/lib.rs +++ b/litellm-rust/crates/gateway/src/lib.rs @@ -16,7 +16,6 @@ use litellm_gateway_inference::{Gateway, ModelList}; use litellm_http::{ ClientVariant, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver, }; -use litellm_llms::base_llm::ocr::settings::OcrSettings; use litellm_secrets::source::EnvironmentSecrets; use litellm_tracing::ByteChunk; use uuid::Uuid; @@ -27,20 +26,12 @@ pub fn build_inference(config: &Config) -> Result, litellm_http::Er let client = pool.client(&http, ClientVariant::Provider)?; let secrets = Arc::new(EnvironmentSecrets::python_compatible(client)); let resources = CoreResources::new(pool); - let ocr = resources.ocr_client( - &http, - Default::default(), - OcrSettings::default(), - secrets.clone(), - )?; - - Ok(Arc::new(Gateway { + Ok(Arc::new(Gateway::new( resources, http, secrets, - models: ModelList::from_model_list(&config.model_list), - ocr, - })) + ModelList::from_model_list(&config.model_list), + )?)) } pub fn router(inference: Arc, config: &Config) -> Router { diff --git a/litellm-rust/crates/host-python/AGENTS.md b/litellm-rust/crates/host-python/AGENTS.md index cadc55a35a7..4b4cfc1a772 100644 --- a/litellm-rust/crates/host-python/AGENTS.md +++ b/litellm-rust/crates/host-python/AGENTS.md @@ -1,10 +1,10 @@ - Target invariants; implementation and runtime validation may lag these rules -- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`ProtocolHost` traits +- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonBinding`, `PythonHostCalls`, `PythonCallHooks` and `PythonOwned` traits - No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features - - The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business - - `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, rewritten in place by the route's `Preflight`, not the caller's dict; a protocol host that projects from it inherits the adapter's rewrites (for the legacy adapter: setup, deployment hooks) and the preflight's (credential inheritance) - - A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is - - A failing `classify` is raised with the native error's text as its `__context__`, never swallowed + - The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the lifecycle's business + - `PythonBinding::decode_request` receives the keyword view the hooks' `prepare_arguments` returned, rewritten in place by the route's `Preflight`, not the caller's dict; a binding that decodes from it inherits the lifecycle's rewrites (for the legacy adapter: setup, deployment hooks) and the preflight's (credential inheritance) + - A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the binding's `map_error`; a Python exception raised inside the call, and a failure in `prepare_arguments` or `transform_response`, is raised as is + - A failing `map_error` is raised with the native error's text as its `__context__`, never swallowed - Use standard PyO3 ownership and conversion APIs - Prefer `Bound<'py, T>` for attached operations/results, `Py` for retention; binding/unbinding does not copy payloads - Use `pythonize` for selected Serde data, never a JSON-text round trip; share conversion with `Pythonized` @@ -15,6 +15,6 @@ - Release exclusive class borrows/locks before Python calls or decrements that can invoke finalizers; expose retained Python edges to GC without calling Python during traversal - Keep coroutine driving in the shared Python driver and the native handle - Driver: `litellm/rust_bridge/lifecycle.py`; handle: `src/handle.rs`; call driver: `src/driver.rs`; native-backed behavior tests: `tests/lifecycle.py` - - Every adapter suspension is awaited inline in the caller's task; `into_future` creates a separate task and cannot satisfy this contract + - Every lifecycle suspension is awaited inline in the caller's task; `into_future` creates a separate task and cannot satisfy this contract - References: [ownership](https://pyo3.rs/v0.29.2/types.html), [conversions](https://pyo3.rs/v0.29.2/conversions/traits.html), [pythonize errors](https://docs.rs/pythonize/0.29.0/src/pythonize/error.rs.html) - [GC](https://pyo3.rs/v0.29.2/class/protocols.html#garbage-collector-integration), [re-entry](https://pyo3.rs/v0.29.2/class/call.html), [parallelism](https://pyo3.rs/v0.29.2/parallelism.html), [async delivery source](https://docs.rs/pyo3-async-runtimes/0.29.0/src/pyo3_async_runtimes/generic.rs.html) diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs deleted file mode 100644 index 83ed6416d6e..00000000000 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ /dev/null @@ -1,168 +0,0 @@ -use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest}; -use litellm_host::protocol::Protocol; -use pyo3::exceptions::PyRuntimeError; -use pyo3::gc::{PyTraverseError, PyVisit}; -use pyo3::prelude::*; -use pyo3::types::PyDict; - -pub fn missing_state() -> PyErr { - PyRuntimeError::new_err("missing native call state") -} - -/// The SDK's request policy, run by the driver on the keyword view `begin` returned and -/// before the protocol host projects from it. It rewrites that view in place, so the -/// lifecycle that returned it sees the rewrite too; a rejection fails the call as a host -/// failure, so the lifecycle still observes it. -pub type Preflight = fn(Python<'_>, &Bound<'_, PyDict>) -> PyResult<()>; - -/// What an adapter step produced: either the value the driver asked for, or a Python -/// awaitable the driver hands back to the caller's task before asking again. -pub enum LifecycleStep { - Await(Py), - Arguments(Py), - Wire(Box), - Response(Py), - Done, -} - -/// What a lifecycle observes: the driver's start, the machine's own events, and one -/// terminal event carrying the public value the caller receives. -pub enum LifecycleEvent<'a> { - Started { - start_time: f64, - }, - Machine(&'a MachineEvent), - Succeeded { - timing: Timing, - response: &'a Py, - }, - Failed { - timing: Timing, - origin: FailureOrigin, - error: &'a PyErr, - }, -} - -/// One consumer of a call's lifecycle on the Python side. The driver calls the steps in -/// order: `begin` before the machine starts, `before_send` and `emit` while it runs, -/// `after_success` and one terminal `emit` after it completes. Whenever a step returns -/// [`LifecycleStep::Await`], the driver awaits it in the caller's task and continues the -/// same step through `resume`. -/// -/// A step that fails with an ordinary exception fails the call with that exception, -/// except on a terminal event, where the adapter is expected to report and swallow its -/// own errors. An exception that is not a `PyException`, such as a cancellation, ends -/// the call without further dispatch. -pub trait PythonLifecycle: Send + Sync { - fn begin( - &mut self, - py: Python<'_>, - arguments: Py, - started_at: f64, - ) -> PyResult; - - fn before_send( - &mut self, - py: Python<'_>, - wire: Box, - context: &RequestContext, - ) -> PyResult; - - fn after_success( - &mut self, - py: Python<'_>, - response: Py, - timing: Timing, - ) -> PyResult; - - fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult; - - /// The call streams and its stream was handed to the caller. The caller is not - /// inside an await here, so this step and `delivered` cannot suspend. - fn opened(&mut self, py: Python<'_>) -> PyResult<()>; - - /// One chunk of an open stream is about to reach the caller. - fn delivered(&mut self, py: Python<'_>, chunk: &Py) -> PyResult<()>; - - fn resume(&mut self, py: Python<'_>, result: PyResult>) -> PyResult; - - fn close(&mut self, py: Python<'_>); - - fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; -} - -/// Why a custom operation the host answered did not produce a result: the route's own code -/// rejected it, which the route classifies like any other native failure, or Python code -/// raised, which reaches the caller as it was raised. -#[derive(Debug)] -pub enum InvokeError { - Native(E), - Python(PyErr), -} - -impl From for InvokeError { - fn from(error: PyErr) -> Self { - Self::Python(error) - } -} - -/// The Python side of one protocol: answers its custom operations, builds the public -/// response and classifies native failures into public exceptions. -pub trait ProtocolHost: Send + Sync { - type Protocol: Protocol; - - /// The public exception a native failure maps to, kept as a value until the driver - /// raises it. - type Failure: Into; - - /// Projects the call's request. `arguments` is the keyword view the lifecycle's - /// `begin` produced, not the caller's own dict, so the projection inherits whatever - /// that adapter rewrote. - fn project( - &mut self, - py: Python<'_>, - arguments: &Bound<'_, PyDict>, - ) -> Result< - ::Projection, - InvokeError<::Error>, - >; - - /// Answers `op` through its reply. - fn invoke( - &mut self, - py: Python<'_>, - op: ::Op, - ) -> Result<(), InvokeError<::Error>>; - - fn complete( - &mut self, - py: Python<'_>, - response: ::Response, - ) -> PyResult>; - - /// What the stream carries at hand-off, as the caller's stream receives it. - fn head( - &mut self, - py: Python<'_>, - head: ::StreamHead, - ) -> PyResult>; - - /// One streamed chunk as the caller receives it. - fn chunk( - &mut self, - py: Python<'_>, - chunk: ::Chunk, - ) -> PyResult>; - - fn classify( - &self, - py: Python<'_>, - error: ::Error, - ) -> PyResult; - - fn host_error(error: &PyErr) -> ::Error; - - fn close(&mut self, py: Python<'_>); - - fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; -} diff --git a/litellm-rust/crates/host-python/src/binding.rs b/litellm-rust/crates/host-python/src/binding.rs new file mode 100644 index 00000000000..91fbc8c0e52 --- /dev/null +++ b/litellm-rust/crates/host-python/src/binding.rs @@ -0,0 +1,51 @@ +use crate::{InvokeError, PythonOwned}; +use litellm_host::protocol::Protocol; +use pyo3::prelude::*; +use pyo3::types::PyDict; + +/// Converts requests, responses, stream values and errors at the Python boundary. +pub trait PythonBinding: PythonOwned { + type Protocol: Protocol; + + /// The public exception a native failure maps to, kept as a value until the driver + /// raises it. + type Failure: Into; + + /// Decodes the keyword view returned by `prepare_arguments`, including preflight rewrites. + fn decode_request( + &mut self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> Result< + ::Request, + InvokeError<::Error>, + >; + + fn encode_response( + &mut self, + py: Python<'_>, + response: ::Response, + ) -> PyResult>; + + /// What the stream carries at hand-off, as the caller's stream receives it. + fn encode_stream_head( + &mut self, + py: Python<'_>, + head: ::StreamHead, + ) -> PyResult>; + + /// One streamed chunk as the caller receives it. + fn encode_chunk( + &mut self, + py: Python<'_>, + chunk: ::Chunk, + ) -> PyResult>; + + fn map_error( + &self, + py: Python<'_>, + error: ::Error, + ) -> PyResult; + + fn host_error(error: &PyErr) -> ::Error; +} diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 50eae1e0225..123840151cb 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -1,41 +1,28 @@ -use std::sync::Arc; -use std::task::Poll; - -use futures_util::future::{AbortHandle, Abortable}; +use crate::PythonHostCalls; +use litellm_host::call::HostedCompletion; use litellm_host::event::WireRequest; use litellm_host::event::{FailureOrigin, Timing, epoch_seconds}; -use litellm_host::host::{Demand, HostOp, HostStep, Reply}; use litellm_host::machine::{HostFailure, Machine, MachineStep}; -use litellm_host::protocol::Protocol; +use litellm_host::protocol::HookRequest; +use litellm_host::protocol::StreamDelivery; +use litellm_host::protocol::{Demand, Protocol, Reply, Suspension}; use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; use pyo3::types::PyDict; -use tokio::sync::Mutex; -use crate::adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, Preflight, ProtocolHost, PythonLifecycle, - missing_state, -}; -use crate::execution::{poll_async_value, run_async_value, run_sync_value}; use crate::handle::{Execution, ExecutionBody, ExecutionStep}; +use crate::hooks::{HookEvent, HookResume, HookStep, Preflight, PythonCallHooks}; +use crate::native::{NativeMachine, NativePoll}; +use crate::{InvokeError, PythonBinding, missing_state}; -type ProtocolOf = ::Protocol; +type ProtocolOf = ::Protocol; type ErrorOf = as Protocol>::Error; type ResponseOf = as Protocol>::Response; -type NativeStep = MachineStep, ResponseOf>; +type NativeStep = MachineStep, HostedCompletion>>; type NativeResult = Result, ErrorOf>; type Interruption = Option>>; - -type MachineResult = Result< - MachineStep<::Protocol, ::Complete>, - <::Protocol as Protocol>::Error, ->; - -struct MachineState { - machine: M, - result: Option>, -} +type StartMachine = Box::Request) -> M + Send + Sync>; enum Stage { Begin, @@ -46,19 +33,18 @@ enum Stage { Failed(Py), } -enum Expect { +enum EventNext { Started, - Arguments, - Wire(Reply), Emitted(Reply<()>), - Response, Terminal, } -enum Pending { +enum Pending { Native, - Adapter(Expect), - /// The stream handed to the caller waits for its next read or its close. + Arguments(HookResume>), + Wire(HookResume>, Reply), + Response(HookResume>), + Event(HookResume, EventNext), Consumer(Reply), } @@ -72,77 +58,71 @@ fn answered(answer: Result<(), InvokeError>) -> PyResult> { } } -enum Next { +enum Next> { Return(ExecutionStep), - Continue(HostStep, Py>), + Continue(NativePoll>), } -struct PythonDriver +struct PythonDriver where - H: ProtocolHost, - M: Machine> + 'static, + L: PythonCallHooks + 'static, + H: PythonBinding + PythonHostCalls, + M: Machine + 'static, + M::Complete: Into>>, { - host: H, - adapter: Box, + binding: H, + hooks: L, preflight: Preflight, - machine: Option>>>, + native: NativeMachine, + start: Option, M>>, + closed: bool, arguments: Option>, started_at: f64, ended_at: Option, stage: Stage, - pending: Option, - native_abort: Option, + pending: Option>, interrupted: Option>, - asynchronous: bool, } /// 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 adapter's `begin` returned, before the host projects from it. -pub fn run_call( +/// the hooks' `prepare_arguments` returned, before the binding decodes the request. +pub fn run_call( py: Python<'_>, - machine: M, - host: H, - adapter: Box, + start: impl FnOnce(::Request) -> M + Send + Sync + 'static, + binding: H, + hooks: L, preflight: Preflight, arguments: Py, asynchronous: bool, ) -> PyResult> where - H: ProtocolHost + 'static, - M: Machine> + 'static, + L: PythonCallHooks + 'static, + H: PythonBinding + PythonHostCalls + 'static, + M: Machine + 'static, + M::Complete: Into>>, { let mut driver = PythonDriver { - host, - adapter, + binding, + hooks, preflight, - machine: Some(Arc::new(Mutex::new(MachineState { - machine, - result: None, - }))), + native: NativeMachine::new(asynchronous), + start: Some(Box::new(start)), + closed: false, arguments: Some(arguments), started_at: 0.0, ended_at: None, stage: Stage::Begin, pending: None, - native_abort: None, interrupted: None, - asynchronous, }; if asynchronous { - let execution = Py::new(py, Execution::new(driver))?; - return py - .import("litellm.rust_bridge.lifecycle")? - .getattr("drive")? - .call1((execution,)) - .map(Bound::unbind); + return Execution::new(driver).into_coroutine(py).map(Bound::unbind); } match driver.resume(None)? { ExecutionStep::Return(value) => Ok(value), - ExecutionStep::Open(head) => py - .import("litellm.rust_bridge.lifecycle")? - .getattr("SyncStream")? - .call1((Py::new(py, Execution::suspended(driver))?, head)) + ExecutionStep::Open(head) => Execution::suspended(driver) + .into_sync_stream(py, head) .map(Bound::unbind), ExecutionStep::Await(_) | ExecutionStep::Yield(_) => { Err(PyRuntimeError::new_err("sync call suspended")) @@ -154,10 +134,12 @@ fn is_cancellation(py: Python<'_>, error: &PyErr) -> bool { !error.is_instance_of::(py) } -impl PythonDriver +impl PythonDriver where - H: ProtocolHost, - M: Machine> + 'static, + L: PythonCallHooks + 'static, + H: PythonBinding + PythonHostCalls, + M: Machine + 'static, + M::Complete: Into>>, { fn timing(&self) -> Timing { Timing { @@ -174,17 +156,17 @@ where match (self.pending.take(), result) { (None, None) => { self.started_at = epoch_seconds(); - let started = LifecycleEvent::Started { + let started = HookEvent::Started { start_time: self.started_at, }; - match self.adapter.emit(py, started) { - Ok(step) => self.on_adapter(py, step, Expect::Started), - Err(error) => self.adapter_failed(py, error), + match self.hooks.on_event(py, started) { + Ok(step) => self.on_event(py, step, EventNext::Started), + Err(error) => self.hook_failed(py, error), } } (Some(Pending::Native), Some(Ok(_))) => { - let result = self.take_native_result()?; - self.run_steps(py, HostStep::Ready(result)) + let result = self.native.take_result()?; + self.run_steps(py, NativePoll::Ready(result)) } (Some(Pending::Native), Some(Err(error))) => self.interrupt(py, error), (Some(Pending::Consumer(reply)), Some(read)) => { @@ -195,63 +177,138 @@ where }); self.resume_machine(py, None) } - (Some(Pending::Adapter(expect)), Some(result)) => { - match self.adapter.resume(py, result) { - Ok(step) => self.on_adapter(py, step, expect), - Err(error) => self.adapter_failed(py, error), + (Some(Pending::Arguments(resume)), Some(result)) => { + let step = resume(&mut self.hooks, py, result); + match step { + Ok(step) => self.on_arguments(py, step), + Err(error) => self.hook_failed(py, error), + } + } + (Some(Pending::Wire(resume, reply)), Some(result)) => { + let step = resume(&mut self.hooks, py, result); + match step { + Ok(step) => self.on_wire(py, step, reply), + Err(error) => self.hook_failed(py, error), + } + } + (Some(Pending::Response(resume)), Some(result)) => { + let step = resume(&mut self.hooks, py, result); + match step { + Ok(step) => self.on_response(py, step), + Err(error) => self.hook_failed(py, error), + } + } + (Some(Pending::Event(resume, next)), Some(result)) => { + let step = resume(&mut self.hooks, py, result); + match step { + Ok(step) => self.on_event(py, step, next), + Err(error) => self.hook_failed(py, error), } } _ => Err(missing_state()), } } - fn on_adapter( + fn on_arguments( &mut self, py: Python<'_>, - step: LifecycleStep, - expect: Expect, + step: HookStep>, ) -> PyResult { - if let LifecycleStep::Await(awaitable) = step { - self.pending = Some(Pending::Adapter(expect)); - return Ok(ExecutionStep::Await(awaitable)); - } - match (expect, step) { - (Expect::Started, LifecycleStep::Done) => self.begin(py), - (Expect::Arguments, LifecycleStep::Arguments(arguments)) => { + match step { + HookStep::Await(awaitable, resume) => { + self.pending = Some(Pending::Arguments(resume)); + Ok(ExecutionStep::Await(awaitable)) + } + HookStep::Ready(arguments) => { if let Err(error) = (self.preflight)(py, arguments.bind(py)) { - return self.adapter_failed(py, error); + return self.hook_failed(py, error); } + let decoded = self.binding.decode_request(py, arguments.bind(py)); self.arguments = Some(arguments); + let request = match decoded { + Ok(request) => request, + Err(InvokeError::Native(error)) => return self.machine_failed(py, error), + Err(InvokeError::Python(error)) => { + return self.failure(py, error, FailureOrigin::Call); + } + }; + let start = self.start.take().ok_or_else(missing_state)?; + self.native.start(start(request)); self.stage = Stage::Call; self.resume_machine(py, None) } - (Expect::Wire(reply), LifecycleStep::Wire(wire)) => { + } + } + + fn on_wire( + &mut self, + py: Python<'_>, + step: HookStep>, + reply: Reply, + ) -> PyResult { + match step { + HookStep::Await(awaitable, resume) => { + self.pending = Some(Pending::Wire(resume, reply)); + Ok(ExecutionStep::Await(awaitable)) + } + HookStep::Ready(wire) => { reply.send(*wire); self.resume_machine(py, None) } - (Expect::Emitted(reply), LifecycleStep::Done) => { - reply.send(()); - self.resume_machine(py, None) + } + } + + fn on_response( + &mut self, + py: Python<'_>, + step: HookStep>, + ) -> PyResult { + match step { + HookStep::Await(awaitable, resume) => { + self.pending = Some(Pending::Response(resume)); + Ok(ExecutionStep::Await(awaitable)) } - (Expect::Response, LifecycleStep::Response(response)) => self.succeeded(py, response), - (Expect::Terminal, LifecycleStep::Done) => 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()), + HookStep::Ready(response) => self.succeeded(py, response), + } + } + + fn on_event( + &mut self, + py: Python<'_>, + step: HookStep, + next: EventNext, + ) -> PyResult { + match step { + HookStep::Await(awaitable, resume) => { + self.pending = Some(Pending::Event(resume, next)); + Ok(ExecutionStep::Await(awaitable)) + } + HookStep::Ready(()) => match next { + EventNext::Started => self.begin(py), + EventNext::Emitted(reply) => { + 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())) + } + _ => Err(missing_state()), + }, }, - _ => Err(missing_state()), } } fn begin(&mut self, py: Python<'_>) -> PyResult { let arguments = self.arguments.take().ok_or_else(missing_state)?; - match self.adapter.begin(py, arguments, self.started_at) { - Ok(step) => self.on_adapter(py, step, Expect::Arguments), - Err(error) => self.adapter_failed(py, error), + match self.hooks.prepare_arguments(py, arguments, self.started_at) { + Ok(step) => self.on_arguments(py, step), + Err(error) => self.hook_failed(py, error), } } - fn adapter_failed(&mut self, py: Python<'_>, error: PyErr) -> PyResult { + fn hook_failed(&mut self, py: Python<'_>, error: PyErr) -> PyResult { match self.stage { Stage::Begin | Stage::AfterSuccess => self.failure(py, error, FailureOrigin::Host), Stage::Call | Stage::Streaming => self.interrupt(py, error), @@ -264,22 +321,22 @@ where py: Python<'_>, interruption: Interruption, ) -> PyResult { - let step = self.resume_core(py, interruption)?; + let step = self.native.resume(py, interruption)?; self.run_steps(py, step) } fn run_steps( &mut self, py: Python<'_>, - mut step: HostStep, Py>, + mut step: NativePoll>, ) -> PyResult { loop { let result = match step { - HostStep::Suspend(awaitable) => { + NativePoll::Suspend(awaitable) => { self.pending = Some(Pending::Native); return Ok(ExecutionStep::Await(awaitable)); } - HostStep::Ready(result) => result, + NativePoll::Ready(result) => result, }; step = match self.handle_native(py, result)? { Next::Return(step) => return Ok(step), @@ -291,58 +348,54 @@ where /// Answers one machine step: performs the op it asked for, or finishes the call. fn handle_native(&mut self, py: Python<'_>, result: NativeResult) -> PyResult> { let op = match result { - Ok(MachineStep::Host(op)) => op, + Ok(MachineStep::Suspended(op)) => op, Ok(MachineStep::Complete(response)) => { return self.completed(py, response).map(Next::Return); } Err(error) => return self.machine_failed(py, error).map(Next::Return), }; let answered = match op { - HostOp::Project(reply) => { - let arguments = self.arguments.as_ref().ok_or_else(missing_state)?; - let projected = self.host.project(py, arguments.bind(py)); - answered(projected.map(|projection| reply.send(projection))) - } - HostOp::Custom(op) => answered(self.host.invoke(py, op)), - HostOp::BeforeSend { + Suspension::HostCall(op) => answered(self.binding.handle_host_call(py, op)), + Suspension::Hook(HookRequest::BeforeProviderRequest { wire, context, reply, - } => match self.adapter.before_send(py, wire, &context) { - Ok(LifecycleStep::Wire(wire)) => { + }) => match self.hooks.before_provider_request(py, wire, &context) { + Ok(HookStep::Ready(wire)) => { reply.send(*wire); Ok(Ok(())) } - Ok(LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(Expect::Wire(reply))); + Ok(HookStep::Await(awaitable, resume)) => { + self.pending = Some(Pending::Wire(resume, reply)); return Ok(Next::Return(ExecutionStep::Await(awaitable))); } - Ok(_) => return Err(missing_state()), Err(error) => Err(error), }, - HostOp::Open(head, reply) => return self.opened(py, head, reply).map(Next::Return), - HostOp::Deliver(chunk, reply) => { + Suspension::Stream(StreamDelivery::Open(head, reply)) => { + return self.opened(py, head, reply).map(Next::Return); + } + Suspension::Stream(StreamDelivery::Chunk(chunk, reply)) => { return self.delivered(py, chunk, reply).map(Next::Return); } - HostOp::Emit(event, reply) => { - match self.adapter.emit(py, LifecycleEvent::Machine(&event)) { - Ok(LifecycleStep::Done) => { + Suspension::Hook(HookRequest::Event(event, reply)) => { + match self.hooks.on_event(py, HookEvent::Machine(&event)) { + Ok(HookStep::Ready(())) => { reply.send(()); Ok(Ok(())) } - Ok(LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(Expect::Emitted(reply))); + Ok(HookStep::Await(awaitable, resume)) => { + self.pending = Some(Pending::Event(resume, EventNext::Emitted(reply))); return Ok(Next::Return(ExecutionStep::Await(awaitable))); } - Ok(_) => return Err(missing_state()), Err(error) => Err(error), } } }; match answered { - Ok(Ok(())) => self.resume_core(py, None).map(Next::Continue), + Ok(Ok(())) => self.native.resume(py, None).map(Next::Continue), Ok(Err(native)) => self - .resume_core(py, Some(HostFailure::Error(native))) + .native + .resume(py, Some(HostFailure::Error(native))) .map(Next::Continue), Err(error) => self.interrupt(py, error).map(Next::Return), } @@ -355,11 +408,11 @@ where reply: Reply, ) -> PyResult { self.stage = Stage::Streaming; - let head = match self.host.head(py, head) { + let head = match self.binding.encode_stream_head(py, head) { Ok(head) => head, Err(error) => return self.interrupt(py, error), }; - match self.adapter.opened(py) { + match self.hooks.on_stream_open(py) { Ok(()) => { self.pending = Some(Pending::Consumer(reply)); Ok(ExecutionStep::Open(head)) @@ -374,11 +427,11 @@ where chunk: as Protocol>::Chunk, reply: Reply, ) -> PyResult { - let chunk = match self.host.chunk(py, chunk) { + let chunk = match self.binding.encode_chunk(py, chunk) { Ok(chunk) => chunk, Err(error) => return self.interrupt(py, error), }; - match self.adapter.delivered(py, &chunk) { + match self.hooks.on_stream_chunk(py, &chunk) { Ok(()) => { self.pending = Some(Pending::Consumer(reply)); Ok(ExecutionStep::Yield(chunk)) @@ -399,60 +452,19 @@ where self.resume_machine(py, Some(failure)) } - fn resume_core( + fn completed( &mut self, py: Python<'_>, - interruption: Interruption, - ) -> PyResult, Py>> { - let state = Arc::clone(self.machine.as_ref().ok_or_else(missing_state)?); - let future = async move { - let mut state = state.lock().await; - let result = match interruption { - Some(failure) => state - .machine - .interrupt(failure) - .await - .map(MachineStep::Complete), - None => state.machine.resume().await, - }; - state.result = Some(result); - Ok(()) - }; - if self.asynchronous { - let mut future = Box::pin(future); - if let Poll::Ready(()) = poll_async_value(py, future.as_mut())? { - return Ok(HostStep::Ready(self.take_native_result()?)); - } - let (abort, registration) = AbortHandle::new_pair(); - self.native_abort = Some(abort); - Ok(HostStep::Suspend( - run_async_value(py, async move { - Abortable::new(future, registration) - .await - .map_err(|_| PyRuntimeError::new_err("native execution closed"))? - })? - .unbind(), - )) - } else { - run_sync_value(py, future)?; - Ok(HostStep::Ready(self.take_native_result()?)) - } - } - - fn take_native_result(&self) -> PyResult> { - self.machine - .as_ref() - .ok_or_else(missing_state)? - .try_lock() - .map_err(|_| missing_state())? - .result - .take() - .ok_or_else(missing_state) - } - - fn completed(&mut self, py: Python<'_>, response: ResponseOf) -> PyResult { + response: HostedCompletion>, + ) -> PyResult { self.ended_at = Some(epoch_seconds()); - let public = match self.host.complete(py, response) { + let response = match response { + HostedCompletion::Complete(response) => response, + HostedCompletion::StreamEnded | HostedCompletion::Detached => { + return self.succeeded(py, py.None()); + } + }; + let public = match self.binding.encode_response(py, response) { Ok(public) => public, Err(error) => return self.failure(py, error, FailureOrigin::Call), }; @@ -460,8 +472,8 @@ where return self.succeeded(py, public); } self.stage = Stage::AfterSuccess; - match self.adapter.after_success(py, public, self.timing()) { - Ok(step) => self.on_adapter(py, step, Expect::Response), + match self.hooks.transform_response(py, public, self.timing()) { + Ok(step) => self.on_response(py, step), Err(error) => self.failure(py, error, FailureOrigin::Host), } } @@ -479,7 +491,7 @@ where /// fails, that failure is raised with the native error's text as its `__context__`. fn classified(&self, py: Python<'_>, error: ErrorOf) -> PyErr { let native = error.to_string(); - let classifier_error = match self.host.classify(py, error) { + let classifier_error = match self.binding.map_error(py, error) { Ok(failure) => return failure.into(), Err(classifier_error) => classifier_error, }; @@ -488,13 +500,13 @@ where } fn succeeded(&mut self, py: Python<'_>, response: Py) -> PyResult { - let event = LifecycleEvent::Succeeded { + let event = HookEvent::Succeeded { timing: self.timing(), response: &response, }; - let step = self.adapter.emit(py, event)?; + let step = self.hooks.on_event(py, event)?; self.stage = Stage::Succeeded(response); - self.on_adapter(py, step, Expect::Terminal) + self.on_event(py, step, EventNext::Terminal) } fn failure( @@ -507,41 +519,43 @@ where if is_cancellation(py, &error) { return Err(error); } - let event = LifecycleEvent::Failed { + let event = HookEvent::Failed { timing: self.timing(), origin, error: &error, }; - let step = self.adapter.emit(py, event)?; + let step = self.hooks.on_event(py, event)?; self.stage = Stage::Failed(error.into_value(py)); - self.on_adapter(py, step, Expect::Terminal) + self.on_event(py, step, EventNext::Terminal) } fn clear(&mut self) { - if let Some(abort) = self.native_abort.take() { - abort.abort(); - } - if self.machine.take().is_some() { + if !self.closed { + self.closed = true; + self.native.close(); + self.start = None; Python::attach(|py| { - self.adapter.close(py); - self.host.close(py); + self.hooks.close(py); + self.binding.close(py); }); } } } -impl ExecutionBody for PythonDriver +impl ExecutionBody for PythonDriver where - H: ProtocolHost, - M: Machine> + 'static, + L: PythonCallHooks + 'static, + H: PythonBinding + PythonHostCalls, + M: Machine + 'static, + M::Complete: Into>>, { fn resume(&mut self, result: Option>>) -> PyResult { Python::attach(|py| self.drive(py, result)) } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - self.host.traverse(visit)?; - self.adapter.traverse(visit)?; + self.binding.traverse(visit)?; + self.hooks.traverse(visit)?; visit.call(&self.arguments)?; visit.call(&self.interrupted)?; match &self.stage { @@ -552,10 +566,12 @@ where } } -impl Drop for PythonDriver +impl Drop for PythonDriver where - H: ProtocolHost, - M: Machine> + 'static, + L: PythonCallHooks + 'static, + H: PythonBinding + PythonHostCalls, + M: Machine + 'static, + M::Complete: Into>>, { fn drop(&mut self) { self.clear(); @@ -572,6 +588,8 @@ mod tests { use pyo3::types::PyDict; use super::*; + use crate::PythonOwned; + use litellm_host::hooks::RouteHooks; static PYTHON_GLOBALS: Mutex<()> = Mutex::new(()); @@ -622,8 +640,8 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri impl Protocol for Synthetic { type Response = String; type Error = Error; - type Projection = String; - type Op = (&'static str, Reply); + type Request = String; + type HostCall = (&'static str, Reply); type Chunk = std::convert::Infallible; type StreamHead = std::convert::Infallible; } @@ -664,9 +682,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri Answer, RaisePython, RejectNatively, + RejectRequestNatively, + RaiseRequestPython, } - struct SyntheticHost { + struct SyntheticBinding { log: Log, op: OpScript, classifier_fails: bool, @@ -683,55 +703,62 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - impl SyntheticHost { + impl SyntheticBinding { fn answer(&self, value: impl FnOnce() -> String) -> Result> { match self.op { OpScript::Answer => Ok(value()), - OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()), - OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))), + OpScript::RaisePython | OpScript::RaiseRequestPython => { + Err(PyValueError::new_err("op failed").into()) + } + OpScript::RejectNatively | OpScript::RejectRequestNatively => { + Err(InvokeError::Native(Error("op rejected".into()))) + } } } } - impl ProtocolHost for SyntheticHost { + impl PythonBinding for SyntheticBinding { type Protocol = Synthetic; type Failure = Classified; - fn project( + fn decode_request( &mut self, _: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> Result> { self.log.push("project"); - self.answer(|| format!("project:{}", arguments.len())) + match self.op { + OpScript::RejectRequestNatively | OpScript::RaiseRequestPython => { + self.answer(String::new) + } + _ => Ok(format!("project:{}", arguments.len())), + } } - fn invoke( + fn encode_stream_head( &mut self, _: Python<'_>, - (op, reply): (&'static str, Reply), - ) -> Result<(), InvokeError> { - self.log.push(format!("op:{op}")); - self.answer(|| op.to_string()) - .map(|answer| reply.send(answer)) - } - - fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult> { + head: std::convert::Infallible, + ) -> PyResult> { match head {} } - fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { + fn encode_chunk( + &mut self, + _: Python<'_>, + chunk: std::convert::Infallible, + ) -> PyResult> { match chunk {} } - fn complete(&mut self, py: Python<'_>, response: String) -> PyResult> { + fn encode_response(&mut self, py: Python<'_>, response: String) -> PyResult> { self.log.push("complete"); Ok(pyo3::types::PyString::new(py, &response) .into_any() .unbind()) } - fn classify(&self, _: Python<'_>, error: Error) -> PyResult { + fn map_error(&self, _: Python<'_>, error: Error) -> PyResult { self.log.push(format!("classify:{error}")); if self.classifier_fails { return Err(pyo3::exceptions::PyTypeError::new_err("classifier failed")); @@ -742,110 +769,120 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn host_error(error: &PyErr) -> Error { Error(error.to_string()) } + } + impl PythonHostCalls for SyntheticBinding { + fn handle_host_call( + &mut self, + _: Python<'_>, + (op, reply): (&'static str, Reply), + ) -> Result<(), InvokeError> { + self.log.push(format!("op:{op}")); + self.answer(|| op.to_string()) + .map(|answer| reply.send(answer)) + } + } + + impl PythonOwned for SyntheticBinding { fn close(&mut self, _: Python<'_>) { self.log.push("host.close"); } - fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { Ok(()) } } #[derive(Clone, Copy)] - enum AdapterScript { + enum HookScript { Plain, FailBegin, ReplaceResponse, FailAfterSuccess, } - struct SyntheticAdapter { + struct SyntheticHooks { log: Log, - script: AdapterScript, + script: HookScript, } - impl PythonLifecycle for SyntheticAdapter { - fn begin( + impl PythonCallHooks for SyntheticHooks { + fn prepare_arguments( &mut self, _: Python<'_>, arguments: Py, _: f64, - ) -> PyResult { + ) -> PyResult>> { self.log.push("begin"); - if matches!(self.script, AdapterScript::FailBegin) { + if matches!(self.script, HookScript::FailBegin) { return Err(PyValueError::new_err("begin failed")); } - Ok(LifecycleStep::Arguments(arguments)) + Ok(HookStep::Ready(arguments)) } - fn before_send( + fn before_provider_request( &mut self, _: Python<'_>, wire: Box, _: &RequestContext, - ) -> PyResult { - self.log.push("before_send"); - Ok(LifecycleStep::Wire(Box::new(WireRequest { + ) -> PyResult>> { + self.log.push("before_provider_request"); + Ok(HookStep::Ready(Box::new(WireRequest { url: "rewritten".into(), ..*wire }))) } - fn after_success( + fn transform_response( &mut self, py: Python<'_>, response: Py, _: Timing, - ) -> PyResult { + ) -> PyResult>> { self.log.push("after_success"); match self.script { - AdapterScript::ReplaceResponse => Ok(LifecycleStep::Response( + HookScript::ReplaceResponse => Ok(HookStep::Ready( "replaced".into_pyobject(py)?.into_any().unbind(), )), - AdapterScript::FailAfterSuccess => { - Err(PyValueError::new_err("after_success failed")) - } - AdapterScript::Plain | AdapterScript::FailBegin => { - Ok(LifecycleStep::Response(response)) - } + HookScript::FailAfterSuccess => Err(PyValueError::new_err("after_success failed")), + HookScript::Plain | HookScript::FailBegin => Ok(HookStep::Ready(response)), } } - fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult { + fn on_event( + &mut self, + py: Python<'_>, + event: HookEvent<'_>, + ) -> PyResult> { self.log.push(match event { - LifecycleEvent::Started { .. } => "started".into(), - LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => { + HookEvent::Started { .. } => "started".into(), + HookEvent::Machine(MachineEvent::ResponseReceived { raw }) => { format!("response:{}", raw.body) } - LifecycleEvent::Succeeded { response, .. } => { + HookEvent::Succeeded { response, .. } => { format!("succeeded:{}", response.bind(py)) } - LifecycleEvent::Failed { origin, error, .. } => { + HookEvent::Failed { origin, error, .. } => { format!("failed:{origin:?}:{}", error.value(py)) } }); - Ok(LifecycleStep::Done) + Ok(HookStep::Ready(())) } - fn opened(&mut self, _: Python<'_>) -> PyResult<()> { + fn on_stream_open(&mut self, _: Python<'_>) -> PyResult<()> { self.log.push("opened"); Ok(()) } - fn delivered(&mut self, _: Python<'_>, _: &Py) -> PyResult<()> { + fn on_stream_chunk(&mut self, _: Python<'_>, _: &Py) -> PyResult<()> { self.log.push("delivered"); Ok(()) } + } - fn resume(&mut self, _: Python<'_>, _: PyResult>) -> PyResult { - Err(missing_state()) - } - + impl PythonOwned for SyntheticHooks { fn close(&mut self, _: Python<'_>) { self.log.push("adapter.close"); } - fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { Ok(()) } @@ -853,15 +890,15 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_scripted( py: Python<'_>, - machine: CallMachine, + machine: impl FnOnce(String) -> CallMachine + Send + Sync + 'static, op: OpScript, - script: AdapterScript, + script: HookScript, asynchronous: bool, ) -> (PyResult>, Vec) { run_hosted( py, machine, - SyntheticHost { + SyntheticBinding { log: Log::default(), op, classifier_fails: false, @@ -873,9 +910,9 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_hosted( py: Python<'_>, - machine: CallMachine, - host: SyntheticHost, - script: AdapterScript, + machine: impl FnOnce(String) -> CallMachine + Send + Sync + 'static, + host: SyntheticBinding, + script: HookScript, asynchronous: bool, ) -> (PyResult>, Vec) { run_preflighted(py, machine, host, script, no_preflight, asynchronous) @@ -887,14 +924,14 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_preflighted( py: Python<'_>, - machine: CallMachine, - host: SyntheticHost, - script: AdapterScript, + machine: impl FnOnce(String) -> CallMachine + Send + Sync + 'static, + host: SyntheticBinding, + script: HookScript, preflight: Preflight, asynchronous: bool, ) -> (PyResult>, Vec) { let log = Log(host.log.0.clone()); - let adapter = SyntheticAdapter { + let adapter = SyntheticHooks { log: Log(log.0.clone()), script, }; @@ -904,7 +941,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py, machine, host, - Box::new(adapter), + adapter, preflight, arguments.unbind(), asynchronous, @@ -925,24 +962,83 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri (result, log.entries()) } - /// Answers to projection, to the route op and to `before_send` all reach the + /// Answers to projection, to the route op and to `before_provider_request` all reach the /// response, so a driver that misroutes a reply changes what the call returns. - fn success_machine() -> CallMachine { - CallMachine::new(|host| { - Box::pin(async move { - let projected = host.project().await?; - let signed = host.custom_op(|reply| ("sign", reply)).await?; - let wire = host.before_send(wire(), context()).await?; - host.emit(MachineEvent::ResponseReceived { - raw: RawResponse { body: "raw".into() }, + fn success_machine() -> impl FnOnce(String) -> CallMachine + Send + Sync { + move |projected| { + CallMachine::new(move |host| { + Box::pin(async move { + let signed = host.services.call(|reply| ("sign", reply)).await?; + let wire = host + .hooks + .before_provider_request(wire(), context()) + .await?; + host.hooks + .on_event(MachineEvent::ResponseReceived { + raw: RawResponse { body: "raw".into() }, + }) + .await?; + Ok(format!("{projected}|{signed}|{}", wire.url)) }) - .await?; - Ok(format!("{projected}|{signed}|{}", wire.url)) }) - }) + } } - #[test] + #[rstest::rstest] + #[case::preparation(HookScript::FailBegin, OpScript::Answer)] + #[case::native_decode(HookScript::Plain, OpScript::RejectRequestNatively)] + #[case::python_decode(HookScript::Plain, OpScript::RaiseRequestPython)] + fn startup_failures_release_the_factory_without_constructing_a_machine( + #[case] script: HookScript, + #[case] op: OpScript, + #[values(false, true)] asynchronous: bool, + ) { + use std::sync::atomic::{AtomicBool, Ordering}; + struct Release(Arc); + impl Drop for Release { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + let constructed = Arc::new(AtomicBool::new(false)); + let did_construct = constructed.clone(); + let released = Arc::new(AtomicBool::new(false)); + let release = Release(released.clone()); + let (result, log) = run_scripted( + py, + move |request| { + let _release = release; + did_construct.store(true, Ordering::SeqCst); + success_machine()(request) + }, + op, + script, + asynchronous, + ); + assert!(result.is_err()); + assert!(!constructed.load(Ordering::SeqCst)); + assert!(released.load(Ordering::SeqCst)); + assert_eq!( + log.iter() + .filter(|event| event.starts_with("failed:")) + .count(), + 1 + ); + assert_eq!( + log.iter().filter(|event| *event == "adapter.close").count(), + 1 + ); + assert_eq!(log.iter().filter(|event| *event == "host.close").count(), 1); + }); + } + + #[rstest::rstest] fn success_runs_every_step_in_order_and_returns_the_public_response() { let _guard = PYTHON_GLOBALS .lock() @@ -955,7 +1051,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py, success_machine(), OpScript::Answer, - AdapterScript::Plain, + HookScript::Plain, asynchronous, ); assert_eq!( @@ -969,7 +1065,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "begin", "project", "op:sign", - "before_send", + "before_provider_request", "response:raw", "complete", "after_success", @@ -987,19 +1083,19 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri impl Protocol for Streaming { type Response = (); type Error = Error; - type Projection = (); - type Op = std::convert::Infallible; + type Request = (); + type HostCall = std::convert::Infallible; type Chunk = &'static str; type StreamHead = Vec<(&'static str, &'static str)>; } - struct StreamingHost; + struct StreamingBinding; - impl ProtocolHost for StreamingHost { + impl PythonBinding for StreamingBinding { type Protocol = Streaming; type Failure = Classified; - fn project( + fn decode_request( &mut self, _: Python<'_>, _: &Bound<'_, PyDict>, @@ -1007,15 +1103,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri Ok(()) } - fn invoke( - &mut self, - _: Python<'_>, - op: std::convert::Infallible, - ) -> Result<(), InvokeError> { - match op {} - } - - fn head( + fn encode_stream_head( &mut self, py: Python<'_>, head: Vec<(&'static str, &'static str)>, @@ -1029,44 +1117,117 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri Ok(hidden.into_any().unbind()) } - fn chunk(&mut self, py: Python<'_>, chunk: &'static str) -> PyResult> { + fn encode_chunk(&mut self, py: Python<'_>, chunk: &'static str) -> PyResult> { Ok(pyo3::types::PyString::new(py, chunk).into_any().unbind()) } - fn complete(&mut self, py: Python<'_>, (): ()) -> PyResult> { + fn encode_response(&mut self, py: Python<'_>, (): ()) -> PyResult> { Ok(py.None()) } - fn classify(&self, _: Python<'_>, error: Error) -> PyResult { + fn map_error(&self, _: Python<'_>, error: Error) -> PyResult { Ok(Classified(error.0)) } fn host_error(error: &PyErr) -> Error { Error(error.to_string()) } + } + impl PythonHostCalls for StreamingBinding { + fn handle_host_call( + &mut self, + _: Python<'_>, + op: std::convert::Infallible, + ) -> Result<(), InvokeError> { + match op {} + } + } + + impl PythonOwned for StreamingBinding { fn close(&mut self, _: Python<'_>) {} - fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { Ok(()) } } - fn streaming_machine() -> CallMachine { - CallMachine::new(|host| { - Box::pin(async move { - host.project().await?; - if host.open(vec![("request-id", "req_1")]).await? == Demand::Detached { - return Ok(()); - } - for chunk in ["first", "second"] { - if host.deliver(chunk).await? == Demand::Detached { - break; - } - } - Ok(()) + fn streaming_machine() + -> impl FnOnce(()) -> litellm_host::call::HostedMachine + Send + Sync { + move |()| { + litellm_host::call::hosted_call((), |(), _, _| async { + Ok(litellm_host::call::CallOutput::Stream { + head: vec![("request-id", "req_1")], + chunks: Box::pin(futures_util::stream::iter([Ok("first"), Ok("second")])), + }) }) - }) + } + } + + #[rstest::rstest] + #[case::sync(false)] + #[case::asynchronous(true)] + fn explicitly_closing_a_stream_dispatches_success_for_delivered_chunks( + #[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 handed = run_call( + py, + streaming_machine(), + StreamingBinding, + SyntheticHooks { + log: Log(log.0.clone()), + script: HookScript::Plain, + }, + no_preflight, + PyDict::new(py).unbind(), + asynchronous, + ) + .unwrap(); + let stream = if asynchronous { + let stop = handed.call_method1(py, "send", (py.None(),)).unwrap_err(); + stop.value(py).getattr("value").unwrap() + } else { + handed.into_bound(py) + }; + if asynchronous { + for method in ["__anext__", "aclose", "aclose"] { + let stop = stream + .call_method0(method) + .unwrap() + .call_method1("send", (py.None(),)) + .unwrap_err(); + assert!(stop.is_instance_of::(py)); + } + } else { + assert_eq!( + stream + .call_method0("__next__") + .unwrap() + .extract::() + .unwrap(), + "first" + ); + stream.call_method0("close").unwrap(); + stream.call_method0("close").unwrap(); + } + assert_eq!( + log.entries(), + [ + "started", + "begin", + "opened", + "delivered", + "succeeded:None", + "adapter.close" + ] + ); + }); } /// Drives a `Stream` (async) or `SyncStream` to completion from a sync test. @@ -1093,7 +1254,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri .collect() } - #[test] + #[rstest::rstest] fn a_stream_carries_its_head_as_hidden_params_before_the_first_chunk() { let _guard = PYTHON_GLOBALS .lock() @@ -1103,15 +1264,15 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri install_lifecycle_module(py); for asynchronous in [false, true] { let log = Log::default(); - let adapter = SyntheticAdapter { + let adapter = SyntheticHooks { log: Log(log.0.clone()), - script: AdapterScript::Plain, + script: HookScript::Plain, }; let handed = run_call( py, streaming_machine(), - StreamingHost, - Box::new(adapter), + StreamingBinding, + adapter, no_preflight, PyDict::new(py).unbind(), asynchronous, @@ -1140,16 +1301,13 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - fn failing_machine() -> CallMachine { - CallMachine::new(|host| { - Box::pin(async move { - host.project().await?; - Err(Error("provider exploded".into())) - }) - }) + fn failing_machine() -> impl FnOnce(String) -> CallMachine + Send + Sync { + move |_| { + CallMachine::new(|_| Box::pin(async move { Err(Error("provider exploded".into())) })) + } } - #[test] + #[rstest::rstest] fn a_native_failure_is_classified_once_and_reported_classified() { let _guard = PYTHON_GLOBALS .lock() @@ -1162,7 +1320,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py, failing_machine(), OpScript::Answer, - AdapterScript::Plain, + HookScript::Plain, asynchronous, ); let error = result.unwrap_err(); @@ -1184,7 +1342,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - #[test] + #[rstest::rstest] fn a_native_rejection_from_a_host_operation_is_classified_once() { let _guard = PYTHON_GLOBALS .lock() @@ -1195,7 +1353,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py, success_machine(), OpScript::RejectNatively, - AdapterScript::Plain, + HookScript::Plain, false, ); assert_eq!( @@ -1208,6 +1366,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "started", "begin", "project", + "op:sign", "classify:op rejected", "failed:Call:classified: op rejected", "adapter.close", @@ -1217,7 +1376,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - #[test] + #[rstest::rstest] fn a_python_exception_from_a_host_operation_is_reported_as_raised() { let _guard = PYTHON_GLOBALS .lock() @@ -1228,7 +1387,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py, success_machine(), OpScript::RaisePython, - AdapterScript::Plain, + HookScript::Plain, false, ); let error = result.unwrap_err(); @@ -1240,6 +1399,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "started", "begin", "project", + "op:sign", "failed:Call:op failed", "adapter.close", "host.close", @@ -1248,7 +1408,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - #[test] + #[rstest::rstest] fn a_failing_classifier_surfaces_with_the_native_error_as_context() { let _guard = PYTHON_GLOBALS .lock() @@ -1258,12 +1418,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let (result, log) = run_hosted( py, failing_machine(), - SyntheticHost { + SyntheticBinding { log: Log::default(), op: OpScript::Answer, classifier_fails: true, }, - AdapterScript::Plain, + HookScript::Plain, false, ); let error = result.unwrap_err(); @@ -1287,7 +1447,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - #[test] + #[rstest::rstest] fn begin_failures_are_host_failures_without_provider_mapping() { let _guard = PYTHON_GLOBALS .lock() @@ -1298,7 +1458,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py, success_machine(), OpScript::Answer, - AdapterScript::FailBegin, + HookScript::FailBegin, false, ); let error = result.unwrap_err(); @@ -1330,7 +1490,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri arguments.set_item("api_key", "inherited") } - #[test] + #[rstest::rstest] fn a_preflight_rejection_is_the_callers_error_and_the_machine_never_starts() { let _guard = PYTHON_GLOBALS .lock() @@ -1342,12 +1502,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let (result, log) = run_preflighted( py, success_machine(), - SyntheticHost { + SyntheticBinding { log: Log::default(), op: OpScript::Answer, classifier_fails: false, }, - AdapterScript::Plain, + HookScript::Plain, rejecting_preflight, asynchronous, ); @@ -1368,7 +1528,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - #[test] + #[rstest::rstest] fn the_host_projects_from_the_keyword_view_the_preflight_rewrote() { let _guard = PYTHON_GLOBALS .lock() @@ -1380,12 +1540,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let (result, _) = run_preflighted( py, success_machine(), - SyntheticHost { + SyntheticBinding { log: Log::default(), op: OpScript::Answer, classifier_fails: false, }, - AdapterScript::Plain, + HookScript::Plain, inheriting_preflight, asynchronous, ); @@ -1397,7 +1557,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - #[test] + #[rstest::rstest] fn the_adapters_finalized_response_is_what_the_call_returns_and_reports() { let _guard = PYTHON_GLOBALS .lock() @@ -1410,7 +1570,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py, success_machine(), OpScript::Answer, - AdapterScript::ReplaceResponse, + HookScript::ReplaceResponse, asynchronous, ); assert_eq!(result.unwrap().extract::(py).unwrap(), "replaced"); @@ -1420,7 +1580,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - #[test] + #[rstest::rstest] fn a_failure_while_finalizing_fails_the_call_instead_of_succeeding() { let _guard = PYTHON_GLOBALS .lock() @@ -1433,7 +1593,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py, success_machine(), OpScript::Answer, - AdapterScript::FailAfterSuccess, + HookScript::FailAfterSuccess, asynchronous, ); let error = result.unwrap_err(); @@ -1452,7 +1612,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - #[test] + #[rstest::rstest] fn cancellation_ends_the_call_without_terminal_dispatch() { let _guard = PYTHON_GLOBALS .lock() @@ -1460,10 +1620,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri crate::initialize_python(); Python::attach(|py| { struct Cancelling(Log); - impl ProtocolHost for Cancelling { + impl PythonBinding for Cancelling { type Protocol = Synthetic; type Failure = Classified; - fn project( + fn decode_request( &mut self, _: Python<'_>, _: &Bound<'_, PyDict>, @@ -1471,37 +1631,44 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri self.0.push("project"); Err(pyo3::exceptions::asyncio::CancelledError::new_err(()).into()) } - fn invoke( - &mut self, - _: Python<'_>, - _: (&'static str, Reply), - ) -> Result<(), InvokeError> { - Err(missing_state().into()) - } - fn head( + + fn encode_stream_head( &mut self, _: Python<'_>, head: std::convert::Infallible, ) -> PyResult> { match head {} } - fn chunk( + fn encode_chunk( &mut self, _: Python<'_>, chunk: std::convert::Infallible, ) -> PyResult> { match chunk {} } - fn complete(&mut self, _: Python<'_>, _: String) -> PyResult> { + fn encode_response(&mut self, _: Python<'_>, _: String) -> PyResult> { Err(missing_state()) } - fn classify(&self, _: Python<'_>, error: Error) -> PyResult { + fn map_error(&self, _: Python<'_>, error: Error) -> PyResult { self.0.push("classify"); Ok(Classified(error.0)) } fn host_error(error: &PyErr) -> Error { Error(error.to_string()) } + } + + impl PythonHostCalls for Cancelling { + fn handle_host_call( + &mut self, + _: Python<'_>, + _: (&'static str, Reply), + ) -> Result<(), InvokeError> { + Err(missing_state().into()) + } + } + + impl PythonOwned for Cancelling { fn close(&mut self, _: Python<'_>) {} fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { Ok(()) @@ -1509,15 +1676,15 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } let log = Log::default(); let host = Cancelling(Log(log.0.clone())); - let adapter = SyntheticAdapter { + let adapter = SyntheticHooks { log: Log(log.0.clone()), - script: AdapterScript::Plain, + script: HookScript::Plain, }; let error = run_call( py, success_machine(), host, - Box::new(adapter), + adapter, no_preflight, PyDict::new(py).unbind(), false, @@ -1531,7 +1698,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } - #[test] + #[rstest::rstest] fn python_driver_preserves_inline_await_and_native_ownership() { let _guard = PYTHON_GLOBALS .lock() @@ -1644,7 +1811,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri Execution::new(ErrorBody(Some(error.unbind()))) } - #[test] + #[rstest::rstest] fn retained_exception_frames_are_collectable() { crate::initialize_python(); Python::attach(|py| { @@ -1684,7 +1851,7 @@ assert reference() is None }); } - #[test] + #[rstest::rstest] fn coroutine_collects_cycles_retained_by_bridge_host() { crate::initialize_python(); Python::attach(|py| { diff --git a/litellm-rust/crates/host-python/src/error.rs b/litellm-rust/crates/host-python/src/error.rs new file mode 100644 index 00000000000..d6e186860d7 --- /dev/null +++ b/litellm-rust/crates/host-python/src/error.rs @@ -0,0 +1,20 @@ +use pyo3::prelude::*; + +/// Why a custom operation the host answered did not produce a result: the route's own code +/// rejected it, which the route classifies like any other native failure, or Python code +/// raised, which reaches the caller as it was raised. +#[derive(Debug)] +pub enum InvokeError { + Native(E), + Python(PyErr), +} + +impl From for InvokeError { + fn from(error: PyErr) -> Self { + Self::Python(error) + } +} + +pub fn missing_state() -> PyErr { + pyo3::exceptions::PyRuntimeError::new_err("missing native call state") +} diff --git a/litellm-rust/crates/host-python/src/handle.rs b/litellm-rust/crates/host-python/src/handle.rs index 24adfd404d7..91d003f32b9 100644 --- a/litellm-rust/crates/host-python/src/handle.rs +++ b/litellm-rust/crates/host-python/src/handle.rs @@ -31,6 +31,10 @@ pub struct Execution { state: ExecutionState, } +fn lifecycle(py: Python<'_>) -> PyResult> { + py.import("litellm.rust_bridge.lifecycle") +} + impl Execution { pub fn new(body: impl ExecutionBody + 'static) -> Self { Self { @@ -38,6 +42,21 @@ impl Execution { } } + pub fn into_coroutine(self, py: Python<'_>) -> PyResult> { + let execution = Py::new(py, self)?; + lifecycle(py)?.getattr("drive")?.call1((execution,)) + } + + pub(crate) fn into_sync_stream( + self, + py: Python<'_>, + head: Py, + ) -> PyResult> { + lifecycle(py)? + .getattr("SyncStream")? + .call1((Py::new(py, self)?, head)) + } + /// An execution already started elsewhere and now waiting for its next input. pub fn suspended(body: impl ExecutionBody + 'static) -> Self { Self { @@ -79,11 +98,7 @@ impl Execution { ExecutionStep::Yield(value) => ("Yield", value, true), ExecutionStep::Return(value) => ("Complete", value, false), }; - let step = py - .import("litellm.rust_bridge.lifecycle")? - .getattr(tag)? - .call1((value,))? - .unbind(); + let step = lifecycle(py)?.getattr(tag)?.call1((value,))?.unbind(); Ok((step, suspended)) })) .map_err(panic_to_pyerr) diff --git a/litellm-rust/crates/host-python/src/hooks.rs b/litellm-rust/crates/host-python/src/hooks.rs new file mode 100644 index 00000000000..5a0f95d9c89 --- /dev/null +++ b/litellm-rust/crates/host-python/src/hooks.rs @@ -0,0 +1,79 @@ +use crate::PythonOwned; +use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest}; +use pyo3::prelude::*; +use pyo3::types::PyDict; + +/// The SDK's request policy, run by the driver on the keyword view `prepare_arguments` returned and +/// before the binding decodes from it. It rewrites that view in place, so the +/// hooks that returned it see the rewrite too; a rejection fails the call as a host +/// failure, so the hooks still observe it. +pub type Preflight = fn(Python<'_>, &Bound<'_, PyDict>) -> PyResult<()>; + +/// What a hook step produced: either the value the driver asked for, or a Python +/// awaitable the driver hands back to the caller's task before asking again. +pub type HookResume = fn(&mut L, Python<'_>, PyResult>) -> PyResult>; + +pub enum HookStep { + Await(Py, HookResume), + 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, + }, +} + +/// 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>>; + + fn before_provider_request( + &mut self, + py: Python<'_>, + wire: Box, + context: &RequestContext, + ) -> PyResult>>; + + 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<()>; +} diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 2f9e37fe968..cbb7fe3507b 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -1,38 +1,44 @@ //! The CPython runtime adapter: value marshalling, interpreter detachment, the tokio and //! asyncio glue, and the driver that runs a native [`Machine`](litellm_host::machine::Machine) -//! against a Python protocol host and a Python lifecycle. Everything here is Python-specific by +//! against a Python binding, host services and active call hooks. Everything here is Python-specific by //! construction; another host language gets its own crate of the same shape. -mod adapter; mod argument; +mod binding; mod callable; mod driver; -mod execution; +mod error; mod file_reader; mod fork_gate; mod gil; mod handle; +mod hooks; mod marshal; +mod native; +mod owned; +mod runtime; +mod services; -pub use adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, Preflight, ProtocolHost, PythonLifecycle, - missing_state, -}; pub use argument::lookup; +pub use binding::PythonBinding; pub use callable::wrap_failure; pub use driver::run_call; -pub use execution::{ - ForkedAfterNativeRuntimeStarted, ProcessReservedForForking, enter_native, poll_async_value, - reserve_process_for_forking, run_async, run_async_value, run_sync, run_sync_value, - runtime_started, -}; +pub use error::{InvokeError, missing_state}; pub use file_reader::{FileContent, PythonFileReader, py_bytes}; pub use fork_gate::RuntimeAlreadyStarted; pub use gil::{PythonContext, attach_blocking, release_count, release_gil}; pub use handle::{Execution, ExecutionBody, ExecutionStep}; +pub use hooks::{HookEvent, HookResume, HookStep, Preflight, PythonCallHooks}; pub use marshal::{ Pythonized, from_py, from_py_argument, json_loads, json_object_field, panic_to_pyerr, to_py, }; +pub use owned::PythonOwned; +pub use runtime::{ + ForkedAfterNativeRuntimeStarted, ProcessReservedForForking, enter_native, poll_async_value, + reserve_process_for_forking, run_async, run_async_value, run_sync, run_sync_value, + runtime_started, +}; +pub use services::PythonHostCalls; /// Starts the interpreter and imports the standard modules the tests share, once, so /// parallel test threads never race a first import of `asyncio`. diff --git a/litellm-rust/crates/host-python/src/native.rs b/litellm-rust/crates/host-python/src/native.rs new file mode 100644 index 00000000000..65d2ffc2c7a --- /dev/null +++ b/litellm-rust/crates/host-python/src/native.rs @@ -0,0 +1,134 @@ +use std::sync::Arc; +use std::task::Poll; + +use futures_util::future::{AbortHandle, Abortable}; +use litellm_host::call::HostedCompletion; +use litellm_host::machine::{HostFailure, Machine, MachineStep}; +use litellm_host::protocol::Protocol; +use pyo3::exceptions::PyRuntimeError; +use pyo3::prelude::*; +use tokio::sync::Mutex; + +use crate::missing_state; +use crate::runtime::{poll_async_value, run_async_value, run_sync_value}; + +type NativeResult = Result< + MachineStep< + ::Protocol, + HostedCompletion<<::Protocol as Protocol>::Response>, + >, + <::Protocol as Protocol>::Error, +>; + +type MachineResult = Result< + MachineStep<::Protocol, ::Complete>, + <::Protocol as Protocol>::Error, +>; + +struct MachineState { + machine: M, + result: Option>, +} + +pub(super) enum NativePoll { + Ready(T), + Suspend(Py), +} + +pub(super) struct NativeMachine { + state: Option>>>, + abort: Option, + asynchronous: bool, +} + +impl NativeMachine +where + M::Complete: Into::Response>>, +{ + pub(super) fn new(asynchronous: bool) -> Self { + Self { + state: None, + abort: None, + asynchronous, + } + } + + pub(super) fn start(&mut self, machine: M) { + self.state = Some(Arc::new(Mutex::new(MachineState { + machine, + result: None, + }))); + } + + pub(super) fn resume( + &mut self, + py: Python<'_>, + interruption: Option::Error>>, + ) -> PyResult>> { + let state = Arc::clone(self.state.as_ref().ok_or_else(missing_state)?); + let future = async move { + let mut state = state.lock().await; + let result = match interruption { + Some(failure) => state + .machine + .interrupt(failure) + .await + .map(MachineStep::Complete), + None => state.machine.resume().await, + }; + state.result = Some(result); + Ok(()) + }; + if self.asynchronous { + let mut future = Box::pin(future); + if let Poll::Ready(()) = poll_async_value(py, future.as_mut())? { + return Ok(NativePoll::Ready(self.take_result()?)); + } + let (abort, registration) = AbortHandle::new_pair(); + self.abort = Some(abort); + Ok(NativePoll::Suspend( + run_async_value(py, async move { + Abortable::new(future, registration) + .await + .map_err(|_| PyRuntimeError::new_err("native execution closed"))? + })? + .unbind(), + )) + } else { + run_sync_value(py, future)?; + Ok(NativePoll::Ready(self.take_result()?)) + } + } + + pub(super) fn take_result(&self) -> PyResult> { + self.state + .as_ref() + .ok_or_else(missing_state)? + .try_lock() + .map_err(|_| missing_state())? + .result + .take() + .ok_or_else(missing_state) + .map(|result| { + result.map(|step| match step { + MachineStep::Suspended(op) => MachineStep::Suspended(op), + MachineStep::Complete(response) => MachineStep::Complete(response.into()), + }) + }) + } +} + +impl NativeMachine { + pub(super) fn close(&mut self) { + if let Some(abort) = self.abort.take() { + abort.abort(); + } + self.state = None; + } +} + +impl Drop for NativeMachine { + fn drop(&mut self) { + self.close(); + } +} diff --git a/litellm-rust/crates/host-python/src/owned.rs b/litellm-rust/crates/host-python/src/owned.rs new file mode 100644 index 00000000000..aa0b9a23300 --- /dev/null +++ b/litellm-rust/crates/host-python/src/owned.rs @@ -0,0 +1,9 @@ +use pyo3::{ + gc::{PyTraverseError, PyVisit}, + prelude::*, +}; + +pub trait PythonOwned: Send + Sync { + fn close(&mut self, py: Python<'_>); + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; +} diff --git a/litellm-rust/crates/host-python/src/execution.rs b/litellm-rust/crates/host-python/src/runtime.rs similarity index 100% rename from litellm-rust/crates/host-python/src/execution.rs rename to litellm-rust/crates/host-python/src/runtime.rs diff --git a/litellm-rust/crates/host-python/src/services.rs b/litellm-rust/crates/host-python/src/services.rs new file mode 100644 index 00000000000..0a06118926e --- /dev/null +++ b/litellm-rust/crates/host-python/src/services.rs @@ -0,0 +1,11 @@ +use crate::{InvokeError, PythonOwned}; +use litellm_host::protocol::Protocol; +use pyo3::prelude::*; + +pub trait PythonHostCalls: PythonOwned { + fn handle_host_call( + &mut self, + py: Python<'_>, + call: P::HostCall, + ) -> Result<(), InvokeError>; +} diff --git a/litellm-rust/crates/host/AGENTS.md b/litellm-rust/crates/host/AGENTS.md new file mode 100644 index 00000000000..47f3a15575b --- /dev/null +++ b/litellm-rust/crates/host/AGENTS.md @@ -0,0 +1,19 @@ +`litellm-host` defines typed calls, execution hooks, host services, stream delivery and the resumable machine. HTTP and Python drivers interpret the same suspension protocol in their own runtimes + +| Responsibility | Shared contract | HTTP | Python | +| --- | --- | --- | --- | +| Input and output conversion | `Protocol::{Request, Response, StreamHead, Chunk, Error}` | Typed input; `ResponseEncoder` and `StreamEncoder` produce HTTP values | `PythonBinding` decodes prepared arguments, encodes public values and maps native errors | +| Host services | `Protocol::HostCall`, `HostServices::call` | `HostCallHandler` answers typed calls; `()` handles protocols without host calls | `PythonHostCalls` invokes retained Python objects; it may share an owner with the binding | +| Active hooks | `RouteHooks::{before_provider_request, on_event}` | Request interception and fallible execution callbacks | `PythonCallHooks` also prepares arguments, transforms public responses and receives stream callbacks | +| Passive observation | `lifecycle::CallObserver` | Start and terminal observation, retained by the response body | Public Python callbacks remain active hooks with their existing failure policy | +| Runtime driving | `Machine`, `Suspension::{HostCall, Hook, Stream}` | Body polling controls demand | Native polling and inline caller-task Python awaits control progress | + +`hosted_call(request, execute)` starts with a typed request. Its route closure receives separate `HostServices` and `ChannelHooks`; it returns `CallOutput`. Hosted-call plumbing alone forwards the returned stream through demand replies. Lower-level `CallMachine` users receive a `CallContext` containing separately named services, hooks and stream delivery + +Core route constructors prepare their dependencies and return a closure accepting the typed request. Python starts that closure only after argument preparation, preflight and decoding succeed. These steps remain inside the driver's terminal and error handling. Decoding may retain objects for subsequent host service calls + +Each driver owns terminal dispatch. Hooks can change or fail execution; passive observers return no result. HTTP observes success after response conversion or stream exhaustion, failure on errors, and cancellation on body drop. Python preserves exception identity and maps native failures once. Explicit Python stream close reports success for delivered chunks; cancellation stops further callback dispatch + +`in_process::Host` is an assembly of services, hooks, stream consumer and optional observer. It is not a trait mirroring every suspension. Use `run_hosted` to preserve the distinction between stream completion and detachment + +Keep API policy in gateway-inference and python-bridge, and legacy callback policy in callbacks-legacy-python. Python bindings and hooks expose retained references through `PythonOwned`, with idempotent close and GC traversal. Runtime machinery stays in driver, native, handle and runtime modules diff --git a/litellm-rust/crates/host/Cargo.toml b/litellm-rust/crates/host/Cargo.toml index bbbed68f345..51015e99906 100644 --- a/litellm-rust/crates/host/Cargo.toml +++ b/litellm-rust/crates/host/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +futures-util.workspace = true litellm-auth.workspace = true litellm-coroutine.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/host/src/call.rs b/litellm-rust/crates/host/src/call.rs new file mode 100644 index 00000000000..6468cf319ae --- /dev/null +++ b/litellm-rust/crates/host/src/call.rs @@ -0,0 +1,65 @@ +use std::future::Future; + +use futures_util::{TryStreamExt, stream::BoxStream}; + +use crate::{ + machine::{CallMachine, ChannelHooks, HostServices, MachineFault}, + protocol::{Demand, Protocol}, +}; + +pub enum CallOutput { + Complete(Response), + Stream { + head: Head, + chunks: BoxStream<'static, Result>, + }, +} + +#[derive(Debug, PartialEq, Eq)] +pub enum HostedCompletion { + Complete(Response), + StreamEnded, + Detached, +} + +impl From for HostedCompletion { + fn from(response: Response) -> Self { + Self::Complete(response) + } +} + +pub type OutputOf

= CallOutput< +

::Response, +

::StreamHead, +

::Chunk, +

::Error, +>; + +pub type HostedMachine

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

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

, ChannelHooks

) -> Fut + Send + 'static, + Fut: Future, P::Error>> + Send + 'static, +{ + CallMachine::new(move |host| { + Box::pin(async move { + match execute(request, host.services, host.hooks).await? { + CallOutput::Complete(response) => Ok(HostedCompletion::Complete(response)), + CallOutput::Stream { head, mut chunks } => { + if host.stream.open_stream(head).await? == Demand::Detached { + return Ok(HostedCompletion::Detached); + } + while let Some(chunk) = chunks.try_next().await? { + if host.stream.send_chunk(chunk).await? == Demand::Detached { + return Ok(HostedCompletion::Detached); + } + } + Ok(HostedCompletion::StreamEnded) + } + } + }) + }) +} diff --git a/litellm-rust/crates/host/src/event.rs b/litellm-rust/crates/host/src/event.rs index 182dab657d3..dc888ddf7e7 100644 --- a/litellm-rust/crates/host/src/event.rs +++ b/litellm-rust/crates/host/src/event.rs @@ -73,4 +73,7 @@ pub enum CallEvent { 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 14b0f1ea08a..fdf6e05bfa9 100644 --- a/litellm-rust/crates/host/src/hooks.rs +++ b/litellm-rust/crates/host/src/hooks.rs @@ -1,51 +1,38 @@ use std::future::Future; -use crate::{ - event::{MachineEvent, RequestContext, WireRequest}, - machine::{HostChannel, MachineFault}, - protocol::Protocol, -}; +use crate::event::{MachineEvent, RequestContext, WireRequest}; /// What a route reaches for mid-call: the send-time rewrite and the events it reports. /// Python's `logging_obj.pre_call` and `post_call`, in that order. pub trait RouteHooks: Send + Sync { - fn before_send( + fn observer(&self) -> Option> { + None + } + + fn before_provider_request( &self, wire: WireRequest, context: RequestContext, ) -> impl Future> + Send; - fn emit(&self, event: MachineEvent) -> 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_send(&self, wire: WireRequest, _: RequestContext) -> Result { + async fn before_provider_request( + &self, + wire: WireRequest, + _: RequestContext, + ) -> Result { Ok(wire) } - async fn emit(&self, _: MachineEvent) -> Result<(), E> { + async fn on_event(&self, _: MachineEvent) -> Result<(), E> { Ok(()) } } -impl RouteHooks for HostChannel -where - R::Error: From, -{ - async fn before_send( - &self, - wire: WireRequest, - context: RequestContext, - ) -> Result { - HostChannel::before_send(self, wire, context).await - } - - async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { - HostChannel::emit(self, event).await - } -} - #[cfg(test)] mod tests { use std::convert::Infallible; @@ -53,10 +40,13 @@ mod tests { use serde_json::json; use super::*; + use crate::protocol::HookRequest; use crate::{ event::RawResponse, - host::HostOp, + machine::MachineFault, machine::{CallMachine, Machine, MachineStep}, + protocol::Protocol, + protocol::Suspension, }; struct Unit; @@ -67,8 +57,8 @@ mod tests { impl Protocol for Unit { type Response = (WireRequest, ()); type Error = Fault; - type Projection = (); - type Op = Infallible; + type Request = (); + type HostCall = Infallible; type Chunk = Infallible; type StreamHead = Infallible; } @@ -97,13 +87,19 @@ mod tests { } } + #[rstest::rstest] #[tokio::test] async fn the_channel_yields_each_hook_as_its_op_and_returns_the_answer() { let mut machine = CallMachine::::new(|channel| { Box::pin(async move { - let sent = RouteHooks::before_send(&channel, wire("prepared"), context()).await?; - RouteHooks::emit( - &channel, + let sent = RouteHooks::before_provider_request( + &channel.hooks, + wire("prepared"), + context(), + ) + .await?; + RouteHooks::on_event( + &channel.hooks, MachineEvent::ResponseReceived { raw: RawResponse { body: "raw".into() }, }, @@ -113,9 +109,13 @@ mod tests { }) }); - let Ok(MachineStep::Host(HostOp::BeforeSend { wire, reply, .. })) = machine.resume().await + let Ok(MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest { + wire, + reply, + .. + }))) = machine.resume().await else { - panic!("before_send yields BeforeSend"); + panic!("before_provider_request yields BeforeSend"); }; assert_eq!(wire.url, "prepared"); reply.send(WireRequest { @@ -123,8 +123,10 @@ mod tests { ..*wire }); - let Ok(MachineStep::Host(HostOp::Emit(event, reply))) = machine.resume().await else { - panic!("emit yields Emit"); + let Ok(MachineStep::Suspended(Suspension::Hook(HookRequest::Event(event, reply)))) = + machine.resume().await + else { + panic!("on_event yields Emit"); }; assert!(matches!(event, MachineEvent::ResponseReceived { .. })); reply.send(()); @@ -135,9 +137,10 @@ mod tests { assert_eq!(sent.url, "rewritten"); } + #[rstest::rstest] #[tokio::test] async fn no_hooks_pass_the_wire_request_through() { - let sent = RouteHooks::::before_send(&(), wire("prepared"), context()) + let sent = RouteHooks::::before_provider_request(&(), wire("prepared"), context()) .await .unwrap(); assert_eq!(sent.url, "prepared"); diff --git a/litellm-rust/crates/host/src/host.rs b/litellm-rust/crates/host/src/host.rs deleted file mode 100644 index 9714b9470a3..00000000000 --- a/litellm-rust/crates/host/src/host.rs +++ /dev/null @@ -1,68 +0,0 @@ -use std::future::Future; - -pub use litellm_coroutine::{Abandoned, Answer, Reply, reply}; - -use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; -use crate::protocol::Protocol; - -/// One suspension point of a native call, performed by the host and answered through the -/// [`Reply`] it carries. -pub enum HostOp { - /// The first op of every call: the caller's request as the host projects it. - Project(Reply), - Custom(R::Op), - BeforeSend { - wire: Box, - context: Box, - reply: Reply, - }, - Emit(MachineEvent, Reply<()>), - /// The response streams: the host hands the caller a stream and answers once the - /// caller asks for the first chunk or goes away. - Open(R::StreamHead, Reply), - /// The next chunk of an open stream, answered once the caller asks for the one after. - Deliver(R::Chunk, Reply), -} - -/// Whether the caller of a streamed call still reads it. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum Demand { - More, - Detached, -} - -/// A host answer that is either available now or arrives once the host's own -/// suspension (a Python awaitable, for example) resolves. -pub enum HostStep { - Ready(V), - Suspend(S), -} - -/// An in-process host: answers custom operations and observes the call without leaving -/// the Rust runtime. Language hosts implement their own driver instead. -pub trait Host: Send + Sync { - fn project(&self) -> impl Future> + Send; - - /// Answers `op` through its reply, or fails the call. - fn custom_op(&self, op: R::Op) -> impl Future> + Send; - - fn before_send( - &self, - wire: WireRequest, - _context: &RequestContext, - ) -> impl Future> + Send { - async move { Ok(wire) } - } - - fn emit(&self, _event: &CallEvent) -> impl Future> + Send { - async { Ok(()) } - } - - fn open(&self, _head: R::StreamHead) -> impl Future> + Send { - async { Ok(Demand::More) } - } - - fn deliver(&self, _chunk: R::Chunk) -> impl Future> + Send { - async { Ok(Demand::More) } - } -} diff --git a/litellm-rust/crates/host/src/in_process.rs b/litellm-rust/crates/host/src/in_process.rs new file mode 100644 index 00000000000..ab2025fb289 --- /dev/null +++ b/litellm-rust/crates/host/src/in_process.rs @@ -0,0 +1,323 @@ +use crate::{ + event::{CallEvent, FailureOrigin, Timing, epoch_seconds}, + hooks::RouteHooks, + lifecycle::CallObserver, + machine::{HostFailure, Machine, MachineStep}, + protocol::{Demand, HookRequest, Protocol, StreamDelivery, Suspension}, + services::HostCallHandler, +}; +use std::future::Future; + +pub trait StreamConsumer: 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/lib.rs b/litellm-rust/crates/host/src/lib.rs index 1df68941fa3..da130002449 100644 --- a/litellm-rust/crates/host/src/lib.rs +++ b/litellm-rust/crates/host/src/lib.rs @@ -2,13 +2,16 @@ //! //! A host is whatever sits on the far side of the language boundary: CPython today, //! another runtime later. Core runs each route on a [`machine::CallMachine`] and never learns -//! which host is on the other end. The machine yields [`host::HostOp`]s; a driver answers -//! each through the typed [`host::Reply`] it carries, observes [`event::CallEvent`]s and +//! which host is on the other end. The machine yields [`protocol::Suspension`]s; a driver answers +//! each through the typed [`protocol::Reply`] it carries, observes [`event::CallEvent`]s and //! may rewrite the wire request before it is sent. +pub mod call; pub mod event; pub mod hooks; -pub mod host; +pub mod in_process; +pub mod lifecycle; pub mod machine; pub mod protocol; -pub mod run; + +pub mod services; diff --git a/litellm-rust/crates/host/src/lifecycle.rs b/litellm-rust/crates/host/src/lifecycle.rs new file mode 100644 index 00000000000..25ea4bff8fc --- /dev/null +++ b/litellm-rust/crates/host/src/lifecycle.rs @@ -0,0 +1,121 @@ +use std::{future::Future, sync::Arc}; + +use futures_util::TryStreamExt; + +use crate::{ + call::CallOutput, + event::{CallEvent, FailureOrigin, Timing, epoch_seconds}, +}; + +pub trait CallObserver: Send + Sync { + fn observe(&self, event: CallEvent); +} + +struct CallGuard { + observer: Option>, + started_at: f64, +} + +impl CallGuard { + fn new(observer: Arc) -> Self { + let started_at = epoch_seconds(); + observer.observe(CallEvent::Started { + start_time: started_at, + }); + Self { + observer: Some(observer), + started_at, + } + } + + fn timing(&self) -> Timing { + Timing { + start_time: self.started_at, + end_time: epoch_seconds(), + } + } + + fn finish(mut self, failed: bool) { + if let Some(observer) = self.observer.take() { + observer.observe(if failed { + CallEvent::Failed { + timing: self.timing(), + origin: FailureOrigin::Call, + } + } else { + CallEvent::Succeeded { + timing: self.timing(), + } + }); + } + } +} + +impl Drop for CallGuard { + fn drop(&mut self) { + if let Some(observer) = self.observer.take() { + observer.observe(CallEvent::Cancelled { + timing: self.timing(), + }); + } + } +} + +pub async fn observe_call( + observer: Option>, + execute: impl Future, E>>, +) -> Result, E> +where + C: Send + 'static, + E: Send + 'static, +{ + let Some(observer) = observer else { + return execute.await; + }; + let guard = CallGuard::new(observer); + match execute.await { + Err(error) => { + guard.finish(true); + Err(error) + } + Ok(CallOutput::Complete(response)) => { + guard.finish(false); + Ok(CallOutput::Complete(response)) + } + Ok(CallOutput::Stream { head, chunks }) => { + let stream = futures_util::stream::try_unfold( + (chunks, guard), + |(mut chunks, guard)| async move { + match chunks.try_next().await { + Ok(Some(chunk)) => Ok(Some((chunk, (chunks, guard)))), + Ok(None) => { + guard.finish(false); + Ok(None) + } + Err(error) => { + guard.finish(true); + Err(error) + } + } + }, + ); + Ok(CallOutput::Stream { + head, + chunks: Box::pin(stream), + }) + } + } +} + +pub async fn observe_unary( + observer: Option>, + execute: impl Future>, +) -> Result { + let Some(observer) = observer else { + return execute.await; + }; + let guard = CallGuard::new(observer); + let result = execute.await; + guard.finish(result.is_err()); + result +} diff --git a/litellm-rust/crates/host/src/machine/auth.rs b/litellm-rust/crates/host/src/machine/auth.rs index df8c866ce4e..21ba1dc570f 100644 --- a/litellm-rust/crates/host/src/machine/auth.rs +++ b/litellm-rust/crates/host/src/machine/auth.rs @@ -1,18 +1,18 @@ use std::sync::Arc; -use super::{HostChannel, MachineFault}; -use crate::{host::Reply, protocol::Protocol}; +use super::{HostServices, MachineFault}; +use crate::{protocol::Protocol, protocol::Reply}; use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; /// A protocol whose host can mint credentials on the call's behalf. pub trait TokenProtocol: Protocol { - fn acquire_token_op(reply: Reply) -> Self::Op; + 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: HostChannel, + channel: HostServices, } impl std::fmt::Debug for HostTokenProvider { @@ -26,7 +26,7 @@ where R: TokenProtocol, R::Error: From + std::fmt::Display, { - pub fn handle(channel: HostChannel) -> TokenProviderHandle { + pub fn handle(channel: HostServices) -> TokenProviderHandle { TokenProviderHandle::new(Arc::new(Self { channel })) } } @@ -39,7 +39,7 @@ where fn acquire(&self) -> TokenFuture<'_> { Box::pin(async move { self.channel - .custom_op(R::acquire_token_op) + .call(R::acquire_token_op) .await .map_err(|error| Error::CredentialAcquisition(error.to_string().into())) }) diff --git a/litellm-rust/crates/host/src/machine/call_machine.rs b/litellm-rust/crates/host/src/machine/call_machine.rs index af0bc50fbe6..3d57b024a55 100644 --- a/litellm-rust/crates/host/src/machine/call_machine.rs +++ b/litellm-rust/crates/host/src/machine/call_machine.rs @@ -1,7 +1,9 @@ //! The one machine every route runs on: the route's provider future as a -//! [`Coroutine`] that yields [`HostOp`]s, each answered through its own typed reply. No +//! [`Coroutine`] that yields [`Suspension`]s, each answered through its own typed reply. No //! task is spawned; dropping the machine drops the in-flight call. +use crate::protocol::HookRequest; +use crate::protocol::StreamDelivery; use std::{future::Future, pin::Pin}; use litellm_coroutine::{Co, Coroutine, CoroutineState, ResumeError}; @@ -9,8 +11,7 @@ use litellm_coroutine::{Co, Coroutine, CoroutineState, ResumeError}; use super::{HostFailure, Interrupted, Machine, MachineStep, Step}; use crate::{ event::{MachineEvent, RequestContext, WireRequest}, - host::{Demand, HostOp, Reply}, - protocol::Protocol, + protocol::{Demand, Protocol, Reply, Suspension}, }; /// The machine's own failures, distinct from anything the provider call reports. @@ -22,99 +23,146 @@ pub enum MachineFault { Protocol(ResumeError), } -pub type ExecuteFuture = - Pin::Response, ::Error>> + Send>>; +pub type ExecuteFuture::Response> = + Pin::Error>> + Send>>; -/// The provider side of the machine: how the in-flight call reaches its host. -pub struct HostChannel { - co: Co>, +pub struct CallContext { + pub services: HostServices

, + pub hooks: ChannelHooks

, + pub stream: StreamSender

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

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

where - R::Error: From, + P::Error: From, { - async fn yield_( + async fn request_reply( &self, - ask: impl FnOnce(Reply) -> HostOp + Send, - ) -> Result { - self.co - .yield_(ask) + request: impl FnOnce(Reply) -> Suspension

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

); + +impl Clone for HostServices

{ + fn clone(&self) -> Self { + Self(self.0.clone()) } +} - /// Asks the host to perform the custom operation `ask` builds around its reply, as in - /// `host.custom_op(OcrOp::AcquireAzureAdToken)`. - pub async fn custom_op( +impl HostServices

+where + P::Error: From, +{ + pub async fn call( &self, - ask: impl FnOnce(Reply) -> R::Op + Send, - ) -> Result { - self.yield_(|reply| HostOp::Custom(ask(reply))).await + request: impl FnOnce(Reply) -> P::HostCall + Send, + ) -> Result { + self.0 + .request_reply(|reply| Suspension::HostCall(request(reply))) + .await } +} - pub async fn before_send( +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.yield_(|reply| HostOp::BeforeSend { - wire: Box::new(wire), - context: Box::new(context), - reply, - }) - .await + ) -> Result { + self.0 + .request_reply(|reply| { + Suspension::Hook(HookRequest::BeforeProviderRequest { + wire: Box::new(wire), + context: Box::new(context), + reply, + }) + }) + .await } - pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { - self.yield_(|reply| HostOp::Emit(event, reply)).await - } - - pub async fn open(&self, head: R::StreamHead) -> Result { - self.yield_(|reply| HostOp::Open(head, reply)).await - } - - pub async fn deliver(&self, chunk: R::Chunk) -> Result { - self.yield_(|reply| HostOp::Deliver(chunk, reply)).await + async fn on_event(&self, event: MachineEvent) -> Result<(), P::Error> { + self.0 + .request_reply(|reply| Suspension::Hook(HookRequest::Event(event, reply))) + .await } } -type CallCoroutine = - Coroutine, Result<::Response, ::Error>>; +pub struct StreamSender(Channel

); -pub struct CallMachine { - coroutine: CallCoroutine, +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 + } } -impl CallMachine +type CallCoroutine = Coroutine, Result::Error>>; + +pub struct CallMachine::Response> { + coroutine: CallCoroutine, +} + +impl CallMachine where R::Error: From, { - pub fn new(execute: impl FnOnce(HostChannel) -> ExecuteFuture + Send + 'static) -> Self { + pub fn new( + execute: impl FnOnce(CallContext) -> ExecuteFuture + Send + 'static, + ) -> Self { Self { - coroutine: Coroutine::new(|co| execute(HostChannel { co })), + coroutine: Coroutine::new(|co| { + let channel = Channel(co); + execute(CallContext { + services: HostServices(channel.clone()), + hooks: ChannelHooks(channel.clone()), + stream: StreamSender(channel), + }) + }), } } } -impl Machine for CallMachine +impl Machine for CallMachine where R::Error: From, { type Protocol = R; - type Complete = R::Response; + type Complete = C; fn resume(&mut self) -> Step<'_, Self> { Box::pin(async move { @@ -124,7 +172,7 @@ where .await .map_err(MachineFault::Protocol)? { - CoroutineState::Yielded(op) => Ok(MachineStep::Host(op)), + CoroutineState::Yielded(op) => Ok(MachineStep::Suspended(op)), CoroutineState::Complete(outcome) => outcome.map(MachineStep::Complete), } }) diff --git a/litellm-rust/crates/host/src/machine/mod.rs b/litellm-rust/crates/host/src/machine/mod.rs index 0c7501633fa..d307c9cec7d 100644 --- a/litellm-rust/crates/host/src/machine/mod.rs +++ b/litellm-rust/crates/host/src/machine/mod.rs @@ -5,13 +5,14 @@ use std::future::Future; use std::pin::Pin; pub use auth::{HostTokenProvider, TokenProtocol}; -pub use call_machine::{CallMachine, ExecuteFuture, HostChannel, MachineFault}; +pub use call_machine::{ + CallContext, CallMachine, ChannelHooks, ExecuteFuture, HostServices, MachineFault, StreamSender, +}; -use crate::host::HostOp; -use crate::protocol::Protocol; +use crate::protocol::{Protocol, Suspension}; pub enum MachineStep { - Host(HostOp), + Suspended(Suspension), Complete(C), } diff --git a/litellm-rust/crates/host/src/protocol.rs b/litellm-rust/crates/host/src/protocol.rs index a7c0f3470b2..1715d185c04 100644 --- a/litellm-rust/crates/host/src/protocol.rs +++ b/litellm-rust/crates/host/src/protocol.rs @@ -1,17 +1,38 @@ -/// One public call surface: what a completed call produces, how it fails, what the host -/// projects the caller's request into, and the protocol-specific operations only its host -/// can perform mid-call (token acquisition, for one). +pub use litellm_coroutine::{Abandoned, Answer, Reply, reply}; + +use crate::event::{MachineEvent, RequestContext, WireRequest}; + pub trait Protocol: Send + Sync + 'static { + type Request: Send + 'static; type Response: Send + 'static; type Error: Clone + Send + Sync + 'static; - /// The caller's request as the host projects it, answered once before anything else. - type Projection: Send + 'static; - /// Each operation carries the [`Reply`](crate::host::Reply) its answer goes through. - /// A protocol with no operations of its own uses `Infallible`. - type Op: Send + 'static; - /// One piece of a streamed response, handed to the caller as it arrives. A protocol - /// that never streams uses `Infallible`. + type HostCall: Send + 'static; type Chunk: Send + 'static; - /// What the call knows once a streamed response starts, before its first chunk. type StreamHead: Send + 'static; } + +pub enum Suspension { + HostCall(P::HostCall), + Hook(HookRequest), + Stream(StreamDelivery

), +} + +pub enum HookRequest { + BeforeProviderRequest { + wire: Box, + context: Box, + reply: Reply, + }, + Event(MachineEvent, Reply<()>), +} + +pub enum StreamDelivery { + Open(P::StreamHead, Reply), + Chunk(P::Chunk, Reply), +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Demand { + More, + Detached, +} diff --git a/litellm-rust/crates/host/src/run.rs b/litellm-rust/crates/host/src/run.rs deleted file mode 100644 index baa3b58e058..00000000000 --- a/litellm-rust/crates/host/src/run.rs +++ /dev/null @@ -1,207 +0,0 @@ -use crate::event::{CallEvent, FailureOrigin, Timing, epoch_seconds}; -use crate::host::{Host, HostOp}; -use crate::machine::{HostFailure, Machine, MachineStep}; -use crate::protocol::Protocol; - -/// Drives a machine to completion against an in-process host and emits exactly one -/// terminal event. -pub async fn run( - mut machine: M, - host: &H, -) -> Result::Error> -where - M: Machine, - H: Host, -{ - let start_time = epoch_seconds(); - let _ = host.emit(&CallEvent::Started { start_time }).await; - let outcome = loop { - let op = match machine.resume().await { - Ok(MachineStep::Complete(complete)) => break Ok(complete), - Ok(MachineStep::Host(op)) => op, - Err(error) => break Err(error), - }; - if let Err(error) = perform(host, op).await { - break machine.interrupt(HostFailure::Error(error)).await; - } - }; - let timing = Timing { - start_time, - end_time: epoch_seconds(), - }; - let terminal = match &outcome { - Ok(_) => CallEvent::Succeeded { timing }, - Err(_) => CallEvent::Failed { - timing, - origin: FailureOrigin::Call, - }, - }; - let _ = host.emit(&terminal).await; - outcome -} - -async fn perform>(host: &H, op: HostOp) -> Result<(), R::Error> { - match op { - HostOp::Project(reply) => host - .project() - .await - .map(|projection| reply.send(projection)), - HostOp::Custom(op) => host.custom_op(op).await, - HostOp::BeforeSend { - wire, - context, - reply, - } => host - .before_send(*wire, &context) - .await - .map(|wire| reply.send(wire)), - HostOp::Emit(event, reply) => host - .emit(&CallEvent::Machine(event)) - .await - .map(|()| reply.send(())), - HostOp::Open(head, reply) => host.open(head).await.map(|demand| reply.send(demand)), - HostOp::Deliver(chunk, reply) => host.deliver(chunk).await.map(|demand| reply.send(demand)), - } -} - -#[cfg(test)] -mod tests { - use std::sync::Mutex; - - use super::*; - use crate::host::Reply; - use crate::machine::{CallMachine, MachineFault}; - - struct Unit; - - impl Protocol for Unit { - type Response = (); - type Error = &'static str; - type Projection = (); - type Op = (&'static str, Reply<()>); - type Chunk = std::convert::Infallible; - type StreamHead = std::convert::Infallible; - } - - impl From for &'static str { - fn from(_: MachineFault) -> Self { - "machine fault" - } - } - - #[derive(Default)] - struct Recording { - seen: Mutex>, - fail: Option<&'static str>, - } - - impl Host for Recording { - async fn project(&self) -> Result<(), &'static str> { - self.seen.lock().unwrap().push("project".into()); - Ok(()) - } - - async fn custom_op( - &self, - (op, reply): (&'static str, Reply<()>), - ) -> Result<(), &'static str> { - self.seen.lock().unwrap().push(format!("op:{op}")); - if self.fail == Some(op) { - return Err("host failed"); - } - reply.send(()); - Ok(()) - } - - async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> { - self.seen.lock().unwrap().push(match event { - CallEvent::Started { .. } => "started".into(), - CallEvent::Succeeded { .. } => "succeeded".into(), - CallEvent::Failed { .. } => "failed".into(), - other => format!("{other:?}"), - }); - Ok(()) - } - } - - fn scripted( - ops: &'static [&'static str], - outcome: Result<(), &'static str>, - ) -> CallMachine { - CallMachine::new(move |host| { - Box::pin(async move { - host.project().await?; - for op in ops { - host.custom_op(|reply| (*op, reply)).await?; - } - outcome - }) - }) - } - - #[tokio::test] - async fn forwards_every_op_then_emits_one_succeeded() { - let host = Recording::default(); - let outcome = run(scripted(&["sign", "send"], Ok(())), &host).await; - assert_eq!(outcome, Ok(())); - assert_eq!( - *host.seen.lock().unwrap(), - ["started", "project", "op:sign", "op:send", "succeeded"] - ); - } - - #[tokio::test] - async fn errors_and_host_failures_each_emit_failed_once() { - let host = Recording::default(); - let outcome = run(scripted(&[], Err("boom")), &host).await; - assert_eq!(outcome, Err("boom")); - assert_eq!(*host.seen.lock().unwrap(), ["started", "project", "failed"]); - - let host = Recording { - fail: Some("send"), - ..Recording::default() - }; - let outcome = run(scripted(&["sign", "send", "never"], Ok(())), &host).await; - assert_eq!(outcome, Err("host failed")); - assert_eq!( - *host.seen.lock().unwrap(), - ["started", "project", "op:sign", "op:send", "failed"] - ); - } - - struct StartTimes(Mutex>); - - impl Host for StartTimes { - async fn project(&self) -> Result<(), &'static str> { - Ok(()) - } - - async fn custom_op( - &self, - (_, reply): (&'static str, Reply<()>), - ) -> Result<(), &'static str> { - reply.send(()); - Ok(()) - } - - async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> { - if let CallEvent::Started { start_time } - | CallEvent::Succeeded { - timing: Timing { start_time, .. }, - } = event - { - self.0.lock().unwrap().push(*start_time); - } - Err("observer failed") - } - } - - #[tokio::test] - async fn started_opens_the_call_at_the_terminal_start_time_and_cannot_fail_it() { - let host = StartTimes(Mutex::default()); - assert_eq!(run(scripted(&["send"], Ok(())), &host).await, Ok(())); - let times = host.0.lock().unwrap(); - assert_eq!(times.len(), 2); - assert_eq!(times[0], times[1]); - } -} diff --git a/litellm-rust/crates/host/src/services.rs b/litellm-rust/crates/host/src/services.rs new file mode 100644 index 00000000000..48e046ea69c --- /dev/null +++ b/litellm-rust/crates/host/src/services.rs @@ -0,0 +1,15 @@ +use crate::protocol::Protocol; +use std::{convert::Infallible, future::Future}; + +pub trait HostCallHandler: Send + Sync { + fn handle_host_call( + &self, + call: P::HostCall, + ) -> impl Future> + Send; +} + +impl> HostCallHandler

for () { + async fn handle_host_call(&self, call: Infallible) -> Result<(), P::Error> { + match call {} + } +} diff --git a/litellm-rust/crates/host/tests/call.rs b/litellm-rust/crates/host/tests/call.rs new file mode 100644 index 00000000000..3d94f9046e4 --- /dev/null +++ b/litellm-rust/crates/host/tests/call.rs @@ -0,0 +1,224 @@ +use litellm_host::protocol::StreamDelivery; +use std::{ + convert::Infallible, + sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, + }, +}; + +use futures_util::{StreamExt, stream}; +use litellm_host::{ + call::{CallOutput, HostedCompletion, hosted_call}, + event::CallEvent, + lifecycle::{CallObserver, observe_call, observe_unary}, + machine::{Machine, MachineFault, MachineStep}, + protocol::{Demand, Protocol, Suspension}, +}; +use rstest::{fixture, rstest}; + +#[derive(Debug, Clone)] +struct TestError; + +struct TestProtocol; + +impl Protocol for TestProtocol { + type Response = &'static str; + type Error = TestError; + type Request = usize; + type HostCall = Infallible; + type Chunk = usize; + type StreamHead = &'static str; +} + +impl From for TestError { + fn from(_: MachineFault) -> Self { + Self + } +} + +#[rstest] +#[case::end(None, 3, HostedCompletion::StreamEnded)] +#[case::detach_at_open(Some(0), 0, HostedCompletion::Detached)] +#[case::detach_after_chunk(Some(1), 1, HostedCompletion::Detached)] +#[tokio::test] +async fn delivery_obeys_demand_and_distinguishes_detachment( + #[case] detach_after: Option, + #[case] expected_polls: usize, + #[case] expected: HostedCompletion<&'static str>, +) { + let polls = Arc::new(AtomicUsize::new(0)); + let stream_polls = polls.clone(); + let mut machine = hosted_call::(3, move |count, _, _| async move { + let chunks = stream::iter((0..count).map(Ok)) + .inspect(move |_| { + stream_polls.fetch_add(1, Ordering::SeqCst); + }) + .boxed(); + Ok(CallOutput::Stream { + head: "headers", + chunks, + }) + }); + let MachineStep::Suspended(Suspension::Stream(StreamDelivery::Open(head, reply))) = + machine.resume().await.unwrap() + else { + panic!() + }; + assert_eq!(head, "headers"); + assert_eq!(polls.load(Ordering::SeqCst), 0); + reply.send(if detach_after == Some(0) { + Demand::Detached + } else { + Demand::More + }); + let mut delivered = Vec::new(); + let completed = loop { + match machine.resume().await.unwrap() { + MachineStep::Suspended(Suspension::Stream(StreamDelivery::Chunk(chunk, reply))) => { + delivered.push(chunk); + assert_eq!(polls.load(Ordering::SeqCst), delivered.len()); + reply.send(if detach_after == Some(delivered.len()) { + Demand::Detached + } else { + Demand::More + }); + } + MachineStep::Complete(result) => break result, + _ => panic!("unexpected operation"), + } + }; + assert_eq!(completed, expected); + assert_eq!(polls.load(Ordering::SeqCst), expected_polls); + assert_eq!(delivered, (0..expected_polls).collect::>()); +} + +#[derive(Default)] +struct Observer(Mutex>); + +impl CallObserver for Observer { + fn observe(&self, event: CallEvent) { + self.0.lock().unwrap().push(event); + } +} + +#[fixture] +fn observer() -> Arc { + Arc::new(Observer::default()) +} + +type Output = CallOutput<(), (), usize, &'static str>; + +#[rstest] +#[case::success(false)] +#[case::failure(true)] +#[tokio::test] +async fn unary_calls_emit_one_terminal_event(observer: Arc, #[case] fail: bool) { + let expected = if fail { Err("provider") } else { Ok(7) }; + assert_eq!( + observe_unary(Some(observer.clone()), async { expected }).await, + expected + ); + let events = observer.0.lock().unwrap(); + assert_eq!(events.len(), 2); + assert!(matches!(events[0], CallEvent::Started { .. })); + assert_eq!(matches!(events[1], CallEvent::Failed { .. }), fail); + assert_eq!(matches!(events[1], CallEvent::Succeeded { .. }), !fail); +} + +#[rstest] +#[case::end(false)] +#[case::error(true)] +#[tokio::test] +async fn streams_finish_only_when_consumed(observer: Arc, #[case] fail: bool) { + let chunks = stream::iter([Ok(1), if fail { Err("provider") } else { Ok(2) }]).boxed(); + let output = observe_call(Some(observer.clone()), async { + Ok::(CallOutput::Stream { head: (), chunks }) + }) + .await + .unwrap(); + assert_eq!(observer.0.lock().unwrap().len(), 1); + let CallOutput::Stream { mut chunks, .. } = output else { + panic!() + }; + assert_eq!(chunks.next().await, Some(Ok(1))); + assert_eq!(observer.0.lock().unwrap().len(), 1); + let last = chunks.next().await; + if fail { + assert_eq!(last, Some(Err("provider"))); + } else { + assert_eq!(last, Some(Ok(2))); + assert_eq!(chunks.next().await, None); + } + drop(chunks); + let events = observer.0.lock().unwrap(); + assert_eq!(events.len(), 2); + assert_eq!(matches!(events[1], CallEvent::Failed { .. }), fail); + assert_eq!(matches!(events[1], CallEvent::Succeeded { .. }), !fail); +} + +#[rstest] +#[tokio::test] +async fn dropping_a_stream_cancels_without_success(observer: Arc) { + let chunks = stream::pending().boxed(); + let output = observe_call(Some(observer.clone()), async { + Ok::(CallOutput::Stream { head: (), chunks }) + }) + .await + .unwrap(); + drop(output); + let events = observer.0.lock().unwrap(); + assert_eq!(events.len(), 2); + assert!(matches!(events[1], CallEvent::Cancelled { .. })); +} + +#[rstest] +#[tokio::test] +async fn cancelling_provider_execution_releases_the_lifecycle(observer: Arc) { + let mut call = Box::pin(observe_unary( + Some(observer.clone()), + std::future::pending::>(), + )); + assert!(futures_util::poll!(&mut call).is_pending()); + drop(call); + let events = observer.0.lock().unwrap(); + assert_eq!(events.len(), 2); + assert!(matches!(events[1], CallEvent::Cancelled { .. })); +} + +#[rstest] +#[tokio::test] +async fn a_hosted_detachment_is_reported_as_cancellation(observer: Arc) { + 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/llms/src/base_llm/ocr/handler.rs b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs index 0148ca2841b..6b34fc3ef7c 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs @@ -26,7 +26,7 @@ use litellm_secrets::source::SecretSource; /// The route's view of one call, handed to provider code that has to reach the /// caller's hooks mid-flight (guardrails on the outgoing body, raw response events). pub trait CallHooks: Send + Sync { - fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result>; + fn before_provider_request(&self, wire: WireRequest) -> BoxFuture<'_, Result>; fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), E>>; } @@ -226,7 +226,7 @@ pub async fn transform_request_body( )?; config.validate_request_body(&composed)?; let changed = hooks - .before_send(wire_request(url, headers, composed)) + .before_provider_request(wire_request(url, headers, composed)) .await?; if !changed.body.is_object() { return Err(Error::RequestField { @@ -278,7 +278,9 @@ pub async fn guardrail_document( let body = serde_json::to_value(&request.document).map_err(|_| Error::RequestField { path: "document".into(), })?; - let changed = hooks.before_send(wire_request(url, headers, body)).await?; + let changed = hooks + .before_provider_request(wire_request(url, headers, body)) + .await?; let document = decode_request_value(changed.body, "guardrail.document")?; Ok((document, changed.headers)) } @@ -304,6 +306,7 @@ mod tests { use super::*; + #[rstest::rstest] #[tokio::test] async fn request_timeout_has_an_http_408_status() { let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); diff --git a/litellm-rust/crates/python-bridge/AGENTS.md b/litellm-rust/crates/python-bridge/AGENTS.md index e2b56e390f3..507224e3446 100644 --- a/litellm-rust/crates/python-bridge/AGENTS.md +++ b/litellm-rust/crates/python-bridge/AGENTS.md @@ -7,7 +7,7 @@ - Python, Rust SDK and gateway use one lifecycle-bearing core route entrypoint; provider helpers stay private, never bridge-accessible transport drivers - Built-in provider/config/secret/auth/document preparation stays in Rust; caller-authored callbacks and focused Python-file reads run only at core-selected points - Target GIL-enabled CPython explicitly with `#[pymodule(gil_used = true)]`; detach Rust-only work - - GIL and tokio invariants, each pinned by a test in `host-python` (`execution.rs`, + - GIL and tokio invariants, each pinned by a test in `host-python` (`runtime.rs`, `gil.rs`) so a regression fails there before it deadlocks a proxy: - Never hold the GIL while waiting on the runtime. A sync entrypoint releases it with `release_gil` around `block_on`, because every task that attaches would otherwise wait diff --git a/litellm-rust/crates/python-bridge/src/cache/semantic.rs b/litellm-rust/crates/python-bridge/src/cache/semantic.rs index eb67f5594e2..f05635c03df 100644 --- a/litellm-rust/crates/python-bridge/src/cache/semantic.rs +++ b/litellm-rust/crates/python-bridge/src/cache/semantic.rs @@ -207,8 +207,5 @@ impl ExecutionBody for SemanticExecution { } pub(super) fn drive(py: Python<'_>, body: SemanticExecution) -> PyResult> { - let execution = Py::new(py, Execution::new(body))?; - py.import("litellm.rust_bridge.lifecycle")? - .getattr("drive")? - .call1((execution,)) + Execution::new(body).into_coroutine(py) } diff --git a/litellm-rust/crates/python-bridge/src/http.rs b/litellm-rust/crates/python-bridge/src/http.rs index e74b9d198a0..a5c2166c18e 100644 --- a/litellm-rust/crates/python-bridge/src/http.rs +++ b/litellm-rust/crates/python-bridge/src/http.rs @@ -117,6 +117,15 @@ pub(crate) fn call_config( Ok(resolution.config) } +pub(crate) fn provider_client( + py: Python<'_>, + kwargs: &Bound<'_, PyDict>, + asynchronous: bool, +) -> PyResult> { + let config = call_config(py, kwargs, asynchronous)?; + Ok(pool().client(&config, ClientVariant::Provider)) +} + pub(crate) fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult { let config = call_config(py, &PyDict::new(py), true)?; pool().client(&config, variant).map_err(client_error) diff --git a/litellm-rust/crates/python-bridge/src/logger/tests.rs b/litellm-rust/crates/python-bridge/src/logger/tests.rs index b65e7d37023..a9b7ef41c08 100644 --- a/litellm-rust/crates/python-bridge/src/logger/tests.rs +++ b/litellm-rust/crates/python-bridge/src/logger/tests.rs @@ -12,8 +12,8 @@ struct DiagnosticMachine; impl Protocol for DiagnosticMachine { type Response = (); type Error = String; - type Projection = (); - type Op = (); + type Request = (); + type HostCall = (); type Chunk = (); type StreamHead = (); } diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index e134f2deeb5..32369890dea 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -1,13 +1,10 @@ use crate::logger::{run_async, run_sync}; use litellm_core::audio_transcription::{ - Error, audio_transcription as run_audio_transcription, types::AudioTranscriptionRequest, + AudioTranscriptionRoute, Error, types::AudioTranscriptionRequest, }; use litellm_host_python::from_py_argument; -use litellm_http::HttpClientConfig; -use litellm_secrets::source::SecretSource; use pyo3::{prelude::*, types::PyDict}; use serde_json::{Map, Value}; -use std::sync::Arc; use crate::{ errors::route_error_to_pyerr, @@ -15,8 +12,8 @@ use crate::{ }; async fn execute( - config: HttpClientConfig, - secrets: Arc, + http: Result, + secrets: std::sync::Arc, audio: Value, optional_params: Map, options: RouteOptions, @@ -29,11 +26,8 @@ async fn execute( extra_headers, timeout, } = options; - run_audio_transcription( - crate::http::resources(), - &config, - secrets.as_ref(), - AudioTranscriptionRequest { + AudioTranscriptionRoute::new(http?, crate::http::resources().auth.clone(), secrets) + .execute(AudioTranscriptionRequest { model: &model, audio, api_key: api_key.as_deref(), @@ -42,9 +36,8 @@ async fn execute( extra_headers, optional_params, timeout, - }, - ) - .await + }) + .await } #[pyfunction] @@ -72,12 +65,13 @@ pub(crate) fn transcription( extra_headers, timeout: optional_timeout(timeout_seconds), }; - let config = crate::http::call_config(py, &PyDict::new(py), false)?; + let http = crate::http::provider_client(py, &PyDict::new(py), false)?; + let secrets = crate::secrets::source(py)?; run_sync( py, execute( - config, - crate::secrets::source(py)?, + http, + secrets, audio, optional_params.unwrap_or_default(), options, @@ -111,12 +105,13 @@ pub(crate) fn atranscription<'py>( extra_headers, timeout: optional_timeout(timeout_seconds), }; - let config = crate::http::call_config(py, &PyDict::new(py), true)?; + let http = crate::http::provider_client(py, &PyDict::new(py), true)?; + let secrets = crate::secrets::source(py)?; run_async( py, execute( - config, - crate::secrets::source(py)?, + http, + secrets, audio, optional_params.unwrap_or_default(), options, 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 6962d7a145d..1f84545d142 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -1,15 +1,11 @@ -use litellm_secrets::source::SecretSource; use pyo3::types::{PyDict, PyTuple}; -use std::sync::Arc; use crate::errors::RustBridgeDeclined; use crate::logger::{run_async, run_sync}; use litellm_core::chat_completions::{ - Error, chat_completions as run_chat_completions, chat_completions_decline_reason, - types::ChatCompletionsRequest, + ChatCompletionsRoute, Error, chat_completions_decline_reason, types::ChatCompletionsRequest, }; use litellm_host_python::from_py_argument; -use litellm_http::HttpClientConfig; use litellm_types::utils::ChatCompletionsResponse; use pyo3::prelude::*; use serde_json::{Map, Value}; @@ -23,8 +19,8 @@ use crate::{ }; async fn execute( - config: HttpClientConfig, - secrets: Arc, + http: Result, + secrets: std::sync::Arc, messages: Vec, optional_params: Map, options: RouteOptions, @@ -37,22 +33,21 @@ async fn execute( extra_headers, timeout, } = options; - run_chat_completions( - crate::http::resources(), - &config, - secrets.as_ref(), - ChatCompletionsRequest { - model: &model, - messages: Value::Array(messages), - optional_params, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }, - ) - .await + ChatCompletionsRoute::new(http?, crate::http::resources().auth.clone(), secrets) + .execute( + ChatCompletionsRequest { + model: &model, + messages: Value::Array(messages), + optional_params, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + timeout, + }, + &(), + ) + .await } #[pyfunction] @@ -97,12 +92,13 @@ pub(crate) fn chat_completions( extra_headers, timeout: optional_timeout(timeout_seconds), }; - let config = crate::http::call_config(py, &PyDict::new(py), false)?; + let http = crate::http::provider_client(py, &PyDict::new(py), false)?; + let secrets = crate::secrets::source(py)?; run_sync( py, execute( - config, - crate::secrets::source(py)?, + http, + secrets, messages, optional_params.unwrap_or_default(), options, @@ -136,12 +132,13 @@ pub(crate) fn achat_completions<'py>( extra_headers, timeout: optional_timeout(timeout_seconds), }; - let config = crate::http::call_config(py, &PyDict::new(py), true)?; + let http = crate::http::provider_client(py, &PyDict::new(py), true)?; + let secrets = crate::secrets::source(py)?; run_async( py, execute( - config, - crate::secrets::source(py)?, + http, + secrets, messages, optional_params.unwrap_or_default(), options, diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index 9d7969fee15..d05dbb75d73 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -1,11 +1,12 @@ +use litellm_host_python::{PythonHostCalls, PythonOwned}; use std::convert::Infallible; use bytes::Bytes; use litellm_core::messages::{ Error, MessagesCall, MessagesShaping, messages_body, - route::{Messages, MessagesOutput, MessagesStreamHead}, + route::{Messages, MessagesStreamHead}, }; -use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py}; +use litellm_host_python::{InvokeError, PythonBinding, from_py, lookup, to_py}; use litellm_http::transport::Error as TransportError; use litellm_types::utils::ProviderSpecificHeaders; use pyo3::{ @@ -221,11 +222,11 @@ impl MessagesPythonHost { } } -impl ProtocolHost for MessagesPythonHost { +impl PythonBinding for MessagesPythonHost { type Protocol = Messages; type Failure = PyErr; - fn project( + fn decode_request( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, @@ -235,33 +236,35 @@ impl ProtocolHost for MessagesPythonHost { .map_err(InvokeError::Native) } - fn invoke(&mut self, _: Python<'_>, op: Infallible) -> Result<(), InvokeError> { - match op {} + fn encode_response( + &mut self, + py: Python<'_>, + response: Box< + litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse, + >, + ) -> PyResult> { + py.import(ROUTE_HOST_MODULE)? + .getattr("response")? + .call1((to_py(py, response.as_ref())?,)) + .map(Bound::unbind) } - fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult> { - match response { - MessagesOutput::Message(message) => py - .import(ROUTE_HOST_MODULE)? - .getattr("response")? - .call1((to_py(py, message.as_ref())?,)) - .map(Bound::unbind), - MessagesOutput::Streamed => Ok(py.None()), - } - } - - fn head(&mut self, py: Python<'_>, head: MessagesStreamHead) -> PyResult> { + fn encode_stream_head( + &mut self, + py: Python<'_>, + head: MessagesStreamHead, + ) -> PyResult> { py.import(ROUTE_HOST_MODULE)? .getattr("stream_hidden_params")? .call1((to_py(py, &head.headers)?,)) .map(Bound::unbind) } - fn chunk(&mut self, py: Python<'_>, chunk: Bytes) -> PyResult> { + fn encode_chunk(&mut self, py: Python<'_>, chunk: Bytes) -> PyResult> { Ok(PyBytes::new(py, &chunk).into_any().unbind()) } - fn classify(&self, py: Python<'_>, error: Error) -> PyResult { + fn map_error(&self, py: Python<'_>, error: Error) -> PyResult { if let Error::Secret(source) = &error && let Some(original) = crate::secrets::python_error(py, source.source_error()) { @@ -273,9 +276,20 @@ impl ProtocolHost for MessagesPythonHost { fn host_error(error: &PyErr) -> Error { Error::InvalidRequest(error.to_string().into()) } +} +impl PythonHostCalls for MessagesPythonHost { + fn handle_host_call( + &mut self, + _: Python<'_>, + op: Infallible, + ) -> Result<(), InvokeError> { + match op {} + } +} + +impl PythonOwned for MessagesPythonHost { fn close(&mut self, _: Python<'_>) {} - fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.request) } 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 96ca9eebecb..f288b71dbde 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -4,7 +4,6 @@ use host::MessagesPythonHost; use litellm_callbacks_legacy_python::{ LegacySurface, PassThroughStream, PublicCall, run_legacy_call, }; -use litellm_core::messages::route::messages_machine; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, @@ -26,15 +25,17 @@ fn run_messages( kwargs: Bound<'_, PyDict>, asynchronous: bool, ) -> PyResult> { - let secrets = crate::secrets::source(py)?; - let config = crate::http::call_config(py, &kwargs, asynchronous)?; - let machine = messages_machine(crate::http::resources(), &config, secrets) - .map_err(crate::http::client_error)?; + let route = litellm_core::messages::MessagesRoute::new( + crate::http::provider_client(py, &kwargs, asynchronous)? + .map_err(crate::http::client_error)?, + crate::http::resources().auth.clone(), + crate::secrets::source(py)?, + ); run_legacy_call( py, SURFACE, PublicCall::capture(&request, &args, &kwargs)?, - crate::logger::LoggedMachine::new(machine), + move |request| crate::logger::LoggedMachine::new(route.machine(request)), MessagesPythonHost::new(request.unbind()), crate::preflight::sdk_preflight, asynchronous, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index dc01ced15a0..03a982f8117 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -1,6 +1,7 @@ use litellm_auth::ResolvedCredential; -use litellm_core::ocr::route::{Ocr, OcrOp, OcrProjection}; -use litellm_host_python::{InvokeError, ProtocolHost, missing_state, to_py}; +use litellm_core::ocr::route::{Ocr, OcrCall, OcrOp}; +use litellm_host_python::{InvokeError, PythonBinding, missing_state, to_py}; +use litellm_host_python::{PythonHostCalls, PythonOwned}; use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}; use pyo3::{ exceptions::{PyBaseException, PyException}, @@ -51,18 +52,14 @@ impl OcrPythonHost { .acquire(py) } - fn projection( - &mut self, - py: Python<'_>, - arguments: &Bound<'_, PyDict>, - ) -> PyResult { + fn projection(&mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { let OcrHostData::Unprojected = self.data else { return Err(missing_state()); }; let (request, handles) = project_request(self.request.bind(py), arguments)?; let caller_token = handles.azure_ad_token_provider.is_some(); self.data = OcrHostData::Projected(Box::new(handles)); - Ok(OcrProjection { + Ok(OcrCall { request, caller_token, }) @@ -88,44 +85,47 @@ impl OcrPythonHost { } } -impl ProtocolHost for OcrPythonHost { +impl PythonBinding for OcrPythonHost { type Protocol = Ocr; type Failure = PyErr; - fn project( + fn decode_request( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - ) -> Result> { + ) -> Result> { self.projection(py, arguments) .map_err(|error| InvokeError::Python(self.map_failure(py, error))) } - fn invoke(&mut self, py: Python<'_>, op: OcrOp) -> Result<(), InvokeError> { - match op { - OcrOp::AcquireAzureAdToken(reply) => self - .acquire_azure_ad_token(py) - .map(|token| reply.send(token)) - .map_err(|error| InvokeError::Python(self.map_failure(py, error))), - } - } - - fn complete(&mut self, py: Python<'_>, response: LiteLLMOcrResponse) -> PyResult> { + fn encode_response( + &mut self, + py: Python<'_>, + response: LiteLLMOcrResponse, + ) -> PyResult> { py.import("litellm.rust_bridge.ocr.route_host")? .getattr("response")? .call1((to_py(py, &response)?,)) .map(Bound::unbind) } - fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult> { + fn encode_stream_head( + &mut self, + _: Python<'_>, + head: std::convert::Infallible, + ) -> PyResult> { match head {} } - fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { + fn encode_chunk( + &mut self, + _: Python<'_>, + chunk: std::convert::Infallible, + ) -> PyResult> { match chunk {} } - fn classify(&self, py: Python<'_>, error: Error) -> PyResult { + fn map_error(&self, py: Python<'_>, error: Error) -> PyResult { if let Error::Secret(source) = &error && let Some(original) = crate::secrets::python_error(py, source) { @@ -137,11 +137,23 @@ impl ProtocolHost for OcrPythonHost { fn host_error(error: &PyErr) -> Error { Error::InvalidRequest(error.to_string()) } +} +impl PythonHostCalls for OcrPythonHost { + fn handle_host_call(&mut self, py: Python<'_>, op: OcrOp) -> Result<(), InvokeError> { + match op { + OcrOp::AcquireAzureAdToken(reply) => self + .acquire_azure_ad_token(py) + .map(|token| reply.send(token)) + .map_err(|error| InvokeError::Python(self.map_failure(py, error))), + } + } +} + +impl PythonOwned for OcrPythonHost { fn close(&mut self, _: Python<'_>) { self.data = OcrHostData::Released; } - fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.request)?; if let OcrHostData::Projected(handles) = &self.data @@ -199,12 +211,13 @@ del provider .cast_into::() .unwrap(); let mut host = OcrPythonHost::new(py.None()); - assert!(host.project(py, &kwargs).unwrap().caller_token); + assert!(host.decode_request(py, &kwargs).unwrap().caller_token); locals.del_item("kwargs").unwrap(); drop(kwargs); - let (reply, _) = litellm_host::host::reply(); + let (reply, _) = litellm_host::protocol::reply(); assert_eq!( - host.invoke(py, OcrOp::AcquireAzureAdToken(reply)).is_ok(), + host.handle_host_call(py, OcrOp::AcquireAzureAdToken(reply)) + .is_ok(), succeeds ); let alive = || { 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 d7c54e996ff..82ee3277f97 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -5,7 +5,7 @@ mod project; use host::OcrPythonHost; use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; -use litellm_core::ocr::{provider_config, route::ocr_machine}; +use litellm_core::ocr::provider_config; use litellm_core_utils::settings::ProcessEnvironment; use litellm_host_python::to_py; use litellm_llms::base_llm::ocr::settings::OcrSettings; @@ -18,7 +18,6 @@ use crate::{ coercion::FieldSpec, http, python_settings::{PythonSettings, Snapshot}, - secrets, }; const VERTEX_PROJECT: FieldSpec> = @@ -48,16 +47,22 @@ fn run_ocr( kwargs: Bound<'_, PyDict>, asynchronous: bool, ) -> PyResult> { - let secrets = secrets::source(py)?; let config = http::call_config(py, &kwargs, asynchronous)?; - let client = http::resources() - .ocr_client(&config, http::url_policy(py)?, ocr_settings(py)?, secrets) - .map_err(http::client_error)?; + let client = litellm_llms::base_llm::ocr::handler::OcrClient::new( + &http::resources().pool, + &config, + http::url_policy(py)?, + http::resources().auth.clone(), + ocr_settings(py)?, + crate::secrets::source(py)?, + ) + .map_err(http::client_error)?; + let route = litellm_core::ocr::OcrRoute::new(client); run_legacy_call( py, if asynchronous { ASYNC_SURFACE } else { SURFACE }, PublicCall::capture(&request, &args, &kwargs)?, - crate::logger::LoggedMachine::new(ocr_machine(client)), + move |request| crate::logger::LoggedMachine::new(route.machine(request)), OcrPythonHost::new(request.unbind()), crate::preflight::sdk_preflight, asynchronous, @@ -127,7 +132,7 @@ mod tests { use crate::python_settings::PythonSettings; - #[test] + #[rstest::rstest] fn provider_defaults_distinguish_falsey_values_and_exact_true() { Python::initialize(); Python::attach(|py| { diff --git a/litellm-rust/crates/python-bridge/src/secrets/mod.rs b/litellm-rust/crates/python-bridge/src/secrets/mod.rs index 28d2d177fa1..edf6aead53a 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/mod.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/mod.rs @@ -12,7 +12,7 @@ mod vault; use std::sync::Arc; pub(crate) use error::python_error; -use litellm_secrets::source::SecretSource; +use litellm_secrets::source::{EnvironmentSecrets, SecretSource}; use pyo3::prelude::*; use python::PythonSecrets; use resolved::ResolvedSecrets; @@ -23,15 +23,15 @@ const NATIVE: FieldSpec = FieldSpec::new("native", |field| field.schema_bo /// Where a Rust route reads provider secrets from. Python's `get_secret_str` until a /// `SecretManagerRule` in `catalog.py` moves the configured system off `PYTHON_ONLY`, then the -/// native secret manager. +/// native secret manager. A bare extension module without the litellm package reads the process +/// environment. pub(crate) fn source(py: Python<'_>) -> PyResult> { - let Some(settings) = PythonSettings::SecretManager.read_or_unset(py)? else { - let client = crate::http::host_client(py, litellm_http::ClientVariant::Provider)?; - return Ok(Arc::new( - litellm_secrets::source::EnvironmentSecrets::python_compatible(client), - )); + let Some(snapshot) = PythonSettings::SecretManager.read_or_unset(py)? else { + return Ok(Arc::new(EnvironmentSecrets::python_compatible( + crate::http::host_client(py, litellm_http::ClientVariant::Provider)?, + ))); }; - if settings.read(&NATIVE)? { + if snapshot.read(&NATIVE)? { let context = litellm_host_python::PythonContext::capture(py)?; let client = crate::http::host_client(py, litellm_http::ClientVariant::Provider)?; return Ok(Arc::new(ResolvedSecrets::new(