191 lines
7.0 KiB
Rust
191 lines
7.0 KiB
Rust
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<Arc<AppState>>,
|
|
Query(params): Query<HashMap<String, String>>,
|
|
) -> 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::<u64>().ok())
|
|
.map(tokio::time::Duration::from_millis)
|
|
.or_else(|| {
|
|
params
|
|
.get("timeout")
|
|
.and_then(|secs| secs.parse::<u64>().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<Arc<AppState>>,
|
|
axum::Json(event): axum::Json<GenericEvent>,
|
|
) -> 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<Arc<AppState>>,
|
|
Query(params): Query<HashMap<String, String>>,
|
|
) -> 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<GenericEvent> 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();
|
|
}
|
|
}
|