use klieo_core::agent::AgentEvent;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use tokio::sync::broadcast;
const PROGRESS_BUFFER: usize = 256;
#[derive(Clone, Default)]
pub(crate) struct ProgressHub {
channels: Arc<Mutex<HashMap<String, broadcast::Sender<AgentEvent>>>>,
}
impl ProgressHub {
pub fn register(&self, run_id: &str) -> broadcast::Sender<AgentEvent> {
let (sender, _receiver) = broadcast::channel(PROGRESS_BUFFER);
self.lock().insert(run_id.to_string(), sender.clone());
sender
}
pub fn deregister(&self, run_id: &str) {
self.lock().remove(run_id);
}
pub fn subscribe(&self, run_id: &str) -> Option<broadcast::Receiver<AgentEvent>> {
self.lock().get(run_id).map(broadcast::Sender::subscribe)
}
fn lock(&self) -> std::sync::MutexGuard<'_, HashMap<String, broadcast::Sender<AgentEvent>>> {
self.channels.lock().expect("progress hub mutex poisoned")
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn subscribe_reflects_register_and_deregister() {
let hub = ProgressHub::default();
assert!(hub.subscribe("run-1").is_none());
let _sender = hub.register("run-1");
assert!(hub.subscribe("run-1").is_some());
hub.deregister("run-1");
assert!(hub.subscribe("run-1").is_none());
}
#[test]
fn registered_sender_delivers_to_subscriber() {
let hub = ProgressHub::default();
let sender = hub.register("run-1");
let mut receiver = hub.subscribe("run-1").unwrap();
sender.send(AgentEvent::LlmCallStarted).unwrap();
assert!(matches!(
receiver.try_recv(),
Ok(AgentEvent::LlmCallStarted)
));
}
}