laminar_core/cluster/control/
tls.rs1use std::sync::atomic::{AtomicU8, Ordering};
5use std::sync::OnceLock;
6
7use sha2::{Digest as _, Sha256};
8use tonic::transport::{Certificate, ClientTlsConfig, Endpoint, Identity, ServerTlsConfig};
9
10const TRANSPORT_UNUSED: u8 = 0;
11const TRANSPORT_PLAINTEXT: u8 = 1;
12const TRANSPORT_TLS_INSTALLING: u8 = 2;
13const TRANSPORT_TLS: u8 = 3;
14
15pub struct ClusterTls {
17 server: ServerTlsConfig,
18 client: ClientTlsConfig,
19 fingerprint: [u8; 32],
20}
21
22impl ClusterTls {
23 #[must_use]
26 pub fn from_pem(cert: &[u8], key: &[u8], ca: &[u8], server_name: &str) -> Self {
27 let _ = tokio_rustls::rustls::crypto::aws_lc_rs::default_provider().install_default();
28 let fingerprint = tls_material_fingerprint(cert, key, ca, server_name);
29 let identity = Identity::from_pem(cert, key);
30 let ca = Certificate::from_pem(ca);
31 let server = ServerTlsConfig::new()
32 .identity(identity.clone())
33 .client_ca_root(ca.clone());
34 let client = ClientTlsConfig::new()
35 .ca_certificate(ca)
36 .identity(identity)
37 .domain_name(server_name.to_string());
38 Self {
39 server,
40 client,
41 fingerprint,
42 }
43 }
44}
45
46fn tls_material_fingerprint(cert: &[u8], key: &[u8], ca: &[u8], server_name: &str) -> [u8; 32] {
47 let mut digest = Sha256::new();
48 digest.update(b"laminardb-cluster-tls-v1");
49 for field in [cert, key, ca, server_name.as_bytes()] {
50 let length = u64::try_from(field.len()).expect("TLS material length fits u64");
51 digest.update(length.to_be_bytes());
52 digest.update(field);
53 }
54 digest.finalize().into()
55}
56
57struct ClusterTlsState {
58 mode: AtomicU8,
59 tls: OnceLock<ClusterTls>,
60}
61
62impl ClusterTlsState {
63 const fn new() -> Self {
64 Self {
65 mode: AtomicU8::new(TRANSPORT_UNUSED),
66 tls: OnceLock::new(),
67 }
68 }
69
70 fn install(&self, tls: ClusterTls) -> Result<(), String> {
71 loop {
72 match self.mode.load(Ordering::Acquire) {
73 TRANSPORT_UNUSED => {
74 if self
75 .mode
76 .compare_exchange(
77 TRANSPORT_UNUSED,
78 TRANSPORT_TLS_INSTALLING,
79 Ordering::AcqRel,
80 Ordering::Acquire,
81 )
82 .is_err()
83 {
84 continue;
85 }
86 let fingerprint = tls.fingerprint;
87 let result = match self.tls.set(tls) {
88 Ok(()) => Ok(()),
89 Err(_)
90 if self
91 .tls
92 .get()
93 .is_some_and(|tls| tls.fingerprint == fingerprint) =>
94 {
95 Ok(())
96 }
97 Err(_) => {
98 Err("cluster TLS state already contains different material".to_string())
99 }
100 };
101 self.mode.store(TRANSPORT_TLS, Ordering::Release);
102 return result;
103 }
104 TRANSPORT_PLAINTEXT => {
105 return Err(
106 "cluster TLS cannot be installed after plaintext was selected".into(),
107 );
108 }
109 TRANSPORT_TLS_INSTALLING => std::hint::spin_loop(),
110 TRANSPORT_TLS => {
111 return if self
112 .tls
113 .get()
114 .is_some_and(|installed| installed.fingerprint == tls.fingerprint)
115 {
116 Ok(())
117 } else {
118 Err("different cluster TLS material is already installed".into())
119 };
120 }
121 _ => unreachable!("cluster transport mode is internal"),
122 }
123 }
124 }
125
126 fn claim_plaintext(&self) -> Result<(), String> {
127 loop {
128 match self.mode.load(Ordering::Acquire) {
129 TRANSPORT_UNUSED => {
130 if self
131 .mode
132 .compare_exchange(
133 TRANSPORT_UNUSED,
134 TRANSPORT_PLAINTEXT,
135 Ordering::AcqRel,
136 Ordering::Acquire,
137 )
138 .is_err()
139 {
140 continue;
141 }
142 return Ok(());
143 }
144 TRANSPORT_PLAINTEXT => return Ok(()),
145 TRANSPORT_TLS_INSTALLING | TRANSPORT_TLS => {
146 return Err(
147 "cluster plaintext cannot be claimed after TLS installation has begun"
148 .into(),
149 );
150 }
151 _ => unreachable!("cluster transport mode is internal"),
152 }
153 }
154 }
155
156 fn transport_tls(&self) -> Option<&ClusterTls> {
157 loop {
158 match self.mode.load(Ordering::Acquire) {
159 TRANSPORT_UNUSED => {
160 if self
161 .mode
162 .compare_exchange(
163 TRANSPORT_UNUSED,
164 TRANSPORT_PLAINTEXT,
165 Ordering::AcqRel,
166 Ordering::Acquire,
167 )
168 .is_err()
169 {
170 continue;
171 }
172 return None;
173 }
174 TRANSPORT_PLAINTEXT => return None,
175 TRANSPORT_TLS_INSTALLING => std::hint::spin_loop(),
176 TRANSPORT_TLS => {
177 return Some(
178 self.tls
179 .get()
180 .expect("TLS mode is published only after its material"),
181 );
182 }
183 _ => unreachable!("cluster transport mode is internal"),
184 }
185 }
186 }
187}
188
189static CLUSTER_TLS: ClusterTlsState = ClusterTlsState::new();
190
191pub fn set_cluster_tls(tls: ClusterTls) -> Result<(), String> {
199 CLUSTER_TLS.install(tls)
200}
201
202pub fn claim_cluster_plaintext() -> Result<(), String> {
209 CLUSTER_TLS.claim_plaintext()
210}
211
212pub(crate) fn server_tls() -> Option<&'static ServerTlsConfig> {
214 CLUSTER_TLS.transport_tls().map(|tls| &tls.server)
215}
216
217pub(crate) fn client_endpoint(host_port: &str) -> Result<Endpoint, String> {
220 let tls = CLUSTER_TLS.transport_tls();
221 let scheme = if tls.is_some() { "https" } else { "http" };
222 let endpoint = Endpoint::from_shared(format!("{scheme}://{host_port}"))
228 .map_err(|e| e.to_string())?
229 .connect_timeout(std::time::Duration::from_secs(3))
230 .http2_keep_alive_interval(std::time::Duration::from_secs(3))
231 .keep_alive_timeout(std::time::Duration::from_secs(3))
232 .keep_alive_while_idle(true);
233 match tls {
234 Some(t) => endpoint
235 .tls_config(t.client.clone())
236 .map_err(|e| e.to_string()),
237 None => Ok(endpoint),
238 }
239}
240
241#[cfg(test)]
242mod tests {
243 use super::*;
244
245 fn tls(seed: u8) -> ClusterTls {
246 ClusterTls::from_pem(
247 &[b'c', seed],
248 &[b'k', seed],
249 &[b'a', seed],
250 &format!("cluster-{seed}"),
251 )
252 }
253
254 #[test]
255 fn install_before_use_freezes_tls_and_identical_repeat_is_idempotent() {
256 let state = ClusterTlsState::new();
257 state.install(tls(1)).unwrap();
258 assert!(state.transport_tls().is_some());
259 state.install(tls(1)).unwrap();
260 }
261
262 #[test]
263 fn conflicting_repeat_is_rejected() {
264 let state = ClusterTlsState::new();
265 state.install(tls(1)).unwrap();
266 let error = state.install(tls(2)).unwrap_err();
267 assert!(error.contains("different cluster TLS material"), "{error}");
268 assert!(state.transport_tls().is_some());
269 }
270
271 #[test]
272 fn plaintext_claim_is_idempotent_and_rejects_late_tls_install() {
273 let state = ClusterTlsState::new();
274 state.claim_plaintext().unwrap();
275 state.claim_plaintext().unwrap();
276 assert!(state.transport_tls().is_none());
277 let error = state.install(tls(1)).unwrap_err();
278 assert!(error.contains("after plaintext was selected"), "{error}");
279 assert!(state.transport_tls().is_none());
280 }
281
282 #[test]
283 fn plaintext_claim_rejects_tls_installing_or_installed() {
284 let installing = ClusterTlsState::new();
285 installing
286 .mode
287 .store(TRANSPORT_TLS_INSTALLING, Ordering::Release);
288 let error = installing.claim_plaintext().unwrap_err();
289 assert!(error.contains("TLS installation has begun"), "{error}");
290
291 let installed = ClusterTlsState::new();
292 installed.install(tls(1)).unwrap();
293 let error = installed.claim_plaintext().unwrap_err();
294 assert!(error.contains("TLS installation has begun"), "{error}");
295 assert!(installed.transport_tls().is_some());
296 }
297
298 #[test]
299 fn plaintext_transport_use_rejects_late_tls_install() {
300 let state = ClusterTlsState::new();
301 assert!(state.transport_tls().is_none());
302 let error = state.install(tls(1)).unwrap_err();
303 assert!(error.contains("after plaintext was selected"), "{error}");
304 }
305
306 #[test]
307 fn material_fingerprint_is_length_framed() {
308 let left = tls_material_fingerprint(b"ab", b"c", b"d", "e");
309 let right = tls_material_fingerprint(b"a", b"bc", b"d", "e");
310 assert_ne!(left, right);
311 }
312}