diff --git a/proto b/proto index 0509e3eb..cd1afb72 160000 --- a/proto +++ b/proto @@ -1 +1 @@ -Subproject commit 0509e3ebf2508aca829e81ae2bb5423d748a49a8 +Subproject commit cd1afb727774b51118c470fafbb5efd808e4d405 diff --git a/src/grpc.rs b/src/grpc.rs index 1c35a869..93ffc097 100644 --- a/src/grpc.rs +++ b/src/grpc.rs @@ -40,7 +40,6 @@ use crate::{ // connected clients type ClientMap = HashMap>>; - // requests awaiting a `CoreResponse`, keyed by the id stamped into the outgoing `CoreRequest` type ResultMap = HashMap>; @@ -275,15 +274,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 { diff --git a/src/handlers/mfa_config.rs b/src/handlers/mfa_config.rs new file mode 100644 index 00000000..31c97032 --- /dev/null +++ b/src/handlers/mfa_config.rs @@ -0,0 +1,99 @@ +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. +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 52bd068a..ada1f175 100644 --- a/src/handlers/mod.rs +++ b/src/handlers/mod.rs @@ -8,15 +8,16 @@ use axum::{ http::{HeaderName, header::FORWARDED, request::Parts}, }; use axum_client_ip::{RightmostForwarded, RightmostXForwardedFor}; -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; @@ -63,6 +64,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 2cc8f290..2ffae12a 100644 --- a/src/http.rs +++ b/src/http.rs @@ -50,7 +50,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, }; @@ -80,7 +80,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 AppState { @@ -420,26 +420,31 @@ async fn build_tls_config(cert_pem: &str, key_pem: &str) -> anyhow::Result) -> Router, - tls_active: Arc, -) -> anyhow::Result { - // Collect all API routes into a separate router to scope API-only middleware. - let api_router = Router::new().nest( +/// All `/api/v1` routes without middleware, so tests can drive them with a fake Core. +pub(crate) fn api_router() -> 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)), - ); - let mut api_router = apply_rate_limit(api_router); + ) +} + +pub(crate) fn build_router( + state: AppState, + apply_rate_limit: impl FnOnce(Router) -> Router, + tls_active: Arc, +) -> anyhow::Result { + // Collect all API routes into a separate router to scope API-only middleware. + let mut api_router = apply_rate_limit(api_router()); api_router = api_router.layer(middleware::from_fn_with_state( state.clone(), ensure_configured, diff --git a/src/tests/cookies.rs b/src/tests/cookies.rs index 997cd8c8..8259beee 100644 --- a/src/tests/cookies.rs +++ b/src/tests/cookies.rs @@ -1,10 +1,10 @@ -use std::sync::{Arc, RwLock, atomic::AtomicBool}; +use std::sync::{Arc, atomic::AtomicBool}; use axum::{ body::Body, http::{Request, StatusCode, header}, }; -use axum_extra::extract::cookie::{Cookie, Key, SameSite}; +use axum_extra::extract::cookie::{Cookie, SameSite}; use tokio::sync::mpsc; use tonic::Status; use tower::ServiceExt; @@ -16,7 +16,7 @@ use crate::{ AuthInfoResponse, CoreRequest, EnrollmentStartResponse, PasswordResetStartResponse, core_response, }, - tests::support::{test_proxy_server, test_public_settings}, + tests::support::{cookie_key, test_proxy_server, test_public_settings}, }; /// A router wired to a `ProxyServer` whose Core responses the test drives by hand. @@ -29,7 +29,7 @@ struct TestApp { /// Build a router whose cookie `Secure` attribute reflects `public_url`. Passing `None` leaves /// the state at its default, standing in for a Core that never sent `PublicSettings`. fn test_app(public_url: Option<&str>) -> TestApp { - let cookie_key = Arc::new(RwLock::new(Some(Key::generate()))); + let cookie_key = cookie_key(); let server = test_proxy_server(Arc::clone(&cookie_key)); if public_url.is_some() { server diff --git a/src/tests/mfa_config.rs b/src/tests/mfa_config.rs new file mode 100644 index 00000000..d4f25a16 --- /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::support::{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 91e3e0fe..29186f8c 100644 --- a/src/tests/mod.rs +++ b/src/tests/mod.rs @@ -1,3 +1,4 @@ mod cookies; +mod mfa_config; mod mtls; pub(crate) mod support; diff --git a/src/tests/mtls.rs b/src/tests/mtls.rs index b38eb92b..a2c695f2 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::support::{cookie_key, test_proxy_server}; +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 = test_proxy_server(cookie_key()); server.configure(make_tls_config(certs)); // Find a free port, drop the listener, pass the addr to run(). @@ -212,7 +189,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 = test_proxy_server(cookie_key()); // configure() is deliberately NOT called. let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0); let result = server diff --git a/src/tests/support.rs b/src/tests/support.rs index f506f5f2..7bceb2ec 100644 --- a/src/tests/support.rs +++ b/src/tests/support.rs @@ -1,14 +1,30 @@ //! Shared fixtures for the proxy's unit and handler tests. use std::{ - path::PathBuf, + 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, proto::PublicSettings}; +use crate::{ + grpc::ProxyServer, + http::{AppState, api_router}, + proto::{CoreError, PublicSettings, core_request, core_response}, +}; + +pub(crate) fn cookie_key() -> Arc>> { + Arc::new(RwLock::new(Some(Key::generate()))) +} /// A `ProxyServer` with throwaway channels, suitable for driving handlers in tests. pub(crate) fn test_proxy_server(cookie_key: Arc>>) -> ProxyServer { @@ -18,7 +34,7 @@ pub(crate) fn test_proxy_server(cookie_key: Arc>>) -> ProxySe let (_, logs_rx) = mpsc::channel(1); ProxyServer::new( cookie_key, - PathBuf::new(), + temp_dir(), reset_tx, https_cert_tx, clear_https_tx, @@ -36,3 +52,63 @@ pub(crate) fn test_public_settings(public_url: Option<&str>) -> PublicSettings { public_url: public_url.map(str::to_owned), } } + +/// An API router backed by a fake Core that answers every request through `core`. +pub(crate) 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 = test_proxy_server(Arc::clone(&cookie_key)); + let mut requests = server.register_test_client(); + 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.resolve_test_response(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. +pub(crate) 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) +}