use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use tokio::time::timeout;
use tracing::{debug, info, warn};
use crate::config::{ChallengeConfig, DnsConfig};
use crate::dns::{HickoryResolver, Resolver, resolver_addr};
pub mod dns_01;
pub mod http_01;
pub mod tls_alpn_01;
pub const HTTP_01: &str = "http-01";
pub const DNS_01: &str = "dns-01";
pub const TLS_ALPN_01: &str = "tls-alpn-01";
pub const KNOWN_TYPES: &[&str] = &[HTTP_01, DNS_01, TLS_ALPN_01];
#[async_trait]
pub trait ChallengeValidator: Send + Sync {
fn typ(&self) -> &'static str;
async fn validate(&self, ctx: &ValidationContext<'_>) -> Result<(), ChallengeError>;
}
#[derive(Debug)]
pub struct ValidationContext<'a> {
pub identifier: &'a str,
pub wildcard: bool,
pub token: &'a str,
pub key_authorization: &'a str,
pub challenge_id: &'a str,
}
#[derive(Debug)]
pub enum ChallengeError {
Connection(String),
Dns(String),
IncorrectResponse(String),
Tls(String),
Unauthorized(String),
Internal(String),
}
impl ChallengeError {
#[must_use]
pub fn kind(&self) -> &'static str {
match self {
Self::Connection(_) => "connection",
Self::Dns(_) => "dns",
Self::IncorrectResponse(_) => "incorrectResponse",
Self::Tls(_) => "tls",
Self::Unauthorized(_) => "unauthorized",
Self::Internal(_) => "internal",
}
}
#[must_use]
pub fn detail(&self) -> &str {
match self {
Self::Connection(detail)
| Self::Dns(detail)
| Self::IncorrectResponse(detail)
| Self::Tls(detail)
| Self::Unauthorized(detail)
| Self::Internal(detail) => detail,
}
}
}
pub struct ChallengeRegistry {
validators: Vec<Arc<dyn ChallengeValidator>>,
enabled: Vec<String>,
bypass: bool,
timeout: Duration,
}
impl std::fmt::Debug for ChallengeRegistry {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ChallengeRegistry")
.field("enabled", &self.enabled)
.field("bypass", &self.bypass)
.field("timeout", &self.timeout)
.finish()
}
}
impl Default for ChallengeRegistry {
fn default() -> Self {
Self {
validators: Vec::new(),
enabled: vec![HTTP_01.to_string()],
bypass: true,
timeout: Duration::from_secs(5),
}
}
}
impl ChallengeRegistry {
#[must_use]
pub fn new(
validators: Vec<Arc<dyn ChallengeValidator>>,
enabled: Vec<String>,
bypass: bool,
timeout: Duration,
) -> Self {
Self {
validators,
enabled,
bypass,
timeout,
}
}
#[must_use]
pub fn enabled_types(&self) -> &[String] {
&self.enabled
}
#[must_use]
pub fn is_bypassed(&self) -> bool {
self.bypass
}
pub fn types_for(&self, wildcard: bool) -> Vec<&str> {
self.enabled
.iter()
.map(String::as_str)
.filter(|typ| !wildcard || *typ == DNS_01)
.collect()
}
pub async fn validate(
&self,
typ: &str,
ctx: &ValidationContext<'_>,
) -> Result<(), ChallengeError> {
if self.bypass {
debug!(
event = "challenge_bypassed",
outcome = "success",
typ,
identifier = ctx.identifier,
challenge_id = ctx.challenge_id,
);
return Ok(());
}
let validator = self
.validators
.iter()
.find(|validator| validator.typ() == typ)
.ok_or_else(|| {
ChallengeError::Internal(format!("no validator registered for {typ}"))
})?;
match timeout(self.timeout, validator.validate(ctx)).await {
Ok(result) => result.inspect_err(|error| {
warn!(
event = "challenge_validation_failed",
outcome = "failure",
typ,
identifier = ctx.identifier,
challenge_id = ctx.challenge_id,
kind = error.kind(),
detail = %error.detail(),
);
}),
Err(_) => {
warn!(
event = "challenge_validation_timeout",
outcome = "failure",
typ,
identifier = ctx.identifier,
challenge_id = ctx.challenge_id,
timeout_ms = crate::millis(self.timeout),
);
Err(ChallengeError::Connection(format!(
"{typ} validation of {} timed out after {}ms",
ctx.identifier,
self.timeout.as_millis()
)))
}
}
}
}
fn validate_enabled(enabled: &[String]) -> anyhow::Result<()> {
if enabled.is_empty() {
anyhow::bail!(
"challenge.enabled is empty: authorizations would carry no challenges, \
so no client could ever prove control of a name"
);
}
for name in enabled {
if !KNOWN_TYPES.contains(&name.as_str()) {
anyhow::bail!("unknown challenge type: {name}");
}
}
Ok(())
}
pub fn build_resolver(addr: Option<std::net::SocketAddr>) -> anyhow::Result<Arc<dyn Resolver>> {
Ok(Arc::new(match addr {
Some(addr) => HickoryResolver::from_address_uncached(addr)
.map_err(|error| anyhow::anyhow!("dns.resolver: {error}"))?,
None => HickoryResolver::from_system_uncached()
.map_err(|error| anyhow::anyhow!("challenge: {error}"))?,
}))
}
pub fn from_config(
cfg: &ChallengeConfig,
dns: &DnsConfig,
proxies: Arc<crate::proxy::OutboundProxies>,
) -> anyhow::Result<Arc<ChallengeRegistry>> {
validate_enabled(&cfg.enabled)?;
let addr = resolver_addr(dns)?;
let timeout = Duration::from_millis(cfg.timeout_ms);
if cfg.bypass {
warn!(
event = "challenge_validation_bypassed",
outcome = "advisory",
enabled = ?cfg.enabled,
"challenge.bypass is on: triggering a challenge marks it valid with no network \
check, so any client that can reach this server can obtain a certificate for \
any name (set challenge.bypass = false)"
);
return Ok(Arc::new(ChallengeRegistry::new(
Vec::new(),
cfg.enabled.clone(),
true,
timeout,
)));
}
let resolver = build_resolver(addr)?;
let outbound = crate::http_client::Outbound::new(resolver.clone(), proxies);
let mut validators: Vec<Arc<dyn ChallengeValidator>> = Vec::with_capacity(cfg.enabled.len());
for name in &cfg.enabled {
let validator: Arc<dyn ChallengeValidator> = match name.as_str() {
HTTP_01 => Arc::new(http_01::Http01Validator::from_config(
&cfg.http_01,
outbound.clone(),
)?),
DNS_01 => Arc::new(dns_01::Dns01Validator::from_config(resolver.clone())),
TLS_ALPN_01 => Arc::new(tls_alpn_01::TlsAlpn01Validator::from_config(
&cfg.tls_alpn_01,
outbound.clone(),
)?),
other => anyhow::bail!("unknown challenge type: {other}"),
};
validators.push(validator);
}
info!(
event = "challenge_validation_enabled",
outcome = "success",
enabled = ?cfg.enabled,
timeout_ms = cfg.timeout_ms,
);
Ok(Arc::new(ChallengeRegistry::new(
validators,
cfg.enabled.clone(),
false,
timeout,
)))
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
struct StubValidator {
typ: &'static str,
outcome: Option<&'static str>,
hang: bool,
calls: Arc<AtomicUsize>,
}
impl StubValidator {
fn passing(typ: &'static str) -> Self {
Self {
typ,
outcome: None,
hang: false,
calls: Arc::new(AtomicUsize::new(0)),
}
}
}
#[async_trait]
impl ChallengeValidator for StubValidator {
fn typ(&self) -> &'static str {
self.typ
}
async fn validate(&self, _ctx: &ValidationContext<'_>) -> Result<(), ChallengeError> {
self.calls.fetch_add(1, Ordering::SeqCst);
if self.hang {
tokio::time::sleep(Duration::from_secs(3600)).await;
}
match self.outcome {
Some(detail) => Err(ChallengeError::IncorrectResponse(detail.to_string())),
None => Ok(()),
}
}
}
fn context<'a>(identifier: &'a str, key_authorization: &'a str) -> ValidationContext<'a> {
ValidationContext {
identifier,
wildcard: false,
token: "tok",
key_authorization,
challenge_id: "chall-1",
}
}
fn cfg(enabled: &[&str], bypass: bool) -> ChallengeConfig {
ChallengeConfig {
enabled: enabled
.iter()
.map(std::string::ToString::to_string)
.collect(),
bypass,
..ChallengeConfig::default()
}
}
#[test]
fn the_default_config_validates_and_offers_http_01_alone() {
let registry = from_config(
&ChallengeConfig::default(),
&DnsConfig::default(),
crate::testutil::no_proxies(),
)
.unwrap();
assert!(!registry.is_bypassed());
assert_eq!(registry.enabled_types(), [HTTP_01.to_string()]);
assert_eq!(registry.validators.len(), 1);
}
#[test]
fn bypassing_constructs_no_validators() {
let registry = from_config(
&cfg(&["http-01"], true),
&DnsConfig::default(),
crate::testutil::no_proxies(),
)
.unwrap();
assert!(registry.is_bypassed());
assert!(registry.validators.is_empty());
}
#[test]
fn the_test_default_bypasses_where_the_configured_default_does_not() {
let configured = from_config(
&ChallengeConfig::default(),
&DnsConfig::default(),
crate::testutil::no_proxies(),
)
.unwrap();
let direct = ChallengeRegistry::default();
assert!(!configured.is_bypassed());
assert!(direct.is_bypassed());
assert_eq!(configured.enabled_types(), direct.enabled_types());
assert_eq!(configured.timeout, direct.timeout);
}
#[test]
fn an_unknown_type_is_a_startup_error_even_when_bypassing() {
for bypass in [true, false] {
let error = from_config(
&cfg(&["http-01", "htttp-01"], bypass),
&DnsConfig::default(),
crate::testutil::no_proxies(),
)
.unwrap_err()
.to_string();
assert!(
error.contains("unknown challenge type") && error.contains("htttp-01"),
"{error}"
);
}
}
#[test]
fn an_empty_enabled_list_is_a_startup_error() {
let error = from_config(
&cfg(&[], true),
&DnsConfig::default(),
crate::testutil::no_proxies(),
)
.unwrap_err()
.to_string();
assert!(error.contains("challenge.enabled is empty"), "{error}");
}
#[test]
fn an_invalid_dns_resolver_is_a_startup_error_even_when_bypassing() {
let dns = DnsConfig {
resolver: Some("not-a-socket-address".to_string()),
};
for bypass in [true, false] {
let error = from_config(
&cfg(&["http-01"], bypass),
&dns,
crate::testutil::no_proxies(),
)
.unwrap_err()
.to_string();
assert!(error.contains("dns.resolver"), "{error}");
}
}
#[test]
fn a_valid_dns_resolver_is_used_to_build_the_real_validators() {
let dns = DnsConfig {
resolver: Some("127.0.0.1:5300".to_string()),
};
let registry = from_config(
&cfg(&["dns-01"], false),
&dns,
crate::testutil::no_proxies(),
)
.unwrap();
assert!(!registry.is_bypassed());
}
#[tokio::test]
async fn bypass_accepts_every_type_without_a_validator() {
let registry = ChallengeRegistry::default();
for typ in KNOWN_TYPES {
assert!(
registry
.validate(typ, &context("example.com", "tok.thumb"))
.await
.is_ok()
);
}
}
#[tokio::test]
async fn validate_dispatches_on_the_challenge_type() {
let http = StubValidator::passing(HTTP_01);
let dns = StubValidator {
outcome: Some("no matching TXT record"),
..StubValidator::passing(DNS_01)
};
let (http_calls, dns_calls) = (http.calls.clone(), dns.calls.clone());
let registry = ChallengeRegistry::new(
vec![Arc::new(http), Arc::new(dns)],
vec![HTTP_01.to_string(), DNS_01.to_string()],
false,
Duration::from_secs(5),
);
let ctx = context("example.com", "tok.thumb");
assert!(registry.validate(HTTP_01, &ctx).await.is_ok());
assert!(matches!(
registry.validate(DNS_01, &ctx).await,
Err(ChallengeError::IncorrectResponse(_))
));
assert_eq!(http_calls.load(Ordering::SeqCst), 1);
assert_eq!(dns_calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_type_without_a_validator_is_internal() {
let registry = ChallengeRegistry::new(
vec![Arc::new(StubValidator::passing(HTTP_01))],
vec![HTTP_01.to_string()],
false,
Duration::from_secs(5),
);
assert!(matches!(
registry
.validate(DNS_01, &context("example.com", "tok.thumb"))
.await,
Err(ChallengeError::Internal(detail)) if detail.contains("dns-01")
));
}
#[tokio::test]
async fn a_wedged_validator_times_out_rather_than_hanging() {
let registry = ChallengeRegistry::new(
vec![Arc::new(StubValidator {
hang: true,
..StubValidator::passing(HTTP_01)
})],
vec![HTTP_01.to_string()],
false,
Duration::from_millis(10),
);
assert!(matches!(
registry
.validate(HTTP_01, &context("example.com", "tok.thumb"))
.await,
Err(ChallengeError::Connection(detail)) if detail.contains("timed out")
));
}
#[test]
fn types_for_restricts_a_wildcard_to_dns_01() {
let all = ChallengeRegistry::new(
Vec::new(),
KNOWN_TYPES
.iter()
.map(std::string::ToString::to_string)
.collect(),
true,
Duration::from_secs(5),
);
assert_eq!(all.types_for(false), KNOWN_TYPES);
assert_eq!(all.types_for(true), [DNS_01]);
let without_dns = ChallengeRegistry::default();
assert_eq!(without_dns.types_for(false), [HTTP_01]);
assert!(without_dns.types_for(true).is_empty());
}
#[test]
fn challenge_error_labels_and_details() {
let errors = [
ChallengeError::Connection("a".into()),
ChallengeError::Dns("b".into()),
ChallengeError::IncorrectResponse("c".into()),
ChallengeError::Tls("d".into()),
ChallengeError::Unauthorized("e".into()),
ChallengeError::Internal("f".into()),
];
let kinds: Vec<_> = errors.iter().map(ChallengeError::kind).collect();
assert_eq!(
kinds,
[
"connection",
"dns",
"incorrectResponse",
"tls",
"unauthorized",
"internal"
]
);
let details: Vec<_> = errors.iter().map(ChallengeError::detail).collect();
assert_eq!(details, ["a", "b", "c", "d", "e", "f"]);
}
#[test]
fn debug_shows_the_policy_not_the_validators() {
let rendered = format!("{:?}", ChallengeRegistry::default());
assert!(rendered.contains("http-01"), "{rendered}");
assert!(rendered.contains("bypass: true"), "{rendered}");
}
}