use crate::AppState; use crate::state::GenericEvent; use axum::extract::{Query, State}; use axum::response::IntoResponse; use axum::response::sse::{Event, KeepAlive, Sse}; use std::collections::HashMap; use std::sync::Arc; use tokio_stream::StreamExt; use tokio_stream::wrappers::BroadcastStream; pub async fn wait_for_event_handler( State(state): State>, Query(params): Query>, ) -> impl IntoResponse { let topic = params.get("topic").cloned(); let session_id = params.get("session_id").cloned(); let timeout_duration = params .get("timeout_ms") .and_then(|ms| ms.parse::().ok()) .map(tokio::time::Duration::from_millis) .or_else(|| { params .get("timeout") .and_then(|secs| secs.parse::().ok()) .map(tokio::time::Duration::from_secs) }); let mut rx = state.handler.state.event_bus_tx.subscribe(); let wait_future = async { loop { match rx.recv().await { Ok(event) => { let topic_matches = topic.as_ref().is_none_or(|t| t == &event.topic); let session_matches = session_id .as_ref() .is_none_or(|s| Some(s) == event.session_id.as_ref()); if topic_matches && session_matches { return Ok(event); } } Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => { tracing::warn!("Event bus receiver lagged by {} messages; continuing wait.", skipped); continue; } Err(tokio::sync::broadcast::error::RecvError::Closed) => { return Err("Event bus closed"); } } } }; match timeout_duration { Some(dur) => match tokio::time::timeout(dur, wait_future).await { Ok(Ok(event)) => (axum::http::StatusCode::OK, axum::Json(event)).into_response(), Ok(Err(err)) => ( axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(GenericEvent { topic: "error".to_string(), session_id: None, payload: serde_json::json!({ "error": err }), }), ) .into_response(), Err(_) => ( axum::http::StatusCode::REQUEST_TIMEOUT, axum::Json(GenericEvent { topic: "timeout".to_string(), session_id: None, payload: serde_json::json!({ "error": "timed out waiting for event", "topic": topic }), }), ) .into_response(), }, None => match wait_future.await { Ok(event) => (axum::http::StatusCode::OK, axum::Json(event)).into_response(), Err(err) => ( axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(GenericEvent { topic: "error".to_string(), session_id: None, payload: serde_json::json!({ "error": err }), }), ) .into_response(), }, } } pub async fn post_event_handler( State(state): State>, axum::Json(event): axum::Json, ) -> impl IntoResponse { let _ = state.handler.state.event_bus_tx.send(event); axum::Json(serde_json::json!({"status": "ok"})) } /// Real-time Server-Sent Events (SSE) stream for agent execution events and observability (ADR-0110) pub async fn sse_events_handler( State(state): State>, Query(params): Query>, ) -> impl IntoResponse { let topic_filter = params.get("topic").cloned(); let session_filter = params.get("session_id").cloned(); let rx = state.handler.state.event_bus_tx.subscribe(); let stream = BroadcastStream::new(rx).filter_map(move |msg| match msg { Ok(event) => { let topic_matches = topic_filter.as_ref().is_none_or(|t| t == &event.topic); let session_matches = session_filter .as_ref() .is_none_or(|s| Some(s) == event.session_id.as_ref()); if topic_matches && session_matches { let json_data = serde_json::to_string(&event).unwrap_or_default(); Some(Ok::<_, std::convert::Infallible>( Event::default().event(event.topic).data(json_data), )) } else { None } } Err(_) => None, }); Sse::new(stream).keep_alive(KeepAlive::default()) } #[cfg(test)] mod tests { use super::*; use crate::router::MemoryHandler; use crate::state::MemoryState; use axum::extract::Query; use axum::extract::State; use std::collections::HashMap; use std::sync::RwLock; use std::sync::atomic::AtomicUsize; use tempfile::tempdir; #[tokio::test] async fn test_events_wait_and_post() { let dir = tempdir().unwrap(); let mem_state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); let (shutdown_tx, _) = tokio::sync::oneshot::channel(); let app_state = Arc::new(AppState { handler: Arc::new(MemoryHandler::new(mem_state.clone())), clients: RwLock::new(HashMap::new()), next_id: AtomicUsize::new(1), shutdown_tx: std::sync::Mutex::new(Some(shutdown_tx)), }); // Start wait_for_event in a background task let app_state_clone = app_state.clone(); let mut params = HashMap::new(); params.insert("topic".to_string(), "test_topic".to_string()); params.insert("session_id".to_string(), "123".to_string()); let wait_task = tokio::spawn(async move { let res = wait_for_event_handler(State(app_state_clone), Query(params)).await; // axum::Json is returned, we need to extract it somehow, but just returning is enough for testing res }); // Yield slightly to ensure the wait task has subscribed tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; // Post an event that shouldn't match let unmatched_event = GenericEvent { topic: "wrong_topic".to_string(), session_id: Some("123".to_string()), payload: serde_json::json!({}), }; post_event_handler(State(app_state.clone()), axum::Json(unmatched_event)).await; // Post the matching event let matched_event = GenericEvent { topic: "test_topic".to_string(), session_id: Some("123".to_string()), payload: serde_json::json!({"foo": "bar"}), }; post_event_handler(State(app_state.clone()), axum::Json(matched_event)).await; // Wait for the wait task to complete let _ = wait_task.await.unwrap(); } }