Skip to main content

vector_core/tls/
reload.rs

1//! A TLS acceptor that can be swapped at runtime.
2//!
3//! A [`TlsAcceptorReloader`] wraps the acceptor a bound [`MaybeTlsListener`](super::MaybeTlsListener)
4//! serves. Handing the same reloader to
5//! [`MaybeTlsSettings::bind_reloadable`](super::MaybeTlsSettings::bind_reloadable) lets a background
6//! task swap in freshly built material with [`reload`](TlsAcceptorReloader::reload); each new
7//! connection then handshakes with the latest acceptor while in-flight connections keep what they
8//! negotiated.
9
10use std::sync::{Arc, Weak};
11
12use arc_swap::ArcSwap;
13use openssl::ssl::SslAcceptor;
14
15use super::{MaybeTlsSettings, TlsSettings};
16
17/// A cloneable handle to a server TLS acceptor that can be swapped at runtime.
18///
19/// Connections accepted after [`reload`](Self::reload) use the new acceptor; connections already
20/// established keep whatever they negotiated at handshake time.
21#[derive(Clone)]
22pub struct TlsAcceptorReloader {
23    acceptor: Arc<ArcSwap<SslAcceptor>>,
24}
25
26impl TlsAcceptorReloader {
27    /// Wrap an initial acceptor in a swappable cell.
28    pub(super) fn new(acceptor: SslAcceptor) -> Self {
29        Self {
30            acceptor: Arc::new(ArcSwap::from_pointee(acceptor)),
31        }
32    }
33
34    /// The shared cell the bound listener reads from on each accept.
35    pub(super) fn shared(&self) -> Arc<ArcSwap<SslAcceptor>> {
36        Arc::clone(&self.acceptor)
37    }
38
39    /// Swap in a freshly built acceptor from `settings`. New connections pick it up
40    /// immediately; the previous acceptor is dropped once its last in-flight handshake completes.
41    pub fn reload(&self, settings: &TlsSettings) -> crate::tls::Result<()> {
42        self.acceptor.store(Arc::new(settings.acceptor()?));
43        Ok(())
44    }
45
46    /// Downgrade to a [`WeakTlsAcceptorReloader`] that does not keep the served acceptor alive.
47    pub fn downgrade(&self) -> WeakTlsAcceptorReloader {
48        WeakTlsAcceptorReloader {
49            acceptor: Arc::downgrade(&self.acceptor),
50        }
51    }
52}
53
54/// A non-owning handle to a served TLS acceptor, obtained from [`TlsAcceptorReloader::downgrade`].
55#[derive(Clone)]
56pub struct WeakTlsAcceptorReloader {
57    acceptor: Weak<ArcSwap<SslAcceptor>>,
58}
59
60impl WeakTlsAcceptorReloader {
61    /// Return the live [`TlsAcceptorReloader`], or `None` once the bound listener has been dropped.
62    pub fn upgrade(&self) -> Option<TlsAcceptorReloader> {
63        self.acceptor
64            .upgrade()
65            .map(|acceptor| TlsAcceptorReloader { acceptor })
66    }
67}
68
69impl MaybeTlsSettings {
70    /// Build a reloadable acceptor handle for server use, or `None` when TLS is disabled.
71    ///
72    /// Pass the returned handle to [`bind_reloadable`](Self::bind_reloadable) so the bound listener
73    /// serves it, and keep a clone to call [`reload`](TlsAcceptorReloader::reload) when the
74    /// certificate material rotates.
75    pub fn reloadable_acceptor(&self) -> crate::tls::Result<Option<TlsAcceptorReloader>> {
76        match self {
77            Self::Tls(tls) => Ok(Some(TlsAcceptorReloader::new(tls.acceptor()?))),
78            Self::Raw(()) => Ok(None),
79        }
80    }
81}
82
83#[cfg(test)]
84mod test {
85    use std::{net::SocketAddr, pin::Pin};
86
87    use openssl::{
88        asn1::Asn1Time,
89        bn::{BigNum, MsbOption},
90        hash::MessageDigest,
91        nid::Nid,
92        pkey::PKey,
93        rsa::Rsa,
94        ssl::{SslConnector, SslMethod, SslVerifyMode},
95        x509::{X509, X509NameBuilder},
96    };
97
98    use crate::tls::{MaybeTls, MaybeTlsSettings, TlsConfig, TlsEnableableConfig};
99
100    #[test]
101    fn no_reloadable_acceptor_without_tls() {
102        assert!(
103            MaybeTlsSettings::Raw(())
104                .reloadable_acceptor()
105                .unwrap()
106                .is_none(),
107            "plaintext settings have no acceptor to reload"
108        );
109    }
110
111    #[tokio::test]
112    async fn reloadable_acceptor_swaps_and_detects_shutdown() {
113        let settings =
114            MaybeTlsSettings::from_config(Some(&TlsEnableableConfig::test_config()), true).unwrap();
115        let tls = match &settings {
116            MaybeTls::Tls(tls) => tls.clone(),
117            MaybeTls::Raw(()) => panic!("expected TLS to be enabled"),
118        };
119
120        let reloader = settings
121            .reloadable_acceptor()
122            .unwrap()
123            .expect("tls enabled, so an acceptor should exist");
124        let weak = reloader.downgrade();
125
126        // Binding takes over the reloader's (sole) strong reference to the served acceptor.
127        let addr = "127.0.0.1:0".parse().unwrap();
128        let listener = settings
129            .bind_reloadable(&addr, Some(reloader))
130            .await
131            .unwrap();
132
133        weak.upgrade()
134            .expect("listener alive, so the weak handle upgrades")
135            .reload(&tls)
136            .unwrap();
137
138        // Once the listener (the last strong owner) is dropped, the weak handle no longer upgrades.
139        drop(listener);
140        assert!(
141            weak.upgrade().is_none(),
142            "weak handle must not upgrade after the listener is dropped"
143        );
144    }
145
146    /// End-to-end: bind a reloadable TLS listener, complete a real handshake and confirm the served
147    /// leaf certificate, then reload with a different certificate and confirm a fresh connection is
148    /// served the new one.
149    #[tokio::test]
150    async fn served_certificate_changes_after_reload() {
151        let (crt_a, key_a) = self_signed("old.example");
152        let (crt_b, key_b) = self_signed("new.example");
153
154        let settings = server_settings(&crt_a, &key_a);
155        let reloader = settings
156            .reloadable_acceptor()
157            .unwrap()
158            .expect("tls enabled, so an acceptor should exist");
159
160        let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
161        let mut listener = settings
162            .bind_reloadable(&addr, Some(reloader.clone()))
163            .await
164            .unwrap();
165        let local_addr = listener.local_addr().unwrap();
166
167        // Accept and complete the server side of each handshake until the test drops the listener.
168        let server = tokio::spawn(async move {
169            while let Ok(mut stream) = listener.accept().await {
170                // The client may drop as soon as it has the cert, so a handshake error is expected.
171                stream.handshake().await.ok();
172            }
173        });
174
175        // Before any reload, the original certificate is served.
176        assert_eq!(served_common_name(local_addr).await, "old.example");
177
178        // Reload with a different certificate...
179        let settings_b = server_settings(&crt_b, &key_b);
180        let tls_b = match &settings_b {
181            MaybeTls::Tls(tls) => tls.clone(),
182            MaybeTls::Raw(()) => unreachable!(),
183        };
184        reloader.reload(&tls_b).unwrap();
185
186        // ...and a new connection is served the rotated certificate.
187        assert_eq!(served_common_name(local_addr).await, "new.example");
188
189        server.abort();
190    }
191
192    /// Connect as a TLS client (trusting any server cert) and return the CN of the presented leaf.
193    async fn served_common_name(addr: SocketAddr) -> String {
194        let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
195
196        let mut builder = SslConnector::builder(SslMethod::tls()).unwrap();
197        builder.set_verify(SslVerifyMode::NONE);
198        let mut config = builder.build().configure().unwrap();
199        config.set_verify_hostname(false);
200        let ssl = config.into_ssl("localhost").unwrap();
201
202        let mut stream = tokio_openssl::SslStream::new(ssl, tcp).unwrap();
203        Pin::new(&mut stream).connect().await.unwrap();
204
205        let cert = stream
206            .ssl()
207            .peer_certificate()
208            .expect("server presents a certificate");
209        cert.subject_name()
210            .entries_by_nid(Nid::COMMONNAME)
211            .next()
212            .unwrap()
213            .data()
214            .to_string()
215            .unwrap()
216    }
217
218    fn server_settings(crt_pem: &str, key_pem: &str) -> MaybeTlsSettings {
219        // `crt_file`/`key_file` accept inline PEM (detected by the `-----BEGIN ` marker), so no
220        // temp files are needed.
221        let config = TlsEnableableConfig {
222            enabled: Some(true),
223            options: TlsConfig {
224                crt_file: Some(crt_pem.into()),
225                key_file: Some(key_pem.into()),
226                ..Default::default()
227            },
228        };
229        MaybeTlsSettings::from_config(Some(&config), true).unwrap()
230    }
231
232    /// Generate a self-signed certificate/key pair (PEM) with the given common name.
233    fn self_signed(common_name: &str) -> (String, String) {
234        let key = PKey::from_rsa(Rsa::generate(2048).unwrap()).unwrap();
235
236        let mut name = X509NameBuilder::new().unwrap();
237        name.append_entry_by_text("CN", common_name).unwrap();
238        let name = name.build();
239
240        let mut serial = BigNum::new().unwrap();
241        serial.rand(128, MsbOption::MAYBE_ZERO, false).unwrap();
242
243        let mut builder = X509::builder().unwrap();
244        builder.set_version(2).unwrap();
245        builder
246            .set_serial_number(&serial.to_asn1_integer().unwrap())
247            .unwrap();
248        builder.set_subject_name(&name).unwrap();
249        builder.set_issuer_name(&name).unwrap();
250        builder.set_pubkey(&key).unwrap();
251        builder
252            .set_not_before(&Asn1Time::days_from_now(0).unwrap())
253            .unwrap();
254        builder
255            .set_not_after(&Asn1Time::days_from_now(1).unwrap())
256            .unwrap();
257        builder.sign(&key, MessageDigest::sha256()).unwrap();
258        let cert = builder.build();
259
260        (
261            String::from_utf8(cert.to_pem().unwrap()).unwrap(),
262            String::from_utf8(key.private_key_to_pem_pkcs8().unwrap()).unwrap(),
263        )
264    }
265}