1use std::collections::{BTreeSet, HashMap};
56use std::fmt;
57use std::future::Future;
58use std::io;
59use std::net::SocketAddr;
60use std::pin::Pin;
61use std::sync::{Arc, Mutex};
62use std::time::Duration;
63
64use mongreldb_consensus::identity::{MongrelRaft, MongrelRaftConfig, RaftNodeId};
65use mongreldb_consensus::network::{AppendRpcError, RaftTransport, SnapshotRpcError, VoteRpcError};
66use mongreldb_protocol::envelope::{
67 EnvelopeError, ProtocolEnvelope, CHECKSUM_LEN, HEADER_LEN, MAX_MESSAGE_PAYLOAD_BYTES,
68};
69use mongreldb_types::ids::NodeId;
70use openraft::error::{NetworkError, RPCError, Timeout, Unreachable};
71use openraft::raft::{
72 AppendEntriesRequest, AppendEntriesResponse, InstallSnapshotRequest, InstallSnapshotResponse,
73 VoteRequest, VoteResponse,
74};
75use openraft::RPCTypes;
76use rustls::pki_types::{CertificateDer, ServerName};
77use serde::{Deserialize, Serialize};
78use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
79use tokio::net::{TcpListener, TcpStream};
80use tokio::sync::{watch, OwnedSemaphorePermit, Semaphore};
81use tokio::task::{JoinHandle, JoinSet};
82
83pub const RAFT_MSG_ERROR: u32 = 0;
87pub const RAFT_MSG_APPEND_ENTRIES_REQUEST: u32 = 1;
89pub const RAFT_MSG_APPEND_ENTRIES_RESPONSE: u32 = 2;
91pub const RAFT_MSG_VOTE_REQUEST: u32 = 3;
93pub const RAFT_MSG_VOTE_RESPONSE: u32 = 4;
95pub const RAFT_MSG_INSTALL_SNAPSHOT_REQUEST: u32 = 5;
97pub const RAFT_MSG_INSTALL_SNAPSHOT_RESPONSE: u32 = 6;
99pub const RAFT_MSG_TRIGGER_ELECTION_REQUEST: u32 = 7;
101pub const RAFT_MSG_TRIGGER_ELECTION_RESPONSE: u32 = 8;
103pub const INTERNAL_MSG_REQUEST: u32 = 100;
106pub const INTERNAL_MSG_RESPONSE: u32 = 101;
108
109pub const NODE_CERT_DOMAIN: &str = "mongreldb.cluster";
111
112#[derive(Debug, thiserror::Error)]
114pub enum TransportError {
115 #[error("no route to node {0}")]
117 NoRoute(RaftNodeId),
118 #[error("transport I/O error: {0}")]
120 Io(#[from] io::Error),
121 #[error("operation timed out after {0:?}")]
123 Timeout(Duration),
124 #[error("TLS error: {0}")]
126 Tls(String),
127 #[error("protocol envelope error: {0}")]
130 Envelope(#[from] EnvelopeError),
131 #[error("frame payload of {actual} bytes exceeds the {limit} byte bound")]
134 FrameTooLarge {
135 actual: usize,
137 limit: usize,
139 },
140 #[error("peer authentication failed: {0}")]
142 PeerAuthentication(String),
143 #[error("protocol violation: {0}")]
146 ProtocolViolation(String),
147 #[error("remote node error: {0}")]
150 Remote(String),
151}
152
153#[derive(Clone, Debug)]
155pub struct TransportConfig {
156 pub connect_timeout: Duration,
158 pub rpc_timeout: Duration,
161 pub snapshot_timeout: Duration,
164 pub connect_attempts: usize,
167 pub reconnect_backoff: Duration,
169 pub max_frame_bytes: usize,
173 pub max_connections: usize,
176 pub handshake_timeout: Duration,
178 pub shutdown_grace: Duration,
181}
182
183impl Default for TransportConfig {
184 fn default() -> Self {
185 Self {
186 connect_timeout: Duration::from_millis(2_000),
187 rpc_timeout: Duration::from_millis(1_000),
188 snapshot_timeout: Duration::from_millis(10_000),
189 connect_attempts: 3,
190 reconnect_backoff: Duration::from_millis(25),
191 max_frame_bytes: 16 * 1024 * 1024,
192 max_connections: 256,
193 handshake_timeout: Duration::from_millis(3_000),
194 shutdown_grace: Duration::from_millis(5_000),
195 }
196 }
197}
198
199#[derive(Clone, Debug, PartialEq, Eq)]
201pub struct PeerEndpoint {
202 pub address: String,
204 pub tls_node_id: Option<NodeId>,
208}
209
210impl PeerEndpoint {
211 pub fn plaintext(address: impl Into<String>) -> Self {
213 Self {
214 address: address.into(),
215 tls_node_id: None,
216 }
217 }
218
219 pub fn mtls(address: impl Into<String>, node_id: NodeId) -> Self {
221 Self {
222 address: address.into(),
223 tls_node_id: Some(node_id),
224 }
225 }
226}
227
228type InternalHandlerMap = HashMap<(RaftNodeId, u32), Arc<dyn InternalRpcHandler>>;
233
234#[derive(Clone, Default)]
235pub struct TransportRegistry {
236 nodes: Arc<Mutex<HashMap<RaftNodeId, MongrelRaft>>>,
237 internal_handlers: Arc<Mutex<InternalHandlerMap>>,
238}
239
240pub type InternalRpcFuture<'a> = Pin<Box<dyn Future<Output = Result<Vec<u8>, String>> + Send + 'a>>;
242
243pub trait InternalRpcHandler: Send + Sync {
249 fn handle<'a>(&'a self, body: &'a [u8]) -> InternalRpcFuture<'a>;
251}
252
253impl TransportRegistry {
254 pub fn new() -> Self {
256 Self::default()
257 }
258
259 pub fn attach(&self, node_id: RaftNodeId, raft: MongrelRaft) {
261 self.lock().insert(node_id, raft);
262 }
263
264 pub fn detach(&self, node_id: RaftNodeId) {
266 self.lock().remove(&node_id);
267 }
268
269 pub fn get(&self, node_id: RaftNodeId) -> Option<MongrelRaft> {
271 self.lock().get(&node_id).cloned()
272 }
273
274 pub fn attach_internal(
276 &self,
277 node_id: RaftNodeId,
278 service_id: u32,
279 handler: Arc<dyn InternalRpcHandler>,
280 ) {
281 self.internal_handlers
282 .lock()
283 .expect("internal handler registry lock poisoned")
284 .insert((node_id, service_id), handler);
285 }
286
287 pub fn detach_internal(&self, node_id: RaftNodeId, service_id: u32) {
289 self.internal_handlers
290 .lock()
291 .expect("internal handler registry lock poisoned")
292 .remove(&(node_id, service_id));
293 }
294
295 pub fn internal_handler(
297 &self,
298 node_id: RaftNodeId,
299 service_id: u32,
300 ) -> Option<Arc<dyn InternalRpcHandler>> {
301 self.internal_handlers
302 .lock()
303 .expect("internal handler registry lock poisoned")
304 .get(&(node_id, service_id))
305 .cloned()
306 }
307
308 pub fn len(&self) -> usize {
310 self.lock().len()
311 }
312
313 pub fn is_empty(&self) -> bool {
315 self.len() == 0
316 }
317
318 fn lock(&self) -> std::sync::MutexGuard<'_, HashMap<RaftNodeId, MongrelRaft>> {
319 self.nodes.lock().expect("transport registry lock poisoned")
320 }
321}
322
323impl fmt::Debug for TransportRegistry {
324 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
325 f.debug_struct("TransportRegistry")
326 .field("attached", &self.len())
327 .finish()
328 }
329}
330
331pub fn node_cert_name(node_id: &NodeId) -> String {
339 let mut hex = String::with_capacity(32);
340 for byte in node_id.as_bytes() {
341 hex.push(char::from_digit(u32::from(byte >> 4), 16).expect("hex digit"));
342 hex.push(char::from_digit(u32::from(byte & 0x0f), 16).expect("hex digit"));
343 }
344 format!("node-{hex}.{NODE_CERT_DOMAIN}")
345}
346
347#[derive(Clone)]
355pub struct TlsConfig {
356 client_config: Arc<rustls::ClientConfig>,
357 server_config: Arc<rustls::ServerConfig>,
358 allowed_identities: Arc<BTreeSet<String>>,
359}
360
361impl TlsConfig {
362 pub fn from_pems(
366 ca_cert_pem: &str,
367 node_cert_pem: &str,
368 node_key_pem: &str,
369 allowed_node_ids: &[NodeId],
370 ) -> Result<Self, TransportError> {
371 let ca_certs = parse_pem_certificates("ca_cert_pem", ca_cert_pem)?;
372 let mut roots = rustls::RootCertStore::empty();
373 for cert in ca_certs {
374 roots.add(cert).map_err(|e| {
375 TransportError::Tls(format!("cluster CA certificate rejected: {e}"))
376 })?;
377 }
378 let roots = Arc::new(roots);
379 let node_certs = parse_pem_certificates("node_cert_pem", node_cert_pem)?;
380 let key = rustls_pemfile::private_key(&mut io::BufReader::new(node_key_pem.as_bytes()))
381 .map_err(|e| TransportError::Tls(format!("node_key_pem unreadable: {e}")))?
382 .ok_or(TransportError::Tls(
383 "node_key_pem contains no private key".to_owned(),
384 ))?;
385 if allowed_node_ids.is_empty() {
386 return Err(TransportError::Tls(
387 "allowed_node_ids is empty; an empty list admits no peers".to_owned(),
388 ));
389 }
390 let allowed_identities = Arc::new(
391 allowed_node_ids
392 .iter()
393 .map(node_cert_name)
394 .collect::<BTreeSet<_>>(),
395 );
396
397 let provider = Arc::new(rustls::crypto::ring::default_provider());
400 let client_verifier = rustls::server::WebPkiClientVerifier::builder_with_provider(
401 roots.clone(),
402 provider.clone(),
403 )
404 .build()
405 .map_err(|e| TransportError::Tls(format!("client-certificate verifier: {e}")))?;
406 let server_config = rustls::ServerConfig::builder_with_provider(provider.clone())
407 .with_protocol_versions(&[&rustls::version::TLS13])
408 .map_err(|e| TransportError::Tls(format!("TLS 1.3 server configuration: {e}")))?
409 .with_client_cert_verifier(client_verifier)
410 .with_single_cert(node_certs.clone(), key.clone_key())
411 .map_err(|e| TransportError::Tls(format!("node certificate/key rejected: {e}")))?;
412 let client_config = rustls::ClientConfig::builder_with_provider(provider)
413 .with_protocol_versions(&[&rustls::version::TLS13])
414 .map_err(|e| TransportError::Tls(format!("TLS 1.3 client configuration: {e}")))?
415 .with_root_certificates(roots)
416 .with_client_auth_cert(node_certs, key)
417 .map_err(|e| TransportError::Tls(format!("node certificate/key rejected: {e}")))?;
418 Ok(Self {
419 client_config: Arc::new(client_config),
420 server_config: Arc::new(server_config),
421 allowed_identities,
422 })
423 }
424
425 pub fn from_trust(trust: &crate::bootstrap::TrustConfig) -> Result<Self, TransportError> {
427 Self::from_pems(
428 &trust.ca_cert_pem,
429 &trust.node_cert_pem,
430 &trust.node_key_pem,
431 &trust.allowed_node_ids,
432 )
433 }
434
435 pub fn client_config(&self) -> Arc<rustls::ClientConfig> {
437 self.client_config.clone()
438 }
439
440 pub fn server_config(&self) -> Arc<rustls::ServerConfig> {
442 self.server_config.clone()
443 }
444
445 pub fn allows_identity(&self, identity: &str) -> bool {
448 self.allowed_identities.contains(identity)
449 }
450}
451
452impl fmt::Debug for TlsConfig {
453 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
454 f.debug_struct("TlsConfig")
455 .field("allowed_identities", &self.allowed_identities)
456 .finish_non_exhaustive()
457 }
458}
459
460fn parse_pem_certificates(
463 field: &'static str,
464 pem: &str,
465) -> Result<Vec<CertificateDer<'static>>, TransportError> {
466 let certs: Vec<CertificateDer<'static>> =
467 rustls_pemfile::certs(&mut io::BufReader::new(pem.as_bytes()))
468 .collect::<Result<_, _>>()
469 .map_err(|e| TransportError::Tls(format!("{field} unreadable: {e}")))?;
470 if certs.is_empty() {
471 return Err(TransportError::Tls(format!(
472 "{field} contains no certificates"
473 )));
474 }
475 Ok(certs)
476}
477
478#[derive(Clone, Debug)]
484pub enum TransportSecurity {
485 Mtls(TlsConfig),
488 PlaintextForTesting,
495}
496
497fn peer_server_name(peer: &PeerEndpoint) -> Result<ServerName<'static>, TransportError> {
499 let Some(node_id) = peer.tls_node_id else {
500 return Err(TransportError::PeerAuthentication(format!(
501 "mTLS endpoint {} has no node identity to authenticate against",
502 peer.address
503 )));
504 };
505 let name = node_cert_name(&node_id);
506 ServerName::try_from(name.clone())
507 .map_err(|e| TransportError::Tls(format!("invalid peer server name {name}: {e}")))
508}
509
510fn verify_peer_identity(
514 peer_certificates: Option<&[CertificateDer<'_>]>,
515 tls: &TlsConfig,
516) -> Result<(), TransportError> {
517 let Some(certificates) = peer_certificates else {
518 return Err(TransportError::PeerAuthentication(
519 "peer presented no certificate".to_owned(),
520 ));
521 };
522 let Some(leaf) = certificates.first() else {
523 return Err(TransportError::PeerAuthentication(
524 "peer presented an empty certificate chain".to_owned(),
525 ));
526 };
527 let identities = der::certificate_identities(leaf.as_ref()).map_err(|e| {
528 TransportError::PeerAuthentication(format!("peer certificate identity unreadable: {e}"))
529 })?;
530 if identities
531 .iter()
532 .any(|identity| tls.allows_identity(identity))
533 {
534 Ok(())
535 } else {
536 Err(TransportError::PeerAuthentication(format!(
537 "peer certificate identities {identities:?} are not admitted to this cluster"
538 )))
539 }
540}
541
542mod der {
551 struct Tlv<'a> {
553 tag: u8,
554 content: &'a [u8],
555 encoded_len: usize,
557 }
558
559 const TAG_SEQUENCE: u8 = 0x30;
560 const TAG_SET: u8 = 0x31;
561 const TAG_OCTET_STRING: u8 = 0x04;
562 const TAG_OID: u8 = 0x06;
563 const TAG_UTF8_STRING: u8 = 0x0c;
564 const TAG_PRINTABLE_STRING: u8 = 0x13;
565 const TAG_IA5_STRING: u8 = 0x16;
566 const TAG_EXTENSIONS: u8 = 0xa3;
568 const TAG_SAN_DNS_NAME: u8 = 0x82;
570
571 const OID_COMMON_NAME: &[u8] = &[0x55, 0x04, 0x03];
573 const OID_SUBJECT_ALT_NAME: &[u8] = &[0x55, 0x1d, 0x11];
575
576 fn read_tlv(buf: &[u8]) -> Result<Tlv<'_>, String> {
578 if buf.len() < 2 {
579 return Err("truncated TLV header".to_owned());
580 }
581 let tag = buf[0];
582 if tag & 0x1f == 0x1f {
583 return Err("multi-byte tags are outside the certificate subset".to_owned());
584 }
585 let first_len = buf[1];
586 let (content_len, header_len) = if first_len & 0x80 == 0 {
587 (usize::from(first_len), 2)
588 } else {
589 let len_bytes = usize::from(first_len & 0x7f);
590 if len_bytes == 0 || len_bytes > 4 {
593 return Err("invalid DER length".to_owned());
594 }
595 if buf.len() < 2 + len_bytes {
596 return Err("truncated DER length".to_owned());
597 }
598 let mut len = 0usize;
599 for &byte in &buf[2..2 + len_bytes] {
600 len = len
601 .checked_mul(256)
602 .and_then(|len| len.checked_add(usize::from(byte)))
603 .ok_or_else(|| "DER length overflow".to_owned())?;
604 }
605 (len, 2 + len_bytes)
606 };
607 if buf.len() < header_len + content_len {
608 return Err("truncated TLV content".to_owned());
609 }
610 Ok(Tlv {
611 tag,
612 content: &buf[header_len..header_len + content_len],
613 encoded_len: header_len + content_len,
614 })
615 }
616
617 fn expect_directory_string(tlv: &Tlv<'_>) -> Result<String, String> {
618 match tlv.tag {
619 TAG_UTF8_STRING | TAG_PRINTABLE_STRING | TAG_IA5_STRING => {
620 std::str::from_utf8(tlv.content)
621 .map(str::to_owned)
622 .map_err(|_| "non-UTF8 directory string".to_owned())
623 }
624 other => Err(format!("unexpected string tag 0x{other:02x}")),
625 }
626 }
627
628 pub fn certificate_identities(cert_der: &[u8]) -> Result<Vec<String>, String> {
631 let certificate = read_tlv(cert_der)?;
632 if certificate.tag != TAG_SEQUENCE || certificate.encoded_len != cert_der.len() {
633 return Err("certificate is not exactly one SEQUENCE".to_owned());
634 }
635 let tbs = read_tlv(certificate.content)?;
636 if tbs.tag != TAG_SEQUENCE {
637 return Err("tbsCertificate is not a SEQUENCE".to_owned());
638 }
639 let mut identities = Vec::new();
640 let mut rest = tbs.content;
645 let mut sequences_seen = 0u32;
646 while !rest.is_empty() {
647 let tlv = read_tlv(rest)?;
648 rest = &rest[tlv.encoded_len..];
649 match tlv.tag {
650 TAG_SEQUENCE => {
651 sequences_seen += 1;
652 if sequences_seen == 4 {
653 identities.extend(subject_common_names(tlv.content)?);
654 }
655 }
656 TAG_EXTENSIONS => {
657 identities.extend(extension_san_dns_names(tlv.content)?);
658 }
659 _ => {}
660 }
661 }
662 Ok(identities)
663 }
664
665 fn subject_common_names(rdn_sequence: &[u8]) -> Result<Vec<String>, String> {
667 let mut names = Vec::new();
668 let mut rest = rdn_sequence;
669 while !rest.is_empty() {
670 let rdn = read_tlv(rest)?;
671 rest = &rest[rdn.encoded_len..];
672 if rdn.tag != TAG_SET {
673 return Err("RDN is not a SET".to_owned());
674 }
675 let mut rdn_rest = rdn.content;
676 while !rdn_rest.is_empty() {
677 let attribute = read_tlv(rdn_rest)?;
678 rdn_rest = &rdn_rest[attribute.encoded_len..];
679 if attribute.tag != TAG_SEQUENCE {
680 return Err("AttributeTypeAndValue is not a SEQUENCE".to_owned());
681 }
682 let oid = read_tlv(attribute.content)?;
683 if oid.tag != TAG_OID {
684 return Err("attribute type is not an OID".to_owned());
685 }
686 if oid.content == OID_COMMON_NAME {
687 let value = read_tlv(&attribute.content[oid.encoded_len..])?;
688 names.push(expect_directory_string(&value)?);
689 }
690 }
691 }
692 Ok(names)
693 }
694
695 fn extension_san_dns_names(explicit_content: &[u8]) -> Result<Vec<String>, String> {
698 let extensions = read_tlv(explicit_content)?;
699 if extensions.tag != TAG_SEQUENCE || extensions.encoded_len != explicit_content.len() {
700 return Err("malformed extensions".to_owned());
701 }
702 let mut names = Vec::new();
703 let mut rest = extensions.content;
704 while !rest.is_empty() {
705 let extension = read_tlv(rest)?;
706 rest = &rest[extension.encoded_len..];
707 if extension.tag != TAG_SEQUENCE {
708 return Err("extension is not a SEQUENCE".to_owned());
709 }
710 let oid = read_tlv(extension.content)?;
711 if oid.tag != TAG_OID {
712 return Err("extension id is not an OID".to_owned());
713 }
714 if oid.content != OID_SUBJECT_ALT_NAME {
715 continue;
716 }
717 let mut ext_rest = &extension.content[oid.encoded_len..];
720 let mut san_bytes = None;
721 while !ext_rest.is_empty() {
722 let child = read_tlv(ext_rest)?;
723 ext_rest = &ext_rest[child.encoded_len..];
724 if child.tag == TAG_OCTET_STRING {
725 san_bytes = Some(child.content);
726 }
727 }
728 let Some(san_bytes) = san_bytes else {
729 return Err("subjectAltName extension has no extnValue".to_owned());
730 };
731 let general_names = read_tlv(san_bytes)?;
732 if general_names.tag != TAG_SEQUENCE || general_names.encoded_len != san_bytes.len() {
733 return Err("malformed GeneralNames".to_owned());
734 }
735 let mut names_rest = general_names.content;
736 while !names_rest.is_empty() {
737 let name = read_tlv(names_rest)?;
738 names_rest = &names_rest[name.encoded_len..];
739 if name.tag == TAG_SAN_DNS_NAME {
740 names.push(
741 std::str::from_utf8(name.content)
742 .map(str::to_owned)
743 .map_err(|_| "non-UTF8 dNSName".to_owned())?,
744 );
745 }
746 }
747 }
748 Ok(names)
749 }
750}
751
752#[derive(Debug, Serialize, Deserialize)]
760struct RpcPayload<T> {
761 target: RaftNodeId,
762 rpc: T,
763}
764
765#[derive(Debug, Serialize, Deserialize)]
766struct InternalRpcPayload {
767 target: RaftNodeId,
768 service_id: u32,
769 body: Vec<u8>,
770}
771
772async fn write_frame<S: AsyncWrite + Unpin>(
774 stream: &mut S,
775 message_type: u32,
776 payload: Vec<u8>,
777 timeout: Duration,
778) -> Result<(), TransportError> {
779 match tokio::time::timeout(timeout, write_frame_inner(stream, message_type, payload)).await {
780 Ok(result) => result,
781 Err(_) => Err(TransportError::Timeout(timeout)),
782 }
783}
784
785async fn write_frame_inner<S: AsyncWrite + Unpin>(
786 stream: &mut S,
787 message_type: u32,
788 payload: Vec<u8>,
789) -> Result<(), TransportError> {
790 if payload.len() > MAX_MESSAGE_PAYLOAD_BYTES {
791 return Err(TransportError::FrameTooLarge {
792 actual: payload.len(),
793 limit: MAX_MESSAGE_PAYLOAD_BYTES,
794 });
795 }
796 let envelope = ProtocolEnvelope::new(message_type, payload);
797 stream.write_all(&envelope.encode()).await?;
798 stream.flush().await?;
799 Ok(())
800}
801
802async fn read_frame<S: AsyncRead + Unpin>(
807 stream: &mut S,
808 max_frame_bytes: usize,
809 timeout: Duration,
810) -> Result<Option<ProtocolEnvelope>, TransportError> {
811 match tokio::time::timeout(timeout, read_frame_inner(stream, max_frame_bytes)).await {
812 Ok(result) => result,
813 Err(_) => Err(TransportError::Timeout(timeout)),
814 }
815}
816
817async fn read_frame_inner<S: AsyncRead + Unpin>(
818 stream: &mut S,
819 max_frame_bytes: usize,
820) -> Result<Option<ProtocolEnvelope>, TransportError> {
821 let mut header = [0u8; HEADER_LEN];
822 let first = stream.read(&mut header[..1]).await?;
824 if first == 0 {
825 return Ok(None);
826 }
827 stream.read_exact(&mut header[1..]).await?;
828 let payload_len = u32::from_le_bytes(header[8..12].try_into().expect("header slice")) as usize;
829 let bound = max_frame_bytes.min(MAX_MESSAGE_PAYLOAD_BYTES);
830 if payload_len > bound {
831 return Err(TransportError::FrameTooLarge {
832 actual: payload_len,
833 limit: bound,
834 });
835 }
836 let mut frame = Vec::with_capacity(HEADER_LEN + payload_len + CHECKSUM_LEN);
837 frame.extend_from_slice(&header);
838 frame.resize(HEADER_LEN + payload_len + CHECKSUM_LEN, 0);
839 stream.read_exact(&mut frame[HEADER_LEN..]).await?;
840 Ok(Some(ProtocolEnvelope::decode(&frame)?))
841}
842
843pub struct TcpTransport {
856 config: TransportConfig,
857 security: Arc<std::sync::RwLock<TransportSecurity>>,
860 peers: Arc<Mutex<HashMap<RaftNodeId, PeerEndpoint>>>,
861 registry: TransportRegistry,
862}
863
864impl TcpTransport {
865 pub fn new(config: TransportConfig, security: TransportSecurity) -> Self {
867 Self {
868 config,
869 security: Arc::new(std::sync::RwLock::new(security)),
870 peers: Arc::new(Mutex::new(HashMap::new())),
871 registry: TransportRegistry::new(),
872 }
873 }
874
875 pub fn reload_security(&self, security: TransportSecurity) {
878 *self
879 .security
880 .write()
881 .expect("transport security lock poisoned") = security;
882 }
883
884 pub fn security(&self) -> TransportSecurity {
886 self.security
887 .read()
888 .expect("transport security lock poisoned")
889 .clone()
890 }
891
892 pub fn security_handle(&self) -> Arc<std::sync::RwLock<TransportSecurity>> {
894 Arc::clone(&self.security)
895 }
896
897 pub fn registry(&self) -> TransportRegistry {
900 self.registry.clone()
901 }
902
903 pub fn config(&self) -> &TransportConfig {
905 &self.config
906 }
907
908 pub fn upsert_peer(&self, node_id: RaftNodeId, endpoint: PeerEndpoint) {
910 self.peers
911 .lock()
912 .expect("peer directory lock poisoned")
913 .insert(node_id, endpoint);
914 }
915
916 pub fn remove_peer(&self, node_id: RaftNodeId) -> Option<PeerEndpoint> {
919 self.peers
920 .lock()
921 .expect("peer directory lock poisoned")
922 .remove(&node_id)
923 }
924
925 pub fn peer(&self, node_id: RaftNodeId) -> Option<PeerEndpoint> {
927 self.peers
928 .lock()
929 .expect("peer directory lock poisoned")
930 .get(&node_id)
931 .cloned()
932 }
933
934 pub async fn internal_rpc(
941 &self,
942 target: RaftNodeId,
943 service_id: u32,
944 body: Vec<u8>,
945 ) -> Result<Vec<u8>, TransportError> {
946 let payload = serde_json::to_vec(&InternalRpcPayload {
947 target,
948 service_id,
949 body,
950 })
951 .map_err(|error| {
952 TransportError::ProtocolViolation(format!("unencodable internal request: {error}"))
953 })?;
954 self.round_trip(
955 target,
956 INTERNAL_MSG_REQUEST,
957 INTERNAL_MSG_RESPONSE,
958 payload,
959 self.config.rpc_timeout,
960 )
961 .await
962 }
963
964 async fn round_trip(
967 &self,
968 target: RaftNodeId,
969 request_type: u32,
970 expected_response_type: u32,
971 payload: Vec<u8>,
972 timeout: Duration,
973 ) -> Result<Vec<u8>, TransportError> {
974 let peer = self.peer(target).ok_or(TransportError::NoRoute(target))?;
975 let attempts = self.config.connect_attempts.max(1);
976 let mut attempt = 0usize;
977 let mut stream = loop {
978 match self.connect_once(&peer).await {
979 Ok(stream) => break stream,
980 Err(error) => {
981 attempt += 1;
982 if attempt >= attempts || !matches!(error, TransportError::Io(_)) {
985 return Err(error);
986 }
987 let shift = u32::try_from(attempt - 1).unwrap_or(u32::MAX).min(10);
988 tokio::time::sleep(self.config.reconnect_backoff * (1u32 << shift)).await;
989 }
990 }
991 };
992 match &mut stream {
993 ClientStream::Plain(stream) => {
994 self.round_trip_on(
995 stream,
996 request_type,
997 expected_response_type,
998 payload,
999 timeout,
1000 )
1001 .await
1002 }
1003 ClientStream::Tls(stream) => {
1004 self.round_trip_on(
1005 stream,
1006 request_type,
1007 expected_response_type,
1008 payload,
1009 timeout,
1010 )
1011 .await
1012 }
1013 }
1014 }
1015
1016 async fn round_trip_on<S: AsyncRead + AsyncWrite + Unpin>(
1017 &self,
1018 stream: &mut S,
1019 request_type: u32,
1020 expected_response_type: u32,
1021 payload: Vec<u8>,
1022 timeout: Duration,
1023 ) -> Result<Vec<u8>, TransportError> {
1024 write_frame(stream, request_type, payload, timeout).await?;
1025 let Some(frame) = read_frame(stream, self.config.max_frame_bytes, timeout).await? else {
1026 return Err(TransportError::Io(io::Error::new(
1027 io::ErrorKind::UnexpectedEof,
1028 "peer closed the connection without answering",
1029 )));
1030 };
1031 if frame.message_type == RAFT_MSG_ERROR {
1032 let message = serde_json::from_slice(&frame.payload)
1033 .unwrap_or_else(|_| "remote dispatch failure".to_owned());
1034 return Err(TransportError::Remote(message));
1035 }
1036 if frame.message_type != expected_response_type {
1037 return Err(TransportError::ProtocolViolation(format!(
1038 "expected response type {expected_response_type}, got {}",
1039 frame.message_type
1040 )));
1041 }
1042 Ok(frame.payload)
1043 }
1044
1045 async fn connect_once(&self, peer: &PeerEndpoint) -> Result<ClientStream, TransportError> {
1048 let tcp = match tokio::time::timeout(
1049 self.config.connect_timeout,
1050 TcpStream::connect(&peer.address),
1051 )
1052 .await
1053 {
1054 Ok(Ok(tcp)) => tcp,
1055 Ok(Err(error)) => return Err(error.into()),
1056 Err(_) => return Err(TransportError::Timeout(self.config.connect_timeout)),
1057 };
1058 let _ = tcp.set_nodelay(true);
1059 let security = self
1060 .security
1061 .read()
1062 .expect("transport security lock poisoned")
1063 .clone();
1064 match security {
1065 TransportSecurity::PlaintextForTesting => Ok(ClientStream::Plain(tcp)),
1066 TransportSecurity::Mtls(tls) => {
1067 let server_name = peer_server_name(peer)?;
1068 let connector = tokio_rustls::TlsConnector::from(tls.client_config());
1069 match tokio::time::timeout(
1070 self.config.handshake_timeout,
1071 connector.connect(server_name, tcp),
1072 )
1073 .await
1074 {
1075 Ok(Ok(stream)) => Ok(ClientStream::Tls(Box::new(stream))),
1076 Ok(Err(error)) => Err(TransportError::Tls(format!(
1077 "TLS handshake with {} failed: {error}",
1078 peer.address
1079 ))),
1080 Err(_) => Err(TransportError::Timeout(self.config.handshake_timeout)),
1081 }
1082 }
1083 }
1084 }
1085
1086 fn encode_payload<T: Serialize>(
1087 target: RaftNodeId,
1088 rpc: &T,
1089 ) -> Result<Vec<u8>, TransportError> {
1090 serde_json::to_vec(&RpcPayload { target, rpc })
1091 .map_err(|e| TransportError::ProtocolViolation(format!("unencodable request: {e}")))
1092 }
1093
1094 fn append_error(from: RaftNodeId, target: RaftNodeId, error: TransportError) -> AppendRpcError {
1095 match error {
1096 TransportError::Timeout(timeout) => RPCError::Timeout(Timeout {
1097 action: RPCTypes::AppendEntries,
1098 id: from,
1099 target,
1100 timeout,
1101 }),
1102 error @ (TransportError::NoRoute(_)
1103 | TransportError::Io(_)
1104 | TransportError::Tls(_)) => RPCError::Unreachable(Unreachable::new(&error)),
1105 error => RPCError::Network(NetworkError::new(&error)),
1106 }
1107 }
1108
1109 fn vote_error(from: RaftNodeId, target: RaftNodeId, error: TransportError) -> VoteRpcError {
1110 match error {
1111 TransportError::Timeout(timeout) => RPCError::Timeout(Timeout {
1112 action: RPCTypes::Vote,
1113 id: from,
1114 target,
1115 timeout,
1116 }),
1117 error @ (TransportError::NoRoute(_)
1118 | TransportError::Io(_)
1119 | TransportError::Tls(_)) => RPCError::Unreachable(Unreachable::new(&error)),
1120 error => RPCError::Network(NetworkError::new(&error)),
1121 }
1122 }
1123
1124 fn snapshot_error(
1125 from: RaftNodeId,
1126 target: RaftNodeId,
1127 error: TransportError,
1128 ) -> SnapshotRpcError {
1129 match error {
1130 TransportError::Timeout(timeout) => RPCError::Timeout(Timeout {
1131 action: RPCTypes::InstallSnapshot,
1132 id: from,
1133 target,
1134 timeout,
1135 }),
1136 error @ (TransportError::NoRoute(_)
1137 | TransportError::Io(_)
1138 | TransportError::Tls(_)) => RPCError::Unreachable(Unreachable::new(&error)),
1139 error => RPCError::Network(NetworkError::new(&error)),
1140 }
1141 }
1142}
1143
1144enum ClientStream {
1147 Plain(TcpStream),
1148 Tls(Box<tokio_rustls::client::TlsStream<TcpStream>>),
1149}
1150
1151impl RaftTransport for TcpTransport {
1152 async fn append_entries(
1153 &self,
1154 from: RaftNodeId,
1155 target: RaftNodeId,
1156 rpc: AppendEntriesRequest<MongrelRaftConfig>,
1157 ) -> Result<AppendEntriesResponse<RaftNodeId>, AppendRpcError> {
1158 let result = async {
1159 let payload = Self::encode_payload(target, &rpc)?;
1160 self.round_trip(
1161 target,
1162 RAFT_MSG_APPEND_ENTRIES_REQUEST,
1163 RAFT_MSG_APPEND_ENTRIES_RESPONSE,
1164 payload,
1165 self.config.rpc_timeout,
1166 )
1167 .await
1168 }
1169 .await;
1170 match result {
1171 Ok(bytes) => serde_json::from_slice(&bytes).map_err(|e| {
1172 RPCError::Network(NetworkError::new(&TransportError::ProtocolViolation(
1173 format!("undecodable AppendEntries response: {e}"),
1174 )))
1175 }),
1176 Err(error) => Err(Self::append_error(from, target, error)),
1177 }
1178 }
1179
1180 async fn vote(
1181 &self,
1182 from: RaftNodeId,
1183 target: RaftNodeId,
1184 rpc: VoteRequest<RaftNodeId>,
1185 ) -> Result<VoteResponse<RaftNodeId>, VoteRpcError> {
1186 let result = async {
1187 let payload = Self::encode_payload(target, &rpc)?;
1188 self.round_trip(
1189 target,
1190 RAFT_MSG_VOTE_REQUEST,
1191 RAFT_MSG_VOTE_RESPONSE,
1192 payload,
1193 self.config.rpc_timeout,
1194 )
1195 .await
1196 }
1197 .await;
1198 match result {
1199 Ok(bytes) => serde_json::from_slice(&bytes).map_err(|e| {
1200 RPCError::Network(NetworkError::new(&TransportError::ProtocolViolation(
1201 format!("undecodable Vote response: {e}"),
1202 )))
1203 }),
1204 Err(error) => Err(Self::vote_error(from, target, error)),
1205 }
1206 }
1207
1208 async fn install_snapshot(
1209 &self,
1210 from: RaftNodeId,
1211 target: RaftNodeId,
1212 rpc: InstallSnapshotRequest<MongrelRaftConfig>,
1213 ) -> Result<InstallSnapshotResponse<RaftNodeId>, SnapshotRpcError> {
1214 let result = async {
1215 let payload = Self::encode_payload(target, &rpc)?;
1216 self.round_trip(
1217 target,
1218 RAFT_MSG_INSTALL_SNAPSHOT_REQUEST,
1219 RAFT_MSG_INSTALL_SNAPSHOT_RESPONSE,
1220 payload,
1221 self.config.snapshot_timeout,
1222 )
1223 .await
1224 }
1225 .await;
1226 match result {
1227 Ok(bytes) => serde_json::from_slice(&bytes).map_err(|e| {
1228 RPCError::Network(NetworkError::new(&TransportError::ProtocolViolation(
1229 format!("undecodable InstallSnapshot response: {e}"),
1230 )))
1231 }),
1232 Err(error) => Err(Self::snapshot_error(from, target, error)),
1233 }
1234 }
1235
1236 fn attach(&self, node_id: RaftNodeId, raft: MongrelRaft) {
1237 self.registry.attach(node_id, raft);
1238 }
1239
1240 fn detach(&self, node_id: RaftNodeId) {
1241 self.registry.detach(node_id);
1242 }
1243
1244 async fn trigger_election(
1245 &self,
1246 target: RaftNodeId,
1247 ) -> Result<(), mongreldb_consensus::network::TransportError> {
1248 let payload = Self::encode_payload(target, &()).map_err(|error| {
1249 mongreldb_consensus::network::TransportError::Fault(error.to_string())
1250 })?;
1251 self.round_trip(
1252 target,
1253 RAFT_MSG_TRIGGER_ELECTION_REQUEST,
1254 RAFT_MSG_TRIGGER_ELECTION_RESPONSE,
1255 payload,
1256 self.config.rpc_timeout,
1257 )
1258 .await
1259 .map(|_| ())
1260 .map_err(|error| match error {
1261 TransportError::NoRoute(node) => {
1262 mongreldb_consensus::network::TransportError::NoRoute(node)
1263 }
1264 other => mongreldb_consensus::network::TransportError::Fault(other.to_string()),
1267 })
1268 }
1269}
1270
1271impl fmt::Debug for TcpTransport {
1272 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1273 f.debug_struct("TcpTransport")
1274 .field("config", &self.config)
1275 .field(
1276 "security",
1277 &*self.security.read().unwrap_or_else(|e| e.into_inner()),
1278 )
1279 .field("peers", &self.peers.lock().map(|peers| peers.len()))
1280 .field("registry", &self.registry)
1281 .finish()
1282 }
1283}
1284
1285pub struct TransportServer {
1302 local_addr: SocketAddr,
1303 shutdown: Option<watch::Sender<bool>>,
1304 task: Option<JoinHandle<()>>,
1305 #[allow(dead_code)]
1309 security: Arc<std::sync::RwLock<TransportSecurity>>,
1310}
1311
1312impl TransportServer {
1313 pub async fn bind(
1317 address: &str,
1318 security: TransportSecurity,
1319 registry: TransportRegistry,
1320 config: TransportConfig,
1321 ) -> Result<Self, TransportError> {
1322 Self::bind_shared(
1323 address,
1324 Arc::new(std::sync::RwLock::new(security)),
1325 registry,
1326 config,
1327 )
1328 .await
1329 }
1330
1331 pub async fn bind_shared(
1334 address: &str,
1335 security: Arc<std::sync::RwLock<TransportSecurity>>,
1336 registry: TransportRegistry,
1337 config: TransportConfig,
1338 ) -> Result<Self, TransportError> {
1339 let listener = TcpListener::bind(address).await?;
1340 let local_addr = listener.local_addr()?;
1341 let (shutdown_tx, shutdown_rx) = watch::channel(false);
1342 let task = tokio::spawn(Self::accept_loop(
1343 listener,
1344 Arc::clone(&security),
1345 registry,
1346 config,
1347 shutdown_rx,
1348 ));
1349 Ok(Self {
1350 local_addr,
1351 shutdown: Some(shutdown_tx),
1352 task: Some(task),
1353 security,
1354 })
1355 }
1356
1357 pub fn local_addr(&self) -> SocketAddr {
1359 self.local_addr
1360 }
1361
1362 pub async fn shutdown(mut self) {
1365 if let Some(shutdown) = self.shutdown.take() {
1366 let _ = shutdown.send(true);
1367 }
1368 if let Some(task) = self.task.take() {
1369 let _ = task.await;
1370 }
1371 }
1372
1373 async fn accept_loop(
1374 listener: TcpListener,
1375 security: Arc<std::sync::RwLock<TransportSecurity>>,
1376 registry: TransportRegistry,
1377 config: TransportConfig,
1378 mut shutdown: watch::Receiver<bool>,
1379 ) {
1380 let permits = Arc::new(Semaphore::new(config.max_connections));
1381 let mut connections: JoinSet<()> = JoinSet::new();
1382 loop {
1383 tokio::select! {
1384 biased;
1385 _ = shutdown.changed() => break,
1386 accepted = listener.accept() => {
1387 match accepted {
1388 Ok((stream, _peer)) => {
1389 match permits.clone().try_acquire_owned() {
1390 Ok(permit) => {
1391 let current = security
1392 .read()
1393 .expect("transport security lock poisoned")
1394 .clone();
1395 connections.spawn(Self::serve_connection(
1396 stream,
1397 current,
1398 registry.clone(),
1399 config.clone(),
1400 permit,
1401 ));
1402 }
1403 Err(_) => drop(stream),
1407 }
1408 }
1409 Err(_) => {
1410 tokio::time::sleep(Duration::from_millis(10)).await;
1412 }
1413 }
1414 }
1415 }
1416 }
1417 let drain = async { while connections.join_next().await.is_some() {} };
1418 if tokio::time::timeout(config.shutdown_grace, drain)
1419 .await
1420 .is_err()
1421 {
1422 connections.abort_all();
1423 while connections.join_next().await.is_some() {}
1424 }
1425 }
1426
1427 async fn serve_connection(
1428 stream: TcpStream,
1429 security: TransportSecurity,
1430 registry: TransportRegistry,
1431 config: TransportConfig,
1432 _permit: OwnedSemaphorePermit,
1433 ) {
1434 let _ = stream.set_nodelay(true);
1435 match security {
1436 TransportSecurity::PlaintextForTesting => {
1437 Self::serve_stream(stream, ®istry, &config).await;
1438 }
1439 TransportSecurity::Mtls(tls) => {
1440 let acceptor = tokio_rustls::TlsAcceptor::from(tls.server_config());
1441 let handshake =
1442 tokio::time::timeout(config.handshake_timeout, acceptor.accept(stream)).await;
1443 let Ok(Ok(mut tls_stream)) = handshake else {
1444 return;
1445 };
1446 let authenticated = {
1447 let (_, connection) = tls_stream.get_ref();
1448 verify_peer_identity(connection.peer_certificates(), &tls)
1449 };
1450 if authenticated.is_err() {
1451 return;
1454 }
1455 Self::serve_stream(&mut tls_stream, ®istry, &config).await;
1456 }
1457 }
1458 }
1459
1460 async fn serve_stream<S: AsyncRead + AsyncWrite + Unpin>(
1465 mut stream: S,
1466 registry: &TransportRegistry,
1467 config: &TransportConfig,
1468 ) {
1469 let frame =
1472 match read_frame(&mut stream, config.max_frame_bytes, config.snapshot_timeout).await {
1473 Ok(Some(frame)) => frame,
1474 Ok(None) | Err(_) => return,
1475 };
1476 let (response_type, response_payload) =
1477 Self::dispatch(registry, frame.message_type, &frame.payload).await;
1478 let _ = write_frame(
1479 &mut stream,
1480 response_type,
1481 response_payload,
1482 config.rpc_timeout,
1483 )
1484 .await;
1485 }
1486
1487 async fn dispatch(
1491 registry: &TransportRegistry,
1492 message_type: u32,
1493 payload: &[u8],
1494 ) -> (u32, Vec<u8>) {
1495 match message_type {
1496 RAFT_MSG_APPEND_ENTRIES_REQUEST => {
1497 Self::dispatch_rpc(
1498 registry,
1499 RAFT_MSG_APPEND_ENTRIES_RESPONSE,
1500 payload,
1501 |raft, rpc| async move { raft.append_entries(rpc).await },
1502 )
1503 .await
1504 }
1505 RAFT_MSG_VOTE_REQUEST => {
1506 Self::dispatch_rpc(
1507 registry,
1508 RAFT_MSG_VOTE_RESPONSE,
1509 payload,
1510 |raft, rpc| async move { raft.vote(rpc).await },
1511 )
1512 .await
1513 }
1514 RAFT_MSG_INSTALL_SNAPSHOT_REQUEST => {
1515 Self::dispatch_rpc(
1516 registry,
1517 RAFT_MSG_INSTALL_SNAPSHOT_RESPONSE,
1518 payload,
1519 |raft, rpc| async move { raft.install_snapshot(rpc).await },
1520 )
1521 .await
1522 }
1523 RAFT_MSG_TRIGGER_ELECTION_REQUEST => {
1524 let request = serde_json::from_slice::<RpcPayload<()>>(payload);
1525 match request {
1526 Err(error) => {
1527 Self::error_frame(format!("undecodable request payload: {error}"))
1528 }
1529 Ok(request) => match registry.get(request.target) {
1530 None => Self::error_frame(format!(
1531 "node {} is not attached to this transport",
1532 request.target
1533 )),
1534 Some(raft) => match raft.trigger().elect().await {
1535 Ok(()) => (RAFT_MSG_TRIGGER_ELECTION_RESPONSE, Vec::new()),
1536 Err(error) => Self::error_frame(format!("raft core error: {error}")),
1537 },
1538 },
1539 }
1540 }
1541 INTERNAL_MSG_REQUEST => {
1542 let request = serde_json::from_slice::<InternalRpcPayload>(payload);
1543 match request {
1544 Err(error) => {
1545 Self::error_frame(format!("undecodable internal request payload: {error}"))
1546 }
1547 Ok(request) => {
1548 match registry.internal_handler(request.target, request.service_id) {
1549 None => Self::error_frame(format!(
1550 "node {} has no internal RPC handler for service {}",
1551 request.target, request.service_id
1552 )),
1553 Some(handler) => match handler.handle(&request.body).await {
1554 Ok(response) => (INTERNAL_MSG_RESPONSE, response),
1555 Err(error) => Self::error_frame(format!(
1556 "internal RPC handler failed: {error}"
1557 )),
1558 },
1559 }
1560 }
1561 }
1562 }
1563 unknown => Self::error_frame(format!("unknown raft message type {unknown}")),
1564 }
1565 }
1566
1567 async fn dispatch_rpc<Req, Resp, E, F, Fut>(
1573 registry: &TransportRegistry,
1574 response_type: u32,
1575 payload: &[u8],
1576 call: F,
1577 ) -> (u32, Vec<u8>)
1578 where
1579 Req: for<'de> Deserialize<'de>,
1580 Resp: Serialize,
1581 E: fmt::Display,
1582 F: FnOnce(MongrelRaft, Req) -> Fut,
1583 Fut: std::future::Future<Output = Result<Resp, openraft::error::RaftError<RaftNodeId, E>>>,
1584 {
1585 let request = serde_json::from_slice::<RpcPayload<Req>>(payload);
1586 match request {
1587 Err(error) => Self::error_frame(format!("undecodable request payload: {error}")),
1588 Ok(request) => match registry.get(request.target) {
1589 None => Self::error_frame(format!(
1590 "node {} is not attached to this transport",
1591 request.target
1592 )),
1593 Some(raft) => match call(raft, request.rpc).await {
1594 Ok(response) => match serde_json::to_vec(&response) {
1595 Ok(bytes) => (response_type, bytes),
1596 Err(error) => Self::error_frame(format!("unencodable response: {error}")),
1597 },
1598 Err(error) => Self::error_frame(format!("raft core error: {error}")),
1599 },
1600 },
1601 }
1602 }
1603
1604 fn error_frame(message: String) -> (u32, Vec<u8>) {
1605 (
1606 RAFT_MSG_ERROR,
1607 serde_json::to_vec(&message).expect("string serialization is total"),
1608 )
1609 }
1610}
1611
1612impl Drop for TransportServer {
1613 fn drop(&mut self) {
1614 if let Some(task) = &self.task {
1616 task.abort();
1617 }
1618 }
1619}
1620
1621impl fmt::Debug for TransportServer {
1622 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1623 f.debug_struct("TransportServer")
1624 .field("local_addr", &self.local_addr)
1625 .finish_non_exhaustive()
1626 }
1627}
1628
1629#[cfg(test)]
1630mod tests {
1631 use super::*;
1632
1633 const FIXTURES: &str = concat!(env!("CARGO_MANIFEST_DIR"), "/tests/fixtures");
1634
1635 fn fixture(name: &str) -> String {
1636 std::fs::read_to_string(format!("{FIXTURES}/{name}")).expect("fixture is checked in")
1637 }
1638
1639 fn fixture_bytes(name: &str) -> Vec<u8> {
1640 std::fs::read(format!("{FIXTURES}/{name}")).expect("fixture is checked in")
1641 }
1642
1643 fn node_id(byte: u8) -> NodeId {
1644 NodeId::from_bytes([byte; 16])
1645 }
1646
1647 fn node1_tls(allowed: &[NodeId]) -> TlsConfig {
1648 TlsConfig::from_pems(
1649 &fixture("ca.crt.pem"),
1650 &fixture("node1.crt.pem"),
1651 &fixture("node1.key.pem"),
1652 allowed,
1653 )
1654 .expect("fixture TLS config loads")
1655 }
1656
1657 #[test]
1660 fn node_cert_name_matches_the_fixture_scheme() {
1661 assert_eq!(
1662 node_cert_name(&node_id(1)),
1663 "node-01010101010101010101010101010101.mongreldb.cluster"
1664 );
1665 assert_eq!(
1666 node_cert_name(&node_id(0xab)),
1667 "node-abababababababababababababababab.mongreldb.cluster"
1668 );
1669 }
1670
1671 #[test]
1674 fn tls_config_loads_fixture_material() {
1675 let allowed = vec![node_id(1), node_id(2)];
1676 let tls = node1_tls(&allowed);
1677 assert!(tls.allows_identity(&node_cert_name(&node_id(1))));
1678 assert!(tls.allows_identity(&node_cert_name(&node_id(2))));
1679 assert!(!tls.allows_identity(&node_cert_name(&node_id(3))));
1680 }
1681
1682 #[test]
1683 fn tls_config_rejects_unparsable_material() {
1684 let allowed = vec![node_id(1)];
1685 for (ca, cert, key) in [
1686 (
1687 "not pem",
1688 fixture("node1.crt.pem").as_str(),
1689 fixture("node1.key.pem").as_str(),
1690 ),
1691 (
1692 fixture("ca.crt.pem").as_str(),
1693 "not pem",
1694 fixture("node1.key.pem").as_str(),
1695 ),
1696 (
1697 fixture("ca.crt.pem").as_str(),
1698 fixture("node1.crt.pem").as_str(),
1699 "not pem",
1700 ),
1701 ] {
1702 assert!(
1703 TlsConfig::from_pems(ca, cert, key, &allowed).is_err(),
1704 "garbage material must fail closed"
1705 );
1706 }
1707 assert!(TlsConfig::from_pems(
1709 &fixture("ca.crt.pem"),
1710 &fixture("node1.crt.pem"),
1711 &fixture("node1.key.pem"),
1712 &[],
1713 )
1714 .is_err());
1715 }
1716
1717 #[test]
1718 fn tls_config_loads_from_trust_config() {
1719 let trust = crate::bootstrap::TrustConfig::from_pems(
1720 fixture("ca.crt.pem"),
1721 fixture("node1.crt.pem"),
1722 fixture("node1.key.pem"),
1723 vec![node_id(1)],
1724 )
1725 .expect("trust config validates");
1726 let tls = TlsConfig::from_trust(&trust).expect("TLS config from trust");
1727 assert!(tls.allows_identity(&node_cert_name(&node_id(1))));
1728 }
1729
1730 fn pem_cert_der(name: &str) -> Vec<u8> {
1733 let pem = fixture_bytes(name);
1734 let mut reader = io::BufReader::new(pem.as_slice());
1735 let certificates: Vec<_> = rustls_pemfile::certs(&mut reader)
1736 .collect::<Result<_, _>>()
1737 .expect("parsable certificate");
1738 certificates
1739 .into_iter()
1740 .next()
1741 .expect("one certificate")
1742 .to_vec()
1743 }
1744
1745 #[test]
1746 fn der_extracts_node_identity_from_cn_and_san() {
1747 let identities = der::certificate_identities(&pem_cert_der("node1.crt.pem")).unwrap();
1748 assert!(
1749 identities.contains(&node_cert_name(&node_id(1))),
1750 "expected the node identity in {identities:?}"
1751 );
1752 assert_eq!(identities.len(), 2, "{identities:?}");
1754 }
1755
1756 #[test]
1757 fn der_extracts_ca_cn_without_san() {
1758 let identities = der::certificate_identities(&pem_cert_der("ca.crt.pem")).unwrap();
1759 assert_eq!(identities, vec!["mongreldb-test-ca".to_owned()]);
1760 }
1761
1762 #[test]
1763 fn der_fails_closed_on_malformed_input() {
1764 assert!(der::certificate_identities(&[]).is_err());
1765 assert!(der::certificate_identities(b"junk").is_err());
1766 let mut truncated = pem_cert_der("node1.crt.pem");
1767 truncated.truncate(truncated.len() / 2);
1768 assert!(der::certificate_identities(&truncated).is_err());
1769 let mut garbage = pem_cert_der("node1.crt.pem");
1770 for byte in &mut garbage[10..20] {
1771 *byte ^= 0xff;
1772 }
1773 assert!(der::certificate_identities(&garbage).is_err());
1774 }
1775
1776 const TEST_MAX_FRAME: usize = 1024;
1779
1780 async fn write_and_read(
1781 bytes: &[u8],
1782 max_frame: usize,
1783 ) -> Result<Option<ProtocolEnvelope>, TransportError> {
1784 let (mut writer, mut reader) = tokio::io::duplex(bytes.len() + 64);
1785 writer.write_all(bytes).await.unwrap();
1786 drop(writer);
1787 read_frame(&mut reader, max_frame, Duration::from_secs(5)).await
1788 }
1789
1790 #[tokio::test]
1791 async fn frame_round_trip() {
1792 let (mut writer, mut reader) = tokio::io::duplex(4096);
1793 write_frame(
1794 &mut writer,
1795 RAFT_MSG_VOTE_REQUEST,
1796 b"hello".to_vec(),
1797 Duration::from_secs(5),
1798 )
1799 .await
1800 .unwrap();
1801 let frame = read_frame(&mut reader, TEST_MAX_FRAME, Duration::from_secs(5))
1802 .await
1803 .unwrap()
1804 .expect("a frame");
1805 assert_eq!(frame.message_type, RAFT_MSG_VOTE_REQUEST);
1806 assert_eq!(frame.payload, b"hello");
1807 drop(writer);
1809 assert_eq!(
1810 read_frame(&mut reader, TEST_MAX_FRAME, Duration::from_secs(5))
1811 .await
1812 .unwrap(),
1813 None
1814 );
1815 }
1816
1817 #[tokio::test]
1818 async fn unknown_protocol_version_fails_closed() {
1819 let mut envelope = ProtocolEnvelope::new(1, vec![1, 2, 3]);
1820 envelope.protocol_version = 99;
1821 envelope.payload_crc32 =
1822 ProtocolEnvelope::checksum(envelope.protocol_version, 1, &envelope.payload);
1823 let error = write_and_read(&envelope.encode(), TEST_MAX_FRAME)
1824 .await
1825 .unwrap_err();
1826 assert!(
1827 matches!(
1828 error,
1829 TransportError::Envelope(EnvelopeError::UnsupportedVersion { found: 99, .. })
1830 ),
1831 "unexpected error: {error}"
1832 );
1833 }
1834
1835 #[tokio::test]
1836 async fn oversize_length_prefix_fails_before_reading_the_payload() {
1837 let mut header = Vec::new();
1840 header.extend_from_slice(&1u32.to_le_bytes());
1841 header.extend_from_slice(&1u32.to_le_bytes());
1842 header.extend_from_slice(&(TEST_MAX_FRAME as u32 + 1).to_le_bytes());
1843 let error = write_and_read(&header, TEST_MAX_FRAME).await.unwrap_err();
1844 assert!(
1845 matches!(
1846 error,
1847 TransportError::FrameTooLarge {
1848 actual,
1849 limit
1850 } if actual == TEST_MAX_FRAME + 1 && limit == TEST_MAX_FRAME
1851 ),
1852 "unexpected error: {error}"
1853 );
1854 let mut header = Vec::new();
1857 header.extend_from_slice(&1u32.to_le_bytes());
1858 header.extend_from_slice(&1u32.to_le_bytes());
1859 header.extend_from_slice(&(MAX_MESSAGE_PAYLOAD_BYTES as u32 + 1).to_le_bytes());
1860 let error = write_and_read(&header, usize::MAX).await.unwrap_err();
1861 assert!(
1862 matches!(error, TransportError::FrameTooLarge { .. }),
1863 "unexpected error: {error}"
1864 );
1865 }
1866
1867 #[tokio::test]
1868 async fn truncated_frame_fails_closed() {
1869 let envelope = ProtocolEnvelope::new(3, vec![9u8; 64]);
1870 let bytes = envelope.encode();
1871 let error = write_and_read(&bytes[..bytes.len() - 5], TEST_MAX_FRAME)
1873 .await
1874 .unwrap_err();
1875 assert!(
1876 matches!(error, TransportError::Io(ref e) if e.kind() == io::ErrorKind::UnexpectedEof),
1877 "unexpected error: {error}"
1878 );
1879 let error = write_and_read(&bytes[..4], TEST_MAX_FRAME)
1881 .await
1882 .unwrap_err();
1883 assert!(
1884 matches!(error, TransportError::Io(ref e) if e.kind() == io::ErrorKind::UnexpectedEof),
1885 "unexpected error: {error}"
1886 );
1887 }
1888
1889 #[tokio::test]
1890 async fn checksum_mismatch_fails_closed() {
1891 let envelope = ProtocolEnvelope::new(2, vec![7u8; 32]);
1892 let mut bytes = envelope.encode();
1893 bytes[HEADER_LEN] ^= 0x01;
1894 let error = write_and_read(&bytes, TEST_MAX_FRAME).await.unwrap_err();
1895 assert!(
1896 matches!(
1897 error,
1898 TransportError::Envelope(EnvelopeError::ChecksumMismatch)
1899 ),
1900 "unexpected error: {error}"
1901 );
1902 }
1903
1904 #[tokio::test]
1905 async fn write_rejects_oversize_payload() {
1906 let (mut writer, _reader) = tokio::io::duplex(64);
1907 let error = write_frame(
1908 &mut writer,
1909 1,
1910 vec![0u8; MAX_MESSAGE_PAYLOAD_BYTES + 1],
1911 Duration::from_secs(5),
1912 )
1913 .await
1914 .unwrap_err();
1915 assert!(matches!(error, TransportError::FrameTooLarge { .. }));
1916 }
1917
1918 #[test]
1921 fn registry_attach_detach_lifecycle() {
1922 let registry = TransportRegistry::new();
1923 assert!(registry.is_empty());
1924 assert!(registry.get(42).is_none());
1925 registry.detach(42); assert!(registry.is_empty());
1927 }
1928
1929 #[test]
1932 fn reload_security_swaps_material_for_new_handshakes() {
1933 let transport = TcpTransport::new(
1934 TransportConfig::default(),
1935 TransportSecurity::PlaintextForTesting,
1936 );
1937 assert!(matches!(
1938 transport.security(),
1939 TransportSecurity::PlaintextForTesting
1940 ));
1941 let node = node_id(1);
1942 let tls = node1_tls(&[node]);
1943 transport.reload_security(TransportSecurity::Mtls(tls));
1944 assert!(matches!(transport.security(), TransportSecurity::Mtls(_)));
1945 let handle = transport.security_handle();
1947 assert!(matches!(
1948 &*handle.read().expect("lock"),
1949 TransportSecurity::Mtls(_)
1950 ));
1951 }
1952
1953 #[test]
1954 fn peer_directory_crud() {
1955 let transport = TcpTransport::new(
1956 TransportConfig::default(),
1957 TransportSecurity::PlaintextForTesting,
1958 );
1959 assert_eq!(transport.peer(1), None);
1960 transport.upsert_peer(1, PeerEndpoint::plaintext("127.0.0.1:9001"));
1961 assert_eq!(
1962 transport.peer(1),
1963 Some(PeerEndpoint::plaintext("127.0.0.1:9001"))
1964 );
1965 transport.upsert_peer(1, PeerEndpoint::plaintext("127.0.0.1:9002"));
1966 assert_eq!(
1967 transport.peer(1).map(|peer| peer.address),
1968 Some("127.0.0.1:9002".to_owned())
1969 );
1970 assert!(transport.remove_peer(1).is_some());
1971 assert_eq!(transport.peer(1), None);
1972 }
1973
1974 #[test]
1975 fn mtls_endpoint_requires_an_identity() {
1976 let transport = TcpTransport::new(
1977 TransportConfig::default(),
1978 TransportSecurity::Mtls(node1_tls(&[node_id(1), node_id(2)])),
1979 );
1980 let error = super::peer_server_name(&PeerEndpoint::plaintext("127.0.0.1:1")).unwrap_err();
1983 assert!(
1984 matches!(error, TransportError::PeerAuthentication(_)),
1985 "unexpected error: {error}"
1986 );
1987 assert!(super::peer_server_name(&PeerEndpoint::mtls("127.0.0.1:1", node_id(2))).is_ok());
1988 drop(transport);
1989 }
1990}