otel_bootstrap/
rotating_mtls.rs1use std::error::Error;
13use std::future::Future;
14use std::pin::Pin;
15use std::sync::Arc;
16use std::task::{Context, Poll};
17use std::time::Duration;
18
19use hyper_util::rt::TokioIo;
20use rustls::RootCertStore;
21use rustls::pki_types::pem::PemObject as _;
22use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName};
23use tokio::net::TcpStream;
24use tonic::codegen::http::Uri;
25use tonic::transport::{Channel, Endpoint};
26
27use crate::MtlsMaterial;
28
29type BoxError = Box<dyn Error + Send + Sync>;
30
31pub trait CertSource: Send + Sync + 'static {
38 fn current(&self) -> Option<MtlsMaterial>;
43}
44
45#[derive(Clone)]
49pub struct StaticCertSource(MtlsMaterial);
50
51impl StaticCertSource {
52 #[must_use]
54 pub fn new(material: MtlsMaterial) -> Self {
55 Self(material)
56 }
57}
58
59impl CertSource for StaticCertSource {
60 fn current(&self) -> Option<MtlsMaterial> {
61 Some(self.0.clone())
62 }
63}
64
65fn client_config(material: &MtlsMaterial) -> Result<rustls::ClientConfig, BoxError> {
69 let mut roots = RootCertStore::empty();
70 for ca in CertificateDer::pem_slice_iter(&material.trust_bundle_pem) {
71 roots.add(ca?)?;
72 }
73 if roots.is_empty() {
74 return Err("mTLS trust bundle holds no certificates".into());
75 }
76 let chain = CertificateDer::pem_slice_iter(&material.client_cert_chain_pem)
77 .collect::<Result<Vec<_>, _>>()?;
78 let key = PrivateKeyDer::from_pem_slice(&material.client_key_pem)?;
79
80 let mut config = rustls::ClientConfig::builder_with_provider(Arc::new(
81 rustls::crypto::aws_lc_rs::default_provider(),
82 ))
83 .with_safe_default_protocol_versions()?
84 .with_root_certificates(roots)
85 .with_client_auth_cert(chain, key)?;
86 config.alpn_protocols = vec![b"h2".to_vec()];
89 Ok(config)
90}
91
92#[derive(Clone)]
95pub(crate) struct RotatingConnector {
96 source: Arc<dyn CertSource>,
97}
98
99type TlsIo = TokioIo<tokio_rustls::client::TlsStream<TcpStream>>;
100
101impl RotatingConnector {
102 async fn connect(source: Arc<dyn CertSource>, uri: Uri) -> Result<TlsIo, BoxError> {
103 let host = uri.host().ok_or("OTLP endpoint has no host")?.to_owned();
104 let port = uri.port_u16().unwrap_or(443);
105 let material = source
106 .current()
107 .ok_or("mTLS material is not available yet")?;
108 let config = client_config(&material)?;
109 let tcp = TcpStream::connect((host.as_str(), port)).await?;
110 tcp.set_nodelay(true)?;
111 let name = ServerName::try_from(host)?;
112 let tls = tokio_rustls::TlsConnector::from(Arc::new(config))
113 .connect(name, tcp)
114 .await?;
115 Ok(TokioIo::new(tls))
116 }
117}
118
119impl tower::Service<Uri> for RotatingConnector {
120 type Response = TlsIo;
121 type Error = BoxError;
122 type Future = Pin<Box<dyn Future<Output = Result<TlsIo, BoxError>> + Send>>;
123
124 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
125 Poll::Ready(Ok(()))
126 }
127
128 fn call(&mut self, uri: Uri) -> Self::Future {
129 Box::pin(Self::connect(Arc::clone(&self.source), uri))
130 }
131}
132
133pub(crate) fn channel(
139 endpoint: &str,
140 source: Arc<dyn CertSource>,
141 timeout: Option<Duration>,
142) -> Result<Channel, BoxError> {
143 let parsed: Uri = endpoint.parse()?;
144 let authority = parsed
145 .authority()
146 .ok_or("OTLP endpoint has no host")?
147 .as_str();
148 let mut ep = Endpoint::from_shared(format!("http://{authority}"))?
149 .connect_timeout(Duration::from_secs(10));
150 if let Some(t) = timeout {
151 ep = ep.timeout(t);
152 }
153 Ok(ep.connect_with_connector_lazy(RotatingConnector { source }))
154}
155
156#[cfg(test)]
157mod tests {
158 use super::*;
159 use std::sync::Mutex;
160 use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
161 use tokio::net::TcpListener;
162 use tower::Service as _;
163
164 const CA: &str = include_str!("../tests/fixtures/mtls/ca.pem");
168 const STRANGER_CA: &str = include_str!("../tests/fixtures/mtls/stranger-ca.pem");
169 const SERVER_CERT: &str = include_str!("../tests/fixtures/mtls/server.pem");
170 const SERVER_KEY: &str = include_str!("../tests/fixtures/mtls/server.key");
171 const CLIENT_ONE_CERT: &str = include_str!("../tests/fixtures/mtls/client-one.pem");
172 const CLIENT_ONE_KEY: &str = include_str!("../tests/fixtures/mtls/client-one.key");
173 const CLIENT_TWO_CERT: &str = include_str!("../tests/fixtures/mtls/client-two.pem");
174 const CLIENT_TWO_KEY: &str = include_str!("../tests/fixtures/mtls/client-two.key");
175
176 fn material(cert: &str, key: &str, bundle: &str) -> MtlsMaterial {
177 MtlsMaterial {
178 client_cert_chain_pem: cert.as_bytes().to_vec(),
179 client_key_pem: key.as_bytes().to_vec(),
180 trust_bundle_pem: bundle.as_bytes().to_vec(),
181 }
182 }
183
184 fn client_one() -> MtlsMaterial {
185 material(CLIENT_ONE_CERT, CLIENT_ONE_KEY, CA)
186 }
187
188 fn client_two() -> MtlsMaterial {
189 material(CLIENT_TWO_CERT, CLIENT_TWO_KEY, CA)
190 }
191
192 async fn server() -> (u16, Arc<Mutex<Vec<Vec<u8>>>>) {
195 let mut roots = RootCertStore::empty();
196 roots
197 .add(CertificateDer::from_pem_slice(CA.as_bytes()).unwrap())
198 .unwrap();
199 let verifier = rustls::server::WebPkiClientVerifier::builder_with_provider(
200 Arc::new(roots),
201 Arc::new(rustls::crypto::aws_lc_rs::default_provider()),
202 )
203 .build()
204 .unwrap();
205 let mut config = rustls::ServerConfig::builder_with_provider(Arc::new(
206 rustls::crypto::aws_lc_rs::default_provider(),
207 ))
208 .with_safe_default_protocol_versions()
209 .unwrap()
210 .with_client_cert_verifier(verifier)
211 .with_single_cert(
212 CertificateDer::pem_slice_iter(SERVER_CERT.as_bytes())
213 .collect::<Result<Vec<_>, _>>()
214 .unwrap(),
215 PrivateKeyDer::from_pem_slice(SERVER_KEY.as_bytes()).unwrap(),
216 )
217 .unwrap();
218 config.alpn_protocols = vec![b"h2".to_vec()];
219 let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(config));
220
221 let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
222 let port = listener.local_addr().unwrap().port();
223 let seen = Arc::new(Mutex::new(Vec::new()));
224 let seen_task = Arc::clone(&seen);
225 tokio::spawn(async move {
226 loop {
227 let (tcp, _) = listener.accept().await.unwrap();
228 let acceptor = acceptor.clone();
229 let seen = Arc::clone(&seen_task);
230 tokio::spawn(async move {
231 if let Ok(mut tls) = acceptor.accept(tcp).await {
232 let leaf = tls
233 .get_ref()
234 .1
235 .peer_certificates()
236 .and_then(|c| c.first())
237 .map(|c| c.as_ref().to_vec());
238 if let Some(leaf) = leaf {
239 seen.lock().unwrap().push(leaf);
240 }
241 let mut b = [0_u8; 1];
242 let _ = tls.read(&mut b).await;
243 let _ = tls.shutdown().await;
244 }
245 });
246 }
247 });
248 (port, seen)
249 }
250
251 struct Swappable(Mutex<Option<MtlsMaterial>>);
252 impl CertSource for Swappable {
253 fn current(&self) -> Option<MtlsMaterial> {
254 self.0.lock().unwrap().clone()
255 }
256 }
257
258 fn uri(port: u16) -> Uri {
259 format!("http://localhost:{port}").parse().unwrap()
260 }
261
262 fn pem_to_der(pem: &[u8]) -> Vec<u8> {
263 CertificateDer::from_pem_slice(pem).unwrap().to_vec()
264 }
265
266 #[tokio::test]
267 async fn each_new_connection_presents_the_material_current_at_that_moment() {
268 let (port, seen) = server().await;
269 let source = Arc::new(Swappable(Mutex::new(Some(client_one()))));
270 let mut connector = RotatingConnector {
271 source: Arc::clone(&source) as Arc<dyn CertSource>,
272 };
273
274 std::future::poll_fn(|cx| connector.poll_ready(cx))
275 .await
276 .expect("connector is always ready");
277 drop(connector.call(uri(port)).await.expect("first connection"));
278
279 *source.0.lock().unwrap() = Some(client_two());
280 drop(connector.call(uri(port)).await.expect("rotated connection"));
281
282 tokio::time::timeout(Duration::from_secs(5), async {
283 while seen.lock().unwrap().len() < 2 {
284 tokio::time::sleep(Duration::from_millis(10)).await;
285 }
286 })
287 .await
288 .expect("server saw both clients");
289 let seen = seen.lock().unwrap();
290 assert_eq!(seen[0], pem_to_der(CLIENT_ONE_CERT.as_bytes()));
291 assert_eq!(seen[1], pem_to_der(CLIENT_TWO_CERT.as_bytes()));
292 }
293
294 #[tokio::test]
295 async fn a_source_not_ready_at_first_connect_fails_that_attempt_then_connects_once_ready() {
296 let (port, seen) = server().await;
297 let source = Arc::new(Swappable(Mutex::new(None)));
298 let mut connector = RotatingConnector {
299 source: Arc::clone(&source) as Arc<dyn CertSource>,
300 };
301
302 let err = connector
303 .call(uri(port))
304 .await
305 .err()
306 .expect("an unready source must fail the attempt");
307 assert!(err.to_string().contains("not available yet"), "{err}");
308 assert!(seen.lock().unwrap().is_empty());
309
310 *source.0.lock().unwrap() = Some(client_one());
311 drop(
312 connector
313 .call(uri(port))
314 .await
315 .expect("connects once ready"),
316 );
317 }
318
319 #[tokio::test]
320 async fn a_bundle_of_several_certificates_trusts_each_of_them() {
321 let (port, _seen) = server().await;
322 let bundle = format!("{STRANGER_CA}{CA}");
325 let mut connector = RotatingConnector {
326 source: Arc::new(StaticCertSource::new(material(
327 CLIENT_ONE_CERT,
328 CLIENT_ONE_KEY,
329 &bundle,
330 ))),
331 };
332 drop(
333 connector
334 .call(uri(port))
335 .await
336 .expect("second anchor trusted"),
337 );
338 }
339
340 #[tokio::test]
341 async fn a_server_the_bundle_does_not_vouch_for_is_refused() {
342 let (port, _seen) = server().await;
343 let mut connector = RotatingConnector {
344 source: Arc::new(StaticCertSource::new(material(
345 CLIENT_ONE_CERT,
346 CLIENT_ONE_KEY,
347 STRANGER_CA,
348 ))),
349 };
350 assert!(connector.call(uri(port)).await.is_err());
351 }
352
353 #[test]
354 fn malformed_material_is_an_error_not_a_panic() {
355 let bad = MtlsMaterial {
356 client_cert_chain_pem: b"x".to_vec(),
357 client_key_pem: b"x".to_vec(),
358 trust_bundle_pem: Vec::new(),
359 };
360 assert!(
361 client_config(&bad)
362 .unwrap_err()
363 .to_string()
364 .contains("no certificates")
365 );
366 }
367
368 #[test]
369 fn static_source_returns_its_material() {
370 let m = client_one();
371 let got = StaticCertSource::new(m.clone()).current().unwrap();
372 assert_eq!(got.client_cert_chain_pem, m.client_cert_chain_pem);
373 }
374
375 #[tokio::test]
376 async fn channel_rejects_an_endpoint_without_a_host() {
377 let source: Arc<dyn CertSource> = Arc::new(Swappable(Mutex::new(None)));
378 assert!(channel("not a uri", Arc::clone(&source), None).is_err());
379 assert!(
380 channel(
381 "https://collector.example:4320",
382 source,
383 Some(Duration::from_secs(3))
384 )
385 .is_ok()
386 );
387 }
388}