1use std::sync::Arc;
22
23use async_trait::async_trait;
24use serde::{Deserialize, Serialize};
25
26use crate::envelope::KeyEnvelope;
27use crate::kv::{KvError, KvStore};
28
29pub use boatramp_types::cert::CertStatus;
33
34#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
37pub struct StoredCert {
38 #[serde(default = "crate::schema_version")]
40 pub version: u32,
41 pub chain_pem: String,
43 pub key_pem: String,
45 pub not_after_unix: u64,
47}
48
49impl StoredCert {
50 pub fn new(
52 chain_pem: impl Into<String>,
53 key_pem: impl Into<String>,
54 not_after_unix: u64,
55 ) -> Self {
56 Self {
57 version: crate::SCHEMA_VERSION,
58 chain_pem: chain_pem.into(),
59 key_pem: key_pem.into(),
60 not_after_unix,
61 }
62 }
63}
64
65#[derive(Debug)]
67pub enum CertError {
68 Kv(KvError),
70 Decode(String),
72 Issue(String),
74 Envelope(String),
76}
77
78impl std::fmt::Display for CertError {
79 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
80 match self {
81 Self::Kv(e) => write!(f, "cert store kv error: {e}"),
82 Self::Decode(m) => write!(f, "cert decode error: {m}"),
83 Self::Issue(m) => write!(f, "cert issuance error: {m}"),
84 Self::Envelope(m) => write!(f, "cert key envelope error: {m}"),
85 }
86 }
87}
88
89impl std::error::Error for CertError {}
90
91impl From<KvError> for CertError {
92 fn from(e: KvError) -> Self {
93 Self::Kv(e)
94 }
95}
96
97pub fn cert_key(domain: &str) -> String {
99 format!("cert/{domain}")
100}
101
102#[async_trait]
104pub trait CertStore: Send + Sync {
105 async fn get(&self, domain: &str) -> Result<Option<StoredCert>, CertError>;
107 async fn put(&self, domain: &str, cert: &StoredCert) -> Result<(), CertError>;
109}
110
111pub struct KvCertStore {
120 kv: Arc<dyn KvStore>,
121 envelope: Option<Arc<dyn KeyEnvelope>>,
122}
123
124impl KvCertStore {
125 pub fn new(kv: Arc<dyn KvStore>) -> Self {
127 Self { kv, envelope: None }
128 }
129
130 pub fn with_envelope(kv: Arc<dyn KvStore>, envelope: Arc<dyn KeyEnvelope>) -> Self {
132 Self {
133 kv,
134 envelope: Some(envelope),
135 }
136 }
137}
138
139#[async_trait]
140impl CertStore for KvCertStore {
141 async fn get(&self, domain: &str) -> Result<Option<StoredCert>, CertError> {
142 let Some(raw) = self.kv.get(&cert_key(domain)).await? else {
143 return Ok(None);
144 };
145 let mut cert: StoredCert =
146 serde_json::from_slice(&raw).map_err(|e| CertError::Decode(e.to_string()))?;
147 if let Some(envelope) = &self.envelope {
148 let wrapped =
150 hex::decode(cert.key_pem.trim()).map_err(|e| CertError::Envelope(e.to_string()))?;
151 let plaintext = envelope
152 .unwrap(&wrapped)
153 .await
154 .map_err(|e| CertError::Envelope(e.to_string()))?;
155 cert.key_pem =
156 String::from_utf8(plaintext).map_err(|e| CertError::Envelope(e.to_string()))?;
157 }
158 Ok(Some(cert))
159 }
160
161 async fn put(&self, domain: &str, cert: &StoredCert) -> Result<(), CertError> {
162 let to_store = if let Some(envelope) = &self.envelope {
164 let wrapped = envelope
165 .wrap(cert.key_pem.as_bytes())
166 .await
167 .map_err(|e| CertError::Envelope(e.to_string()))?;
168 StoredCert {
169 key_pem: hex::encode(wrapped),
170 ..cert.clone()
171 }
172 } else {
173 cert.clone()
174 };
175 let json = serde_json::to_vec(&to_store).map_err(|e| CertError::Decode(e.to_string()))?;
176 self.kv.put(&cert_key(domain), json).await?;
177 Ok(())
178 }
179}
180
181pub async fn ensure_cert<F, Fut, E>(
191 store: &dyn CertStore,
192 domain: &str,
193 is_leader: bool,
194 now_unix: u64,
195 renew_before_secs: u64,
196 issue: F,
197) -> Result<Option<StoredCert>, CertError>
198where
199 F: FnOnce() -> Fut,
200 Fut: std::future::Future<Output = Result<StoredCert, E>>,
201 E: std::fmt::Display,
202{
203 let existing = store.get(domain).await?;
204 let fresh = existing
205 .as_ref()
206 .is_some_and(|c| c.not_after_unix > now_unix.saturating_add(renew_before_secs));
207 if fresh {
208 return Ok(existing);
209 }
210 if !is_leader {
211 return Ok(existing);
213 }
214 let cert = issue().await.map_err(|e| CertError::Issue(e.to_string()))?;
215 store.put(domain, &cert).await?;
216 Ok(Some(cert))
217}
218
219#[cfg(test)]
220mod tests {
221 use super::*;
222 use crate::envelope::EnvelopeError;
223 use crate::kv::MemoryKv;
224 use std::sync::atomic::{AtomicUsize, Ordering};
225
226 fn store() -> KvCertStore {
227 KvCertStore::new(Arc::new(MemoryKv::new()))
228 }
229
230 struct ReverseEnvelope;
233
234 #[async_trait]
235 impl KeyEnvelope for ReverseEnvelope {
236 async fn wrap(&self, plaintext: &[u8]) -> Result<Vec<u8>, EnvelopeError> {
237 let mut out = vec![0xEE];
238 out.extend(plaintext.iter().rev());
239 Ok(out)
240 }
241 async fn unwrap(&self, wrapped: &[u8]) -> Result<Vec<u8>, EnvelopeError> {
242 match wrapped.split_first() {
243 Some((0xEE, rest)) => Ok(rest.iter().rev().copied().collect()),
244 _ => Err(EnvelopeError::new("not a ReverseEnvelope blob")),
245 }
246 }
247 }
248
249 #[tokio::test]
252 async fn envelope_wraps_the_key_at_rest_and_reads_recover_it() {
253 let kv: Arc<dyn KvStore> = Arc::new(MemoryKv::new());
254 let s = KvCertStore::with_envelope(kv.clone(), Arc::new(ReverseEnvelope));
255 let cert = StoredCert::new("CHAIN", "SECRET-KEY-PEM", 9999);
256 s.put("blog", &cert).await.unwrap();
257
258 let raw = kv.get(&cert_key("blog")).await.unwrap().unwrap();
260 let raw_str = String::from_utf8_lossy(&raw);
261 assert!(
262 !raw_str.contains("SECRET-KEY-PEM"),
263 "the private key must not be stored in cleartext"
264 );
265 assert!(raw_str.contains("CHAIN"), "the chain stays clear");
266
267 let got = s.get("blog").await.unwrap().unwrap();
269 assert_eq!(got, cert);
270 }
271
272 #[tokio::test]
274 async fn wrapped_key_is_unreadable_without_the_envelope() {
275 let kv: Arc<dyn KvStore> = Arc::new(MemoryKv::new());
276 KvCertStore::with_envelope(kv.clone(), Arc::new(ReverseEnvelope))
277 .put("blog", &StoredCert::new("C", "K", 1))
278 .await
279 .unwrap();
280 let plain = KvCertStore::new(kv);
283 let got = plain.get("blog").await.unwrap().unwrap();
284 assert_ne!(
285 got.key_pem, "K",
286 "cleartext read must not yield the real key"
287 );
288 }
289
290 fn issuer(
292 calls: &AtomicUsize,
293 not_after: u64,
294 ) -> impl FnOnce() -> std::future::Ready<Result<StoredCert, String>> + '_ {
295 move || {
296 calls.fetch_add(1, Ordering::SeqCst);
297 std::future::ready(Ok(StoredCert::new("CHAIN", "KEY", not_after)))
298 }
299 }
300
301 #[tokio::test]
302 async fn kv_cert_store_round_trips() {
303 let s = store();
304 assert!(s.get("blog.example.com").await.unwrap().is_none());
305 let cert = StoredCert::new("chain", "key", 1000);
306 s.put("blog.example.com", &cert).await.unwrap();
307 assert_eq!(s.get("blog.example.com").await.unwrap(), Some(cert));
308 }
309
310 #[tokio::test]
311 async fn leader_issues_once_then_serves_from_store() {
312 let s = store();
313 let calls = AtomicUsize::new(0);
314 let c = ensure_cert(&s, "d", true, 100, 50, issuer(&calls, 10_000))
316 .await
317 .unwrap();
318 assert!(c.is_some());
319 assert_eq!(calls.load(Ordering::SeqCst), 1);
320 let c2 = ensure_cert(&s, "d", true, 200, 50, issuer(&calls, 10_000))
322 .await
323 .unwrap();
324 assert_eq!(c2.unwrap().chain_pem, "CHAIN");
325 assert_eq!(
326 calls.load(Ordering::SeqCst),
327 1,
328 "fresh cert must not re-issue"
329 );
330 }
331
332 #[tokio::test]
333 async fn follower_never_issues_but_serves_replicated() {
334 let s = store();
335 let calls = AtomicUsize::new(0);
336 let c = ensure_cert(&s, "d", false, 100, 50, issuer(&calls, 10_000))
338 .await
339 .unwrap();
340 assert!(c.is_none());
341 assert_eq!(
342 calls.load(Ordering::SeqCst),
343 0,
344 "a follower must not call the CA"
345 );
346 s.put("d", &StoredCert::new("CHAIN", "KEY", 10_000))
348 .await
349 .unwrap();
350 let c = ensure_cert(&s, "d", false, 200, 50, issuer(&calls, 10_000))
352 .await
353 .unwrap();
354 assert_eq!(c.unwrap().chain_pem, "CHAIN");
355 assert_eq!(calls.load(Ordering::SeqCst), 0);
356 }
357
358 #[tokio::test]
359 async fn leader_renews_near_expiry() {
360 let s = store();
361 let calls = AtomicUsize::new(0);
362 s.put("d", &StoredCert::new("OLD", "KEY", 1000))
364 .await
365 .unwrap();
366 let c = ensure_cert(&s, "d", true, 900, 200, issuer(&calls, 99_999))
367 .await
368 .unwrap();
369 assert_eq!(
370 calls.load(Ordering::SeqCst),
371 1,
372 "near-expiry cert must renew"
373 );
374 assert_eq!(c.unwrap().not_after_unix, 99_999);
375 }
376}