Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion proto
15 changes: 10 additions & 5 deletions src/grpc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,6 @@ use crate::{

// connected clients
type ClientMap = HashMap<SocketAddr, mpsc::UnboundedSender<Result<CoreRequest, Status>>>;

// requests awaiting a `CoreResponse`, keyed by the id stamped into the outgoing `CoreRequest`
type ResultMap = HashMap<u64, oneshot::Sender<core_response::Payload>>;

Expand Down Expand Up @@ -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 {
Expand Down
99 changes: 99 additions & 0 deletions src/handlers/mfa_config.rs
Original file line number Diff line number Diff line change
@@ -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<AppState> {
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<AppState>,
device_info: DeviceInfo,
Json(req): Json<MfaConfigStartRequest>,
) -> Result<Json<MfaConfigStartResponse>, 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<AppState>,
device_info: DeviceInfo,
Json(req): Json<MfaConfigSendCodeRequest>,
) -> 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<AppState>,
device_info: DeviceInfo,
Json(req): Json<MfaConfigAuthorizeRequest>,
) -> Result<Json<MfaConfigAuthorizeResponse>, 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<AppState>,
device_info: DeviceInfo,
Json(req): Json<CodeMfaSetupStartRequest>,
) -> Result<Json<CodeMfaSetupStartResponse>, ApiError> {
code_mfa_setup_start(&state, device_info, req).await
}

#[instrument(level = "debug", skip(state, req))]
async fn finish_mfa_setup(
State(state): State<AppState>,
device_info: DeviceInfo,
Json(req): Json<CodeMfaSetupFinishRequest>,
) -> Result<Json<CodeMfaSetupFinishResponse>, ApiError> {
code_mfa_setup_finish(&state, device_info, req).await
}
13 changes: 11 additions & 2 deletions src/handlers/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<String, ApiError> {
cookie_jar
.get(ENROLLMENT_COOKIE_NAME)
.map(|cookie| cookie.value().to_string())
.ok_or_else(|| ApiError::Unauthorized(String::new()))
}

impl<S> FromRequestParts<S> for DeviceInfo
where
S: Send + Sync,
Expand Down
124 changes: 72 additions & 52 deletions src/handlers/register_mfa.rs
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -18,6 +18,55 @@ pub(crate) fn router() -> Router<AppState> {
.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<Json<CodeMfaSetupStartResponse>, 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<Json<CodeMfaSetupFinishResponse>, 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,
Expand All @@ -29,31 +78,17 @@ async fn register_code_mfa_start(
device_info: DeviceInfo,
cookie_jar: PrivateCookieJar,
Json(req): Json<RegisterMfaCodeStartRequest>,
) -> Result<Json<CodeMfaSetupStartResponse>, 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<Json<CodeMfaSetupStartResponse>, 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)]
Expand All @@ -68,31 +103,16 @@ async fn register_code_mfa_finish(
device_info: DeviceInfo,
cookie_jar: PrivateCookieJar,
Json(req): Json<RegisterMfaCodeFinishRequest>,
) -> Result<Json<CodeMfaSetupFinishResponse>, 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<Json<CodeMfaSetupFinishResponse>, 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
}
27 changes: 16 additions & 11 deletions src/http.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
};

Expand Down Expand Up @@ -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<RwLock<Option<Key>>>,
pub(crate) cookie_key: Arc<RwLock<Option<Key>>>,
}

impl AppState {
Expand Down Expand Up @@ -420,26 +420,31 @@ async fn build_tls_config(cert_pem: &str, key_pem: &str) -> anyhow::Result<Rustl
.context("Failed to build HTTPS TLS configuration from PEM")
}

pub(crate) fn build_router(
state: AppState,
apply_rate_limit: impl FnOnce(Router<AppState>) -> Router<AppState>,
tls_active: Arc<AtomicBool>,
) -> anyhow::Result<Router> {
// 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<AppState> {
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<AppState>) -> Router<AppState>,
tls_active: Arc<AtomicBool>,
) -> anyhow::Result<Router> {
// 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,
Expand Down
Loading
Loading