use std::sync::Arc;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use crate::envelope::KeyEnvelope;
use crate::kv::{KvError, KvStore};
pub use boatramp_types::cert::CertStatus;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct StoredCert {
#[serde(default = "crate::schema_version")]
pub version: u32,
pub chain_pem: String,
pub key_pem: String,
pub not_after_unix: u64,
}
impl StoredCert {
pub fn new(
chain_pem: impl Into<String>,
key_pem: impl Into<String>,
not_after_unix: u64,
) -> Self {
Self {
version: crate::SCHEMA_VERSION,
chain_pem: chain_pem.into(),
key_pem: key_pem.into(),
not_after_unix,
}
}
}
#[derive(Debug)]
pub enum CertError {
Kv(KvError),
Decode(String),
Issue(String),
Envelope(String),
}
impl std::fmt::Display for CertError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Kv(e) => write!(f, "cert store kv error: {e}"),
Self::Decode(m) => write!(f, "cert decode error: {m}"),
Self::Issue(m) => write!(f, "cert issuance error: {m}"),
Self::Envelope(m) => write!(f, "cert key envelope error: {m}"),
}
}
}
impl std::error::Error for CertError {}
impl From<KvError> for CertError {
fn from(e: KvError) -> Self {
Self::Kv(e)
}
}
pub fn cert_key(domain: &str) -> String {
format!("cert/{domain}")
}
#[async_trait]
pub trait CertStore: Send + Sync {
async fn get(&self, domain: &str) -> Result<Option<StoredCert>, CertError>;
async fn put(&self, domain: &str, cert: &StoredCert) -> Result<(), CertError>;
}
pub struct KvCertStore {
kv: Arc<dyn KvStore>,
envelope: Option<Arc<dyn KeyEnvelope>>,
}
impl KvCertStore {
pub fn new(kv: Arc<dyn KvStore>) -> Self {
Self { kv, envelope: None }
}
pub fn with_envelope(kv: Arc<dyn KvStore>, envelope: Arc<dyn KeyEnvelope>) -> Self {
Self {
kv,
envelope: Some(envelope),
}
}
}
#[async_trait]
impl CertStore for KvCertStore {
async fn get(&self, domain: &str) -> Result<Option<StoredCert>, CertError> {
let Some(raw) = self.kv.get(&cert_key(domain)).await? else {
return Ok(None);
};
let mut cert: StoredCert =
serde_json::from_slice(&raw).map_err(|e| CertError::Decode(e.to_string()))?;
if let Some(envelope) = &self.envelope {
let wrapped =
hex::decode(cert.key_pem.trim()).map_err(|e| CertError::Envelope(e.to_string()))?;
let plaintext = envelope
.unwrap(&wrapped)
.await
.map_err(|e| CertError::Envelope(e.to_string()))?;
cert.key_pem =
String::from_utf8(plaintext).map_err(|e| CertError::Envelope(e.to_string()))?;
}
Ok(Some(cert))
}
async fn put(&self, domain: &str, cert: &StoredCert) -> Result<(), CertError> {
let to_store = if let Some(envelope) = &self.envelope {
let wrapped = envelope
.wrap(cert.key_pem.as_bytes())
.await
.map_err(|e| CertError::Envelope(e.to_string()))?;
StoredCert {
key_pem: hex::encode(wrapped),
..cert.clone()
}
} else {
cert.clone()
};
let json = serde_json::to_vec(&to_store).map_err(|e| CertError::Decode(e.to_string()))?;
self.kv.put(&cert_key(domain), json).await?;
Ok(())
}
}
pub async fn ensure_cert<F, Fut, E>(
store: &dyn CertStore,
domain: &str,
is_leader: bool,
now_unix: u64,
renew_before_secs: u64,
issue: F,
) -> Result<Option<StoredCert>, CertError>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = Result<StoredCert, E>>,
E: std::fmt::Display,
{
let existing = store.get(domain).await?;
let fresh = existing
.as_ref()
.is_some_and(|c| c.not_after_unix > now_unix.saturating_add(renew_before_secs));
if fresh {
return Ok(existing);
}
if !is_leader {
return Ok(existing);
}
let cert = issue().await.map_err(|e| CertError::Issue(e.to_string()))?;
store.put(domain, &cert).await?;
Ok(Some(cert))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::envelope::EnvelopeError;
use crate::kv::MemoryKv;
use std::sync::atomic::{AtomicUsize, Ordering};
fn store() -> KvCertStore {
KvCertStore::new(Arc::new(MemoryKv::new()))
}
struct ReverseEnvelope;
#[async_trait]
impl KeyEnvelope for ReverseEnvelope {
async fn wrap(&self, plaintext: &[u8]) -> Result<Vec<u8>, EnvelopeError> {
let mut out = vec![0xEE];
out.extend(plaintext.iter().rev());
Ok(out)
}
async fn unwrap(&self, wrapped: &[u8]) -> Result<Vec<u8>, EnvelopeError> {
match wrapped.split_first() {
Some((0xEE, rest)) => Ok(rest.iter().rev().copied().collect()),
_ => Err(EnvelopeError::new("not a ReverseEnvelope blob")),
}
}
}
#[tokio::test]
async fn envelope_wraps_the_key_at_rest_and_reads_recover_it() {
let kv: Arc<dyn KvStore> = Arc::new(MemoryKv::new());
let s = KvCertStore::with_envelope(kv.clone(), Arc::new(ReverseEnvelope));
let cert = StoredCert::new("CHAIN", "SECRET-KEY-PEM", 9999);
s.put("blog", &cert).await.unwrap();
let raw = kv.get(&cert_key("blog")).await.unwrap().unwrap();
let raw_str = String::from_utf8_lossy(&raw);
assert!(
!raw_str.contains("SECRET-KEY-PEM"),
"the private key must not be stored in cleartext"
);
assert!(raw_str.contains("CHAIN"), "the chain stays clear");
let got = s.get("blog").await.unwrap().unwrap();
assert_eq!(got, cert);
}
#[tokio::test]
async fn wrapped_key_is_unreadable_without_the_envelope() {
let kv: Arc<dyn KvStore> = Arc::new(MemoryKv::new());
KvCertStore::with_envelope(kv.clone(), Arc::new(ReverseEnvelope))
.put("blog", &StoredCert::new("C", "K", 1))
.await
.unwrap();
let plain = KvCertStore::new(kv);
let got = plain.get("blog").await.unwrap().unwrap();
assert_ne!(
got.key_pem, "K",
"cleartext read must not yield the real key"
);
}
fn issuer(
calls: &AtomicUsize,
not_after: u64,
) -> impl FnOnce() -> std::future::Ready<Result<StoredCert, String>> + '_ {
move || {
calls.fetch_add(1, Ordering::SeqCst);
std::future::ready(Ok(StoredCert::new("CHAIN", "KEY", not_after)))
}
}
#[tokio::test]
async fn kv_cert_store_round_trips() {
let s = store();
assert!(s.get("blog.example.com").await.unwrap().is_none());
let cert = StoredCert::new("chain", "key", 1000);
s.put("blog.example.com", &cert).await.unwrap();
assert_eq!(s.get("blog.example.com").await.unwrap(), Some(cert));
}
#[tokio::test]
async fn leader_issues_once_then_serves_from_store() {
let s = store();
let calls = AtomicUsize::new(0);
let c = ensure_cert(&s, "d", true, 100, 50, issuer(&calls, 10_000))
.await
.unwrap();
assert!(c.is_some());
assert_eq!(calls.load(Ordering::SeqCst), 1);
let c2 = ensure_cert(&s, "d", true, 200, 50, issuer(&calls, 10_000))
.await
.unwrap();
assert_eq!(c2.unwrap().chain_pem, "CHAIN");
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"fresh cert must not re-issue"
);
}
#[tokio::test]
async fn follower_never_issues_but_serves_replicated() {
let s = store();
let calls = AtomicUsize::new(0);
let c = ensure_cert(&s, "d", false, 100, 50, issuer(&calls, 10_000))
.await
.unwrap();
assert!(c.is_none());
assert_eq!(
calls.load(Ordering::SeqCst),
0,
"a follower must not call the CA"
);
s.put("d", &StoredCert::new("CHAIN", "KEY", 10_000))
.await
.unwrap();
let c = ensure_cert(&s, "d", false, 200, 50, issuer(&calls, 10_000))
.await
.unwrap();
assert_eq!(c.unwrap().chain_pem, "CHAIN");
assert_eq!(calls.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn leader_renews_near_expiry() {
let s = store();
let calls = AtomicUsize::new(0);
s.put("d", &StoredCert::new("OLD", "KEY", 1000))
.await
.unwrap();
let c = ensure_cert(&s, "d", true, 900, 200, issuer(&calls, 99_999))
.await
.unwrap();
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"near-expiry cert must renew"
);
assert_eq!(c.unwrap().not_after_unix, 99_999);
}
}