use std::sync::Arc;
use std::time::Duration;
use base64::prelude::*;
use http_body_util::{BodyExt, Full, Limited};
use hyper::body::Bytes;
use hyper::header::{HeaderValue, LOCATION, RETRY_AFTER};
use hyper::{Method, Request, StatusCode};
use ring::rand::SystemRandom;
use ring::signature::{ECDSA_P256_SHA256_FIXED_SIGNING, EcdsaKeyPair, KeyPair};
use serde::Deserialize;
use serde_json::{Value, json};
use tracing::debug;
use url::Url;
use crate::http_client::{MAX_RESPONSE_BYTES, error_excerpt};
#[derive(Debug, thiserror::Error)]
pub enum UpstreamError {
#[error("upstream URL invalid: {0}")]
Url(String),
#[error("upstream transport failed: {0}")]
Transport(String),
#[error("upstream protocol error: {0}")]
Protocol(String),
#[error("upstream returned {status} {typ}: {detail}")]
Problem {
status: u16,
typ: String,
detail: String,
},
#[error("outbound JWS signing failed: {0}")]
Jws(String),
}
impl UpstreamError {
pub fn is_bad_csr(&self) -> bool {
matches!(self, UpstreamError::Problem { typ, .. } if typ.ends_with(":badCSR"))
}
pub fn is_already_revoked(&self) -> bool {
matches!(self, UpstreamError::Problem { typ, .. } if typ.ends_with(":alreadyRevoked"))
}
pub fn is_external_account_required(&self) -> bool {
matches!(self, UpstreamError::Problem { typ, .. } if typ.ends_with(":externalAccountRequired"))
}
fn is_bad_nonce(&self) -> bool {
matches!(self, UpstreamError::Problem { typ, .. } if typ.ends_with(":badNonce"))
}
}
#[derive(Debug, Clone, Deserialize)]
pub struct Directory {
#[serde(rename = "newNonce")]
pub new_nonce: String,
#[serde(rename = "newAccount")]
pub new_account: String,
#[serde(rename = "newOrder")]
pub new_order: String,
#[serde(rename = "revokeCert")]
pub revoke_cert: Option<String>,
#[serde(rename = "renewalInfo")]
pub renewal_info: Option<String>,
}
#[derive(Debug)]
pub struct AcmeResponse {
pub status: StatusCode,
pub body: Bytes,
pub location: Option<String>,
pub retry_after: Option<u64>,
pub nonce: Option<String>,
}
impl AcmeResponse {
pub fn json<T: serde::de::DeserializeOwned>(&self) -> Result<T, UpstreamError> {
serde_json::from_slice(&self.body).map_err(|error| {
UpstreamError::Protocol(format!("response was not the expected JSON: {error}"))
})
}
pub fn text(&self) -> Result<String, UpstreamError> {
String::from_utf8(self.body.to_vec())
.map_err(|_| UpstreamError::Protocol("response body was not UTF-8".to_string()))
}
}
pub struct AccountKey {
pair: EcdsaKeyPair,
rng: SystemRandom,
spki_der: Vec<u8>,
}
impl AccountKey {
pub fn from_pkcs8(pkcs8: &[u8]) -> Result<Self, UpstreamError> {
let rng = SystemRandom::new();
let pair = EcdsaKeyPair::from_pkcs8(&ECDSA_P256_SHA256_FIXED_SIGNING, pkcs8, &rng)
.map_err(|error| UpstreamError::Jws(format!("account key unusable: {error}")))?;
let spki_der = spki_from_p256_public(pair.public_key().as_ref())?;
Ok(Self {
pair,
rng,
spki_der,
})
}
pub fn jwk(&self) -> Value {
let point = self.pair.public_key().as_ref();
json!({
"crv": "P-256",
"kty": "EC",
"x": BASE64_URL_SAFE_NO_PAD.encode(&point[1..33]),
"y": BASE64_URL_SAFE_NO_PAD.encode(&point[33..65]),
})
}
pub fn spki_der(&self) -> &[u8] {
&self.spki_der
}
fn sign(&self, input: &[u8]) -> Result<Vec<u8>, UpstreamError> {
self.pair
.sign(&self.rng, input)
.map(|sig| sig.as_ref().to_vec())
.map_err(|error| UpstreamError::Jws(format!("signing failed: {error}")))
}
}
fn spki_from_p256_public(point: &[u8]) -> Result<Vec<u8>, UpstreamError> {
if point.len() != 65 || point[0] != 0x04 {
return Err(UpstreamError::Jws(
"P-256 public key was not an uncompressed 65-byte point".to_string(),
));
}
const PREFIX: &[u8] = &[
0x30, 0x59, 0x30, 0x13, 0x06, 0x07, 0x2a, 0x86, 0x48, 0xce, 0x3d, 0x02, 0x01, 0x06, 0x08, 0x2a, 0x86, 0x48, 0xce, 0x3d, 0x03, 0x01,
0x07, 0x03, 0x42, 0x00, ];
let mut der = Vec::with_capacity(PREFIX.len() + point.len());
der.extend_from_slice(PREFIX);
der.extend_from_slice(point);
Ok(der)
}
pub enum Signer<'a> {
Jwk,
Kid(&'a str),
}
impl std::fmt::Debug for AcmeClient {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("AcmeClient")
.field("directory", &self.directory)
.field("timeout", &self.timeout)
.finish()
}
}
pub struct AcmeClient {
directory: Directory,
tls: Arc<rustls::ClientConfig>,
outbound: crate::http_client::Outbound,
timeout: Duration,
}
impl AcmeClient {
pub async fn discover(
directory_url: &str,
outbound: crate::http_client::Outbound,
timeout: Duration,
) -> Result<Self, UpstreamError> {
let tls = Arc::new(crate::http_client::webpki_tls_config());
let url = Url::parse(directory_url)
.map_err(|error| UpstreamError::Url(format!("{directory_url}: {error}")))?;
let response = request(&tls, &outbound, Method::GET, &url, None, timeout).await?;
if !response.status.is_success() {
return Err(problem_from(&response));
}
let directory: Directory = response.json()?;
debug!(event = "upstream_directory_discovered", outcome = "success", upstream_url = %directory_url);
Ok(Self {
directory,
outbound,
tls,
timeout,
})
}
pub fn directory(&self) -> &Directory {
&self.directory
}
async fn nonce(&self) -> Result<String, UpstreamError> {
let url = self.parse(&self.directory.new_nonce)?;
let response = request(
&self.tls,
&self.outbound,
Method::HEAD,
&url,
None,
self.timeout,
)
.await?;
response.nonce.ok_or_else(|| {
UpstreamError::Protocol("newNonce response carried no Replay-Nonce".to_string())
})
}
fn parse(&self, url: &str) -> Result<Url, UpstreamError> {
Url::parse(url).map_err(|error| UpstreamError::Url(format!("{url}: {error}")))
}
pub async fn post(
&self,
key: &AccountKey,
signer: &Signer<'_>,
url: &str,
payload: Option<&Value>,
) -> Result<AcmeResponse, UpstreamError> {
match self.post_once(key, signer, url, payload).await {
Err(error) if error.is_bad_nonce() => {
debug!(event = "upstream_bad_nonce_retry", outcome = "progress", upstream_url = %url);
self.post_once(key, signer, url, payload).await
}
other => other,
}
}
async fn post_once(
&self,
key: &AccountKey,
signer: &Signer<'_>,
url: &str,
payload: Option<&Value>,
) -> Result<AcmeResponse, UpstreamError> {
let nonce = self.nonce().await?;
let parsed = self.parse(url)?;
let protected = match signer {
Signer::Jwk => json!({
"alg": "ES256", "jwk": key.jwk(), "nonce": nonce, "url": url,
}),
Signer::Kid(kid) => json!({
"alg": "ES256", "kid": kid, "nonce": nonce, "url": url,
}),
};
let protected_b64 = BASE64_URL_SAFE_NO_PAD.encode(
serde_json::to_vec(&protected)
.map_err(|error| UpstreamError::Jws(error.to_string()))?,
);
let payload_b64 = match payload {
Some(value) => BASE64_URL_SAFE_NO_PAD.encode(
serde_json::to_vec(value).map_err(|error| UpstreamError::Jws(error.to_string()))?,
),
None => String::new(),
};
let signature = key.sign(format!("{protected_b64}.{payload_b64}").as_bytes())?;
let body = json!({
"protected": protected_b64,
"payload": payload_b64,
"signature": BASE64_URL_SAFE_NO_PAD.encode(signature),
});
let body =
serde_json::to_vec(&body).map_err(|error| UpstreamError::Jws(error.to_string()))?;
let response = request(
&self.tls,
&self.outbound,
Method::POST,
&parsed,
Some(Bytes::from(body)),
self.timeout,
)
.await?;
if response.status.is_success() {
Ok(response)
} else {
Err(problem_from(&response))
}
}
pub async fn get(
&self,
key: &AccountKey,
kid: &str,
url: &str,
) -> Result<AcmeResponse, UpstreamError> {
self.post(key, &Signer::Kid(kid), url, None).await
}
pub async fn get_unsigned(&self, url: &str) -> Result<AcmeResponse, UpstreamError> {
let parsed = self.parse(url)?;
let response = request(
&self.tls,
&self.outbound,
Method::GET,
&parsed,
None,
self.timeout,
)
.await?;
if response.status.is_success() {
Ok(response)
} else {
Err(problem_from(&response))
}
}
}
fn problem_from(response: &AcmeResponse) -> UpstreamError {
#[derive(Deserialize)]
struct ProblemDoc {
#[serde(rename = "type")]
typ: Option<String>,
detail: Option<String>,
}
let status = response.status.as_u16();
match serde_json::from_slice::<ProblemDoc>(&response.body) {
Ok(doc) => UpstreamError::Problem {
status,
typ: doc.typ.unwrap_or_else(|| "about:blank".to_string()),
detail: doc.detail.unwrap_or_default(),
},
Err(_) => UpstreamError::Problem {
status,
typ: "about:blank".to_string(),
detail: error_excerpt(&response.body),
},
}
}
async fn request(
tls: &Arc<rustls::ClientConfig>,
outbound: &crate::http_client::Outbound,
method: Method,
url: &Url,
body: Option<Bytes>,
timeout: Duration,
) -> Result<AcmeResponse, UpstreamError> {
tokio::time::timeout(timeout, request_inner(tls, outbound, method, url, body))
.await
.map_err(|_| UpstreamError::Transport(format!("timed out after {timeout:?}")))?
}
async fn request_inner(
tls: &Arc<rustls::ClientConfig>,
outbound: &crate::http_client::Outbound,
method: Method,
url: &Url,
body: Option<Bytes>,
) -> Result<AcmeResponse, UpstreamError> {
let endpoint = crate::http_client::Endpoint::from_url(url).map_err(UpstreamError::Url)?;
let connection = outbound
.connect(&endpoint, tls)
.await
.map_err(UpstreamError::Transport)?;
let has_body = body.is_some();
let request = Request::builder()
.method(method)
.uri(connection.request_target(url))
.header(hyper::header::HOST, endpoint.authority())
.header(hyper::header::USER_AGENT, "acme-proxy")
.header(
hyper::header::CONTENT_TYPE,
if has_body {
"application/jose+json"
} else {
"application/json"
},
)
.body(Full::new(body.unwrap_or_default()))
.map_err(|error| UpstreamError::Transport(error.to_string()))?;
send(connection, request).await
}
async fn send(
mut connection: crate::http_client::Connection<Full<Bytes>>,
request: Request<Full<Bytes>>,
) -> Result<AcmeResponse, UpstreamError> {
let response = connection
.send_request(request)
.await
.map_err(|error| UpstreamError::Transport(error.to_string()))?;
let status = response.status();
let location = header_string(response.headers().get(LOCATION));
let nonce = header_string(response.headers().get("replay-nonce"));
let retry_after = header_string(response.headers().get(RETRY_AFTER))
.and_then(|value| value.trim().parse::<u64>().ok());
let body = Limited::new(response.into_body(), MAX_RESPONSE_BYTES)
.collect()
.await
.map_err(|error| UpstreamError::Transport(format!("response body: {error}")))?
.to_bytes();
Ok(AcmeResponse {
status,
body,
location,
retry_after,
nonce,
})
}
fn header_string(value: Option<&HeaderValue>) -> Option<String> {
value
.and_then(|value| value.to_str().ok())
.map(str::to_string)
}
#[cfg(test)]
mod tests {
use super::*;
fn test_resolver() -> std::sync::Arc<dyn crate::dns::Resolver> {
std::sync::Arc::new(crate::dns::HickoryResolver::from_system_uncached().unwrap())
}
use crate::signer::relay::testsrv::{self, Script};
fn pkcs8() -> Vec<u8> {
rcgen::KeyPair::generate_for(&rcgen::PKCS_ECDSA_P256_SHA256)
.unwrap()
.serialize_der()
}
fn key() -> AccountKey {
AccountKey::from_pkcs8(&pkcs8()).unwrap()
}
const TIMEOUT: Duration = Duration::from_secs(5);
#[test]
fn an_account_key_exposes_a_jwk_and_a_matching_spki() {
let key = key();
let jwk = key.jwk();
assert_eq!(jwk["kty"], "EC");
assert_eq!(jwk["crv"], "P-256");
let thumbprint = crate::extractors::acme::jwk_thumbprint(key.spki_der()).unwrap();
let canonical = format!(
r#"{{"crv":"P-256","kty":"EC","x":"{}","y":"{}"}}"#,
jwk["x"].as_str().unwrap(),
jwk["y"].as_str().unwrap()
);
let expected = BASE64_URL_SAFE_NO_PAD
.encode(ring::digest::digest(&ring::digest::SHA256, canonical.as_bytes()).as_ref());
assert_eq!(thumbprint, expected);
}
#[test]
fn a_non_p256_key_is_refused() {
let ed25519 = rcgen::KeyPair::generate_for(&rcgen::PKCS_ED25519)
.unwrap()
.serialize_der();
assert!(matches!(
AccountKey::from_pkcs8(&ed25519),
Err(UpstreamError::Jws(_))
));
}
#[test]
fn spki_encoding_rejects_a_malformed_point() {
assert!(spki_from_p256_public(&[0x04, 0x01]).is_err());
assert!(spki_from_p256_public(&[0x02; 65]).is_err());
}
#[test]
fn upstream_errors_classify_the_three_cases_the_caller_branches_on() {
let bad_csr = UpstreamError::Problem {
status: 403,
typ: "urn:ietf:params:acme:error:badCSR".to_string(),
detail: String::new(),
};
assert!(bad_csr.is_bad_csr());
assert!(!bad_csr.is_already_revoked());
assert!(!bad_csr.is_bad_nonce());
let revoked = UpstreamError::Problem {
status: 400,
typ: "urn:ietf:params:acme:error:alreadyRevoked".to_string(),
detail: String::new(),
};
assert!(revoked.is_already_revoked());
assert!(!revoked.is_bad_csr());
let nonce = UpstreamError::Problem {
status: 400,
typ: "urn:ietf:params:acme:error:badNonce".to_string(),
detail: String::new(),
};
assert!(nonce.is_bad_nonce());
let transport = UpstreamError::Transport("boom".to_string());
assert!(!transport.is_bad_csr() && !transport.is_already_revoked());
}
#[test]
fn errors_display_their_detail() {
assert!(UpstreamError::Url("bad".into()).to_string().contains("bad"));
assert!(
UpstreamError::Transport("refused".into())
.to_string()
.contains("refused")
);
assert!(
UpstreamError::Protocol("garbage".into())
.to_string()
.contains("garbage")
);
assert!(
UpstreamError::Jws("nope".into())
.to_string()
.contains("nope")
);
let rendered = UpstreamError::Problem {
status: 429,
typ: "urn:ietf:params:acme:error:rateLimited".to_string(),
detail: "slow down".to_string(),
}
.to_string();
assert!(
rendered.contains("429") && rendered.contains("slow down"),
"{rendered}"
);
}
#[tokio::test]
async fn discover_reads_the_directory() {
let upstream = testsrv::start(Script::default()).await;
let client = AcmeClient::discover(
&upstream.directory_url(),
crate::testutil::outbound_with(test_resolver()),
TIMEOUT,
)
.await
.unwrap();
assert_eq!(
client.directory().new_order,
format!("{}/newOrder", upstream.base)
);
assert!(client.directory().revoke_cert.is_some());
assert!(client.directory().renewal_info.is_some());
}
#[tokio::test]
async fn discover_fails_on_an_unusable_url() {
assert!(matches!(
AcmeClient::discover(
"not a url",
crate::testutil::outbound_with(test_resolver()),
TIMEOUT
)
.await,
Err(UpstreamError::Url(_))
));
assert!(matches!(
AcmeClient::discover(
"ftp://example.invalid/dir",
crate::testutil::outbound_with(test_resolver()),
TIMEOUT
)
.await,
Err(UpstreamError::Url(_))
));
}
#[tokio::test]
async fn discover_fails_when_nothing_is_listening() {
let port = {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
listener.local_addr().unwrap().port()
};
let error = AcmeClient::discover(
&format!("http://127.0.0.1:{port}/directory"),
crate::testutil::outbound_with(test_resolver()),
TIMEOUT,
)
.await
.unwrap_err();
assert!(matches!(error, UpstreamError::Transport(_)), "{error:?}");
}
#[tokio::test]
async fn every_signed_post_fetches_a_fresh_nonce() {
let upstream = testsrv::start(Script::default()).await;
let client = AcmeClient::discover(
&upstream.directory_url(),
crate::testutil::outbound_with(test_resolver()),
TIMEOUT,
)
.await
.unwrap();
let key = key();
for _ in 0..3 {
client
.post(
&key,
&Signer::Jwk,
&client.directory().new_account.clone(),
Some(&json!({})),
)
.await
.unwrap();
}
assert_eq!(upstream.nonce_fetches(), 3);
}
#[tokio::test]
async fn a_bad_nonce_is_retried_once() {
let upstream = testsrv::start(Script {
bad_nonce_once: true,
..Script::default()
})
.await;
let client = AcmeClient::discover(
&upstream.directory_url(),
crate::testutil::outbound_with(test_resolver()),
TIMEOUT,
)
.await
.unwrap();
let response = client
.post(
&key(),
&Signer::Jwk,
&client.directory().new_account.clone(),
Some(&json!({})),
)
.await
.expect("the retry must absorb a single badNonce");
assert_eq!(response.status, 201);
assert_eq!(upstream.nonce_fetches(), 2);
}
#[tokio::test]
async fn a_created_response_carries_its_location() {
let upstream = testsrv::start(Script::default()).await;
let client = AcmeClient::discover(
&upstream.directory_url(),
crate::testutil::outbound_with(test_resolver()),
TIMEOUT,
)
.await
.unwrap();
let response = client
.post(
&key(),
&Signer::Jwk,
&client.directory().new_account.clone(),
Some(&json!({})),
)
.await
.unwrap();
assert_eq!(
response.location.as_deref(),
Some(format!("{}/acct/1", upstream.base).as_str())
);
}
#[tokio::test]
async fn post_as_get_sends_an_empty_payload() {
let upstream = testsrv::start(Script::default()).await;
let client = AcmeClient::discover(
&upstream.directory_url(),
crate::testutil::outbound_with(test_resolver()),
TIMEOUT,
)
.await
.unwrap();
let response = client
.get(&key(), "kid-1", &format!("{}/order/1", upstream.base))
.await
.unwrap();
assert_eq!(response.status, 200);
assert_eq!(upstream.order_polls(), 1);
}
#[tokio::test]
async fn an_error_response_becomes_a_problem() {
let upstream = testsrv::start(Script::default()).await;
let client = AcmeClient::discover(
&upstream.directory_url(),
crate::testutil::outbound_with(test_resolver()),
TIMEOUT,
)
.await
.unwrap();
let error = client
.get(&key(), "kid-1", &format!("{}/nope", upstream.base))
.await
.unwrap_err();
match error {
UpstreamError::Problem { status, typ, .. } => {
assert_eq!(status, 404);
assert!(typ.ends_with(":malformed"), "{typ}");
}
other => panic!("expected a problem document, got {other:?}"),
}
}
#[test]
fn a_non_problem_error_body_still_carries_the_status() {
let response = AcmeResponse {
status: StatusCode::BAD_GATEWAY,
body: Bytes::from_static(b"<html>proxy error</html>"),
location: None,
retry_after: None,
nonce: None,
};
match problem_from(&response) {
UpstreamError::Problem {
status,
typ,
detail,
} => {
assert_eq!(status, 502);
assert_eq!(typ, "about:blank");
assert!(detail.contains("proxy error"), "{detail}");
}
other => panic!("expected a problem, got {other:?}"),
}
}
#[test]
fn response_helpers_decode_json_and_text() {
let response = AcmeResponse {
status: StatusCode::OK,
body: Bytes::from_static(br#"{"status":"valid"}"#),
location: None,
retry_after: None,
nonce: None,
};
let value: Value = response.json().unwrap();
assert_eq!(value["status"], "valid");
assert_eq!(response.text().unwrap(), r#"{"status":"valid"}"#);
let invalid = AcmeResponse {
status: StatusCode::OK,
body: Bytes::from_static(b"not json"),
location: None,
retry_after: None,
nonce: None,
};
assert!(invalid.json::<Value>().is_err());
let binary = AcmeResponse {
status: StatusCode::OK,
body: Bytes::from_static(&[0xff, 0xfe]),
location: None,
retry_after: None,
nonce: None,
};
assert!(binary.text().is_err());
}
#[tokio::test]
async fn a_silent_server_times_out() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
std::mem::forget(stream);
});
let error = AcmeClient::discover(
&format!("http://127.0.0.1:{port}/directory"),
crate::testutil::outbound_with(test_resolver()),
Duration::from_millis(150),
)
.await
.unwrap_err();
assert!(
matches!(&error, UpstreamError::Transport(detail) if detail.contains("timed out")),
"{error:?}"
);
}
#[tokio::test]
async fn get_unsigned_reaches_an_endpoint_that_takes_no_jws() {
let upstream = testsrv::start(Script::default()).await;
let client = AcmeClient::discover(
&upstream.directory_url(),
crate::testutil::outbound_with(test_resolver()),
TIMEOUT,
)
.await
.unwrap();
let response = client
.get_unsigned(&upstream.directory_url())
.await
.unwrap();
assert_eq!(response.status, 200);
assert!(
client
.get_unsigned(&format!("{}/nope", upstream.base))
.await
.is_err()
);
}
}