use std::sync::{
Arc,
atomic::{AtomicU64, Ordering},
};
use futures_util::{Stream, StreamExt};
use tokio::{sync::Notify, task::AbortHandle};
use crate::{
Error, Result,
agent::types::{ConversationResponse, ConversationStreamEvent},
};
pub async fn drive_conversation_stream<S, F>(
mut stream: S,
mut on_event: F,
) -> Result<ConversationResponse>
where
S: Stream<Item = Result<ConversationStreamEvent>> + Send + Unpin,
F: FnMut(ConversationStreamEvent) + Send,
{
let mut final_response = None;
while let Some(event) = stream.next().await {
let event = event?;
match &event {
ConversationStreamEvent::WorkflowFinished(resp)
| ConversationStreamEvent::HumanInteractionRequired(resp) => {
final_response = Some(resp.clone());
}
_ => {}
}
on_event(event);
}
final_response.ok_or(Error::ConversationStreamEnded)
}
pub struct ConversationStreamIter(std::sync::mpsc::Receiver<Result<ConversationStreamEvent>>);
impl Iterator for ConversationStreamIter {
type Item = Result<ConversationStreamEvent>;
fn next(&mut self) -> Option<Self::Item> {
self.0.recv().ok()
}
}
pub fn conversation_stream_iter(
stream: impl Stream<Item = Result<ConversationStreamEvent>> + Send + 'static,
) -> ConversationStreamIter {
let (tx, rx) = std::sync::mpsc::channel();
let mut stream = Box::pin(stream);
crate::runtime_handle().spawn(async move {
while let Some(item) = stream.next().await {
if tx.send(item).is_err() {
break; }
}
});
ConversationStreamIter(rx)
}
pub struct ConversationStreamSubscription {
demand: Arc<(AtomicU64, Notify)>,
abort: AbortHandle,
}
impl ConversationStreamSubscription {
pub fn spawn<S, F1, F2, F3>(stream: S, on_next: F1, on_error: F2, on_complete: F3) -> Self
where
S: Stream<Item = Result<ConversationStreamEvent>> + Send + 'static,
F1: Fn(ConversationStreamEvent) + Send + Sync + 'static,
F2: FnOnce(Error) + Send + 'static,
F3: FnOnce() + Send + 'static,
{
let demand = Arc::new((AtomicU64::new(0), Notify::new()));
let demand2 = demand.clone();
let handle = crate::runtime_handle().spawn(async move {
let mut stream = Box::pin(stream);
loop {
while demand2.0.load(Ordering::Acquire) == 0 {
demand2.1.notified().await;
}
match stream.next().await {
Some(Ok(event)) => {
demand2.0.fetch_sub(1, Ordering::AcqRel);
on_next(event);
}
Some(Err(err)) => {
on_error(err);
break;
}
None => {
on_complete();
break;
}
}
}
});
Self {
demand,
abort: handle.abort_handle(),
}
}
pub fn request(&self, n: u64) {
self.demand.0.fetch_add(n, Ordering::AcqRel);
self.demand.1.notify_one();
}
pub fn cancel(&self) {
self.abort.abort();
}
}
#[cfg(test)]
mod tests {
use futures_util::stream;
use super::*;
use crate::agent::types::{ChatFinishedPayload, ChatStartedPayload, Interrupt};
#[tokio::test]
async fn drive_conversation_stream_terminates_on_human_interaction_required() {
let interrupt_resp = ConversationResponse::from_stream_interrupt(
Some(("ct_1".to_string(), "1".to_string())),
Interrupt {
node_id: "n_ask_human".to_string(),
tool_call_id: "call_1".to_string(),
questions: vec![],
message_id: 1,
chat_id: 1,
},
);
let events: Vec<Result<ConversationStreamEvent>> = vec![
Ok(ConversationStreamEvent::ChatStarted(ChatStartedPayload {
chat_uid: "ct_1".to_string(),
message_id: "1".to_string(),
})),
Ok(ConversationStreamEvent::HumanInteractionRequired(
interrupt_resp,
)),
Ok(ConversationStreamEvent::ChatFinished(
ChatFinishedPayload::default(),
)),
];
let mut seen = 0;
let resp = drive_conversation_stream(stream::iter(events), |_| seen += 1)
.await
.unwrap();
assert_eq!(seen, 3);
assert_eq!(
resp.status,
crate::agent::types::ConversationStatus::Interrupted
);
assert!(resp.interrupt.is_some());
}
}