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/// Extract the `run_id` from an event, if the variant carries one.
67fn 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
86/// `GET /api/v1/events` -- Server-Sent Events stream.
87///
88/// Streams domain events in real time. Supports optional filtering:
89/// - `?run_id=<uuid>` -- only events for that run
90/// - `?types=run_status_changed,step_completed` -- only those event types
91///
92/// Each SSE message has:
93/// - `event:` set to the event type (e.g. `run_status_changed`)
94/// - `data:` JSON-serialized event payload
95///
96/// A keep-alive comment is sent every 30 seconds.
97///
98/// # Errors
99///
100/// Returns 401 if the request is not authenticated.
101pub 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    /// Start a real TCP server and return (address, sender, auth header).
214    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    /// Connect to the SSE endpoint with auth and return a line reader.
230    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    /// Read all available data from the SSE stream until `needle` is found
249    /// in the accumulated text, or timeout.
250    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        // Connect without auth header
350        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}