use std::io;
use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel};
use super::EventSink;
use crate::event::Event;
#[derive(Debug, Clone, PartialEq)]
pub struct TaggedEvent<T> {
pub tag: T,
pub event: Event,
}
pub struct EventFanIn<T> {
sender: UnboundedSender<TaggedEvent<T>>,
receiver: UnboundedReceiver<TaggedEvent<T>>,
}
impl<T> std::fmt::Debug for EventFanIn<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EventFanIn").finish_non_exhaustive()
}
}
impl<T> Default for EventFanIn<T> {
fn default() -> Self {
Self::new()
}
}
impl<T> EventFanIn<T> {
pub fn new() -> Self {
let (sender, receiver) = unbounded_channel();
Self { sender, receiver }
}
pub fn sink(&self, tag: T) -> TaggedSink<T> {
TaggedSink {
tag,
sender: self.sender.clone(),
}
}
pub fn into_events(self) -> MergedEvents<T> {
MergedEvents {
receiver: self.receiver,
}
}
}
pub struct TaggedSink<T> {
tag: T,
sender: UnboundedSender<TaggedEvent<T>>,
}
impl<T> std::fmt::Debug for TaggedSink<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TaggedSink").finish_non_exhaustive()
}
}
impl<T> TaggedSink<T> {
pub fn tag(&self) -> &T {
&self.tag
}
}
impl<T: Clone + Send + 'static> EventSink for TaggedSink<T> {
fn emit(&mut self, event: Event) -> io::Result<()> {
self.sender
.send(TaggedEvent {
tag: self.tag.clone(),
event,
})
.map_err(|_| {
io::Error::new(
io::ErrorKind::BrokenPipe,
"the merged event stream was dropped",
)
})
}
}
pub struct MergedEvents<T> {
receiver: UnboundedReceiver<TaggedEvent<T>>,
}
impl<T> std::fmt::Debug for MergedEvents<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MergedEvents").finish_non_exhaustive()
}
}
impl<T> MergedEvents<T> {
pub async fn recv(&mut self) -> Option<TaggedEvent<T>> {
self.receiver.recv().await
}
pub fn blocking_recv(&mut self) -> Option<TaggedEvent<T>> {
self.receiver.blocking_recv()
}
pub fn drain(&mut self) -> Vec<TaggedEvent<T>> {
let mut arrived = Vec::new();
while let Ok(tagged) = self.receiver.try_recv() {
arrived.push(tagged);
}
arrived
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
use crate::event::RunOutcome;
fn delta(text: &str) -> Event {
Event::AssistantDelta {
text: text.to_string(),
}
}
#[test]
fn tags_say_which_run_an_event_came_from() {
let fan = EventFanIn::new();
let mut planner = fan.sink("planner");
let mut reviewer = fan.sink("reviewer");
let mut merged = fan.into_events();
planner.emit(delta("plan")).expect("emits");
reviewer.emit(delta("review")).expect("emits");
planner.emit(delta("more plan")).expect("emits");
let arrived: Vec<_> = merged
.drain()
.into_iter()
.map(|tagged| (tagged.tag, tagged.event))
.collect();
assert_eq!(
arrived,
vec![
("planner", delta("plan")),
("reviewer", delta("review")),
("planner", delta("more plan")),
]
);
assert_eq!(planner.tag(), &"planner");
}
#[test]
fn a_sink_travels_to_another_thread_and_the_consumer_needs_no_runtime() {
let fan = EventFanIn::new();
let mut sink = fan.sink("worker");
let mut merged = fan.into_events();
let run = std::thread::spawn(move || {
sink.emit(delta("hello")).expect("emits");
sink.emit(Event::RunFinished {
outcome: RunOutcome::Ok,
stopped_by: None,
})
.expect("emits");
});
let mut seen = Vec::new();
while let Some(tagged) = merged.blocking_recv() {
seen.push(tagged);
}
run.join().expect("the emitting thread finishes");
assert_eq!(seen.len(), 2);
assert!(seen.iter().all(|tagged| tagged.tag == "worker"));
assert_eq!(seen[0].event, delta("hello"));
}
#[test]
fn a_fan_in_nobody_minted_from_is_a_stream_that_is_already_over() {
let mut merged = EventFanIn::<&str>::new().into_events();
assert!(merged.blocking_recv().is_none());
}
#[tokio::test]
async fn the_stream_ends_when_the_last_sink_is_dropped() {
let fan = EventFanIn::new();
let first = fan.sink("first");
let second = fan.sink("second");
let mut merged = fan.into_events();
drop(first);
assert!(
tokio::time::timeout(Duration::from_millis(50), merged.recv())
.await
.is_err(),
"a live sink keeps the stream open"
);
drop(second);
assert!(
merged.recv().await.is_none(),
"the last sink going closes the stream"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn each_run_keeps_its_own_order_under_concurrent_emission() {
const RUNS: usize = 4;
const EVENTS: usize = 250;
let fan = EventFanIn::new();
let sinks: Vec<_> = (0..RUNS).map(|run| fan.sink(run)).collect();
let mut merged = fan.into_events();
let emitters: Vec<_> = sinks
.into_iter()
.map(|mut sink| {
tokio::spawn(async move {
for step in 0..EVENTS {
sink.emit(delta(&step.to_string())).expect("emits");
tokio::task::yield_now().await;
}
})
})
.collect();
let mut seen: Vec<Vec<String>> = vec![Vec::new(); RUNS];
while let Some(tagged) = merged.recv().await {
let Event::AssistantDelta { text } = tagged.event else {
panic!("only deltas were emitted");
};
seen[tagged.tag].push(text);
}
for emitter in emitters {
emitter.await.expect("emitter finishes");
}
let expected: Vec<String> = (0..EVENTS).map(|step| step.to_string()).collect();
for (run, texts) in seen.iter().enumerate() {
assert_eq!(texts, &expected, "run {run} arrived out of order");
}
}
#[test]
fn a_consumer_that_never_reads_does_not_stall_the_run_feeding_it() {
const EVENTS: usize = 10_000;
let fan = EventFanIn::new();
let mut sink = fan.sink("chatty");
let mut merged = fan.into_events();
for step in 0..EVENTS {
sink.emit(delta(&step.to_string())).expect("emits");
}
drop(sink);
assert_eq!(
merged.drain().len(),
EVENTS,
"nothing is lost while nobody is reading"
);
}
#[test]
fn a_dropped_consumer_fails_the_sink_rather_than_blocking_it() {
let fan = EventFanIn::new();
let mut sink = fan.sink("orphan");
drop(fan.into_events());
let error = sink
.emit(delta("nobody is listening"))
.expect_err("the consumer is gone");
assert_eq!(error.kind(), io::ErrorKind::BrokenPipe);
assert!(sink.emit(delta("still nobody")).is_err());
}
}