use super::{AgentOperation, PhaseType};
use async_trait::async_trait;
use serde::Serialize;
use std::sync::Arc;
use std::time::Instant;
use tokio::sync::{broadcast, RwLock};
use tracing::{debug, error};
#[derive(Debug, Clone, Serialize)]
#[serde(tag = "type")]
pub enum ProgressEvent {
#[serde(rename = "workflow_started")]
WorkflowStarted {
job_id: String,
total_items: usize,
timestamp: std::time::SystemTime,
},
#[serde(rename = "item_complete")]
ItemComplete {
item_id: String,
total_completed: usize,
percentage: f64,
},
#[serde(rename = "agent_update")]
AgentUpdate {
agent_index: usize,
#[serde(skip)]
operation: AgentOperation,
#[serde(skip)]
timestamp: Instant,
},
#[serde(rename = "phase_change")]
PhaseChange {
#[serde(skip)]
phase: PhaseType,
#[serde(skip)]
timestamp: Instant,
},
#[serde(rename = "error")]
Error {
message: String,
failed_count: usize,
},
#[serde(rename = "message")]
Message(String),
#[serde(rename = "workflow_completed")]
WorkflowCompleted {
job_id: String,
total_processed: usize,
total_failed: usize,
duration_secs: f64,
},
#[serde(rename = "metrics")]
Metrics {
items_per_second: f64,
agent_utilization: f64,
estimated_remaining_secs: Option<f64>,
},
}
#[derive(Debug)]
pub struct StreamError {
message: String,
}
impl std::fmt::Display for StreamError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "Stream error: {}", self.message)
}
}
impl std::error::Error for StreamError {}
impl StreamError {
pub fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
}
#[async_trait]
pub trait ProgressStreamConsumer: Send + Sync {
async fn consume(&mut self, event: ProgressEvent) -> Result<(), StreamError>;
async fn on_start(&mut self) -> Result<(), StreamError> {
Ok(())
}
async fn on_end(&mut self) -> Result<(), StreamError> {
Ok(())
}
fn name(&self) -> &str;
}
pub struct ProgressStreamer {
sender: broadcast::Sender<ProgressEvent>,
consumers: Arc<RwLock<Vec<Box<dyn ProgressStreamConsumer>>>>,
is_active: Arc<RwLock<bool>>,
}
impl ProgressStreamer {
pub fn new() -> Self {
let (sender, _) = broadcast::channel(1000);
Self {
sender,
consumers: Arc::new(RwLock::new(Vec::new())),
is_active: Arc::new(RwLock::new(true)),
}
}
pub fn subscribe(&self) -> broadcast::Receiver<ProgressEvent> {
self.sender.subscribe()
}
pub async fn add_consumer(&self, mut consumer: Box<dyn ProgressStreamConsumer>) {
if let Err(e) = consumer.on_start().await {
error!("Consumer {} failed to start: {}", consumer.name(), e);
return;
}
let mut consumers = self.consumers.write().await;
consumers.push(consumer);
}
pub async fn stream_event(&self, event: ProgressEvent) {
if !*self.is_active.read().await {
return;
}
if let Err(e) = self.sender.send(event.clone()) {
debug!("No subscribers for progress event: {}", e);
}
let mut consumers = self.consumers.write().await;
let mut failed_indices = Vec::new();
for (index, consumer) in consumers.iter_mut().enumerate() {
if let Err(e) = consumer.consume(event.clone()).await {
error!("Consumer {} failed: {}", consumer.name(), e);
failed_indices.push(index);
}
}
for index in failed_indices.into_iter().rev() {
consumers.remove(index);
}
}
pub async fn stop(&self) {
*self.is_active.write().await = false;
let mut consumers = self.consumers.write().await;
for consumer in consumers.iter_mut() {
if let Err(e) = consumer.on_end().await {
error!("Consumer {} failed to end: {}", consumer.name(), e);
}
}
consumers.clear();
}
pub async fn is_active(&self) -> bool {
*self.is_active.read().await
}
pub async fn consumer_count(&self) -> usize {
self.consumers.read().await.len()
}
}
impl Default for ProgressStreamer {
fn default() -> Self {
Self::new()
}
}
pub struct JsonLinesConsumer {
writer: tokio::io::BufWriter<tokio::fs::File>,
event_count: usize,
}
impl JsonLinesConsumer {
pub async fn new(path: impl AsRef<std::path::Path>) -> Result<Self, StreamError> {
let file = tokio::fs::File::create(path)
.await
.map_err(|e| StreamError::new(format!("Failed to create file: {}", e)))?;
Ok(Self {
writer: tokio::io::BufWriter::new(file),
event_count: 0,
})
}
}
#[async_trait]
impl ProgressStreamConsumer for JsonLinesConsumer {
async fn consume(&mut self, event: ProgressEvent) -> Result<(), StreamError> {
use tokio::io::AsyncWriteExt;
let json = serde_json::to_string(&event)
.map_err(|e| StreamError::new(format!("Failed to serialize: {}", e)))?;
self.writer
.write_all(json.as_bytes())
.await
.map_err(|e| StreamError::new(format!("Write failed: {}", e)))?;
self.writer
.write_all(b"\n")
.await
.map_err(|e| StreamError::new(format!("Write failed: {}", e)))?;
self.event_count += 1;
if self.event_count.is_multiple_of(100) {
self.writer
.flush()
.await
.map_err(|e| StreamError::new(format!("Flush failed: {}", e)))?;
}
Ok(())
}
async fn on_end(&mut self) -> Result<(), StreamError> {
use tokio::io::AsyncWriteExt;
self.writer
.flush()
.await
.map_err(|e| StreamError::new(format!("Final flush failed: {}", e)))?;
Ok(())
}
fn name(&self) -> &str {
"JSON Lines Consumer"
}
}
pub struct WebSocketConsumer {
endpoint: String,
#[allow(dead_code)]
client_id: String,
}
impl WebSocketConsumer {
pub fn new(endpoint: impl Into<String>, client_id: impl Into<String>) -> Self {
Self {
endpoint: endpoint.into(),
client_id: client_id.into(),
}
}
}
#[async_trait]
impl ProgressStreamConsumer for WebSocketConsumer {
async fn consume(&mut self, event: ProgressEvent) -> Result<(), StreamError> {
debug!("Would send to WebSocket {}: {:?}", self.endpoint, event);
Ok(())
}
fn name(&self) -> &str {
"WebSocket Consumer"
}
}
pub struct MetricsAggregator {
events: Vec<ProgressEvent>,
start_time: Option<Instant>,
}
impl MetricsAggregator {
pub fn new() -> Self {
Self {
events: Vec::new(),
start_time: None,
}
}
pub fn get_metrics(&self) -> AggregatedMetrics {
let total_events = self.events.len();
let completed_items = self
.events
.iter()
.filter(|e| matches!(e, ProgressEvent::ItemComplete { .. }))
.count();
let errors = self
.events
.iter()
.filter(|e| matches!(e, ProgressEvent::Error { .. }))
.count();
let duration = self.start_time.map(|s| s.elapsed()).unwrap_or_default();
AggregatedMetrics {
total_events,
completed_items,
errors,
duration,
}
}
}
impl Default for MetricsAggregator {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl ProgressStreamConsumer for MetricsAggregator {
async fn consume(&mut self, event: ProgressEvent) -> Result<(), StreamError> {
if self.start_time.is_none() {
self.start_time = Some(Instant::now());
}
self.events.push(event);
Ok(())
}
fn name(&self) -> &str {
"Metrics Aggregator"
}
}
#[derive(Debug, Clone)]
pub struct AggregatedMetrics {
pub total_events: usize,
pub completed_items: usize,
pub errors: usize,
pub duration: std::time::Duration,
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_progress_streamer() {
let streamer = ProgressStreamer::new();
let aggregator = Box::new(MetricsAggregator::new());
streamer.add_consumer(aggregator).await;
streamer
.stream_event(ProgressEvent::Message("Test".to_string()))
.await;
streamer
.stream_event(ProgressEvent::ItemComplete {
item_id: "item1".to_string(),
total_completed: 1,
percentage: 10.0,
})
.await;
assert_eq!(streamer.consumer_count().await, 1);
streamer.stop().await;
assert!(!streamer.is_active().await);
assert_eq!(streamer.consumer_count().await, 0);
}
#[tokio::test]
async fn test_metrics_aggregator() {
let mut aggregator = MetricsAggregator::new();
aggregator
.consume(ProgressEvent::ItemComplete {
item_id: "1".to_string(),
total_completed: 1,
percentage: 50.0,
})
.await
.unwrap();
aggregator
.consume(ProgressEvent::Error {
message: "Test error".to_string(),
failed_count: 1,
})
.await
.unwrap();
let metrics = aggregator.get_metrics();
assert_eq!(metrics.total_events, 2);
assert_eq!(metrics.completed_items, 1);
assert_eq!(metrics.errors, 1);
}
#[tokio::test]
async fn test_broadcast_subscription() {
let streamer = ProgressStreamer::new();
let mut receiver = streamer.subscribe();
let event = ProgressEvent::Message("Test broadcast".to_string());
streamer.stream_event(event.clone()).await;
let received = receiver.recv().await.unwrap();
matches!(received, ProgressEvent::Message(msg) if msg == "Test broadcast");
}
}