Fix principal auth gap regressions

This commit is contained in:
Bryan Helmkamp 2026-05-02 10:02:12 -04:00
parent 29c45498b0
commit f6b8d1acdb
No known key found for this signature in database
28 changed files with 661 additions and 401 deletions

View file

@ -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

View file

@ -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 |

View file

@ -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();

View file

@ -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;

View file

@ -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]

View file

@ -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,
})
}
}

View file

@ -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(

View file

@ -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 },
}
}
}

View file

@ -15,7 +15,9 @@ pub struct QueueInterviewer {
impl QueueInterviewer {
#[must_use]
pub fn new(answers: VecDeque<Answer>) -> Self {
Self::with_actor(answers, Principal::system(SystemActorKind::Engine))
Self::with_actor(answers, Principal::System {
system_kind: SystemActorKind::Engine,
})
}
#[must_use]

View file

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

View file

@ -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<Arc<AppState>>,
Extension(auth_mode): Extension<AuthMode>,
Query(params): Query<CliStartParams>,
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<Arc<AppState>>,
Extension(auth_mode): Extension<AuthMode>,
Query(params): Query<CliResumeParams>,
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<Arc<AppState>>,
Extension(auth_mode): Extension<AuthMode>,
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<Arc<AppState>>,
Extension(auth_mode): Extension<AuthMode>,
RequestAuth(auth_slot): RequestAuth,
headers: HeaderMap,
body: Result<Json<CliTokenRequest>, 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<SessionCookie> {
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"));

View file

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

View file

@ -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<String> {
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;
}

View file

@ -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<VerifiedAuth, ApiError> {
if token.starts_with("fabro_refresh_") {
if token.starts_with(REFRESH_TOKEN_PREFIX) {
info!(
path = %parts.uri.path(),
"Refresh token presented at protected endpoint"

View file

@ -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<Mutex<RequestAuthContext>>);
// 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<Arc<AppState>> 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<Arc<AppState>> 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<Arc<AppState>> 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<Arc<AppState>> for RequireCommandLog {
let stream = stream
.parse::<CommandOutputStream>()
.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::<AuthContextSlot>()
.map_or_else(Principal::anonymous, |slot| slot.snapshot().principal)
.map_or_else(RequestAuthContext::initial, AuthContextSlot::snapshot)
}
pub(crate) fn require_user(slot: &AuthContextSlot) -> Result<UserPrincipal, ApiError> {
@ -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<String, ApiError> {
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"));
}
}

View file

@ -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(),

View file

@ -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<SessionCoo
Some(session)
}
fn session_cookie_present(headers: &HeaderMap) -> 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<RequestAuthContext> {
pub(crate) fn auth_context_from_session(session: &SessionCookie) -> Option<RequestAuthContext> {
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<RequestAuthConte
))
}
fn invalid_auth_context() -> RequestAuthContext {
RequestAuthContext::rejected(AuthStatus::Invalid, Some("unauthorized"))
}
fn read_private_oauth_state(headers: &HeaderMap, key: &Key) -> Option<OAuthStateCookie> {
let jar = parse_cookie_header(headers);
jar.private(key)
@ -337,17 +332,17 @@ async fn login_dev_token(
payload: Result<Json<DevTokenLoginRequest>, 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<OAuthCallbackParams>,
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, ""))

View file

@ -56,7 +56,11 @@ fn event_actor(payload: &serde_json::Value) -> Option<fabro_types::Principal> {
.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)]

View file

@ -62,7 +62,11 @@ fn interaction_actor(payload: &Value) -> Option<Principal> {
.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]

View file

@ -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,

View file

@ -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,
};

View file

@ -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<String>,
pub idp_subject: Option<String>,
pub login: Option<String>,
pub run_id: Option<String>,
pub delivery_id: Option<String>,
pub team_id: Option<String>,
pub user_id: Option<String>,
}
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<String>) -> Self {
Self::Slack {
team_id,
user_id,
user_name,
}
}
#[must_use]
pub fn agent(
session_id: Option<String>,
parent_session_id: Option<String>,
model: Option<String>,
) -> 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");
}
}

View file

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

View file

@ -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

View file

@ -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<Principal>,
},
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<Principal> {
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<Principal> {
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,
})
);
}

View file

@ -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(),

View file

@ -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<Principal> {
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<String> {
outcome
.context_updates
@ -206,16 +218,19 @@ impl RunLifecycle<WorkflowGraph> 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<WorkflowGraph> 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,
);

View file

@ -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]