vector_core/tls/
reload.rs1use std::sync::{Arc, Weak};
11
12use arc_swap::ArcSwap;
13use openssl::ssl::SslAcceptor;
14
15use super::{MaybeTlsSettings, TlsSettings};
16
17#[derive(Clone)]
22pub struct TlsAcceptorReloader {
23 acceptor: Arc<ArcSwap<SslAcceptor>>,
24}
25
26impl TlsAcceptorReloader {
27 pub(super) fn new(acceptor: SslAcceptor) -> Self {
29 Self {
30 acceptor: Arc::new(ArcSwap::from_pointee(acceptor)),
31 }
32 }
33
34 pub(super) fn shared(&self) -> Arc<ArcSwap<SslAcceptor>> {
36 Arc::clone(&self.acceptor)
37 }
38
39 pub fn reload(&self, settings: &TlsSettings) -> crate::tls::Result<()> {
42 self.acceptor.store(Arc::new(settings.acceptor()?));
43 Ok(())
44 }
45
46 pub fn downgrade(&self) -> WeakTlsAcceptorReloader {
48 WeakTlsAcceptorReloader {
49 acceptor: Arc::downgrade(&self.acceptor),
50 }
51 }
52}
53
54#[derive(Clone)]
56pub struct WeakTlsAcceptorReloader {
57 acceptor: Weak<ArcSwap<SslAcceptor>>,
58}
59
60impl WeakTlsAcceptorReloader {
61 pub fn upgrade(&self) -> Option<TlsAcceptorReloader> {
63 self.acceptor
64 .upgrade()
65 .map(|acceptor| TlsAcceptorReloader { acceptor })
66 }
67}
68
69impl MaybeTlsSettings {
70 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 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 drop(listener);
140 assert!(
141 weak.upgrade().is_none(),
142 "weak handle must not upgrade after the listener is dropped"
143 );
144 }
145
146 #[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 let server = tokio::spawn(async move {
169 while let Ok(mut stream) = listener.accept().await {
170 stream.handshake().await.ok();
172 }
173 });
174
175 assert_eq!(served_common_name(local_addr).await, "old.example");
177
178 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 assert_eq!(served_common_name(local_addr).await, "new.example");
188
189 server.abort();
190 }
191
192 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 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 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}