diff --git a/lib/crates/fabro-server/src/server.rs b/lib/crates/fabro-server/src/server.rs index 9714a5d39..4075faae0 100644 --- a/lib/crates/fabro-server/src/server.rs +++ b/lib/crates/fabro-server/src/server.rs @@ -1536,16 +1536,8 @@ async fn attach_events( }; let stream = - BroadcastStream::new(state.global_event_tx.subscribe()).filter_map(move |result| { - match result { - Ok(event) => { - if !event_matches_run_filter(&event, run_filter.as_ref()) { - return None; - } - sse_event_from_store(&event).map(Ok::) - } - Err(_) => None, - } + filtered_global_events(state.global_event_tx.subscribe(), run_filter).filter_map(|event| { + sse_event_from_store(&event).map(Ok::) }); Sse::new(stream) @@ -1553,6 +1545,16 @@ async fn attach_events( .into_response() } +fn filtered_global_events( + event_rx: broadcast::Receiver, + run_filter: Option>, +) -> impl tokio_stream::Stream { + BroadcastStream::new(event_rx).filter_map(move |result| match result { + Ok(event) if event_matches_run_filter(&event, run_filter.as_ref()) => Some(event), + Ok(_) | Err(_) => None, + }) +} + struct PrunePlan { run_ids: Vec, rows: Vec, @@ -7342,11 +7344,13 @@ mod tests { use axum::body::Body; use axum::http::{Request, header}; + use chrono::Utc; use fabro_interview::{AnswerValue, ControlInterviewer, Interviewer, Question, QuestionType}; use fabro_model::Provider; use fabro_types::settings::ServerAuthMethod; use fabro_types::{InterviewQuestionRecord, InterviewQuestionType, RunBlobId, RunId, fixtures}; use serde_json::json; + use tokio_stream::StreamExt as _; use tower::ServiceExt; use super::*; @@ -8137,6 +8141,27 @@ allowed_usernames = ["octocat"] run_store.append_event(&payload).await.unwrap(); } + fn test_event_envelope(seq: u32, run_id: RunId, body: EventBody) -> EventEnvelope { + EventEnvelope { + seq, + event: RunEvent { + id: format!("evt-{seq}"), + ts: Utc::now(), + run_id, + node_id: None, + node_label: None, + stage_id: None, + parallel_group_id: None, + parallel_branch_id: None, + session_id: None, + parent_session_id: None, + tool_call_id: None, + actor: None, + body, + }, + } + } + #[tokio::test] async fn test_model_unknown_returns_404() { let app = test_app_with(); @@ -11187,6 +11212,36 @@ timeout = "30s" assert!(matches!(sandbox_id, "sb-first" | "sb-second")); } + #[tokio::test] + async fn filtered_global_events_streams_only_matching_run_ids() { + let run_one = fixtures::RUN_1; + let run_two = fixtures::RUN_2; + let (event_tx, _) = broadcast::channel(8); + + let stream = filtered_global_events(event_tx.subscribe(), Some(HashSet::from([run_one]))); + + event_tx + .send(test_event_envelope( + 1, + run_two, + EventBody::RunQueued(Default::default()), + )) + .unwrap(); + event_tx + .send(test_event_envelope( + 2, + run_one, + EventBody::RunQueued(Default::default()), + )) + .unwrap(); + drop(event_tx); + + let events = stream.collect::>().await; + assert_eq!(events.len(), 1); + assert_eq!(events[0].seq, 2); + assert_eq!(events[0].event.run_id, run_one); + } + #[test] fn validate_github_slug_accepts_real_names() { assert!(super::validate_github_slug("owner", "anthropic", 39).is_ok()); diff --git a/lib/crates/fabro-server/tests/it/api/system.rs b/lib/crates/fabro-server/tests/it/api/system.rs index bcf066220..184eb0625 100644 --- a/lib/crates/fabro-server/tests/it/api/system.rs +++ b/lib/crates/fabro-server/tests/it/api/system.rs @@ -4,7 +4,6 @@ )] use std::path::PathBuf; -use std::time::Duration; use axum::body::Body; use axum::http::{Request, StatusCode}; @@ -13,9 +12,7 @@ use fabro_types::RunId; use fabro_types::settings::SettingsLayer; use fabro_types::settings::interp::InterpString; use fabro_types::settings::server::{ServerLayer, ServerStorageLayer}; -use http_body_util::BodyExt; use tempfile::tempdir; -use tokio::time::timeout; use tower::ServiceExt; use crate::helpers::{ @@ -263,22 +260,20 @@ async fn prune_runs_supports_dry_run_and_deletion() { } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] -async fn attach_events_streams_only_matching_run_ids() { +async fn attach_events_returns_sse_stream() { let (_temp, settings, _storage_dir) = temp_storage_settings(); let app = test_app_with_scheduler(test_app_state_with_options(settings, 5)); - - let run_one = create_run(&app, minimal_manifest_json_with_dry_run(MINIMAL_DOT)).await; - let run_two = create_run(&app, minimal_manifest_json_with_dry_run(MINIMAL_DOT)).await; + let run_id = RunId::new(); let request = Request::builder() .method("GET") - .uri(api(&format!("/attach?run_id={run_one}"))) + .uri(api(&format!("/attach?run_id={run_id}"))) .body(Body::empty()) .unwrap(); let response = checked_response( app.clone().oneshot(request).await.unwrap(), StatusCode::OK, - format!("GET /api/v1/attach?run_id={run_one}"), + format!("GET /api/v1/attach?run_id={run_id}"), ) .await; let content_type = response @@ -288,27 +283,4 @@ async fn attach_events_streams_only_matching_run_ids() { .to_str() .unwrap(); assert!(content_type.contains("text/event-stream")); - - start_run(&app, &run_one).await; - start_run(&app, &run_two).await; - - let mut body = response.into_body(); - let mut sse_data = String::new(); - while let Ok(Some(Ok(frame))) = timeout(Duration::from_secs(2), body.frame()).await { - if let Some(data) = frame.data_ref() { - sse_data.push_str(&String::from_utf8_lossy(data)); - if sse_data.contains(&run_one) { - break; - } - } - } - - assert!( - sse_data.contains(&run_one), - "expected filtered stream data: {sse_data}" - ); - assert!( - !sse_data.contains(&run_two), - "filtered stream should exclude non-matching run ids: {sse_data}" - ); }