litellm/litellm-rust/crates/host/tests/call.rs
devin-ai-integration[bot] e4190d86a6
refactor(rust): centralize host execution and compose callbacks (#43515)
* refactor(rust): extract litellm-host-native as the shared Rust host driver

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

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

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

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

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

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

* auth update

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

* style(rust): keep host driver imports formatted

* chores

* mostly relocation

* refactor(rust): separate interceptors from queued observers

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

* docs: define Python host boundaries and migration plan

* refactor: enforce Python host and bridge boundaries

* refactor(rust): separate operations from callback composition

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

---------

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-09-28 19:20:27 +00:00

217 lines
6.7 KiB
Rust

use litellm_host::protocol::StreamDelivery;
use std::{
convert::Infallible,
ops::ControlFlow,
sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
},
};
use futures_util::{StreamExt, stream};
use litellm_host::{
call::{CallOutput, HostedCompletion, hosted_call},
lifecycle::{CallEvent, CallObserver, observe_call, observe_unary},
machine::{Machine, MachineFault, MachineStep},
protocol::{HostRequest, Protocol},
};
use rstest::{fixture, rstest};
#[derive(Debug, Clone)]
struct TestError;
struct TestProtocol;
impl Protocol for TestProtocol {
type Response = &'static str;
type Error = TestError;
type Request = usize;
type HostCall = Infallible;
type Chunk = usize;
type StreamHead = &'static str;
}
impl From<MachineFault> for TestError {
fn from(_: MachineFault) -> Self {
Self
}
}
#[rstest]
#[case::end(None, 3, HostedCompletion::StreamEnded)]
#[case::detach_at_open(Some(0), 0, HostedCompletion::Detached)]
#[case::detach_after_chunk(Some(1), 1, HostedCompletion::Detached)]
#[tokio::test]
async fn delivery_obeys_demand_and_distinguishes_detachment(
#[case] detach_after: Option<usize>,
#[case] expected_polls: usize,
#[case] expected: HostedCompletion<&'static str>,
) {
let polls = Arc::new(AtomicUsize::new(0));
let stream_polls = polls.clone();
let mut machine =
hosted_call::<TestProtocol, _, _>(3, None, move |count, _, _, _observations| async move {
let chunks = stream::iter((0..count).map(Ok))
.inspect(move |_| {
stream_polls.fetch_add(1, Ordering::SeqCst);
})
.boxed();
Ok(CallOutput::Stream {
head: "headers",
chunks,
})
});
let MachineStep::Suspended(HostRequest::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) {
ControlFlow::Break(())
} else {
ControlFlow::Continue(())
});
let mut delivered = Vec::new();
let completed = loop {
match machine.resume().await.unwrap() {
MachineStep::Suspended(HostRequest::Stream(StreamDelivery::Chunk(chunk, reply))) => {
delivered.push(chunk);
assert_eq!(polls.load(Ordering::SeqCst), delivered.len());
reply.send(if detach_after == Some(delivered.len()) {
ControlFlow::Break(())
} else {
ControlFlow::Continue(())
});
}
MachineStep::Complete(result) => break result,
_ => panic!("unexpected operation"),
}
};
assert_eq!(completed, expected);
assert_eq!(polls.load(Ordering::SeqCst), expected_polls);
assert_eq!(delivered, (0..expected_polls).collect::<Vec<_>>());
}
#[derive(Default)]
struct Observer(Observations);
struct Observations {
sender: litellm_host::observation::ObservationSender,
receiver: Mutex<tokio::sync::mpsc::Receiver<CallEvent>>,
recorded: Mutex<Vec<CallEvent>>,
}
impl Default for Observations {
fn default() -> Self {
let (sender, receiver) = litellm_host::observation::observation_channel(
std::num::NonZeroUsize::new(128).unwrap(),
);
Self {
sender,
receiver: Mutex::new(receiver),
recorded: Mutex::new(Vec::new()),
}
}
}
impl Observations {
fn lock(&self) -> std::sync::LockResult<std::sync::MutexGuard<'_, Vec<CallEvent>>> {
let mut events = self.recorded.lock()?;
let mut receiver = self.receiver.lock().unwrap();
while let Ok(event) = receiver.try_recv() {
events.push(event);
}
Ok(events)
}
}
impl CallObserver for Observer {
fn observe(&self, event: CallEvent) {
self.0.sender.emit(event);
}
}
#[fixture]
fn observer() -> Arc<Observer> {
Arc::new(Observer::default())
}
type Output = CallOutput<(), (), usize, &'static str>;
#[rstest]
#[case::success(false)]
#[case::failure(true)]
#[tokio::test]
async fn unary_calls_emit_one_terminal_event(observer: Arc<Observer>, #[case] fail: bool) {
let expected = if fail { Err("provider") } else { Ok(7) };
assert_eq!(
observe_unary(Some(observer.0.sender.clone()), async { expected }).await,
expected
);
let events = observer.0.lock().unwrap();
assert_eq!(events.len(), 2);
assert!(matches!(events[0], CallEvent::Started { .. }));
assert_eq!(matches!(events[1], CallEvent::Failed { .. }), fail);
assert_eq!(matches!(events[1], CallEvent::Succeeded { .. }), !fail);
}
#[rstest]
#[case::end(false)]
#[case::error(true)]
#[tokio::test]
async fn streams_finish_only_when_consumed(observer: Arc<Observer>, #[case] fail: bool) {
let chunks = stream::iter([Ok(1), if fail { Err("provider") } else { Ok(2) }]).boxed();
let output = observe_call(Some(observer.0.sender.clone()), async {
Ok::<Output, _>(CallOutput::Stream { head: (), chunks })
})
.await
.unwrap();
assert_eq!(observer.0.lock().unwrap().len(), 1);
let CallOutput::Stream { mut chunks, .. } = output else {
panic!()
};
assert_eq!(chunks.next().await, Some(Ok(1)));
assert_eq!(observer.0.lock().unwrap().len(), 1);
let last = chunks.next().await;
if fail {
assert_eq!(last, Some(Err("provider")));
} else {
assert_eq!(last, Some(Ok(2)));
assert_eq!(chunks.next().await, None);
}
drop(chunks);
let events = observer.0.lock().unwrap();
assert_eq!(events.len(), 2);
assert_eq!(matches!(events[1], CallEvent::Failed { .. }), fail);
assert_eq!(matches!(events[1], CallEvent::Succeeded { .. }), !fail);
}
#[rstest]
#[tokio::test]
async fn dropping_a_stream_cancels_without_success(observer: Arc<Observer>) {
let chunks = stream::pending().boxed();
let output = observe_call(Some(observer.0.sender.clone()), async {
Ok::<Output, _>(CallOutput::Stream { head: (), chunks })
})
.await
.unwrap();
drop(output);
let events = observer.0.lock().unwrap();
assert_eq!(events.len(), 2);
assert!(matches!(events[1], CallEvent::Cancelled { .. }));
}
#[rstest]
#[tokio::test]
async fn cancelling_provider_execution_releases_the_lifecycle(observer: Arc<Observer>) {
let mut call = Box::pin(observe_unary(
Some(observer.0.sender.clone()),
std::future::pending::<Result<(), ()>>(),
));
assert!(futures_util::poll!(&mut call).is_pending());
drop(call);
let events = observer.0.lock().unwrap();
assert_eq!(events.len(), 2);
assert!(matches!(events[1], CallEvent::Cancelled { .. }));
}