mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
* 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>
217 lines
6.7 KiB
Rust
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 { .. }));
|
|
}
|