use crate::types::{AssistantMessage, AssistantMessageEvent};
use tokio::sync::{mpsc, oneshot};
pub struct AssistantMessageEventStream {
rx: mpsc::UnboundedReceiver<AssistantMessageEvent>,
result_rx: oneshot::Receiver<AssistantMessage>,
done: bool,
}
impl AssistantMessageEventStream {
pub async fn next(&mut self) -> Option<AssistantMessageEvent> {
if self.done {
return None;
}
let event = self.rx.recv().await?;
if event.is_terminal() {
self.done = true;
}
Some(event)
}
pub async fn result(self) -> Result<AssistantMessage, RecvError> {
self.result_rx.await.map_err(|_| RecvError)
}
pub fn split(
self,
) -> (
mpsc::UnboundedReceiver<AssistantMessageEvent>,
oneshot::Receiver<AssistantMessage>,
) {
(self.rx, self.result_rx)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RecvError;
impl std::fmt::Display for RecvError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"assistant-message event stream ended without a terminal Done/Error event"
)
}
}
impl std::error::Error for RecvError {}
pub struct AssistantMessageEventStreamProducer {
tx: mpsc::UnboundedSender<AssistantMessageEvent>,
result_tx: Option<oneshot::Sender<AssistantMessage>>,
done: bool,
}
impl AssistantMessageEventStreamProducer {
pub fn push(&mut self, event: AssistantMessageEvent) -> bool {
if self.done {
return true;
}
if event.is_terminal() {
self.done = true;
if let AssistantMessageEvent::Done { message, .. } = &event {
self.fulfill_result(message.clone());
} else if let AssistantMessageEvent::Error { error, .. } = &event {
self.fulfill_result(error.clone());
}
}
match self.tx.send(event) {
Ok(()) => true,
Err(_) => false,
}
}
fn fulfill_result(&mut self, message: AssistantMessage) {
if let Some(rx) = self.result_tx.take() {
let _ = rx.send(message);
}
}
pub fn is_done(&self) -> bool {
self.done
}
pub fn close(self) {
drop(self);
}
}
pub fn create_assistant_message_event_stream() -> (
AssistantMessageEventStreamProducer,
AssistantMessageEventStream,
) {
let (tx, rx) = mpsc::unbounded_channel::<AssistantMessageEvent>();
let (result_tx, result_rx) = oneshot::channel::<AssistantMessage>();
(
AssistantMessageEventStreamProducer {
tx,
result_tx: Some(result_tx),
done: false,
},
AssistantMessageEventStream {
rx,
result_rx,
done: false,
},
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{Api, DoneReason};
use std::sync::Arc;
fn empty_partial() -> Arc<AssistantMessage> {
Arc::new(AssistantMessage::empty(Api::Faux, "faux", "faux", 0))
}
#[tokio::test]
async fn drain_deltas_then_done() {
let (mut prod, mut stream) = create_assistant_message_event_stream();
let partial = empty_partial();
prod.push(AssistantMessageEvent::Start { partial: partial.clone() });
prod.push(AssistantMessageEvent::TextStart {
content_index: 0,
partial: partial.clone(),
});
prod.push(AssistantMessageEvent::TextDelta {
content_index: 0,
delta: "hi".into(),
partial: partial.clone(),
});
prod.push(AssistantMessageEvent::TextEnd {
content_index: 0,
content: "hi".into(),
partial: partial.clone(),
});
let mut final_msg = (*partial).clone();
final_msg.stop_reason = crate::types::StopReason::Stop;
prod.push(AssistantMessageEvent::Done {
reason: DoneReason::Stop,
message: final_msg.clone(),
});
let mut tags = Vec::new();
while let Some(ev) = stream.next().await {
tags.push(ev.type_tag());
}
assert_eq!(
tags,
vec!["start", "text_start", "text_delta", "text_end", "done"]
);
let result = stream.result().await.unwrap();
assert!(matches!(result.stop_reason, crate::types::StopReason::Stop));
}
#[tokio::test]
async fn error_path_resolves_to_error_message() {
let (mut prod, stream) = create_assistant_message_event_stream();
let err_msg = AssistantMessage::terminal(
Api::Faux,
"faux",
"faux",
crate::types::StopReason::Aborted,
"cancelled",
0,
);
prod.push(AssistantMessageEvent::Error {
reason: crate::types::ErrorReason::Aborted,
error: err_msg.clone(),
});
let result = stream.result().await.unwrap();
assert!(matches!(result.stop_reason, crate::types::StopReason::Aborted));
assert_eq!(result.error_message.as_deref(), Some("cancelled"));
}
#[tokio::test]
async fn push_after_terminal_is_noop() {
let (mut prod, mut stream) = create_assistant_message_event_stream();
let msg = AssistantMessage::terminal(
Api::Faux,
"faux",
"faux",
crate::types::StopReason::Stop,
"",
0,
);
prod.push(AssistantMessageEvent::Done {
reason: DoneReason::Stop,
message: msg,
});
prod.push(AssistantMessageEvent::Start { partial: empty_partial() });
let first = stream.next().await.unwrap();
assert!(matches!(first, AssistantMessageEvent::Done { .. }));
assert!(stream.next().await.is_none());
}
}