1use std::convert::Infallible;
4use std::str::FromStr;
5use std::time::Duration;
6
7use axum::extract::{Query, State};
8use axum::response::sse::{Event as SseEvent, KeepAlive, Sse};
9use futures_util::stream::{Stream, StreamExt};
10use serde::Deserialize;
11use serde::de::{self, Deserializer};
12use tokio_stream::wrappers::BroadcastStream;
13use uuid::Uuid;
14
15use crate::state::AppState;
16use ironflow_auth::extractor::Authenticated;
17use ironflow_engine::notify::Event;
18
19pub use ironflow_store::entities::EventKind;
20
21fn deserialize_comma_event_kinds<'de, D>(
23 deserializer: D,
24) -> Result<Option<Vec<EventKind>>, D::Error>
25where
26 D: Deserializer<'de>,
27{
28 let opt: Option<String> = Option::deserialize(deserializer)?;
29 match opt {
30 None => Ok(None),
31 Some(raw) => {
32 let kinds: Result<Vec<EventKind>, _> = raw
33 .split(',')
34 .map(|s| s.trim())
35 .filter(|s| !s.is_empty())
36 .map(EventKind::from_str)
37 .collect();
38 kinds.map(Some).map_err(de::Error::custom)
39 }
40 }
41}
42
43#[derive(Debug, Deserialize)]
58pub struct EventsQuery {
59 pub run_id: Option<Uuid>,
61 #[serde(default, deserialize_with = "deserialize_comma_event_kinds")]
63 pub types: Option<Vec<EventKind>>,
64}
65
66pub async fn events(
82 _auth: Authenticated,
83 State(state): State<AppState>,
84 Query(query): Query<EventsQuery>,
85) -> Sse<impl Stream<Item = Result<SseEvent, Infallible>>> {
86 let receiver = state.event_sender.subscribe();
87 let type_filter = query.types;
88
89 let stream = BroadcastStream::new(receiver).filter_map(move |result: Result<Event, _>| {
90 let type_filter = type_filter.clone();
91 let run_id_filter = query.run_id;
92 async move {
93 let event = result.ok()?;
94
95 if let Some(ref rid) = run_id_filter
96 && event.run_id() != Some(*rid)
97 {
98 return None;
99 }
100
101 if let Some(ref kinds) = type_filter {
102 let event_type = event.event_type();
103 if !kinds.iter().any(|k| k.as_str() == event_type) {
104 return None;
105 }
106 }
107
108 let data = serde_json::to_string(&event).ok()?;
109 let sse_event = SseEvent::default().event(event.event_type()).data(data);
110
111 Some(Ok::<_, Infallible>(sse_event))
112 }
113 });
114
115 Sse::new(stream).keep_alive(KeepAlive::new().interval(Duration::from_secs(30)))
116}
117
118#[cfg(test)]
119mod tests {
120 use std::collections::HashMap;
121 use std::sync::Arc;
122 use std::time::Duration;
123
124 use axum::Router;
125 use axum::routing::get;
126 use chrono::Utc;
127 use ironflow_auth::jwt::AccessToken;
128 use ironflow_core::providers::claude::ClaudeCodeProvider;
129 use ironflow_engine::engine::Engine;
130 use ironflow_engine::notify::{Event, RunStatusChangedEvent, UserSignedInEvent};
131 use ironflow_store::memory::InMemoryStore;
132 use ironflow_store::models::RunStatus;
133 use rust_decimal::Decimal;
134 use tokio::io::AsyncBufReadExt;
135 use tokio::io::BufReader;
136 use tokio::net::TcpListener;
137 use tokio::sync::broadcast;
138 use tokio::time::{sleep, timeout};
139 use uuid::Uuid;
140
141 use super::events;
142 use crate::state::AppState;
143
144 fn test_state() -> AppState {
145 let store = Arc::new(InMemoryStore::new());
146 let provider = Arc::new(ClaudeCodeProvider::new());
147 let engine = Arc::new(Engine::new(store.clone(), provider));
148 let jwt_config = Arc::new(ironflow_auth::jwt::JwtConfig {
149 secret: "test-secret".to_string(),
150 access_token_ttl_secs: 900,
151 refresh_token_ttl_secs: 604800,
152 cookie_domain: None,
153 cookie_secure: false,
154 });
155 let (event_sender, _) = broadcast::channel::<Event>(16);
156 AppState::new(
157 store,
158 engine,
159 jwt_config,
160 "test-worker-token".to_string(),
161 event_sender,
162 )
163 }
164
165 fn sample_run_event(run_id: Uuid) -> Event {
166 Event::RunStatusChanged(RunStatusChangedEvent {
167 run_id,
168 workflow_name: "deploy".to_string(),
169 from: RunStatus::Running,
170 to: RunStatus::Completed,
171 error: None,
172 cost_usd: Decimal::ZERO,
173 duration_ms: 1000,
174 labels: HashMap::new(),
175 at: Utc::now(),
176 })
177 }
178
179 fn sample_user_event() -> Event {
180 Event::UserSignedIn(UserSignedInEvent {
181 user_id: Uuid::now_v7(),
182 username: "alice".to_string(),
183 at: Utc::now(),
184 })
185 }
186
187 fn make_auth_token(state: &AppState) -> String {
188 let user_id = Uuid::now_v7();
189 let token = AccessToken::for_user(user_id, "testuser", false, &state.jwt_config).unwrap();
190 format!("Bearer {}", token.0)
191 }
192
193 async fn start_sse_server(state: AppState) -> (String, broadcast::Sender<Event>, String) {
195 let sender = state.event_sender.clone();
196 let auth = make_auth_token(&state);
197 let app = Router::new()
198 .route("/events", get(events))
199 .with_state(state);
200
201 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
202 let addr = listener.local_addr().unwrap().to_string();
203 tokio::spawn(async move {
204 axum::serve(listener, app).await.unwrap();
205 });
206 (addr, sender, auth)
207 }
208
209 async fn connect_sse(addr: &str, query: &str, auth: &str) -> BufReader<tokio::net::TcpStream> {
211 let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
212 let (reader, mut writer) = stream.into_split();
213
214 use tokio::io::AsyncWriteExt;
215 writer
216 .write_all(
217 format!(
218 "GET /events{query} HTTP/1.1\r\nHost: {addr}\r\nAccept: text/event-stream\r\nAuthorization: {auth}\r\n\r\n"
219 )
220 .as_bytes(),
221 )
222 .await
223 .unwrap();
224
225 BufReader::new(reader.reunite(writer).unwrap())
226 }
227
228 async fn read_until_contains(
231 reader: &mut BufReader<tokio::net::TcpStream>,
232 needle: &str,
233 dur: Duration,
234 ) -> String {
235 let mut accumulated = String::new();
236 let result = timeout(dur, async {
237 loop {
238 let mut line = String::new();
239 let n = reader.read_line(&mut line).await.unwrap();
240 if n == 0 {
241 break;
242 }
243 accumulated.push_str(&line);
244 if accumulated.contains(needle) {
245 break;
246 }
247 }
248 })
249 .await;
250 if result.is_err() {
251 panic!("timeout waiting for '{needle}' in SSE stream. Data so far:\n{accumulated}");
252 }
253 accumulated
254 }
255
256 async fn wait_for_subscriber(sender: &broadcast::Sender<Event>) {
262 let ready = timeout(Duration::from_secs(5), async {
263 while sender.receiver_count() == 0 {
264 sleep(Duration::from_millis(5)).await;
265 }
266 })
267 .await;
268 assert!(ready.is_ok(), "SSE handler did not subscribe in time");
269 }
270
271 #[tokio::test]
272 async fn sse_stream_receives_events() {
273 let state = test_state();
274 let (addr, sender, auth) = start_sse_server(state).await;
275 let mut reader = connect_sse(&addr, "", &auth).await;
276
277 wait_for_subscriber(&sender).await;
278
279 let run_id = Uuid::now_v7();
280 sender.send(sample_run_event(run_id)).unwrap();
281
282 let text =
283 read_until_contains(&mut reader, &run_id.to_string(), Duration::from_secs(5)).await;
284
285 assert!(text.contains("run_status_changed"));
286 assert!(text.contains(&run_id.to_string()));
287 }
288
289 #[tokio::test]
290 async fn sse_filters_by_run_id() {
291 let state = test_state();
292 let (addr, sender, auth) = start_sse_server(state).await;
293
294 let target_run = Uuid::now_v7();
295 let other_run = Uuid::now_v7();
296
297 let mut reader = connect_sse(&addr, &format!("?run_id={target_run}"), &auth).await;
298 wait_for_subscriber(&sender).await;
299
300 sender.send(sample_run_event(other_run)).unwrap();
301 sender.send(sample_run_event(target_run)).unwrap();
302
303 let text =
304 read_until_contains(&mut reader, &target_run.to_string(), Duration::from_secs(5)).await;
305
306 assert!(text.contains(&target_run.to_string()));
307 assert!(!text.contains(&other_run.to_string()));
308 }
309
310 #[tokio::test]
311 async fn sse_filters_by_event_type() {
312 let state = test_state();
313 let (addr, sender, auth) = start_sse_server(state).await;
314
315 let mut reader = connect_sse(&addr, "?types=user_signed_in", &auth).await;
316 wait_for_subscriber(&sender).await;
317
318 let run_id = Uuid::now_v7();
319 sender.send(sample_run_event(run_id)).unwrap();
320 sender.send(sample_user_event()).unwrap();
321
322 let text = read_until_contains(&mut reader, "user_signed_in", Duration::from_secs(5)).await;
323
324 assert!(text.contains("user_signed_in"));
325 assert!(!text.contains("run_status_changed"));
326 }
327
328 #[tokio::test]
329 async fn sse_returns_correct_content_type() {
330 let state = test_state();
331 let (addr, _sender, auth) = start_sse_server(state).await;
332 let mut reader = connect_sse(&addr, "", &auth).await;
333
334 let text =
335 read_until_contains(&mut reader, "text/event-stream", Duration::from_secs(5)).await;
336
337 assert!(text.contains("text/event-stream"));
338 }
339
340 #[tokio::test]
341 async fn sse_rejects_unauthenticated() {
342 let state = test_state();
343 let (addr, _sender, _auth) = start_sse_server(state).await;
344 let stream = tokio::net::TcpStream::connect(&addr).await.unwrap();
346 let (reader, mut writer) = stream.into_split();
347
348 use tokio::io::AsyncWriteExt;
349 writer
350 .write_all(
351 format!(
352 "GET /events HTTP/1.1\r\nHost: {addr}\r\nAccept: text/event-stream\r\n\r\n"
353 )
354 .as_bytes(),
355 )
356 .await
357 .unwrap();
358
359 let mut buf_reader = BufReader::new(reader.reunite(writer).unwrap());
360 let text = read_until_contains(&mut buf_reader, "401", Duration::from_secs(5)).await;
361
362 assert!(text.contains("401"));
363 }
364}