Skip to main content

laminar_core/cluster/control/
tls.rs

1//! Process-wide transport mode for the cluster control plane (barrier and
2//! shuffle), resolved once at startup with at most one TLS identity.
3
4use 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
15/// TLS material shared by every control-plane server and client in this process.
16pub struct ClusterTls {
17    server: ServerTlsConfig,
18    client: ClientTlsConfig,
19    fingerprint: [u8; 32],
20}
21
22impl ClusterTls {
23    /// Build mTLS configs from PEM: this node's `cert`+`key`, the `ca` that signed
24    /// every peer cert, and the `server_name` SAN peers are verified against.
25    #[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
191/// Install process-wide control-plane TLS before any cluster transport is constructed.
192///
193/// Reinstalling byte-identical material is idempotent. Different material, or installation after
194/// plaintext was explicitly selected or used by a transport, fails closed.
195///
196/// # Errors
197/// Returns an error when transport mode is already plaintext or different TLS material is active.
198pub fn set_cluster_tls(tls: ClusterTls) -> Result<(), String> {
199    CLUSTER_TLS.install(tls)
200}
201
202/// Select process-wide plaintext before any cluster transport is constructed.
203///
204/// Repeated plaintext claims are idempotent. A claim fails once TLS installation begins.
205///
206/// # Errors
207/// Returns an error when TLS installation has already begun or completed.
208pub fn claim_cluster_plaintext() -> Result<(), String> {
209    CLUSTER_TLS.claim_plaintext()
210}
211
212/// Server config for the shared control-plane / shuffle listeners.
213pub(crate) fn server_tls() -> Option<&'static ServerTlsConfig> {
214    CLUSTER_TLS.transport_tls().map(|tls| &tls.server)
215}
216
217/// Client endpoint for `host_port`, using control-plane TLS + `https` when
218/// installed, plaintext `http` otherwise.
219pub(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    // HTTP/2 keepalive + connect timeout so a half-open conn (peer machine gone with no
223    // TCP RST — kill-9 on loopback RST-closes promptly, but a true network/host failure
224    // does not) flips dead within ~6s instead of blocking sends until the OS TCP timeout
225    // (minutes). Sub-`align_shuffle_barriers` ALIGN_TIMEOUT (8s) so the driver errors and
226    // the next send reconnects before alignment gives up.
227    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}