refactor(auth): simplify canonical-origin validation and OAuth handlers

Inline the validate_canonical_origin wrapper, move reload-failure logging
to the single caller with accurate wording, drop a hand-rolled tracing
capture layer from tests, and migrate web_auth OAuth handlers to
state.canonical_origin() so the is_empty/resolve guards fall out.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Bryan Helmkamp 2026-04-22 00:15:48 -04:00
parent 537a5125cb
commit 0a587d11cb
No known key found for this signature in database
4 changed files with 24 additions and 120 deletions

View file

@ -3,13 +3,6 @@ use url::Url;
use crate::server::EnvLookup;
pub(crate) fn validate_canonical_origin(
resolved: &ResolvedServerSettings,
env_lookup: &EnvLookup,
) -> Result<(), String> {
resolve_canonical_origin(resolved, env_lookup).map(|_| ())
}
pub(crate) fn resolve_canonical_origin(
resolved: &ResolvedServerSettings,
env_lookup: &EnvLookup,
@ -22,8 +15,7 @@ pub(crate) fn resolve_canonical_origin(
.value;
let parsed = Url::parse(&value).map_err(|_| canonical_origin_error(&value))?;
let scheme = parsed.scheme();
if !matches!(scheme, "http" | "https") || parsed.host_str().is_none() {
if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() {
return Err(canonical_origin_error(&value));
}

View file

@ -27,7 +27,7 @@ use tokio::time::interval;
use tracing::{error, info, warn};
use crate::bind::{self, Bind, BindRequest};
use crate::canonical_origin::validate_canonical_origin;
use crate::canonical_origin::resolve_canonical_origin;
use crate::github_webhooks::{TailscaleFunnelManager, WEBHOOK_ROUTE, WEBHOOK_SECRET_ENV};
use crate::ip_allowlist::{GitHubMetaResolver, IpAllowlistConfig, resolve_ip_allowlist_config};
use crate::jwt_auth::resolve_auth_mode_with_lookup;
@ -511,8 +511,7 @@ where
build_artifact_object_store(&resolved_server_settings)?;
let artifact_store = fabro_store::ArtifactStore::new(artifact_object_store, artifact_prefix);
let env_lookup: EnvLookup = Arc::new(|name| std::env::var(name).ok());
validate_canonical_origin(&resolved_server_settings, &env_lookup)
.map_err(anyhow::Error::msg)?;
resolve_canonical_origin(&resolved_server_settings, &env_lookup).map_err(anyhow::Error::msg)?;
let state = build_app_state(AppStateConfig {
settings: Arc::clone(&shared_settings),
registry_factory_override: None,
@ -610,8 +609,13 @@ where
.expect("config lock poisoned");
*cfg != effective
};
if changed && state_for_poll.replace_settings(effective).is_ok() {
info!("Server config reloaded");
if changed {
match state_for_poll.replace_settings(effective) {
Ok(()) => info!("Server config reloaded"),
Err(err) => {
warn!(error = %err, "Rejected reloaded server config, keeping previous");
}
}
}
}
Err(e) => {

View file

@ -108,7 +108,7 @@ use ulid::Ulid;
use crate::auth::{self, GithubEndpoints, auth_translation_middleware, demo_routing_middleware};
use crate::bind::Bind;
use crate::canonical_origin::{resolve_canonical_origin, validate_canonical_origin};
use crate::canonical_origin::resolve_canonical_origin;
use crate::error::ApiError;
use crate::github_webhooks::{
WEBHOOK_ROUTE, WEBHOOK_SECRET_ENV, parse_event_metadata, verify_signature,
@ -793,15 +793,7 @@ impl AppState {
.join("\n")
)
})?);
let resolved_ref = Arc::as_ref(&resolved);
if let Err(error) = validate_canonical_origin(resolved_ref, &self.env_lookup) {
let error = anyhow::anyhow!(error);
warn!(
error = %error,
"Failed to resolve reloaded server config, keeping previous"
);
return Err(error);
}
resolve_canonical_origin(&resolved, &self.env_lookup).map_err(anyhow::Error::msg)?;
*self.settings.write().expect("settings lock poisoned") = settings;
*self
@ -7317,7 +7309,6 @@ mod tests {
use std::path::{Path, PathBuf};
#[cfg(unix)]
use std::process::Stdio;
use std::sync::{Arc as StdArc, Mutex as StdMutex};
use axum::body::Body;
use axum::http::{Request, header};
@ -7327,10 +7318,6 @@ mod tests {
use fabro_types::{InterviewQuestionRecord, InterviewQuestionType, RunBlobId, RunId, fixtures};
use serde_json::json;
use tower::ServiceExt;
use tracing::field::{Field, Visit};
use tracing::{Event, Subscriber};
use tracing_subscriber::layer::{Context, SubscriberExt};
use tracing_subscriber::{Layer, Registry};
use super::*;
use crate::github_webhooks::compute_signature;
@ -7438,54 +7425,6 @@ mod tests {
})
}
#[derive(Debug)]
struct LogCapture {
level: tracing::Level,
target: String,
fields: Vec<(String, String)>,
}
#[derive(Default)]
struct LogCaptureVisitor {
fields: Vec<(String, String)>,
}
impl Visit for LogCaptureVisitor {
fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
self.fields
.push((field.name().to_string(), format!("{value:?}")));
}
}
struct LogCaptureLayer {
events: StdArc<StdMutex<Vec<LogCapture>>>,
}
impl<S: Subscriber> Layer<S> for LogCaptureLayer {
fn on_event(&self, event: &Event<'_>, _ctx: Context<'_, S>) {
if event.metadata().target() != "fabro_server::server" {
return;
}
let mut visitor = LogCaptureVisitor::default();
event.record(&mut visitor);
self.events.lock().unwrap().push(LogCapture {
level: *event.metadata().level(),
target: event.metadata().target().to_string(),
fields: visitor.fields,
});
}
}
fn capture_logs<T>(f: impl FnOnce() -> T) -> (T, StdArc<StdMutex<Vec<LogCapture>>>) {
let events = StdArc::new(StdMutex::new(Vec::<LogCapture>::new()));
let subscriber = Registry::default().with(LogCaptureLayer {
events: StdArc::clone(&events),
});
let result = tracing::subscriber::with_default(subscriber, f);
(result, events)
}
fn canonical_origin_settings(url: &str) -> SettingsLayer {
fabro_config::parse_settings_layer(&format!(
r#"
@ -7538,11 +7477,9 @@ type = "http"
},
);
let (result, logs) = capture_logs(|| {
state.replace_settings(canonical_origin_settings("{{ env.FABRO_WEB_URL }}"))
});
let err = result.expect_err("invalid canonical origin should be rejected");
let err = state
.replace_settings(canonical_origin_settings("{{ env.FABRO_WEB_URL }}"))
.expect_err("invalid canonical origin should be rejected");
assert!(
err.to_string()
.contains("server.web.url is required and must be an absolute http(s) URL"),
@ -7552,18 +7489,6 @@ type = "http"
state.canonical_origin().unwrap(),
"http://valid.example.com".to_string()
);
let logs = logs.lock().unwrap();
assert!(logs.iter().any(|event| {
event.level == tracing::Level::WARN
&& event.target == "fabro_server::server"
&& event.fields.iter().any(|(name, value)| {
name == "message"
&& value.contains(
"Failed to resolve reloaded server config, keeping previous",
)
})
}));
}
}

View file

@ -7,7 +7,7 @@ use axum::routing::{get, post};
use axum::{Extension, Json, Router};
use cookie::time::Duration;
use cookie::{Cookie, CookieJar, Key, SameSite};
use fabro_types::settings::{InterpString, ServerAuthMethod};
use fabro_types::settings::ServerAuthMethod;
use fabro_types::{IdpIdentity, RunAuthMethod};
use fabro_util::dev_token::validate_dev_token_format;
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, utf8_percent_encode};
@ -257,10 +257,6 @@ fn callback_error_redirect(
}
}
fn resolve_interp(state: &AppState, value: &InterpString) -> anyhow::Result<String> {
state.resolve_interp(value)
}
fn auth_methods_from_mode(auth_mode: &AuthMode) -> Vec<String> {
match auth_mode {
AuthMode::Enabled(config) => config
@ -385,7 +381,7 @@ async fn login_github(
json!({"error": "GitHub App client_id is not configured"}),
);
};
let client_id = match resolve_interp(state.as_ref(), client_id) {
let client_id = match state.resolve_interp(client_id) {
Ok(client_id) => client_id,
Err(err) => {
warn!(error = %err, "OAuth login failed: client_id could not be resolved");
@ -395,23 +391,13 @@ async fn login_github(
);
}
};
let web_url = match resolve_interp(state.as_ref(), &settings.web.url) {
let web_url = match state.canonical_origin() {
Ok(web_url) => web_url,
Err(err) => {
warn!(error = %err, "OAuth login failed: server.web.url could not be resolved");
return json_response(
StatusCode::CONFLICT,
json!({"error": format!("server.web.url could not be resolved: {err}")}),
);
warn!(error = %err, "OAuth login failed: server.web.url is invalid");
return json_response(StatusCode::CONFLICT, json!({"error": err}));
}
};
if web_url.is_empty() {
warn!("OAuth login failed: server.web.url not configured");
return json_response(
StatusCode::CONFLICT,
json!({"error": "server.web.url is not configured"}),
);
}
let state_token = format!("fabro-{}", ulid::Ulid::new());
let authorize_url = fabro_http::Url::parse_with_params(
@ -537,7 +523,7 @@ async fn callback_github(
json!({"error": "GitHub App client_id is not configured"}),
);
};
let client_id = match resolve_interp(state.as_ref(), client_id) {
let client_id = match state.resolve_interp(client_id) {
Ok(client_id) => client_id,
Err(err) => {
error!(error = %err, "OAuth callback failed: client_id could not be resolved");
@ -554,14 +540,11 @@ async fn callback_github(
json!({"error": "GITHUB_APP_CLIENT_SECRET is not configured"}),
);
};
let web_url = match resolve_interp(state.as_ref(), &settings.web.url) {
let web_url = match state.canonical_origin() {
Ok(web_url) => web_url,
Err(err) => {
error!(error = %err, "OAuth callback failed: server.web.url could not be resolved");
return json_response(
StatusCode::CONFLICT,
json!({"error": format!("server.web.url could not be resolved: {err}")}),
);
error!(error = %err, "OAuth callback failed: server.web.url is invalid");
return json_response(StatusCode::CONFLICT, json!({"error": err}));
}
};