use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use crate::error::Result;
use crate::harness::context::{RunConfig, RunContext};
use crate::harness::events::{EventListener, EventRecord, EventSink};
use crate::harness::message::Message;
use crate::harness::middleware::AgentRun;
use crate::harness::runtime::AgentHarness;
use super::AgentLoopResult;
#[derive(Clone, Debug)]
pub enum AgentStreamItem {
Event(EventRecord),
Completed(Box<AgentRun>),
Failed(String),
}
struct ChannelListener {
tx: tokio::sync::mpsc::UnboundedSender<EventRecord>,
}
impl EventListener for ChannelListener {
fn on_event(&self, record: &EventRecord) {
let _ = self.tx.send(record.clone());
}
}
struct ChannelListenerGuard {
events: EventSink,
listener: Arc<dyn EventListener>,
}
impl Drop for ChannelListenerGuard {
fn drop(&mut self) {
let _ = self.events.unsubscribe(&self.listener);
}
}
fn terminal_item(result: Result<AgentLoopResult>) -> AgentStreamItem {
match result {
Ok(loop_result) => AgentStreamItem::Completed(Box::new(loop_result.run)),
Err(error) => AgentStreamItem::Failed(error.to_string()),
}
}
enum Phase<'a> {
Running {
run_fut: Pin<Box<dyn Future<Output = Result<AgentLoopResult>> + Send + 'a>>,
listener_guard: ChannelListenerGuard,
},
Draining {
terminal: AgentStreamItem,
listener_guard: ChannelListenerGuard,
},
Done,
}
impl<State: Send + Sync, Ctx: Send + Sync + 'static> AgentHarness<State, Ctx> {
pub fn invoke_stream<'a>(
&'a self,
state: &'a State,
ctx_data: Ctx,
config: RunConfig,
input: Vec<Message>,
) -> impl futures::Stream<Item = AgentStreamItem> + Send + 'a {
let ctx = RunContext::new(config, ctx_data);
self.invoke_stream_in_context(state, ctx, input)
}
pub fn invoke_stream_in_context<'a>(
&'a self,
state: &'a State,
ctx: RunContext<Ctx>,
input: Vec<Message>,
) -> impl futures::Stream<Item = AgentStreamItem> + Send + 'a {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
let listener: Arc<dyn EventListener> = Arc::new(ChannelListener { tx });
ctx.events.subscribe(listener.clone());
let listener_guard = ChannelListenerGuard {
events: ctx.events.clone(),
listener,
};
let run_fut: Pin<Box<dyn Future<Output = Result<AgentLoopResult>> + Send + 'a>> =
Box::pin(self.invoke_streaming_in_context_with_status(state, ctx, input));
futures::stream::unfold(
(
Phase::Running {
run_fut,
listener_guard,
},
rx,
),
|(phase, mut rx)| async move {
match phase {
Phase::Running {
mut run_fut,
listener_guard,
} => {
tokio::select! {
biased;
maybe = rx.recv() => match maybe {
Some(record) => {
Some((
AgentStreamItem::Event(record),
(
Phase::Running {
run_fut,
listener_guard,
},
rx,
),
))
}
None => {
let terminal = terminal_item(run_fut.await);
drop(listener_guard);
Some((terminal, (Phase::Done, rx)))
}
},
result = &mut run_fut => {
let terminal = terminal_item(result);
match rx.try_recv() {
Ok(record) => Some((
AgentStreamItem::Event(record),
(
Phase::Draining {
terminal,
listener_guard,
},
rx,
),
)),
Err(_) => {
drop(listener_guard);
Some((terminal, (Phase::Done, rx)))
}
}
}
}
}
Phase::Draining {
terminal,
listener_guard,
} => match rx.try_recv() {
Ok(record) => Some((
AgentStreamItem::Event(record),
(
Phase::Draining {
terminal,
listener_guard,
},
rx,
),
)),
Err(_) => {
drop(listener_guard);
Some((terminal, (Phase::Done, rx)))
}
},
Phase::Done => None,
}
},
)
}
}