use crate::input::OwnedInputs;
use crate::{Error, Result};
use litert_lm_edge_sys as ffi;
use std::ffi::{c_char, c_void, CStr};
use std::marker::PhantomData;
use std::ptr::NonNull;
use std::sync::mpsc::{self, Receiver, Sender, TryRecvError};
use std::time::{Duration, Instant};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum StreamEvent {
Chunk(String),
Final,
Error(String),
}
pub struct TextStream<'session> {
receiver: Receiver<StreamEvent>,
state: Option<NonNull<StreamState>>,
session: NonNull<ffi::LiteRtLmSession>,
terminal_received: bool,
_session: PhantomData<&'session mut ffi::LiteRtLmSession>,
}
struct StreamState {
sender: Sender<StreamEvent>,
input: OwnedInputs,
}
pub(crate) fn start_text_stream<'session>(
session: NonNull<ffi::LiteRtLmSession>,
input: OwnedInputs,
_session_lifetime: PhantomData<&'session mut ffi::LiteRtLmSession>,
) -> Result<TextStream<'session>> {
let (sender, receiver) = mpsc::channel();
let state = Box::new(StreamState { sender, input });
let state = NonNull::new(Box::into_raw(state)).expect("Box::into_raw never returns null");
let code = unsafe {
let state_ref = state.as_ref();
ffi::litert_lm_session_generate_content_stream(
session.as_ptr(),
state_ref.input.as_ffi().as_ptr(),
state_ref.input.as_ffi().len(),
Some(stream_callback),
state.as_ptr().cast(),
)
};
if code != 0 {
unsafe { drop(Box::from_raw(state.as_ptr())) };
return Err(Error::StartStream(code));
}
Ok(TextStream {
receiver,
state: Some(state),
session,
terminal_received: false,
_session: PhantomData,
})
}
impl Iterator for TextStream<'_> {
type Item = StreamEvent;
fn next(&mut self) -> Option<Self::Item> {
let event = self.receiver.recv().ok()?;
self.mark_terminal(&event);
Some(event)
}
}
impl Drop for TextStream<'_> {
fn drop(&mut self) {
let Some(state) = self.state.take() else {
return;
};
if !self.terminal_received {
unsafe { ffi::litert_lm_session_cancel_process(self.session.as_ptr()) };
if !self.drain_until_terminal(Duration::from_secs(30)) {
return;
}
}
unsafe { drop(Box::from_raw(state.as_ptr())) };
}
}
impl TextStream<'_> {
pub fn recv_timeout(&mut self, timeout: Duration) -> Result<Option<StreamEvent>> {
let event = match self.receiver.recv_timeout(timeout) {
Ok(event) => event,
Err(mpsc::RecvTimeoutError::Timeout) => return Ok(None),
Err(mpsc::RecvTimeoutError::Disconnected) => {
return Err(Error::InvalidResponse(
"stream callback disconnected".to_owned(),
))
}
};
self.mark_terminal(&event);
Ok(Some(event))
}
pub fn try_recv(&mut self) -> Result<Option<StreamEvent>> {
let event = match self.receiver.try_recv() {
Ok(event) => event,
Err(TryRecvError::Empty) => return Ok(None),
Err(TryRecvError::Disconnected) => {
return Err(Error::InvalidResponse(
"stream callback disconnected".to_owned(),
))
}
};
self.mark_terminal(&event);
Ok(Some(event))
}
fn mark_terminal(&mut self, event: &StreamEvent) {
if matches!(event, StreamEvent::Final | StreamEvent::Error(_)) {
self.terminal_received = true;
}
}
fn drain_until_terminal(&mut self, timeout: Duration) -> bool {
let deadline = Instant::now() + timeout;
while !self.terminal_received {
let now = Instant::now();
if now >= deadline {
return false;
}
match self.receiver.recv_timeout(deadline - now) {
Ok(event) => self.mark_terminal(&event),
Err(mpsc::RecvTimeoutError::Timeout) => return false,
Err(mpsc::RecvTimeoutError::Disconnected) => return true,
}
}
true
}
}
unsafe extern "C" fn stream_callback(
callback_data: *mut c_void,
chunk: *const ffi::LiteRtLmStreamChunk,
) {
if chunk.is_null() {
return;
}
let (text, is_final, error) = unsafe {
(
ffi::litert_lm_stream_chunk_get_text(chunk),
ffi::litert_lm_stream_chunk_is_final(chunk),
ffi::litert_lm_stream_chunk_get_error(chunk),
)
};
unsafe { dispatch_stream_event(callback_data, text, is_final, error) };
}
unsafe fn dispatch_stream_event(
callback_data: *mut c_void,
chunk: *const c_char,
is_final: bool,
error_msg: *const c_char,
) {
if callback_data.is_null() {
return;
}
let state = callback_data.cast::<StreamState>();
if !error_msg.is_null() {
let message = unsafe { CStr::from_ptr(error_msg) }
.to_string_lossy()
.into_owned();
let _ = unsafe { (*state).sender.send(StreamEvent::Error(message)) };
return;
}
if !chunk.is_null() {
let text = unsafe { CStr::from_ptr(chunk) }
.to_string_lossy()
.into_owned();
if !text.is_empty() {
let _ = unsafe { (*state).sender.send(StreamEvent::Chunk(text)) };
}
}
if is_final {
let _ = unsafe { (*state).sender.send(StreamEvent::Final) };
}
}
#[cfg(test)]
pub(crate) mod tests {
use super::*;
use std::ffi::CString;
fn callback_for_test(
state: &mut StreamState,
chunk: Option<&str>,
is_final: bool,
error: Option<&str>,
) {
let chunk = chunk.map(CString::new).transpose().unwrap();
let error = error.map(CString::new).transpose().unwrap();
unsafe {
dispatch_stream_event(
(state as *mut StreamState).cast(),
chunk
.as_ref()
.map_or(std::ptr::null(), |value| value.as_ptr()),
is_final,
error
.as_ref()
.map_or(std::ptr::null(), |value| value.as_ptr()),
);
}
}
#[test]
fn callback_sends_chunk_and_final() {
let (sender, receiver) = mpsc::channel();
let input = OwnedInputs::new(&[crate::InputData::Text("hello".to_owned())]).unwrap();
let mut state = StreamState { sender, input };
callback_for_test(&mut state, Some("a"), false, None);
callback_for_test(&mut state, Some("b"), true, None);
assert_eq!(receiver.recv().unwrap(), StreamEvent::Chunk("a".to_owned()));
assert_eq!(receiver.recv().unwrap(), StreamEvent::Chunk("b".to_owned()));
assert_eq!(receiver.recv().unwrap(), StreamEvent::Final);
}
#[test]
fn callback_sends_error() {
let (sender, receiver) = mpsc::channel();
let input = OwnedInputs::new(&[crate::InputData::Text("hello".to_owned())]).unwrap();
let mut state = StreamState { sender, input };
callback_for_test(&mut state, None, false, Some("failed"));
assert_eq!(
receiver.recv().unwrap(),
StreamEvent::Error("failed".to_owned())
);
}
}