From f6b8d1acdb826bbce2aacbf8e19c8b30b33b6983 Mon Sep 17 00:00:00 2001 From: Bryan Helmkamp Date: Sat, 2 May 2026 10:02:12 -0400 Subject: [PATCH] Fix principal auth gap regressions --- .github/workflows/typescript.yml | 1 + docs/internal/logging-strategy.md | 2 + .../fabro-api/tests/principal_round_trip.rs | 38 ++- .../fabro-cli/src/commands/run/attach.rs | 2 +- .../fabro-cli/src/commands/run/runner.rs | 25 +- .../fabro-interview/src/auto_approve.rs | 4 +- lib/crates/fabro-interview/src/callback.rs | 7 +- lib/crates/fabro-interview/src/lib.rs | 2 +- lib/crates/fabro-interview/src/queue.rs | 4 +- lib/crates/fabro-interview/src/replay.rs | 7 +- lib/crates/fabro-server/src/auth/cli_flow.rs | 251 +++++++++++++++--- lib/crates/fabro-server/src/auth/mod.rs | 2 + lib/crates/fabro-server/src/auth/translate.rs | 4 +- lib/crates/fabro-server/src/jwt_auth.rs | 4 +- .../fabro-server/src/principal_middleware.rs | 105 ++++---- lib/crates/fabro-server/src/server.rs | 168 +++++++++--- lib/crates/fabro-server/src/web_auth.rs | 21 +- lib/crates/fabro-slack/src/dispatch.rs | 6 +- lib/crates/fabro-slack/src/interaction.rs | 19 +- lib/crates/fabro-types/src/event_envelope.rs | 10 +- lib/crates/fabro-types/src/lib.rs | 2 +- lib/crates/fabro-types/src/principal.rs | 235 +++++----------- lib/crates/fabro-types/src/run_event/mod.rs | 13 +- lib/crates/fabro-workflow/src/error.rs | 1 + lib/crates/fabro-workflow/src/event.rs | 72 +++-- .../fabro-workflow/src/handler/human.rs | 8 +- .../fabro-workflow/src/lifecycle/event.rs | 32 ++- .../src/pipeline/execute/tests.rs | 17 +- 28 files changed, 661 insertions(+), 401 deletions(-) diff --git a/.github/workflows/typescript.yml b/.github/workflows/typescript.yml index dd537b762..0e50054a3 100644 --- a/.github/workflows/typescript.yml +++ b/.github/workflows/typescript.yml @@ -46,6 +46,7 @@ jobs: - uses: oven-sh/setup-bun@0c5077e51419868618aeaa5fe8019c62421857d6 # v2.2.0 - run: bun install --frozen-lockfile - run: cd apps/fabro-web && bun run typecheck + - run: cd lib/packages/fabro-api-client && bun run typecheck test: name: Test diff --git a/docs/internal/logging-strategy.md b/docs/internal/logging-strategy.md index ee64476c9..b720d491f 100644 --- a/docs/internal/logging-strategy.md +++ b/docs/internal/logging-strategy.md @@ -123,6 +123,8 @@ Fields are key-value pairs that make events queryable. Include enough context th | `idp_issuer`, `idp_subject` | Canonical user identity for authenticated user requests | For HTTP request logs, use the request `Principal` projection rather than hand-assembled auth strings. User identity fields are present only for `Principal::User`; worker and webhook requests use their variant-specific fields (`run_id`, `delivery_id`). + +Server auth intentionally exposes a mutable `RequestAuth` context slot for public auth routes and guard extractors such as `RequiredUser` / `RequireRunScoped` for protected routes. There is no loose `RequestPrincipal` extractor; route-facing extractors should enforce the route's auth contract while the slot supplies the final HTTP log fields. | `input_tokens` | Token count for LLM input | | `output_tokens` | Token count for LLM output | diff --git a/lib/crates/fabro-api/tests/principal_round_trip.rs b/lib/crates/fabro-api/tests/principal_round_trip.rs index ab60d0ace..cba92bf4d 100644 --- a/lib/crates/fabro-api/tests/principal_round_trip.rs +++ b/lib/crates/fabro-api/tests/principal_round_trip.rs @@ -43,7 +43,9 @@ fn principal_round_trips_representative_json() { #[test] fn principal_system_uses_system_kind_field() { - let principal = Principal::system(SystemActorKind::Watchdog); + let principal = Principal::System { + system_kind: SystemActorKind::Watchdog, + }; assert_eq!( serde_json::to_value(principal).unwrap(), @@ -62,16 +64,26 @@ fn principal_round_trips_every_variant_through_api_type() { "octocat".to_string(), AuthMethod::Github, ), - Principal::worker(fixtures::RUN_1), - Principal::webhook("delivery-1".to_string()), - Principal::slack("T1".to_string(), "U1".to_string(), Some("ada".to_string())), - Principal::agent( - Some("ses_agent".to_string()), - Some("ses_parent".to_string()), - Some("gpt-5.4".to_string()), - ), - Principal::system(SystemActorKind::Watchdog), - Principal::anonymous(), + Principal::Worker { + run_id: fixtures::RUN_1, + }, + Principal::Webhook { + delivery_id: "delivery-1".to_string(), + }, + Principal::Slack { + team_id: "T1".to_string(), + user_id: "U1".to_string(), + user_name: Some("ada".to_string()), + }, + Principal::Agent { + session_id: Some("ses_agent".to_string()), + parent_session_id: Some("ses_parent".to_string()), + model: Some("gpt-5.4".to_string()), + }, + Principal::System { + system_kind: SystemActorKind::Watchdog, + }, + Principal::Anonymous, ]; for principal in variants { @@ -93,7 +105,9 @@ fn run_provenance_subject_round_trips_as_principal() { name: Some("fabro-cli".to_string()), version: Some("0.1.0".to_string()), }), - subject: Some(Principal::worker(fixtures::RUN_1)), + subject: Some(Principal::Worker { + run_id: fixtures::RUN_1, + }), }; let json = serde_json::to_value(&provenance).unwrap(); diff --git a/lib/crates/fabro-cli/src/commands/run/attach.rs b/lib/crates/fabro-cli/src/commands/run/attach.rs index e9dda77fa..fe97b1d5e 100644 --- a/lib/crates/fabro-cli/src/commands/run/attach.rs +++ b/lib/crates/fabro-cli/src/commands/run/attach.rs @@ -245,7 +245,7 @@ async fn handle_pending_server_interview( } hide_progress(progress_ui, json_output); - let interviewer = ConsoleInterviewer::new(styles, fabro_types::Principal::anonymous()); + let interviewer = ConsoleInterviewer::new(styles, fabro_types::Principal::Anonymous); let submission = fabro_interview::Interviewer::ask(&interviewer, api_question_to_question(&question)).await; let answer = submission.answer; diff --git a/lib/crates/fabro-cli/src/commands/run/runner.rs b/lib/crates/fabro-cli/src/commands/run/runner.rs index 947951247..f5aebdc82 100644 --- a/lib/crates/fabro-cli/src/commands/run/runner.rs +++ b/lib/crates/fabro-cli/src/commands/run/runner.rs @@ -501,7 +501,9 @@ fn update_worker_title_from_event(event: &RunEvent) { fn stamp_system_worker(mut event: RunEvent) -> RunEvent { if event.actor.is_none() { - event.actor = Some(Principal::worker(event.run_id)); + event.actor = Some(Principal::Worker { + run_id: event.run_id, + }); } event } @@ -758,7 +760,12 @@ mod tests { fn stamp_system_worker_fills_missing_actor_only() { let stamped = stamp_system_worker(running_event(None)); - assert_eq!(stamped.actor, Some(Principal::worker(fixtures::RUN_1))); + assert_eq!( + stamped.actor, + Some(Principal::Worker { + run_id: fixtures::RUN_1, + }) + ); let existing_actor = test_user_principal("octocat"); let stamped = stamp_system_worker(running_event(Some(existing_actor.clone()))); @@ -796,8 +803,18 @@ mod tests { let first = first.lock().await; let second = second.lock().await; - assert_eq!(first[0].actor, Some(Principal::worker(fixtures::RUN_1))); - assert_eq!(second[0].actor, Some(Principal::worker(fixtures::RUN_1))); + assert_eq!( + first[0].actor, + Some(Principal::Worker { + run_id: fixtures::RUN_1, + }) + ); + assert_eq!( + second[0].actor, + Some(Principal::Worker { + run_id: fixtures::RUN_1, + }) + ); } #[tokio::test] diff --git a/lib/crates/fabro-interview/src/auto_approve.rs b/lib/crates/fabro-interview/src/auto_approve.rs index 93290241c..faeba9a7b 100644 --- a/lib/crates/fabro-interview/src/auto_approve.rs +++ b/lib/crates/fabro-interview/src/auto_approve.rs @@ -17,7 +17,9 @@ impl AutoApproveInterviewer { #[must_use] pub fn engine() -> Self { - Self::new(Principal::system(SystemActorKind::Engine)) + Self::new(Principal::System { + system_kind: SystemActorKind::Engine, + }) } } diff --git a/lib/crates/fabro-interview/src/callback.rs b/lib/crates/fabro-interview/src/callback.rs index b3d30af1d..2159a7eeb 100644 --- a/lib/crates/fabro-interview/src/callback.rs +++ b/lib/crates/fabro-interview/src/callback.rs @@ -11,7 +11,12 @@ pub struct CallbackInterviewer { impl CallbackInterviewer { pub fn new(callback: impl Fn(Question) -> Answer + Send + Sync + 'static) -> Self { - Self::with_actor(Principal::system(SystemActorKind::Engine), callback) + Self::with_actor( + Principal::System { + system_kind: SystemActorKind::Engine, + }, + callback, + ) } pub fn with_actor( diff --git a/lib/crates/fabro-interview/src/lib.rs b/lib/crates/fabro-interview/src/lib.rs index 5aae86815..c242cf14a 100644 --- a/lib/crates/fabro-interview/src/lib.rs +++ b/lib/crates/fabro-interview/src/lib.rs @@ -169,7 +169,7 @@ impl AnswerSubmission { pub fn system(answer: Answer, system_kind: SystemActorKind) -> Self { Self { answer, - actor: Principal::system(system_kind), + actor: Principal::System { system_kind }, } } } diff --git a/lib/crates/fabro-interview/src/queue.rs b/lib/crates/fabro-interview/src/queue.rs index f15edc7d0..7289fc326 100644 --- a/lib/crates/fabro-interview/src/queue.rs +++ b/lib/crates/fabro-interview/src/queue.rs @@ -15,7 +15,9 @@ pub struct QueueInterviewer { impl QueueInterviewer { #[must_use] pub fn new(answers: VecDeque) -> Self { - Self::with_actor(answers, Principal::system(SystemActorKind::Engine)) + Self::with_actor(answers, Principal::System { + system_kind: SystemActorKind::Engine, + }) } #[must_use] diff --git a/lib/crates/fabro-interview/src/replay.rs b/lib/crates/fabro-interview/src/replay.rs index 90d428142..94a30eca0 100644 --- a/lib/crates/fabro-interview/src/replay.rs +++ b/lib/crates/fabro-interview/src/replay.rs @@ -31,10 +31,9 @@ impl Interviewer for ReplayInterviewer { async fn ask(&self, _question: Question) -> AnswerSubmission { let mut submissions = self.submissions.lock().expect("answers lock poisoned"); if submissions.is_empty() { - AnswerSubmission::new( - Answer::interrupted(), - Principal::system(SystemActorKind::Engine), - ) + AnswerSubmission::new(Answer::interrupted(), Principal::System { + system_kind: SystemActorKind::Engine, + }) } else { submissions.remove(0) } diff --git a/lib/crates/fabro-server/src/auth/cli_flow.rs b/lib/crates/fabro-server/src/auth/cli_flow.rs index c6f26493c..fd8c98368 100644 --- a/lib/crates/fabro-server/src/auth/cli_flow.rs +++ b/lib/crates/fabro-server/src/auth/cli_flow.rs @@ -27,17 +27,18 @@ use tracing::{info, warn}; use url::{Host, Url}; use crate::auth::browser_shell::browser_shell; -use crate::auth::{self, AuthCode, ConsumeOutcome, JwtSubject, RefreshToken}; +use crate::auth::{self, AuthCode, ConsumeOutcome, JwtSubject, REFRESH_TOKEN_PREFIX, RefreshToken}; use crate::jwt_auth::{AuthMode, ConfiguredAuth}; -use crate::principal_middleware::{AuthStatus, RequestAuth, RequestAuthContext}; +use crate::principal_middleware::{AuthContextSlot, AuthStatus, RequestAuth, RequestAuthContext}; use crate::server::AppState; -use crate::web_auth::{SessionCookie, read_private_session}; +use crate::web_auth::{ + SessionCookie, auth_context_from_session, read_private_session, session_cookie_present, +}; const CLI_FLOW_COOKIE_NAME: &str = "fabro_cli_flow"; const QUERY_VALUE_ENCODE_SET: &AsciiSet = &NON_ALPHANUMERIC.remove(b'_').remove(b'-'); const ACCESS_TOKEN_TTL_MINUTES: i64 = 10; const REFRESH_TOKEN_TTL_DAYS: i64 = 30; -const REFRESH_TOKEN_PREFIX: &str = "fabro_refresh_"; const GITHUB_NOT_CONFIGURED: &str = "GitHub login is not configured for this server."; const DEV_TOKEN_LOGIN_INSTRUCTIONS: &str = concat!( "This server uses dev-token auth.\n\n", @@ -120,8 +121,14 @@ async fn start( State(state): State>, Extension(auth_mode): Extension, Query(params): Query, + RequestAuth(auth_slot): RequestAuth, headers: HeaderMap, ) -> Response { + let session_key = state.session_key(); + let session = session_key + .as_ref() + .and_then(|session_key| stamp_cli_session_auth_context(&auth_slot, &headers, session_key)); + let Some(redirect_uri) = params .redirect_uri .as_deref() @@ -146,7 +153,7 @@ async fn start( ); } - let Some(session_key) = state.session_key() else { + let Some(session_key) = session_key else { return redirect_with_error( &redirect_uri, state_token, @@ -171,8 +178,6 @@ async fn start( "Invalid PKCE parameters", ); } - let session = read_private_session(&headers, &session_key); - let secure = session_cookie_secure(state.as_ref()); let mut jar = CookieJar::new(); add_cli_flow_cookie( @@ -199,13 +204,19 @@ async fn resume( State(state): State>, Extension(auth_mode): Extension, Query(params): Query, + RequestAuth(auth_slot): RequestAuth, headers: HeaderMap, ) -> Response { + let session_key = state.session_key(); + let session = session_key + .as_ref() + .and_then(|session_key| stamp_cli_session_auth_context(&auth_slot, &headers, session_key)); + if !github_enabled(&auth_mode) { return static_error_page(GITHUB_NOT_CONFIGURED); } - let Some(session_key) = state.session_key() else { + let Some(session_key) = session_key else { return static_error_page(GITHUB_NOT_CONFIGURED); }; let Some(flow) = read_private_cli_flow(&headers, &session_key) else { @@ -238,7 +249,6 @@ async fn resume( return response; } - let session = read_private_session(&headers, &session_key); let Some(session) = eligible_session(session.as_ref()) else { let mut jar = CookieJar::new(); remove_cli_flow_cookie(&mut jar, &session_key, secure); @@ -258,8 +268,14 @@ async fn resume( async fn confirm_resume( State(state): State>, Extension(auth_mode): Extension, + RequestAuth(auth_slot): RequestAuth, headers: HeaderMap, ) -> Response { + let session_key = state.session_key(); + let session = session_key + .as_ref() + .and_then(|session_key| stamp_cli_session_auth_context(&auth_slot, &headers, session_key)); + if !github_enabled(&auth_mode) { return static_error_page(GITHUB_NOT_CONFIGURED); } @@ -268,7 +284,7 @@ async fn confirm_resume( return static_error_page(INVALID_CONFIRMATION_REQUEST); } - let Some(session_key) = state.session_key() else { + let Some(session_key) = session_key else { return static_error_page(GITHUB_NOT_CONFIGURED); }; let Some(flow) = read_private_cli_flow(&headers, &session_key) else { @@ -287,7 +303,6 @@ async fn confirm_resume( }; let secure = session_cookie_secure(state.as_ref()); - let session = read_private_session(&headers, &session_key); let Some(session) = eligible_session(session.as_ref()) else { let mut jar = CookieJar::new(); remove_cli_flow_cookie(&mut jar, &session_key, secure); @@ -318,6 +333,7 @@ async fn confirm_resume( async fn token( State(state): State>, Extension(auth_mode): Extension, + RequestAuth(auth_slot): RequestAuth, headers: HeaderMap, body: Result, JsonRejection>, ) -> Response { @@ -325,21 +341,24 @@ async fn token( return github_auth_not_configured(); }; let Ok(Json(body)) = body else { - return oauth_error( + return oauth_invalid( + &auth_slot, StatusCode::BAD_REQUEST, "invalid_request", "Invalid request", ); }; let Some(code) = body.code.as_deref() else { - return oauth_error( + return oauth_invalid( + &auth_slot, StatusCode::BAD_REQUEST, "invalid_request", "Invalid request", ); }; let Some(code_verifier) = body.code_verifier.as_deref() else { - return oauth_error( + return oauth_invalid( + &auth_slot, StatusCode::BAD_REQUEST, "invalid_request", "Invalid request", @@ -350,14 +369,16 @@ async fn token( .as_deref() .and_then(canonical_loopback_redirect_uri) else { - return oauth_error( + return oauth_invalid( + &auth_slot, StatusCode::BAD_REQUEST, "invalid_request", "Invalid request", ); }; if body.grant_type.as_deref() != Some("authorization_code") { - return oauth_error( + return oauth_invalid( + &auth_slot, StatusCode::BAD_REQUEST, "invalid_request", "Invalid request", @@ -386,7 +407,8 @@ async fn token( ); } }) else { - return oauth_error( + return oauth_invalid( + &auth_slot, StatusCode::BAD_REQUEST, "invalid_code", "Invalid authorization code", @@ -394,7 +416,8 @@ async fn token( }; if pkce_challenge(code_verifier) != entry.code_challenge { - return oauth_error( + return oauth_invalid( + &auth_slot, StatusCode::BAD_REQUEST, "pkce_verification_failed", "PKCE verification failed", @@ -403,14 +426,20 @@ async fn token( if canonical_loopback_redirect_uri(&entry.redirect_uri).as_deref() != Some(redirect_uri.as_str()) { - return oauth_error( + return oauth_invalid( + &auth_slot, StatusCode::BAD_REQUEST, "redirect_uri_mismatch", "Redirect URI mismatch", ); } if !login_allowed(state.as_ref(), &entry.login) { - return oauth_error(StatusCode::FORBIDDEN, "unauthorized", "Login not permitted"); + return oauth_invalid( + &auth_slot, + StatusCode::FORBIDDEN, + "unauthorized", + "Login not permitted", + ); } let Some(jwt_key) = config.jwt_key.as_ref() else { @@ -482,6 +511,7 @@ async fn token( ); log_cli_auth_tokens_issued(&entry.login, &entry.email); + auth_slot.replace(refresh_user_context(&refresh_row)); Json(CliTokenResponse { access_token, @@ -513,7 +543,7 @@ async fn refresh( ); } RefreshCredential::Invalid => { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); return oauth_error(StatusCode::UNAUTHORIZED, "unauthorized", "Unauthorized"); } RefreshCredential::Present(secret) => secret, @@ -583,7 +613,7 @@ async fn refresh( let (old, new_row) = match outcome { ConsumeOutcome::NotFound | ConsumeOutcome::Expired => { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); if auth_tokens.was_recently_replay_revoked(&secret_hash, now) { return oauth_error( StatusCode::UNAUTHORIZED, @@ -598,7 +628,7 @@ async fn refresh( ); } ConsumeOutcome::Reused(old) => { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); auth_tokens.mark_refresh_token_replay(secret_hash, now); if let Err(err) = auth_tokens.delete_chain(old.chain_id).await { warn!(error = %err, chain_id = %old.chain_id, "Failed to revoke replayed refresh token chain"); @@ -614,7 +644,7 @@ async fn refresh( }; if !login_allowed(state.as_ref(), &old.login) { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); if let Err(err) = auth_tokens.delete_chain(old.chain_id).await { warn!(error = %err, chain_id = %old.chain_id, "Failed to revoke deauthorized refresh token chain"); } @@ -661,7 +691,7 @@ async fn logout( let secret = match refresh_credential_from_headers(&headers) { RefreshCredential::Missing => return StatusCode::NO_CONTENT.into_response(), RefreshCredential::Invalid => { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); return StatusCode::NO_CONTENT.into_response(); } RefreshCredential::Present(secret) => secret, @@ -705,7 +735,7 @@ async fn logout( } log_cli_refresh_chain_logged_out(&refresh_token.login, &refresh_token.email); } else { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); } StatusCode::NO_CONTENT.into_response() @@ -748,6 +778,22 @@ fn eligible_session(session: Option<&SessionCookie>) -> Option<&SessionCookie> { .filter(|session| session.identity.is_some() && session.auth_method == AuthMethod::Github) } +fn stamp_cli_session_auth_context( + auth_slot: &AuthContextSlot, + headers: &HeaderMap, + session_key: &Key, +) -> Option { + let session = read_private_session(headers, session_key); + if let Some(session) = eligible_session(session.as_ref()) { + if let Some(context) = auth_context_from_session(session) { + auth_slot.replace(context); + } + } else if session_cookie_present(headers) { + auth_slot.replace(RequestAuthContext::invalid()); + } + session +} + fn valid_state_token(state: &str) -> bool { (16..=512).contains(&state.len()) && state @@ -926,10 +972,6 @@ fn refresh_user_context(refresh_token: &RefreshToken) -> RequestAuthContext { ) } -fn invalid_auth_context() -> RequestAuthContext { - RequestAuthContext::rejected(AuthStatus::Invalid, Some("unauthorized")) -} - fn hash_refresh_secret(secret: &str) -> [u8; 32] { Sha256::digest(secret.as_bytes()).into() } @@ -995,6 +1037,19 @@ fn oauth_error( .into_response() } +fn oauth_invalid( + auth_slot: &AuthContextSlot, + status: StatusCode, + error: &'static str, + error_description: &'static str, +) -> Response { + auth_slot.replace(RequestAuthContext::rejected( + AuthStatus::Invalid, + Some(error), + )); + oauth_error(status, error, error_description) +} + fn random_secret() -> String { let mut bytes = [0_u8; 32]; OsRng.try_fill_bytes(&mut bytes).expect("OS RNG"); @@ -1336,6 +1391,25 @@ client_id = "github-client-id" .to_string() } + fn cli_flow_cookie(key: &Key) -> String { + let mut jar = cookie::CookieJar::new(); + add_cli_flow_cookie( + &mut jar, + key, + &CliFlowCookie { + redirect_uri: "http://127.0.0.1:4444/callback".to_string(), + state: "abcdefghijklmnop".to_string(), + code_challenge: "challenge".to_string(), + }, + true, + ); + jar.delta() + .next() + .expect("flow cookie should exist") + .encoded() + .to_string() + } + fn pkce_challenge(verifier: &str) -> String { URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes())) } @@ -1562,6 +1636,65 @@ client_id = "github-client-id" }); } + #[tokio::test] + async fn cli_session_routes_stamp_public_auth_context() { + let key = test_cookie_key(); + let (app, _state, captured) = test_router_with_auth_capture( + github_settings("https://fabro.example"), + github_auth_mode(), + ); + let session_cookie = github_session_cookie(&key); + let flow_cookie = cli_flow_cookie(&key); + let cookie_header = format!("{session_cookie}; {flow_cookie}"); + + let start = app + .clone() + .oneshot( + Request::builder() + .uri("/auth/cli/start?redirect_uri=http://127.0.0.1:4444/callback&state=abcdefghijklmnop&code_challenge=challenge&code_challenge_method=S256") + .header(header::COOKIE, session_cookie) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(start.status(), StatusCode::SEE_OTHER); + + let resume = app + .clone() + .oneshot( + Request::builder() + .uri("/auth/cli/resume") + .header(header::COOKIE, cookie_header.clone()) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resume.status(), StatusCode::OK); + + let confirm = app + .oneshot( + Request::builder() + .method("POST") + .uri("/auth/cli/resume") + .header(header::COOKIE, cookie_header) + .header(header::ORIGIN, "https://fabro.example") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(confirm.status(), StatusCode::SEE_OTHER); + + let contexts = captured.lock().expect("captured auth contexts").clone(); + assert_eq!(contexts.len(), 3); + for context in contexts { + assert_eq!(context.auth_status, AuthStatus::Authenticated); + assert_eq!(context.principal.display(), "octocat"); + } + } + #[tokio::test] async fn resume_with_github_session_renders_confirmation_page() { let key = test_cookie_key(); @@ -1903,6 +2036,64 @@ client_id = "github-client-id" assert_eq!(refresh.user_agent, "fabro-cli/0.1"); } + #[tokio::test] + async fn token_stamps_public_auth_context() { + let (app, state, captured) = test_router_with_auth_capture( + github_settings("https://fabro.example"), + github_auth_mode(), + ); + insert_auth_code(state.as_ref(), "auth-code-auth-context", "test-verifier").await; + + let response = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/auth/cli/token") + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from( + json!({ + "grant_type": "authorization_code", + "code": "auth-code-auth-context", + "code_verifier": "test-verifier", + "redirect_uri": "http://127.0.0.1:4444/callback" + }) + .to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/auth/cli/token") + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from( + json!({ + "grant_type": "authorization_code", + "code": "missing-code", + "code_verifier": "test-verifier", + "redirect_uri": "http://127.0.0.1:4444/callback" + }) + .to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + + let contexts = captured.lock().expect("captured auth contexts").clone(); + assert_eq!(contexts[0].auth_status, AuthStatus::Authenticated); + assert_eq!(contexts[0].principal.display(), "octocat"); + assert_eq!(contexts[1].auth_status, AuthStatus::Invalid); + assert_eq!(contexts[1].auth_error_code, Some("invalid_code")); + } + #[tokio::test] async fn token_wrong_verifier_fails_and_burns_code() { let (app, state) = test_router(github_settings("https://fabro.example")); diff --git a/lib/crates/fabro-server/src/auth/mod.rs b/lib/crates/fabro-server/src/auth/mod.rs index 31d3ef083..449d29a8b 100644 --- a/lib/crates/fabro-server/src/auth/mod.rs +++ b/lib/crates/fabro-server/src/auth/mod.rs @@ -5,6 +5,8 @@ mod jwt; mod keys; mod translate; +pub(crate) const REFRESH_TOKEN_PREFIX: &str = "fabro_refresh_"; + pub(crate) use browser_shell::browser_shell; pub(crate) use cli_flow::web_routes; pub(crate) use fabro_store::{AuthCode, ConsumeOutcome, RefreshToken}; diff --git a/lib/crates/fabro-server/src/auth/translate.rs b/lib/crates/fabro-server/src/auth/translate.rs index 28b3594f7..e17d8fc0f 100644 --- a/lib/crates/fabro-server/src/auth/translate.rs +++ b/lib/crates/fabro-server/src/auth/translate.rs @@ -9,7 +9,7 @@ use fabro_types::{AuthMethod, IdpIdentity}; use fabro_util::dev_token::validate_dev_token_format; use tracing::trace; -use crate::auth::{self, JwtSubject}; +use crate::auth::{self, JwtSubject, REFRESH_TOKEN_PREFIX}; use crate::jwt_auth::{AuthMode, ConfiguredAuth, dev_token_matches}; use crate::server::AppState; use crate::web_auth::{self, SessionCookie}; @@ -58,7 +58,7 @@ pub(crate) async fn auth_translation_middleware( } fn translate_bearer_token(token: &str, config: &ConfiguredAuth) -> Option { - if token.starts_with("fabro_refresh_") || !token.starts_with("fabro_dev_") { + if token.starts_with(REFRESH_TOKEN_PREFIX) || !token.starts_with("fabro_dev_") { return None; } diff --git a/lib/crates/fabro-server/src/jwt_auth.rs b/lib/crates/fabro-server/src/jwt_auth.rs index f25b496ba..c94baff62 100644 --- a/lib/crates/fabro-server/src/jwt_auth.rs +++ b/lib/crates/fabro-server/src/jwt_auth.rs @@ -11,6 +11,8 @@ use sha2::Sha256; #[cfg(test)] use tracing::info; +#[cfg(test)] +use crate::auth::REFRESH_TOKEN_PREFIX; use crate::auth::{self, JwtError, JwtSigningKey, KeyDeriveError}; use crate::error::ApiError; @@ -265,7 +267,7 @@ fn authenticate_bearer( token: &str, config: &ConfiguredAuth, ) -> Result { - if token.starts_with("fabro_refresh_") { + if token.starts_with(REFRESH_TOKEN_PREFIX) { info!( path = %parts.uri.path(), "Refresh token presented at protected endpoint" diff --git a/lib/crates/fabro-server/src/principal_middleware.rs b/lib/crates/fabro-server/src/principal_middleware.rs index 55222ca1a..0bd3efd42 100644 --- a/lib/crates/fabro-server/src/principal_middleware.rs +++ b/lib/crates/fabro-server/src/principal_middleware.rs @@ -9,8 +9,9 @@ use axum::response::{IntoResponse, Response}; use fabro_types::{CommandOutputStream, Principal, RunBlobId, RunId, StageId, UserPrincipal}; use jsonwebtoken::dangerous::insecure_decode; use serde::Deserialize; +use strum::IntoStaticStr; -use crate::auth::JwtError; +use crate::auth::{JwtError, REFRESH_TOKEN_PREFIX}; use crate::error::ApiError; use crate::jwt_auth::{self, AuthMode, ConfiguredAuth}; use crate::server::{AppState, parse_blob_id_path, parse_run_id_path, parse_stage_id_path}; @@ -32,7 +33,8 @@ pub(crate) struct UserProfile { pub user_url: String, } -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, IntoStaticStr)] +#[strum(serialize_all = "snake_case")] pub(crate) enum AuthStatus { Missing, Invalid, @@ -43,6 +45,9 @@ pub(crate) enum AuthStatus { #[derive(Clone)] pub(crate) struct AuthContextSlot(pub(crate) Arc>); +// Route handlers intentionally use either this slot handle or a guard extractor +// such as RequiredUser/RequireRunScoped. A loose RequestPrincipal extractor +// would make it easy to read a principal without enforcing the route's guard. pub(crate) struct RequestAuth(pub(crate) AuthContextSlot); pub(crate) struct RequiredUser(pub(crate) UserPrincipal); @@ -70,7 +75,7 @@ impl RequestAuthContext { #[must_use] pub(crate) fn initial() -> Self { Self { - principal: Principal::anonymous(), + principal: Principal::Anonymous, auth_status: AuthStatus::Missing, auth_error_code: None, user_profile: None, @@ -90,23 +95,23 @@ impl RequestAuthContext { #[must_use] pub(crate) fn rejected(status: AuthStatus, code: Option<&'static str>) -> Self { Self { - principal: Principal::anonymous(), + principal: Principal::Anonymous, auth_status: status, auth_error_code: code, user_profile: None, } } + + #[must_use] + pub(crate) fn invalid() -> Self { + Self::rejected(AuthStatus::Invalid, Some("unauthorized")) + } } impl AuthStatus { #[must_use] pub(crate) fn as_str(self) -> &'static str { - match self { - Self::Missing => "missing", - Self::Invalid => "invalid", - Self::Expired => "expired", - Self::Authenticated => "authenticated", - } + self.into() } } @@ -163,7 +168,7 @@ impl FromRequestParts> for RequireRunScoped { .await .map_err(IntoResponse::into_response)?; let run_id = parse_run_id_path(&id)?; - require_worker_or_user_for_run(&principal_from_parts(parts), &run_id) + require_worker_or_user_for_run(&auth_context_from_parts(parts), &run_id) .map_err(IntoResponse::into_response)?; Ok(Self(run_id)) } @@ -181,7 +186,7 @@ impl FromRequestParts> for RequireRunBlob { .map_err(IntoResponse::into_response)?; let run_id = parse_run_id_path(&id)?; let blob_id = parse_blob_id_path(&blob_id)?; - require_worker_or_user_for_run(&principal_from_parts(parts), &run_id) + require_worker_or_user_for_run(&auth_context_from_parts(parts), &run_id) .map_err(IntoResponse::into_response)?; Ok(Self(run_id, blob_id)) } @@ -199,7 +204,7 @@ impl FromRequestParts> for RequireStageArtifact { .map_err(IntoResponse::into_response)?; let run_id = parse_run_id_path(&id)?; let stage_id = parse_stage_id_path(&stage_id)?; - require_worker_or_user_for_run(&principal_from_parts(parts), &run_id) + require_worker_or_user_for_run(&auth_context_from_parts(parts), &run_id) .map_err(IntoResponse::into_response)?; Ok(Self(run_id, stage_id)) } @@ -221,7 +226,7 @@ impl FromRequestParts> for RequireCommandLog { let stream = stream .parse::() .map_err(|_| ApiError::bad_request("Invalid command log stream.").into_response())?; - require_worker_or_user_for_run(&principal_from_parts(parts), &run_id) + require_worker_or_user_for_run(&auth_context_from_parts(parts), &run_id) .map_err(IntoResponse::into_response)?; Ok(Self(run_id, stage_id, stream)) } @@ -247,11 +252,11 @@ pub(crate) async fn principal_middleware( next.run(req).await } -fn principal_from_parts(parts: &Parts) -> Principal { +fn auth_context_from_parts(parts: &Parts) -> RequestAuthContext { parts .extensions .get::() - .map_or_else(Principal::anonymous, |slot| slot.snapshot().principal) + .map_or_else(RequestAuthContext::initial, AuthContextSlot::snapshot) } pub(crate) fn require_user(slot: &AuthContextSlot) -> Result { @@ -280,45 +285,15 @@ pub(crate) fn require_authenticated_user( } } -#[allow( - dead_code, - reason = "Worker-only route migration is staged behind the shared context." -)] -pub(crate) fn require_worker_for_run( - principal: &Principal, +fn require_worker_or_user_for_run( + context: &RequestAuthContext, route_run_id: &RunId, ) -> Result<(), ApiError> { - match principal { - Principal::Worker { run_id } if run_id == route_run_id => Ok(()), - Principal::Worker { .. } => Err(ApiError::forbidden()), - _ => Err(ApiError::unauthorized()), - } -} - -#[allow( - dead_code, - reason = "Run-scoped route migration is staged behind the shared context." -)] -pub(crate) fn require_worker_or_user_for_run( - principal: &Principal, - route_run_id: &RunId, -) -> Result<(), ApiError> { - match principal { + match &context.principal { Principal::User(_) => Ok(()), Principal::Worker { run_id } if run_id == route_run_id => Ok(()), Principal::Worker { .. } => Err(ApiError::forbidden()), - _ => Err(ApiError::unauthorized()), - } -} - -#[allow( - dead_code, - reason = "Webhook route stamps the slot inline; guard is for future webhook consumers." -)] -pub(crate) fn require_webhook(principal: &Principal) -> Result { - match principal { - Principal::Webhook { delivery_id } => Ok(delivery_id.clone()), - _ => Err(ApiError::unauthorized()), + _ => Err(auth_rejection(context.auth_status, context.auth_error_code)), } } @@ -336,7 +311,7 @@ fn classify_request(req: &Request, state: &AppState) -> RequestAuthContext { Some(Ok(token)) => token, }; - if token.starts_with("fabro_refresh_") { + if token.starts_with(REFRESH_TOKEN_PREFIX) { return rejected(AuthStatus::Invalid, Some("unauthorized")); } if !jwt_auth::looks_like_jwt(token) { @@ -350,7 +325,7 @@ fn classify_request(req: &Request, state: &AppState) -> RequestAuthContext { if issuer == WORKER_TOKEN_ISSUER { return match worker_token::decode_worker_token(token, state.worker_token_keys()) { - Ok(run_id) => authenticated(Principal::worker(run_id), None), + Ok(run_id) => authenticated(Principal::Worker { run_id }, None), Err(JwtError::AccessTokenExpired) => { rejected(AuthStatus::Expired, Some("access_token_expired")) } @@ -552,7 +527,7 @@ mod tests { let context = classify_request(&request, state.as_ref()); assert_eq!(context.auth_status, AuthStatus::Authenticated); - assert_eq!(context.principal, Principal::worker(run_id)); + assert_eq!(context.principal, Principal::Worker { run_id }); } #[test] @@ -618,6 +593,28 @@ mod tests { assert_eq!(context.auth_status, AuthStatus::Missing); assert_eq!(context.auth_error_code, None); - assert_eq!(context.principal, Principal::anonymous()); + assert_eq!(context.principal, Principal::Anonymous); + } + + #[test] + fn run_scoped_guard_preserves_expired_auth_error_code() { + let context = + RequestAuthContext::rejected(AuthStatus::Expired, Some("access_token_expired")); + + let err = require_worker_or_user_for_run(&context, &RunId::new()).unwrap_err(); + + assert_eq!(err.status(), StatusCode::UNAUTHORIZED); + assert_eq!(err.code(), Some("access_token_expired")); + } + + #[test] + fn run_scoped_guard_preserves_invalid_auth_error_code() { + let context = + RequestAuthContext::rejected(AuthStatus::Invalid, Some("access_token_invalid")); + + let err = require_worker_or_user_for_run(&context, &RunId::new()).unwrap_err(); + + assert_eq!(err.status(), StatusCode::UNAUTHORIZED); + assert_eq!(err.code(), Some("access_token_invalid")); } } diff --git a/lib/crates/fabro-server/src/server.rs b/lib/crates/fabro-server/src/server.rs index 3bb8d4f1c..e35d951b6 100644 --- a/lib/crates/fabro-server/src/server.rs +++ b/lib/crates/fabro-server/src/server.rs @@ -1098,7 +1098,7 @@ async fn http_log_middleware(mut req: axum_extract::Request, next: Next) -> Resp let status = response.status().as_u16(); let latency_ms = start.elapsed().as_millis(); let auth_context = auth_slot.snapshot(); - let principal_fields = auth_context.principal.log_fields(); + let principal_kind = auth_context.principal.kind(); let auth_status = auth_context.auth_status.as_str(); macro_rules! emit_http_log { @@ -1110,7 +1110,7 @@ async fn http_log_middleware(mut req: axum_extract::Request, next: Next) -> Resp status, latency_ms, request_id = %request_id, - principal_kind = principal_fields.principal_kind, + principal_kind, auth_status, auth_error_code, $($field = $value,)* @@ -1123,7 +1123,7 @@ async fn http_log_middleware(mut req: axum_extract::Request, next: Next) -> Resp status, latency_ms, request_id = %request_id, - principal_kind = principal_fields.principal_kind, + principal_kind, auth_status, $($field = $value,)* "HTTP response" @@ -1135,45 +1135,25 @@ async fn http_log_middleware(mut req: axum_extract::Request, next: Next) -> Resp macro_rules! emit_principal_http_log { ($level:ident) => {{ match &auth_context.principal { - Principal::User(_) => emit_http_log!( + Principal::User(user) => emit_http_log!( $level, - user_auth_method = principal_fields - .user_auth_method - .expect("user auth method field"), - idp_issuer = principal_fields - .idp_issuer - .as_deref() - .expect("user idp issuer field"), - idp_subject = principal_fields - .idp_subject - .as_deref() - .expect("user idp subject field"), - login = principal_fields.login.as_deref().expect("user login field"), + user_auth_method = user.auth_method.as_str(), + idp_issuer = user.identity.issuer(), + idp_subject = user.identity.subject(), + login = user.login.as_str(), ), - Principal::Worker { .. } => emit_http_log!( + Principal::Worker { run_id } => { + emit_http_log!($level, run_id = run_id.to_string().as_str(),) + } + Principal::Webhook { delivery_id } => { + emit_http_log!($level, delivery_id = delivery_id.as_str(),) + } + Principal::Slack { + team_id, user_id, .. + } => emit_http_log!( $level, - run_id = principal_fields - .run_id - .as_deref() - .expect("worker run id field"), - ), - Principal::Webhook { .. } => emit_http_log!( - $level, - delivery_id = principal_fields - .delivery_id - .as_deref() - .expect("webhook delivery id field"), - ), - Principal::Slack { .. } => emit_http_log!( - $level, - team_id = principal_fields - .team_id - .as_deref() - .expect("slack team id field"), - user_id = principal_fields - .user_id - .as_deref() - .expect("slack user id field"), + team_id = team_id.as_str(), + user_id = user_id.as_str(), ), Principal::Agent { .. } | Principal::System { .. } | Principal::Anonymous => { emit_http_log!($level) @@ -1415,7 +1395,7 @@ async fn github_webhook( .and_then(|value| value.to_str().ok()) else { auth_slot.replace(RequestAuthContext { - principal: Principal::anonymous(), + principal: Principal::Anonymous, auth_status: AuthStatus::Invalid, auth_error_code: Some("unauthorized"), user_profile: None, @@ -1426,7 +1406,7 @@ async fn github_webhook( if !verify_signature(&secret, &body, signature) { auth_slot.replace(RequestAuthContext { - principal: Principal::anonymous(), + principal: Principal::Anonymous, auth_status: AuthStatus::Invalid, auth_error_code: Some("unauthorized"), user_profile: None, @@ -1436,7 +1416,9 @@ async fn github_webhook( } auth_slot.replace(RequestAuthContext { - principal: Principal::webhook(delivery_id.to_string()), + principal: Principal::Webhook { + delivery_id: delivery_id.to_string(), + }, auth_status: AuthStatus::Authenticated, auth_error_code: None, user_profile: None, @@ -8585,6 +8567,32 @@ methods = ["dev-token"] (guard, events) } + fn captured_field<'a>(event: &'a CapturedTracingEvent, name: &str) -> Option<&'a str> { + event + .fields + .iter() + .find_map(|(field_name, value)| (field_name == name).then_some(value.as_str())) + } + + fn assert_log_field(event: &CapturedTracingEvent, name: &str, expected: &str) { + let actual = captured_field(event, name) + .unwrap_or_else(|| panic!("expected log field {name}; fields were {:?}", event.fields)); + let debug_expected = format!("{expected:?}"); + assert!( + actual == expected || actual == debug_expected, + "expected field {name} to be {expected:?}, got {actual:?}; fields were {:?}", + event.fields + ); + } + + fn assert_log_field_absent(event: &CapturedTracingEvent, name: &str) { + assert!( + captured_field(event, name).is_none(), + "expected log field {name} to be absent; fields were {:?}", + event.fields + ); + } + macro_rules! response_json { ($response:expr, $expected:expr) => { fabro_test::expect_axum_json($response, $expected, concat!(file!(), ":", line!())) @@ -8633,6 +8641,80 @@ methods = ["dev-token"] assert!(!field_names.contains(&"run_id")); } + #[tokio::test(flavor = "current_thread")] + async fn http_log_records_user_principal_fields() { + let (_state, app) = jwt_auth_app(); + let bearer = issue_test_user_jwt(); + let (_guard, events) = capture_server_logs(); + + let response = app + .oneshot(bearer_request(Method::GET, "/runs", &bearer, Body::empty())) + .await + .unwrap(); + assert_status!(response, StatusCode::OK).await; + + let events = events.lock().expect("captured log events").clone(); + assert_eq!(events.len(), 1); + let event = &events[0]; + assert_log_field(event, "principal_kind", "user"); + assert_log_field(event, "auth_status", "authenticated"); + assert_log_field(event, "user_auth_method", "github"); + assert_log_field(event, "idp_issuer", "https://github.com"); + assert_log_field(event, "idp_subject", "12345"); + assert_log_field(event, "login", "octocat"); + assert_log_field_absent(event, "auth_error_code"); + } + + #[tokio::test(flavor = "current_thread")] + async fn http_log_records_worker_principal_fields() { + let (_state, app) = jwt_auth_app(); + let user_bearer = issue_test_user_jwt(); + let run_id = create_run_with_bearer(&app, &user_bearer).await; + let worker_bearer = issue_test_worker_token(&run_id); + let (_guard, events) = capture_server_logs(); + + let response = app + .oneshot(bearer_request( + Method::GET, + &format!("/runs/{run_id}/state"), + &worker_bearer, + Body::empty(), + )) + .await + .unwrap(); + assert_status!(response, StatusCode::OK).await; + + let events = events.lock().expect("captured log events").clone(); + assert_eq!(events.len(), 1); + let event = &events[0]; + assert_log_field(event, "principal_kind", "worker"); + assert_log_field(event, "auth_status", "authenticated"); + assert_log_field(event, "run_id", &run_id.to_string()); + assert_log_field_absent(event, "auth_error_code"); + } + + #[tokio::test(flavor = "current_thread")] + async fn http_log_records_webhook_principal_fields() { + let body = br#"{"repository":{"full_name":"owner/repo"},"action":"opened"}"#; + let signature = compute_signature(TEST_WEBHOOK_SECRET.as_bytes(), body); + let app = webhook_test_app(dev_token_auth_mode()); + let (_guard, events) = capture_server_logs(); + + let response = app + .oneshot(webhook_request(Some(&signature), None, body)) + .await + .unwrap(); + assert_status!(response, StatusCode::OK).await; + + let events = events.lock().expect("captured log events").clone(); + assert_eq!(events.len(), 1); + let event = &events[0]; + assert_log_field(event, "principal_kind", "webhook"); + assert_log_field(event, "auth_status", "authenticated"); + assert_log_field(event, "delivery_id", "delivery-1"); + assert_log_field_absent(event, "auth_error_code"); + } + #[allow( clippy::needless_pass_by_value, reason = "Test helper mirrors the public build_router convenience API." @@ -8665,6 +8747,7 @@ methods = ["dev-token"] let mut builder = Request::builder() .method("POST") .uri(api("/webhooks/github")) + .header("x-github-delivery", "delivery-1") .header("x-github-event", "pull_request"); if let Some(sig) = signature { builder = builder.header("x-hub-signature-256", sig); @@ -10119,6 +10202,7 @@ allowed_usernames = ["octocat"] failure: FailureDetail::new("try again", FailureCategory::TransientInfra), will_retry: true, duration_ms: 10, + actor: None, }, workflow_event::Event::StageRetrying { node_id: "work".to_string(), diff --git a/lib/crates/fabro-server/src/web_auth.rs b/lib/crates/fabro-server/src/web_auth.rs index 4ca94394d..8fcccfbed 100644 --- a/lib/crates/fabro-server/src/web_auth.rs +++ b/lib/crates/fabro-server/src/web_auth.rs @@ -21,8 +21,7 @@ use tracing::{debug, error, info, warn}; use crate::auth::{GithubEndpoints, browser_shell}; use crate::jwt_auth::{AuthMode, auth_method_name, dev_token_matches}; use crate::principal_middleware::{ - AuthStatus, RequestAuth, RequestAuthContext, RequiredUser, UserProfile, - require_authenticated_user, + RequestAuth, RequestAuthContext, RequiredUser, UserProfile, require_authenticated_user, }; use crate::server::AppState; @@ -162,13 +161,13 @@ pub fn read_private_session(headers: &HeaderMap, key: &Key) -> Option bool { +pub(crate) fn session_cookie_present(headers: &HeaderMap) -> bool { parse_cookie_header(headers) .get(SESSION_COOKIE_NAME) .is_some() } -fn auth_context_from_session(session: &SessionCookie) -> Option { +pub(crate) fn auth_context_from_session(session: &SessionCookie) -> Option { let identity = session.identity.clone()?; let principal = Principal::user(identity, session.login.clone(), session.auth_method); Some(RequestAuthContext::authenticated( @@ -182,10 +181,6 @@ fn auth_context_from_session(session: &SessionCookie) -> Option RequestAuthContext { - RequestAuthContext::rejected(AuthStatus::Invalid, Some("unauthorized")) -} - fn read_private_oauth_state(headers: &HeaderMap, key: &Key) -> Option { let jar = parse_cookie_header(headers); jar.private(key) @@ -337,17 +332,17 @@ async fn login_dev_token( payload: Result, JsonRejection>, ) -> Response { let Ok(Json(payload)) = payload else { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); return json_response(StatusCode::UNAUTHORIZED, json!({"error": "Unauthorized"})); }; let expected = dev_token_from_mode(&auth_mode); let Some(expected) = expected else { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); return json_response(StatusCode::UNAUTHORIZED, json!({"error": "Unauthorized"})); }; if !validate_dev_token_format(&payload.token) || !dev_token_matches(&payload.token, &expected) { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); return json_response(StatusCode::UNAUTHORIZED, json!({"error": "Unauthorized"})); } @@ -492,7 +487,7 @@ async fn callback_github( Query(params): Query, headers: HeaderMap, ) -> Response { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); if !auth_method_enabled(&auth_mode, ServerAuthMethod::Github) { return json_response(StatusCode::UNAUTHORIZED, json!({"error": "Unauthorized"})); @@ -819,7 +814,7 @@ async fn logout( auth_slot.replace(context); } } else if session_cookie_present(&headers) { - auth_slot.replace(invalid_auth_context()); + auth_slot.replace(RequestAuthContext::invalid()); } jar.private_mut(&key).remove( Cookie::build((SESSION_COOKIE_NAME, "")) diff --git a/lib/crates/fabro-slack/src/dispatch.rs b/lib/crates/fabro-slack/src/dispatch.rs index 26d57b1b8..3e9e940cd 100644 --- a/lib/crates/fabro-slack/src/dispatch.rs +++ b/lib/crates/fabro-slack/src/dispatch.rs @@ -56,7 +56,11 @@ fn event_actor(payload: &serde_json::Value) -> Option { .to_string(); let user_id = event["user"].as_str()?.to_string(); let user_name = event["user_name"].as_str().map(str::to_string); - Some(fabro_types::Principal::slack(team_id, user_id, user_name)) + Some(fabro_types::Principal::Slack { + team_id, + user_id, + user_name, + }) } #[cfg(test)] diff --git a/lib/crates/fabro-slack/src/interaction.rs b/lib/crates/fabro-slack/src/interaction.rs index 448fa6fbe..3754c64e2 100644 --- a/lib/crates/fabro-slack/src/interaction.rs +++ b/lib/crates/fabro-slack/src/interaction.rs @@ -62,7 +62,11 @@ fn interaction_actor(payload: &Value) -> Option { .as_str() .or_else(|| user["username"].as_str()) .map(str::to_string); - Some(Principal::slack(team_id, user_id, user_name)) + Some(Principal::Slack { + team_id, + user_id, + user_name, + }) } /// Extract selected checkbox values from `payload.state.values`. @@ -105,14 +109,11 @@ mod tests { assert_eq!(result.run_id, "run-1"); assert_eq!(result.qid, "q-1"); assert_eq!(result.answer.value, AnswerValue::Yes); - assert_eq!( - result.actor, - fabro_types::Principal::slack( - "T123".to_string(), - "U123".to_string(), - Some("ada".to_string()) - ) - ); + assert_eq!(result.actor, fabro_types::Principal::Slack { + team_id: "T123".to_string(), + user_id: "U123".to_string(), + user_name: Some("ada".to_string()), + }); } #[test] diff --git a/lib/crates/fabro-types/src/event_envelope.rs b/lib/crates/fabro-types/src/event_envelope.rs index 8ee812c48..cac7dc13f 100644 --- a/lib/crates/fabro-types/src/event_envelope.rs +++ b/lib/crates/fabro-types/src/event_envelope.rs @@ -72,11 +72,11 @@ mod tests { session_id: Some("ses_42".to_string()), parent_session_id: Some("ses_root".to_string()), tool_call_id: Some("tool_call_xyz".to_string()), - actor: Some(Principal::agent( - Some("ses_42".to_string()), - Some("ses_root".to_string()), - Some("claude-sonnet".to_string()), - )), + actor: Some(Principal::Agent { + session_id: Some("ses_42".to_string()), + parent_session_id: Some("ses_root".to_string()), + model: Some("claude-sonnet".to_string()), + }), body: EventBody::RunCompleted(RunCompletedProps { duration_ms: 100, artifact_count: 1, diff --git a/lib/crates/fabro-types/src/lib.rs b/lib/crates/fabro-types/src/lib.rs index d04249069..1abb953a7 100644 --- a/lib/crates/fabro-types/src/lib.rs +++ b/lib/crates/fabro-types/src/lib.rs @@ -53,7 +53,7 @@ pub use interview::{InterviewQuestionRecord, QuestionType}; pub use outcome::{ FailureCategory, FailureDetail, NodeResult, Outcome, OutcomeMeta, StageOutcome, StageState, }; -pub use principal::{AuthMethod, Principal, PrincipalLogFields, SystemActorKind, UserPrincipal}; +pub use principal::{AuthMethod, Principal, SystemActorKind, UserPrincipal}; pub use pull_request::{ PullRequestDetail, PullRequestGithubDetail, PullRequestRecord, PullRequestRef, PullRequestUser, }; diff --git a/lib/crates/fabro-types/src/principal.rs b/lib/crates/fabro-types/src/principal.rs index a983f7b87..618e8d5dc 100644 --- a/lib/crates/fabro-types/src/principal.rs +++ b/lib/crates/fabro-types/src/principal.rs @@ -1,4 +1,5 @@ use serde::{Deserialize, Serialize}; +use strum::{Display, IntoStaticStr}; use crate::{IdpIdentity, RunId}; @@ -39,34 +40,23 @@ pub enum Principal { Anonymous, } -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, IntoStaticStr)] #[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] pub enum AuthMethod { Github, DevToken, } -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Display)] #[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] pub enum SystemActorKind { Engine, Watchdog, Timeout, } -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct PrincipalLogFields { - pub principal_kind: &'static str, - pub user_auth_method: Option<&'static str>, - pub idp_issuer: Option, - pub idp_subject: Option, - pub login: Option, - pub run_id: Option, - pub delivery_id: Option, - pub team_id: Option, - pub user_id: Option, -} - impl Principal { #[must_use] pub fn user(identity: IdpIdentity, login: String, auth_method: AuthMethod) -> Self { @@ -77,48 +67,6 @@ impl Principal { }) } - #[must_use] - pub fn worker(run_id: RunId) -> Self { - Self::Worker { run_id } - } - - #[must_use] - pub fn webhook(delivery_id: String) -> Self { - Self::Webhook { delivery_id } - } - - #[must_use] - pub fn slack(team_id: String, user_id: String, user_name: Option) -> Self { - Self::Slack { - team_id, - user_id, - user_name, - } - } - - #[must_use] - pub fn agent( - session_id: Option, - parent_session_id: Option, - model: Option, - ) -> Self { - Self::Agent { - session_id, - parent_session_id, - model, - } - } - - #[must_use] - pub fn system(system_kind: SystemActorKind) -> Self { - Self::System { system_kind } - } - - #[must_use] - pub fn anonymous() -> Self { - Self::Anonymous - } - #[must_use] pub fn user_identity(&self) -> Option<&IdpIdentity> { match self { @@ -127,6 +75,19 @@ impl Principal { } } + #[must_use] + pub fn kind(&self) -> &'static str { + match self { + Self::User(_) => "user", + Self::Worker { .. } => "worker", + Self::Webhook { .. } => "webhook", + Self::Slack { .. } => "slack", + Self::Agent { .. } => "agent", + Self::System { .. } => "system", + Self::Anonymous => "anonymous", + } + } + #[must_use] pub fn display(&self) -> String { match self { @@ -148,104 +109,16 @@ impl Principal { .. } => session_id.clone(), Self::Agent { .. } => "agent".to_string(), - Self::System { system_kind } => format!("system:{system_kind:?}").to_lowercase(), + Self::System { system_kind } => format!("system:{system_kind}"), Self::Anonymous => "anonymous".to_string(), } } - - #[must_use] - pub fn log_fields(&self) -> PrincipalLogFields { - match self { - Self::User(user) => PrincipalLogFields { - principal_kind: "user", - user_auth_method: Some(user.auth_method.as_str()), - idp_issuer: Some(user.identity.issuer().to_string()), - idp_subject: Some(user.identity.subject().to_string()), - login: Some(user.login.clone()), - run_id: None, - delivery_id: None, - team_id: None, - user_id: None, - }, - Self::Worker { run_id } => PrincipalLogFields { - principal_kind: "worker", - user_auth_method: None, - idp_issuer: None, - idp_subject: None, - login: None, - run_id: Some(run_id.to_string()), - delivery_id: None, - team_id: None, - user_id: None, - }, - Self::Webhook { delivery_id } => PrincipalLogFields { - principal_kind: "webhook", - user_auth_method: None, - idp_issuer: None, - idp_subject: None, - login: None, - run_id: None, - delivery_id: Some(delivery_id.clone()), - team_id: None, - user_id: None, - }, - Self::Slack { - team_id, user_id, .. - } => PrincipalLogFields { - principal_kind: "slack", - user_auth_method: None, - idp_issuer: None, - idp_subject: None, - login: None, - run_id: None, - delivery_id: None, - team_id: Some(team_id.clone()), - user_id: Some(user_id.clone()), - }, - Self::Agent { .. } => PrincipalLogFields { - principal_kind: "agent", - user_auth_method: None, - idp_issuer: None, - idp_subject: None, - login: None, - run_id: None, - delivery_id: None, - team_id: None, - user_id: None, - }, - Self::System { .. } => PrincipalLogFields { - principal_kind: "system", - user_auth_method: None, - idp_issuer: None, - idp_subject: None, - login: None, - run_id: None, - delivery_id: None, - team_id: None, - user_id: None, - }, - Self::Anonymous => PrincipalLogFields { - principal_kind: "anonymous", - user_auth_method: None, - idp_issuer: None, - idp_subject: None, - login: None, - run_id: None, - delivery_id: None, - team_id: None, - user_id: None, - }, - } - } } impl AuthMethod { #[must_use] pub fn as_str(self) -> &'static str { - match self { - Self::Github => "github", - Self::DevToken => "dev_token", - } + self.into() } } @@ -253,7 +126,7 @@ impl AuthMethod { mod tests { use serde_json::json; - use super::{AuthMethod, Principal, SystemActorKind}; + use super::{AuthMethod, Principal, SystemActorKind, UserPrincipal}; use crate::{IdpIdentity, fixtures}; fn identity() -> IdpIdentity { @@ -280,7 +153,9 @@ mod tests { #[test] fn system_principal_uses_system_kind_field() { - let principal = Principal::system(SystemActorKind::Watchdog); + let principal = Principal::System { + system_kind: SystemActorKind::Watchdog, + }; assert_eq!( serde_json::to_value(&principal).unwrap(), @@ -295,16 +170,26 @@ mod tests { fn round_trips_all_variants() { let variants = [ Principal::user(identity(), "octocat".to_string(), AuthMethod::Github), - Principal::worker(fixtures::RUN_1), - Principal::webhook("delivery-1".to_string()), - Principal::slack("T1".to_string(), "U1".to_string(), Some("ada".to_string())), - Principal::agent( - Some("session".to_string()), - Some("parent".to_string()), - Some("gpt".to_string()), - ), - Principal::system(SystemActorKind::Engine), - Principal::anonymous(), + Principal::Worker { + run_id: fixtures::RUN_1, + }, + Principal::Webhook { + delivery_id: "delivery-1".to_string(), + }, + Principal::Slack { + team_id: "T1".to_string(), + user_id: "U1".to_string(), + user_name: Some("ada".to_string()), + }, + Principal::Agent { + session_id: Some("session".to_string()), + parent_session_id: Some("parent".to_string()), + model: Some("gpt".to_string()), + }, + Principal::System { + system_kind: SystemActorKind::Engine, + }, + Principal::Anonymous, ]; for principal in variants { @@ -315,21 +200,25 @@ mod tests { } #[test] - fn projects_log_fields() { - let user = Principal::user(identity(), "octocat".to_string(), AuthMethod::DevToken); - let fields = user.log_fields(); + fn auth_method_as_str_matches_serde() { + assert_eq!(AuthMethod::Github.as_str(), "github"); + assert_eq!(AuthMethod::DevToken.as_str(), "dev_token"); + } - assert_eq!(fields.principal_kind, "user"); - assert_eq!(fields.user_auth_method, Some("dev_token")); - assert_eq!(fields.idp_issuer.as_deref(), Some("https://github.com")); - assert_eq!(fields.idp_subject.as_deref(), Some("12345")); - assert_eq!(fields.login.as_deref(), Some("octocat")); + #[test] + fn system_actor_kind_displays_snake_case() { + assert_eq!(SystemActorKind::Engine.to_string(), "engine"); + assert_eq!(SystemActorKind::Watchdog.to_string(), "watchdog"); + assert_eq!(SystemActorKind::Timeout.to_string(), "timeout"); + } - let worker = Principal::worker(fixtures::RUN_1); - assert_eq!(worker.log_fields().principal_kind, "worker"); - assert_eq!( - worker.log_fields().run_id, - Some(fixtures::RUN_1.to_string()) - ); + #[test] + fn user_principal_kind_is_user() { + let principal = Principal::User(UserPrincipal { + identity: identity(), + login: "octocat".to_string(), + auth_method: AuthMethod::Github, + }); + assert_eq!(principal.kind(), "user"); } } diff --git a/lib/crates/fabro-types/src/run_event/mod.rs b/lib/crates/fabro-types/src/run_event/mod.rs index e8bdfb27c..b65eb38b1 100644 --- a/lib/crates/fabro-types/src/run_event/mod.rs +++ b/lib/crates/fabro-types/src/run_event/mod.rs @@ -1015,14 +1015,11 @@ mod tests { ); assert_eq!(parsed.tool_call_id.as_deref(), Some("call_1")); let actor = parsed.actor.as_ref().expect("actor present"); - assert_eq!( - actor, - &Principal::agent( - Some("ses_child".to_string()), - Some("ses_parent".to_string()), - Some("claude-sonnet".to_string()), - ) - ); + assert_eq!(actor, &Principal::Agent { + session_id: Some("ses_child".to_string()), + parent_session_id: Some("ses_parent".to_string()), + model: Some("claude-sonnet".to_string()), + }); let serialized = parsed.to_value().unwrap(); assert_eq!(serialized["stage_id"], value["stage_id"]); diff --git a/lib/crates/fabro-workflow/src/error.rs b/lib/crates/fabro-workflow/src/error.rs index 2cbee58f8..297b696e1 100644 --- a/lib/crates/fabro-workflow/src/error.rs +++ b/lib/crates/fabro-workflow/src/error.rs @@ -1829,6 +1829,7 @@ mod tests { failure: failure.clone(), will_retry: false, duration_ms: 0, + actor: None, }; // 4. Verify classification survived all the way through diff --git a/lib/crates/fabro-workflow/src/event.rs b/lib/crates/fabro-workflow/src/event.rs index 0b66d39cf..378679f2d 100644 --- a/lib/crates/fabro-workflow/src/event.rs +++ b/lib/crates/fabro-workflow/src/event.rs @@ -220,6 +220,8 @@ pub enum Event { failure: FailureDetail, will_retry: bool, duration_ms: u64, + #[serde(default, skip_serializing_if = "Option::is_none")] + actor: Option, }, StageRetrying { node_id: String, @@ -1507,7 +1509,7 @@ fn stored_event_fields(event: &Event, scope: Option<&StageScope>) -> StoredEvent fn stored_event_fields_for_variant(event: &Event) -> StoredEventFields { match event { Event::RunCreated { provenance, .. } => StoredEventFields { - actor: provenance.as_ref().and_then(actor_from_provenance), + actor: provenance.as_ref().and_then(|p| p.subject.clone()), ..StoredEventFields::default() }, Event::RunCancelRequested { actor } @@ -1520,7 +1522,6 @@ fn stored_event_fields_for_variant(event: &Event) -> StoredEventFields { ..StoredEventFields::default() }, Event::StageCompleted { node_id, name, .. } - | Event::StageFailed { node_id, name, .. } | Event::StageStarted { node_id, name, .. } | Event::StageRetrying { node_id, name, .. } => { let node_id_str = node_id.clone(); @@ -1531,6 +1532,21 @@ fn stored_event_fields_for_variant(event: &Event) -> StoredEventFields { ..StoredEventFields::default() } } + Event::StageFailed { + node_id, + name, + actor, + .. + } => { + let node_id_str = node_id.clone(); + let node_label = default_node_label(Some(&node_id_str), Some(name.clone())); + StoredEventFields { + node_id: Some(node_id_str), + node_label, + actor: actor.clone(), + ..StoredEventFields::default() + } + } Event::ParallelStarted { node_id, visit, .. } | Event::ParallelCompleted { node_id, visit, .. } => { let node_id_str = node_id.clone(); @@ -1614,17 +1630,15 @@ fn stored_event_fields_for_variant(event: &Event) -> StoredEventFields { } Event::StallWatchdogTimeout { node, .. } => { let mut fields = node_stored_fields(Some(node.clone())); - fields.actor = Some(Principal::system(SystemActorKind::Watchdog)); + fields.actor = Some(Principal::System { + system_kind: SystemActorKind::Watchdog, + }); fields } _ => StoredEventFields::default(), } } -fn actor_from_provenance(provenance: &RunProvenance) -> Option { - provenance.subject.clone() -} - fn agent_tool_call_id(event: &AgentEvent) -> Option<&str> { match event { AgentEvent::ToolCallStarted { tool_call_id, .. } @@ -1639,18 +1653,18 @@ fn agent_actor_for_event( parent_session_id: Option<&str>, ) -> Option { match event { - AgentEvent::AssistantMessage { model, .. } => Some(Principal::agent( - session_id.map(str::to_string), - parent_session_id.map(str::to_string), - Some(model.clone()), - )), + AgentEvent::AssistantMessage { model, .. } => Some(Principal::Agent { + session_id: session_id.map(str::to_string), + parent_session_id: parent_session_id.map(str::to_string), + model: Some(model.clone()), + }), AgentEvent::ToolCallStarted { .. } | AgentEvent::ToolCallOutputDelta { .. } - | AgentEvent::ToolCallCompleted { .. } => Some(Principal::agent( - session_id.map(str::to_string), - parent_session_id.map(str::to_string), - None, - )), + | AgentEvent::ToolCallCompleted { .. } => Some(Principal::Agent { + session_id: session_id.map(str::to_string), + parent_session_id: parent_session_id.map(str::to_string), + model: None, + }), _ => None, } } @@ -3336,6 +3350,7 @@ mod tests { ), will_retry: true, duration_ms: 5000, + actor: None, }); assert_eq!(stored.event_name(), "stage.failed"); @@ -3731,7 +3746,11 @@ mod tests { assert_eq!(stored.tool_call_id.as_deref(), Some("call_abc")); assert_eq!( stored.actor, - Some(Principal::agent(Some("ses_1".to_string()), None, None)) + Some(Principal::Agent { + session_id: Some("ses_1".to_string()), + parent_session_id: None, + model: None, + }) ); assert_eq!(stored.parallel_group_id, Some(StageId::new("fanout", 2))); assert_eq!( @@ -3981,14 +4000,11 @@ mod tests { parent_session_id: None, }); let actor = stored.actor.as_ref().expect("actor set"); - assert_eq!( - actor, - &Principal::agent( - Some("ses_agent".to_string()), - None, - Some("claude-sonnet".to_string()), - ) - ); + assert_eq!(actor, &Principal::Agent { + session_id: Some("ses_agent".to_string()), + parent_session_id: None, + model: Some("claude-sonnet".to_string()), + }); } #[test] @@ -4002,7 +4018,9 @@ mod tests { assert_eq!(stored.node_id.as_deref(), Some("code")); assert_eq!( stored.actor, - Some(Principal::system(SystemActorKind::Watchdog)) + Some(Principal::System { + system_kind: SystemActorKind::Watchdog, + }) ); } diff --git a/lib/crates/fabro-workflow/src/handler/human.rs b/lib/crates/fabro-workflow/src/handler/human.rs index 924f098eb..5e3f7bb9d 100644 --- a/lib/crates/fabro-workflow/src/handler/human.rs +++ b/lib/crates/fabro-workflow/src/handler/human.rs @@ -293,7 +293,9 @@ impl Handler for HumanHandler { self.emit( &services.run.emitter, &Event::InterviewTimeout { - actor: Some(Principal::system(SystemActorKind::Timeout)), + actor: Some(Principal::System { + system_kind: SystemActorKind::Timeout, + }), question_id: question_id.clone(), question: question_text, stage: node.id.clone(), @@ -334,7 +336,9 @@ impl Handler for HumanHandler { self.emit( &services.run.emitter, &Event::InterviewInterrupted { - actor: Some(Principal::system(SystemActorKind::Engine)), + actor: Some(Principal::System { + system_kind: SystemActorKind::Engine, + }), question_id: question_id.clone(), question: question_text, stage: node.id.clone(), diff --git a/lib/crates/fabro-workflow/src/lifecycle/event.rs b/lib/crates/fabro-workflow/src/lifecycle/event.rs index 4dc95e512..0c8cc9880 100644 --- a/lib/crates/fabro-workflow/src/lifecycle/event.rs +++ b/lib/crates/fabro-workflow/src/lifecycle/event.rs @@ -10,7 +10,7 @@ use fabro_core::lifecycle::{ }; use fabro_core::outcome::NodeResult; use fabro_core::state::ExecutionState; -use fabro_types::RunId; +use fabro_types::{Principal, RunId, SystemActorKind}; use super::circuit_breaker::CircuitBreakerLifecycle; use super::git::GitCheckpointResult; @@ -67,6 +67,18 @@ fn snapshot_failure_signatures( (loop_sigs, restart_sigs) } +fn actor_for_stage_failure(failure: &FailureDetail) -> Option { + if failure.category == FailureCategory::TransientInfra + && failure.message.starts_with("handler timed out after ") + { + Some(Principal::System { + system_kind: SystemActorKind::Timeout, + }) + } else { + None + } +} + fn response_from_outcome(node_id: &str, outcome: &Outcome) -> Option { outcome .context_updates @@ -206,16 +218,19 @@ impl RunLifecycle for EventLifecycle { let scope = stage_scope_for(state, &gv.id); let duration_ms = crate::millis_u64(ctx.result.duration); + let failure = outcome.failure.clone().unwrap_or_else(|| { + FailureDetail::new("handler failed", FailureCategory::TransientInfra) + }); + let actor = actor_for_stage_failure(&failure); self.emitter.emit_scoped( &Event::StageFailed { node_id: gv.id.clone(), name: gv.label().to_string(), index: stage_index, - failure: outcome.failure.clone().unwrap_or_else(|| { - FailureDetail::new("handler failed", FailureCategory::TransientInfra) - }), + failure, will_retry: true, duration_ms, + actor, }, &scope, ); @@ -254,16 +269,19 @@ impl RunLifecycle for EventLifecycle { snapshot_failure_signatures(&self.circuit_breaker); if outcome.status.is_failure() { + let failure = outcome.failure.clone().unwrap_or_else(|| { + FailureDetail::new("handler failed", FailureCategory::Deterministic) + }); + let actor = actor_for_stage_failure(&failure); self.emitter.emit_scoped( &Event::StageFailed { node_id: gv.id.clone(), name: gv.label().to_string(), index: stage_index, - failure: outcome.failure.clone().unwrap_or_else(|| { - FailureDetail::new("handler failed", FailureCategory::Deterministic) - }), + failure, will_retry: false, duration_ms, + actor, }, &scope, ); diff --git a/lib/crates/fabro-workflow/src/pipeline/execute/tests.rs b/lib/crates/fabro-workflow/src/pipeline/execute/tests.rs index 942727558..ff70db28c 100644 --- a/lib/crates/fabro-workflow/src/pipeline/execute/tests.rs +++ b/lib/crates/fabro-workflow/src/pipeline/execute/tests.rs @@ -17,7 +17,7 @@ use fabro_hooks::HookSettings; use fabro_interview::AutoApproveInterviewer; use fabro_sandbox::SandboxSpec; use fabro_store::Database; -use fabro_types::{RunId, WorkflowSettings, fixtures, format_blob_ref}; +use fabro_types::{Principal, RunId, SystemActorKind, WorkflowSettings, fixtures, format_blob_ref}; use object_store::memory::InMemory; use super::*; @@ -833,6 +833,21 @@ async fn timeout_causes_fail_status_record() { assert_eq!(status.outcome, StageOutcome::Failed { retry_requested: false, }); + + let events = executed.engine.run.run_store.list_events().await.unwrap(); + let stage_failed = events + .iter() + .map(|envelope| &envelope.event) + .find(|event| { + event.event_name() == "stage.failed" && event.node_id.as_deref() == Some("work") + }) + .expect("work stage failed event should be persisted"); + assert_eq!( + stage_failed.actor, + Some(Principal::System { + system_kind: SystemActorKind::Timeout, + }) + ); } #[tokio::test]