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
66fn event_run_id(event: &Event) -> Option<Uuid> {
68 match event {
69 Event::RunCreated { run_id, .. }
70 | Event::RunStatusChanged { run_id, .. }
71 | Event::RunFailed { run_id, .. }
72 | Event::RunBudgetExceeded { run_id, .. }
73 | Event::StepCompleted { run_id, .. }
74 | Event::StepFailed { run_id, .. }
75 | Event::ApprovalRequested { run_id, .. }
76 | Event::ApprovalGranted { run_id, .. }
77 | Event::ApprovalRejected { run_id, .. }
78 | Event::LogLine { run_id, .. }
79 | Event::RetryForced { run_id, .. } => Some(*run_id),
80 Event::UserSignedIn { .. } | Event::UserSignedUp { .. } | Event::UserSignedOut { .. } => {
81 None
82 }
83 }
84}
85
86pub async fn events(
102 _auth: Authenticated,
103 State(state): State<AppState>,
104 Query(query): Query<EventsQuery>,
105) -> Sse<impl Stream<Item = Result<SseEvent, Infallible>>> {
106 let receiver = state.event_sender.subscribe();
107 let type_filter = query.types;
108
109 let stream = BroadcastStream::new(receiver).filter_map(move |result: Result<Event, _>| {
110 let type_filter = type_filter.clone();
111 let run_id_filter = query.run_id;
112 async move {
113 let event = result.ok()?;
114
115 if let Some(ref rid) = run_id_filter
116 && event_run_id(&event) != Some(*rid)
117 {
118 return None;
119 }
120
121 if let Some(ref kinds) = type_filter {
122 let event_type = event.event_type();
123 if !kinds.iter().any(|k| k.as_str() == event_type) {
124 return None;
125 }
126 }
127
128 let data = serde_json::to_string(&event).ok()?;
129 let sse_event = SseEvent::default().event(event.event_type()).data(data);
130
131 Some(Ok::<_, Infallible>(sse_event))
132 }
133 });
134
135 Sse::new(stream).keep_alive(KeepAlive::new().interval(Duration::from_secs(30)))
136}
137
138#[cfg(test)]
139mod tests {
140 use std::collections::HashMap;
141 use std::sync::Arc;
142 use std::time::Duration;
143
144 use axum::Router;
145 use axum::routing::get;
146 use chrono::Utc;
147 use ironflow_auth::jwt::AccessToken;
148 use ironflow_core::providers::claude::ClaudeCodeProvider;
149 use ironflow_engine::engine::Engine;
150 use ironflow_engine::notify::Event;
151 use ironflow_store::memory::InMemoryStore;
152 use ironflow_store::models::RunStatus;
153 use rust_decimal::Decimal;
154 use tokio::io::AsyncBufReadExt;
155 use tokio::io::BufReader;
156 use tokio::net::TcpListener;
157 use tokio::sync::broadcast;
158 use tokio::time::{sleep, timeout};
159 use uuid::Uuid;
160
161 use super::events;
162 use crate::state::AppState;
163
164 fn test_state() -> AppState {
165 let store = Arc::new(InMemoryStore::new());
166 let provider = Arc::new(ClaudeCodeProvider::new());
167 let engine = Arc::new(Engine::new(store.clone(), provider));
168 let jwt_config = Arc::new(ironflow_auth::jwt::JwtConfig {
169 secret: "test-secret".to_string(),
170 access_token_ttl_secs: 900,
171 refresh_token_ttl_secs: 604800,
172 cookie_domain: None,
173 cookie_secure: false,
174 });
175 let (event_sender, _) = broadcast::channel::<Event>(16);
176 AppState::new(
177 store,
178 engine,
179 jwt_config,
180 "test-worker-token".to_string(),
181 event_sender,
182 )
183 }
184
185 fn sample_run_event(run_id: Uuid) -> Event {
186 Event::RunStatusChanged {
187 run_id,
188 workflow_name: "deploy".to_string(),
189 from: RunStatus::Running,
190 to: RunStatus::Completed,
191 error: None,
192 cost_usd: Decimal::ZERO,
193 duration_ms: 1000,
194 labels: HashMap::new(),
195 at: Utc::now(),
196 }
197 }
198
199 fn sample_user_event() -> Event {
200 Event::UserSignedIn {
201 user_id: Uuid::now_v7(),
202 username: "alice".to_string(),
203 at: Utc::now(),
204 }
205 }
206
207 fn make_auth_token(state: &AppState) -> String {
208 let user_id = Uuid::now_v7();
209 let token = AccessToken::for_user(user_id, "testuser", false, &state.jwt_config).unwrap();
210 format!("Bearer {}", token.0)
211 }
212
213 async fn start_sse_server(state: AppState) -> (String, broadcast::Sender<Event>, String) {
215 let sender = state.event_sender.clone();
216 let auth = make_auth_token(&state);
217 let app = Router::new()
218 .route("/events", get(events))
219 .with_state(state);
220
221 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
222 let addr = listener.local_addr().unwrap().to_string();
223 tokio::spawn(async move {
224 axum::serve(listener, app).await.unwrap();
225 });
226 (addr, sender, auth)
227 }
228
229 async fn connect_sse(addr: &str, query: &str, auth: &str) -> BufReader<tokio::net::TcpStream> {
231 let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
232 let (reader, mut writer) = stream.into_split();
233
234 use tokio::io::AsyncWriteExt;
235 writer
236 .write_all(
237 format!(
238 "GET /events{query} HTTP/1.1\r\nHost: {addr}\r\nAccept: text/event-stream\r\nAuthorization: {auth}\r\n\r\n"
239 )
240 .as_bytes(),
241 )
242 .await
243 .unwrap();
244
245 BufReader::new(reader.reunite(writer).unwrap())
246 }
247
248 async fn read_until_contains(
251 reader: &mut BufReader<tokio::net::TcpStream>,
252 needle: &str,
253 dur: Duration,
254 ) -> String {
255 let mut accumulated = String::new();
256 let result = timeout(dur, async {
257 loop {
258 let mut line = String::new();
259 let n = reader.read_line(&mut line).await.unwrap();
260 if n == 0 {
261 break;
262 }
263 accumulated.push_str(&line);
264 if accumulated.contains(needle) {
265 break;
266 }
267 }
268 })
269 .await;
270 if result.is_err() {
271 panic!("timeout waiting for '{needle}' in SSE stream. Data so far:\n{accumulated}");
272 }
273 accumulated
274 }
275
276 #[tokio::test]
277 async fn sse_stream_receives_events() {
278 let state = test_state();
279 let (addr, sender, auth) = start_sse_server(state).await;
280 let mut reader = connect_sse(&addr, "", &auth).await;
281
282 sleep(Duration::from_millis(50)).await;
283
284 let run_id = Uuid::now_v7();
285 sender.send(sample_run_event(run_id)).unwrap();
286
287 let text =
288 read_until_contains(&mut reader, &run_id.to_string(), Duration::from_secs(5)).await;
289
290 assert!(text.contains("run_status_changed"));
291 assert!(text.contains(&run_id.to_string()));
292 }
293
294 #[tokio::test]
295 async fn sse_filters_by_run_id() {
296 let state = test_state();
297 let (addr, sender, auth) = start_sse_server(state).await;
298
299 let target_run = Uuid::now_v7();
300 let other_run = Uuid::now_v7();
301
302 let mut reader = connect_sse(&addr, &format!("?run_id={target_run}"), &auth).await;
303 sleep(Duration::from_millis(50)).await;
304
305 sender.send(sample_run_event(other_run)).unwrap();
306 sender.send(sample_run_event(target_run)).unwrap();
307
308 let text =
309 read_until_contains(&mut reader, &target_run.to_string(), Duration::from_secs(5)).await;
310
311 assert!(text.contains(&target_run.to_string()));
312 assert!(!text.contains(&other_run.to_string()));
313 }
314
315 #[tokio::test]
316 async fn sse_filters_by_event_type() {
317 let state = test_state();
318 let (addr, sender, auth) = start_sse_server(state).await;
319
320 let mut reader = connect_sse(&addr, "?types=user_signed_in", &auth).await;
321 sleep(Duration::from_millis(50)).await;
322
323 let run_id = Uuid::now_v7();
324 sender.send(sample_run_event(run_id)).unwrap();
325 sender.send(sample_user_event()).unwrap();
326
327 let text = read_until_contains(&mut reader, "user_signed_in", Duration::from_secs(5)).await;
328
329 assert!(text.contains("user_signed_in"));
330 assert!(!text.contains("run_status_changed"));
331 }
332
333 #[tokio::test]
334 async fn sse_returns_correct_content_type() {
335 let state = test_state();
336 let (addr, _sender, auth) = start_sse_server(state).await;
337 let mut reader = connect_sse(&addr, "", &auth).await;
338
339 let text =
340 read_until_contains(&mut reader, "text/event-stream", Duration::from_secs(5)).await;
341
342 assert!(text.contains("text/event-stream"));
343 }
344
345 #[tokio::test]
346 async fn sse_rejects_unauthenticated() {
347 let state = test_state();
348 let (addr, _sender, _auth) = start_sse_server(state).await;
349 let stream = tokio::net::TcpStream::connect(&addr).await.unwrap();
351 let (reader, mut writer) = stream.into_split();
352
353 use tokio::io::AsyncWriteExt;
354 writer
355 .write_all(
356 format!(
357 "GET /events HTTP/1.1\r\nHost: {addr}\r\nAccept: text/event-stream\r\n\r\n"
358 )
359 .as_bytes(),
360 )
361 .await
362 .unwrap();
363
364 let mut buf_reader = BufReader::new(reader.reunite(writer).unwrap());
365 let text = read_until_contains(&mut buf_reader, "401", Duration::from_secs(5)).await;
366
367 assert!(text.contains("401"));
368 }
369}