use std::sync::Arc;
use async_trait::async_trait;
use ring::digest;
use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
use rustls::{ClientConfig, DigitallySignedStruct, SignatureScheme};
use subtle::ConstantTimeEq;
use tokio_rustls::TlsConnector;
use tracing::{debug, info};
use x509_parser::prelude::*;
use super::{ChallengeError, ChallengeValidator, TLS_ALPN_01, ValidationContext};
use crate::config::TlsAlpnConfig;
const ACME_TLS_ALPN: &[u8] = b"acme-tls/1";
const ACME_IDENTIFIER_OID: &[u64] = &[1, 3, 6, 1, 5, 5, 7, 1, 31];
#[async_trait]
pub trait TlsAlpnProbe: Send + Sync {
async fn peer_certificate(&self, identifier: &str, port: u16) -> Result<Vec<u8>, ProbeError>;
}
#[derive(Debug)]
pub enum ProbeError {
Connect(String),
Tls(String),
}
#[derive(Debug)]
struct AcceptAnyServerCert {
schemes: Vec<SignatureScheme>,
}
impl ServerCertVerifier for AcceptAnyServerCert {
fn verify_server_cert(
&self,
_end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp_response: &[u8],
_now: UnixTime,
) -> Result<ServerCertVerified, rustls::Error> {
Ok(ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, rustls::Error> {
Ok(HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, rustls::Error> {
Ok(HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
self.schemes.clone()
}
}
pub fn accept_any_client_config(alpn: &[&[u8]]) -> anyhow::Result<Arc<ClientConfig>> {
let provider = rustls::crypto::ring::default_provider();
let schemes = provider
.signature_verification_algorithms
.supported_schemes();
let mut config = ClientConfig::builder_with_provider(Arc::new(provider))
.with_safe_default_protocol_versions()
.map_err(|error| anyhow::anyhow!("building the TLS client configuration: {error}"))?
.dangerous()
.with_custom_certificate_verifier(Arc::new(AcceptAnyServerCert { schemes }))
.with_no_client_auth();
config.alpn_protocols = alpn.iter().map(|protocol| protocol.to_vec()).collect();
Ok(Arc::new(config))
}
pub struct RustlsProbe {
config: Arc<ClientConfig>,
outbound: crate::http_client::Outbound,
}
impl RustlsProbe {
pub fn new(outbound: crate::http_client::Outbound) -> anyhow::Result<Self> {
Ok(Self {
config: accept_any_client_config(&[ACME_TLS_ALPN])?,
outbound,
})
}
}
#[async_trait]
impl TlsAlpnProbe for RustlsProbe {
async fn peer_certificate(&self, identifier: &str, port: u16) -> Result<Vec<u8>, ProbeError> {
let server_name = ServerName::try_from(identifier.to_string()).map_err(|error| {
ProbeError::Connect(format!("{identifier} is not a valid SNI name: {error}"))
})?;
let endpoint = crate::http_client::Endpoint::tls(identifier, port);
let stream = self
.outbound
.connect_stream(&endpoint)
.await
.map_err(|error| {
ProbeError::Connect(format!("connecting to {identifier}:{port}: {error}"))
})?;
let stream = TlsConnector::from(self.config.clone())
.connect(server_name, stream)
.await
.map_err(|error| {
ProbeError::Tls(format!("TLS handshake with {identifier}:{port}: {error}"))
})?;
let (_, connection) = stream.get_ref();
match connection.alpn_protocol() {
Some(ACME_TLS_ALPN) => {}
Some(other) => {
return Err(ProbeError::Tls(format!(
"{identifier}:{port} negotiated ALPN {:?} instead of acme-tls/1",
String::from_utf8_lossy(other)
)));
}
None => {
return Err(ProbeError::Tls(format!(
"{identifier}:{port} did not negotiate the acme-tls/1 ALPN protocol"
)));
}
}
let leaf = connection
.peer_certificates()
.and_then(<[CertificateDer<'_>]>::first)
.ok_or_else(|| {
ProbeError::Tls(format!("{identifier}:{port} presented no certificate"))
})?;
Ok(leaf.as_ref().to_vec())
}
}
pub struct TlsAlpn01Validator {
probe: Arc<dyn TlsAlpnProbe>,
port: u16,
}
impl std::fmt::Debug for TlsAlpn01Validator {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("TlsAlpn01Validator")
.field("port", &self.port)
.finish_non_exhaustive()
}
}
impl TlsAlpn01Validator {
pub fn from_config(
cfg: &TlsAlpnConfig,
outbound: crate::http_client::Outbound,
) -> anyhow::Result<Self> {
let probe = Arc::new(
RustlsProbe::new(outbound)
.map_err(|error| anyhow::anyhow!("challenge.tls_alpn_01: {error}"))?,
);
info!(
event = "challenge_tls_alpn_01_loaded",
outcome = "success",
port = cfg.port
);
Ok(Self::with_probe(cfg, probe))
}
pub fn with_probe(cfg: &TlsAlpnConfig, probe: Arc<dyn TlsAlpnProbe>) -> Self {
Self {
probe,
port: cfg.port,
}
}
}
#[async_trait]
impl ChallengeValidator for TlsAlpn01Validator {
fn typ(&self) -> &'static str {
TLS_ALPN_01
}
async fn validate(&self, ctx: &ValidationContext<'_>) -> Result<(), ChallengeError> {
let leaf = self
.probe
.peer_certificate(ctx.identifier, self.port)
.await
.map_err(|error| match error {
ProbeError::Connect(detail) => ChallengeError::Connection(detail),
ProbeError::Tls(detail) => ChallengeError::Tls(detail),
})?;
verify_acme_identifier(&leaf, ctx.identifier, ctx.key_authorization).inspect(|()| {
debug!(
event = "challenge_tls_alpn_01_matched",
outcome = "success",
identifier = ctx.identifier,
challenge_id = ctx.challenge_id,
);
})
}
}
pub(crate) fn verify_acme_identifier(
leaf_der: &[u8],
identifier: &str,
key_authorization: &str,
) -> Result<(), ChallengeError> {
let (_, certificate) = X509Certificate::from_der(leaf_der).map_err(|error| {
ChallengeError::Tls(format!(
"responder certificate could not be parsed: {error}"
))
})?;
let sans = certificate
.subject_alternative_name()
.ok()
.flatten()
.map(|extension| extension.value.general_names.as_slice())
.unwrap_or_default();
match sans {
[GeneralName::DNSName(name)] if name.eq_ignore_ascii_case(identifier) => {}
[] => {
return Err(ChallengeError::IncorrectResponse(
"responder certificate has no subject alternative name".to_string(),
));
}
[GeneralName::DNSName(name)] => {
return Err(ChallengeError::IncorrectResponse(format!(
"responder certificate is for {name}, not {identifier}"
)));
}
other => {
return Err(ChallengeError::IncorrectResponse(format!(
"responder certificate must carry exactly one dNSName, found {}",
other.len()
)));
}
}
let oid = der_parser::oid::Oid::from(ACME_IDENTIFIER_OID)
.map_err(|error| ChallengeError::Internal(format!("acmeIdentifier OID: {error:?}")))?;
let extension = certificate
.get_extension_unique(&oid)
.map_err(|error| {
ChallengeError::IncorrectResponse(format!(
"responder certificate has a malformed acmeIdentifier extension: {error}"
))
})?
.ok_or_else(|| {
ChallengeError::IncorrectResponse(
"responder certificate has no acmeIdentifier extension".to_string(),
)
})?;
if !extension.critical {
return Err(ChallengeError::IncorrectResponse(
"the acmeIdentifier extension must be critical".to_string(),
));
}
let payload = match extension.value {
[0x04, 0x20, rest @ ..] if rest.len() == 32 => rest,
_ => {
return Err(ChallengeError::IncorrectResponse(
"the acmeIdentifier extension is not a 32-octet OCTET STRING".to_string(),
));
}
};
let expected = digest::digest(&digest::SHA256, key_authorization.as_bytes());
if payload.ct_eq(expected.as_ref()).into() {
Ok(())
} else {
Err(ChallengeError::IncorrectResponse(
"the acmeIdentifier extension does not match the key authorization".to_string(),
))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dns::Resolver;
use rcgen::{CertificateParams, CustomExtension, KeyPair, SanType};
const KEY_AUTH: &str = "token-value.thumbprint-value";
fn responder_cert(sans: Vec<SanType>, digest_of: &str, critical: bool) -> Vec<u8> {
let key_pair = KeyPair::generate().unwrap();
let mut params = CertificateParams::default();
params.subject_alt_names = sans;
let hash = digest::digest(&digest::SHA256, digest_of.as_bytes());
let mut value = vec![0x04, 0x20];
value.extend_from_slice(hash.as_ref());
let mut extension = CustomExtension::from_oid_content(ACME_IDENTIFIER_OID, value);
extension.set_criticality(critical);
params.custom_extensions = vec![extension];
params
.self_signed(&key_pair)
.unwrap()
.der()
.as_ref()
.to_vec()
}
fn dns_san(name: &str) -> SanType {
SanType::DnsName(name.to_string().try_into().unwrap())
}
fn valid_cert() -> Vec<u8> {
responder_cert(vec![dns_san("example.com")], KEY_AUTH, true)
}
struct StubProbe {
outcome: Result<Vec<u8>, &'static str>,
tls_error: bool,
}
impl StubProbe {
fn serving(der: Vec<u8>) -> Self {
Self {
outcome: Ok(der),
tls_error: false,
}
}
}
#[async_trait]
impl TlsAlpnProbe for StubProbe {
async fn peer_certificate(
&self,
_identifier: &str,
_port: u16,
) -> Result<Vec<u8>, ProbeError> {
match &self.outcome {
Ok(der) => Ok(der.clone()),
Err(detail) if self.tls_error => Err(ProbeError::Tls(detail.to_string())),
Err(detail) => Err(ProbeError::Connect(detail.to_string())),
}
}
}
fn validator(probe: StubProbe) -> TlsAlpn01Validator {
TlsAlpn01Validator::with_probe(&TlsAlpnConfig::default(), Arc::new(probe))
}
fn context(identifier: &str) -> ValidationContext<'_> {
ValidationContext {
identifier,
wildcard: false,
token: "token-value",
key_authorization: KEY_AUTH,
challenge_id: "chall-1",
}
}
#[tokio::test]
async fn a_conforming_responder_certificate_passes() {
assert!(
validator(StubProbe::serving(valid_cert()))
.validate(&context("example.com"))
.await
.is_ok()
);
}
#[test]
fn the_digest_must_be_of_this_challenge_s_key_authorization() {
let der = responder_cert(vec![dns_san("example.com")], "some.other-key-auth", true);
assert!(matches!(
verify_acme_identifier(&der, "example.com", KEY_AUTH),
Err(ChallengeError::IncorrectResponse(detail))
if detail.contains("does not match the key authorization")
));
}
#[test]
fn a_non_critical_extension_is_refused() {
let der = responder_cert(vec![dns_san("example.com")], KEY_AUTH, false);
assert!(matches!(
verify_acme_identifier(&der, "example.com", KEY_AUTH),
Err(ChallengeError::IncorrectResponse(detail)) if detail.contains("must be critical")
));
}
#[test]
fn a_certificate_without_the_extension_is_refused() {
let key_pair = KeyPair::generate().unwrap();
let mut params = CertificateParams::default();
params.subject_alt_names = vec![dns_san("example.com")];
let der = params
.self_signed(&key_pair)
.unwrap()
.der()
.as_ref()
.to_vec();
assert!(matches!(
verify_acme_identifier(&der, "example.com", KEY_AUTH),
Err(ChallengeError::IncorrectResponse(detail))
if detail.contains("no acmeIdentifier extension")
));
}
#[test]
fn the_certificate_must_name_the_identifier_being_validated() {
let der = responder_cert(vec![dns_san("other.example")], KEY_AUTH, true);
assert!(matches!(
verify_acme_identifier(&der, "example.com", KEY_AUTH),
Err(ChallengeError::IncorrectResponse(detail))
if detail.contains("other.example") && detail.contains("example.com")
));
}
#[test]
fn the_dns_name_comparison_ignores_case() {
let der = responder_cert(vec![dns_san("EXAMPLE.com")], KEY_AUTH, true);
assert!(verify_acme_identifier(&der, "example.com", KEY_AUTH).is_ok());
}
#[test]
fn extra_subject_alternative_names_are_refused() {
let der = responder_cert(
vec![dns_san("example.com"), dns_san("victim.example")],
KEY_AUTH,
true,
);
assert!(matches!(
verify_acme_identifier(&der, "example.com", KEY_AUTH),
Err(ChallengeError::IncorrectResponse(detail)) if detail.contains("exactly one dNSName")
));
let with_ip = responder_cert(
vec![
dns_san("example.com"),
SanType::IpAddress("10.0.0.1".parse().unwrap()),
],
KEY_AUTH,
true,
);
assert!(matches!(
verify_acme_identifier(&with_ip, "example.com", KEY_AUTH),
Err(ChallengeError::IncorrectResponse(_))
));
}
#[test]
fn a_certificate_with_no_subject_alternative_name_is_refused() {
let key_pair = KeyPair::generate().unwrap();
let der = CertificateParams::default()
.self_signed(&key_pair)
.unwrap()
.der()
.as_ref()
.to_vec();
assert!(matches!(
verify_acme_identifier(&der, "example.com", KEY_AUTH),
Err(ChallengeError::IncorrectResponse(detail))
if detail.contains("no subject alternative name")
));
}
#[test]
fn an_unparsable_certificate_is_a_tls_error() {
assert!(matches!(
verify_acme_identifier(&[0xde, 0xad, 0xbe, 0xef], "example.com", KEY_AUTH),
Err(ChallengeError::Tls(_))
));
}
#[tokio::test]
async fn probe_failures_keep_their_kind() {
let refused = validator(StubProbe {
outcome: Err("connection refused"),
tls_error: false,
});
assert!(matches!(
refused.validate(&context("example.com")).await,
Err(ChallengeError::Connection(_))
));
let alerted = validator(StubProbe {
outcome: Err("handshake failure"),
tls_error: true,
});
assert!(matches!(
alerted.validate(&context("example.com")).await,
Err(ChallengeError::Tls(_))
));
}
#[test]
fn reports_its_challenge_type() {
assert_eq!(
validator(StubProbe::serving(valid_cert())).typ(),
"tls-alpn-01"
);
}
#[test]
fn the_client_config_advertises_alpn_without_a_global_provider() {
let config = accept_any_client_config(&[ACME_TLS_ALPN]).unwrap();
assert_eq!(config.alpn_protocols, vec![ACME_TLS_ALPN.to_vec()]);
let plain = accept_any_client_config(&[]).unwrap();
assert!(plain.alpn_protocols.is_empty());
}
mod loopback {
use super::*;
use rustls::ServerConfig;
use rustls::pki_types::PrivateKeyDer;
use rustls::server::{ClientHello, ResolvesServerCert};
use rustls::sign::CertifiedKey;
use std::net::SocketAddr;
use tokio::net::TcpListener;
use tokio_rustls::TlsAcceptor;
struct UnreachableResolver;
#[async_trait]
impl Resolver for UnreachableResolver {
async fn reverse(&self, _ip: std::net::IpAddr) -> Result<Vec<String>, String> {
unreachable!()
}
async fn forward(&self, _name: &str) -> Result<Vec<std::net::IpAddr>, String> {
unreachable!("a literal 127.0.0.1 must short-circuit before this is called")
}
async fn txt(&self, _name: &str) -> Result<Vec<String>, String> {
unreachable!()
}
}
fn probe() -> RustlsProbe {
RustlsProbe::new(crate::testutil::outbound_with(Arc::new(
UnreachableResolver,
)))
.unwrap()
}
#[derive(Debug)]
struct FixedCert(Arc<CertifiedKey>);
impl ResolvesServerCert for FixedCert {
fn resolve(&self, _hello: ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
Some(self.0.clone())
}
}
async fn serve_once(
der: Vec<u8>,
key: PrivateKeyDer<'static>,
alpn: &[&[u8]],
) -> SocketAddr {
let provider = rustls::crypto::ring::default_provider();
let signing_key = provider.key_provider.load_private_key(key).unwrap();
let certified = CertifiedKey::new(vec![CertificateDer::from(der)], signing_key);
let mut config = ServerConfig::builder_with_provider(Arc::new(provider))
.with_safe_default_protocol_versions()
.unwrap()
.with_no_client_auth()
.with_cert_resolver(Arc::new(FixedCert(Arc::new(certified))));
config.alpn_protocols = alpn.iter().map(|p| p.to_vec()).collect();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let acceptor = TlsAcceptor::from(Arc::new(config));
tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let _ = acceptor.accept(stream).await;
});
address
}
fn responder(identifier: &str, key_auth: &str) -> (Vec<u8>, PrivateKeyDer<'static>) {
let key_pair = KeyPair::generate().unwrap();
let key = PrivateKeyDer::try_from(key_pair.serialize_der()).unwrap();
let mut params = CertificateParams::default();
params.subject_alt_names = vec![dns_san(identifier)];
let hash = digest::digest(&digest::SHA256, key_auth.as_bytes());
let mut value = vec![0x04, 0x20];
value.extend_from_slice(hash.as_ref());
let mut extension = CustomExtension::from_oid_content(ACME_IDENTIFIER_OID, value);
extension.set_criticality(true);
params.custom_extensions = vec![extension];
let der = params
.self_signed(&key_pair)
.unwrap()
.der()
.as_ref()
.to_vec();
(der, key)
}
#[tokio::test]
async fn a_conforming_responder_is_probed_end_to_end() {
let (der, key) = responder("example.com", KEY_AUTH);
let address = serve_once(der.clone(), key, &[ACME_TLS_ALPN]).await;
let leaf = probe()
.peer_certificate("127.0.0.1", address.port())
.await
.expect("the critical acmeIdentifier extension must not break the handshake");
assert_eq!(leaf, der);
assert!(verify_acme_identifier(&leaf, "example.com", KEY_AUTH).is_ok());
}
#[tokio::test]
async fn a_server_without_the_alpn_protocol_is_a_tls_error() {
let (der, key) = responder("example.com", KEY_AUTH);
let address = serve_once(der, key, &[]).await;
let error = probe()
.peer_certificate("127.0.0.1", address.port())
.await
.unwrap_err();
assert!(matches!(error, ProbeError::Tls(_)), "{error:?}");
}
#[tokio::test]
async fn a_closed_port_is_a_connect_error() {
let port = {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
listener.local_addr().unwrap().port()
};
let error = probe()
.peer_certificate("127.0.0.1", port)
.await
.unwrap_err();
assert!(matches!(error, ProbeError::Connect(_)), "{error:?}");
}
#[tokio::test]
async fn an_identifier_that_is_not_a_valid_sni_name_is_refused() {
let error = probe()
.peer_certificate("not a hostname", 443)
.await
.unwrap_err();
match error {
ProbeError::Connect(detail) => {
assert!(detail.contains("not a valid SNI name"), "{detail}")
}
other => panic!("expected a Connect error, got {other:?}"),
}
}
#[tokio::test]
async fn from_config_builds_a_real_probe() {
let validator = TlsAlpn01Validator::from_config(
&TlsAlpnConfig { port: 8443 },
crate::testutil::outbound_with(Arc::new(UnreachableResolver)),
)
.expect("the real probe must build");
let rendered = format!("{validator:?}");
assert!(rendered.contains("TlsAlpn01Validator"), "{rendered}");
assert!(rendered.contains("8443"), "{rendered}");
}
}
#[test]
fn an_extension_that_is_not_a_32_octet_octet_string_is_refused() {
let key_pair = KeyPair::generate().unwrap();
let mut params = CertificateParams::default();
params.subject_alt_names = vec![dns_san("example.com")];
let mut value = vec![0x04, 0x10];
value.extend_from_slice(&[0xab; 16]);
let mut extension = CustomExtension::from_oid_content(ACME_IDENTIFIER_OID, value);
extension.set_criticality(true);
params.custom_extensions = vec![extension];
let der = params
.self_signed(&key_pair)
.unwrap()
.der()
.as_ref()
.to_vec();
match verify_acme_identifier(&der, "example.com", KEY_AUTH) {
Err(ChallengeError::IncorrectResponse(detail)) => {
assert!(detail.contains("32-octet OCTET STRING"), "{detail}")
}
other => panic!("expected an IncorrectResponse, got {other:?}"),
}
}
}