From deba5d05c21a2dc5b815de1126d1970723e731cb Mon Sep 17 00:00:00 2001 From: Aleksander <170264518+t-aleksander@users.noreply.github.com> Date: Wed, 9 Sep 2026 14:27:01 +0200 Subject: [PATCH 1/5] mfa setup grpc endpoint --- src/grpc.rs | 74 +++++++++++----- src/handlers/mfa_config.rs | 101 ++++++++++++++++++++++ src/handlers/mod.rs | 13 ++- src/handlers/register_mfa.rs | 124 +++++++++++++++----------- src/http.rs | 36 ++++---- src/tests/mfa_config.rs | 163 +++++++++++++++++++++++++++++++++++ src/tests/mod.rs | 109 +++++++++++++++++++++++ src/tests/mtls.rs | 33 ++----- 8 files changed, 536 insertions(+), 117 deletions(-) create mode 100644 src/handlers/mfa_config.rs create mode 100644 src/tests/mfa_config.rs diff --git a/src/grpc.rs b/src/grpc.rs index 2f294473..e4f520d6 100644 --- a/src/grpc.rs +++ b/src/grpc.rs @@ -40,6 +40,7 @@ use crate::{ // connected clients type ClientMap = HashMap>>; +type ResultMap = Arc>>>; #[derive(Debug, Clone, Default)] pub struct TlsConfig { @@ -54,7 +55,7 @@ pub struct TlsConfig { pub(crate) struct ProxyServer { current_id: Arc, clients: Arc>, - results: Arc>>>, + results: ResultMap, pub(crate) connected: Arc, pub(crate) core_version: Arc>>, /// Whether the password reset option is displayed on the Edge home page. @@ -195,15 +196,21 @@ impl ProxyServer { device_info: Some(device_info), payload: Some(payload), }; - if let Err(err) = client_tx.send(Ok(res)) { - error!("Failed to send CoreRequest: {err}"); - return Err(ApiError::Unexpected("Failed to send CoreRequest".into())); - } + // Core can answer before this thread resumes, so the receiver must be registered + // before the request leaves. let (tx, rx) = oneshot::channel(); self.results .write() .expect("Failed to acquire lock on results hashmap when sending CoreRequest") .insert(id, tx); + if let Err(err) = client_tx.send(Ok(res)) { + error!("Failed to send CoreRequest: {err}"); + self.results + .write() + .expect("Failed to acquire lock on results hashmap when sending CoreRequest") + .remove(&id); + return Err(ApiError::Unexpected("Failed to send CoreRequest".into())); + } self.connected.store(true, Ordering::Relaxed); Ok(rx) } else { @@ -247,6 +254,39 @@ impl Clone for ProxyServer { } } +/// Hands a Core response to the HTTP handler that awaits request `id`. +fn deliver(results: &ResultMap, id: u64, payload: core_response::Payload) { + let maybe_tx = results + .write() + .expect("Failed to acquire lock on results hashmap when processing response") + .remove(&id); + if let Some(tx) = maybe_tx { + if let Err(err) = tx.send(payload) { + error!("Failed to send message to rx {:?}", err.type_id()); + } + } else { + error!("Missing receiver for response #{id}"); + } +} + +#[cfg(test)] +impl ProxyServer { + /// Registers a fake Core connection and returns the stream of requests sent to it. + pub(crate) fn connect_fake_core(&self) -> mpsc::UnboundedReceiver> { + let (tx, rx) = mpsc::unbounded_channel(); + self.clients + .write() + .unwrap() + .insert(SocketAddr::from(([127, 0, 0, 1], 1)), tx); + rx + } + + /// Answers request `id` as Core would over the bidi stream. + pub(crate) fn respond(&self, id: u64, payload: core_response::Payload) { + deliver(&self.results, id, payload); + } +} + #[tonic::async_trait] impl proxy_server::Proxy for ProxyServer { type BidiStream = UnboundedReceiverStream>; @@ -317,9 +357,7 @@ impl proxy_server::Proxy for ProxyServer { if let Err(err) = https_cert_tx.send((certs.cert_pem, certs.key_pem)) { - error!( - "Failed to broadcast HTTPS certificates: {err}" - ); + error!("Failed to broadcast HTTPS certificates: {err}"); } } core_response::Payload::ClearHttpsCerts(_) => { @@ -339,16 +377,7 @@ impl proxy_server::Proxy for ProxyServer { Ordering::Relaxed, ); } - other => { - let maybe_rx = results.write().expect("Failed to acquire lock on results hashmap when processing response").remove(&response.id); - if let Some(rx) = maybe_rx { - if let Err(err) = rx.send(other) { - error!("Failed to send message to rx {:?}", err.type_id()); - } - } else { - error!("Missing receiver for response #{}", response.id); - } - } + other => deliver(&results, response.id, other), } } } @@ -364,8 +393,13 @@ impl proxy_server::Proxy for ProxyServer { } info!("Defguard core client disconnected: {address}"); connected.store(false, Ordering::Relaxed); - clients.write().expect("Failed to acquire lock on clients hashmap when removing \ - disconnected client").remove(&address); + clients + .write() + .expect( + "Failed to acquire lock on clients hashmap when removing \ + disconnected client", + ) + .remove(&address); } .instrument(tracing::Span::current()), ); diff --git a/src/handlers/mfa_config.rs b/src/handlers/mfa_config.rs new file mode 100644 index 00000000..df731aff --- /dev/null +++ b/src/handlers/mfa_config.rs @@ -0,0 +1,101 @@ +use axum::{Json, Router, extract::State, routing::post}; + +use super::register_mfa::{code_mfa_setup_finish, code_mfa_setup_start}; +use crate::{ + error::ApiError, + handlers::get_core_response, + http::AppState, + proto::{ + CodeMfaSetupFinishRequest, CodeMfaSetupFinishResponse, CodeMfaSetupStartRequest, + CodeMfaSetupStartResponse, DeviceInfo, MfaConfigAuthorizeRequest, + MfaConfigAuthorizeResponse, MfaConfigSendCodeRequest, MfaConfigStartRequest, + MfaConfigStartResponse, core_request, core_response, + }, +}; + +/// MFA factor configuration for an enrolled desktop client. +/// +/// The client is not a browser, so every request carries its token in the JSON body. +pub(crate) fn router() -> Router { + Router::new() + .route("/start", post(start_mfa_config)) + .route("/send-code", post(send_mfa_config_code)) + .route("/authorize", post(authorize_mfa_config)) + .route("/setup/start", post(start_mfa_setup)) + .route("/setup/finish", post(finish_mfa_setup)) +} + +#[instrument(level = "debug", skip(state, req))] +async fn start_mfa_config( + State(state): State, + device_info: DeviceInfo, + Json(req): Json, +) -> Result, ApiError> { + info!("Starting MFA configuration for device {}", req.pubkey); + let rx = state + .grpc_server + .send(core_request::Payload::MfaConfigStart(req), device_info)?; + let payload = get_core_response(rx, None).await?; + if let core_response::Payload::MfaConfigStart(response) = payload { + Ok(Json(response)) + } else { + error!("Received invalid gRPC response type, expected MfaConfigStart"); + Err(ApiError::InvalidResponseType) + } +} + +#[instrument(level = "debug", skip(state, req))] +async fn send_mfa_config_code( + State(state): State, + device_info: DeviceInfo, + Json(req): Json, +) -> Result<(), ApiError> { + info!("Sending MFA configuration email code"); + let rx = state + .grpc_server + .send(core_request::Payload::MfaConfigSendCode(req), device_info)?; + let payload = get_core_response(rx, None).await?; + if let core_response::Payload::MfaConfigSendCode(_) = payload { + Ok(()) + } else { + error!("Received invalid gRPC response type, expected MfaConfigSendCode"); + Err(ApiError::InvalidResponseType) + } +} + +#[instrument(level = "debug", skip(state, req))] +async fn authorize_mfa_config( + State(state): State, + device_info: DeviceInfo, + Json(req): Json, +) -> Result, ApiError> { + info!("Authorizing MFA configuration session"); + let rx = state + .grpc_server + .send(core_request::Payload::MfaConfigAuthorize(req), device_info)?; + let payload = get_core_response(rx, None).await?; + if let core_response::Payload::MfaConfigAuthorize(response) = payload { + Ok(Json(response)) + } else { + error!("Received invalid gRPC response type, expected MfaConfigAuthorize"); + Err(ApiError::InvalidResponseType) + } +} + +#[instrument(level = "debug", skip(state, req))] +async fn start_mfa_setup( + State(state): State, + device_info: DeviceInfo, + Json(req): Json, +) -> Result, ApiError> { + code_mfa_setup_start(&state, device_info, req).await +} + +#[instrument(level = "debug", skip(state, req))] +async fn finish_mfa_setup( + State(state): State, + device_info: DeviceInfo, + Json(req): Json, +) -> Result, ApiError> { + code_mfa_setup_finish(&state, device_info, req).await +} diff --git a/src/handlers/mod.rs b/src/handlers/mod.rs index 187de9ba..32d307a6 100644 --- a/src/handlers/mod.rs +++ b/src/handlers/mod.rs @@ -2,15 +2,16 @@ use std::time::Duration; use axum::{extract::FromRequestParts, http::request::Parts}; use axum_client_ip::{InsecureClientIp, LeftmostXForwardedFor}; -use axum_extra::{TypedHeader, headers::UserAgent}; +use axum_extra::{TypedHeader, extract::PrivateCookieJar, headers::UserAgent}; use tokio::{sync::oneshot::Receiver, time}; use tonic::Code; use super::proto::DeviceInfo; -use crate::{error::ApiError, proto::core_response::Payload}; +use crate::{error::ApiError, http::ENROLLMENT_COOKIE_NAME, proto::core_response::Payload}; pub(crate) mod desktop_client_mfa; pub(crate) mod enrollment; +pub(crate) mod mfa_config; pub(crate) mod mobile_client; pub(crate) mod password_reset; pub(crate) mod polling; @@ -21,6 +22,14 @@ const CORE_RESPONSE_TIMEOUT: Duration = Duration::from_secs(5); const CLIENT_VERSION_HEADER: &str = "defguard-client-version"; const CLIENT_PLATFORM_HEADER: &str = "defguard-client-platform"; +/// Reads the enrollment token from the private cookie set at enrollment start. +pub(super) fn enrollment_token(cookie_jar: &PrivateCookieJar) -> Result { + cookie_jar + .get(ENROLLMENT_COOKIE_NAME) + .map(|cookie| cookie.value().to_string()) + .ok_or_else(|| ApiError::Unauthorized(String::new())) +} + impl FromRequestParts for DeviceInfo where S: Send + Sync, diff --git a/src/handlers/register_mfa.rs b/src/handlers/register_mfa.rs index 24663904..d241e395 100644 --- a/src/handlers/register_mfa.rs +++ b/src/handlers/register_mfa.rs @@ -1,11 +1,11 @@ -use axum::{Json, Router, extract::State, response::IntoResponse, routing::post}; +use axum::{Json, Router, extract::State, routing::post}; use axum_extra::extract::PrivateCookieJar; use serde::Deserialize; use crate::{ error::ApiError, - handlers::get_core_response, - http::{AppState, ENROLLMENT_COOKIE_NAME}, + handlers::{enrollment_token, get_core_response}, + http::AppState, proto::{ CodeMfaSetupFinishRequest, CodeMfaSetupFinishResponse, CodeMfaSetupStartRequest, CodeMfaSetupStartResponse, DeviceInfo, MfaMethod, core_request, core_response, @@ -18,6 +18,55 @@ pub(crate) fn router() -> Router { .route("/code/finish", post(register_code_mfa_finish)) } +/// Forwards a code MFA setup start to Core. +/// +/// `req.token` is either an enrollment token or an authorized MFA config session token. +pub(super) async fn code_mfa_setup_start( + state: &AppState, + device_info: DeviceInfo, + req: CodeMfaSetupStartRequest, +) -> Result, ApiError> { + debug!("Code MFA setup started"); + reject_non_code_method(req.method)?; + + let rx = state + .grpc_server + .send(core_request::Payload::CodeMfaSetupStart(req), device_info)?; + let payload = get_core_response(rx, None).await?; + match payload { + core_response::Payload::CodeMfaSetupStartResponse(response) => Ok(Json(response)), + _ => Err(ApiError::InvalidResponseType), + } +} + +/// Forwards a code MFA setup finish to Core. See [`code_mfa_setup_start`] for the token. +pub(super) async fn code_mfa_setup_finish( + state: &AppState, + device_info: DeviceInfo, + req: CodeMfaSetupFinishRequest, +) -> Result, ApiError> { + reject_non_code_method(req.method)?; + + let rx = state + .grpc_server + .send(core_request::Payload::CodeMfaSetupFinish(req), device_info)?; + let payload = get_core_response(rx, None).await?; + match payload { + core_response::Payload::CodeMfaSetupFinishResponse(response) => Ok(Json(response)), + _ => Err(ApiError::InvalidResponseType), + } +} + +/// Code MFA setup only knows how to deliver a code by email or TOTP. +fn reject_non_code_method(method: i32) -> Result<(), ApiError> { + if method == MfaMethod::Email as i32 || method == MfaMethod::Totp as i32 { + Ok(()) + } else { + error!("Requested method not supported"); + Err(ApiError::BadRequest("Method not supported.".to_string())) + } +} + #[derive(Debug, Clone, Deserialize)] struct RegisterMfaCodeStartRequest { pub method: MfaMethod, @@ -29,31 +78,17 @@ async fn register_code_mfa_start( device_info: DeviceInfo, cookie_jar: PrivateCookieJar, Json(req): Json, -) -> Result, impl IntoResponse> { - debug!("Register code MFA started"); - let token = cookie_jar - .get(ENROLLMENT_COOKIE_NAME) - .ok_or_else(|| ApiError::Unauthorized(String::new()))? - .value() - .to_string(); - - if req.method != MfaMethod::Email && req.method != MfaMethod::Totp { - error!("Requested method not supported"); - return Err(ApiError::BadRequest("Method not supported.".to_string())); - } - - let rx = state.grpc_server.send( - core_request::Payload::CodeMfaSetupStart(CodeMfaSetupStartRequest { +) -> Result, ApiError> { + let token = enrollment_token(&cookie_jar)?; + code_mfa_setup_start( + &state, + device_info, + CodeMfaSetupStartRequest { token, method: req.method.into(), - }), - device_info, - )?; - let payload = get_core_response(rx, None).await?; - match payload { - core_response::Payload::CodeMfaSetupStartResponse(response) => Ok(Json(response)), - _ => Err(ApiError::InvalidResponseType), - } + }, + ) + .await } #[derive(Debug, Clone, Deserialize)] @@ -68,31 +103,16 @@ async fn register_code_mfa_finish( device_info: DeviceInfo, cookie_jar: PrivateCookieJar, Json(req): Json, -) -> Result, impl IntoResponse> { - let token = cookie_jar - .get(ENROLLMENT_COOKIE_NAME) - .ok_or_else(|| ApiError::Unauthorized(String::new()))? - .value() - .to_string(); - - let code = req.code; - let method = req.method; - - if method != MfaMethod::Totp && method != MfaMethod::Email { - return Err(ApiError::BadRequest("Method not supported".to_string())); - } - - let rx = state.grpc_server.send( - core_request::Payload::CodeMfaSetupFinish(CodeMfaSetupFinishRequest { - token, - code, - method: method as i32, - }), +) -> Result, ApiError> { + let token = enrollment_token(&cookie_jar)?; + code_mfa_setup_finish( + &state, device_info, - )?; - let payload = get_core_response(rx, None).await?; - match payload { - core_response::Payload::CodeMfaSetupFinishResponse(response) => Ok(Json(response)), - _ => Err(ApiError::InvalidResponseType), - } + CodeMfaSetupFinishRequest { + token, + code: req.code, + method: req.method as i32, + }, + ) + .await } diff --git a/src/http.rs b/src/http.rs index ce3b8776..06b457b0 100644 --- a/src/http.rs +++ b/src/http.rs @@ -46,7 +46,7 @@ use crate::{ enterprise::handlers::{desktop_client_posture, openid_login}, error::ApiError, grpc::{ProxyServer, TlsConfig}, - handlers::{desktop_client_mfa, enrollment, password_reset, polling}, + handlers::{desktop_client_mfa, enrollment, mfa_config, password_reset, polling}, setup::ProxySetupServer, }; @@ -73,7 +73,7 @@ pub use crate::setup::{CORE_CLIENT_CERT_NAME, GRPC_CA_CERT_NAME, GRPC_CERT_NAME, #[derive(Clone)] pub(crate) struct AppState { pub(crate) grpc_server: ProxyServer, - cookie_key: Arc>>, + pub(crate) cookie_key: Arc>>, } impl FromRef for Key { @@ -366,6 +366,24 @@ async fn build_tls_config(cert_pem: &str, key_pem: &str) -> anyhow::Result Router { + Router::new().nest( + "/api/v1", + Router::new() + .nest("/enrollment", enrollment::router()) + .nest("/password-reset", password_reset::router()) + .nest("/client-mfa", desktop_client_mfa::router()) + .nest("/mfa-config", mfa_config::router()) + .nest("/openid", openid_login::router()) + .nest("/posture", desktop_client_posture::router()) + .route("/poll", post(polling::info)) + .route("/health", get(healthcheck)) + .route("/health-grpc", get(healthcheckgrpc)) + .route("/info", get(app_info)), + ) +} + pub async fn run_server( env_config: EnvConfig, tls_config: Option, @@ -515,19 +533,7 @@ pub async fn run_server( }; // Collect all API routes into a separate router to scope API-only middleware. - let mut api_router = Router::new().nest( - "/api/v1", - Router::new() - .nest("/enrollment", enrollment::router()) - .nest("/password-reset", password_reset::router()) - .nest("/client-mfa", desktop_client_mfa::router()) - .nest("/openid", openid_login::router()) - .nest("/posture", desktop_client_posture::router()) - .route("/poll", post(polling::info)) - .route("/health", get(healthcheck)) - .route("/health-grpc", get(healthcheckgrpc)) - .route("/info", get(app_info)), - ); + let mut api_router = api_router(); if let Some(conf) = governor_conf { api_router = api_router.layer(GovernorLayer::new(conf)); } diff --git a/src/tests/mfa_config.rs b/src/tests/mfa_config.rs new file mode 100644 index 00000000..a576ac50 --- /dev/null +++ b/src/tests/mfa_config.rs @@ -0,0 +1,163 @@ +//! Drives the MFA config flow through the HTTP API against a scripted fake Core. + +use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, +}; + +use axum::http::StatusCode; +use serde_json::json; + +use super::{app_with_fake_core, post_json}; +use crate::proto::{ + CodeMfaSetupFinishResponse, CodeMfaSetupStartResponse, CoreError, MfaConfigAuthorizeResponse, + MfaConfigSendCodeResponse, MfaConfigStartResponse, MfaMethod, core_request, core_response, +}; + +const SESSION_TOKEN: &str = "mfa-config-session"; +const TOTP: i32 = MfaMethod::Totp as i32; +const EMAIL: i32 = MfaMethod::Email as i32; + +/// Fake Core for a user with no factor: the email fallback authorizes, then TOTP is set up. +fn fallback_then_totp( + steps: Arc, +) -> impl Fn(core_request::Payload) -> core_response::Payload { + move |request| { + steps.fetch_add(1, Ordering::Relaxed); + match request { + core_request::Payload::MfaConfigStart(req) => { + assert_eq!(req.token, "polling-token"); + assert_eq!(req.pubkey, "device-pubkey"); + core_response::Payload::MfaConfigStart(MfaConfigStartResponse { + session_token: SESSION_TOKEN.into(), + available_methods: vec![], + email_fallback: true, + deadline_timestamp: 1_800_000_000, + }) + } + core_request::Payload::MfaConfigSendCode(req) => { + assert_eq!(req.session_token, SESSION_TOKEN); + core_response::Payload::MfaConfigSendCode(MfaConfigSendCodeResponse {}) + } + core_request::Payload::MfaConfigAuthorize(req) => { + assert_eq!(req.session_token, SESSION_TOKEN); + assert_eq!(req.method, EMAIL); + assert_eq!(req.code, "123456"); + core_response::Payload::MfaConfigAuthorize(MfaConfigAuthorizeResponse { + deadline_timestamp: 1_800_003_600, + }) + } + core_request::Payload::CodeMfaSetupStart(req) => { + assert_eq!(req.token, SESSION_TOKEN); + assert_eq!(req.method, TOTP); + core_response::Payload::CodeMfaSetupStartResponse(CodeMfaSetupStartResponse { + totp_secret: Some("JBSWY3DPEHPK3PXP".into()), + }) + } + core_request::Payload::CodeMfaSetupFinish(req) => { + assert_eq!(req.token, SESSION_TOKEN); + assert_eq!(req.method, TOTP); + assert_eq!(req.code, "654321"); + core_response::Payload::CodeMfaSetupFinishResponse(CodeMfaSetupFinishResponse { + recovery_codes: vec!["aaaa-bbbb".into(), "cccc-dddd".into()], + }) + } + _ => panic!("unexpected request to Core"), + } + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn test_mfa_config_flow_forwards_session_token() { + let steps = Arc::new(AtomicUsize::new(0)); + let app = app_with_fake_core(fallback_then_totp(Arc::clone(&steps))); + + let (status, body) = post_json( + &app, + "/api/v1/mfa-config/start", + &json!({ "token": "polling-token", "pubkey": "device-pubkey" }), + ) + .await; + assert_eq!(status, StatusCode::OK, "{body}"); + assert_eq!(body["session_token"], SESSION_TOKEN); + assert_eq!(body["email_fallback"], true); + assert_eq!(body["available_methods"], json!([])); + let session_token = body["session_token"].as_str().unwrap().to_owned(); + + let (status, body) = post_json( + &app, + "/api/v1/mfa-config/send-code", + &json!({ "session_token": session_token }), + ) + .await; + assert_eq!(status, StatusCode::OK, "{body}"); + + let (status, body) = post_json( + &app, + "/api/v1/mfa-config/authorize", + &json!({ "session_token": session_token, "method": EMAIL, "code": "123456" }), + ) + .await; + assert_eq!(status, StatusCode::OK, "{body}"); + assert_eq!(body["deadline_timestamp"], 1_800_003_600); + + let (status, body) = post_json( + &app, + "/api/v1/mfa-config/setup/start", + &json!({ "token": session_token, "method": TOTP }), + ) + .await; + assert_eq!(status, StatusCode::OK, "{body}"); + assert_eq!(body["totp_secret"], "JBSWY3DPEHPK3PXP"); + + let (status, body) = post_json( + &app, + "/api/v1/mfa-config/setup/finish", + &json!({ "token": session_token, "method": TOTP, "code": "654321" }), + ) + .await; + assert_eq!(status, StatusCode::OK, "{body}"); + assert_eq!(body["recovery_codes"], json!(["aaaa-bbbb", "cccc-dddd"])); + + assert_eq!( + steps.load(Ordering::Relaxed), + 5, + "Core must see every step once" + ); +} + +#[tokio::test] +async fn test_mfa_setup_rejects_unsupported_method_before_core() { + let app = app_with_fake_core(|_| panic!("Core must not be called")); + + for path in [ + "/api/v1/mfa-config/setup/start", + "/api/v1/mfa-config/setup/finish", + ] { + let (status, _) = post_json( + &app, + path, + &json!({ "token": SESSION_TOKEN, "method": MfaMethod::Oidc as i32, "code": "1" }), + ) + .await; + assert_eq!(status, StatusCode::BAD_REQUEST, "{path}"); + } +} + +#[tokio::test] +async fn test_mfa_config_core_error_maps_to_http_status() { + let app = app_with_fake_core(|_| { + core_response::Payload::CoreError(CoreError { + status_code: tonic::Code::Unauthenticated as i32, + message: "invalid token".into(), + }) + }); + + let (status, body) = post_json( + &app, + "/api/v1/mfa-config/authorize", + &json!({ "session_token": "stale", "method": EMAIL, "code": "000000" }), + ) + .await; + assert_eq!(status, StatusCode::UNAUTHORIZED, "{body}"); +} diff --git a/src/tests/mod.rs b/src/tests/mod.rs index cd5c8686..3d2598a3 100644 --- a/src/tests/mod.rs +++ b/src/tests/mod.rs @@ -1 +1,110 @@ +use std::{ + env::temp_dir, + panic::{AssertUnwindSafe, catch_unwind}, + sync::{Arc, RwLock}, +}; + +use axum::{ + Router, + body::{Body, to_bytes}, + http::{Request, StatusCode, header}, +}; +use axum_extra::extract::cookie::Key; +use serde::Serialize; +use tokio::sync::{Mutex, broadcast, mpsc}; +use tower::ServiceExt; + +use crate::{ + grpc::ProxyServer, + http::{AppState, api_router}, + proto::{CoreError, core_request, core_response}, +}; + +mod mfa_config; mod mtls; + +pub(super) fn cookie_key() -> Arc>> { + Arc::new(RwLock::new(Some(Key::generate()))) +} + +pub(super) fn build_proxy_server(cookie_key: Arc>>) -> ProxyServer { + let (reset_tx, _) = broadcast::channel(1); + let (https_cert_tx, _) = broadcast::channel(1); + let (clear_https_tx, _) = broadcast::channel(1); + let (_, logs_rx) = mpsc::channel(1); + ProxyServer::new( + cookie_key, + temp_dir(), + reset_tx, + https_cert_tx, + clear_https_tx, + None, + Arc::new(Mutex::new(logs_rx)), + false, + ) +} + +/// The API router backed by a fake Core that answers every request with `core`. +/// +/// `core` receives each request payload and returns the response payload. It runs on a +/// separate task, so a handler can await the response like in production. A panic inside +/// `core` would only kill that task, so it is turned into an internal `CoreError` that +/// carries the panic message back through the handler under test. +pub(super) fn app_with_fake_core(core: F) -> Router +where + F: Fn(core_request::Payload) -> core_response::Payload + Send + 'static, +{ + let cookie_key = cookie_key(); + let server = build_proxy_server(Arc::clone(&cookie_key)); + let mut requests = server.connect_fake_core(); + let responder = server.clone(); + tokio::spawn(async move { + while let Some(Ok(request)) = requests.recv().await { + let payload = request.payload.expect("request without payload"); + let response = catch_unwind(AssertUnwindSafe(|| core(payload))).unwrap_or_else(|err| { + let message = err + .downcast_ref::() + .cloned() + .or_else(|| err.downcast_ref::<&str>().map(ToString::to_string)) + .unwrap_or_default(); + core_response::Payload::CoreError(CoreError { + status_code: tonic::Code::Internal as i32, + message: format!("fake Core panicked: {message}"), + }) + }); + responder.respond(request.id, response); + } + }); + api_router().with_state(AppState { + grpc_server: server, + cookie_key, + }) +} + +/// Sends a JSON POST like a desktop client and returns the status with the JSON body. +/// +/// The `X-Forwarded-For` header satisfies the `DeviceInfo` extractor without a socket. +pub(super) async fn post_json( + app: &Router, + path: &str, + body: &T, +) -> (StatusCode, serde_json::Value) { + let request = Request::builder() + .method("POST") + .uri(path) + .header(header::CONTENT_TYPE, "application/json") + .header("X-Forwarded-For", "10.0.0.1") + .body(Body::from(serde_json::to_vec(body).unwrap())) + .unwrap(); + let response = app.clone().oneshot(request).await.unwrap(); + let status = response.status(); + let bytes = to_bytes(response.into_body(), usize::MAX).await.unwrap(); + let json = if bytes.is_empty() { + serde_json::Value::Null + } else { + serde_json::from_slice(&bytes).unwrap_or(serde_json::Value::String( + String::from_utf8_lossy(&bytes).into_owned(), + )) + }; + (status, json) +} diff --git a/src/tests/mtls.rs b/src/tests/mtls.rs index be5b6b6d..a6905cca 100644 --- a/src/tests/mtls.rs +++ b/src/tests/mtls.rs @@ -1,11 +1,8 @@ use std::{ - env::temp_dir, net::{IpAddr, Ipv4Addr, SocketAddr, TcpListener}, - sync::{Arc, RwLock}, time::Duration, }; -use axum_extra::extract::cookie::Key; use defguard_certs::{ CertificateAuthority, Csr, PemLabel, cert_der_to_pem, der_to_pem, generate_key_pair, }; @@ -14,7 +11,7 @@ use rustls::crypto::aws_lc_rs; use tokio::{ net::TcpStream, spawn, - sync::{Mutex, broadcast, mpsc, oneshot}, + sync::oneshot, time::{Instant, sleep}, }; use tonic::{ @@ -22,10 +19,8 @@ use tonic::{ transport::{Certificate, Channel, ClientTlsConfig, Endpoint, Identity}, }; -use crate::{ - grpc::{ProxyServer, TlsConfig}, - proto::proxy_client::ProxyClient, -}; +use super::{build_proxy_server, cookie_key}; +use crate::{grpc::TlsConfig, proto::proxy_client::ProxyClient}; struct TestCerts { /// PEM-encoded CA certificate (used as the trust root for both server and client validation). @@ -104,24 +99,6 @@ fn make_tls_config(certs: &TestCerts) -> TlsConfig { } } -fn build_proxy_server() -> ProxyServer { - let (reset_tx, _) = broadcast::channel(1); - let (https_cert_tx, _) = broadcast::channel(1); - let (clear_https_tx, _) = broadcast::channel(1); - let (_, logs_rx) = mpsc::channel(1); - let cookie_key = Arc::new(RwLock::new(Some(Key::generate()))); - ProxyServer::new( - cookie_key, - temp_dir(), - reset_tx, - https_cert_tx, - clear_https_tx, - None, - Arc::new(Mutex::new(logs_rx)), - false, - ) -} - /// Install the rustls AWS-LC crypto provider for the process. /// /// Must be called before any TLS code runs. Safe to call from multiple tests - @@ -137,7 +114,7 @@ fn init_crypto() { /// Waits until the server is accepting TCP connections before returning, so /// callers do not need a fixed sleep to avoid startup races. async fn spawn_test_proxy(certs: &TestCerts) -> (SocketAddr, oneshot::Sender<()>) { - let server = build_proxy_server(); + let server = build_proxy_server(cookie_key()); server.configure(make_tls_config(certs)); // Find a free port, drop the listener, pass the addr to run(). @@ -211,7 +188,7 @@ async fn call_bidi(client: &mut ProxyClient) -> Status { /// `run()` must return `Err` immediately when no `TlsConfig` has been set. #[tokio::test] async fn run_errors_without_tls_config() { - let server = build_proxy_server(); + let server = build_proxy_server(cookie_key()); // configure() is deliberately NOT called. let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0); let result = server From 8f714a8c4b042f2bce5942e8c67ccbd462b3f1b1 Mon Sep 17 00:00:00 2001 From: Aleksander <170264518+t-aleksander@users.noreply.github.com> Date: Thu, 10 Sep 2026 09:18:11 +0200 Subject: [PATCH 2/5] cleanup --- src/tests/mod.rs | 8 +------- 1 file changed, 1 insertion(+), 7 deletions(-) diff --git a/src/tests/mod.rs b/src/tests/mod.rs index 3d2598a3..0fe88f12 100644 --- a/src/tests/mod.rs +++ b/src/tests/mod.rs @@ -44,12 +44,6 @@ pub(super) fn build_proxy_server(cookie_key: Arc>>) -> ProxyS ) } -/// The API router backed by a fake Core that answers every request with `core`. -/// -/// `core` receives each request payload and returns the response payload. It runs on a -/// separate task, so a handler can await the response like in production. A panic inside -/// `core` would only kill that task, so it is turned into an internal `CoreError` that -/// carries the panic message back through the handler under test. pub(super) fn app_with_fake_core(core: F) -> Router where F: Fn(core_request::Payload) -> core_response::Payload + Send + 'static, @@ -83,7 +77,7 @@ where /// Sends a JSON POST like a desktop client and returns the status with the JSON body. /// -/// The `X-Forwarded-For` header satisfies the `DeviceInfo` extractor without a socket. +/// The `X-Forwarded-For` header satisfies the `DeviceInfo` extractor. pub(super) async fn post_json( app: &Router, path: &str, From 583d166287c3674d6703c0cc2c4d0570f9101285 Mon Sep 17 00:00:00 2001 From: Aleksander <170264518+t-aleksander@users.noreply.github.com> Date: Thu, 10 Sep 2026 12:06:58 +0200 Subject: [PATCH 3/5] cleanup --- src/handlers/mfa_config.rs | 2 -- 1 file changed, 2 deletions(-) diff --git a/src/handlers/mfa_config.rs b/src/handlers/mfa_config.rs index df731aff..31c97032 100644 --- a/src/handlers/mfa_config.rs +++ b/src/handlers/mfa_config.rs @@ -14,8 +14,6 @@ use crate::{ }; /// MFA factor configuration for an enrolled desktop client. -/// -/// The client is not a browser, so every request carries its token in the JSON body. pub(crate) fn router() -> Router { Router::new() .route("/start", post(start_mfa_config)) From 22a3246ba939a4f74fa73a7bbfe25a4715b181f9 Mon Sep 17 00:00:00 2001 From: Aleksander <170264518+t-aleksander@users.noreply.github.com> Date: Thu, 10 Sep 2026 12:32:26 +0200 Subject: [PATCH 4/5] proto --- proto | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/proto b/proto index 0509e3eb..f0218228 160000 --- a/proto +++ b/proto @@ -1 +1 @@ -Subproject commit 0509e3ebf2508aca829e81ae2bb5423d748a49a8 +Subproject commit f0218228d68c7d77de1b40b0721bbba978ed846a From 96b9d7119649f0ac7a7a4c4260b094984225a83b Mon Sep 17 00:00:00 2001 From: Aleksander <170264518+t-aleksander@users.noreply.github.com> Date: Mon, 14 Sep 2026 10:30:48 +0200 Subject: [PATCH 5/5] proto --- proto | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/proto b/proto index f0218228..cd1afb72 160000 --- a/proto +++ b/proto @@ -1 +1 @@ -Subproject commit f0218228d68c7d77de1b40b0721bbba978ed846a +Subproject commit cd1afb727774b51118c470fafbb5efd808e4d405