use std::any::Any;
use std::sync::Arc;
use axum::body::Body;
use axum::http::{HeaderValue, Request, header};
use axum::{
Router,
extract::DefaultBodyLimit,
middleware,
middleware::Next,
response::{IntoResponse, Redirect, Response},
routing::{get, post},
};
use tower_http::catch_panic::CatchPanicLayer;
use tower_http::set_header::SetResponseHeaderLayer;
use tracing::{Span, info};
use crate::{handlers, middlewares};
use acme_proxy_core::config::Config;
use acme_proxy_core::error::Problem;
use acme_proxy_core::routes;
use acme_proxy_jobs::metrics;
use acme_proxy_net::challenge;
use acme_proxy_signer as signer;
use acme_proxy_store::db::Database;
use crate::profile::Profile;
#[derive(Clone)]
pub struct AppState {
pub database: Arc<Database>,
pub config: Arc<Config>,
pub profile: Arc<Profile>,
pub audit: Arc<acme_proxy_jobs::auditor::Auditor>,
pub jobs: acme_proxy_jobs::jobs::JobQueue,
}
fn http01_stores(profiles: &[Arc<Profile>]) -> Vec<Arc<dyn signer::Http01TokenStore>> {
let mut stores: Vec<Arc<dyn signer::Http01TokenStore>> = Vec::new();
for profile in profiles {
if let Some(store) = profile.signer_info.http01_tokens()
&& !stores.iter().any(|existing| Arc::ptr_eq(existing, &store))
{
stores.push(store);
}
}
stores
}
pub fn security_headers() -> (
SetResponseHeaderLayer<HeaderValue>,
SetResponseHeaderLayer<HeaderValue>,
SetResponseHeaderLayer<HeaderValue>,
) {
(
SetResponseHeaderLayer::overriding(
header::STRICT_TRANSPORT_SECURITY,
HeaderValue::from_static("max-age=31536000; includeSubDomains"),
),
SetResponseHeaderLayer::overriding(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
),
SetResponseHeaderLayer::overriding(
header::X_FRAME_OPTIONS,
HeaderValue::from_static("DENY"),
),
)
}
pub fn panic_message(err: &(dyn Any + Send)) -> &str {
err.downcast_ref::<&'static str>()
.copied()
.or_else(|| err.downcast_ref::<String>().map(String::as_str))
.unwrap_or("a handler panicked")
}
fn acme_panic_response(err: Box<dyn Any + Send + 'static>) -> Response {
tracing::error!(
event = "request_handler_panicked",
outcome = "failure",
listener = "acme",
error = %panic_message(err.as_ref()),
);
Problem::server_internal("Internal server error").into_response()
}
pub fn catch_panic_acme() -> CatchPanicLayer<fn(Box<dyn Any + Send + 'static>) -> Response> {
CatchPanicLayer::custom(acme_panic_response as fn(Box<dyn Any + Send + 'static>) -> Response)
}
pub fn build_app(
database: Arc<Database>,
config: Arc<Config>,
profiles: Vec<Arc<Profile>>,
audit: Arc<acme_proxy_jobs::auditor::Auditor>,
metrics: Arc<metrics::Metrics>,
jobs: acme_proxy_jobs::jobs::JobQueue,
) -> Router {
let mut root = Router::new()
.route("/", get(|| async { Redirect::temporary("/health") }))
.route("/health", get(handlers::get_health_check));
let stores = http01_stores(&profiles);
if !stores.is_empty() {
info!(
event = "http_01_responder_mounted",
outcome = "advisory",
path = challenge::http_01::WELL_KNOWN_PREFIX,
stores = stores.len(),
"a reverse proxy must forward or redirect \
http://<identifier>:80/.well-known/acme-challenge/ here for the upstream to reach it"
);
root = root.merge(
Router::new()
.route(
&format!("{}{{token}}", challenge::http_01::WELL_KNOWN_PREFIX),
get(handlers::get_challenge_file),
)
.with_state(handlers::Http01Stores(Arc::new(stores))),
);
}
let mut acme = Router::new();
for profile in &profiles {
let path = profile.path.clone();
acme = acme.nest(
&path,
build_router(
database.clone(),
config.clone(),
profile.clone(),
audit.clone(),
jobs.clone(),
),
);
}
let server = &config.server;
let acme = acme
.layer(middleware::from_fn_with_state(
middlewares::admission::Admission::new(
server.max_concurrent_requests,
server.admission_wait_ms,
server.request_timeout_ms,
),
middlewares::admission::admission_middleware,
))
.layer(DefaultBodyLimit::max(server.max_body_bytes));
let app = root.merge(acme);
let app = app.layer(catch_panic_acme());
let app = if config.metrics.enabled {
app.layer(middleware::from_fn_with_state(
metrics,
middlewares::metrics::record_request,
))
} else {
app
};
app.layer(security_headers())
.layer(middleware::from_fn(
middlewares::access::add_access_middleware,
))
}
pub fn metrics_app(metrics: Arc<metrics::Metrics>) -> Router {
Router::new()
.route("/metrics", get(handlers::get_metrics))
.with_state(handlers::MetricsState(metrics))
.layer(middleware::from_fn(
middlewares::access::add_access_middleware,
))
}
pub fn build_router(
database: Arc<Database>,
config: Arc<Config>,
profile: Arc<Profile>,
audit: Arc<acme_proxy_jobs::auditor::Auditor>,
jobs: acme_proxy_jobs::jobs::JobQueue,
) -> Router {
let filter = profile.filter.clone();
let state = AppState {
database: database.clone(),
config,
profile: profile.clone(),
audit,
jobs,
};
let profile_name = profile.name.clone();
let index_link =
HeaderValue::from_str(&format!("<{}/directory>;rel=\"index\"", profile.base_url));
let router = Router::<AppState>::new()
.route(
routes::DIRECTORY,
get(handlers::get_directory).post(handlers::post_directory),
)
.route(
routes::NEW_NONCE,
get(handlers::get_new_nonce)
.head(handlers::head_new_nonce)
.post(handlers::post_new_nonce),
)
.route(routes::NEW_ACCOUNT, post(handlers::post_new_account))
.route("/acct/{id}", post(handlers::post_account))
.route("/acct/{id}/orders", post(handlers::post_account_orders))
.route(routes::KEY_CHANGE, post(handlers::post_key_change))
.route(routes::NEW_ORDER, post(handlers::post_new_order))
.route("/order/{id}", post(handlers::post_order))
.route("/order/{id}/finalize", post(handlers::post_finalize))
.route("/authz/{id}", post(handlers::post_authz))
.route("/chall/{id}", post(handlers::post_challenge))
.route("/certificate/{id}", post(handlers::post_certificate))
.route(routes::REVOKE_CERT, post(handlers::post_revoke_cert))
.route(
&format!("{}/{{id}}", routes::RENEWAL_INFO),
get(handlers::get_renewal_info),
)
.route(routes::CRL, get(handlers::get_crl))
.route(routes::CA_CHAIN, get(handlers::get_ca_chain))
.method_not_allowed_fallback(|| async {
Problem::method_not_allowed("This resource must be read with POST-as-GET")
})
.fallback(|| async { Problem::not_found("No such resource") })
.with_state(state)
.layer(middleware::from_fn_with_state(
filter,
middlewares::filter::add_filter_middleware,
))
.layer(middleware::from_fn_with_state(
database.clone(),
middlewares::nonce::add_nonce_middleware,
));
let router = match index_link {
Ok(value) => router.layer(middleware::from_fn_with_state(
value,
middlewares::index_link::add_index_link_middleware,
)),
Err(error) => {
tracing::error!(
event = "request_index_link_header_invalid",
outcome = "failure",
base_url = %profile.base_url,
error = %error,
);
router
}
};
router.layer(middleware::from_fn(
move |request: Request<Body>, next: Next| {
let name = profile_name.clone();
async move {
Span::current().record("profile", &*name);
next.run(request).await
}
},
))
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::to_bytes;
use axum::http::StatusCode;
use axum::routing::get;
use tower::ServiceExt;
#[test]
fn panic_message_covers_every_payload_shape() {
assert_eq!(panic_message(&"boom"), "boom");
assert_eq!(panic_message(&String::from("boom")), "boom");
assert_eq!(panic_message(&0u8), "a handler panicked");
}
#[tokio::test]
async fn acme_panic_response_is_a_problem_document() {
let response = acme_panic_response(Box::new("secret internal detail"));
assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(
response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok()),
Some("application/problem+json"),
);
let body = to_bytes(response.into_body(), 64 * 1024).await.unwrap();
let problem: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(problem["type"], "urn:ietf:params:acme:error:serverInternal");
assert_eq!(problem["status"], 500);
assert!(
!body_contains(&body, "secret internal detail"),
"the panic message must not reach the client",
);
}
fn body_contains(bytes: &[u8], needle: &str) -> bool {
std::str::from_utf8(bytes)
.map(|s| s.contains(needle))
.unwrap_or(false)
}
async fn boom() -> &'static str {
panic!("this handler panics on purpose")
}
fn app() -> Router {
Router::new()
.route("/ok", get(|| async { "ok" }))
.route("/boom", get(boom))
.layer(catch_panic_acme())
}
#[tokio::test]
async fn a_panicking_route_answers_a_problem_document() {
let response = app()
.oneshot(Request::get("/boom").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(
response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok()),
Some("application/problem+json"),
);
}
#[tokio::test]
async fn the_layer_is_transparent_on_the_happy_path() {
let response = app()
.oneshot(Request::get("/ok").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
}