test(server): make global attach filter deterministic

This commit is contained in:
Bryan Helmkamp 2026-04-22 16:04:33 -04:00
parent 090e1022ed
commit 27cc1bb6c2
No known key found for this signature in database
2 changed files with 69 additions and 42 deletions

View file

@ -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::<Event, std::convert::Infallible>)
}
Err(_) => None,
}
filtered_global_events(state.global_event_tx.subscribe(), run_filter).filter_map(|event| {
sse_event_from_store(&event).map(Ok::<Event, std::convert::Infallible>)
});
Sse::new(stream)
@ -1553,6 +1545,16 @@ async fn attach_events(
.into_response()
}
fn filtered_global_events(
event_rx: broadcast::Receiver<EventEnvelope>,
run_filter: Option<HashSet<RunId>>,
) -> impl tokio_stream::Stream<Item = EventEnvelope> {
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<RunId>,
rows: Vec<PruneRunEntry>,
@ -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::<Vec<_>>().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());

View file

@ -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}"
);
}