Skip to main content

ironflow_api/routes/
events.rs

1//! SSE endpoint for real-time event streaming.
2
3use 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
21/// Deserialize a comma-separated string into `Option<Vec<EventKind>>`.
22fn 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/// Query parameters for the SSE events endpoint.
44///
45/// Both fields are optional. When set, only matching events are streamed.
46///
47/// # Examples
48///
49/// ```
50/// use ironflow_api::routes::events::EventsQuery;
51///
52/// let query = EventsQuery {
53///     run_id: None,
54///     types: None,
55/// };
56/// ```
57#[derive(Debug, Deserialize)]
58pub struct EventsQuery {
59    /// Only stream events related to this run.
60    pub run_id: Option<Uuid>,
61    /// Comma-separated list of event types to include (e.g. `?types=run_status_changed,step_completed`).
62    #[serde(default, deserialize_with = "deserialize_comma_event_kinds")]
63    pub types: Option<Vec<EventKind>>,
64}
65
66/// `GET /api/v1/events` -- Server-Sent Events stream.
67///
68/// Streams domain events in real time. Supports optional filtering:
69/// - `?run_id=<uuid>` -- only events for that run
70/// - `?types=run_status_changed,step_completed` -- only those event types
71///
72/// Each SSE message has:
73/// - `event:` set to the event type (e.g. `run_status_changed`)
74/// - `data:` JSON-serialized event payload
75///
76/// A keep-alive comment is sent every 30 seconds.
77///
78/// # Errors
79///
80/// Returns 401 if the request is not authenticated.
81pub 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    /// Start a real TCP server and return (address, sender, auth header).
194    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    /// Connect to the SSE endpoint with auth and return a line reader.
210    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    /// Read all available data from the SSE stream until `needle` is found
229    /// in the accumulated text, or timeout.
230    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    /// Wait until the SSE handler has subscribed to the broadcast channel.
257    ///
258    /// Sending before a receiver exists makes `broadcast::Sender::send` return
259    /// `SendError`, so a fixed sleep is racy on slow runners. Poll the receiver
260    /// count instead.
261    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        // Connect without auth header
345        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}